// Package actionlog records what CSM did to a host, as opposed to what it
// found. Findings already have one stream (the SIEM audit log); actions were
// spread across a firewall log, a web UI log, a mail-freeze log and, for
// everything the daemon did on its own, nothing at all.
//
// One record per action, one file, one schema. Each record names the
// privileged operation it belongs to (the IDs in internal/privops), so an
// operator can read the capability matrix and the action log with the same
// vocabulary, and carries the evidence that makes the action reviewable: the
// exact argv when CSM ran a command, and the file's digest before and after
// when it changed a file.
package actionlog
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/safepath"
)
// SchemaVersion is bumped only on an incompatible change. Additive fields do
// not bump it, so a parser can pin on v and ignore unknown keys.
const SchemaVersion = 1
// Result is the outcome of an action.
type Result string
const (
// Applied means the action changed the host.
Applied Result = "applied"
// DryRun means CSM decided to act and did not, because a dry-run or
// observe setting was in force. The record says what it would have done.
DryRun Result = "dry_run"
// Failed means the action was attempted and did not complete.
Failed Result = "failed"
// Refused means a safety rule stopped the action before it ran.
Refused Result = "refused"
)
// Actor is who started the action.
type Actor string
const (
// Daemon means CSM acted on its own.
Daemon Actor = "daemon"
// CLI means an operator ran a command.
CLI Actor = "cli"
// WebUI means an operator used the dashboard or the API.
WebUI Actor = "webui"
)
// FileState is a file's identity at one moment. Digest is a SHA-256 of the
// content; an empty digest with Exists false records a file that was not there.
type FileState struct {
Exists bool `json:"exists"`
Digest string `json:"sha256,omitempty"`
Size int64 `json:"size,omitempty"`
Mode string `json:"mode,omitempty"`
UID uint32 `json:"uid,omitempty"`
GID uint32 `json:"gid,omitempty"`
}
// Record is one action.
type Record struct {
V int `json:"v"`
Timestamp time.Time `json:"ts"`
Hostname string `json:"hostname,omitempty"`
// Op is the privileged-operation ID from internal/privops.
Op string `json:"op"`
// Action distinguishes changes within one operation, such as block and unblock.
Action string `json:"action,omitempty"`
// Actor and ActorDetail say who asked for it. ActorDetail carries the
// operator's source address for a web UI action, or the command name for
// a CLI action.
Actor Actor `json:"actor"`
ActorDetail string `json:"actor_detail,omitempty"`
// FindingID ties the action to the finding that caused it, using the same
// ID the SIEM audit log emits.
FindingID string `json:"finding_id,omitempty"`
IncidentID string `json:"incident_id,omitempty"`
// ActionID and ActionVersion identify a durable lifecycle event across retries.
ActionID string `json:"action_id,omitempty"`
ActionVersion uint64 `json:"action_version,omitempty"`
// UndoOf links a typed undo to its original action.
UndoOf string `json:"undo_of,omitempty"`
// Target is what was acted on: a path, an address, a message ID.
Target string `json:"target"`
Account string `json:"account,omitempty"`
Reason string `json:"reason,omitempty"`
// Command is the exact argv when CSM ran a program. Empty when the action
// was performed through system calls.
Command []string `json:"command,omitempty"`
// Before and After are the target file's state around the change.
Before *FileState `json:"before,omitempty"`
After *FileState `json:"after,omitempty"`
Result Result `json:"result"`
Error string `json:"error,omitempty"`
// Undo is the command that reverses the action, when one exists.
Undo string `json:"undo,omitempty"`
// RecoveryPath identifies retained file content and its metadata sidecar.
RecoveryPath string `json:"recovery_path,omitempty"`
}
// Sink writes records. The daemon installs a file sink at startup; tests
// install their own.
type Sink interface {
Write(Record) error
}
var (
mu sync.RWMutex
sink Sink
host string
byActor = Daemon
)
// SetSink installs the process-wide sink. Passing nil disables recording,
// which is what a CLI that has not opened the log does.
func SetSink(s Sink, hostname string) {
mu.Lock()
defer mu.Unlock()
sink, host = s, hostname
}
// SetDefaultActor declares which process is recording. The daemon leaves it at
// Daemon; a CLI that opens the log sets CLI. Actions an operator starts through
// the web UI run inside the daemon, so they record as Daemon here and carry the
// operator's address in the web UI's own action log.
func SetDefaultActor(a Actor) {
mu.Lock()
defer mu.Unlock()
byActor = a
}
// DefaultActor returns the process-wide actor for call sites that have no more
// specific attribution.
func DefaultActor() Actor {
mu.RLock()
defer mu.RUnlock()
return byActor
}
// A stuck filesystem or a broken sink must not hold a response worker forever.
// Bound both the wait and outstanding writes; normal writes complete before
// returning, including in short-lived CLI processes.
const writeTimeout = 250 * time.Millisecond
// Write records one action without propagating sink errors or panics. Recording
// is best effort: saturation or an unresponsive sink can cost an action record.
func Write(r Record) {
mu.RLock()
s, h, a := sink, host, byActor
mu.RUnlock()
if s == nil {
return
}
actionWrites.write(s, prepareRecord(r, h, a))
}
func prepareRecord(r Record, h string, a Actor) Record {
r.V = SchemaVersion
if r.Timestamp.IsZero() {
r.Timestamp = time.Now().UTC()
}
if r.Hostname == "" {
r.Hostname = h
}
if r.Actor == "" {
r.Actor = a
}
// A timed-out write can outlive the caller's buffers.
r.Command = append([]string(nil), r.Command...)
if r.Before != nil {
before := *r.Before
r.Before = &before
}
if r.After != nil {
after := *r.After
r.After = &after
}
return r
}
// ErrDurableUnavailable means no acknowledging sink is installed.
var ErrDurableUnavailable = errors.New("durable action log sink unavailable")
// ErrDurableUnacknowledged means delivery did not finish within the bounded
// wait. The sink may still complete; retry using the same action identity.
var ErrDurableUnacknowledged = errors.New("durable action log delivery unacknowledged")
// WriteDurable acknowledges only a sink's durable write. Delivery is at least
// once: a timeout can leave an append running after this call returns.
func WriteDurable(r Record) error {
mu.RLock()
s, h, a := sink, host, byActor
mu.RUnlock()
durable, ok := s.(interface{ WriteDurable(Record) error })
if !ok {
return ErrDurableUnavailable
}
acknowledgement := make(chan error, 1)
actionWrites.write(durableWriteAdapter{write: durable.WriteDurable, acknowledgement: acknowledgement}, prepareRecord(r, h, a))
select {
case err := <-acknowledgement:
return err
default:
return ErrDurableUnacknowledged
}
}
type durableWriteAdapter struct {
write func(Record) error
acknowledgement chan<- error
}
func (s durableWriteAdapter) Write(r Record) error {
err := s.write(r)
s.acknowledgement <- err
return err
}
// maxFileSize is the rotation threshold, matching the firewall and web UI
// logs this stream consolidates.
const maxFileSize = 10 * 1024 * 1024
// FileSink appends JSON lines to a file, rotating it once at the threshold.
type FileSink struct {
resolve func() string
mu sync.Mutex
path string
onErr func(error)
syncFile func(*os.File) error
}
// NewFileSink returns a sink writing to the file resolve names. The path is
// resolved on the first record and remembered, so installing a sink in a
// process that never records an action costs nothing: a CLI that only reads
// does not load config to work out where the log lives. onErr reports write
// failures; pass nil to discard them.
func NewFileSink(resolve func() string, onErr func(error)) *FileSink {
return &FileSink{resolve: resolve, onErr: onErr}
}
// logPath returns the resolved path, resolving it once. The caller holds mu.
func (f *FileSink) logPath() string {
if f.path == "" {
f.path = f.resolve()
}
return f.path
}
// DefaultPath is where the daemon keeps the action log.
func DefaultPath(logDir string) string { return filepath.Join(logDir, "actions.jsonl") }
func (f *FileSink) Write(r Record) error {
err := f.write(r, false)
// Reporting outside the file lock lets a callback inspect or replace the
// sink without deadlocking a completed action.
if err != nil {
f.report(err)
}
return err
}
// WriteDurable appends and syncs the record and its directory entries while
// holding the same cross-process lock as ordinary writes and rotation.
func (f *FileSink) WriteDurable(r Record) error {
err := f.write(r, true)
if err != nil {
f.report(err)
}
return err
}
func (f *FileSink) sync(file *os.File) error {
if f.syncFile != nil {
return f.syncFile(file)
}
return file.Sync()
}
func (f *FileSink) write(r Record, durable bool) error {
// Cleaning explanations and command errors may contain attacker-controlled
// content. Keep individual lines readable by the bounded history reader.
r.Reason = boundedDetail(r.Reason)
r.Error = boundedDetail(r.Error)
data, marshalErr := json.Marshal(r)
if marshalErr != nil {
return marshalErr
}
data = append(data, '\n')
f.mu.Lock()
defer f.mu.Unlock()
path := f.logPath()
if err := os.MkdirAll(filepath.Dir(path), 0750); err != nil {
return err
}
// The daemon and CLI are separate writers. Lock a stable sidecar inode so
// rotation cannot move another writer's newly opened log out from under it.
lock, lockErr := openLogFile(path+".lock", os.O_RDWR|os.O_CREATE, 0640)
if lockErr != nil {
return lockErr
}
defer func() { _ = lock.Close() }()
// #nosec G115 -- an open file descriptor fits in int on supported Unix hosts.
if err := unix.Flock(int(lock.Fd()), unix.LOCK_EX); err != nil {
return err
}
defer func() { _ = unix.Flock(int(lock.Fd()), unix.LOCK_UN) }() // #nosec G115 -- open file descriptor.
if info, statErr := os.Lstat(path); statErr == nil {
if !info.Mode().IsRegular() {
return fmt.Errorf("action log is not a regular file: %s", path)
}
if info.Size() > maxFileSize {
if err := os.Rename(path, path+".1"); err != nil {
return err
}
}
}
fh, err := openLogFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0640)
if err != nil {
return err
}
if _, err := fh.Write(data); err != nil {
_ = fh.Close()
return err
}
if durable {
if err := f.sync(fh); err != nil {
_ = fh.Close()
return err
}
}
if err := fh.Close(); err != nil {
return err
}
if durable {
return f.syncDirectories(filepath.Dir(path))
}
return nil
}
// Sync the whole ancestor chain because MkdirAll may have created multiple
// directories, and a retry must also acknowledge earlier uncertain creation.
func (f *FileSink) syncDirectories(dir string) error {
absolute, err := filepath.Abs(dir)
if err != nil {
return err
}
for {
// #nosec G304 -- operator-configured log directory, opened only for syncing.
directory, err := os.Open(absolute)
if err != nil {
return err
}
syncErr := f.sync(directory)
closeErr := directory.Close()
if syncErr != nil {
return syncErr
}
if closeErr != nil {
return closeErr
}
parent := filepath.Dir(absolute)
if parent == absolute {
return nil
}
absolute = parent
}
}
func boundedDetail(value string) string {
const limit = 4096
if len(value) > limit {
return strings.Clone(value[:limit]) + " [truncated]"
}
return value
}
func openLogFile(path string, flags int, mode os.FileMode) (*os.File, error) {
// #nosec G304 G302 -- operator-configured log path; no symlinks or special
// files. 0640 allows the log shipper's group to read the stream.
f, err := os.OpenFile(path, flags|unix.O_NOFOLLOW|unix.O_NONBLOCK, mode)
if err != nil {
return nil, err
}
info, err := f.Stat()
if err == nil && !info.Mode().IsRegular() {
err = fmt.Errorf("action log is not a regular file: %s", path)
}
if err != nil {
_ = f.Close()
return nil, err
}
return f, nil
}
// Read visits fixed snapshots of the rotated and current logs, oldest first.
// Both files are pinned before releasing the rotation lock, so a slow reader
// neither loses a rotated file nor holds up action writers.
func Read(path string, visit func(io.Reader) error) error {
for {
retry, err := readSnapshot(path, visit)
if !retry {
return err
}
}
}
func readSnapshot(path string, visit func(io.Reader) error) (bool, error) {
lock, err := openLogFile(path+".lock", os.O_RDONLY, 0)
if err != nil && !os.IsNotExist(err) {
return false, err
}
if lock != nil {
defer func() { _ = lock.Close() }()
// #nosec G115 -- an open file descriptor fits in int on supported Unix hosts.
if err := unix.Flock(int(lock.Fd()), unix.LOCK_SH); err != nil {
return false, err
}
}
var files []*os.File
defer func() {
for _, f := range files {
_ = f.Close()
}
}()
var readers []io.Reader
for _, name := range []string{path + ".1", path} {
f, openErr := openLogFile(name, os.O_RDONLY, 0)
if os.IsNotExist(openErr) {
continue
}
if openErr != nil {
return false, openErr
}
files = append(files, f)
info, statErr := f.Stat()
if statErr != nil {
return false, statErr
}
readers = append(readers, io.NewSectionReader(f, 0, info.Size()))
}
if lock == nil {
// The first writer may have created its lock while we opened the
// files. Retry under that lock before exposing a mixed snapshot.
if _, err := os.Lstat(path + ".lock"); err == nil {
return true, nil
} else if !os.IsNotExist(err) {
return false, err
}
} else if err := unix.Flock(int(lock.Fd()), unix.LOCK_UN); err != nil { // #nosec G115 -- open file descriptor.
return false, err
}
for _, reader := range readers {
if err := visit(reader); err != nil {
return false, err
}
}
return false, nil
}
func (f *FileSink) report(err error) {
if f.onErr != nil {
f.onErr(err)
}
}
// Metadata captures identity without opening or hashing the target. It is safe
// to use before a security action: evidence collection must not delay removal.
func Metadata(path string) *FileState {
info, err := os.Lstat(path)
if err != nil {
return &FileState{}
}
return FromInfo(info)
}
// FromInfo describes metadata already pinned by the operation itself.
func FromInfo(info os.FileInfo) *FileState {
if info == nil {
return &FileState{}
}
st := &FileState{Exists: true, Size: info.Size(), Mode: info.Mode().String()}
fillOwner(st, info)
return st
}
// Stat hashes only a bounded regular file reached without following symlinks.
// No digest is better than a digest of a different inode or a truncated prefix.
func Stat(path string) *FileState {
st := Metadata(path)
if !st.Exists {
return st
}
absolute, err := filepath.Abs(path)
if err != nil {
return st
}
dir, err := safepath.OpenDirNoFollow(filepath.Dir(absolute))
if err != nil {
return st
}
defer func() { _ = dir.Close() }()
f, err := dir.OpenFile(filepath.Base(absolute), os.O_RDONLY, 0)
if err != nil {
return st
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return st
}
st = FromInfo(info)
if !info.Mode().IsRegular() || info.Size() > maxDigestBytes {
return st
}
h := sha256.New()
n, err := io.Copy(h, io.LimitReader(f, maxDigestBytes+1))
if err != nil || n != info.Size() || n > maxDigestBytes {
return st
}
after, err := f.Stat()
if err == nil && after.Size() == info.Size() && after.ModTime().Equal(info.ModTime()) {
st.Digest = hex.EncodeToString(h.Sum(nil))
}
return st
}
const maxDigestBytes = 64 * 1024 * 1024
// ContentState captures the exact bytes read or written through a pinned file.
func ContentState(info os.FileInfo, data []byte) *FileState {
st := FromInfo(info)
st.Size = int64(len(data))
if len(data) <= maxDigestBytes {
sum := sha256.Sum256(data)
st.Digest = hex.EncodeToString(sum[:])
}
return st
}
// Describe renders a record as one operator-readable line, used by `csm
// actions` and by the daemon log.
func (r Record) Describe() string {
line := fmt.Sprintf("%s %s %s target=%q", r.Timestamp.UTC().Format(time.RFC3339), r.Op, r.Result, r.Target)
if r.Action != "" {
line += fmt.Sprintf(" action=%q", r.Action)
}
if r.Account != "" {
line += fmt.Sprintf(" account=%q", r.Account)
}
if len(r.Command) > 0 {
line += " command=" + fmt.Sprintf("%q", r.Command)
}
if r.Before != nil && r.After != nil {
line += fmt.Sprintf(" sha256 %s -> %s", shortDigest(r.Before), shortDigest(r.After))
}
if r.Error != "" {
line += fmt.Sprintf(" error=%q", r.Error)
}
return line
}
func shortDigest(s *FileState) string {
switch {
case s == nil || !s.Exists:
return "absent"
case s.Digest == "":
return "unhashed"
default:
return s.Digest[:min(12, len(s.Digest))]
}
}
//go:build unix
package actionlog
import (
"os"
"syscall"
)
func fillOwner(st *FileState, info os.FileInfo) {
if sys, ok := info.Sys().(*syscall.Stat_t); ok {
st.UID, st.GID = sys.Uid, sys.Gid
}
}
package actionlog
import (
"log"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type writePool struct {
slots chan struct{}
stats *queuehealth.Tracker
}
// writeStallBudget is how long a sink write may run before the row reports a
// stall. The caller's own wait budget is far shorter, so a batch of actions
// serialising behind one file lock is normal rather than a degradation.
const writeStallBudget = time.Minute
var actionWrites = newWritePool(64)
func newWritePool(capacity int) *writePool {
return &writePool{
slots: make(chan struct{}, capacity),
stats: queuehealth.NewSharedCapacity(capacity, writeStallBudget),
}
}
// QueueStatus preserves process-wide work and loss across sink changes. A
// caller deadline does not discard a record that its sink may still write.
func QueueStatus(now time.Time) queuehealth.Status {
return actionWrites.stats.Snapshot(now)
}
func (p *writePool) write(s Sink, r Record) {
timer := time.NewTimer(writeTimeout)
defer timer.Stop()
select {
case p.slots <- struct{}{}:
case <-timer.C:
p.stats.Lose(time.Now(), 1)
return
}
ticket := p.stats.Begin(time.Now())
ticket.Start(time.Now())
done := make(chan struct{})
go func() {
failed := true
defer func() {
v := recover()
if failed {
p.stats.Lose(time.Now(), 1)
}
defer func() {
<-p.slots
ticket.Finish(time.Now())
close(done)
}()
if v != nil {
log.Printf("action log sink panicked: %v", v)
}
}()
failed = s.Write(r) != nil
}()
select {
case <-done:
case <-timer.C:
}
}
package alert
import (
"crypto/sha256"
"encoding/json"
"fmt"
"net"
"os"
"path/filepath"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/processctx"
"github.com/pidginhost/csm/internal/queuehealth"
)
const alertDispatchFailuresMetric = "csm_alert_dispatch_failures_total"
// alertDispatchFailures counts individual channel send failures (email,
// webhook, phpanel) so an operator can see when alerts are silently failing
// to deliver instead of the daemon looking healthy while findings never
// reach anyone.
var alertDispatchFailures = metrics.NewCounter(
alertDispatchFailuresMetric,
"Alert deliveries that failed (email/webhook/phpanel). Sustained growth means findings are being detected but not reaching operators -- check SMTP/webhook reachability and credentials.",
)
func init() {
metrics.MustRegister(alertDispatchFailuresMetric, alertDispatchFailures)
}
func addDispatchError(errs *[]error, err error) {
*errs = append(*errs, err)
alertDispatchFailures.Inc()
}
// Severity levels for findings.
type Severity int
const (
Warning Severity = iota
High Severity = iota
Critical Severity = iota
)
func (s Severity) String() string {
switch s {
case Warning:
return "WARNING"
case High:
return "HIGH"
case Critical:
return "CRITICAL"
}
return "UNKNOWN"
}
// Finding represents a single security check result.
type Finding struct {
queueTicket queuehealth.Ticket
Severity Severity `json:"severity"`
// DemotedFrom retains the severity an automatically demoted finding came
// from, so a later positive re-check can restore it. It is deliberately not
// part of Key(): a finding's identity must not change when its severity
// does, or every dismissal and dedup entry keyed to it would be orphaned.
DemotedFrom Severity `json:"demoted_from,omitempty"`
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details,omitempty"`
// CoverageScope identifies a scanner-owned unit that can be retired after
// complete coverage. Empty legacy scopes require whole-check completion.
// It is opaque, contains no credentials, and does not change dedup identity.
CoverageScope string `json:"coverage_scope,omitempty"`
// DedupKey, when set, pins the finding's dedup identity (Key and
// Fingerprint) regardless of Message/Details content. For findings whose
// details embed volatile values (pids, byte counts) that would otherwise
// mint a new identity for the same ongoing condition on every scan.
DedupKey string `json:"dedup_key,omitempty"`
FilePath string `json:"file_path,omitempty"`
ProcessInfo string `json:"process_info,omitempty"` // "pid=N cmd=name uid=N" from fanotify
PID int `json:"pid,omitempty"` // structured PID for auto-response
// Content fingerprint for re-verifiable content findings (PHP heuristics,
// signature, YARA). Set at emit time by content-family checks; empty for
// all other finding kinds and for findings emitted before this field
// existed. The re-verification re-check uses it to tell a superseded-
// heuristic false positive (identical bytes, current logic no longer
// flags) from a file edited after detection (never auto-cleared).
ContentSHA256 string `json:"content_sha256,omitempty"`
// DetectLogic is the checks.ContentDetectionVersion() token in effect when
// this content finding was emitted. Optional; used for sweep gating and
// audit explainability.
DetectLogic string `json:"detect_logic,omitempty"`
// ScanCarryForward marks an unchanged snapshot re-emitted only because the
// current scan could not examine its path. It is process-local provenance for
// the atomic latest-state merge, not part of the public finding contract.
ScanCarryForward bool `json:"-"`
// AutoFileResponseEvaluated records process-local delivery provenance.
// A detector or scan already considered automatic file remediation, so
// the alert dispatcher must not retry it, including a refused attempt.
// New detections and findings read from storage get a fresh evaluation.
AutoFileResponseEvaluated bool `json:"-"`
// PHP-relay structured fields (Stage 1 email_php_relay_abuse). All optional;
// zero values mean "this finding does not carry that dimension".
Path string `json:"path,omitempty"` // path1 trigger label: "header" | "volume" | "volume_account" | "fanout" | "baseline" | "reputation"
MsgIDs []string `json:"msg_ids,omitempty"` // sample of in-flight msgIDs (auto-action acts on the live snapshot, not this list)
ScriptKey string `json:"script_key,omitempty"` // host:path from X-PHP-Script
SourceIP string `json:"source_ip,omitempty"` // IP after "for " in X-PHP-Script
CPUser string `json:"cp_user,omitempty"` // cPanel user from spool -H line 2
// RelayTotal is the trigger count for the PHP-relay path that fired
// (qualifying/volume/fanout/account-window count). RelayBreakdown lists
// the scripts that contributed, with per-script hit counts and a sample
// subject. Both optional; volume_account carries RelayTotal with no
// breakdown (account log-tail path has no trustworthy script key).
RelayTotal int `json:"relay_total,omitempty"`
RelayBreakdown []RelayScriptHit `json:"relay_breakdown,omitempty"`
// Tenant context (added v2.12.0). Optional - populated when the check
// has enough info to attribute the finding to a specific tenant within
// a multi-tenant host. Empty strings render as omitted JSON keys so
// existing webhook consumers see no diff.
TenantID string `json:"tenant_id,omitempty"`
Domain string `json:"domain,omitempty"`
Mailbox string `json:"mailbox,omitempty"`
// SprayTargets carries the per-account targets for aggregate auth
// findings. It is internal-only so API payloads keep the public Finding
// contract while the incident correlator can count distinct targets.
SprayTargets []string `json:"-"`
// CIDRs carries the collapsed offending subnets for subnet-scoped
// findings (http_asn_crawl). Internal-only so the public Finding/webhook
// contract is unchanged; the subnet auto-response reads this, never the
// Message/Details text.
CIDRs []string `json:"-"`
// Process context (Phase 1 process-ancestry enrichment). Optional.
// Populated by exec/connection live monitors when cache or enricher
// has data. Omitted from JSON when nil so existing webhook consumers
// see no diff.
Process *processctx.ProcessContext `json:"process,omitempty"`
Timestamp time.Time `json:"timestamp"`
// FirstSeen is when this condition was first observed, as opposed to
// when it was last reported. The latest-state merge carries it across
// re-reports; a scan that finds the same condition again refreshes
// Timestamp but not this. Correlation reads it so a months-old finding
// re-emitted by every scan cannot keep re-entering a recent-activity
// window. Zero on findings that never went through the merge, and on
// rows stored before the field existed; callers fall back to Timestamp.
FirstSeen time.Time `json:"first_seen,omitzero"`
// Full-scan quarantine outcome (Phase 2). Set ONLY on findings produced by a
// `--full --quarantine` job; empty for all report-only findings so existing
// consumers see no JSON diff.
RemediationStatus string `json:"remediation_status,omitempty"` // "quarantined" | "cleaned" | "left_for_review" | "failed"
RemediationDetail string `json:"remediation_detail,omitempty"` // action description or error
}
// RelayScriptHit is one script's contribution to a PHP-relay finding.
type RelayScriptHit struct {
ScriptKey string `json:"script_key"` // "host:/path" from X-PHP-Script
Hits int `json:"hits"` // messages counted in the path window
LastSeen time.Time `json:"last_seen"`
SampleSubject string `json:"sample_subject,omitempty"` // attacker-controlled; render escaped
}
func (f Finding) String() string {
ts := f.Timestamp.Format("2006-01-02 15:04:05")
s := fmt.Sprintf("[%s] %s - %s", f.Severity, f.Check, f.Message)
if f.Details != "" {
s += "\n " + strings.ReplaceAll(f.Details, "\n", "\n ")
}
if f.ProcessInfo != "" {
s += fmt.Sprintf("\n Process: %s", f.ProcessInfo)
}
s += fmt.Sprintf("\n Time: %s", ts)
return s
}
// Key returns a unique key for deduplication.
func (f Finding) Key() string {
if f.DedupKey != "" {
// Keep explicit identities in a leading namespace. Putting the marker
// after Check would still collide with an ordinary finding from that
// check whose Message happens to start with "dedup:".
return fmt.Sprintf("dedup:%s:%s", f.Check, f.DedupKey)
}
if key := f.sourceIPKey(); key != "" {
return key
}
if f.Details == "" {
return fmt.Sprintf("%s:%s", f.Check, f.Message)
}
h := sha256.Sum256([]byte(f.Details))
return fmt.Sprintf("%s:%s:%x", f.Check, f.Message, h[:4])
}
// Fingerprint returns the content hash used by alert-state deduplication.
func (f Finding) Fingerprint() string {
if f.DedupKey != "" {
h := sha256.Sum256([]byte(f.Key()))
return fmt.Sprintf("%x", h[:8])
}
if key := f.sourceIPKey(); key != "" {
h := sha256.Sum256([]byte(key))
return fmt.Sprintf("%x", h[:8])
}
h := sha256.Sum256([]byte(fmt.Sprintf("%s:%s:%s", f.Check, f.Message, f.Details)))
return fmt.Sprintf("%x", h[:8])
}
func (f Finding) sourceIPKey() string {
severityScoped := false
switch f.Check {
case "admin_panel_bruteforce", "wp_login_bruteforce", "wp_user_enumeration", "xmlrpc_abuse",
"http_request_flood", "http_scanner_profile", "http_claimed_bot_unverified", "http_ua_spoof",
"ftp_bruteforce":
case "php_shield_webshell", "php_shield_block", "php_shield_eval":
// One scanner sweeping a shared host hits every account in the same
// second, which produced one alert per site instead of one per scanner.
// The collapse is per severity: the same check name carries both the
// Warning "observed a parameter" and the Critical "blocked a webshell",
// and a block from an address that was merely observed earlier is a new
// event, not a repeat, or it is never alerted, stored or correlated.
severityScoped = true
default:
return ""
}
ip := normalizeFindingIP(f.SourceIP)
if ip == "" {
ip = sourceIPFromFindingMessage(f.Message)
}
if ip == "" {
return ""
}
if severityScoped {
return fmt.Sprintf("%s:ip:%s:%s", f.Check, ip, f.Severity)
}
return fmt.Sprintf("%s:ip:%s", f.Check, ip)
}
func normalizeFindingIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if host, _, err := net.SplitHostPort(raw); err == nil {
raw = host
}
raw = strings.Trim(raw, "[]")
ip := net.ParseIP(raw)
if ip == nil {
return ""
}
return ip.String()
}
func sourceIPFromFindingMessage(msg string) string {
for _, sep := range []string{" from ", ": "} {
idx := strings.LastIndex(msg, sep)
if idx < 0 {
continue
}
rest := msg[idx+len(sep):]
fields := strings.Fields(rest)
if len(fields) == 0 {
continue
}
candidate := strings.TrimRight(fields[0], ",:;)([]")
if ip := normalizeFindingIP(candidate); ip != "" {
return ip
}
}
return ""
}
// SplitEmail returns (localpart, domain) from an email address. Returns
// ("", "") when the input doesn't look like an email.
func SplitEmail(addr string) (localpart, domain string) {
at := strings.LastIndexByte(addr, '@')
if at <= 0 || at == len(addr)-1 {
return "", ""
}
return addr[:at], addr[at+1:]
}
// Deduplicate removes findings with the same Check+Message, keeping the first.
func Deduplicate(findings []Finding) []Finding {
seen := make(map[string]bool)
var result []Finding
for _, f := range findings {
key := f.Key()
if !seen[key] {
seen[key] = true
result = append(result, f)
}
}
return result
}
// FormatAlert formats a list of findings into a human-readable alert body.
// Sensitive data (passwords, tokens) is redacted before sending.
func FormatAlert(hostname string, findings []Finding) string {
var b strings.Builder
critCount := 0
highCount := 0
warnCount := 0
for _, f := range findings {
switch f.Severity {
case Critical:
critCount++
case High:
highCount++
case Warning:
warnCount++
}
}
fmt.Fprintf(&b, "SECURITY ALERT - %s\n", hostname)
fmt.Fprintf(&b, "Timestamp: %s\n", time.Now().Format("2006-01-02 15:04:05 MST"))
fmt.Fprintf(&b, "Findings: %d critical, %d high, %d warning\n", critCount, highCount, warnCount)
b.WriteString(strings.Repeat("─", 60) + "\n\n")
for _, sev := range []Severity{Critical, High, Warning} {
for _, f := range findings {
if f.Severity == sev {
b.WriteString(SanitizeFinding(f).String())
b.WriteString("\n\n")
}
}
}
b.WriteString(strings.Repeat("─", 60) + "\n")
b.WriteString("CSM - Continuous Security Monitor\n")
return b.String()
}
// SanitizeFinding returns a copy with recognized credentials redacted from
// Message and Details. Call it at output and persistence boundaries so detection
// and identity calculations can still use the original finding. Other fields
// are unchanged; nested data is not modified.
func SanitizeFinding(f Finding) Finding {
f.Message = redactSensitive(f.Message)
f.Details = redactSensitive(f.Details)
return f
}
// redactSensitive replaces password values and tokens in text with [REDACTED].
func redactSensitive(s string) string {
if s == "" {
return s
}
s = redactCredentialFields(s)
// Normalize command-line text first: NUL-delimited arguments can expose
// session keywords once the argument separators become spaces.
s = RedactCommandLine(s)
// Gate each log line separately: a finding can also contain unrelated
// prose with NEW/PURGE and colons whose evidence must survive.
var lines strings.Builder
for line := range strings.SplitAfterSeq(s, "\n") {
lines.WriteString(redactSessionLogLine(line))
}
return lines.String()
}
// Scan the original text once so every field is covered without searching
// replacement markers or changing byte offsets. Log envelopes can quote a whole
// request, so command-line tokenization alone cannot find these nested fields.
func redactCredentialFields(s string) string {
lower := lowerASCII(s)
var b strings.Builder
last := 0
for i := 0; i < len(s); {
prefixLen := 0
for _, prefix := range []string{
"password=", "pass=", "passwd=", "new_password=",
"old_password=", "confirmpassword=", "token_value=", "api_token=",
} {
if strings.HasPrefix(lower[i:], prefix) {
prefixLen = len(prefix)
break
}
}
if prefixLen == 0 {
i++
continue
}
start := i + prefixLen
end := start
var quote byte
if end < len(s) && (s[end] == '\'' || s[end] == '"') {
quote = s[end]
end++
}
for end < len(s) {
c := s[end]
if c == '\\' && end+1 < len(s) {
end += 2
continue
}
if quote != 0 {
end++
if c == quote {
quote = 0
}
continue
}
if strings.ContainsRune(" &\t\n\r\v\f\x00\"',", rune(c)) {
break
}
end++
}
if end > start && s[start:end] != redactedToken {
b.WriteString(s[last:start])
b.WriteString(redactedToken)
last = end
}
i = end
}
if last == 0 {
return s
}
b.WriteString(s[last:])
return b.String()
}
// Credential names are ASCII. Unicode case folding can change byte lengths
// (including invalid UTF-8 from logs), invalidating offsets into the input.
func lowerASCII(s string) string {
b := []byte(s)
for i, c := range b {
if c >= 'A' && c <= 'Z' {
b[i] = c + ('a' - 'A')
}
}
return string(b)
}
func redactSessionLogLine(s string) string {
if !containsSessionLogTag(s) {
return s
}
// Scan the original text only, keeping replacements out of the search.
// Account names survive; only the credential after the colon is masked.
var b strings.Builder
last := 0
for i := 0; i < len(s); i++ {
keywordLen := 0
switch {
case strings.HasPrefix(s[i:], " NEW "):
keywordLen = len(" NEW ")
case strings.HasPrefix(s[i:], " PURGE "):
keywordLen = len(" PURGE ")
default:
continue
}
fieldStart := i + keywordLen
fieldEnd := fieldStart
for fieldEnd < len(s) && !strings.ContainsRune(" \t\n\r", rune(s[fieldEnd])) {
fieldEnd++
}
colon := strings.IndexByte(s[fieldStart:fieldEnd], ':')
// Keep the keyword's trailing space searchable: the malformed field
// may itself be NEW or PURGE, followed by a valid account:session.
i = fieldStart - 2
if colon < 0 {
continue
}
tokenStart := fieldStart + colon + 1
if tokenStart == fieldEnd {
continue
}
if s[tokenStart:fieldEnd] != redactedToken {
b.WriteString(s[last:tokenStart])
b.WriteString(redactedToken)
last = fieldEnd
}
i = fieldEnd - 1
}
if last == 0 {
return s
}
b.WriteString(s[last:])
return b.String()
}
// cPanel session logs use both frontend service names and the shared server
// daemon name. DAV and security purge logs also carry account:session pairs.
func containsSessionLogTag(s string) bool {
for _, tag := range []string{"[cpaneld]", "[webmaild]", "[whostmgr]", "[whostmgrd]", "[cpsrvd]", "[cpdavd]", "[security]"} {
if strings.Contains(s, tag) {
return true
}
}
return false
}
func filterChecks(findings []Finding, disabledChecks []string) []Finding {
if len(findings) == 0 || len(disabledChecks) == 0 {
return findings
}
disabled := make(map[string]bool, len(disabledChecks))
for _, check := range disabledChecks {
check = config.CanonicalCheckName(strings.TrimSpace(check))
if check != "" {
disabled[check] = true
}
}
if len(disabled) == 0 {
return findings
}
filtered := make([]Finding, 0, len(findings))
for _, f := range findings {
if !disabled[config.CanonicalCheckName(f.Check)] {
filtered = append(filtered, f)
}
}
return filtered
}
func buildSubject(hostname string, findings []Finding) string {
subject := fmt.Sprintf("[CSM] %s - %d security finding(s)", hostname, len(findings))
for _, f := range findings {
if f.Severity == Critical {
return fmt.Sprintf("[CSM] CRITICAL - %s - %d finding(s)", hostname, len(findings))
}
}
return subject
}
// rateLimitState tracks alerts sent per hour.
type rateLimitState struct {
Hour string `json:"hour"`
Count int `json:"count"`
}
// FindingBus is set by the daemon at startup to the broadcast.Bus that
// passive observers (e.g. SSE subscribers) drain. nil-safe: Dispatch
// only publishes if non-nil. Importing the broadcast package directly
// would create an import cycle (broadcast imports alert), so this is
// declared as an interface satisfied by *broadcast.Bus.
var FindingBus interface {
Publish(Finding)
}
// ReportHook, when set by the daemon at startup, is called once per
// deduplicated finding so the abuse reporter can consider it for submission to
// a central abuse database or collector. It must not block. Declared as a func
// to avoid an import cycle (the reporting package imports alert for the
// Finding type).
//
// Install or clear it with SetReportHook so Dispatch reads a consistent value.
var ReportHook func(Finding)
var reportHookMu sync.RWMutex
// SetReportHook installs or clears the abuse-reporting hook used by Dispatch.
func SetReportHook(h func(Finding)) {
reportHookMu.Lock()
ReportHook = h
reportHookMu.Unlock()
}
func currentReportHook() func(Finding) {
reportHookMu.RLock()
h := ReportHook
reportHookMu.RUnlock()
return h
}
func callReportHook(f Finding) {
h := currentReportHook()
if h == nil {
return
}
defer func() {
if r := recover(); r != nil {
fmt.Fprintln(os.Stderr, "alert: report hook panic")
}
}()
h(f)
}
// CentralHook, when set by the daemon, is called once per deduplicated finding
// so the central-intel consumer can escalate (challenge/block) when the
// finding's IP is in the verified central scored-set. A finding firing on an IP
// is the node's own local signal, so this is the local-corroboration path.
// Must not block. Install or clear with SetCentralHook.
var CentralHook func(Finding)
var centralHookMu sync.RWMutex
// SetCentralHook installs or clears the central-intel hook used by Dispatch.
func SetCentralHook(h func(Finding)) {
centralHookMu.Lock()
CentralHook = h
centralHookMu.Unlock()
}
func currentCentralHook() func(Finding) {
centralHookMu.RLock()
h := CentralHook
centralHookMu.RUnlock()
return h
}
func callCentralHook(f Finding) {
h := currentCentralHook()
if h == nil {
return
}
defer func() {
if r := recover(); r != nil {
fmt.Fprintln(os.Stderr, "alert: central hook panic")
}
}()
h(f)
}
type rateLimitKey struct {
StatePath string
Hour string
}
type rateLimitReservation struct {
key rateLimitKey
active bool
}
var (
rateLimitMu sync.Mutex
rateLimitPending = make(map[rateLimitKey]int)
)
// reserveRateLimit reports whether the per-hour alert budget can absorb
// another send without committing the slot. The in-memory reservation
// prevents concurrent dispatches from all taking the same final slot while
// the outbound channel is still blocked on SMTP or webhook I/O.
func reserveRateLimit(statePath string, maxPerHour int) (*rateLimitReservation, bool) {
rateLimitMu.Lock()
defer rateLimitMu.Unlock()
if maxPerHour <= 0 {
return nil, false
}
currentHour := time.Now().Format("2006-01-02T15")
rlPath := filepath.Join(statePath, "ratelimit.json")
var rl rateLimitState
// #nosec G304 -- filepath.Join(statePath, "ratelimit.json"); statePath from operator config.
data, err := os.ReadFile(rlPath)
if err == nil {
_ = json.Unmarshal(data, &rl)
}
count := 0
if rl.Hour != currentHour {
count = 0
} else {
count = rl.Count
}
key := rateLimitKey{StatePath: statePath, Hour: currentHour}
if count+rateLimitPending[key] >= maxPerHour {
return nil, false
}
rateLimitPending[key]++
return &rateLimitReservation{key: key, active: true}, true
}
func releaseRateLimit(reservation *rateLimitReservation) {
if reservation == nil {
return
}
rateLimitMu.Lock()
defer rateLimitMu.Unlock()
releaseRateLimitLocked(reservation)
}
func releaseRateLimitLocked(reservation *rateLimitReservation) {
if reservation == nil || !reservation.active {
return
}
if pending := rateLimitPending[reservation.key]; pending <= 1 {
delete(rateLimitPending, reservation.key)
} else {
rateLimitPending[reservation.key] = pending - 1
}
reservation.active = false
}
// commitRateLimit records one successful dispatch toward the hourly
// budget. Logs the WriteFile error so a disk-full or perm regression
// surfaces in the daemon log instead of silently letting the counter
// drift from the on-disk record.
func commitRateLimit(statePath string, reservation *rateLimitReservation) {
rateLimitMu.Lock()
defer rateLimitMu.Unlock()
releaseRateLimitLocked(reservation)
currentHour := time.Now().Format("2006-01-02T15")
rlPath := filepath.Join(statePath, "ratelimit.json")
var rl rateLimitState
// #nosec G304 -- filepath.Join under operator-configured statePath.
if data, err := os.ReadFile(rlPath); err == nil {
_ = json.Unmarshal(data, &rl)
}
if rl.Hour != currentHour {
rl = rateLimitState{Hour: currentHour}
}
rl.Count++
newData, err := json.Marshal(rl)
if err != nil {
fmt.Fprintf(os.Stderr, "alert: rate-limit marshal failed: %v\n", err)
return
}
if err := os.WriteFile(rlPath, newData, 0600); err != nil {
fmt.Fprintf(os.Stderr, "alert: rate-limit write failed for %s: %v\n", rlPath, err)
}
}
// checkRateLimit returns true if we can send more alerts this hour.
//
// Deprecated: kept for callers that expect the check-and-increment pattern.
// New code should use reserveRateLimit and commitRateLimit so a failed
// dispatch does not consume the operator's budget.
func checkRateLimit(statePath string, maxPerHour int) bool {
rateLimitMu.Lock()
defer rateLimitMu.Unlock()
rlPath := filepath.Join(statePath, "ratelimit.json")
currentHour := time.Now().Format("2006-01-02T15")
var rl rateLimitState
// #nosec G304 -- filepath.Join(statePath, "ratelimit.json"); statePath from operator config.
data, err := os.ReadFile(rlPath)
if err == nil {
_ = json.Unmarshal(data, &rl)
}
// Reset if new hour
if rl.Hour != currentHour {
rl = rateLimitState{Hour: currentHour, Count: 0}
}
if rl.Count >= maxPerHour {
return false
}
rl.Count++
newData, _ := json.Marshal(rl)
_ = os.WriteFile(rlPath, newData, 0600)
return true
}
func formatDispatchErrors(errs []error) error {
if len(errs) == 0 {
return nil
}
msgs := make([]string, len(errs))
for i, e := range errs {
msgs[i] = e.Error()
}
return fmt.Errorf("alert dispatch errors: %s", strings.Join(msgs, "; "))
}
// FillTimestamps stamps now on every finding that carries no Timestamp.
// Realtime producers build findings without one; a zero time sorts before
// every real event in a SIEM and makes the audit finding id collide across
// occurrences, so each sink boundary fills it in.
func FillTimestamps(findings []Finding, now time.Time) {
for i := range findings {
if findings[i].Timestamp.IsZero() {
findings[i].Timestamp = now
}
}
}
// Dispatch sends alerts via all configured channels without modifying findings.
func Dispatch(cfg *config.Config, findings []Finding) error {
return DispatchWithSources(cfg, findings, nil)
}
// DispatchWithSources audits source observations even when notification policy
// filters them out. Only findings reach notification channels and observers;
// sources add audit records without changing alert or auto-response policy.
// Both inputs remain caller-owned and must already carry the times used by
// actions that reference them. Missing times are filled on copies for ad-hoc use.
func DispatchWithSources(cfg *config.Config, findings, sources []Finding) error {
return dispatchWithSources(cfg, findings, sources, findings, nil)
}
// DispatchWithEnforcement offers central IP enforcement its own finding set
// instead of the notification set. Suppression rules mute notifications but
// must not exempt an attacker from central challenges and blocks.
func DispatchWithEnforcement(cfg *config.Config, findings, sources, enforcement []Finding) error {
return dispatchWithSources(cfg, findings, sources, enforcement, nil)
}
// DispatchWithNotificationFilter keeps operator policy separate from the
// phpanel and SSE data streams. Notification observers, audit sources and
// central enforcement retain their independent policy boundaries.
func DispatchWithNotificationFilter(cfg *config.Config, findings, sources, enforcement []Finding, filter func([]Finding) []Finding) error {
return dispatchWithSources(cfg, findings, sources, enforcement, filter)
}
func dispatchWithSources(cfg *config.Config, findings, sources, enforcement []Finding, notificationFilter func([]Finding) []Finding) error {
// Deduplicate owns a copy, so stamping cannot race with callers sharing
// the input or pin a reused unstamped finding to its first dispatch time.
findings = Deduplicate(findings)
now := auditNow()
FillTimestamps(findings, now)
sources = append([]Finding(nil), sources...)
FillTimestamps(sources, now)
notifications := findings
if notificationFilter != nil {
notifications = notificationFilter(findings)
// Audit every observation, while registered notification observers
// retain the same suppression policy as operator notifications.
sources = append(sources, findings...)
}
// Audit log captures every (deduplicated) finding before
// FilterBlockedAlerts and the rate limiter, so SIEMs see the
// complete picture even when email/webhook are throttled or
// when "this IP is already blocked" suppression hides a finding
// from the operator-facing channels.
emitAuditWithSources(cfg, notifications, sources)
// The central-intel consumer escalates findings whose IP is in the
// verified central scored-set.
enforcement = Deduplicate(enforcement)
FillTimestamps(enforcement, now)
for _, f := range enforcement {
callCentralHook(f)
}
if len(findings) == 0 {
return nil
}
// Publish to passive observers (e.g. SSE subscribers) immediately after
// auditing, before rate-limit and webhook delivery, so subscribers see
// the complete picture even when operator-facing channels are throttled.
if FindingBus != nil {
for _, f := range findings {
FindingBus.Publish(f)
}
}
// Offer every finding to the abuse reporter (it gates and minimizes
// internally, queueing only confirmed-abuse findings for the drain loop).
for _, f := range notifications {
callReportHook(f)
}
var errs []error
// Phpanel consumes this webhook as a signed data-plane stream. Send the
// full deduplicated stream before operator notification suppression and
// rate limiting, otherwise fleet correlation can miss attacker spread.
phpanelWebhook := cfg.Alerts.Webhook.Enabled && cfg.Alerts.Webhook.Type == "phpanel"
if phpanelWebhook {
if err := enqueuePhpanelFindings(cfg, findings); err != nil {
addDispatchError(&errs, err)
}
}
findings = notifications
// Filter out blocked IP alerts if configured
findings = FilterBlockedAlerts(cfg, findings)
if len(findings) == 0 {
return formatDispatchErrors(errs)
}
emailFindings := []Finding(nil)
if cfg.Alerts.Email.Enabled {
emailFindings = filterChecks(findings, cfg.Alerts.Email.DisabledChecks)
}
webhookFindings := []Finding(nil)
if cfg.Alerts.Webhook.Enabled && !phpanelWebhook {
webhookFindings = findings
}
// Only routine findings spend the hourly budget. Urgent ones always go
// out, but they never carry routine findings from the same batch past the
// cap with them.
var reservation *rateLimitReservation
if hasRoutineFinding(emailFindings) || hasRoutineFinding(webhookFindings) {
var ok bool
reservation, ok = reserveRateLimit(cfg.StatePath, cfg.Alerts.MaxPerHour)
if ok {
defer releaseRateLimit(reservation)
} else {
fmt.Fprintf(os.Stderr, "Alert rate limit reached (%d/hour), skipping non-critical alert dispatch\n", cfg.Alerts.MaxPerHour)
emailFindings = urgentFindings(emailFindings)
webhookFindings = urgentFindings(webhookFindings)
}
}
if len(emailFindings) == 0 && len(webhookFindings) == 0 {
return formatDispatchErrors(errs)
}
routineDispatched := false
if len(emailFindings) > 0 {
subject := buildSubject(cfg.Hostname, emailFindings)
body := FormatAlert(cfg.Hostname, emailFindings)
if err := SendEmail(cfg, subject, body); err != nil {
addDispatchError(&errs, fmt.Errorf("email: %w", err))
} else if hasRoutineFinding(emailFindings) {
routineDispatched = true
}
}
if len(webhookFindings) > 0 {
subject := buildSubject(cfg.Hostname, webhookFindings)
body := FormatAlert(cfg.Hostname, webhookFindings)
if err := SendWebhook(cfg, subject, body); err != nil {
addDispatchError(&errs, fmt.Errorf("webhook: %w", err))
} else if hasRoutineFinding(webhookFindings) {
routineDispatched = true
}
}
// An urgent-only email can succeed while the webhook carrying the routine
// findings fails. Spend the slot only if routine findings were delivered.
if routineDispatched && reservation != nil {
commitRateLimit(cfg.StatePath, reservation)
}
return formatDispatchErrors(errs)
}
// bypassesRateLimit reports whether a finding is delivered regardless of the
// hourly budget. Reputation delivery is check-keyed: its surface-based severity
// is presentation metadata and must not make sightings that previously
// bypassed the budget disappear.
func bypassesRateLimit(f Finding) bool {
return f.Severity == Critical || f.Check == "ip_reputation"
}
func hasRoutineFinding(findings []Finding) bool {
for _, f := range findings {
if !bypassesRateLimit(f) {
return true
}
}
return false
}
func urgentFindings(findings []Finding) []Finding {
var out []Finding
for _, f := range findings {
if bypassesRateLimit(f) {
out = append(out, f)
}
}
return out
}
// SendHeartbeat pings a dead man's switch URL.
func SendHeartbeat(cfg *config.Config) {
if !cfg.Alerts.Heartbeat.Enabled || cfg.Alerts.Heartbeat.URL == "" {
return
}
client := httpClient(10 * time.Second)
resp, err := client.Get(cfg.Alerts.Heartbeat.URL)
if err != nil {
fmt.Fprintf(os.Stderr, "Heartbeat failed: %v\n", err)
return
}
defer closeWebhookResponseBody(resp)
}
package alert
import (
"container/list"
"crypto/sha256"
"encoding/hex"
"fmt"
"os"
"runtime/debug"
"strconv"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/config"
)
// Audit-log dispatching is layered on top of the existing email +
// webhook fork in Dispatch(). Sinks live behind a package-level
// manager so the daemon's per-call Dispatch path does not pay the
// cost of opening a JSONL file or dialling syslog on every alert.
//
// The manager keys its sink set by a fingerprint of the relevant
// config sub-block; on hot reload the fingerprint changes and the
// manager closes the old sinks and rebuilds.
var (
auditMu sync.Mutex
auditSinks []*managedAuditSink
auditFingerprint string
auditNow = time.Now
openJSONLAuditSink = func(path string) (AuditSink, error) { return NewJSONLSink(path) }
openSyslogAuditSink = func(cfg SyslogConfig) (AuditSink, error) { return NewSyslogSink(cfg) }
)
type managedAuditSink struct {
name string
open func() (AuditSink, error)
sink AuditSink
retryAt time.Time
retryDelay time.Duration
failed bool
// Keep successful observation IDs across batches and transient sink
// failures. Each destination owns its receipts so a failed destination
// can receive a replay without duplicating the healthy destination.
delivered map[[sha256.Size]byte]*list.Element
recent list.List
}
// Bound replay receipts independently of traffic volume. Evicted observations
// can be emitted again; distinct observations are never dropped by this cache.
const auditReceiptCap = 16384
func (s *managedAuditSink) remember(id [sha256.Size]byte) {
if s.delivered == nil {
s.delivered = make(map[[sha256.Size]byte]*list.Element)
}
if s.recent.Len() == auditReceiptCap {
oldest := s.recent.Back()
delete(s.delivered, oldest.Value.([sha256.Size]byte))
s.recent.Remove(oldest)
}
s.delivered[id] = s.recent.PushFront(id)
}
// auditObservationKey distinguishes original evidence that the legacy action
// ID omits. For example, a process scan can report several PIDs with the same
// message and timestamp. Keep action IDs stable and hash scalar source evidence
// before redaction, without retaining its raw text or mutable enrichment.
func auditObservationKey(f Finding) [sha256.Size]byte {
var data []byte
for _, field := range []string{
f.Timestamp.UTC().Format(time.RFC3339Nano), f.Check, f.Severity.String(),
f.Message, f.FilePath, f.Details, strconv.Itoa(f.PID), f.SourceIP,
f.TenantID, f.Domain, f.Mailbox, f.DedupKey, f.CoverageScope,
} {
// Quote preserves boundaries and invalid UTF-8 from raw log input.
data = strconv.AppendQuote(data, field)
}
return sha256.Sum256(data)
}
// emitAudit records findings before email/webhook throttling. The manager lock
// covers the whole batch, including reconfiguration, so Close cannot invalidate
// a sink held by another dispatcher. Observers run outside the lock because
// they can dispatch findings themselves.
func emitAudit(cfg *config.Config, findings []Finding) {
emitAuditWithSources(cfg, findings, nil)
}
func emitAuditWithSources(cfg *config.Config, findings, sources []Finding) {
if cfg == nil {
return
}
for _, f := range findings {
notifyFindingObservers(f)
}
if len(sources) > 0 {
// Notification dedup uses condition keys. Audit joins need every
// distinct observation, including repeats with a new timestamp.
combined := make([]Finding, 0, len(sources)+len(findings))
seen := make(map[[sha256.Size]byte]bool, len(sources)+len(findings))
for _, batch := range [][]Finding{sources, findings} {
for _, f := range batch {
id := auditObservationKey(f)
if !seen[id] {
seen[id] = true
combined = append(combined, f)
}
}
}
findings = combined
}
auditMu.Lock()
defer auditMu.Unlock()
ensureAuditSinksLocked(cfg)
for _, f := range findings {
ev := NewAuditEvent(cfg.Hostname, f)
id := auditObservationKey(f)
for _, s := range auditSinks {
if receipt, delivered := s.delivered[id]; delivered {
// Retained findings can be replayed every scan amid realtime
// churn. Keep their receipts hot without growing the cache.
s.recent.MoveToFront(receipt)
continue
}
if s.sink == nil {
auditEventsDropped.With(s.name).Inc()
continue
}
if err := s.sink.Emit(ev); err != nil {
auditEventsDropped.With(s.name).Inc()
_ = s.sink.Close()
s.sink = nil
s.fail("emit", err)
} else {
s.retryDelay = 0
s.remember(id)
}
}
}
}
func ensureAuditSinks(cfg *config.Config) {
auditMu.Lock()
defer auditMu.Unlock()
ensureAuditSinksLocked(cfg)
}
func ensureAuditSinksLocked(cfg *config.Config) {
fp := auditConfigFingerprint(cfg)
if fp != auditFingerprint {
closeAuditSinksLocked()
if cfg.Alerts.AuditLog.File.Enabled {
path := cfg.Alerts.AuditLog.File.Path
auditSinks = append(auditSinks, &managedAuditSink{name: "jsonl", open: func() (AuditSink, error) { return openJSONLAuditSink(path) }})
}
if cfg.Alerts.AuditLog.Syslog.Enabled {
sc := SyslogConfig{
Network: cfg.Alerts.AuditLog.Syslog.Network,
Address: cfg.Alerts.AuditLog.Syslog.Address,
Facility: cfg.Alerts.AuditLog.Syslog.Facility,
Hostname: cfg.Hostname,
TLSCAFile: cfg.Alerts.AuditLog.Syslog.TLSCAFile,
}
auditSinks = append(auditSinks, &managedAuditSink{name: "syslog", open: func() (AuditSink, error) { return openSyslogAuditSink(sc) }})
}
auditFingerprint = fp
}
for _, s := range auditSinks {
if s.sink != nil || auditNow().Before(s.retryAt) {
continue
}
sink, err := s.open()
if err != nil {
s.fail("init", err)
continue
}
s.sink = sink
auditSinkDegraded.With(s.name).Set(0)
if s.failed {
fmt.Fprintf(os.Stderr, "[audit-log] %s recovered\n", s.name)
}
s.failed = false
}
}
func (s *managedAuditSink) fail(phase string, err error) {
if s.retryDelay == 0 {
s.retryDelay = time.Second
} else {
s.retryDelay = min(2*s.retryDelay, time.Minute)
}
s.retryAt = auditNow().Add(s.retryDelay)
s.failed = true
auditSinkDegraded.With(s.name).Set(1)
fmt.Fprintf(os.Stderr, "[audit-log] %s %s failed; retry after %s: %v\n", s.name, phase, s.retryDelay, err)
}
// auditConfigFingerprint reduces the audit-log sub-block to a stable
// hash so ensureAuditSinks can detect config changes without a deep
// reflect-based diff. Hostname is included because it appears in
// every emitted event.
func auditConfigFingerprint(cfg *config.Config) string {
h := sha256.New()
_, _ = fmt.Fprintf(h, "host=%s|", cfg.Hostname)
_, _ = fmt.Fprintf(h, "file.enabled=%t|file.path=%s|",
cfg.Alerts.AuditLog.File.Enabled,
cfg.Alerts.AuditLog.File.Path,
)
_, _ = fmt.Fprintf(h, "syslog.enabled=%t|syslog.network=%s|syslog.address=%s|syslog.facility=%s|syslog.tls=%s",
cfg.Alerts.AuditLog.Syslog.Enabled,
cfg.Alerts.AuditLog.Syslog.Network,
cfg.Alerts.AuditLog.Syslog.Address,
cfg.Alerts.AuditLog.Syslog.Facility,
cfg.Alerts.AuditLog.Syslog.TLSCAFile,
)
return hex.EncodeToString(h.Sum(nil))
}
// CloseAuditSinks waits for in-flight emissions and releases active sinks.
// A later dispatch can initialize them again.
func CloseAuditSinks() {
auditMu.Lock()
defer auditMu.Unlock()
closeAuditSinksLocked()
}
func closeAuditSinksLocked() {
for _, s := range auditSinks {
if s.sink != nil {
_ = s.sink.Close()
}
auditSinkDegraded.With(s.name).Set(0)
}
auditSinks = nil
auditFingerprint = ""
}
// resetAuditSinksForTest is the test-only seam to wipe the package
// state between cases. Production code never needs this -- live
// daemons run a single Dispatch path with a single config object.
func resetAuditSinksForTest() {
CloseAuditSinks()
}
// findingObservers registry. Used by the daemon to feed the incident
// correlator without making the alert package depend on internal/incident.
var (
findingObserversMu sync.RWMutex
findingObservers []findingObserver
findingObserverSeq atomic.Uint64
)
type findingObserver struct {
id uint64
fn func(Finding)
}
// RegisterFindingObserver registers fn to be called for every finding
// dispatched through emitAudit. Returns a cancel func that removes the
// observer. Safe for concurrent use; observer panics are recovered so
// one bad observer cannot stop dispatch.
func RegisterFindingObserver(fn func(Finding)) func() {
id := findingObserverSeq.Add(1)
findingObserversMu.Lock()
findingObservers = append(findingObservers, findingObserver{id: id, fn: fn})
findingObserversMu.Unlock()
return func() {
findingObserversMu.Lock()
defer findingObserversMu.Unlock()
out := findingObservers[:0]
for _, o := range findingObservers {
if o.id != id {
out = append(out, o)
}
}
findingObservers = out
}
}
// notifyFindingObservers fans a finding out to every registered observer.
// Each observer runs in a recover scope so a panic in one cannot stop
// dispatch to the rest, the audit-log sinks, or future ones.
func notifyFindingObservers(f Finding) {
findingObserversMu.RLock()
obs := append([]findingObserver(nil), findingObservers...)
findingObserversMu.RUnlock()
for _, o := range obs {
func(o findingObserver) {
defer func() {
if r := recover(); r != nil {
fmt.Fprintf(os.Stderr,
"alert: finding observer id=%d panic for check=%q: %s\n%s",
o.id, f.Check, formatRecoverValue(r), debug.Stack())
}
}()
o.fn(f)
}(o)
}
}
func formatRecoverValue(v any) (out string) {
defer func() {
if recover() != nil {
out = strconv.Quote(fmt.Sprintf("<unprintable panic value of type %T>", v))
}
}()
return strconv.QuoteToASCII(fmt.Sprint(v))
}
package alert
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"sync"
)
// JSONLSink appends one JSON object per finding to a file. Designed
// for SIEM ingest via standard log shippers (Vector, Filebeat,
// Fluentbit). Writes are mutex-serialised so concurrent emits cannot
// interleave bytes inside a single line; logrotate's copytruncate
// rotation works without daemon restart because the file's offset is
// reset by the truncate, which the open fd then writes past.
type JSONLSink struct {
path string
mu sync.Mutex
f *os.File
}
// NewJSONLSink opens (or creates) the JSONL file. Permissions are
// 0640 -- group-readable so an operator running a log shipper under
// a non-root user in the appropriate group can tail it. The parent
// directory is created with 0750 so packaging (logrotate) sees a
// reasonable default.
func NewJSONLSink(path string) (*JSONLSink, error) {
if path == "" {
return nil, errors.New("jsonl sink: path is empty")
}
if err := os.MkdirAll(filepath.Dir(path), 0750); err != nil {
return nil, fmt.Errorf("jsonl sink: creating dir: %w", err)
}
// #nosec G304 G302 -- G304: path comes from cfg.Alerts.AuditLog.File.Path, which is operator-controlled (root-owned daemon config), not attacker input. G302: 0640 is intentional; SIEM log shippers (Vector, Filebeat, Fluentbit) commonly run as a non-root user that needs group-read access. 0600 would force the shipper to run as root.
f, err := os.OpenFile(path, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0640)
if err != nil {
return nil, fmt.Errorf("jsonl sink: opening %s: %w", path, err)
}
return &JSONLSink{path: path, f: f}, nil
}
// Name returns the sink identifier used in error messages.
func (s *JSONLSink) Name() string { return "jsonl" }
// Emit appends one JSON line. The trailing newline is written as part
// of the same Write call so a partial write at EOL boundaries cannot
// leave a half-finished line in the file.
func (s *JSONLSink) Emit(event AuditEvent) error {
line, err := json.Marshal(event)
if err != nil {
return fmt.Errorf("jsonl sink: marshal: %w", err)
}
line = append(line, '\n')
s.mu.Lock()
defer s.mu.Unlock()
if s.f == nil {
return errors.New("jsonl sink: closed")
}
if _, err := s.f.Write(line); err != nil {
return fmt.Errorf("jsonl sink: write: %w", err)
}
return nil
}
// Close releases the file descriptor. Safe to call multiple times.
func (s *JSONLSink) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.f == nil {
return nil
}
err := s.f.Close()
s.f = nil
return err
}
package alert
import "github.com/pidginhost/csm/internal/metrics"
var (
auditSinkDegraded = metrics.NewGaugeVec("csm_audit_sink_degraded",
"Configured audit destination unavailable or last delivery failed; 1 means degraded, 0 means healthy or disabled.", []string{"sink"})
auditEventsDropped = metrics.NewCounterVec("csm_audit_events_dropped_total",
"Audit events with unavailable destinations or failed writes.", []string{"sink"})
)
func init() {
metrics.MustRegister("csm_audit_sink_degraded", auditSinkDegraded)
metrics.MustRegister("csm_audit_events_dropped_total", auditEventsDropped)
}
package alert
import (
"crypto/sha256"
"encoding/hex"
"time"
"github.com/pidginhost/csm/internal/processctx"
)
// AuditSchemaVersion is the value emitted in every AuditEvent's "v"
// field. Frozen contract -- downstream JSONL / syslog parsers pin on
// it. Bump only on incompatible schema changes; additive fields stay
// at the same version.
const AuditSchemaVersion = 1
// AuditEvent is the wire-stable shape every audit-log sink emits. JSON
// keys match what downstream SIEMs expect; fields are added at the
// end so older parsers ignore unknown ones.
type AuditEvent struct {
V int `json:"v"`
Timestamp time.Time `json:"ts"`
FindingID string `json:"finding_id"`
Severity string `json:"severity"`
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details,omitempty"`
FilePath string `json:"file_path,omitempty"`
Hostname string `json:"hostname"`
TenantID string `json:"tenant_id,omitempty"`
Domain string `json:"domain,omitempty"`
Mailbox string `json:"mailbox,omitempty"`
Process *processctx.ProcessContext `json:"process,omitempty"`
}
// AuditSink is what every audit-log destination implements. Emit must
// be safe for concurrent calls; the alert dispatcher fans out to
// multiple sinks per finding.
type AuditSink interface {
// Name identifies the sink for diagnostics (e.g. "jsonl", "syslog").
Name() string
// Emit ships one event. Should return promptly; sinks that need
// long-haul I/O are expected to handle their own buffering.
Emit(event AuditEvent) error
// Close releases any held resources. Safe to call multiple times.
Close() error
}
// NewAuditEvent builds a versioned audit event from a Finding. hostname
// comes from cfg.Hostname (or os.Hostname() fallback); the caller is
// responsible for picking a stable value across emits.
//
// The finding is redacted here because this is the one constructor both
// audit sinks build from. Redaction used to run only while rendering
// the email digest, so secrets a watcher had copied out of a raw log
// line -- cPanel session identifiers, password fields -- were written
// to audit.jsonl and shipped to syslog in the clear.
func NewAuditEvent(hostname string, f Finding) AuditEvent {
// Remediation records hash the original finding, so redaction must not
// change the ID used to join those records to this event.
id := FindingID(f)
f = SanitizeFinding(f)
return AuditEvent{
V: AuditSchemaVersion,
Timestamp: f.Timestamp.UTC(),
FindingID: id,
Severity: f.Severity.String(),
Check: f.Check,
Message: f.Message,
Details: f.Details,
FilePath: f.FilePath,
Hostname: hostname,
TenantID: f.TenantID,
Domain: f.Domain,
Mailbox: f.Mailbox,
Process: f.Process,
}
}
// FindingID hashes the canonical fields of a Finding to a stable
// 16-hex-char ID. Two emits of the same finding (same timestamp + the
// same other fields) produce the same ID, so downstream dedup works
// across re-runs.
//
// The hash inputs use a "|" separator so the byte-for-byte
// concatenation cannot collide via field-boundary ambiguity (e.g. a
// Check name that ends in the same chars another field starts with).
func FindingID(f Finding) string {
h := sha256.New()
_, _ = h.Write([]byte(f.Timestamp.UTC().Format(time.RFC3339Nano)))
_, _ = h.Write([]byte("|"))
_, _ = h.Write([]byte(f.Check))
_, _ = h.Write([]byte("|"))
_, _ = h.Write([]byte(f.Severity.String()))
_, _ = h.Write([]byte("|"))
_, _ = h.Write([]byte(f.Message))
_, _ = h.Write([]byte("|"))
_, _ = h.Write([]byte(f.FilePath))
return hex.EncodeToString(h.Sum(nil))[:16]
}
package alert
import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"fmt"
"net"
"os"
"strings"
"sync"
"time"
)
// SyslogConfig drives the syslog sink. Network is one of "udp",
// "tcp", "unix", "unixgram", or "tls"; Address is host:port (or a
// filesystem path for the unix variants). Facility names follow the
// classic syslog set; default is "local0".
type SyslogConfig struct {
Network string
Address string
Facility string
Hostname string // typically cfg.Hostname; falls back to os.Hostname()
TLSCAFile string // optional CA cert path for tls; empty = system roots
}
const (
syslogDialTimeout = 5 * time.Second
syslogWriteTimeout = 2 * time.Second
)
// SyslogSink is an RFC 5424 syslog client. The wire payload is the
// AuditEvent JSON so SIEMs that parse our JSONL file have a single
// schema regardless of transport. Messages bigger than the legacy
// 1024-byte limit are still sent -- modern receivers (rsyslog,
// syslog-ng) accept the full RFC 5424 max of 8192 bytes; if the
// operator's receiver caps shorter, syslog truncation is the
// expected behaviour.
type SyslogSink struct {
cfg SyslogConfig
priority int // pre-computed PRI byte; severity is OR'd in per emit
mu sync.Mutex
conn net.Conn
}
// facilityCodes is the standard syslog facility number set. local0..7
// is the customary range for application-level audit traffic; we
// default to local0 if the operator leaves the field blank.
var facilityCodes = map[string]int{
"kern": 0, "user": 1, "mail": 2, "daemon": 3, "auth": 4,
"syslog": 5, "lpr": 6, "news": 7, "uucp": 8, "cron": 9,
"authpriv": 10, "ftp": 11,
"local0": 16, "local1": 17, "local2": 18, "local3": 19,
"local4": 20, "local5": 21, "local6": 22, "local7": 23,
}
// NewSyslogSink validates the config, dials the destination, and
// returns a ready-to-emit sink. The dispatch manager reports dial failures
// and retries with backoff. Direct callers can retry a failed write with
// another Emit, which redials when the connection was dropped.
func NewSyslogSink(cfg SyslogConfig) (*SyslogSink, error) {
if cfg.Network == "" || cfg.Address == "" {
return nil, errors.New("syslog sink: network and address required")
}
switch cfg.Network {
case "udp", "tcp", "unix", "unixgram", "tls":
default:
return nil, fmt.Errorf("syslog sink: unknown network %q (want udp|tcp|unix|unixgram|tls)", cfg.Network)
}
facilityName := strings.ToLower(strings.TrimSpace(cfg.Facility))
if facilityName == "" {
facilityName = "local0"
}
facility, ok := facilityCodes[facilityName]
if !ok {
return nil, fmt.Errorf("syslog sink: unknown facility %q", cfg.Facility)
}
if cfg.Hostname == "" {
if h, err := os.Hostname(); err == nil {
cfg.Hostname = h
} else {
cfg.Hostname = "localhost"
}
}
s := &SyslogSink{cfg: cfg, priority: facility * 8}
if err := s.dial(); err != nil {
return nil, err
}
return s, nil
}
// Name identifies the sink in error messages.
func (s *SyslogSink) Name() string { return "syslog" }
// Emit formats the event as RFC 5424 and writes it to the
// destination. Mutex-serialised so concurrent calls do not interleave
// bytes on stream-oriented transports (TCP, TLS).
func (s *SyslogSink) Emit(event AuditEvent) error {
line, err := s.format(event)
if err != nil {
return err
}
s.mu.Lock()
defer s.mu.Unlock()
if s.conn == nil {
if redialErr := s.dialLocked(); redialErr != nil {
return redialErr
}
}
// Emit runs on the single alert-dispatch goroutine. A receiver that stops
// reading fills the socket buffer, and a write with no deadline would then
// hold this mutex forever and stall every alert on the host. Give up and
// drop the connection instead; the next Emit redials.
if err := s.conn.SetWriteDeadline(time.Now().Add(syslogWriteTimeout)); err != nil {
_ = s.conn.Close()
s.conn = nil
return fmt.Errorf("syslog sink: set write deadline: %w", err)
}
if _, err := s.conn.Write(line); err != nil {
_ = s.conn.Close()
s.conn = nil
return fmt.Errorf("syslog sink: write: %w", err)
}
return nil
}
// Close releases the connection.
func (s *SyslogSink) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if s.conn == nil {
return nil
}
err := s.conn.Close()
s.conn = nil
return err
}
func (s *SyslogSink) dial() error {
s.mu.Lock()
defer s.mu.Unlock()
return s.dialLocked()
}
// dialLocked establishes the connection. Caller must hold s.mu.
func (s *SyslogSink) dialLocked() error {
if s.cfg.Network == "tls" {
tlsCfg, err := s.tlsConfig()
if err != nil {
return err
}
// The context bounds the TLS handshake as well as the TCP connect; a
// blackholed receiver would otherwise hold the dispatch goroutine for
// the kernel's SYN retry budget on every redial.
ctx, cancel := context.WithTimeout(context.Background(), syslogDialTimeout)
defer cancel()
dialer := &tls.Dialer{NetDialer: &net.Dialer{Timeout: syslogDialTimeout}, Config: tlsCfg}
conn, err := dialer.DialContext(ctx, "tcp", s.cfg.Address)
if err != nil {
return fmt.Errorf("syslog sink: tls dial %s: %w", s.cfg.Address, err)
}
s.conn = conn
return nil
}
conn, err := net.DialTimeout(s.cfg.Network, s.cfg.Address, syslogDialTimeout)
if err != nil {
return fmt.Errorf("syslog sink: %s dial %s: %w", s.cfg.Network, s.cfg.Address, err)
}
s.conn = conn
return nil
}
func (s *SyslogSink) tlsConfig() (*tls.Config, error) {
cfg := &tls.Config{MinVersion: tls.VersionTLS12}
if s.cfg.TLSCAFile == "" {
return cfg, nil
}
// #nosec G304 -- TLSCAFile is operator-supplied via cfg.Alerts.AuditLog.Syslog.TLSCAFile; the operator owns the daemon config. Not attacker-controlled.
pem, err := os.ReadFile(s.cfg.TLSCAFile)
if err != nil {
return nil, fmt.Errorf("syslog sink: reading TLS CA: %w", err)
}
pool := x509.NewCertPool()
if !pool.AppendCertsFromPEM(pem) {
return nil, fmt.Errorf("syslog sink: TLS CA file %s is not a valid PEM bundle", s.cfg.TLSCAFile)
}
cfg.RootCAs = pool
return cfg, nil
}
// format produces an RFC 5424 line. The MSG body is the JSON-encoded
// event so receivers can parse the structured payload directly.
//
// <PRI>1 TIMESTAMP HOSTNAME APP-NAME PROCID MSGID - MSG
//
// PRI = facility * 8 + severity-as-syslog-level. STRUCTURED-DATA is
// "-" (we lift everything into the JSON body to avoid duplicate
// representation).
func (s *SyslogSink) format(event AuditEvent) ([]byte, error) {
body, err := json.Marshal(event)
if err != nil {
return nil, fmt.Errorf("syslog sink: marshal: %w", err)
}
pri := s.priority + severityToSyslogLevel(event.Severity)
ts := event.Timestamp.UTC().Format(time.RFC3339Nano)
msgID := event.Check
if msgID == "" {
msgID = "-"
}
procID := os.Getpid()
line := fmt.Sprintf("<%d>1 %s %s csm %d %s - %s",
pri, ts, s.cfg.Hostname, procID, msgID, body)
// RFC 5424 over UDP / unixgram is one datagram per message; over
// TCP / TLS / unix-stream the receiver expects either octet
// counting ("nnn ") or LF framing. LF is the common rsyslog
// default; emit it for stream transports.
if s.cfg.Network == "tcp" || s.cfg.Network == "tls" || s.cfg.Network == "unix" {
line += "\n"
}
return []byte(line), nil
}
// severityToSyslogLevel maps CSM severity strings onto the standard
// syslog level codes. Critical -> 2 (crit), High -> 3 (err),
// Warning -> 4 (warning), default -> 6 (info).
func severityToSyslogLevel(s string) int {
switch s {
case "CRITICAL":
return 2
case "HIGH":
return 3
case "WARNING":
return 4
}
return 6
}
package alert
import (
"crypto/tls"
"fmt"
"net"
"net/smtp"
"strings"
"time"
"github.com/pidginhost/csm/internal/config"
)
var emailSendTimeout = 10 * time.Second
var smtpDial = func(timeout time.Duration, addr string) (net.Conn, error) {
return (&net.Dialer{Timeout: timeout}).Dial("tcp", addr)
}
func SendEmail(cfg *config.Config, subject, body string) error {
to := cfg.Alerts.Email.To
from := cfg.Alerts.Email.From
smtpAddr := cfg.Alerts.Email.SMTP
if len(to) == 0 {
return fmt.Errorf("no email recipients configured")
}
msg := fmt.Sprintf("From: %s\r\nTo: %s\r\nSubject: %s\r\nContent-Type: text/plain; charset=UTF-8\r\n\r\n%s",
from,
strings.Join(to, ", "),
subject,
body,
)
host := smtpHost(smtpAddr)
deadline := time.Now().Add(emailSendTimeout)
if err := sendSMTPWithDeadline(smtpAddr, host, from, to, []byte(msg), deadline, true); err != nil {
if fallbackErr := sendSMTPWithDeadline(smtpAddr, host, from, to, []byte(msg), deadline, false); fallbackErr != nil {
return fmt.Errorf("%w (original: %v)", fallbackErr, err)
}
}
return nil
}
func smtpHost(addr string) string {
host, _, err := net.SplitHostPort(addr)
if err == nil {
return strings.Trim(host, "[]")
}
return strings.Split(addr, ":")[0]
}
func sendSMTPWithDeadline(addr, host, from string, to []string, msg []byte, deadline time.Time, tryStartTLS bool) error {
timeout := time.Until(deadline)
if timeout <= 0 {
return fmt.Errorf("smtp send timed out after %s", emailSendTimeout)
}
conn, dialErr := smtpDial(timeout, addr)
if dialErr != nil {
return fmt.Errorf("smtp dial %s: %w", addr, dialErr)
}
if err := conn.SetDeadline(deadline); err != nil {
_ = conn.Close()
return fmt.Errorf("smtp deadline: %w", err)
}
c, clientErr := smtp.NewClient(conn, host)
if clientErr != nil {
_ = conn.Close()
return fmt.Errorf("smtp connect %s: %w", addr, clientErr)
}
defer func() { _ = c.Close() }()
if err := c.Hello(host); err != nil {
return fmt.Errorf("smtp hello: %w", err)
}
if tryStartTLS {
if ok, _ := c.Extension("STARTTLS"); ok {
if err := c.StartTLS(&tls.Config{ServerName: host, MinVersion: tls.VersionTLS12}); err != nil {
return fmt.Errorf("smtp starttls: %w", err)
}
}
}
if err := c.Mail(from); err != nil {
return fmt.Errorf("smtp mail from: %w", err)
}
for _, addr := range to {
if err := c.Rcpt(addr); err != nil {
return fmt.Errorf("smtp rcpt %s: %w", addr, err)
}
}
w, dataErr := c.Data()
if dataErr != nil {
return fmt.Errorf("smtp data: %w", dataErr)
}
if _, writeErr := w.Write(msg); writeErr != nil {
return fmt.Errorf("smtp write: %w", writeErr)
}
if closeErr := w.Close(); closeErr != nil {
return fmt.Errorf("smtp close: %w", closeErr)
}
if quitErr := c.Quit(); quitErr != nil {
return fmt.Errorf("smtp quit: %w", quitErr)
}
return nil
}
package alert
import (
"encoding/json"
"net"
"os"
"path/filepath"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
)
// State files drive suppression of handled attacks and auto-block alerts.
// Parse failures must degrade open so corrupt state never hides
// operator-facing findings; warning logs keep the corruption visible.
// Missing state files stay silent because they are normal before the
// first block.
// FilterBlockedAlerts removes attack alerts answered by a source IP's block
// or challenge, along with the corresponding automatic action alerts.
func FilterBlockedAlerts(cfg *config.Config, findings []Finding) []Finding {
if !cfg.Suppressions.SuppressBlockedAlerts {
return findings
}
// Load only IPs with actual block state. Pending entries have not reached
// the firewall and must remain visible to the operator.
blockedIPs := loadBlockedIPs(cfg.StatePath)
// Also collect IPs and subnets blocked in this batch, and IPs routed
// to the challenge in this batch. A challenge-routed IP is handled
// automatically the same way a blocked one is: it either solves the
// PoW (legitimate) or gets escalated to a block by the port gate, so
// the operator needs no notification either way.
var blockedSubnets []*net.IPNet
challengedIPs := make(map[string]bool)
for _, f := range findings {
switch f.Check {
case "auto_block":
parts := strings.Fields(f.Message)
for i, p := range parts {
if p == "AUTO-BLOCK:" && i+1 < len(parts) {
blockedIPs[parts[i+1]] = true
break
}
if p == "AUTO-BLOCK-SUBNET:" && i+1 < len(parts) {
if _, ipnet, err := net.ParseCIDR(parts[i+1]); err == nil {
blockedSubnets = append(blockedSubnets, ipnet)
}
break
}
}
case "challenge_route":
parts := strings.Fields(f.Message)
for i, p := range parts {
if p == "CHALLENGE:" && i+1 < len(parts) {
challengedIPs[parts[i+1]] = true
break
}
}
}
}
if len(blockedIPs) == 0 && len(blockedSubnets) == 0 && len(challengedIPs) == 0 && ChallengedIPFunc == nil {
return findings
}
// Canonicalize blocked-IP keys via net.ParseIP so equality survives
// notation differences (2001:db8::1 vs 2001:db8:0:0:0:0:0:1,
// IPv4-mapped IPv6). Keys that do not parse as an IP cannot
// exact-match a finding's IP and are dropped.
canonicalBlocked := make(map[string]bool, len(blockedIPs))
for ip := range blockedIPs {
if parsed := net.ParseIP(ip); parsed != nil {
canonicalBlocked[parsed.String()] = true
}
}
canonicalChallenged := make(map[string]bool, len(challengedIPs))
for ip := range challengedIPs {
if parsed := net.ParseIP(ip); parsed != nil {
canonicalChallenged[parsed.String()] = true
}
}
// Filter out alerts for IPs that are handled automatically.
// When suppress_blocked_alerts is on, the operator doesn't want to be
// notified about IPs that are already dealt with - they only want alerts
// that require human action.
policy := currentIPResponsePolicy()
var filtered []Finding
for _, f := range findings {
if f.Check == "auto_block" || f.Check == "challenge_route" {
continue
}
// Suppression keys on the actual block state, never on
// auto-response intent: an enabled block_ips used to drop every
// reputation finding here, which hid exactly the IPs auto-block
// did NOT handle (dry-run, rate-limited queue drops,
// verdict-allowed). Same-batch AUTO-BLOCK findings already feed
// blockedIPs above, so an IP blocked this cycle stays suppressed.
// Structured SourceIP wins when present; older findings fall
// back to the message token. The address is compared
// canonically: substring matching used to let a blocked 1.2.3.4
// suppress a finding about the unrelated 1.2.3.45. A finding
// with no parseable IP is never suppressed (fail open to
// alerting).
findingIP := suppressionIPFromFinding(f)
if findingIP != nil && ipBlocked(findingIP, canonicalBlocked, blockedSubnets) &&
ipResponseAnswers(policy, cfg, f, true) {
continue
}
// A challenge counts only for findings a challenge answers.
// Critical ip_reputation is a browserless mail/SSH/FTP sighting
// and must stay visible until a hard block lands; otherwise an
// old challenge entry hides the exact finding that is supposed
// to escalate it.
if findingIP != nil && ipResponseAnswers(policy, cfg, f, false) &&
ipChallenged(findingIP, canonicalChallenged) {
continue
}
filtered = append(filtered, f)
}
return filtered
}
func ipBlocked(ip net.IP, blocked map[string]bool, subnets []*net.IPNet) bool {
if blocked[ip.String()] {
return true
}
// AUTO-BLOCK-SUBNET: findings from the same batch silence per-IP
// findings for addresses inside that /24.
for _, subnet := range subnets {
if subnet.Contains(ip) {
return true
}
}
return false
}
func ipChallenged(ip net.IP, challenged map[string]bool) bool {
if challenged[ip.String()] {
return true
}
return ChallengedIPFunc != nil && ChallengedIPFunc(ip.String())
}
// ipResponseAnswers reports whether a block (blocked) or a challenge answers
// the finding. With no policy wired it keeps the reputation-only rule.
func ipResponseAnswers(policy IPResponsePolicy, cfg *config.Config, f Finding, blocked bool) bool {
if policy != nil {
return policy(cfg, f, blocked)
}
if f.Check != "ip_reputation" && f.Check != "local_threat_score" {
return false
}
return blocked || challengeHandlesFinding(f)
}
// challengeHandlesFinding mirrors the only finding-level exception to the
// challenge response policy. ip_reputation has one production grading source:
// High is browser-facing HTTP/cPanel traffic, while Critical is a browserless
// vector that resolves to a hard block. local_threat_score remains challenge-
// eligible at Critical severity.
func challengeHandlesFinding(f Finding) bool {
switch f.Check {
case "ip_reputation":
return f.Severity != Critical
case "local_threat_score":
return true
default:
return false
}
}
func suppressionIPFromFinding(f Finding) net.IP {
if strings.TrimSpace(f.SourceIP) != "" {
if normalized := normalizeFindingIP(f.SourceIP); normalized != "" {
return net.ParseIP(normalized)
}
return nil
}
return extractIPFromFindingMessage(f.Message)
}
// extractIPFromFindingMessage scans a finding message for the first token
// that parses as a valid IP address and returns it. Returns nil if no IP
// is found. Used to match reputation findings against blocked IPs and
// blocked-subnet CIDRs.
func extractIPFromFindingMessage(msg string) net.IP {
for _, field := range strings.Fields(msg) {
if ip := parseIPMessageField(field); ip != nil {
return ip
}
}
return nil
}
func parseIPMessageField(field string) net.IP {
if ip := parseIPMessageToken(field); ip != nil {
return ip
}
// Accept key=value tokens such as "ip=5.5.5.5". IP literals never
// contain '=', so taking the value side cannot mangle a real IP.
if idx := strings.LastIndexByte(field, '='); idx >= 0 {
return parseIPMessageToken(field[idx+1:])
}
return nil
}
func parseIPMessageToken(token string) net.IP {
token = strings.TrimSpace(token)
if token == "" {
return nil
}
if ip := parseNormalizedIP(token); ip != nil {
return ip
}
unquoted := strings.Trim(token, "\"'`")
if ip := parseNormalizedIP(unquoted); ip != nil {
return ip
}
withoutTrailing := strings.TrimRight(unquoted, ",;")
if ip := parseNormalizedIP(withoutTrailing); ip != nil {
return ip
}
unwrapped := strings.Trim(withoutTrailing, "()[]{}<>")
if ip := parseNormalizedIP(unwrapped); ip != nil {
return ip
}
if strings.HasSuffix(unwrapped, ":") {
withoutColon := strings.TrimSuffix(unwrapped, ":")
if ip := parseNormalizedIP(withoutColon); ip != nil {
return ip
}
if ip := parseNormalizedIP(strings.Trim(withoutColon, "()[]{}<>")); ip != nil {
return ip
}
}
return nil
}
func parseNormalizedIP(raw string) net.IP {
if normalized := normalizeFindingIP(raw); normalized != "" {
return net.ParseIP(normalized)
}
return nil
}
func loadBlockedIPSource(statePath string, now time.Time, ips map[string]bool) {
// Use injected loader (bbolt-backed) when available.
if BlockedIPsFunc != nil {
for ip, v := range BlockedIPsFunc() {
ips[ip] = v
}
return
}
loadFirewallStateFile(statePath, now, ips)
}
// BlockedIPsFunc is an optional callback that returns currently blocked IPs.
// Set by the daemon (or tests) to provide blocked IPs from bbolt store,
// avoiding a circular import between alert and store packages.
// When nil, loadBlockedIPs falls back to reading flat files.
var BlockedIPsFunc func() map[string]bool
// ChallengedIPFunc reports whether an IP is currently on the challenge
// list. Set by the daemon from the live challenge IP list, mirroring
// BlockedIPsFunc; nil means challenge membership is unknown and no
// challenge-based suppression happens (fail open to alerting).
var ChallengedIPFunc func(ip string) bool
// loadBlockedIPs reads blocked IPs from both the firewall engine state
// and the legacy blocked_ips.json file.
func loadBlockedIPs(statePath string) map[string]bool {
ips := make(map[string]bool)
now := time.Now()
loadBlockedIPSource(statePath, now, ips)
loadBlockFileEntries(statePath, now, ips, nil, blockFileIPsSection)
return ips
}
func loadFirewallStateFile(statePath string, now time.Time, ips map[string]bool) {
fwPath := filepath.Join(statePath, "firewall", "state.json")
fwData, err := os.ReadFile(fwPath) // #nosec G304 -- filepath.Join under operator-configured statePath.
if err != nil {
return
}
var fwState struct {
Blocked []struct {
IP string `json:"ip"`
ExpiresAt time.Time `json:"expires_at"`
} `json:"blocked"`
}
if err := json.Unmarshal(fwData, &fwState); err != nil {
csmlog.Warn("alert filter: firewall state.json unparseable, suppression degraded",
"path", fwPath, "err", err)
return
}
for _, entry := range fwState.Blocked {
if entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt) {
ips[entry.IP] = true
}
}
}
type blockFile struct {
IPs []blockFileIP
Pending []blockFilePendingIP
}
type blockFileIP struct {
IP string `json:"ip"`
ExpiresAt time.Time `json:"expires_at"`
}
type blockFilePendingIP struct {
IP string `json:"ip"`
}
type blockFileSection uint8
const (
blockFileIPsSection blockFileSection = 1 << iota
blockFilePendingSection
)
func loadBlockFileEntries(statePath string, now time.Time, ips map[string]bool, pending map[string]bool, sections blockFileSection) {
bf, ok := loadBlockFile(statePath, sections)
if !ok {
return
}
for _, entry := range bf.IPs {
if ips != nil && (entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt)) {
ips[entry.IP] = true
}
}
for _, entry := range bf.Pending {
if pending != nil {
pending[entry.IP] = true
}
}
}
func loadBlockFile(statePath string, sections blockFileSection) (blockFile, bool) {
blockedPath := filepath.Join(statePath, "blocked_ips.json")
data, err := os.ReadFile(blockedPath) // #nosec G304 -- filepath.Join under operator-configured statePath.
if err != nil {
return blockFile{}, false
}
var raw struct {
IPs json.RawMessage `json:"ips"`
Pending json.RawMessage `json:"pending"`
}
if err := json.Unmarshal(data, &raw); err != nil {
csmlog.Warn("alert filter: blocked_ips.json unparseable, suppression degraded",
"path", blockedPath, "err", err)
return blockFile{}, false
}
var bf blockFile
if sections&blockFileIPsSection != 0 && len(raw.IPs) > 0 && string(raw.IPs) != "null" {
if err := json.Unmarshal(raw.IPs, &bf.IPs); err != nil {
csmlog.Warn("alert filter: blocked_ips.json blocked entries unparseable, suppression degraded",
"path", blockedPath, "err", err)
}
}
if sections&blockFilePendingSection != 0 && len(raw.Pending) > 0 && string(raw.Pending) != "null" {
if err := json.Unmarshal(raw.Pending, &bf.Pending); err != nil {
csmlog.Warn("alert filter: blocked_ips.json pending entries unparseable, suppression degraded",
"path", blockedPath, "err", err)
}
}
return bf, true
}
// IPResponsePolicy reports whether an IP disposition on a finding's source
// fully answers the finding, so suppress_blocked_alerts may drop it. blocked
// is true for a firewall block and false for a challenge.
type IPResponsePolicy func(cfg *config.Config, f Finding, blocked bool) bool
var (
ipResponsePolicyMu sync.RWMutex
ipResponsePolicy IPResponsePolicy
)
// SetIPResponsePolicy installs or clears the policy FilterBlockedAlerts
// consults and returns the previous one. The check classification lives in
// the checks package, which imports this one and registers it at package
// initialization so every dispatch path uses the same policy.
func SetIPResponsePolicy(p IPResponsePolicy) IPResponsePolicy {
ipResponsePolicyMu.Lock()
defer ipResponsePolicyMu.Unlock()
previous := ipResponsePolicy
ipResponsePolicy = p
return previous
}
func currentIPResponsePolicy() IPResponsePolicy {
ipResponsePolicyMu.RLock()
p := ipResponsePolicy
ipResponsePolicyMu.RUnlock()
return p
}
package alert
import (
"bytes"
"encoding/binary"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/queuehealth"
bolt "go.etcd.io/bbolt"
)
const (
phpanelQueueLimit = 100000
phpanelQuarantineLimit = 1000
)
var phpanelQueueBucket = []byte("findings")
var phpanelQuarantineBucket = []byte("quarantine")
type queuedPhpanelFinding struct {
Finding Finding `json:"finding"`
Timestamp time.Time `json:"timestamp"`
}
type quarantinedPhpanelFinding struct {
Payload []byte `json:"payload"`
Error string `json:"error"`
QuarantinedAt time.Time `json:"quarantined_at"`
}
type phpanelDeliveryConfig struct {
hostname string
url string
hmacSecret string
hmacSecretEnv string
}
type phpanelQueue struct {
db *bolt.DB
cfgMu sync.RWMutex
cfg phpanelDeliveryConfig
wake chan struct{}
stop chan struct{}
done chan struct{}
drain sync.Mutex
retryMu sync.Mutex
retryAt time.Time
retryDelay time.Duration
closed bool
mu sync.Mutex
mutation sync.Mutex
limit int
health phpanelQueueHealth
}
var phpanelQueues = struct {
sync.Mutex
byState map[string]*phpanelQueue
}{byState: make(map[string]*phpanelQueue)}
func enqueuePhpanelFindings(cfg *config.Config, findings []Finding) error {
queue, err := phpanelQueueForMode(cfg, false)
if err != nil {
return err
}
queued := make([]queuedPhpanelFinding, 0, len(findings))
for _, finding := range findings {
queued = append(queued, queuedPhpanelFinding{Finding: finding, Timestamp: time.Now().UTC()})
}
dropped, err := queue.enqueueBatch(queued)
if err != nil {
return fmt.Errorf("queueing phpanel webhook: %w", err)
}
select {
case queue.wake <- struct{}{}:
default:
}
if dropped > 0 {
alertDispatchFailures.Add(float64(dropped))
return fmt.Errorf("phpanel webhook queue reached %d entries and dropped %d oldest findings", phpanelQueueLimit, dropped)
}
return nil
}
// ConfigurePhpanelQueue opens and wakes the durable queue during daemon
// startup and safe config reloads. This lets persisted findings resume delivery
// without waiting for a new finding to arrive after a restart.
func ConfigurePhpanelQueue(cfg *config.Config) error {
if cfg == nil || !cfg.Alerts.Webhook.Enabled || cfg.Alerts.Webhook.Type != "phpanel" {
closePhpanelQueue(cfg)
return nil
}
queue, err := phpanelQueueFor(cfg)
if err != nil {
return err
}
select {
case queue.wake <- struct{}{}:
default:
}
return nil
}
func closePhpanelQueue(cfg *config.Config) {
if cfg == nil || cfg.StatePath == "" {
return
}
statePath, err := filepath.Abs(cfg.StatePath)
if err != nil {
return
}
phpanelQueues.Lock()
queue := phpanelQueues.byState[statePath]
delete(phpanelQueues.byState, statePath)
phpanelQueues.Unlock()
if queue != nil {
queue.close()
}
}
func phpanelQueueFor(cfg *config.Config) (*phpanelQueue, error) {
return phpanelQueueForMode(cfg, true)
}
func phpanelQueueForMode(cfg *config.Config, updateExisting bool) (*phpanelQueue, error) {
if cfg.StatePath == "" {
return nil, fmt.Errorf("phpanel webhook requires state_path for its durable queue")
}
statePath, err := filepath.Abs(cfg.StatePath)
if err != nil {
return nil, fmt.Errorf("resolving phpanel queue state path: %w", err)
}
phpanelQueues.Lock()
defer phpanelQueues.Unlock()
if queue := phpanelQueues.byState[statePath]; queue != nil {
if updateExisting {
queue.updateConfig(cfg)
}
return queue, nil
}
if mkdirErr := os.MkdirAll(statePath, 0o700); mkdirErr != nil {
return nil, fmt.Errorf("creating phpanel queue state directory: %w", mkdirErr)
}
db, err := bolt.Open(filepath.Join(statePath, "phpanel-webhook.db"), 0o600, &bolt.Options{Timeout: 2 * time.Second})
if err != nil {
return nil, fmt.Errorf("opening phpanel webhook queue: %w", err)
}
queue, err := newPhpanelQueue(db, phpanelQueueLimit)
if err != nil {
_ = db.Close()
return nil, fmt.Errorf("creating phpanel webhook queue: %w", err)
}
queue.updateConfig(cfg)
phpanelQueues.byState[statePath] = queue
phpanelHealth.Lock()
phpanelHealth.active[queue] = struct{}{}
phpanelHealth.Unlock()
obs.Go("phpanel-webhook-queue", queue.run)
return queue, nil
}
func newPhpanelQueue(db *bolt.DB, limit int) (*phpanelQueue, error) {
q := &phpanelQueue{
db: db, limit: limit,
wake: make(chan struct{}, 1), stop: make(chan struct{}), done: make(chan struct{}),
health: phpanelQueueHealth{
// Two normal drain intervals without completion warrant attention.
stats: queuehealth.NewSharedCapacity(limit, time.Minute), pending: make(map[string]*phpanelWork),
},
}
if err := db.Update(func(tx *bolt.Tx) error {
if _, err := tx.CreateBucketIfNotExists(phpanelQueueBucket); err != nil {
return err
}
_, err := tx.CreateBucketIfNotExists(phpanelQuarantineBucket)
return err
}); err != nil {
return nil, err
}
now := time.Now()
if err := db.View(func(tx *bolt.Tx) error {
return tx.Bucket(phpanelQueueBucket).ForEach(func(key, payload []byte) error {
q.health.pending[string(key)] = &phpanelWork{ticket: q.health.stats.BeginAt(phpanelRecordTime(payload, now), now)}
return nil
})
}); err != nil {
return nil, err
}
return q, nil
}
func phpanelRecordTime(payload []byte, now time.Time) time.Time {
var item queuedPhpanelFinding
if err := json.Unmarshal(payload, &item); err == nil && !item.Timestamp.IsZero() && item.Timestamp.Before(now) {
return item.Timestamp
}
// Damaged or future timestamps supply no reliable elapsed age.
// The record still enters normal delivery/quarantine processing.
return now
}
func (q *phpanelQueue) updateConfig(cfg *config.Config) {
q.cfgMu.Lock()
q.cfg = phpanelDeliveryConfig{
hostname: cfg.Hostname,
url: cfg.Alerts.Webhook.URL,
hmacSecret: cfg.Alerts.Webhook.HMACSecret,
hmacSecretEnv: cfg.Alerts.Webhook.HMACSecretEnv,
}
q.cfgMu.Unlock()
q.retryMu.Lock()
q.retryAt = time.Time{}
q.retryDelay = 0
q.retryMu.Unlock()
}
func (q *phpanelQueue) enqueueBatch(items []queuedPhpanelFinding) (int, error) {
if len(items) == 0 {
return 0, nil
}
now := time.Now()
work := make([]*phpanelWork, len(items))
for i, item := range items {
queuedAt := item.Timestamp
if queuedAt.IsZero() || queuedAt.After(now) {
queuedAt = now
}
work[i] = &phpanelWork{ticket: q.health.stats.BeginAt(queuedAt, now)}
work[i].ticket.Start(now)
}
committed := false
defer func() {
if !committed {
for _, entry := range work {
q.discardWork(entry, time.Now())
}
}
}()
bodies := make([][]byte, 0, len(items))
for _, item := range items {
body, err := json.Marshal(item)
if err != nil {
q.health.enqueueFailed.Store(true)
return 0, err
}
bodies = append(bodies, body)
}
q.mutation.Lock()
defer q.mutation.Unlock()
select {
case <-q.stop:
return 0, fmt.Errorf("phpanel webhook queue is stopped")
default:
}
var added, evicted []string
err := q.db.Update(func(tx *bolt.Tx) error {
bucket := tx.Bucket(phpanelQueueBucket)
count := queuedFindingCount(bucket) + len(bodies)
for _, body := range bodies {
seq, err := bucket.NextSequence()
if err != nil {
return err
}
var key [8]byte
binary.BigEndian.PutUint64(key[:], seq)
if err := bucket.Put(key[:], body); err != nil {
return err
}
added = append(added, string(key[:]))
}
for count > q.limit {
oldest, _ := bucket.Cursor().First()
if oldest == nil {
// The live span is empty, so the count came from somewhere
// other than this bucket. Nothing is left to evict.
break
}
evictedKey := string(oldest)
if err := bucket.Delete(oldest); err != nil {
return err
}
evicted = append(evicted, evictedKey)
count--
}
return nil
})
if err != nil {
q.health.enqueueFailed.Store(true)
return 0, err
}
now = time.Now()
for i, key := range added {
work[i].ticket.Requeue(now)
q.health.pending[key] = work[i]
}
for _, key := range evicted {
entry := q.health.pending[key]
delete(q.health.pending, key)
switch entry {
case nil:
// A record evicted from the queue file with no accounting cannot
// be attributed to a caller; count the finding it carried as lost.
phpanelHealth.losses.Lose(now, 1)
case q.health.active:
entry.evicted = true
default:
q.discardWork(entry, now)
}
}
q.health.enqueueFailed.Store(false)
committed = true
return len(evicted), nil
}
// queuedFindingCount returns the number of live entries without walking every
// page. Deliveries, quarantine, and overflow trimming only ever remove the
// current oldest entry, so live keys stay a contiguous span of the monotonic
// sequence numbers assigned by NextSequence; the count is that span.
func queuedFindingCount(bucket *bolt.Bucket) int {
cursor := bucket.Cursor()
firstKey, _ := cursor.First()
if firstKey == nil {
return 0
}
lastKey, _ := cursor.Last()
// last >= first (bbolt key order) and the live span is bounded by
// phpanelQueueLimit, so the difference always fits in an int.
// #nosec G115 -- bounded span (<= phpanelQueueLimit); cannot overflow int.
return int(binary.BigEndian.Uint64(lastKey)-binary.BigEndian.Uint64(firstKey)) + 1
}
func (q *phpanelQueue) run() {
defer close(q.done)
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-q.wake:
q.drainQueued()
case <-ticker.C:
q.drainQueued()
case <-q.stop:
return
}
}
}
func (q *phpanelQueue) drainQueued() {
q.drain.Lock()
defer q.drain.Unlock()
q.retryMu.Lock()
if time.Now().Before(q.retryAt) {
q.retryMu.Unlock()
return
}
q.retryMu.Unlock()
for {
select {
case <-q.stop:
return
default:
}
key, payload, work, err := q.takeDelivery()
if err != nil {
fmt.Fprintf(os.Stderr, "alert: reading phpanel webhook queue: %v\n", err)
alertDispatchFailures.Inc()
return
}
if key == nil {
return
}
if !q.deliverQueued(key, payload, work) {
return
}
}
}
func (q *phpanelQueue) takeDelivery() ([]byte, []byte, *phpanelWork, error) {
q.mutation.Lock()
defer q.mutation.Unlock()
var key, payload []byte
err := q.db.View(func(tx *bolt.Tx) error {
firstKey, value := tx.Bucket(phpanelQueueBucket).Cursor().First()
if firstKey != nil {
key = append([]byte(nil), firstKey...)
payload = append([]byte(nil), value...)
}
return nil
})
q.health.readFailed.Store(err != nil)
if err != nil || key == nil {
return key, payload, nil, err
}
work := q.health.pending[string(key)]
if work == nil {
// A boundary record absent from the rebuilt map still needs delivery
// and must retain any reliable waiting age across retries.
now := time.Now()
work = &phpanelWork{ticket: q.health.stats.BeginAt(phpanelRecordTime(payload, now), now)}
q.health.pending[string(key)] = work
}
work.ticket.Start(time.Now())
q.health.active = work
return key, payload, work, nil
}
func (q *phpanelQueue) deliverQueued(key, payload []byte, work *phpanelWork) bool {
sent, removed := false, false
attempted := false
defer func() {
if attempted && !sent {
q.health.sendFailed.Store(true)
}
q.finishDelivery(work, sent, removed)
}()
var item queuedPhpanelFinding
if err := json.Unmarshal(payload, &item); err != nil {
if quarantineErr := q.quarantineMalformed(key, payload, err); quarantineErr != nil {
fmt.Fprintf(os.Stderr, "alert: quarantining malformed phpanel webhook: %v\n", quarantineErr)
alertDispatchFailures.Inc()
return false
}
fmt.Fprintf(os.Stderr, "alert: quarantined malformed phpanel webhook: %v\n", err)
alertDispatchFailures.Inc()
return true
}
q.cfgMu.RLock()
delivery := q.cfg
q.cfgMu.RUnlock()
attempted = true
if err := sendQueuedPhpanelWebhookFinding(delivery, item); err != nil {
fmt.Fprintf(os.Stderr, "alert: phpanel webhook delivery failed: %v\n", err)
alertDispatchFailures.Inc()
q.recordRetryFailure()
return false
}
sent = true
q.health.sendFailed.Store(false)
var err error
removed, err = q.removeDelivered(key)
if err != nil {
fmt.Fprintf(os.Stderr, "alert: deleting delivered phpanel webhook: %v\n", err)
alertDispatchFailures.Inc()
return false
}
q.clearRetryFailure()
return true
}
func (q *phpanelQueue) removeDelivered(key []byte) (bool, error) {
q.mutation.Lock()
defer q.mutation.Unlock()
err := q.db.Update(func(tx *bolt.Tx) error { return tx.Bucket(phpanelQueueBucket).Delete(key) })
q.health.removeFailed.Store(err != nil)
if err != nil {
return false, err
}
delete(q.health.pending, string(key))
return true, nil
}
func (q *phpanelQueue) quarantineMalformed(key, payload []byte, decodeErr error) error {
return q.quarantineMalformedWithLimit(key, payload, decodeErr, phpanelQuarantineLimit)
}
func (q *phpanelQueue) quarantineMalformedWithLimit(key, payload []byte, decodeErr error, limit int) error {
if limit <= 0 {
return fmt.Errorf("phpanel quarantine limit must be positive")
}
record, marshalErr := json.Marshal(quarantinedPhpanelFinding{
Payload: payload,
Error: decodeErr.Error(),
QuarantinedAt: time.Now().UTC(),
})
if marshalErr != nil {
return marshalErr
}
q.mutation.Lock()
defer q.mutation.Unlock()
removed := false
updateErr := q.db.Update(func(tx *bolt.Tx) error {
active := tx.Bucket(phpanelQueueBucket)
current := active.Get(key)
if current == nil {
return nil
}
if !bytes.Equal(current, payload) {
return fmt.Errorf("phpanel queue entry changed while being quarantined")
}
quarantine := tx.Bucket(phpanelQuarantineBucket)
if quarantine.Get(key) == nil {
count := quarantine.Stats().KeyN
for count >= limit {
oldest, _ := quarantine.Cursor().First()
if oldest == nil {
break
}
if err := quarantine.Delete(oldest); err != nil {
return err
}
count--
}
}
if err := quarantine.Put(key, record); err != nil {
return err
}
if err := active.Delete(key); err != nil {
return err
}
removed = true
return nil
})
q.health.removeFailed.Store(updateErr != nil)
if updateErr == nil && removed {
work := q.health.pending[string(key)]
delete(q.health.pending, string(key))
switch work {
case nil:
// Quarantined a record this process never accounted for.
phpanelHealth.losses.Lose(time.Now(), 1)
case q.health.active:
work.evicted = true
default:
q.discardWork(work, time.Now())
}
}
return updateErr
}
func (q *phpanelQueue) recordRetryFailure() {
q.retryMu.Lock()
defer q.retryMu.Unlock()
if q.retryDelay == 0 {
q.retryDelay = 30 * time.Second
} else {
q.retryDelay *= 2
if q.retryDelay > 15*time.Minute {
q.retryDelay = 15 * time.Minute
}
}
q.retryAt = time.Now().Add(q.retryDelay)
}
func (q *phpanelQueue) clearRetryFailure() {
q.retryMu.Lock()
q.retryAt = time.Time{}
q.retryDelay = 0
q.retryMu.Unlock()
}
func (q *phpanelQueue) close() {
q.mu.Lock()
if q.closed {
q.mu.Unlock()
return
}
q.closed = true
close(q.stop)
q.mu.Unlock()
<-q.done
q.drain.Lock()
q.mutation.Lock()
phpanelHealth.Lock()
delete(phpanelHealth.active, q)
phpanelHealth.Unlock()
for _, work := range q.health.pending {
// These findings remain durable for a later start.
work.ticket.Finish(time.Now())
}
clear(q.health.pending)
_ = q.db.Close()
q.mutation.Unlock()
q.drain.Unlock()
}
func closePhpanelQueuesForTest() {
ClosePhpanelQueues()
}
// ClosePhpanelQueues stops delivery workers and closes their durable databases.
func ClosePhpanelQueues() {
phpanelQueues.Lock()
queues := make([]*phpanelQueue, 0, len(phpanelQueues.byState))
for _, queue := range phpanelQueues.byState {
queues = append(queues, queue)
}
phpanelQueues.byState = make(map[string]*phpanelQueue)
phpanelQueues.Unlock()
for _, queue := range queues {
queue.close()
}
}
func phpanelQueueDepthForTest(statePath string) int {
absolute, _ := filepath.Abs(statePath)
phpanelQueues.Lock()
queue := phpanelQueues.byState[absolute]
phpanelQueues.Unlock()
if queue == nil {
return 0
}
depth := 0
_ = queue.db.View(func(tx *bolt.Tx) error {
depth = tx.Bucket(phpanelQueueBucket).Stats().KeyN
return nil
})
return depth
}
package alert
import (
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type phpanelWork struct {
ticket queuehealth.Ticket
evicted bool
delivered bool
}
type phpanelQueueHealth struct {
stats *queuehealth.Tracker
pending map[string]*phpanelWork // guarded by phpanelQueue.mutation
active *phpanelWork
enqueueFailed atomic.Bool
readFailed atomic.Bool
removeFailed atomic.Bool
sendFailed atomic.Bool
}
// Factory locking spans database open. This registry only publishes memory
// snapshots, so status remains available during a stalled open, write or send.
var phpanelHealth = struct {
sync.RWMutex
active map[*phpanelQueue]struct{}
losses *queuehealth.Tracker
}{active: make(map[*phpanelQueue]struct{}), losses: queuehealth.New(0, time.Minute)}
// PhpanelQueueStatus includes cumulative loss from queues disabled or replaced
// during this process. It retains no retired queues or state-path labels.
func PhpanelQueueStatus(now time.Time) queuehealth.Status {
phpanelHealth.RLock()
queues := make([]*phpanelQueue, 0, len(phpanelHealth.active))
for q := range phpanelHealth.active {
queues = append(queues, q)
}
phpanelHealth.RUnlock()
status := phpanelHealth.losses.Snapshot(now)
for _, q := range queues {
local := q.queueStatus(now)
status.Depth += local.Depth
status.InFlight += local.InFlight
status.Capacity += local.Capacity
status.LagSeconds = max(status.LagSeconds, local.LagSeconds)
status.ProcessingSeconds = max(status.ProcessingSeconds, local.ProcessingSeconds)
if phpanelReasonRank(local.Reason) > phpanelReasonRank(status.Reason) {
status.Status, status.Reason = local.Status, local.Reason
}
}
return status
}
func phpanelReasonRank(reason string) int {
switch reason {
case "spool_io":
return 6
case "delivery_failed":
return 5
case "backlog_lag":
return 4
case "processing_lag":
return 3
case "queue_full":
return 2
case "dropped_work":
return 1
default:
return 0
}
}
func (q *phpanelQueue) QueueStatuses(now time.Time) map[string]queuehealth.Status {
return map[string]queuehealth.Status{"spool": q.queueStatus(now)}
}
func (q *phpanelQueue) queueStatus(now time.Time) queuehealth.Status {
status := q.health.stats.Snapshot(now)
switch {
case q.health.enqueueFailed.Load() || q.health.readFailed.Load() || q.health.removeFailed.Load():
status.Status, status.Reason = "degraded", "spool_io"
case q.health.sendFailed.Load():
status.Status, status.Reason = "degraded", "delivery_failed"
}
return status
}
func (q *phpanelQueue) discardWork(work *phpanelWork, now time.Time) {
if work.delivered {
work.ticket.Finish(now)
return
}
work.ticket.Reject(now)
phpanelHealth.losses.Lose(now, 1)
}
// The send owns an evicted record until its result is known. Evicting its
// durable copy is not a lost finding if the collector already received it.
func (q *phpanelQueue) finishDelivery(work *phpanelWork, sent, removed bool) {
q.mutation.Lock()
defer q.mutation.Unlock()
now := time.Now()
// An acknowledged send may remain queued after database removal fails.
// Later retries cannot undo the collector's earlier receipt.
work.delivered = work.delivered || sent
switch {
case removed:
work.ticket.Finish(now)
case work.evicted:
q.discardWork(work, now)
default:
work.ticket.Requeue(now)
}
q.health.active = nil
}
package alert
import (
"errors"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var (
ErrQueueTimeout = errors.New("finding queue deadline exceeded")
ErrQueueStopped = errors.New("finding queue stopped")
)
var ingestQueues = struct {
sync.RWMutex
queues map[chan<- Finding]*queuehealth.Tracker
}{queues: make(map[chan<- Finding]*queuehealth.Tracker)}
// RegisterQueue attaches accounting to a daemon's ingest channel for its
// lifetime. Standalone consumers have no daemon health snapshot and continue
// to use the same watcher APIs without a registered ingest channel.
func RegisterQueue(ch chan<- Finding, q *queuehealth.Tracker) func() {
ingestQueues.Lock()
ingestQueues.queues[ch] = q
ingestQueues.Unlock()
return func() {
ingestQueues.Lock()
defer ingestQueues.Unlock()
if ingestQueues.queues[ch] == q {
delete(ingestQueues.queues, ch)
}
}
}
func queuedFinding(ch chan<- Finding, f Finding) Finding {
f.queueTicket = queuehealth.Ticket{}
ingestQueues.RLock()
q := ingestQueues.queues[ch]
ingestQueues.RUnlock()
if q != nil {
f.queueTicket = q.Begin(time.Now())
}
return f
}
// TryEnqueue preserves the realtime producers' nonblocking contract while
// counting every lost finding in the shared ingest queue's health evidence.
func TryEnqueue(ch chan<- Finding, f Finding) bool {
f = queuedFinding(ch, f)
select {
case ch <- f:
return true
default:
f.queueTicket.Reject(time.Now())
return false
}
}
// Enqueue waits for capacity or shutdown. Callers that cannot block use
// TryEnqueue; bounded queues must never spawn a goroutine for each send.
func Enqueue(ch chan<- Finding, f Finding, stop <-chan struct{}) bool {
f = queuedFinding(ch, f)
select {
case ch <- f:
return true
case <-stop:
f.queueTicket.Reject(time.Now())
return false
}
}
// EnqueueWithin admits one finding with a bounded wait. The initial full
// channel is backpressure, not loss: its ticket is rejected only if the
// deadline or shutdown wins before a send succeeds.
func EnqueueWithin(ch chan<- Finding, f Finding, stop <-chan struct{}, timeout time.Duration) error {
f = queuedFinding(ch, f)
select {
case ch <- f:
return nil
default:
}
timer := time.NewTimer(timeout)
defer timer.Stop()
select {
case ch <- f:
return nil
case <-stop:
f.queueTicket.Reject(time.Now())
return ErrQueueStopped
case <-timer.C:
f.queueTicket.Reject(time.Now())
return ErrQueueTimeout
}
}
// RecordQueueLoss covers work abandoned with the rest of a batch before
// individual tickets were created.
func RecordQueueLoss(ch chan<- Finding, count uint64) {
ingestQueues.RLock()
q := ingestQueues.queues[ch]
ingestQueues.RUnlock()
if q != nil {
q.Lose(time.Now(), count)
}
}
func RejectQueued(f Finding) { f.queueTicket.Reject(time.Now()) }
func StartQueued(f Finding) { f.queueTicket.Start(time.Now()) }
// FinishQueued receives the original input batch before filtering or
// deduplication, so discarded duplicates also release their queue tickets.
func FinishQueued(findings []Finding) {
now := time.Now()
for _, f := range findings {
f.queueTicket.Finish(now)
}
}
package alert
import (
"path"
"strings"
)
const redactedToken = "[REDACTED]"
// RedactCommandLine masks credential-bearing arguments in a command line
// (or any text quoting one): the MySQL family's attached -pSECRET, KEY=VALUE
// tokens whose key names a password, token, secret or API key (on options
// and on environment assignments alike), separated --password VALUE forms,
// curl-style user:password pairs, and URL userinfo or query credentials.
// Whitespace is preserved so the rest of the line stays readable.
func RedactCommandLine(s string) string {
if s == "" {
return s
}
display, tokens := splitTokens(s)
if len(tokens) == 0 {
return display
}
mysqlAt := -1
sshpassAt := -1
for i, tok := range tokens {
base := strings.ToLower(path.Base(strings.Trim(tok.text, "\"'")))
if mysqlAt < 0 && mysqlFamily[base] {
mysqlAt = i
}
if sshpassAt < 0 && base == "sshpass" {
sshpassAt = i
}
}
changed := false
redactNext := false
redactNextPair := false
sshpassPending := false
for i := range tokens {
tok := &tokens[i]
text := tok.text
switch {
case redactNext:
redactNext = false
tok.text = redactedToken
changed = changed || tok.text != text
continue
case redactNextPair:
redactNextPair = false
if r, ok := redactPair(text); ok {
tok.text = r
changed = true
continue
}
}
if mysqlAt >= 0 && i > mysqlAt && len(text) > 2 && strings.HasPrefix(text, "-p") && text[2] != '-' {
tok.text = "-p" + redactedToken
changed = true
continue
}
if sshpassAt >= 0 && i == sshpassAt {
sshpassPending = true
}
if sshpassPending && i > sshpassAt && len(text) > 2 && strings.HasPrefix(text, "-p") && text[2] != '-' {
sshpassPending = false
tok.text = "-p" + redactedToken
changed = true
continue
}
if sshpassPending && i > sshpassAt && text == "-p" {
sshpassPending = false
redactNext = true
continue
}
if separatedSecretFlags[strings.ToLower(text)] {
redactNext = true
continue
}
if separatedPairFlags[strings.ToLower(text)] {
redactNextPair = true
continue
}
hasURL := strings.Contains(text, "://")
// An assignment can contain a URL as its value. An equals sign
// inside the URL itself must never make its authority a secret key.
eq := strings.IndexByte(text, '=')
if !hasURL || (eq >= 0 && eq < strings.Index(text, "://")) {
if r, ok := redactAssignments(text); ok {
tok.text = r
changed = true
continue
}
}
if tok.argv && strings.ContainsAny(text, " \t\n\r\v\f") {
if r := RedactCommandLine(text); r != text {
tok.text = r
changed = true
continue
}
}
if hasURL {
if r := redactURLToken(text); r != text {
tok.text = r
changed = true
}
continue
}
if r, ok := redactAssignments(text); ok {
tok.text = r
changed = true
}
}
if changed {
var b strings.Builder
b.Grow(len(display))
last := 0
for _, tok := range tokens {
b.WriteString(display[last:tok.start])
b.WriteString(tok.text)
last = tok.end
}
b.WriteString(display[last:])
display = b.String()
}
// Flattening argv changes how quotes group the displayed text. Redact
// that representation too, after protecting values with known argv
// boundaries. The display has no NULs, so this cannot recurse again.
if strings.IndexByte(s, 0) >= 0 {
return RedactCommandLine(display)
}
return display
}
type cmdToken struct {
start, end int
text string
argv bool
}
func splitTokens(s string) (string, []cmdToken) {
if strings.IndexByte(s, 0) >= 0 {
for strings.HasSuffix(s, "\x00") {
s = strings.TrimSuffix(s, "\x00")
}
display := strings.ReplaceAll(s, "\x00", " ")
var tokens []cmdToken
start := 0
for start <= len(s) {
end := strings.IndexByte(s[start:], 0)
if end < 0 {
end = len(s)
} else {
end += start
}
if end > start {
tokens = append(tokens, cmdToken{start: start, end: end, text: display[start:end], argv: true})
}
if end == len(s) {
break
}
start = end + 1
}
return display, tokens
}
var tokens []cmdToken
start := -1
var quote byte
escaped := false
for i := 0; i < len(s); i++ {
c := s[i]
if escaped {
escaped = false
continue
}
if c == '\\' {
escaped = true
if start < 0 {
start = i
}
continue
}
if quote != 0 {
if c == quote {
quote = 0
}
continue
}
if c == '\'' || c == '"' {
quote = c
if start < 0 {
start = i
}
continue
}
if c == ' ' || c == '\t' || c == '\n' || c == '\r' || c == '\v' || c == '\f' {
if start >= 0 {
tokens = append(tokens, cmdToken{start: start, end: i, text: s[start:i]})
start = -1
}
continue
}
if start < 0 {
start = i
}
}
if start >= 0 {
tokens = append(tokens, cmdToken{start: start, end: len(s), text: s[start:]})
}
return s, tokens
}
var mysqlFamily = map[string]bool{
"mysql": true, "mysqldump": true, "mysqladmin": true, "mysqlimport": true,
"mysqlcheck": true, "mysqlshow": true, "mysqlbinlog": true, "mysqlpump": true,
"mysqlslap": true, "mariadb": true, "mariadb-dump": true, "mariadb-admin": true,
"mariadb-import": true, "mariadb-check": true, "mariadb-show": true,
"mariadb-binlog": true, "mydumper": true, "myloader": true,
}
// separatedSecretFlags take their secret as the following argument.
var separatedSecretFlags = map[string]bool{
"--password": true, "--passwd": true, "--pass": true, "--pw": true, "-pw": true,
"--token": true, "--secret": true, "--api-key": true, "--apikey": true,
"--access-token": true, "--auth-token": true, "--client-secret": true,
}
// separatedPairFlags take a user:password (or user%password) pair next.
var separatedPairFlags = map[string]bool{
"-u": true, "--user": true, "-U": true, "--username": true, "--credentials": true,
}
// redactAssignments applies the KEY=VALUE rules to a token, treating a
// form-encoded token (a=1&b=2) segment by segment so the other fields
// survive. An empty value is left alone.
func redactAssignments(text string) (string, bool) {
segments := strings.Split(text, "&")
changed := false
for i, seg := range segments {
eq := strings.IndexByte(seg, '=')
if eq <= 0 || eq == len(seg)-1 {
continue
}
key := strings.TrimLeft(seg[:eq], "-")
value := seg[eq+1:]
if value == redactedToken {
continue
}
if sensitiveKey(key) {
segments[i] = seg[:eq+1] + redactedToken
changed = true
continue
}
if pairKey(key) {
if r, ok := redactPair(value); ok {
segments[i] = seg[:eq+1] + r
changed = true
}
}
}
if !changed {
return text, false
}
return strings.Join(segments, "&"), true
}
func pairKey(key string) bool {
switch strings.ToLower(key) {
case "u", "user", "username", "credentials", "creds", "userpwd", "auth":
return true
}
return false
}
// redactPair masks the password half of user:password or user%password.
func redactPair(v string) (string, bool) {
sep := strings.IndexAny(v, ":%")
if sep <= 0 || sep == len(v)-1 {
return v, false
}
return v[:sep+1] + redactedToken, true
}
// sensitiveKey reports whether an option or variable name denotes a secret.
func sensitiveKey(key string) bool {
lower := strings.ToLower(key)
parts := strings.FieldsFunc(lower, func(r rune) bool { return r == '_' || r == '-' || r == '.' })
hasKey := false
hasKeyQualifier := false
for i, p := range parts {
switch {
case p == "password", strings.HasSuffix(p, "password"),
p == "passwd", strings.HasSuffix(p, "passwd"),
p == "pass" && i == len(parts)-1, p == "pwd", strings.HasSuffix(p, "pwd"),
p == "secret", strings.HasSuffix(p, "secret"),
p == "token", strings.HasSuffix(p, "token"),
p == "apikey", p == "credential", p == "credentials":
return true
case p == "key":
hasKey = true
case p == "api", p == "access", p == "private", p == "auth", p == "license", p == "signing":
hasKeyQualifier = true
}
}
return hasKey && hasKeyQualifier
}
// redactURLToken masks the password in userinfo and the values of sensitive
// query parameters inside a URL token.
func redactURLToken(tok string) string {
schemeEnd := strings.Index(tok, "://")
if schemeEnd < 0 {
return tok
}
rest := tok[schemeEnd+3:]
authorityEnd := strings.IndexAny(rest, "/?#")
if authorityEnd < 0 {
authorityEnd = len(rest)
}
authority := rest[:authorityEnd]
tail := rest[authorityEnd:]
if at := strings.LastIndexByte(authority, '@'); at > 0 {
userinfo := authority[:at]
if colon := strings.IndexByte(userinfo, ':'); colon >= 0 && colon < len(userinfo)-1 {
authority = userinfo[:colon+1] + redactedToken + authority[at:]
}
}
if q := strings.IndexByte(tail, '?'); q >= 0 {
query := tail[q+1:]
fragment := ""
if h := strings.IndexByte(query, '#'); h >= 0 {
fragment = query[h:]
query = query[:h]
}
params := strings.Split(query, "&")
for i, p := range params {
eq := strings.IndexByte(p, '=')
if eq <= 0 || eq == len(p)-1 {
continue
}
if sensitiveKey(p[:eq]) {
params[i] = p[:eq+1] + redactedToken
}
}
tail = tail[:q+1] + strings.Join(params, "&") + fragment
}
return tok[:schemeEnd+3] + authority + tail
}
package alert
import "github.com/pidginhost/csm/internal/config"
// EmitForTest drives emitAudit with no audit sinks so callers (often in
// other packages' tests) can trigger observer fan-out without setting up
// jsonl/syslog. Lives in a non-test file because Go tests cannot import
// _test.go symbols across packages.
func EmitForTest(f Finding) {
emitAudit(&config.Config{Hostname: "test"}, []Finding{f})
}
package alert
import (
"bytes"
"crypto/hmac"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"net/http"
"os"
"time"
"github.com/pidginhost/csm/internal/config"
)
// webhookTransport is the shared *http.Transport reused across every
// webhook dispatch. Reusing one transport keeps the underlying TCP /
// TLS connections in its keepalive pool so hosts that fire hundreds
// of webhooks per hour avoid a new handshake per alert. Per-call
// timeouts stay configurable via httpClient: each call wraps the
// shared transport in a fresh *http.Client carrying the requested
// timeout. http.DefaultTransport already configures sensible
// defaults; reuse it directly rather than instantiating a separate
// pool that would shadow Go's HTTP/2 / proxy plumbing.
var webhookTransport http.RoundTripper = http.DefaultTransport
const maxWebhookResponseDrainBytes int64 = 512 << 10
const maxWebhookResponseDrainDuration = 250 * time.Millisecond
// httpClient returns a webhook client with the requested timeout
// backed by the shared transport, so the keepalive pool is shared
// across dispatches without losing per-call timeout configurability.
func httpClient(timeout time.Duration) *http.Client {
return &http.Client{Timeout: timeout, Transport: webhookTransport}
}
// SetWebhookTransportForTest lets tests inject a fake RoundTripper.
// Not safe for concurrent calls; tests should set up before parallel
// dispatch and restore after.
func SetWebhookTransportForTest(rt http.RoundTripper) (restore func()) {
prev := webhookTransport
webhookTransport = rt
return func() { webhookTransport = prev }
}
func closeWebhookResponseBody(resp *http.Response) {
if resp == nil || resp.Body == nil {
return
}
if resp.Close || resp.ContentLength == 0 || resp.ContentLength > maxWebhookResponseDrainBytes {
_ = resp.Body.Close()
return
}
done := make(chan struct{})
go func() {
// Read one sentinel byte past the reuse limit so an exactly-at-limit
// response still reaches the underlying EOF and can be pooled.
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, maxWebhookResponseDrainBytes+1))
close(done)
}()
select {
case <-done:
case <-time.After(maxWebhookResponseDrainDuration):
// A slow or streaming response is not worth holding alert dispatch
// open just to preserve a keepalive connection.
}
_ = resp.Body.Close()
}
func SendWebhook(cfg *config.Config, subject, body string) error {
url := cfg.Alerts.Webhook.URL
if url == "" {
return fmt.Errorf("no webhook URL configured")
}
var payload []byte
var err error
switch cfg.Alerts.Webhook.Type {
case "slack":
payload, err = json.Marshal(map[string]string{
"text": fmt.Sprintf("*%s*\n```\n%s\n```", subject, body),
})
case "discord":
payload, err = json.Marshal(map[string]string{
"content": fmt.Sprintf("**%s**\n```\n%s\n```", subject, body),
})
default:
payload, err = json.Marshal(map[string]string{
"subject": subject,
"body": body,
})
}
if err != nil {
return fmt.Errorf("marshaling webhook payload: %w", err)
}
client := httpClient(10 * time.Second)
resp, err := client.Post(url, "application/json", bytes.NewReader(payload))
if err != nil {
return fmt.Errorf("webhook POST: %w", err)
}
defer closeWebhookResponseBody(resp)
if resp.StatusCode >= 400 {
return fmt.Errorf("webhook returned %d", resp.StatusCode)
}
return nil
}
// SendWebhookJSON posts a pre-built JSON payload to the configured webhook
// URL. Senders that need a structured body (not the slack/discord subject
// envelope) use this. No-op when no URL is configured.
func SendWebhookJSON(cfg *config.Config, payload any) error {
url := cfg.Alerts.Webhook.URL
if url == "" {
return nil
}
body, err := json.Marshal(payload)
if err != nil {
return fmt.Errorf("marshaling webhook payload: %w", err)
}
client := httpClient(10 * time.Second)
resp, err := client.Post(url, "application/json", bytes.NewReader(body))
if err != nil {
return fmt.Errorf("webhook POST: %w", err)
}
defer closeWebhookResponseBody(resp)
if resp.StatusCode >= 400 {
return fmt.Errorf("webhook returned %d", resp.StatusCode)
}
return nil
}
// SendPhpanelWebhookFinding posts a single finding to the configured phpanel
// endpoint, signing the body with HMAC-SHA256 in X-CSM-Signature. Stateless;
// caller is responsible for filtering / batching.
func SendPhpanelWebhookFinding(cfg *config.Config, f Finding) error {
delivery := phpanelDeliveryConfig{
hostname: cfg.Hostname,
url: cfg.Alerts.Webhook.URL,
hmacSecret: cfg.Alerts.Webhook.HMACSecret,
hmacSecretEnv: cfg.Alerts.Webhook.HMACSecretEnv,
}
return sendQueuedPhpanelWebhookFinding(delivery, queuedPhpanelFinding{Finding: f, Timestamp: time.Now().UTC()})
}
func sendQueuedPhpanelWebhookFinding(delivery phpanelDeliveryConfig, queued queuedPhpanelFinding) error {
cfg := &config.Config{Hostname: delivery.hostname}
cfg.Alerts.Webhook.URL = delivery.url
cfg.Alerts.Webhook.HMACSecret = delivery.hmacSecret
cfg.Alerts.Webhook.HMACSecretEnv = delivery.hmacSecretEnv
if cfg.Alerts.Webhook.URL == "" {
return fmt.Errorf("phpanel webhook URL not set")
}
secret := phpanelWebhookSecret(cfg)
if secret == "" {
return fmt.Errorf("phpanel webhook HMAC secret not configured")
}
payload := map[string]interface{}{
"hostname": cfg.Hostname,
"timestamp": queued.Timestamp.Format(time.RFC3339),
"finding": queued.Finding,
}
body, err := json.Marshal(payload)
if err != nil {
return fmt.Errorf("marshaling phpanel payload: %w", err)
}
mac := hmac.New(sha256.New, []byte(secret))
mac.Write(body)
sig := "sha256=" + hex.EncodeToString(mac.Sum(nil))
req, err := http.NewRequest(http.MethodPost, cfg.Alerts.Webhook.URL, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-CSM-Signature", sig)
req.Header.Set("X-CSM-Hostname", cfg.Hostname)
req.Header.Set("User-Agent", "csm")
client := httpClient(10 * time.Second)
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("phpanel webhook POST: %w", err)
}
defer closeWebhookResponseBody(resp)
if resp.StatusCode >= 400 {
return fmt.Errorf("phpanel webhook HTTP %d", resp.StatusCode)
}
return nil
}
func phpanelWebhookSecret(cfg *config.Config) string {
if cfg.Alerts.Webhook.HMACSecretEnv != "" {
if v := os.Getenv(cfg.Alerts.Webhook.HMACSecretEnv); v != "" {
return v
}
}
return cfg.Alerts.Webhook.HMACSecret
}
// Package atomicio implements atomic file writes used by state-bearing
// callers (firewall engine, autoblock tracker, etc.). The package is a
// dependency leaf: it imports only the standard library so any caller
// can use AtomicWriteJSON without risking an import cycle through the
// existing state / store packages.
package atomicio
import (
"encoding/json"
"fmt"
"io"
"os"
"path/filepath"
)
// AtomicWriteJSON marshals v to JSON and writes it to path atomically:
// MarshalIndent, write tmp, fsync tmp, rename. Returns the first error
// encountered. On rename failure the tmp file is removed best-effort.
//
// Reserved for state files that callers re-read on the next startup or
// the next tick - a torn write would leave the daemon with stale or
// corrupt state. Hot-path callers that only need a best-effort cache
// dump should not use this helper.
func AtomicWriteJSON(path string, perm os.FileMode, v any) error {
data, err := json.MarshalIndent(v, "", " ")
if err != nil {
return fmt.Errorf("marshal: %w", err)
}
legacyTmp := path + ".tmp"
if removeErr := os.Remove(legacyTmp); removeErr != nil && !os.IsNotExist(removeErr) {
return fmt.Errorf("remove stale tmp: %w", removeErr)
}
return atomicWrite(path, perm, data)
}
// AtomicWrite writes already-serialized bytes to path atomically with the
// same write-tmp, fsync, rename, dir-fsync sequence as AtomicWriteJSON.
func AtomicWrite(path string, perm os.FileMode, data []byte) error {
return atomicWrite(path, perm, data)
}
func atomicWrite(path string, perm os.FileMode, data []byte) error {
dir := filepath.Dir(path)
// #nosec G304 -- caller owns the destination path; tmp lives in
// the same operator-owned state directory.
f, err := os.CreateTemp(dir, "."+filepath.Base(path)+".*.tmp")
if err != nil {
return fmt.Errorf("open tmp: %w", err)
}
tmp := f.Name()
if err := f.Chmod(perm); err != nil {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("chmod tmp: %w", err)
}
if n, err := f.Write(data); err != nil {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("write tmp: %w", err)
} else if n != len(data) {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("write tmp: %w", io.ErrShortWrite)
}
if err := f.Sync(); err != nil {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("fsync tmp: %w", err)
}
if err := f.Close(); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("close tmp: %w", err)
}
if err := os.Rename(tmp, path); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("rename: %w", err)
}
// #nosec G304 -- dir is filepath.Dir of caller-owned path; opened
// read-only solely to fsync the directory after rename so the
// new dentry survives a power-loss.
d, openErr := os.Open(dir)
if openErr != nil {
return fmt.Errorf("open dir: %w", openErr)
}
if err := d.Sync(); err != nil {
_ = d.Close()
return fmt.Errorf("fsync dir: %w", err)
}
if err := d.Close(); err != nil {
return fmt.Errorf("close dir: %w", err)
}
return nil
}
package attackdb
import (
"bufio"
"fmt"
"io"
"maps"
"net"
"os"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/netutil"
"github.com/pidginhost/csm/internal/store"
)
// AttackType categorises observed attacks for grouping and scoring.
type AttackType string
const (
AttackBruteForce AttackType = "brute_force"
AttackWAFBlock AttackType = "waf_block"
AttackWebshell AttackType = "webshell"
AttackPhishing AttackType = "phishing"
AttackC2 AttackType = "c2"
AttackRecon AttackType = "recon"
AttackSPAM AttackType = "spam"
AttackCPanelLogin AttackType = "cpanel_login"
AttackFileUpload AttackType = "file_upload"
// AttackAuthSuccess marks an event that followed a SUCCESSFUL
// authentication. Recorded for correlation; carries no score.
AttackAuthSuccess AttackType = "auth_success"
AttackReputation AttackType = "reputation"
AttackOther AttackType = "other"
)
// attackTypeLabels is how the Web UI names each attack type.
var attackTypeLabels = map[AttackType]string{
AttackBruteForce: "Brute Force",
AttackWAFBlock: "WAF Block",
AttackWebshell: "Web Shell",
AttackPhishing: "Phishing",
AttackC2: "C2 Communication",
AttackRecon: "Reconnaissance",
AttackSPAM: "Spam",
AttackCPanelLogin: "cPanel Login",
AttackFileUpload: "File Upload",
AttackAuthSuccess: "Authenticated Activity",
AttackReputation: "Known Malicious IP",
AttackOther: "Other",
}
// AttackTypeLabels returns the display label of every attack type, keyed
// by the type's string value.
func AttackTypeLabels() map[string]string {
out := make(map[string]string, len(attackTypeLabels))
for typ, label := range attackTypeLabels {
out[string(typ)] = label
}
return out
}
// checkToAttack maps alert.Finding.Check values to attack types.
var checkToAttack = map[string]AttackType{
// Brute force
"wp_login_bruteforce": AttackBruteForce,
"xmlrpc_abuse": AttackBruteForce,
"ftp_bruteforce": AttackBruteForce,
"ssh_login_unknown_ip": AttackBruteForce,
"webmail_bruteforce": AttackBruteForce,
"api_auth_failure": AttackBruteForce,
"api_auth_failure_realtime": AttackBruteForce,
"ftp_auth_failure_realtime": AttackBruteForce,
"email_auth_failure_realtime": AttackBruteForce,
"credential_stuffing": AttackBruteForce,
"pam_bruteforce": AttackBruteForce,
"smtp_bruteforce": AttackBruteForce,
"smtp_probe_abuse": AttackBruteForce,
"smtp_subnet_spray": AttackBruteForce,
"mail_bruteforce": AttackBruteForce,
"mail_subnet_spray": AttackBruteForce,
"mail_account_compromised": AttackBruteForce,
"admin_panel_bruteforce": AttackBruteForce,
// Webshells and malware
"webshell": AttackWebshell,
"new_webshell_file": AttackWebshell,
"obfuscated_php": AttackWebshell,
"suspicious_php_content": AttackWebshell,
"new_php_in_languages": AttackWebshell,
"new_php_in_upgrade": AttackWebshell,
"backdoor_binary": AttackWebshell,
"new_executable_in_config": AttackWebshell,
// Phishing
"phishing_page": AttackPhishing,
"phishing_php": AttackPhishing,
"phishing_iframe": AttackPhishing,
"phishing_redirector": AttackPhishing,
"phishing_credential_log": AttackPhishing,
"phishing_kit_archive": AttackPhishing,
"phishing_directory": AttackPhishing,
// C2 and suspicious processes
"fake_kernel_thread": AttackC2,
"suspicious_process": AttackC2,
"php_suspicious_execution": AttackC2,
"user_outbound_connection": AttackC2,
"exfiltration_paste_site": AttackC2,
// Recon
"wp_user_enumeration": AttackRecon,
// HTTP-layer abuse from a single source IP: URL enumeration, request
// floods, and scanner UA spoofing. Recorded so repeat offenders build
// local reputation; recon-classed (volume-scored, no brute-force bonus).
// "http_distributed_flood" is intentionally excluded: it is an aggregate
// per-vhost finding with no single source IP, so RecordFinding has nothing
// to attribute it to.
"http_scanner_profile": AttackRecon,
"http_request_flood": AttackRecon,
"http_claimed_bot_unverified": AttackRecon,
"http_ua_spoof": AttackRecon,
// SPAM
"mail_per_account": AttackSPAM,
"exim_frozen_realtime": AttackSPAM,
// WAF: no emitted check maps here today. The two names this table once
// listed were never emitted by any release, so WAF blocks have never
// built local reputation through this database; mapping the real
// ModSecurity block names is a scoring decision recorded in the roadmap.
// cPanel/webmail login
// Successful, post-authentication events. They are RECORDED, because they
// are evidence when correlated with other findings on the same account,
// but they carry no attack weight: scoring them made an account owner an
// attacker for using cPanel, FTP or File Manager normally. One successful
// File Manager upload alone added 20 points that never decayed, and the
// resulting score fed the reputation path that kept re-blocking the owner.
//
// cpanel_multi_ip_login stays a real attack type: several addresses inside
// a window is correlation evidence rather than one successful login.
"cpanel_login": AttackAuthSuccess,
"cpanel_login_realtime": AttackAuthSuccess,
"webmail_login_realtime": AttackAuthSuccess,
"ftp_login": AttackAuthSuccess,
"pam_login": AttackAuthSuccess,
"cpanel_file_upload_realtime": AttackAuthSuccess,
"cpanel_multi_ip_login": AttackCPanelLogin,
// Reputation - known malicious IPs from threat database
"ip_reputation": AttackReputation,
// NOTE: "local_threat_score" is intentionally excluded - it is a derived
// finding, not a raw attack. Recording it would create a feedback loop
// that inflates EventCount by +1 every 10-minute cycle.
}
// MappedChecks lists every check name the attack database records, sorted.
// The slice is the caller's own copy.
func MappedChecks() []string {
out := make([]string, 0, len(checkToAttack))
for name := range checkToAttack {
out = append(out, name)
}
sort.Strings(out)
return out
}
// AttackTypeFor reports the attack type a check name records under, and
// whether the name is mapped at all.
func AttackTypeFor(check string) (AttackType, bool) {
kind, ok := checkToAttack[config.CanonicalCheckName(check)]
return kind, ok
}
// Event is a single observed attack incident.
type Event struct {
Timestamp time.Time `json:"ts"`
IP string `json:"ip"`
AttackType AttackType `json:"type"`
CheckName string `json:"check"`
Severity int `json:"sev"`
Account string `json:"account,omitempty"`
Message string `json:"msg,omitempty"`
}
// IPRecord is the per-IP aggregated intelligence record.
type IPRecord struct {
IP string `json:"ip"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
EventCount int `json:"event_count"`
AttackCounts map[AttackType]int `json:"attack_counts"`
Accounts map[string]int `json:"accounts"`
AuthSuccessAccounts map[string]int `json:"auth_success_accounts,omitempty"`
ThreatScore int `json:"threat_score"`
AutoBlocked bool `json:"auto_blocked"`
BruteForceWindowStart time.Time `json:"brute_force_window_start,omitzero"`
BruteForceWindowCount int `json:"brute_force_window_count,omitempty"`
BruteForceSustainedAt time.Time `json:"brute_force_sustained_at,omitzero"`
}
// DB is the in-memory attack database backed by JSON files.
type DB struct {
flushMu sync.Mutex
mu sync.RWMutex
records map[string]*IPRecord
deletedIPs map[string]struct{}
dirtyIPs map[string]struct{}
pendingEvents []Event
eventHealthOnce sync.Once
eventQueue *eventQueue
openEvents func(string) (io.WriteCloser, error)
recordHealthOnce sync.Once
recordQueue *recordQueue
saveRecord func(*store.DB, store.IPRecord) error
deleteRecord func(*store.DB, string) error
writeRecords func(string, []byte) error
dbPath string
dirty bool
stopCh chan struct{}
wg sync.WaitGroup
}
// markDirtyLocked records that ip's record changed and must be persisted on the
// next flush. The caller must hold db.mu. dirty stays as the flush gate so
// Flush keeps its existing "anything to write?" check.
func (db *DB) markDirtyLocked(ip string) {
if db.dirtyIPs == nil {
db.dirtyIPs = make(map[string]struct{})
}
db.dirtyIPs[ip] = struct{}{}
db.dirty = true
db.queueRecordLocked(ip)
}
func (db *DB) markDeletedLocked(ip string) {
if db.deletedIPs == nil {
db.deletedIPs = make(map[string]struct{})
}
db.deletedIPs[ip] = struct{}{}
delete(db.dirtyIPs, ip)
db.dirty = true
db.queueRecordLocked(ip)
}
var (
globalDB *DB
globalMu sync.Mutex
dbInitOnce sync.Once
)
// Init initializes the global attack database.
func Init(statePath string) *DB {
dbInitOnce.Do(func() {
dbPath := statePath + "/attack_db"
_ = os.MkdirAll(dbPath, 0700)
db := &DB{
records: make(map[string]*IPRecord),
deletedIPs: make(map[string]struct{}),
dbPath: dbPath,
stopCh: make(chan struct{}),
}
db.load()
db.pruneExpired()
// Background saver - flush dirty records every 30 seconds
db.wg.Add(1)
go db.backgroundSaver()
globalDB = db
})
return globalDB
}
// SeedFromPermanentBlocklist imports IPs from the threat DB permanent blocklist
// into the attack database. These are IPs that already attacked and were auto-blocked.
// Only imports IPs not already in the attack DB.
func (db *DB) SeedFromPermanentBlocklist(statePath string) int {
path := statePath + "/threat_db/permanent.txt"
// #nosec G304 -- fixed filename under operator-configured statePath.
f, err := os.Open(path)
if err != nil {
return 0
}
defer func() { _ = f.Close() }()
imported := 0
now := time.Now()
scanner := bufio.NewScanner(f)
db.mu.Lock()
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
fields := strings.Fields(line)
ip := fields[0]
if net.ParseIP(ip) == nil {
continue
}
if _, exists := db.records[ip]; exists {
continue // already tracked
}
// Extract reason from comment: "1.2.3.4 # reason [date]"
reason := "auto-blocked (historical)"
if idx := strings.Index(line, "# "); idx > 0 {
reason = strings.TrimSpace(line[idx+2:])
}
db.records[ip] = &IPRecord{
IP: ip,
FirstSeen: now,
LastSeen: now,
EventCount: 1,
AttackCounts: map[AttackType]int{AttackOther: 1},
Accounts: make(map[string]int),
AutoBlocked: true,
}
db.records[ip].ThreatScore = ComputeScore(db.records[ip])
db.queueEventLocked(Event{
Timestamp: now,
IP: ip,
AttackType: AttackOther,
CheckName: "permanent_blocklist_import",
Severity: 2,
Message: truncate(reason, 200),
})
delete(db.deletedIPs, ip)
db.markDirtyLocked(ip)
imported++
}
db.mu.Unlock()
return imported
}
// Global returns the global attack database instance.
func Global() *DB {
globalMu.Lock()
defer globalMu.Unlock()
return globalDB
}
// SetGlobal overrides the global attack database. Mirrors
// store.SetGlobal: production wires globalDB exactly once via Init;
// tests use this to install a pre-seeded DB without touching the
// sync.Once-guarded Init path.
func SetGlobal(db *DB) {
globalMu.Lock()
globalDB = db
globalMu.Unlock()
}
// NewForTest builds a bare in-memory DB pre-populated with the given
// records. No backgroundSaver is started (no goroutines to clean up),
// no disk path is configured. Reserved for unit tests; production
// wiring stays on Init. Records are deep-copied so a later mutation
// to the caller's map (including its nested AttackCounts / Accounts
// maps) cannot bleed into the DB.
func NewForTest(records map[string]*IPRecord) *DB {
db := &DB{
records: make(map[string]*IPRecord, len(records)),
deletedIPs: make(map[string]struct{}),
stopCh: make(chan struct{}),
}
for k, v := range records {
db.records[k] = cloneIPRecord(v)
}
return db
}
// RecordFinding records an attack event from a finding.
// Fire-and-forget: never blocks, never panics.
func (db *DB) RecordFinding(f alert.Finding) {
attackType, ok := AttackTypeFor(f.Check)
if !ok {
return // not an attack-related check
}
ip := extractFindingIP(f)
if ip == "" {
return
}
account := extractFindingAccount(f)
// Attribute the original finding, then redact before truncation can remove
// the service tag or other context needed to recognize a credential.
f = alert.SanitizeFinding(f)
event := Event{
Timestamp: f.Timestamp,
IP: ip,
AttackType: attackType,
CheckName: f.Check,
Severity: int(f.Severity),
Account: account,
Message: truncate(f.Message, 200),
}
now := f.Timestamp
if now.IsZero() {
now = time.Now()
}
db.mu.Lock()
rec, exists := db.records[ip]
if !exists {
rec = &IPRecord{
IP: ip,
FirstSeen: now,
AttackCounts: make(map[AttackType]int),
Accounts: make(map[string]int),
}
db.records[ip] = rec
}
rec.LastSeen = now
rec.EventCount++
rec.AttackCounts[attackType]++
if tracksSustainedBruteScore(f.Check) {
updateBruteForceWindow(rec, now)
}
if account != "" {
rec.Accounts[account]++
if attackType == AttackAuthSuccess {
if rec.AuthSuccessAccounts == nil {
rec.AuthSuccessAccounts = make(map[string]int)
}
rec.AuthSuccessAccounts[account]++
}
}
rec.ThreatScore = computeScoreAt(rec, now)
db.queueEventLocked(event)
delete(db.deletedIPs, ip)
db.markDirtyLocked(ip)
db.mu.Unlock()
}
// MarkBlocked sets the auto-blocked flag on an IP record.
func (db *DB) MarkBlocked(ip string) {
db.mu.Lock()
if rec, ok := db.records[ip]; ok {
rec.AutoBlocked = true
rec.ThreatScore = ComputeScore(rec)
delete(db.deletedIPs, ip)
db.markDirtyLocked(ip)
}
db.mu.Unlock()
}
// LookupIP returns the record for an IP, or nil if not tracked.
func (db *DB) LookupIP(ip string) *IPRecord {
db.mu.RLock()
defer db.mu.RUnlock()
rec, ok := db.records[ip]
if !ok {
return nil
}
return cloneIPRecord(rec)
}
// TopAttackers returns the top N IPs by threat score.
func (db *DB) TopAttackers(n int) []*IPRecord {
db.mu.RLock()
defer db.mu.RUnlock()
all := make([]*IPRecord, 0, len(db.records))
for _, rec := range db.records {
all = append(all, cloneIPRecord(rec))
}
// Sort by threat score descending, then event count
sortRecords(all)
if n > len(all) {
n = len(all)
}
return all[:n]
}
// Flush saves all pending data to disk. Called on daemon shutdown.
func (db *DB) Flush() error {
// Keep snapshots and disk writes in the same order. A command can flush
// alongside the background saver; an older write must not undo its delete.
db.flushMu.Lock()
defer db.flushMu.Unlock()
db.mu.Lock()
events := db.pendingEvents
eventBatch := db.eventHealth().detach()
db.pendingEvents = nil
dirty := db.dirty
db.dirty = false
db.mu.Unlock()
if len(events) > 0 {
db.appendEvents(events, eventBatch)
}
if dirty {
db.saveRecords()
}
return nil
}
// Stop stops the background saver and flushes.
func (db *DB) Stop() {
close(db.stopCh)
db.wg.Wait()
_ = db.Flush()
}
func (db *DB) backgroundSaver() {
defer db.wg.Done()
ticker := time.NewTicker(30 * time.Second)
defer ticker.Stop()
for {
select {
case <-db.stopCh:
return
case <-ticker.C:
_ = db.Flush()
}
}
}
// extractIP pulls an IP address from a finding message.
func extractIP(message string) string {
for _, sep := range []string{" from ", ": ", "accessing server: "} {
if idx := strings.Index(message, sep); idx >= 0 {
rest := message[idx+len(sep):]
fields := strings.Fields(rest)
if len(fields) > 0 {
token := fields[0]
// Strip AbuseIPDB score suffix like "(AbuseIPDB"
if paren := strings.Index(token, "("); paren > 0 {
token = token[:paren]
}
if ip, ok := netutil.ParseIPToken(token); ok {
return ip
}
}
}
}
return ""
}
func extractFindingIP(f alert.Finding) string {
if ip := normalizeRecordIP(f.SourceIP); ip != "" {
return ip
}
return extractIP(f.Message)
}
func normalizeRecordIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if host, _, err := net.SplitHostPort(raw); err == nil {
raw = host
}
raw = strings.Trim(raw, "[]")
ip := net.ParseIP(raw)
if ip == nil {
return ""
}
return ip.String()
}
func extractFindingAccount(f alert.Finding) string {
mailbox := strings.TrimSpace(f.Mailbox)
domain := strings.TrimSpace(f.Domain)
if mailbox != "" {
if strings.Contains(mailbox, "@") || domain == "" {
return mailbox
}
return mailbox + "@" + strings.ToLower(domain)
}
if tenant := strings.TrimSpace(f.TenantID); tenant != "" {
return tenant
}
return extractAccount(f.Message, f.Details)
}
// extractAccount tries to pull a cPanel account name from message or details.
func extractAccount(message, details string) string {
// Check details first: "Account: username"
for _, text := range []string{details, message} {
if idx := strings.Index(text, "Account: "); idx >= 0 {
rest := text[idx+9:]
fields := strings.Fields(rest)
if len(fields) > 0 {
return fields[0]
}
}
}
// Try /home/username/ pattern
if idx := strings.Index(message, "/home/"); idx >= 0 {
rest := message[idx+6:]
if slash := strings.Index(rest, "/"); slash > 0 {
return rest[:slash]
}
}
return ""
}
func truncate(s string, n int) string {
r := []rune(s)
if len(r) <= n {
return s
}
return string(r[:n])
}
func updateBruteForceWindow(rec *IPRecord, ts time.Time) {
if rec.BruteForceWindowStart.IsZero() ||
ts.Before(rec.BruteForceWindowStart) ||
ts.Sub(rec.BruteForceWindowStart) > sustainedBruteForceWindow {
rec.BruteForceWindowStart = ts
rec.BruteForceWindowCount = 1
return
}
rec.BruteForceWindowCount++
if rec.BruteForceWindowCount >= sustainedBruteForceThreshold &&
(rec.BruteForceSustainedAt.IsZero() || !ts.Before(rec.BruteForceSustainedAt)) {
rec.BruteForceSustainedAt = ts
}
}
func tracksSustainedBruteScore(check string) bool {
return check == "email_auth_failure_realtime"
}
// RemoveIP removes an IP from the attack database entirely.
func (db *DB) RemoveIP(ip string) {
db.mu.Lock()
delete(db.records, ip)
db.markDeletedLocked(ip)
db.mu.Unlock()
}
// ForgetIP atomically removes all scoring records for a parsed IP, including
// legacy imports stored under equivalent spellings. Event history is retained.
// The returned records are detached, so later findings cannot change them.
func (db *DB) ForgetIP(ip net.IP) []*IPRecord {
db.mu.Lock()
defer db.mu.Unlock()
var removed []*IPRecord
for key, rec := range db.records {
if ip.Equal(net.ParseIP(key)) {
removed = append(removed, rec)
delete(db.records, key)
db.markDeletedLocked(key)
}
}
return removed
}
// PruneExpired removes records older than 90 days.
func (db *DB) PruneExpired() {
db.pruneExpired()
}
func (db *DB) pruneExpired() {
cutoff := time.Now().Add(-90 * 24 * time.Hour)
db.mu.Lock()
for ip, rec := range db.records {
if rec.LastSeen.Before(cutoff) {
delete(db.records, ip)
db.markDeletedLocked(ip)
}
}
db.mu.Unlock()
}
// TotalIPs returns the number of tracked IPs.
func (db *DB) TotalIPs() int {
db.mu.RLock()
defer db.mu.RUnlock()
return len(db.records)
}
// AllRecords returns a deep-copy snapshot of all records.
func (db *DB) AllRecords() []*IPRecord {
db.mu.RLock()
defer db.mu.RUnlock()
result := make([]*IPRecord, 0, len(db.records))
for _, rec := range db.records {
result = append(result, cloneIPRecord(rec))
}
return result
}
// FormatTopLine returns a summary string for stderr logging.
func (db *DB) FormatTopLine() string {
db.mu.RLock()
defer db.mu.RUnlock()
total := len(db.records)
blocked := 0
for _, r := range db.records {
if r.AutoBlocked {
blocked++
}
}
return fmt.Sprintf("%d IPs tracked, %d auto-blocked", total, blocked)
}
// Snapshots must detach all count maps from concurrent recording.
func cloneIPRecord(rec *IPRecord) *IPRecord {
cp := *rec
cp.AttackCounts = maps.Clone(rec.AttackCounts)
cp.Accounts = maps.Clone(rec.Accounts)
if cp.AttackCounts == nil {
cp.AttackCounts = make(map[AttackType]int)
}
if cp.Accounts == nil {
cp.Accounts = make(map[string]int)
}
cp.AuthSuccessAccounts = maps.Clone(rec.AuthSuccessAccounts)
return &cp
}
package attackdb
import (
"bytes"
"io"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/store"
)
// Event batches are detached together under db.mu. Health keeps only their
// counts and oldest arrival, independently of state locks and persistence I/O.
type eventQueue struct {
mu sync.Mutex
waiting int
oldest time.Time
active *eventBatch
losses *queuehealth.Tracker
uncertain bool
uncertainAt time.Time
}
type eventBatch struct {
queue *eventQueue
count, offered, saved, lost int
progress time.Time
}
func (db *DB) eventHealth() *eventQueue {
db.eventHealthOnce.Do(func() { db.eventQueue = &eventQueue{losses: queuehealth.New(0, time.Minute)} })
return db.eventQueue
}
func (db *DB) queueEventLocked(event Event) {
db.pendingEvents = append(db.pendingEvents, event)
if db.dbPath == "" && store.Global() == nil {
return
}
q := db.eventHealth()
q.mu.Lock()
if q.waiting == 0 {
q.oldest = time.Now()
}
q.waiting++
q.mu.Unlock()
}
func (q *eventQueue) detach() *eventBatch {
q.mu.Lock()
defer q.mu.Unlock()
if q.waiting == 0 {
return nil
}
batch := &eventBatch{queue: q, count: q.waiting, progress: time.Now()}
q.waiting = 0
q.oldest = time.Time{}
q.active = batch
return batch
}
func (b *eventBatch) beginWrite(count int) {
if b == nil {
return
}
b.queue.mu.Lock()
b.offered = count
b.queue.mu.Unlock()
}
// Only complete records offered to the current write can have an unknown
// outcome. Everything still buffered or not yet encoded is a confirmed loss.
func (b *eventBatch) settleInterrupted() {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
defer q.mu.Unlock()
count := b.count - b.saved - b.lost - b.offered
if count > 0 {
b.lost += count
q.losses.Lose(time.Now(), uint64(count))
}
if b.count > b.saved+b.lost {
q.uncertain = true
q.uncertainAt = time.Now()
}
}
func (b *eventBatch) advance(saved, lost int) {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
defer q.mu.Unlock()
b.offered = 0
b.saved += saved
b.lost += lost
now := time.Now()
if lost > 0 {
q.losses.Lose(now, uint64(lost))
}
b.progress = now
}
func (b *eventBatch) uncertainOutcome() {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
q.uncertain = true
q.uncertainAt = time.Now()
q.mu.Unlock()
}
func (b *eventBatch) discardRemaining() {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
defer q.mu.Unlock()
b.offered = 0
remaining := b.count - b.saved - b.lost
if remaining > 0 {
q.losses.Lose(time.Now(), uint64(remaining))
b.lost += remaining
}
}
func (b *eventBatch) finish() {
if b == nil {
return
}
b.settleInterrupted()
q := b.queue
q.mu.Lock()
q.active = nil
q.mu.Unlock()
}
func (db *DB) eventQueueStatus(now time.Time) queuehealth.Status {
q := db.eventHealth()
q.mu.Lock()
defer q.mu.Unlock()
row := q.losses.Snapshot(now)
row.CapacityUnavailable = true
row.DepthUnit = "events"
row.LagBasis = "operation_progress"
row.Depth = q.waiting
row.DroppedLowerBound = q.uncertain
if q.waiting > 0 {
row.LagSeconds = max(0, now.Sub(q.oldest).Seconds())
}
if b := q.active; b != nil {
row.InFlight = b.count
row.ProcessingSeconds = max(0, now.Sub(b.progress).Seconds())
}
// An uncertain write may have lost evidence already accepted; a backlog
// has not. The record queue orders its reasons the same way.
switch {
case !q.uncertainAt.IsZero() && now.Sub(q.uncertainAt) < time.Minute:
row.Status, row.Reason = "degraded", "persistence_uncertain"
case row.LagSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "backlog_lag"
case row.ProcessingSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "processing_lag"
}
return row
}
// The encoder escapes embedded newlines. Only a complete JSONL delimiter
// accepted by the file writer counts as a persisted event; buffered bytes do not.
type eventWriter struct {
writer io.Writer
batch *eventBatch
}
func (w eventWriter) Write(p []byte) (int, error) {
w.batch.beginWrite(bytes.Count(p, []byte{'\n'}))
n, err := w.writer.Write(p)
w.batch.advance(bytes.Count(p[:n], []byte{'\n'}), 0)
return n, err
}
// An encoder error before Write confirms this event was never submitted, even
// if a later file operation exits without returning an observable byte count.
type eventEncoderWriter struct {
writer io.Writer
called bool
}
func (w *eventEncoderWriter) Write(p []byte) (int, error) {
w.called = true
return w.writer.Write(p)
}
package attackdb
import (
"bufio"
"encoding/json"
"fmt"
"io"
"maps"
"os"
"path/filepath"
"time"
"github.com/pidginhost/csm/internal/store"
)
const (
recordsFile = "records.json"
eventsFile = "events.jsonl"
maxEventsBytes = 10 * 1024 * 1024 // 10 MB
)
// load reads IP records from the bbolt store (if available) or from
// the flat-file records.json.
func (db *DB) load() {
if sdb := store.Global(); sdb != nil {
storeRecords := sdb.LoadAllIPRecords()
db.mu.Lock()
defer db.mu.Unlock()
if db.records == nil {
db.records = make(map[string]*IPRecord, len(storeRecords))
}
for ip, sr := range storeRecords {
rec := &IPRecord{
IP: sr.IP,
FirstSeen: sr.FirstSeen,
LastSeen: sr.LastSeen,
EventCount: sr.EventCount,
ThreatScore: sr.ThreatScore,
AutoBlocked: sr.AutoBlocked,
BruteForceWindowStart: sr.BruteForceWindowStart,
BruteForceWindowCount: sr.BruteForceWindowCount,
BruteForceSustainedAt: sr.BruteForceSustainedAt,
AttackCounts: make(map[AttackType]int),
Accounts: make(map[string]int),
AuthSuccessAccounts: maps.Clone(sr.AuthSuccessAccounts),
}
for k, v := range sr.AttackCounts {
rec.AttackCounts[AttackType(k)] = v
}
for k, v := range sr.Accounts {
rec.Accounts[k] = v
}
if normalizeLoadedRecord(rec) {
db.markDirtyLocked(ip)
}
db.records[ip] = rec
}
return
}
// No configured directory means there is nowhere to persist to. Joining
// an empty dbPath yields a relative path, which would read and write
// state in whatever directory the process was started from.
if db.dbPath == "" {
return
}
// Fallback: flat-file records.json.
path := filepath.Join(db.dbPath, recordsFile)
// #nosec G304 -- filepath.Join under operator-configured db.dbPath.
data, err := os.ReadFile(path)
if err != nil {
return
}
var records map[string]*IPRecord
if err := json.Unmarshal(data, &records); err != nil {
fmt.Fprintf(os.Stderr, "attackdb: error loading %s: %v\n", path, err)
return
}
db.mu.Lock()
for ip, rec := range records {
if normalizeLoadedRecord(rec) {
db.markDirtyLocked(ip)
}
}
db.records = records
db.mu.Unlock()
}
func normalizeLoadedRecord(rec *IPRecord) bool {
changed := false
// Empty maps are an in-memory invariant; bbolt omits them, so nil-to-empty
// alone should not dirty every account-free record on each startup.
if rec.AttackCounts == nil {
rec.AttackCounts = make(map[AttackType]int)
}
if rec.Accounts == nil {
rec.Accounts = make(map[string]int)
}
bruteCount := rec.AttackCounts[AttackBruteForce]
if rec.BruteForceWindowCount > bruteCount {
rec.BruteForceWindowCount = bruteCount
changed = true
}
if rec.BruteForceSustainedAt.IsZero() &&
rec.BruteForceWindowCount >= sustainedBruteForceThreshold {
rec.BruteForceSustainedAt = rec.LastSeen
changed = true
}
score := ComputeScore(rec)
if rec.ThreatScore != score {
rec.ThreatScore = score
changed = true
}
return changed
}
// saveRecords writes records to the bbolt store (if available) or to
// the flat-file records.json.
func (db *DB) saveRecords() {
if sdb := store.Global(); sdb != nil {
// Incremental: only records changed since the last flush are
// re-serialized. On a host tracking tens of thousands of IPs the old
// full rewrite cost seconds of CPU on every 30s flush and on shutdown.
// Snapshot and clear the dirty set under the lock so concurrent
// mutations land in the fresh set for the next flush; a write that
// fails is re-marked dirty below so it retries.
db.mu.Lock()
dirty := db.dirtyIPs
db.dirtyIPs = make(map[string]struct{})
records := make([]store.IPRecord, 0, len(dirty))
for ip := range dirty {
rec, ok := db.records[ip]
if !ok {
continue // removed since marked; deletedIPs carries the removal
}
records = append(records, toStoreIPRecord(rec))
}
var deleted []string
for ip := range db.deletedIPs {
deleted = append(deleted, ip)
}
batch := db.detachRecordsLocked()
db.mu.Unlock()
defer batch.finish(db)
save := db.saveRecord
if save == nil {
save = (*store.DB).SaveIPRecord
}
remove := db.deleteRecord
if remove == nil {
remove = (*store.DB).DeleteIPRecord
}
var failed []string
for _, sr := range records {
batch.begin(sr.IP)
err := save(sdb, sr)
batch.result(sr.IP, err)
if err != nil {
fmt.Fprintf(os.Stderr, "attackdb: store save %s: %v\n", sr.IP, err)
failed = append(failed, sr.IP)
}
}
if len(deleted) > 0 {
var removed []string
var failedDeletes []string
for _, ip := range deleted {
batch.begin(ip)
err := remove(sdb, ip)
batch.result(ip, err)
if err != nil {
fmt.Fprintf(os.Stderr, "attackdb: store delete %s: %v\n", ip, err)
failedDeletes = append(failedDeletes, ip)
continue
}
removed = append(removed, ip)
}
if len(removed) > 0 {
db.mu.Lock()
for _, ip := range removed {
delete(db.deletedIPs, ip)
}
db.mu.Unlock()
}
if len(failedDeletes) > 0 {
db.mu.Lock()
db.dirty = true
db.mu.Unlock()
}
}
if len(failed) > 0 {
db.mu.Lock()
for _, ip := range failed {
db.markDirtyLocked(ip)
}
db.mu.Unlock()
}
return
}
// No configured directory means there is nowhere to persist to. Joining
// an empty dbPath yields a relative path, which would read and write
// state in whatever directory the process was started from.
if db.dbPath == "" {
return
}
// Fallback: flat-file records.json. The whole records map is rewritten
// each flush, so removals are reflected by absence and deletedIPs is
// redundant here -- but it must still be drained or it grows for the
// process lifetime on a host with no bbolt store. Snapshot under the same
// lock as the marshal and swap dirtyIPs before disk I/O, so mutations
// during the write land in a fresh set. Failed writes requeue the snapshot.
db.mu.Lock()
data, err := json.Marshal(db.records)
var drained []string
for ip := range db.deletedIPs {
drained = append(drained, ip)
}
flushedDirty := db.dirtyIPs
db.dirtyIPs = make(map[string]struct{})
batch := db.detachRecordsLocked()
db.mu.Unlock()
defer batch.finish(db)
if err != nil {
batch.resultAll(err)
fmt.Fprintf(os.Stderr, "attackdb: error marshaling records: %v\n", err)
db.requeueDirty(flushedDirty, len(drained) > 0)
return
}
path := filepath.Join(db.dbPath, recordsFile)
write := db.writeRecords
if write == nil {
write = writeRecordFile
}
batch.beginAll()
err = write(path, data)
batch.resultAll(err)
if err != nil {
fmt.Fprintf(os.Stderr, "attackdb: %v\n", err)
db.requeueDirty(flushedDirty, len(drained) > 0)
return
}
db.mu.Lock()
for _, ip := range drained {
delete(db.deletedIPs, ip)
}
db.mu.Unlock()
}
func (db *DB) requeueDirty(dirty map[string]struct{}, hasPendingDelete bool) {
if len(dirty) == 0 && !hasPendingDelete {
return
}
db.mu.Lock()
if hasPendingDelete {
db.dirty = true
}
for ip := range dirty {
db.markDirtyLocked(ip)
}
db.mu.Unlock()
}
// toStoreIPRecord projects an in-memory record into the store's persistence
// shape, copying the count maps so the store never aliases live maps.
func toStoreIPRecord(rec *IPRecord) store.IPRecord {
sr := store.IPRecord{
IP: rec.IP,
FirstSeen: rec.FirstSeen,
LastSeen: rec.LastSeen,
EventCount: rec.EventCount,
ThreatScore: rec.ThreatScore,
AutoBlocked: rec.AutoBlocked,
BruteForceWindowStart: rec.BruteForceWindowStart,
BruteForceWindowCount: rec.BruteForceWindowCount,
BruteForceSustainedAt: rec.BruteForceSustainedAt,
AttackCounts: make(map[string]int, len(rec.AttackCounts)),
Accounts: make(map[string]int, len(rec.Accounts)),
AuthSuccessAccounts: maps.Clone(rec.AuthSuccessAccounts),
}
for k, v := range rec.AttackCounts {
sr.AttackCounts[string(k)] = v
}
for k, v := range rec.Accounts {
sr.Accounts[k] = v
}
return sr
}
// appendEvents writes events to the bbolt store (if available) or appends
// to the JSONL file, rotating if needed.
func (db *DB) appendEvents(events []Event, batch *eventBatch) {
defer batch.finish()
if sdb := store.Global(); sdb != nil {
for i, ev := range events {
batch.beginWrite(1)
ts := ev.Timestamp
if ts.IsZero() {
ts = time.Now()
}
se := store.AttackEvent{
Timestamp: ts,
IP: ev.IP,
AttackType: string(ev.AttackType),
CheckName: ev.CheckName,
Severity: ev.Severity,
Account: ev.Account,
Message: ev.Message,
}
if err := sdb.RecordAttackEvent(se, i); err != nil {
batch.advance(0, 1)
fmt.Fprintf(os.Stderr, "attackdb: store event: %v\n", err)
} else {
batch.advance(1, 0)
}
}
return
}
// No configured directory means there is nowhere to persist to. Joining
// an empty dbPath yields a relative path, which would read and write
// state in whatever directory the process was started from.
if db.dbPath == "" {
batch.discardRemaining()
return
}
// Fallback: flat-file JSONL.
path := filepath.Join(db.dbPath, eventsFile)
// Check file size and rotate if needed
if info, err := os.Stat(path); err == nil && info.Size() > maxEventsBytes {
rotateEventsFile(path)
}
open := db.openEvents
if open == nil {
open = openEventsFile
}
f, err := open(path)
if err != nil {
batch.discardRemaining()
fmt.Fprintf(os.Stderr, "attackdb: error opening %s: %v\n", path, err)
return
}
defer func() {
returned := false
defer func() {
if !returned {
batch.uncertainOutcome()
}
}()
err := f.Close()
returned = true
if err != nil {
batch.uncertainOutcome()
}
}()
// Settle buffered and unencoded work before file cleanup, which may block.
// Only complete records offered to an interrupted write remain uncertain.
defer batch.settleInterrupted()
w := bufio.NewWriter(eventWriter{writer: f, batch: batch})
output := &eventEncoderWriter{writer: w}
enc := json.NewEncoder(output)
for _, ev := range events {
output.called = false
if err := enc.Encode(ev); err != nil && !output.called {
batch.advance(0, 1)
}
}
_ = w.Flush()
batch.discardRemaining()
}
// rotateEventsFile keeps the newest half of the file.
func rotateEventsFile(path string) {
// #nosec G304 -- path is filepath.Join under operator-configured db.dbPath.
data, err := os.ReadFile(path)
if err != nil {
return
}
// Find the midpoint newline
mid := len(data) / 2
for mid < len(data) {
if data[mid] == '\n' {
mid++
break
}
mid++
}
if mid >= len(data) {
return
}
tmpPath := path + ".tmp"
// #nosec G703 -- path is db.path, derived from the operator-configured
// statePath at DB open time (see Open / NewDB).
if err := os.WriteFile(tmpPath, data[mid:], 0600); err != nil {
return
}
_ = os.Rename(tmpPath, path)
}
// QueryEvents reads events for a specific IP from the bbolt store (if
// available) or from the JSONL file. Returns the most recent `limit`
// events in reverse chronological order.
func (db *DB) QueryEvents(ip string, limit int) []Event {
if sdb := store.Global(); sdb != nil {
storeEvents := sdb.QueryAttackEvents(ip, limit)
result := make([]Event, len(storeEvents))
for i, se := range storeEvents {
result[i] = Event{
Timestamp: se.Timestamp,
IP: se.IP,
AttackType: AttackType(se.AttackType),
CheckName: se.CheckName,
Severity: se.Severity,
Account: se.Account,
Message: se.Message,
}
}
return result
}
// See the note in saveRecords: an empty dbPath would resolve to a
// relative path and read an unrelated file from the working directory.
if db.dbPath == "" {
return nil
}
// Fallback: flat-file JSONL.
path := filepath.Join(db.dbPath, eventsFile)
// #nosec G304 -- filepath.Join under operator-configured db.dbPath.
f, err := os.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
var all []Event
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
for scanner.Scan() {
var ev Event
if err := json.Unmarshal(scanner.Bytes(), &ev); err != nil {
continue
}
if ev.IP == ip {
all = append(all, ev)
}
}
// Return most recent first
if len(all) > limit && limit > 0 {
all = all[len(all)-limit:]
}
// Reverse
for i, j := 0, len(all)-1; i < j; i, j = i+1, j-1 {
all[i], all[j] = all[j], all[i]
}
return all
}
func openEventsFile(path string) (io.WriteCloser, error) {
// #nosec G304 -- path is eventsFile under the operator-configured db.dbPath.
return os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600)
}
func writeRecordFile(path string, data []byte) error {
tmpPath := path + ".tmp"
if err := os.WriteFile(tmpPath, data, 0600); err != nil {
return fmt.Errorf("error writing %s: %w", tmpPath, err)
}
if err := os.Rename(tmpPath, path); err != nil {
return fmt.Errorf("error renaming %s: %w", path, err)
}
return nil
}
package attackdb
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/store"
)
type recordPhase uint8
const (
recordQueued recordPhase = iota
recordWriting
recordSucceeded
recordFailed
)
type recordDemand struct {
arrival time.Time
failed bool
phase recordPhase
}
type recordQueue struct {
mu sync.Mutex
waiting map[string]recordDemand
active *recordBatch
losses *queuehealth.Tracker
uncertain bool
uncertainAt time.Time
}
type recordBatch struct {
queue *recordQueue
items map[string]recordDemand
progress time.Time
}
func (db *DB) recordHealth() *recordQueue {
db.recordHealthOnce.Do(func() {
db.recordQueue = &recordQueue{waiting: make(map[string]recordDemand), losses: queuehealth.New(0, time.Minute)}
})
return db.recordQueue
}
func (db *DB) queueRecordLocked(ip string) {
if db.dbPath == "" && store.Global() == nil {
return
}
q := db.recordHealth()
q.mu.Lock()
if _, exists := q.waiting[ip]; !exists {
q.waiting[ip] = recordDemand{arrival: time.Now()}
}
q.mu.Unlock()
}
// The caller detaches the actual dirty and deletion snapshot under db.mu.
// Subsequent mutations then own a separate pending generation for the same IP.
func (db *DB) detachRecordsLocked() *recordBatch {
q := db.recordHealth()
q.mu.Lock()
defer q.mu.Unlock()
if len(q.waiting) == 0 {
return nil
}
for ip, demand := range q.waiting {
demand.phase = recordQueued
q.waiting[ip] = demand
}
b := &recordBatch{queue: q, items: q.waiting, progress: time.Now()}
q.waiting = make(map[string]recordDemand)
q.active = b
return b
}
func (b *recordBatch) begin(ip string) {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
if demand, exists := b.items[ip]; exists {
demand.phase = recordWriting
b.items[ip] = demand
}
q.mu.Unlock()
}
func (b *recordBatch) beginAll() {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
for ip, demand := range b.items {
demand.phase = recordWriting
b.items[ip] = demand
}
q.mu.Unlock()
}
func (b *recordBatch) result(ip string, err error) {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
if demand, exists := b.items[ip]; exists {
demand.failed = err != nil
demand.phase = recordSucceeded
if err != nil {
demand.phase = recordFailed
}
b.items[ip] = demand
}
b.progress = time.Now()
q.mu.Unlock()
}
func (b *recordBatch) resultAll(err error) {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
for ip, demand := range b.items {
demand.failed = err != nil
demand.phase = recordSucceeded
if err != nil {
demand.phase = recordFailed
}
b.items[ip] = demand
}
b.progress = time.Now()
q.mu.Unlock()
}
// Reconcile at the real requeue boundary. An older successful delete may also
// satisfy a newer delete, while an intervening record mutation must stay queued.
func (b *recordBatch) finish(db *DB) {
if b == nil {
return
}
q := b.queue
// Publish unreturned I/O before retry reconciliation can wait for db.mu.
q.mu.Lock()
for _, demand := range b.items {
if demand.phase == recordWriting {
q.uncertain = true
q.uncertainAt = time.Now()
}
}
q.mu.Unlock()
db.mu.Lock()
defer db.mu.Unlock()
q.mu.Lock()
defer q.mu.Unlock()
pending := func(ip string) bool {
_, dirty := db.dirtyIPs[ip]
_, deleted := db.deletedIPs[ip]
return dirty || deleted
}
for ip := range q.waiting {
if !pending(ip) {
delete(q.waiting, ip)
}
}
for ip, demand := range b.items {
if !pending(ip) {
if demand.phase == recordQueued || demand.phase == recordFailed {
q.losses.Lose(time.Now(), 1)
}
continue
}
next, exists := q.waiting[ip]
if !exists || demand.failed && demand.arrival.Before(next.arrival) {
next.arrival = demand.arrival
}
next.failed = next.failed || demand.failed
q.waiting[ip] = next
}
q.active = nil
}
func (q *recordQueue) snapshot(now time.Time) queuehealth.Status {
q.mu.Lock()
defer q.mu.Unlock()
row := q.losses.Snapshot(now)
row.CapacityUnavailable = true
row.DepthUnit = "records"
row.LagBasis = "operation_progress"
row.Depth = len(q.waiting)
row.DroppedLowerBound = q.uncertain
failed := false
for _, demand := range q.waiting {
row.LagSeconds = max(row.LagSeconds, now.Sub(demand.arrival).Seconds())
failed = failed || demand.failed
}
if b := q.active; b != nil {
row.InFlight = len(b.items)
row.ProcessingSeconds = max(0, now.Sub(b.progress).Seconds())
for _, demand := range b.items {
failed = failed || demand.failed
}
}
switch {
case !q.uncertainAt.IsZero() && now.Sub(q.uncertainAt) < time.Minute:
row.Status, row.Reason = "degraded", "persistence_uncertain"
case failed:
row.Status, row.Reason = "degraded", "retry_failed"
case row.LagSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "backlog_lag"
case row.ProcessingSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "processing_lag"
}
return row
}
// QueueStatuses never acquires database state or filesystem locks.
func (db *DB) QueueStatuses(now time.Time) map[string]queuehealth.Status {
return map[string]queuehealth.Status{"events": db.eventQueueStatus(now), "records": db.recordHealth().snapshot(now)}
}
package attackdb
import (
"sort"
"time"
)
const (
sustainedBruteForceThreshold = 50
sustainedBruteForceWindow = 30 * time.Minute
)
// ComputeScore returns a 0-100 local threat score from an IPRecord.
//
// Scoring logic:
// - Volume: min(attack event count * 2, 30); authenticated activity excluded
// - Attack type bonuses (non-cumulative per type)
// - Multi-account targeting: +10
// - Auto-blocked floor: 50
// - Hard cap: 100
func ComputeScore(r *IPRecord) int {
return computeScoreAt(r, time.Now())
}
func computeScoreAt(r *IPRecord, now time.Time) int {
score := 0
// Audit events stay in the record and history, but cannot increase the
// volume score or turn successful access into multi-account targeting.
attackEvents := max(0, r.EventCount-r.AttackCounts[AttackAuthSuccess])
vol := attackEvents * 2
if vol > 30 {
vol = 30
}
score += vol
// Attack type bonuses
if r.AttackCounts[AttackC2] > 0 {
score += 35
}
if r.AttackCounts[AttackWebshell] > 0 {
score += 30
}
if r.AttackCounts[AttackPhishing] > 0 {
score += 25
}
if r.AttackCounts[AttackBruteForce] > 0 {
score += 15
}
// The sustained-brute tier is rate-bound and tied to the raw mail-auth
// signal so stale passwords and unrelated brute-force checks cannot become
// block-eligible by slowly accumulating failures over retention.
if hasSustainedBruteForce(r, now) {
score += 30
}
if r.AttackCounts[AttackWAFBlock] > 5 {
score += 10
}
if r.AttackCounts[AttackFileUpload] > 0 {
score += 20
}
// Count accounts with attack evidence, including accounts that also
// have successful activity from this address.
targetedAccounts := 0
for account, count := range r.Accounts {
if count > r.AuthSuccessAccounts[account] {
targetedAccounts++
}
}
if targetedAccounts > 1 && attackEvents > 0 {
score += 10
}
// Auto-blocked floor
if r.AutoBlocked && score < 50 {
score = 50
}
// Cap at 100
if score > 100 {
score = 100
}
return score
}
func hasSustainedBruteForce(r *IPRecord, now time.Time) bool {
if r.AttackCounts[AttackBruteForce] < sustainedBruteForceThreshold ||
r.BruteForceSustainedAt.IsZero() {
return false
}
return !r.BruteForceSustainedAt.Before(now.Add(-sustainedBruteForceWindow))
}
// sortRecords sorts by threat score descending, then event count descending.
func sortRecords(recs []*IPRecord) {
sort.Slice(recs, func(i, j int) bool {
if recs[i].ThreatScore != recs[j].ThreatScore {
return recs[i].ThreatScore > recs[j].ThreatScore
}
return recs[i].EventCount > recs[j].EventCount
})
}
package attackdb
import (
"bufio"
"encoding/json"
"os"
"path/filepath"
"sync"
"time"
"github.com/pidginhost/csm/internal/store"
)
const statsCacheTTL = 30 * time.Second
// AttackStats contains aggregate statistics for the API and dashboard.
type AttackStats struct {
TotalIPs int `json:"total_ips"`
TotalEvents int `json:"total_events"`
Last24hEvents int `json:"last_24h_events"`
Last7dEvents int `json:"last_7d_events"`
BlockedIPs int `json:"blocked_ips"`
ByType map[AttackType]int `json:"by_type"` // lifetime, aggregated from IPRecord.AttackCounts
ByType24h map[AttackType]int `json:"by_type_24h"` // last 24h, aggregated from the events log
TopAttackers []*IPRecord `json:"top_attackers"`
HourlyBuckets [24]int `json:"hourly_buckets"` // last 24h, index 0 = oldest hour
DailyBuckets [7]int `json:"daily_buckets"` // last 7 days, index 0 = oldest day
}
var (
cachedStats AttackStats
cachedStatsTime time.Time
cachedStatsMu sync.Mutex
)
// Stats returns aggregate statistics, cached for 30 seconds to avoid
// re-scanning the full events.jsonl on every API call.
func (db *DB) Stats() AttackStats {
cachedStatsMu.Lock()
if time.Since(cachedStatsTime) < statsCacheTTL {
s := cachedStats
cachedStatsMu.Unlock()
return s
}
cachedStatsMu.Unlock()
stats := db.computeStats()
cachedStatsMu.Lock()
cachedStats = stats
cachedStatsTime = time.Now()
cachedStatsMu.Unlock()
return stats
}
func (db *DB) computeStats() AttackStats {
now := time.Now()
cutoff24h := now.Add(-24 * time.Hour)
cutoff7d := now.Add(-7 * 24 * time.Hour)
db.mu.RLock()
stats := AttackStats{
TotalIPs: len(db.records),
ByType: make(map[AttackType]int),
ByType24h: make(map[AttackType]int),
}
for _, rec := range db.records {
stats.TotalEvents += rec.EventCount
if rec.AutoBlocked {
stats.BlockedIPs++
}
for atype, count := range rec.AttackCounts {
stats.ByType[atype] += count
}
}
db.mu.RUnlock()
// Compute time-based stats from events log
events := db.readAllEvents()
for _, ev := range events {
if ev.Timestamp.After(cutoff24h) {
stats.Last24hEvents++
// Skip events with no attack type — a malformed or legacy JSONL
// line would otherwise surface as a "" key in the JSON response.
if ev.AttackType != "" {
stats.ByType24h[ev.AttackType]++
}
hoursAgo := int(now.Sub(ev.Timestamp).Hours())
if hoursAgo >= 0 && hoursAgo < 24 {
stats.HourlyBuckets[23-hoursAgo]++
}
}
if ev.Timestamp.After(cutoff7d) {
stats.Last7dEvents++
daysAgo := int(now.Sub(ev.Timestamp).Hours() / 24)
if daysAgo >= 0 && daysAgo < 7 {
stats.DailyBuckets[6-daysAgo]++
}
}
}
stats.TopAttackers = db.TopAttackers(10)
return stats
}
// readAllEvents reads all events for stats computation.
// Uses bbolt store when available, falls back to JSONL file.
func (db *DB) readAllEvents() []Event {
if sdb := store.Global(); sdb != nil {
storeEvents := sdb.ReadAllAttackEvents()
events := make([]Event, 0, len(storeEvents))
for _, se := range storeEvents {
events = append(events, Event{
Timestamp: se.Timestamp,
IP: se.IP,
AttackType: AttackType(se.AttackType),
CheckName: se.CheckName,
Severity: se.Severity,
Account: se.Account,
Message: se.Message,
})
}
return events
}
// An unset directory must not let statistics ingest an unrelated log
// from the process's working directory.
if db.dbPath == "" {
return nil
}
// Fallback: flat-file events.jsonl
path := filepath.Join(db.dbPath, eventsFile)
// #nosec G304 -- filepath.Join under operator-configured db.dbPath.
f, err := os.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
var events []Event
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
for scanner.Scan() {
var ev Event
if err := json.Unmarshal(scanner.Bytes(), &ev); err != nil {
continue
}
events = append(events, ev)
}
return events
}
package auditd
import (
"errors"
"fmt"
"os"
"os/exec"
)
type commandRunner func(name string, args ...string) error
func runCommand(name string, args ...string) error {
// #nosec G204 -- production callers supply the fixed augenrules command or
// the absolute executable returned by exec.LookPath, never user input.
return exec.Command(name, args...).Run()
}
const rulesPath = "/etc/audit/rules.d/csm.rules"
const rules = `## Continuous Security Monitor - auditd rules
# Password/auth file changes
-w /etc/shadow -p wa -k csm_shadow_change
-w /etc/passwd -p wa -k csm_passwd_change
-w /etc/group -p wa -k csm_group_change
# SSH config and keys
-w /etc/ssh/sshd_config -p wa -k csm_sshd_change
-w /root/.ssh/authorized_keys -p wa -k csm_root_ssh_keys
# WHM API tokens
-w /var/cpanel/authn/api_tokens_v2/ -p wa -k csm_whm_api_tokens
# Crontab modifications
-w /var/spool/cron/ -p wa -k csm_crontab_change
-w /etc/cron.d/ -p wa -k csm_crond_change
# Password change commands
-w /usr/bin/passwd -p x -k csm_passwd_exec
-w /usr/sbin/chpasswd -p x -k csm_chpasswd_exec
# CSM binary self-protection
-w /opt/csm/csm -p wa -k csm_binary_tamper
-w /etc/csm/csm.yaml -p wa -k csm_config_tamper
-w /opt/csm/csm.yaml -p wa -k csm_config_tamper
# Execution from suspicious locations
-a always,exit -F arch=b64 -S execve -F dir=/tmp -k csm_exec_tmp
-a always,exit -F arch=b64 -S execve -F dir=/dev/shm -k csm_exec_shm
# User account modifications
-w /usr/sbin/useradd -p x -k csm_useradd
-w /usr/sbin/usermod -p x -k csm_usermod
-w /usr/sbin/userdel -p x -k csm_userdel
# AF_ALG socket creation — CVE-2026-31431 "Copy Fail" exploit signature.
# AF_ALG (numeric family 38) is essentially never used by cPanel/PHP
# workloads, so any non-system UID hitting socket(AF_ALG, ...) is suspicious.
# Filter on uid, not auid: service-launched PHP/cPanel workers commonly have
# unset audit login UID while still running as the account user.
# Two rules — b64 covers native 64-bit binaries, b32 closes the i386
# emulation evasion path on x86_64 hosts with 32-bit compat enabled.
-a always,exit -F arch=b64 -S socket -F a0=38 -F uid>=1000 -k csm_af_alg_socket
-a always,exit -F arch=b32 -S socket -F a0=38 -F uid>=1000 -k csm_af_alg_socket
`
func Deploy() error {
return deployRules(rulesPath, runCommand)
}
func deployRules(path string, run commandRunner) error {
// #nosec G306 -- /etc/audit/rules.d/csm.rules is read by the auditd
// tooling (augenrules) on reload; 0640 keeps world-read off.
if err := os.WriteFile(path, []byte(rules), 0640); err != nil {
return err
}
return run("augenrules", "--load")
}
// EnsureDeployed compares the on-disk rules file to the embedded rules
// constant and re-runs Deploy if they differ. Used by the daemon at
// startup so a CSM upgrade that ships new auditd rules does not silently
// remain inactive when the package postinstall did not invoke Deploy.
//
// Returns (redeployed, err): redeployed=true when the file was updated,
// false when it already matched. err is non-nil only when an unexpected
// I/O failure occurred; a missing rules file is treated as "drift" and
// triggers Deploy.
func EnsureDeployed() (bool, error) {
current, err := os.ReadFile(rulesPath)
if err == nil && string(current) == rules {
return false, nil
}
if err != nil && !os.IsNotExist(err) {
return false, err
}
if err := Deploy(); err != nil {
return false, err
}
return true, nil
}
func Remove() error {
return removeRules(rulesPath, exec.LookPath, runCommand)
}
func removeRules(path string, lookPath func(string) (string, error), run commandRunner) error {
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("removing audit rules: %w", err)
}
command, err := lookPath("augenrules")
if errors.Is(err, exec.ErrNotFound) {
return nil
}
if err != nil {
return fmt.Errorf("locating augenrules: %w", err)
}
if err := run(command, "--load"); err != nil {
return fmt.Errorf("reloading audit rules: %w", err)
}
return nil
}
// Package blockdigest batches CSM auto-block events into a per-country
// roll-up so operators learn when IPs from their customers' countries get
// blocked. It is deliberately free of config/geoip/alert/daemon imports:
// the daemon injects a country lookup and email/webhook sinks, which keeps
// the logic unit-testable without the linux build tag.
package blockdigest
import (
"strings"
"sync"
"time"
)
type Bucket string
const (
BucketCustomer Bucket = "customer"
BucketAttacker Bucket = "attacker"
)
// Record is one deduplicated auto-block observation in the current window.
// Domains and URIs are populated only for modsec-category records, from the
// daemon-injected EnrichModSec lookup, so the digest can name which customer
// sites and request paths the blocked IP was hitting.
type Record struct {
TS time.Time
IP string
Country string
Reason string
Bucket Bucket
Category string
Domains []string
URIs []string
}
// Digest is the rolled-up view drained at each interval.
type Digest struct {
Window time.Duration
Countries []string
Total int
CustomerCount int
AttackerCount int
ByCountry map[string]int
ByReason map[string]int
ByCategory map[string]int
Records []Record
}
// Options is the fully-resolved collector configuration. The daemon resolves
// countries, interval, channel sinks, and lookups before constructing.
type Options struct {
Countries []string // effective watch set, upper-cased; empty means all
SendOn string // any | customer
Interval time.Duration
Live bool
MinBlock int
Host string
Version string
CountriesOf func() []string
CountryOf func(ip string) string
// EnrichModSec, when set, returns the customer domains and top request
// URIs a modsec-category blocked IP was hitting. The daemon injects a
// lookup over already-parsed findings so the package stays store-free.
EnrichModSec func(ip string) (domains, uris []string)
Now func() time.Time
EmailSink func(subject, body string) error
WebhookSink func(p WebhookPayload) error
OnError func(channel string, err error)
// DeliveryEnabled only reads current policy; unfinished delivery also
// consults it during cleanup. Nil means every configured sink is eligible.
DeliveryEnabled func(channel string) bool
}
// Collector accumulates observations and drains them into digests.
type Collector struct {
opts Options
queues *digestQueues
mu sync.Mutex
records []Record
lastLive map[string]time.Time
lastLivePruned time.Time
}
const maxBuffered = 5000
func New(opts Options) *Collector {
if opts.Now == nil {
opts.Now = time.Now
}
if opts.CountryOf == nil {
opts.CountryOf = func(string) string { return "" }
}
opts.Countries = append([]string(nil), opts.Countries...)
return &Collector{opts: opts, lastLive: make(map[string]time.Time), queues: newDigestQueues(opts)}
}
func (c *Collector) countriesSnapshot() []string {
if c.opts.CountriesOf != nil {
return normalizeCountries(c.opts.CountriesOf())
}
return append([]string(nil), c.opts.Countries...)
}
// ResolveCountries returns the effective upper-cased watch set: configured
// wins, else trusted_countries, else empty (meaning all countries).
func ResolveCountries(configured, trusted []string) []string {
if out := normalizeCountries(configured); len(out) > 0 {
return out
}
return normalizeCountries(trusted)
}
func normalizeCountries(src []string) []string {
out := make([]string, 0, len(src))
for _, c := range src {
c = strings.ToUpper(strings.TrimSpace(c))
if c != "" {
out = append(out, c)
}
}
return out
}
// classifyBucket maps an auto-block reason to a bucket. Attacker keywords are
// checked first; anything else (including unrecognized reasons) is treated as
// likely-customer so a possible false positive is never hidden from review.
func classifyBucket(reason string) Bucket {
r := strings.ToLower(reason)
attacker := []string{
"rule escalation", "modsecurity", "brute", "mail auth", "web_attack",
"web attack", "account compromise", "command-and-control",
"user-agent spoof", "ua spoof", "bad asn", "credential stuffing",
"credential-stuffing", "credential abuse", "credential-abuse",
"credentials compromised", "waf blocking high-volume attacker",
}
for _, k := range attacker {
if strings.Contains(r, k) {
return BucketAttacker
}
}
if containsWord(r, "c2") {
return BucketAttacker
}
return BucketCustomer
}
// categoryOf groups an auto-block reason into a stable signal class so the
// digest can break blocks down by what tripped them and call out WAF pressure.
// CSM custom rules (900xxx) and OWASP/Comodo CRS escalations collapse to one
// "modsec" class. Reason strings are CSM-internal and stable.
func categoryOf(reason string) string {
r := strings.ToLower(reason)
switch {
case strings.Contains(r, "modsecurity escalation"),
strings.Contains(r, "rule escalation"),
strings.Contains(r, "waf blocking high-volume attacker"):
return "modsec"
case strings.Contains(r, "xml-rpc"):
return "xmlrpc"
case strings.Contains(r, "wordpress login brute"):
return "wp-bruteforce"
case strings.Contains(r, "admin panel brute"):
return "admin-bruteforce"
case strings.Contains(r, "mail account compromise"),
strings.Contains(r, "compromised email account"),
strings.Contains(r, "credentials compromised"),
strings.Contains(r, "outgoing mail hold"):
return "mail-compromise"
case strings.Contains(r, "mail auth"), strings.Contains(r, "smtp brute"):
return "mail-bruteforce"
case strings.Contains(r, "smtp probe"):
return "smtp-probe"
case strings.Contains(r, "ftp brute"):
return "ftp-bruteforce"
case strings.Contains(r, "url scanner profile"):
return "http-scanner"
case strings.Contains(r, "http request flood"):
return "http-flood"
case strings.Contains(r, "user-agent spoof"),
strings.Contains(r, "ua spoof"),
strings.Contains(r, "unverified claimed bot"):
return "ua-spoof"
case strings.Contains(r, "high local threat score"):
return "local-threat"
case strings.Contains(r, "known malicious ip"),
strings.Contains(r, "threat intelligence"),
strings.Contains(r, "command-and-control"),
strings.Contains(r, "bad asn"),
containsWord(r, "c2"):
return "threat-intel"
default:
return "other"
}
}
func containsWord(s, word string) bool {
for start := 0; start < len(s); {
idx := strings.Index(s[start:], word)
if idx < 0 {
return false
}
idx += start
after := idx + len(word)
if (idx == 0 || !isWordByte(s[idx-1])) && (after == len(s) || !isWordByte(s[after])) {
return true
}
start = after
}
return false
}
func isWordByte(b byte) bool {
return b >= 'a' && b <= 'z' || b >= '0' && b <= '9' || b == '_'
}
func (c *Collector) watched(country string) bool {
countries := c.countriesSnapshot()
if len(countries) == 0 {
return true
}
if country == "" {
return false
}
for _, w := range countries {
if strings.EqualFold(country, w) {
return true
}
}
return false
}
// Observe records one auto-block. It geo-classifies, filters to the watch set,
// buckets the reason, and (when Live) may fire an immediate alert.
func (c *Collector) Observe(ip, reason string, ts time.Time) {
country := c.opts.CountryOf(ip)
if !c.watched(country) {
return
}
bucket := classifyBucket(reason)
rec := Record{TS: ts, IP: ip, Country: country, Reason: reason, Bucket: bucket, Category: categoryOf(reason)}
if rec.Category == "modsec" && c.opts.EnrichModSec != nil {
rec.Domains, rec.URIs = c.opts.EnrichModSec(ip)
}
c.mu.Lock()
if len(c.records) < maxBuffered {
c.records = append(c.records, rec)
} else {
// drop-oldest: the digest only needs the interval's worth.
c.records = append(c.records[1:], rec)
}
c.queues.admit()
c.mu.Unlock()
if c.opts.Live {
c.maybeLive(rec)
}
}
// Drain pulls the current window into a Digest (deduped by IP, customer first)
// and clears the buffer.
func (c *Collector) Drain() Digest {
d, batch := c.drain()
batch.finish(true)
return d
}
func (c *Collector) drain() (Digest, *digestBatch) {
c.mu.Lock()
recs := c.records
c.records = nil
batch := c.queues.detach()
c.mu.Unlock()
handedOff := false
defer func() {
if !handedOff {
batch.finish(false)
}
}()
d := Digest{
Window: c.opts.Interval,
Countries: c.countriesSnapshot(),
ByCountry: map[string]int{},
ByReason: map[string]int{},
ByCategory: map[string]int{},
}
selected := make(map[string]Record, len(recs))
order := make([]string, 0, len(recs))
for _, r := range recs {
existing, ok := selected[r.IP]
if !ok {
selected[r.IP] = r
order = append(order, r.IP)
continue
}
if existing.Bucket != BucketCustomer && r.Bucket == BucketCustomer {
selected[r.IP] = r
}
}
var customer, attacker []Record
for _, ip := range order {
r := selected[ip]
d.Total++
d.ByCountry[r.Country]++
d.ByReason[reasonKey(r.Reason)]++
d.ByCategory[r.Category]++
if r.Bucket == BucketCustomer {
d.CustomerCount++
customer = append(customer, r)
} else {
d.AttackerCount++
attacker = append(attacker, r)
}
}
d.Records = make([]Record, 0, len(customer)+len(attacker))
d.Records = append(d.Records, customer...)
d.Records = append(d.Records, attacker...)
batch.coalesce(d.Total)
handedOff = true
return d, batch
}
// reasonKey collapses a reason to its leading phrase (before the first ':')
// so per-reason counts group "ModSecurity escalation: 5+ ..." together.
func reasonKey(reason string) string {
if i := strings.IndexByte(reason, ':'); i > 0 {
return strings.TrimSpace(reason[:i])
}
return strings.TrimSpace(reason)
}
// maybeLive fires an immediate single-record alert for a qualifying block.
// It honors send_on (customer mode only alerts on customer-risk blocks) and
// dedups per IP within one Interval so a re-blocked IP cannot spam.
func (c *Collector) maybeLive(rec Record) {
if c.opts.SendOn == "customer" && rec.Bucket != BucketCustomer {
return
}
now := c.opts.Now()
c.mu.Lock()
c.pruneLastLiveLocked(now)
if last, ok := c.lastLive[rec.IP]; ok && c.opts.Interval > 0 && now.Sub(last) < c.opts.Interval {
c.mu.Unlock()
return
}
c.lastLive[rec.IP] = now
c.mu.Unlock()
// The delivery policy is injected by the caller, so it is consulted
// outside this lock. The dedup decision above already stands.
plan := c.beginDelivery()
defer plan.finish()
d := Digest{
Window: c.opts.Interval, Countries: c.countriesSnapshot(),
Total: 1, ByCountry: map[string]int{rec.Country: 1},
ByReason: map[string]int{reasonKey(rec.Reason): 1},
ByCategory: map[string]int{rec.Category: 1},
Records: []Record{rec},
}
if rec.Bucket == BucketCustomer {
d.CustomerCount = 1
} else {
d.AttackerCount = 1
}
c.deliver(plan, "block_live", d)
}
func (c *Collector) pruneLastLiveLocked(now time.Time) {
interval := c.opts.Interval
if interval <= 0 {
clear(c.lastLive)
c.lastLivePruned = now
return
}
if !c.lastLivePruned.IsZero() && now.Sub(c.lastLivePruned) < interval {
return
}
cutoff := now.Add(-interval)
for ip, last := range c.lastLive {
if !last.After(cutoff) {
delete(c.lastLive, ip)
}
}
c.lastLivePruned = now
}
// tick drains the window and sends a digest when gating allows.
func (c *Collector) tick() {
d, batch := c.drain()
defer batch.finish(false)
if !c.shouldSend(d) {
batch.finish(true)
return
}
c.dispatch("block_digest", d, batch)
}
// Flush sends one final digest regardless of cadence (shutdown path).
func (c *Collector) Flush() { c.tick() }
// Run loops draining on each tick and drains a final digest on stop.
func (c *Collector) Run(stop <-chan struct{}, tick <-chan time.Time) {
c.queues.setStopped(false)
defer c.queues.setStopped(true)
for {
select {
case <-stop:
c.Flush()
return
case _, ok := <-tick:
if !ok {
c.Flush()
return
}
c.tick()
}
}
}
// dispatch delivers a digest through whichever sinks are configured. Alert
// delivery is best-effort and must never block the collector or the auto-block
// path beyond the sink's own timeout; failures are reported through OnError.
func (c *Collector) dispatch(event string, d Digest, batch *digestBatch) {
plan := c.beginDelivery()
defer plan.finish()
// Eligible destinations own notifications before the buffer releases
// the coalesced records. A failing first sink cannot hide the second.
batch.finish(true)
c.deliver(plan, event, d)
}
func (c *Collector) deliver(plan *digestDeliveryPlan, event string, d Digest) {
if job := plan.email; job != nil {
if job.start() {
subject, body := c.renderSubject(d), c.renderBody(d)
job.offered = true
err := c.opts.EmailSink(subject, body)
job.result(err)
if err != nil {
c.reportSinkError("email", err)
}
}
job.finish()
}
if job := plan.webhook; job != nil {
if job.start() {
payload := c.buildPayload(event, d)
job.offered = true
err := c.opts.WebhookSink(payload)
job.result(err)
if err != nil {
c.reportSinkError("webhook", err)
}
}
job.finish()
}
}
func (c *Collector) reportSinkError(channel string, err error) {
if c.opts.OnError != nil {
c.opts.OnError(channel, err)
}
}
package blockdigest
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type digestQueues struct {
mu sync.Mutex
arrivals []time.Time
batches map[*digestBatch]struct{}
losses *queuehealth.Tracker
maxLag time.Duration
stopped bool
sinks map[string]*digestSinkQueue
}
type digestBatch struct {
queue *digestQueues
count int
progress time.Time
finished bool
}
type digestSinkQueue struct {
mu sync.Mutex
tracker *queuehealth.Tracker
uncertain bool
uncertainAt time.Time
}
type digestDelivery struct {
queue *digestSinkQueue
ticket queuehealth.Ticket
offered, returned, finished bool
tracked bool
channel string
deliveryEnabled func(string) bool
}
type digestDeliveryPlan struct{ email, webhook *digestDelivery }
func newDigestQueues(opts Options) *digestQueues {
q := &digestQueues{batches: make(map[*digestBatch]struct{}), losses: queuehealth.New(maxBuffered, time.Minute), maxLag: max(time.Minute, opts.Interval+time.Minute), sinks: make(map[string]*digestSinkQueue)}
if opts.EmailSink != nil {
q.sinks["email"] = &digestSinkQueue{tracker: queuehealth.New(0, time.Minute)}
}
if opts.WebhookSink != nil {
q.sinks["webhook"] = &digestSinkQueue{tracker: queuehealth.New(0, time.Minute)}
}
return q
}
// Admission and detach run under the collector lock alongside the actual
// buffer. Health never takes that lock or calls injected lookups or sinks.
func (q *digestQueues) admit() {
q.mu.Lock()
defer q.mu.Unlock()
if len(q.arrivals) == maxBuffered {
q.arrivals = q.arrivals[1:]
q.losses.Lose(time.Now(), 1)
}
q.arrivals = append(q.arrivals, time.Now())
}
func (q *digestQueues) detach() *digestBatch {
q.mu.Lock()
defer q.mu.Unlock()
b := &digestBatch{queue: q, count: len(q.arrivals), progress: time.Now()}
q.arrivals = nil
q.batches[b] = struct{}{}
return b
}
func (b *digestBatch) coalesce(count int) {
b.queue.mu.Lock()
b.count = count
b.progress = time.Now()
b.queue.mu.Unlock()
}
func (b *digestBatch) finish(completed bool) {
if b == nil {
return
}
q := b.queue
q.mu.Lock()
defer q.mu.Unlock()
if b.finished {
return
}
if !completed {
q.losses.Lose(time.Now(), uint64(b.count)) // #nosec G115 -- batch count is a slice length bounded by maxBuffered; coalescing can only reduce it.
}
b.finished = true
delete(q.batches, b)
}
func (q *digestQueues) setStopped(stopped bool) {
q.mu.Lock()
q.stopped = stopped
q.mu.Unlock()
}
func (q *digestSinkQueue) begin(channel string, enabled func(string) bool) *digestDelivery {
if q == nil {
return nil
}
d := &digestDelivery{queue: q, channel: channel, deliveryEnabled: enabled}
if d.enabled() {
d.enqueue()
}
return d
}
func (c *Collector) beginDelivery() *digestDeliveryPlan {
return &digestDeliveryPlan{
email: c.queues.sinks["email"].begin("email", c.opts.DeliveryEnabled),
webhook: c.queues.sinks["webhook"].begin("webhook", c.opts.DeliveryEnabled),
}
}
func (d *digestDelivery) enabled() bool {
return d.deliveryEnabled == nil || d.deliveryEnabled(d.channel)
}
func (d *digestDelivery) enqueue() {
d.ticket = d.queue.tracker.Begin(time.Now())
d.tracked = true
}
func (d *digestDelivery) start() bool {
if !d.enabled() {
d.returned = true
return false
}
// A reload may enable the second destination while the first is running.
if !d.tracked {
d.enqueue()
}
d.ticket.Start(time.Now())
return true
}
func (d *digestDelivery) result(err error) {
d.returned = true
if err != nil {
d.queue.tracker.Lose(time.Now(), 1)
}
}
func (d *digestDelivery) finish() {
if d == nil || d.finished {
return
}
if !d.returned {
if d.offered {
d.queue.mu.Lock()
d.queue.uncertain = true
d.queue.uncertainAt = time.Now()
d.queue.mu.Unlock()
} else if d.enabled() {
d.queue.tracker.Lose(time.Now(), 1)
}
}
if d.tracked {
d.ticket.Finish(time.Now())
}
d.finished = true
}
func (p *digestDeliveryPlan) finish() { p.email.finish(); p.webhook.finish() }
func (q *digestQueues) recordStatus(now time.Time) queuehealth.Status {
q.mu.Lock()
defer q.mu.Unlock()
row := q.losses.Snapshot(now)
row.DepthUnit = "records"
row.LagBasis = "operation_progress"
row.Depth = len(q.arrivals)
if row.Depth > 0 {
row.LagSeconds = max(0, now.Sub(q.arrivals[0]).Seconds())
}
for b := range q.batches {
row.InFlight += b.count
row.ProcessingSeconds = max(row.ProcessingSeconds, now.Sub(b.progress).Seconds())
}
switch {
case q.stopped && row.Depth > 0:
row.Status, row.Reason = "degraded", "consumer_stopped"
case row.LagSeconds >= q.maxLag.Seconds():
row.Status, row.Reason = "degraded", "backlog_lag"
case row.ProcessingSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "processing_lag"
}
return row
}
func (q *digestSinkQueue) snapshot(now time.Time) queuehealth.Status {
q.mu.Lock()
defer q.mu.Unlock()
row := q.tracker.Snapshot(now)
row.CapacityUnavailable = true
row.DepthUnit = "notifications"
row.DroppedLowerBound = q.uncertain
if !q.uncertainAt.IsZero() && now.Sub(q.uncertainAt) < time.Minute {
row.Status, row.Reason = "degraded", "delivery_uncertain"
}
return row
}
func (c *Collector) QueueStatuses(now time.Time) map[string]queuehealth.Status {
rows := map[string]queuehealth.Status{"records": c.queues.recordStatus(now)}
for name, q := range c.queues.sinks {
rows[name] = q.snapshot(now)
}
return rows
}
package blockdigest
import (
"fmt"
"sort"
"strings"
"time"
)
// WebhookPayload renders for Slack/Mattermost (Text) and programmatic
// receivers (CSM) at once.
type WebhookPayload struct {
Text string `json:"text"`
CSM WebhookCSM `json:"csm"`
}
type WebhookCSM struct {
Event string `json:"event"`
Host string `json:"host"`
Version string `json:"version"`
Window string `json:"window"`
Countries []string `json:"countries"`
Counts WebhookCounts `json:"counts"`
Blocks []WebhookBlock `json:"blocks"`
}
type WebhookCounts struct {
Total int `json:"total"`
Customer int `json:"customer"`
Attacker int `json:"attacker"`
ByCountry map[string]int `json:"by_country"`
ByReason map[string]int `json:"by_reason"`
ByCategory map[string]int `json:"by_category"`
}
type WebhookBlock struct {
IP string `json:"ip"`
Country string `json:"country"`
Reason string `json:"reason"`
Bucket string `json:"bucket"`
Category string `json:"category"`
Domains []string `json:"domains,omitempty"`
URIs []string `json:"uris,omitempty"`
TS string `json:"ts"`
}
const maxAttackerListed = 10
func countriesLabel(countries []string) string {
if len(countries) == 0 {
return "all"
}
return strings.Join(countries, ",")
}
func (c *Collector) renderSubject(d Digest) string {
return fmt.Sprintf("[%s] %d watched-country IPs blocked (%d customer-risk) last %s",
c.opts.Host, d.Total, d.CustomerCount, d.Window)
}
func (c *Collector) renderBody(d Digest) string {
var b strings.Builder
fmt.Fprintln(&b, c.renderSubject(d))
fmt.Fprintf(&b, "Countries: %s\n", countriesLabel(d.Countries))
fmt.Fprintf(&b, "By country: %s\n", sortedCounts(d.ByCountry))
fmt.Fprintf(&b, "By category: %s\n", sortedCounts(d.ByCategory))
fmt.Fprintf(&b, "By reason: %s\n", sortedCounts(d.ByReason))
fmt.Fprintln(&b)
fmt.Fprintln(&b, "LIKELY CUSTOMER (false-positive risk -- review/unblock with: csm firewall remove <ip>):")
wrote := false
for _, r := range d.Records {
if r.Bucket != BucketCustomer {
continue
}
fmt.Fprintf(&b, " %s | %s | %s\n", r.IP, r.Country, r.Reason)
wrote = true
}
if !wrote {
fmt.Fprintln(&b, " (none)")
}
fmt.Fprintln(&b)
fmt.Fprintf(&b, "Attacker blocks (correctly blocked): %d\n", d.AttackerCount)
listed := 0
for _, r := range d.Records {
if r.Bucket != BucketAttacker {
continue
}
if listed >= maxAttackerListed {
fmt.Fprintf(&b, " ... and %d more\n", d.AttackerCount-listed)
break
}
fmt.Fprintf(&b, " %s | %s | %s\n", r.IP, r.Country, r.Reason)
listed++
}
fmt.Fprintln(&b)
modsec := d.ByCategory["modsec"]
fmt.Fprintf(&b, "ModSecurity blocks (WAF escalations): %d\n", modsec)
listed = 0
for _, r := range d.Records {
if r.Category != "modsec" {
continue
}
if listed >= maxAttackerListed {
fmt.Fprintf(&b, " ... and %d more\n", modsec-listed)
break
}
fmt.Fprintf(&b, " %s | %s | %s\n", r.IP, r.Country, r.Reason)
if len(r.Domains) > 0 {
fmt.Fprintf(&b, " targets: %s\n", strings.Join(r.Domains, ", "))
} else {
fmt.Fprintln(&b, " targets: no customer domain recorded")
}
if len(r.URIs) > 0 {
fmt.Fprintf(&b, " top URIs: %s\n", strings.Join(r.URIs, " | "))
}
listed++
}
fmt.Fprintln(&b)
fmt.Fprintln(&b, "Deep per-IP report (successful-hit counts): run the on-host CSM block report.")
return b.String()
}
func sortedCounts(m map[string]int) string {
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
parts := make([]string, 0, len(keys))
for _, k := range keys {
parts = append(parts, fmt.Sprintf("%s=%d", k, m[k]))
}
if len(parts) == 0 {
return "(none)"
}
return strings.Join(parts, " ")
}
func (c *Collector) buildPayload(event string, d Digest) WebhookPayload {
blocks := make([]WebhookBlock, 0, len(d.Records))
for _, r := range d.Records {
blocks = append(blocks, WebhookBlock{
IP: r.IP, Country: r.Country, Reason: r.Reason,
Bucket: string(r.Bucket), Category: r.Category,
Domains: r.Domains, URIs: r.URIs,
TS: r.TS.UTC().Format(time.RFC3339),
})
}
return WebhookPayload{
Text: c.renderBody(d),
CSM: WebhookCSM{
Event: event, Host: c.opts.Host, Version: c.opts.Version,
Window: d.Window.String(), Countries: d.Countries,
Counts: WebhookCounts{
Total: d.Total, Customer: d.CustomerCount, Attacker: d.AttackerCount,
ByCountry: d.ByCountry, ByReason: d.ByReason, ByCategory: d.ByCategory,
},
Blocks: blocks,
},
}
}
// shouldSend gates a digest: send_on=any sends when total meets MinBlock;
// send_on=customer sends only when customer-risk blocks meet MinBlock.
func (c *Collector) shouldSend(d Digest) bool {
if c.opts.SendOn == "customer" {
return d.CustomerCount >= c.opts.MinBlock
}
return d.Total >= c.opts.MinBlock
}
// Package bpf provides the shared scaffolding that BPF-backed live monitors
// across the daemon use: a common Backend interface, backend-kind constants
// for operator config, sentinel errors that distinguish "not built" from
// "kernel unsupported", and a per-feature backend metric.
//
// Real BPF code (program loading, ringbuf consumption, capability probing)
// lives behind the linux && bpf build tag in sibling files. Default builds
// compile stubs that report all capabilities as unavailable.
package bpf
import (
"context"
"errors"
"sync"
"github.com/pidginhost/csm/internal/metrics"
)
// Backend is the shape every BPF-backed live monitor implements. The legacy
// fallback for each feature implements the same interface so the coordinator
// hands the daemon a uniform handle.
type Backend interface {
Mode() string
EventCount() uint64
Run(ctx context.Context)
}
// Backend kind constants are the shared internal values used after each
// feature validates its operator-facing config setting. Individual features
// may keep older public values, such as AF_ALG's "auditd", and map them to
// BackendLegacy internally.
const (
BackendAuto = "auto"
BackendBPF = "bpf"
BackendLegacy = "legacy"
BackendNone = "none"
)
// ErrNotBuilt is returned by feature loaders when CSM was built without the
// bpf build tag. The coordinator treats this identically to a kernel that
// lacks the required BPF program type: log it and fall back to legacy.
var ErrNotBuilt = errors.New("BPF support not compiled in (rebuild with -tags bpf)")
// ErrUnsupported is returned when CSM was built with the bpf tag but the
// running kernel does not accept the requested BPF program type. Distinct
// from ErrNotBuilt so operator logs explain whether the fix is "rebuild" or
// "newer kernel".
var ErrUnsupported = errors.New("kernel does not support requested BPF program type")
var (
backendMetric *metrics.GaugeVec
backendMetricOnce sync.Once
activeMu sync.RWMutex
activeKinds map[string]string
)
// MetricFor returns the shared csm_bpf_backend gauge. The feature argument
// is accepted for call-site readability; all features share a single
// GaugeVec distinguished by the label value, not by separate vec instances.
// Registered exactly once across the process. Pair with SetActive instead
// of calling With directly.
func MetricFor(_ string) *metrics.GaugeVec {
backendMetricOnce.Do(func() {
backendMetric = metrics.NewGaugeVec(
"csm_bpf_backend",
"Active backend for each BPF-backed live monitor; 1 for the selected kind, 0 otherwise.",
[]string{"feature", "kind"},
)
metrics.MustRegister("csm_bpf_backend", backendMetric)
})
return backendMetric
}
// SetActive sets the metric series so that exactly one of {bpf, legacy, none}
// is at 1 and the others at 0 for the given feature. Call from the coordinator
// after backend selection. Also remembers the active kind in-process so
// internal/health can render the matching capability string without
// reaching into the metric registry.
func SetActive(feature, active string) {
g := MetricFor(feature)
for _, k := range []string{BackendBPF, BackendLegacy, BackendNone} {
v := 0.0
if k == active {
v = 1.0
}
g.With(feature, k).Set(v)
}
activeMu.Lock()
if activeKinds == nil {
activeKinds = make(map[string]string)
}
activeKinds[feature] = active
activeMu.Unlock()
}
// ActiveKind returns the kind currently selected for the given feature, or
// "" if SetActive was never called for it. internal/health uses this to
// render per-feature capability strings without importing the metrics
// registry.
func ActiveKind(feature string) string {
activeMu.RLock()
defer activeMu.RUnlock()
return activeKinds[feature]
}
package bpf
import "sync"
// Capabilities reports which BPF program types the running kernel can
// actually load and attach. Populated once at daemon startup by Probe.
//
// Each field maps to a kernel feature, not a CSM feature: a single CSM
// feature (e.g. AF_ALG kernel-side blocking) may need multiple capability bits
// (LSMAttach + Ringbuf).
type Capabilities struct {
LSMAttach bool // BPF LSM programs can attach (kernel >= 5.7 with BPF LSM trampoline)
CgroupSock bool // BPF_PROG_TYPE_CGROUP_SOCK_ADDR can attach to cgroup/connect4 (>= 4.10)
Tracepoint bool // BPF_PROG_TYPE_TRACEPOINT can attach to sched/sched_process_exec (>= 4.7)
Ringbuf bool // BPF_MAP_TYPE_RINGBUF available (>= 5.8)
}
// Any reports whether at least one capability is true. Used by callers that
// only need to know "any BPF surface is usable" before deciding on legacy
// vs auto.
func (c Capabilities) Any() bool {
return c.LSMAttach || c.CgroupSock || c.Tracepoint || c.Ringbuf
}
var (
probeOnce sync.Once
probeResult Capabilities
)
// Probe returns the cached BPF capability result for this process. The first
// call performs the privileged load/attach probes; later calls return the same
// value without touching the kernel again.
func Probe() Capabilities {
probeOnce.Do(func() {
probeResult = probeKernel()
})
return probeResult
}
//go:build !(linux && bpf)
package bpf
// probeKernel is the no-tag stub. It returns zero-value Capabilities so that
// all coordinators on this build see "no BPF surface available" and pick the
// legacy backend.
func probeKernel() Capabilities { return Capabilities{} }
package bpf
// dropEventLogStride is how often the drop path logs after the first drop.
// 256 keeps sustained back-pressure visible without flooding the daemon log.
const dropEventLogStride uint64 = 256
func shouldLogDroppedEvent(dropped uint64) bool {
return dropped == 1 || (dropped > 0 && dropped%dropEventLogStride == 0)
}
// Package broadcast provides a one-to-many publish bus for alert.Finding
// events. Subscribers each get a buffered channel; if a subscriber's
// buffer fills, that subscriber drops the message rather than blocking
// the publisher. Used by the SSE event stream and any other in-process
// passive consumer.
//
// This is intentionally separate from the daemon's primary alert pipeline
// (the unbuffered or large-buffered alertCh that feeds Dispatch). The bus
// is a side-channel for observers that should not influence dispatch.
package broadcast
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
)
// defaultMaxSubscribers caps concurrent subscribers so a flood of event-stream
// connections (each a goroutine plus a buffered channel) cannot exhaust the
// daemon's memory. Generous for an operator dashboard.
const defaultMaxSubscribers = 256
// Bus fans out published findings to every subscriber.
type Bus struct {
mu sync.RWMutex
subscribers map[*Subscription]struct{}
buffer int
maxSubs int
closed bool
stats *queuehealth.Tracker
}
// NewBus constructs a Bus with the given per-subscriber buffer.
// A buffer < 1 falls back to 16.
func NewBus(buffer int) *Bus {
if buffer < 1 {
buffer = 16
}
return &Bus{
subscribers: make(map[*Subscription]struct{}),
buffer: buffer,
maxSubs: defaultMaxSubscribers,
stats: queuehealth.New(0, time.Minute),
}
}
// SetMaxSubscribers overrides the concurrent-subscriber cap. A value < 1 is
// ignored. Safe to call before the bus is in use.
func (b *Bus) SetMaxSubscribers(n int) {
if n < 1 {
return
}
b.mu.Lock()
b.maxSubs = n
b.mu.Unlock()
}
// TrySubscribe is Subscribe with the concurrent-subscriber cap enforced. It
// returns ok=false when the cap is reached so an untrusted caller (the SSE
// endpoint, reachable with a low-trust read token) cannot open unbounded
// long-lived streams. Use this for externally-driven subscriptions; Subscribe
// remains for trusted in-process consumers.
func (b *Bus) TrySubscribe() (*Subscription, bool) {
b.mu.Lock()
defer b.mu.Unlock()
if !b.closed && len(b.subscribers) >= b.maxSubs {
return nil, false
}
return b.subscribe(), true
}
// Subscribe is for trusted in-process consumers. Every received delivery must
// be processed once, and the consumer must Unsubscribe when it stops reading.
func (b *Bus) Subscribe() *Subscription {
b.mu.Lock()
defer b.mu.Unlock()
return b.subscribe()
}
func (b *Bus) subscribe() *Subscription {
sub := &Subscription{events: make(chan Delivery, b.buffer), stats: queuehealth.New(b.buffer, time.Minute)}
if b.closed {
close(sub.events)
} else {
b.subscribers[sub] = struct{}{}
}
return sub
}
// Unsubscribe withdraws demand for unread events, for example when a browser
// tab closes. A delivery already received remains the consumer's responsibility.
func (b *Bus) Unsubscribe(sub *Subscription) {
b.unsubscribe(sub, false)
}
// Abort removes a failed consumer and counts its unread deliveries as lost.
func (b *Bus) Abort(sub *Subscription) {
b.unsubscribe(sub, true)
}
func (b *Bus) unsubscribe(sub *Subscription, failed bool) {
b.mu.Lock()
defer b.mu.Unlock()
if _, ok := b.subscribers[sub]; !ok {
return
}
delete(b.subscribers, sub)
if !b.closed {
close(sub.events)
}
for delivery := range sub.events {
delivery.finish(failed)
}
}
// Publish sends f to every current subscriber. Non-blocking: if a
// subscriber's buffer is full, that delivery is skipped.
func (b *Bus) Publish(f alert.Finding) {
b.mu.RLock()
defer b.mu.RUnlock()
if b.closed {
return
}
for sub := range b.subscribers {
now := time.Now()
delivery := Delivery{finding: f, total: b.stats.Begin(now), local: sub.stats.Begin(now)}
select {
case sub.events <- delivery:
default:
delivery.finish(true)
}
}
}
// Close shuts the bus down and closes every outstanding subscriber channel.
// Idempotent.
func (b *Bus) Close() {
b.mu.Lock()
defer b.mu.Unlock()
if b.closed {
return
}
b.closed = true
// Consumers may still drain buffered deliveries after close. Retain their
// capacity and ownership until they unsubscribe.
for sub := range b.subscribers {
close(sub.events)
}
}
package broadcast
import (
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
)
// Subscription owns one bounded delivery buffer. Events must have one consumer.
type Subscription struct {
events chan Delivery
stats *queuehealth.Tracker
}
func (s *Subscription) Events() <-chan Delivery { return s.events }
// Delivery retains queue ownership through encoding and writing to the client.
// Call Process exactly once for every received delivery, including failures.
type Delivery struct {
finding alert.Finding
total queuehealth.Ticket
local queuehealth.Ticket
}
func (d Delivery) Process(fn func(alert.Finding) error) error {
now := time.Now()
d.total.Start(now)
d.local.Start(now)
failed := true
defer func() { d.finish(failed) }()
err := fn(d.finding)
failed = err != nil
return err
}
func (d Delivery) finish(failed bool) {
now := time.Now()
if failed {
d.total.Reject(now)
d.local.Reject(now)
} else {
d.total.Finish(now)
d.local.Finish(now)
}
}
// QueueStatuses aggregates subscriber work without retaining departed clients
// or exposing client identities. The total tracker keeps in-flight work and
// cumulative loss after removal; local trackers prevent an empty peer from
// hiding a full subscriber buffer. The row is advisory: a client that stops
// reading loses its own copy of findings that are already stored.
func (b *Bus) QueueStatuses(now time.Time) map[string]queuehealth.Status {
b.mu.RLock()
defer b.mu.RUnlock()
status := b.stats.Snapshot(now)
status.Advisory = true
for sub := range b.subscribers {
local := sub.stats.Snapshot(now)
status.Capacity += local.Capacity
if local.Reason == "queue_full" && (status.Reason == "" || status.Reason == "dropped_work") {
status.Status, status.Reason = "degraded", "queue_full"
}
}
return map[string]queuehealth.Status{"deliveries": status}
}
package challenge
import (
"context"
"encoding/json"
"errors"
"fmt"
"net/http"
"net/url"
"strings"
"time"
)
// CaptchaProvider verifies a third-party CAPTCHA token. Implementations
// post the operator's secret + the visitor's response token to the
// provider's siteverify endpoint and return a single bool: did the
// provider accept this submission?
type CaptchaProvider interface {
Name() string
Verify(ctx context.Context, token, remoteIP string) (bool, error)
}
// providerEndpoint is exposed as a package var so tests can point a
// provider at httptest.Server rather than the live siteverify URL.
var providerEndpoint = map[string]string{
"turnstile": "https://challenges.cloudflare.com/turnstile/v0/siteverify",
"hcaptcha": "https://hcaptcha.com/siteverify",
}
// captchaProvider implements both Cloudflare Turnstile and hCaptcha;
// they accept identical request/response shapes (POST form, JSON
// {"success":bool} reply) so a single struct covers both.
type captchaProvider struct {
name string
endpoint string
secret string
client *http.Client
}
// NewCaptchaProvider returns the right provider for the configured
// name. Returns nil + nil when the operator has not enabled CAPTCHA;
// the server treats nil as "feature off".
func NewCaptchaProvider(name, secret string, timeout time.Duration) (CaptchaProvider, error) {
name = strings.ToLower(strings.TrimSpace(name))
if name == "" {
return nil, nil
}
endpoint, ok := providerEndpoint[name]
if !ok {
return nil, fmt.Errorf("unknown captcha provider %q (want turnstile or hcaptcha)", name)
}
if secret == "" {
return nil, fmt.Errorf("captcha provider %q requires secret_key", name)
}
if timeout <= 0 {
timeout = 10 * time.Second
}
return &captchaProvider{
name: name,
endpoint: endpoint,
secret: secret,
client: &http.Client{Timeout: timeout},
}, nil
}
func (p *captchaProvider) Name() string { return p.name }
// Verify posts the operator's secret + the visitor's token to the
// provider. The remoteIP is optional but recommended; both Turnstile
// and hCaptcha accept it for binding the verification to a single
// client. Network errors propagate; a 200 with success=false returns
// (false, nil).
func (p *captchaProvider) Verify(ctx context.Context, token, remoteIP string) (bool, error) {
if token == "" {
return false, errors.New("empty captcha token")
}
form := url.Values{}
form.Set("secret", p.secret)
form.Set("response", token)
if remoteIP != "" {
form.Set("remoteip", remoteIP)
}
// #nosec G704 -- p.endpoint is set from the providerEndpoint package-level map (turnstile / hcaptcha) by NewCaptchaProvider, which rejects unknown names. Not attacker-controlled; SSRF is not possible.
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.endpoint, strings.NewReader(form.Encode()))
if err != nil {
return false, fmt.Errorf("building siteverify request: %w", err)
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
// #nosec G704 -- same as above: the request URL is hardcoded to a known siteverify endpoint, never operator or attacker input.
resp, err := p.client.Do(req)
if err != nil {
return false, fmt.Errorf("siteverify call: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return false, fmt.Errorf("siteverify status %d", resp.StatusCode)
}
var body struct {
Success bool `json:"success"`
}
if err := json.NewDecoder(resp.Body).Decode(&body); err != nil {
return false, fmt.Errorf("decoding siteverify response: %w", err)
}
return body.Success, nil
}
package challenge
import (
"context"
"net"
"strings"
"sync"
"time"
)
// crawlerSuffix names a verified-crawler family along with the PTR
// record suffixes that legitimate hosts in that family use.
type crawlerSuffix struct {
name string
domains []string
}
// builtinCrawlers lists the known canonical reverse-DNS suffixes for
// the supported crawler families. Adding a new family means appending
// to this list and documenting the name in csm.yaml's
// challenge.verified_crawlers.providers.
var builtinCrawlers = map[string]crawlerSuffix{
"googlebot": {name: "googlebot", domains: []string{".googlebot.com.", ".google.com."}},
"bingbot": {name: "bingbot", domains: []string{".search.msn.com."}},
}
// Resolver matches the subset of net.Resolver that CrawlerVerifier
// uses, so tests can swap in a fake without spinning up a real DNS
// server.
type Resolver interface {
LookupAddr(ctx context.Context, addr string) (names []string, err error)
LookupHost(ctx context.Context, host string) (addrs []string, err error)
}
// CrawlerVerifier classifies an IP as a verified search crawler iff
// the IP's reverse-DNS PTR matches one of the configured suffixes AND
// the PTR forward-resolves back to the same IP. The verifier caches
// both positive and negative results; positive cache TTL is the
// configured cacheTTL, negative is one-fifth of that to keep a
// transiently-broken resolver from locking out a legitimate crawler.
type CrawlerVerifier struct {
suffixes []crawlerSuffix
resolver Resolver
posTTL time.Duration
negTTL time.Duration
mu sync.Mutex
cache map[string]cacheEntry
maxSize int
}
// crawlerCacheMaxEntries bounds the verifier cache between the daemon's
// 60-second prune ticks. Without a cap a scan from many unique source IPs
// (every miss inserts a negative entry) grows the map to request-rate x 60s,
// an external memory-pressure lever. At the cap an insert first drops expired
// entries, then evicts the soonest-to-expire entry to make room.
const crawlerCacheMaxEntries = 50000
type cacheEntry struct {
verified bool
expires time.Time
}
// NewCrawlerVerifier builds a verifier with the named crawler families
// enabled. Unknown names are ignored (operators may have configured a
// family this binary does not know about; that is harmless).
func NewCrawlerVerifier(providers []string, cacheTTL time.Duration, resolver Resolver) *CrawlerVerifier {
if cacheTTL <= 0 {
cacheTTL = 15 * time.Minute
}
if resolver == nil {
resolver = net.DefaultResolver
}
enabled := make([]crawlerSuffix, 0, len(providers))
for _, name := range providers {
if c, ok := builtinCrawlers[strings.ToLower(strings.TrimSpace(name))]; ok {
enabled = append(enabled, c)
}
}
return &CrawlerVerifier{
suffixes: enabled,
resolver: resolver,
posTTL: cacheTTL,
negTTL: cacheTTL / 5,
cache: make(map[string]cacheEntry),
maxSize: crawlerCacheMaxEntries,
}
}
// Enabled reports whether at least one crawler family is configured;
// the server uses this to skip the verifier entirely (no DNS round
// trip) when the operator has not opted in.
func (v *CrawlerVerifier) Enabled() bool {
return v != nil && len(v.suffixes) > 0
}
// Verified does the reverse-DNS + forward-confirm dance for ip and
// caches the result. Returns true only when the PTR ends in one of the
// allowed suffixes AND a forward lookup of the PTR includes ip in the
// result set.
func (v *CrawlerVerifier) Verified(ctx context.Context, ip string) bool {
if !v.Enabled() {
return false
}
if hit, ok := v.cacheGet(ip); ok {
return hit
}
verified := v.probe(ctx, ip)
v.cachePut(ip, verified)
return verified
}
func (v *CrawlerVerifier) probe(ctx context.Context, ip string) bool {
names, err := v.resolver.LookupAddr(ctx, ip)
if err != nil || len(names) == 0 {
return false
}
for _, name := range names {
// LookupAddr returns FQDNs with a trailing dot. The suffix
// list also has trailing dots so HasSuffix is unambiguous.
lower := strings.ToLower(name)
if !v.suffixMatches(lower) {
continue
}
addrs, err := v.resolver.LookupHost(ctx, strings.TrimSuffix(lower, "."))
if err != nil {
continue
}
for _, addr := range addrs {
if addr == ip {
return true
}
}
}
return false
}
func (v *CrawlerVerifier) suffixMatches(name string) bool {
for _, s := range v.suffixes {
for _, d := range s.domains {
if strings.HasSuffix(name, d) {
return true
}
}
}
return false
}
func (v *CrawlerVerifier) cacheGet(ip string) (bool, bool) {
v.mu.Lock()
defer v.mu.Unlock()
e, ok := v.cache[ip]
if !ok {
return false, false
}
if time.Now().After(e.expires) {
delete(v.cache, ip)
return false, false
}
return e.verified, true
}
func (v *CrawlerVerifier) cachePut(ip string, verified bool) {
ttl := v.negTTL
if verified {
ttl = v.posTTL
}
v.mu.Lock()
defer v.mu.Unlock()
if _, exists := v.cache[ip]; !exists && v.maxSize > 0 && len(v.cache) >= v.maxSize {
v.evictForInsertLocked()
}
v.cache[ip] = cacheEntry{verified: verified, expires: time.Now().Add(ttl)}
}
// evictForInsertLocked frees a slot when the cache is at capacity. It first
// drops every expired entry; if that reclaimed nothing (all entries still
// live), it evicts the single soonest-to-expire entry. Caller holds v.mu.
func (v *CrawlerVerifier) evictForInsertLocked() {
now := time.Now()
freed := false
for ip, e := range v.cache {
if now.After(e.expires) {
delete(v.cache, ip)
freed = true
}
}
if freed {
return
}
var soonestIP string
var soonest time.Time
for ip, e := range v.cache {
if soonestIP == "" || e.expires.Before(soonest) {
soonestIP = ip
soonest = e.expires
}
}
if soonestIP != "" {
delete(v.cache, soonestIP)
}
}
// cleanExpired drops every entry whose TTL has lapsed. Without this,
// a scan from many IPs leaves the cache full of stale entries until
// each individual IP is queried again. Called from Server.CleanExpired
// on the daemon's 60-second ticker. now is passed in so the caller can
// share a single timestamp across multiple cleanup paths.
func (v *CrawlerVerifier) cleanExpired(now time.Time) {
if v == nil {
return
}
v.mu.Lock()
defer v.mu.Unlock()
for ip, e := range v.cache {
if now.After(e.expires) {
delete(v.cache, ip)
}
}
}
package challenge
import (
"bytes"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/atomicio"
)
// DefaultMapPath is the webserver-readable Apache / LSWS RewriteMap.
// It lives in the service's cache directory rather than state_path
// (mode 0700, private to CSM's bbolt database) or the runtime directory:
// Apache and LSWS validate a txt: RewriteMap at config-parse time and Nginx
// fails on a missing include, so the maps are part of the webserver's
// configuration and must exist while CSM is stopped -- package upgrades,
// restores, reboots. systemd deletes the runtime directory on every stop
// and keeps the cache directory.
const DefaultMapPath = "/var/cache/csm/challenge_ips.txt"
// DefaultNginxMapPath is the webserver-readable Nginx map include.
const DefaultNginxMapPath = "/var/cache/csm/challenge_ips.nginx.map"
// challengeEntry stores the challenge metadata for a single IP.
type challengeEntry struct {
FindingID string
ExpiresAt time.Time
Reason string
NonEscalating bool
}
// ExpiredEntry is returned by ExpiredEntries for escalation.
type ExpiredEntry struct {
FindingID string
IP string
Reason string
}
// IPList manages the set of IPs that should see challenge pages.
// Webserver integrations read its maps to redirect IPs to the challenge server.
type IPList struct {
path string
nginxPath string
nginxReload func() error
ips map[string]challengeEntry
mu sync.Mutex
gate PortGate
}
// NewIPList creates an IP list writer for the webserver-facing map at
// mapPath and clears any entries a previous run left in it.
func NewIPList(mapPath string) *IPList {
l := &IPList{
path: mapPath,
ips: make(map[string]challengeEntry),
}
_ = l.flush()
return l
}
// SetPortGate attaches a PortGate so every Add/Remove also opens or
// closes the kernel-level allow. Nil is a no-op (callers don't have to
// branch on whether the gate is configured). Safe to call before any
// Add/Remove; not safe to swap a non-nil gate for another at runtime.
func (l *IPList) SetPortGate(g PortGate) {
l.mu.Lock()
defer l.mu.Unlock()
l.gate = g
}
// SetNginxMap attaches a second map writer for Nginx stacks. The
// callback runs only when the rendered include content changes.
func (l *IPList) SetNginxMap(path string, reload func() error) {
if strings.TrimSpace(path) == "" {
path = DefaultNginxMapPath
}
l.mu.Lock()
l.nginxPath = path
l.nginxReload = reload
changed := l.flush()
l.mu.Unlock()
l.reloadNginx(changed, reload)
}
// Add marks an IP for challenge with the given reason.
func (l *IPList) Add(ip string, reason string, duration time.Duration) {
l.add(ip, reason, duration, false, "")
}
// AddNonEscalating marks an IP for challenge without timeout-to-block escalation.
func (l *IPList) AddNonEscalating(ip string, reason string, duration time.Duration) {
l.add(ip, reason, duration, true, "")
}
// AddWithFindingID preserves the originating audit identity for timeout
// escalation. It does not change challenge duration or escalation policy.
func (l *IPList) AddWithFindingID(ip, reason string, duration time.Duration, findingID string) {
l.add(ip, reason, duration, false, findingID)
}
func (l *IPList) add(ip string, reason string, duration time.Duration, nonEscalating bool, findingID string) {
l.mu.Lock()
l.ips[ip] = challengeEntry{
ExpiresAt: time.Now().Add(duration),
Reason: reason,
NonEscalating: nonEscalating,
FindingID: findingID,
}
changed := l.flush()
gate := l.gate
reload := l.nginxReload
l.mu.Unlock()
if gate != nil {
if err := gate.Allow(ip, duration); err != nil {
fmt.Fprintf(os.Stderr, "challenge: port-gate allow %s: %v\n", ip, err)
}
}
l.reloadNginx(changed, reload)
}
// Remove stops challenging an IP (passed or manually removed).
func (l *IPList) Remove(ip string) {
l.mu.Lock()
if _, listed := l.ips[ip]; !listed {
l.mu.Unlock()
return
}
delete(l.ips, ip)
changed := l.flush()
gate := l.gate
reload := l.nginxReload
l.mu.Unlock()
// Only revoke the kernel gate for an IP that was actually on the list: the
// gate element is created alongside the entry (Add -> Allow), so an IP that
// was never listed has no element to delete. Without this guard, verified
// crawlers -- which bypass the gate and call Remove on every request -- make
// the gate delete a nonexistent element and log an ENOENT error each time.
if gate != nil {
if err := gate.Revoke(ip); err != nil {
fmt.Fprintf(os.Stderr, "challenge: port-gate revoke %s: %v\n", ip, err)
}
}
l.reloadNginx(changed, reload)
}
// Contains returns true if the IP is currently on the challenge list.
func (l *IPList) Contains(ip string) bool {
l.mu.Lock()
defer l.mu.Unlock()
_, ok := l.ips[ip]
return ok
}
// Count returns the number of IPs currently waiting on a challenge.
func (l *IPList) Count() int {
l.mu.Lock()
defer l.mu.Unlock()
return len(l.ips)
}
// ExpiredEntries removes expired entries and returns those eligible for escalation.
// The caller is expected to hard-block returned IPs.
func (l *IPList) ExpiredEntries() []ExpiredEntry {
l.mu.Lock()
now := time.Now()
var expired []ExpiredEntry
removed := false
for ip, entry := range l.ips {
if now.After(entry.ExpiresAt) {
if !entry.NonEscalating {
expired = append(expired, ExpiredEntry{IP: ip, Reason: entry.Reason, FindingID: entry.FindingID})
}
delete(l.ips, ip)
removed = true
}
}
var changed bool
if removed {
changed = l.flush()
}
reload := l.nginxReload
l.mu.Unlock()
l.reloadNginx(changed, reload)
return expired
}
// CleanExpired removes expired entries without returning them.
// Use ExpiredEntries() instead when escalation is needed.
func (l *IPList) CleanExpired() {
_ = l.ExpiredEntries()
}
// flush writes the IP list to disk in each configured webserver format.
// The caller must hold l.mu. It returns true when the Nginx include
// changed and needs a reload.
func (l *IPList) flush() bool {
ips := sortedIPKeys(l.ips)
var sb strings.Builder
sb.WriteString("# CSM Challenge IP list - auto-generated, do not edit\n")
sb.WriteString("# Format: IP challenge (for Apache RewriteMap)\n")
for _, ip := range ips {
fmt.Fprintf(&sb, "%s challenge\n", ip)
}
if err := writeMapFile(l.path, []byte(sb.String())); err != nil {
return false
}
if strings.TrimSpace(l.nginxPath) == "" {
return false
}
var nginx strings.Builder
nginx.WriteString("# CSM Challenge IP list - auto-generated, do not edit\n")
nginx.WriteString("# Format: IP 1; (for Nginx map include)\n")
for _, ip := range ips {
fmt.Fprintf(&nginx, "%s 1;\n", ip)
}
changed, err := writeMapFileIfChanged(l.nginxPath, []byte(nginx.String()))
return err == nil && changed
}
func sortedIPKeys(ips map[string]challengeEntry) []string {
keys := make([]string, 0, len(ips))
for ip := range ips {
keys = append(keys, ip)
}
sort.Strings(keys)
return keys
}
func writeMapFileIfChanged(path string, data []byte) (bool, error) {
// #nosec G304 -- path is the daemon-owned challenge map under /var/cache/csm
// (DefaultNginxMapPath or operator-set via SetNginxMap), never
// attacker-controlled. Read only to diff the rendered map and skip
// rewrite when content is unchanged.
if current, err := os.ReadFile(path); err == nil && bytes.Equal(current, data) {
return false, nil
}
return true, writeMapFile(path, data)
}
func ensureMapDir(path string) error {
// /var/cache/csm must be world-readable so the webserver user
// (www-data / nobody / lsws) can stat + read the map underneath.
// The directory holds no sensitive data; only CSM-owned IP files
// live inside.
//
// MkdirAll respects the process umask, so on cPanel/CloudLinux
// hosts where csm.service inherits umask 027 the directory ends
// up at 0o750 and the webserver gets EACCES on the RewriteMap.
// Explicit Chmod after creation forces the mode the integration
// requires regardless of umask.
mapDir := filepath.Dir(path)
// #nosec G301 -- world-readable rationale above.
if err := os.MkdirAll(mapDir, 0o755); err != nil {
return err
}
// #nosec G302 -- same world-readable rationale; needed for the
// webserver user to stat into the directory and read the map.
if err := os.Chmod(mapDir, 0o755); err != nil {
return err
}
return nil
}
// EnsureMapFile creates a readable empty map when path is absent and leaves
// existing map contents intact. Apache validates txt: RewriteMap sources even
// when challenge mode is disabled, before an IPList would otherwise create the
// file.
func EnsureMapFile(path string) error {
if err := ensureMapDir(path); err != nil {
return err
}
// #nosec G302 G304 G703 -- path is the daemon-owned, webserver-readable challenge map.
f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644)
if err != nil {
if !errors.Is(err, os.ErrExist) {
return err
}
info, statErr := os.Lstat(path)
if statErr != nil {
return statErr
}
if !info.Mode().IsRegular() {
return fmt.Errorf("challenge map %s is not a regular file", path)
}
// #nosec G302 -- Apache/LSWS must be able to read the daemon-owned map.
return os.Chmod(path, 0o644)
}
cleanup := func() {
_ = f.Close()
_ = os.Remove(path)
}
// OpenFile respects umask, so force the public read bits Apache needs.
// #nosec G302 -- Apache/LSWS must be able to read the daemon-owned map.
if err := f.Chmod(0o644); err != nil {
cleanup()
return err
}
if err := f.Close(); err != nil {
_ = os.Remove(path)
return err
}
return nil
}
func writeMapFile(path string, data []byte) error {
if err := ensureMapDir(path); err != nil {
return err
}
// #nosec G306 -- webservers read this daemon-owned IP map directly.
return atomicio.AtomicWrite(path, 0o644, data)
}
func (l *IPList) reloadNginx(changed bool, reload func() error) {
if !changed || reload == nil {
return
}
if err := reload(); err != nil {
fmt.Fprintf(os.Stderr, "challenge: nginx reload after map update: %v\n", err)
}
}
package challenge
import (
"net"
"strings"
"time"
)
// PortGate locks the challenge listener TCP port to specific source IPs
// via the host firewall. An IP is allowed only while it is on the
// challenge IPList (plus operator infra IPs and loopback). Everything
// else gets dropped at the kernel before the listener sees the SYN, so
// the listener is invisible to port scanners and stays reachable only
// for the visitors the daemon has actually redirected.
//
// Implementations are pluggable so the netlink-backed Linux variant
// can be swapped for a stub on platforms that do not have nftables.
// All methods are safe to call on a nil PortGate (no-op), so callers
// do not need to nil-check at every IPList Add/Remove site.
type PortGate interface {
// Allow opens the gate for the source IP for at most ttl. The
// underlying firewall enforces the TTL via the set's own timeout
// so the entry expires even if Revoke is never called (daemon
// crash, missed expiry). Returns nil on success or when the IP
// cannot be parsed (best-effort; the IPList accepts only validated
// IPs upstream, so a parse miss here is a bug to log, not block).
Allow(ip string, ttl time.Duration) error
// Revoke closes the gate for ip immediately. Safe to call for IPs
// that were never on the gate (no-op).
Revoke(ip string) error
// Close tears down the gate's nftables footprint (chain, sets,
// table). The port reverts to whatever the rest of the host
// firewall would do with it.
Close() error
}
// PortGateConfig wraps the inputs the gate needs to install rules.
type PortGateConfig struct {
ListenAddr string
ListenPort int
InfraCIDRs []string
}
// NewPortGate returns the platform-appropriate gate. On Linux it
// installs a dedicated `csm_chal` inet table; on non-Linux it returns
// nil so callers naturally no-op via the nil PortGate handling on the
// IPList side.
//
// Returns nil + nil when the listen address is loopback because no
// gate is needed (loopback traffic cannot originate from off-host).
// Caller treats nil as "gate not active" and proceeds without it.
func NewPortGate(cfg PortGateConfig) (PortGate, error) {
if isLoopbackListenAddr(cfg.ListenAddr) {
return nil, nil
}
return newPortGate(cfg)
}
// portGateFamily picks which address families the gate should bind to
// based on the listen address. 0.0.0.0 / blank -> v4 only; :: -> dual
// stack; a specific literal IP gates only that family.
type portGateFamily struct {
v4 bool
v6 bool
}
func familyForListenAddr(addr string) portGateFamily {
addr = strings.Trim(strings.TrimSpace(addr), "[]")
switch addr {
case "", "0.0.0.0":
return portGateFamily{v4: true}
case "::":
return portGateFamily{v4: true, v6: true}
}
ip := net.ParseIP(addr)
if ip == nil {
return portGateFamily{v4: true}
}
if ip.To4() != nil {
return portGateFamily{v4: true}
}
return portGateFamily{v6: true}
}
func portGateFamilyAcceptsIP(fam portGateFamily, ip net.IP) bool {
if ip == nil {
return false
}
if ip.To4() != nil {
return fam.v4
}
return fam.v6
}
//go:build linux
package challenge
import (
"errors"
"fmt"
"net"
"sync"
"syscall"
"time"
"github.com/google/nftables"
"github.com/google/nftables/binaryutil"
"github.com/google/nftables/expr"
"golang.org/x/sys/unix"
)
// linuxPortGate owns the nftables `csm_chal` table. The table is kept
// separate from the firewall package's `csm` table so the two can be
// installed and torn down independently (challenge can run with or
// without csm.firewall enabled).
type linuxPortGate struct {
mu sync.Mutex
conn *nftables.Conn
cfg PortGateConfig
family portGateFamily
table *nftables.Table
setChalIPs *nftables.Set
setChalIPs6 *nftables.Set
setInfra *nftables.Set
setInfra6 *nftables.Set
}
func newPortGate(cfg PortGateConfig) (PortGate, error) {
if cfg.ListenPort <= 0 || cfg.ListenPort > 65535 {
return nil, fmt.Errorf("port-gate: invalid listen port %d", cfg.ListenPort)
}
fam := familyForListenAddr(cfg.ListenAddr)
conn, err := nftables.New()
if err != nil {
return nil, fmt.Errorf("port-gate: nftables open: %w", err)
}
g := &linuxPortGate{conn: conn, cfg: cfg, family: fam}
if err := g.install(); err != nil {
return nil, err
}
return g, nil
}
// install lays down the table, sets, chain, and per-port rules. Any
// pre-existing `csm_chal` table is deleted first so a daemon restart
// always converges on a clean rule shape (and stale rules from a
// crashed previous run do not linger).
func (g *linuxPortGate) install() error {
g.mu.Lock()
defer g.mu.Unlock()
g.dropExistingTableLocked()
g.table = g.conn.AddTable(&nftables.Table{
Family: nftables.TableFamilyINet,
Name: "csm_chal",
})
if g.family.v4 {
g.setChalIPs = &nftables.Set{
Table: g.table,
Name: "chal_ips",
KeyType: nftables.TypeIPAddr,
HasTimeout: true,
}
if err := g.conn.AddSet(g.setChalIPs, nil); err != nil {
return fmt.Errorf("port-gate: add chal_ips: %w", err)
}
g.setInfra = &nftables.Set{
Table: g.table,
Name: "chal_infra",
KeyType: nftables.TypeIPAddr,
Interval: true,
}
if err := g.conn.AddSet(g.setInfra, infraElementsV4(g.cfg.InfraCIDRs)); err != nil {
return fmt.Errorf("port-gate: add chal_infra: %w", err)
}
}
if g.family.v6 {
g.setChalIPs6 = &nftables.Set{
Table: g.table,
Name: "chal_ips6",
KeyType: nftables.TypeIP6Addr,
HasTimeout: true,
}
if err := g.conn.AddSet(g.setChalIPs6, nil); err != nil {
return fmt.Errorf("port-gate: add chal_ips6: %w", err)
}
g.setInfra6 = &nftables.Set{
Table: g.table,
Name: "chal_infra6",
KeyType: nftables.TypeIP6Addr,
Interval: true,
}
if err := g.conn.AddSet(g.setInfra6, infraElementsV6(g.cfg.InfraCIDRs)); err != nil {
return fmt.Errorf("port-gate: add chal_infra6: %w", err)
}
}
prio := nftables.ChainPriority(-200)
policy := nftables.ChainPolicyAccept
chain := g.conn.AddChain(&nftables.Chain{
Name: "challenge_gate",
Table: g.table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookInput,
Priority: &prio,
Policy: &policy,
})
g.addAcceptRules(chain)
g.addDropRule(chain)
if err := g.conn.Flush(); err != nil {
return fmt.Errorf("port-gate: install flush: %w", err)
}
return nil
}
// dropExistingTableLocked is idempotent: ListTables + DelTable so a
// re-install does not stack rules on top of a stale chain.
func (g *linuxPortGate) dropExistingTableLocked() {
tables, err := g.conn.ListTables()
if err != nil {
return
}
for _, t := range tables {
if t.Family == nftables.TableFamilyINet && t.Name == "csm_chal" {
g.conn.DelTable(t)
_ = g.conn.Flush()
return
}
}
}
func (g *linuxPortGate) addAcceptRules(chain *nftables.Chain) {
port := portU16(g.cfg.ListenPort)
if g.family.v4 {
// loopback bypass
g.addRule(chain, exprsTCPDportFromV4(port, net.IPv4(127, 0, 0, 0).To4(), net.IPv4Mask(255, 0, 0, 0), expr.VerdictAccept))
g.addRule(chain, exprsTCPDportSetMatchV4(port, g.setInfra, expr.VerdictAccept))
g.addRule(chain, exprsTCPDportSetMatchV4(port, g.setChalIPs, expr.VerdictAccept))
}
if g.family.v6 {
loop6 := net.ParseIP("::1").To16()
mask128 := net.CIDRMask(128, 128)
g.addRule(chain, exprsTCPDportFromV6(port, loop6, mask128, expr.VerdictAccept))
g.addRule(chain, exprsTCPDportSetMatchV6(port, g.setInfra6, expr.VerdictAccept))
g.addRule(chain, exprsTCPDportSetMatchV6(port, g.setChalIPs6, expr.VerdictAccept))
}
}
func (g *linuxPortGate) addDropRule(chain *nftables.Chain) {
port := portU16(g.cfg.ListenPort)
// Any packet that reached this rule with dport == challenge port
// did not match an accept above; drop it.
g.conn.AddRule(&nftables.Rule{
Table: g.table,
Chain: chain,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(port)},
&expr.Verdict{Kind: expr.VerdictDrop},
},
})
}
func (g *linuxPortGate) addRule(chain *nftables.Chain, exprs []expr.Any) {
g.conn.AddRule(&nftables.Rule{Table: g.table, Chain: chain, Exprs: exprs})
}
func (g *linuxPortGate) Allow(ip string, ttl time.Duration) error {
parsed := net.ParseIP(ip)
if parsed == nil {
return fmt.Errorf("port-gate: invalid ip %q", ip)
}
if ttl <= 0 {
ttl = 30 * time.Minute
}
if !portGateFamilyAcceptsIP(g.family, parsed) {
return nil
}
g.mu.Lock()
defer g.mu.Unlock()
if ip4 := parsed.To4(); ip4 != nil && g.setChalIPs != nil {
if err := g.conn.SetAddElements(g.setChalIPs, []nftables.SetElement{
{Key: ip4, Timeout: ttl},
}); err != nil {
return fmt.Errorf("port-gate: add v4 %s: %w", ip, err)
}
return g.conn.Flush()
}
if g.setChalIPs6 != nil {
if err := g.conn.SetAddElements(g.setChalIPs6, []nftables.SetElement{
{Key: parsed.To16(), Timeout: ttl},
}); err != nil {
return fmt.Errorf("port-gate: add v6 %s: %w", ip, err)
}
return g.conn.Flush()
}
// IP family is not gated (e.g., v6 IP on a v4-only listener); silently
// no-op so the IPList Add path does not surface an unactionable error.
return nil
}
func (g *linuxPortGate) Revoke(ip string) error {
parsed := net.ParseIP(ip)
if parsed == nil {
return fmt.Errorf("port-gate: invalid ip %q", ip)
}
if !portGateFamilyAcceptsIP(g.family, parsed) {
return nil
}
g.mu.Lock()
defer g.mu.Unlock()
if ip4 := parsed.To4(); ip4 != nil && g.setChalIPs != nil {
if err := g.conn.SetDeleteElements(g.setChalIPs, []nftables.SetElement{{Key: ip4}}); err != nil {
return fmt.Errorf("port-gate: del v4 %s: %w", ip, err)
}
return ignoreNftNotFound(g.conn.Flush())
}
if g.setChalIPs6 != nil {
if err := g.conn.SetDeleteElements(g.setChalIPs6, []nftables.SetElement{{Key: parsed.To16()}}); err != nil {
return fmt.Errorf("port-gate: del v6 %s: %w", ip, err)
}
return ignoreNftNotFound(g.conn.Flush())
}
return nil
}
// ignoreNftNotFound treats a "no such file or directory" netlink error from a
// set-element delete as success. Gate elements carry a TTL, so the kernel may
// have already expired the element by the time Revoke runs; deleting an absent
// element is a benign no-op, not a failure worth surfacing.
func ignoreNftNotFound(err error) error {
if errors.Is(err, syscall.ENOENT) {
return nil
}
return err
}
func (g *linuxPortGate) Close() error {
g.mu.Lock()
defer g.mu.Unlock()
if g.table == nil {
return nil
}
g.conn.DelTable(g.table)
if err := g.conn.Flush(); err != nil {
return fmt.Errorf("port-gate: close flush: %w", err)
}
g.table = nil
g.setChalIPs = nil
g.setChalIPs6 = nil
g.setInfra = nil
g.setInfra6 = nil
return nil
}
func portU16(p int) uint16 {
if p < 0 || p > 65535 {
return 0
}
// #nosec G115 -- bounds-checked above.
return uint16(p)
}
// infraElementsV4 builds nftables interval-set elements from the
// operator's infra_ips list, keeping only IPv4 entries. The interval
// set wants [start, end) pairs; net.ParseCIDR + binaryutil pack them
// the same way the firewall engine's infra set does.
func infraElementsV4(cidrs []string) []nftables.SetElement {
var out []nftables.SetElement
for _, raw := range cidrs {
ipnet := parseCIDROrIP(raw)
if ipnet == nil {
continue
}
start := ipnet.IP.To4()
if start == nil {
continue
}
end := lastIPv4(ipnet)
out = append(out,
nftables.SetElement{Key: start},
nftables.SetElement{Key: ipv4Inc(end), IntervalEnd: true},
)
}
return out
}
func infraElementsV6(cidrs []string) []nftables.SetElement {
var out []nftables.SetElement
for _, raw := range cidrs {
ipnet := parseCIDROrIP(raw)
if ipnet == nil {
continue
}
if ipnet.IP.To4() != nil {
continue
}
start := ipnet.IP.To16()
end := lastIPv6(ipnet)
out = append(out,
nftables.SetElement{Key: start},
nftables.SetElement{Key: ipv6Inc(end), IntervalEnd: true},
)
}
return out
}
// parseCIDROrIP accepts both "1.2.3.4" and "1.2.3.0/24". A bare IP is
// treated as a /32 (or /128 for v6).
func parseCIDROrIP(raw string) *net.IPNet {
if _, ipnet, err := net.ParseCIDR(raw); err == nil {
return ipnet
}
ip := net.ParseIP(raw)
if ip == nil {
return nil
}
if v4 := ip.To4(); v4 != nil {
return &net.IPNet{IP: v4, Mask: net.CIDRMask(32, 32)}
}
return &net.IPNet{IP: ip.To16(), Mask: net.CIDRMask(128, 128)}
}
func lastIPv4(n *net.IPNet) net.IP {
ip := n.IP.To4()
out := make(net.IP, 4)
for i := 0; i < 4; i++ {
out[i] = ip[i] | ^n.Mask[i]
}
return out
}
func lastIPv6(n *net.IPNet) net.IP {
ip := n.IP.To16()
out := make(net.IP, 16)
for i := 0; i < 16; i++ {
out[i] = ip[i] | ^n.Mask[i]
}
return out
}
func ipv4Inc(ip net.IP) net.IP {
out := make(net.IP, 4)
copy(out, ip.To4())
for i := 3; i >= 0; i-- {
out[i]++
if out[i] != 0 {
return out
}
}
return out
}
func ipv6Inc(ip net.IP) net.IP {
out := make(net.IP, 16)
copy(out, ip.To16())
for i := 15; i >= 0; i-- {
out[i]++
if out[i] != 0 {
return out
}
}
return out
}
// exprsTCPDportFromV4 produces "L4=TCP, src in CIDR, dport=port -> verdict".
// Mask is the 4-byte IPv4 subnet mask.
func exprsTCPDportFromV4(port uint16, network net.IP, mask net.IPMask, verdict expr.VerdictKind) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.NFPROTO_IPV4}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Bitwise{SourceRegister: 1, DestRegister: 1, Len: 4, Mask: mask, Xor: []byte{0, 0, 0, 0}},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: network},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(port)},
&expr.Verdict{Kind: verdict},
}
}
func exprsTCPDportFromV6(port uint16, network net.IP, mask net.IPMask, verdict expr.VerdictKind) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.NFPROTO_IPV6}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 8, Len: 16},
&expr.Bitwise{SourceRegister: 1, DestRegister: 1, Len: 16, Mask: mask, Xor: make([]byte, 16)},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: network},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(port)},
&expr.Verdict{Kind: verdict},
}
}
func exprsTCPDportSetMatchV4(port uint16, set *nftables.Set, verdict expr.VerdictKind) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.NFPROTO_IPV4}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Lookup{SourceRegister: 1, SetName: set.Name, SetID: set.ID},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(port)},
&expr.Verdict{Kind: verdict},
}
}
func exprsTCPDportSetMatchV6(port uint16, set *nftables.Set, verdict expr.VerdictKind) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.IPPROTO_TCP}},
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{unix.NFPROTO_IPV6}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 8, Len: 16},
&expr.Lookup{SourceRegister: 1, SetName: set.Name, SetID: set.ID},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(port)},
&expr.Verdict{Kind: verdict},
}
}
package challenge
import (
"context"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"crypto/tls"
"encoding/hex"
"fmt"
"html"
"net"
"net/http"
"net/url"
"os"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
)
// Server serves challenge pages to gray-listed IPs. Passing the challenge
// only takes the IP off the challenge list and hands the browser a bypass
// cookie; it grants nothing in the firewall. The list is the sole source of
// truth for who has a challenge pending, so every handler answers only for
// IPs on it -- serving the puzzle to anyone else would let an arbitrary
// client turn one cheap hash search into a verified session.
type Server struct {
cfg *config.Config
secret []byte
ipList *IPList
srv *http.Server
trustedProxies map[string]bool
// Track recently verified IPs to prevent replay
verified map[string]time.Time
verifiedMu sync.Mutex
// Optional subsystems wired in by configuration. Any of these may
// be nil; the handlers branch accordingly so a fresh deployment
// without the new blocks behaves exactly as before.
captcha CaptchaProvider
sessionSigner *AdminSessionSigner
crawlers *CrawlerVerifier
// verifySigner signs the csm_verified cookie handed to every visitor
// who passes the PoW/CAPTCHA. It binds the cookie to one IP and an
// expiry so a returning visitor skips the gate for the allow window
// without it being replayable elsewhere. Always constructed; the key
// rotates on restart like sessionSigner.
verifySigner *AdminSessionSigner
// adminFailures tracks failed admin-token submissions per source
// IP for rate-limiting brute-force probes. Sliding window: an IP
// that hits adminMaxFailuresInWindow within adminFailureWindow
// gets locked out (subsequent submissions return 429) until the
// oldest failure ages out.
adminFailures map[string][]time.Time
adminFailuresMu sync.Mutex
}
const (
adminFailureWindow = 5 * time.Minute
adminMaxFailuresInWindow = 5
// verifyCookieTTL is the lifetime of the csm_verified bypass cookie: how
// long a visitor who passed once is waved through if listed again.
verifyCookieTTL = 4 * time.Hour
)
// New creates a challenge server.
func New(cfg *config.Config, ipList *IPList) *Server {
secret := []byte(cfg.Challenge.Secret)
if len(secret) == 0 {
secret = make([]byte, 32)
_, _ = rand.Read(secret)
}
trusted := make(map[string]bool)
for _, p := range cfg.Challenge.TrustedProxies {
p = strings.TrimSpace(p)
if p == "" {
continue
}
trusted[canonicalIP(p)] = true
}
s := &Server{
cfg: cfg,
secret: secret,
ipList: ipList,
trustedProxies: trusted,
verified: make(map[string]time.Time),
adminFailures: make(map[string][]time.Time),
}
// The verify-cookie signer is always on (independent of the optional
// admin-session feature).
if vs, err := NewAdminSessionSigner(verifyCookieTTL); err != nil {
fmt.Fprintf(os.Stderr, "[challenge] verify-cookie signing disabled: %v\n", err)
} else {
s.verifySigner = vs
}
// Optional sub-features. Each is opt-in via its own config block;
// initialization errors degrade to "feature off" with a stderr
// message rather than refusing to start the challenge server.
if name := cfg.Challenge.CaptchaFallback.Provider; name != "" {
p, err := NewCaptchaProvider(name, cfg.Challenge.CaptchaFallback.SecretKey, cfg.Challenge.CaptchaFallback.Timeout)
if err != nil {
fmt.Fprintf(os.Stderr, "[challenge] captcha disabled: %v\n", err)
}
s.captcha = p
}
if cfg.Challenge.VerifiedSession.Enabled {
signer, err := NewAdminSessionSigner(cfg.Challenge.VerifiedSession.TTL)
if err != nil {
fmt.Fprintf(os.Stderr, "[challenge] verified-session disabled: %v\n", err)
}
s.sessionSigner = signer
}
if cfg.Challenge.VerifiedCrawlers.Enabled {
s.crawlers = NewCrawlerVerifier(
cfg.Challenge.VerifiedCrawlers.Providers,
cfg.Challenge.VerifiedCrawlers.CacheTTL,
nil, // net.DefaultResolver
)
}
mux := http.NewServeMux()
mux.HandleFunc("/challenge", s.handleChallenge)
mux.HandleFunc("/challenge/gate", s.handleGate)
mux.HandleFunc("/challenge/verify", s.handleVerify)
mux.HandleFunc("/challenge/captcha-verify", s.handleCaptchaVerify)
mux.HandleFunc("/challenge/admin-token", s.handleAdminToken)
bindAddr := cfg.Challenge.ListenAddr
if bindAddr == "" {
// Defaults are applied during config.Load, but tests construct
// Server directly with an empty config; default to loopback so
// the production safety guarantee (never bind public by default)
// holds even on those paths.
bindAddr = "127.0.0.1"
}
s.srv = &http.Server{
Addr: net.JoinHostPort(bindAddr, strconv.Itoa(cfg.Challenge.ListenPort)),
Handler: mux,
ReadTimeout: 10 * time.Second,
WriteTimeout: 10 * time.Second,
ReadHeaderTimeout: 5 * time.Second,
}
return s
}
// Listen validates configured TLS material and binds the challenge address.
// Daemon startup calls this synchronously before publishing the challenge IP
// list to response routing, so a bind or certificate error fails open to
// alerting instead of treating an unavailable challenge as a handled finding.
func (s *Server) Listen() (net.Listener, error) {
cert, key := s.resolveTLSMaterial()
if cert != "" && key != "" {
if _, err := tls.LoadX509KeyPair(cert, key); err != nil {
return nil, fmt.Errorf("challenge TLS material: %w", err)
}
}
listener, err := net.Listen("tcp", s.srv.Addr)
if err != nil {
return nil, fmt.Errorf("challenge listen %s: %w", s.srv.Addr, err)
}
return listener, nil
}
// Serve begins serving challenge pages on an already-bound listener. Explicit
// challenge TLS makes the listener HTTPS. Direct/public listeners can reuse
// the WebUI TLS pair. Loopback listeners stay plain HTTP by default.
//
// Resolution order:
// 1. challenge.tls_cert + challenge.tls_key (explicit per-service)
// 2. webui.tls_cert + webui.tls_key (direct/public binds only)
// 3. plain HTTP (loopback-only default)
func (s *Server) Serve(listener net.Listener) error {
cert, key := s.resolveTLSMaterial()
if cert == "" || key == "" {
if !isLoopbackListenAddr(s.cfg.Challenge.ListenAddr) {
fmt.Fprintf(os.Stderr,
"[%s] WARNING: public challenge listener has no complete TLS cert/key configured; HSTS-pinned domains will fail with ERR_SSL_PROTOCOL_ERROR\n",
time.Now().Format("2006-01-02 15:04:05"))
}
return s.srv.Serve(listener)
}
return s.srv.ServeTLS(listener, cert, key)
}
// Start binds and serves in one call for standalone users. The daemon uses
// Listen followed by Serve so it can publish challenge routing only after the
// listener is known to be available.
func (s *Server) Start() error {
listener, err := s.Listen()
if err != nil {
return err
}
return s.Serve(listener)
}
// resolveTLSMaterial picks the cert / key pair the challenge listener
// should present. The WebUI fallback is only safe for direct/public binds;
// loopback listeners stay plain HTTP unless challenge TLS is explicitly
// configured.
func (s *Server) resolveTLSMaterial() (cert, key string) {
if c, k := s.cfg.Challenge.TLSCert, s.cfg.Challenge.TLSKey; c != "" && k != "" {
return c, k
}
if isLoopbackListenAddr(s.cfg.Challenge.ListenAddr) {
return "", ""
}
if c, k := s.cfg.WebUI.TLSCert, s.cfg.WebUI.TLSKey; c != "" && k != "" {
return c, k
}
return "", ""
}
func isLoopbackListenAddr(addr string) bool {
addr = strings.TrimSpace(addr)
if addr == "" || strings.EqualFold(addr, "localhost") {
return true
}
ip := net.ParseIP(addr)
return ip != nil && ip.IsLoopback()
}
// Shutdown gracefully stops the server.
func (s *Server) Shutdown() {
_ = s.srv.Close()
}
func (s *Server) handleChallenge(w http.ResponseWriter, r *http.Request) {
ip := s.extractIP(r)
if !s.challengePending(ip) {
http.Error(w, "No challenge pending for this address", http.StatusNotFound)
return
}
// Bypass paths run before the PoW page is generated. Each
// short-circuits to the same markVerified flow, so a passing
// visitor never sees the challenge UI even once.
if s.bypassByAdminCookie(r, ip) || s.bypassByVerifiedCrawler(r.Context(), ip) || s.bypassByVerifyCookie(r, ip) {
s.markVerified(w, r, ip, "")
return
}
nonce := generateNonce()
difficulty := s.cfg.Challenge.Difficulty
token := s.makeToken(ip, nonce)
w.Header().Set("Content-Type", "text/html; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
captchaBlock := s.captchaNoscriptHTML(token, nonce)
// #nosec G705 -- All interpolated values are constrained to non-HTML
// character sets: ip is validated via net.ParseIP / net.SplitHostPort
// in extractIP (never attacker-controlled string); nonce and token
// are hex.EncodeToString output (0-9, a-f only); difficulty is an int.
// captchaBlock is constructed from configured site keys (operator
// supplied) plus hex token/nonce; it never includes attacker input.
// html.EscapeString on ip is defence-in-depth so static analysers can
// see the request-derived value is rendered safe even if a future
// change to extractIP loosened validation.
fmt.Fprintf(w, challengePageHTML, html.EscapeString(ip), captchaBlock, nonce, token, difficulty, difficulty)
}
func (s *Server) handleGate(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
ip := s.extractIP(r)
if s.bypassByAdminCookie(r, ip) || s.bypassByVerifiedCrawler(r.Context(), ip) || s.bypassByVerifyCookie(r, ip) {
w.WriteHeader(http.StatusNoContent)
return
}
if s.ipList == nil || !s.ipList.Contains(ip) {
w.WriteHeader(http.StatusNoContent)
return
}
w.WriteHeader(http.StatusUnauthorized)
}
// bypassByAdminCookie returns true when the visitor presents a valid
// signed-session cookie issued by THIS daemon for THIS IP. Daemon
// restart rotates the signing key, so old cookies fall back to the
// normal PoW flow automatically.
func (s *Server) bypassByAdminCookie(r *http.Request, ip string) bool {
if s.sessionSigner == nil {
return false
}
cookieName := s.cfg.Challenge.VerifiedSession.CookieName
if cookieName == "" {
cookieName = "csm_admin_session"
}
c, err := r.Cookie(cookieName)
if err != nil || c.Value == "" {
return false
}
return s.sessionSigner.Verify(c.Value, ip) == nil
}
// bypassByVerifiedCrawler resolves the visitor's reverse DNS and
// confirms the forward lookup -- skipping the PoW only for traffic
// from the configured crawler families. A spoofed UA from a residential
// IP fails forward-confirm and falls through to PoW.
func (s *Server) bypassByVerifiedCrawler(ctx context.Context, ip string) bool {
if s.crawlers == nil {
return false
}
return s.crawlers.Verified(ctx, ip)
}
func (s *Server) handleVerify(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
ip := s.extractIP(r)
if !s.challengePending(ip) {
http.Error(w, "No challenge pending for this address", http.StatusNotFound)
return
}
nonce := r.FormValue("nonce")
token := r.FormValue("token")
solution := r.FormValue("solution")
// Verify HMAC token
expected := s.makeToken(ip, nonce)
if token != expected {
http.Error(w, "Invalid token", http.StatusForbidden)
return
}
// Verify proof-of-work solution
if !verifyPoW(nonce, solution, s.cfg.Challenge.Difficulty) {
http.Error(w, "Invalid solution", http.StatusForbidden)
return
}
// Prevent replay -- one nonce, one verification.
s.verifiedMu.Lock()
if _, seen := s.verified[nonce]; seen {
s.verifiedMu.Unlock()
http.Error(w, "Token already used", http.StatusForbidden)
return
}
s.verified[nonce] = time.Now()
s.verifiedMu.Unlock()
s.markVerified(w, r, ip, r.FormValue("dest"))
}
// handleCaptchaVerify accepts a provider token, validates it
// server-side, and (on success) puts the visitor through the same
// markVerified flow PoW uses. Available only when the operator has
// configured a CAPTCHA provider.
func (s *Server) handleCaptchaVerify(w http.ResponseWriter, r *http.Request) {
if s.captcha == nil {
http.NotFound(w, r)
return
}
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
ip := s.extractIP(r)
if !s.challengePending(ip) {
http.Error(w, "No challenge pending for this address", http.StatusNotFound)
return
}
nonce := r.FormValue("nonce")
token := r.FormValue("token")
captchaToken := r.FormValue("captcha-token")
// Bind the CAPTCHA submission to the page token before asking the
// provider. The nonce is spent only after the provider accepts, so a
// rejected widget token can be retried from the same page.
if expected := s.makeToken(ip, nonce); token != expected {
http.Error(w, "Invalid token", http.StatusForbidden)
return
}
s.verifiedMu.Lock()
_, seen := s.verified[nonce]
s.verifiedMu.Unlock()
if seen {
http.Error(w, "Token already used", http.StatusForbidden)
return
}
ok, err := s.captcha.Verify(r.Context(), captchaToken, ip)
if err != nil || !ok {
http.Error(w, "CAPTCHA verification failed", http.StatusForbidden)
return
}
s.verifiedMu.Lock()
if _, seen := s.verified[nonce]; seen {
s.verifiedMu.Unlock()
http.Error(w, "Token already used", http.StatusForbidden)
return
}
s.verified[nonce] = time.Now()
s.verifiedMu.Unlock()
s.markVerified(w, r, ip, r.FormValue("dest"))
}
// handleAdminToken issues a signed-session cookie when the operator
// presents the configured admin_secret. Returns 204 with Set-Cookie on
// success, 403 on bad/missing secret, 429 once an IP has burned
// through adminMaxFailuresInWindow failures inside adminFailureWindow.
// The 429 is returned BEFORE the constant-time compare so an attacker
// cannot keep probing once they are throttled.
func (s *Server) handleAdminToken(w http.ResponseWriter, r *http.Request) {
if s.sessionSigner == nil {
http.NotFound(w, r)
return
}
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
ip := s.extractIP(r)
if s.adminRateLimited(ip) {
http.Error(w, "Too many failed attempts; try again later.", http.StatusTooManyRequests)
return
}
presented := r.FormValue("secret")
if !CompareAdminSecret(s.cfg.Challenge.VerifiedSession.AdminSecret, presented) {
s.recordAdminFailure(ip)
http.Error(w, "Forbidden", http.StatusForbidden)
return
}
// Successful auth clears the failure log so a legitimate operator
// who fat-fingered the secret a few times can keep using the
// endpoint after they get it right.
s.clearAdminFailures(ip)
cookieName := s.cfg.Challenge.VerifiedSession.CookieName
if cookieName == "" {
cookieName = "csm_admin_session"
}
cookie := &http.Cookie{
Name: cookieName,
Value: s.sessionSigner.Issue(ip),
Path: "/",
MaxAge: int(s.sessionSigner.TTL().Seconds()),
Secure: true,
HttpOnly: true,
SameSite: http.SameSiteLaxMode,
}
http.SetCookie(w, cookie)
w.WriteHeader(http.StatusNoContent)
}
// adminRateLimited returns true when ip has at least
// adminMaxFailuresInWindow failures in the last adminFailureWindow.
// Also opportunistically prunes aged-out entries so the per-IP slice
// stays bounded.
func (s *Server) adminRateLimited(ip string) bool {
cutoff := time.Now().Add(-adminFailureWindow)
s.adminFailuresMu.Lock()
defer s.adminFailuresMu.Unlock()
pruned := s.adminFailures[ip][:0]
for _, t := range s.adminFailures[ip] {
if t.After(cutoff) {
pruned = append(pruned, t)
}
}
s.adminFailures[ip] = pruned
return len(pruned) >= adminMaxFailuresInWindow
}
func (s *Server) recordAdminFailure(ip string) {
s.adminFailuresMu.Lock()
defer s.adminFailuresMu.Unlock()
s.adminFailures[ip] = append(s.adminFailures[ip], time.Now())
}
func (s *Server) clearAdminFailures(ip string) {
s.adminFailuresMu.Lock()
defer s.adminFailuresMu.Unlock()
delete(s.adminFailures, ip)
}
// challengePending reports whether ip currently has a challenge to answer.
// Without a list nobody does.
func (s *Server) challengePending(ip string) bool {
return s.ipList != nil && s.ipList.Contains(ip)
}
// markVerified is the shared post-success path used by handleVerify,
// handleCaptchaVerify, and the bypass shortcuts in handleChallenge.
// Centralising the side effects (ipList removal, verification cookie,
// redirect render) keeps all four paths in sync; the alternative --
// copy-pasting four times -- is the easiest way to drift the behaviour of
// one path away from the others over time. Nothing here touches the
// firewall: a passed challenge proves a browser, not a trustworthy client,
// so the visitor goes back to the ordinary rules like everyone else.
func (s *Server) markVerified(w http.ResponseWriter, r *http.Request, ip, destOverride string) {
s.ipList.Remove(ip)
// Set verification cookie so a visitor who is listed again inside the
// window skips the puzzle. Secure is always on: CSM is designed to run
// behind HTTPS and the cookie grants a multi-hour bypass of the PoW
// gate, so leaking it over plaintext is never acceptable.
if s.verifySigner != nil {
http.SetCookie(w, &http.Cookie{
Name: "csm_verified",
Value: s.verifySigner.Issue(ip),
Path: "/",
MaxAge: int(verifyCookieTTL.Seconds()),
Secure: true,
HttpOnly: true,
SameSite: http.SameSiteLaxMode,
})
}
dest := sanitizeRedirectDest(destOverride, r.Host)
w.Header().Set("Content-Type", "text/html; charset=utf-8")
fmt.Fprintf(w, `<!DOCTYPE html><html><head><meta http-equiv="refresh" content="2;url=%s">
<style>body{font-family:system-ui;display:flex;justify-content:center;align-items:center;min-height:100vh;margin:0;background:#1a2234;color:#c8d3e0}
.ok{text-align:center}.ok h1{color:#2fb344;font-size:3em}p{font-size:1.2em}</style>
</head><body><div class="ok"><h1>✓</h1><p>Verified - redirecting...</p></div></body></html>`, html.EscapeString(dest))
}
// captchaNoscriptHTML renders the CAPTCHA fallback widget. Returns the
// empty string when no provider is configured, in which case the
// challenge page's <noscript> block falls back to the existing
// "JavaScript is required" message.
//
// token and nonce are guaranteed hex by upstream constructors; siteKey
// is operator-supplied so we escape it to keep a typo from breaking
// the page (an attacker would need write access to csm.yaml to inject
// HTML here, which is already game-over, but defence-in-depth is
// cheap).
func (s *Server) captchaNoscriptHTML(token, nonce string) string {
if s.captcha == nil {
return ""
}
siteKey := html.EscapeString(s.cfg.Challenge.CaptchaFallback.SiteKey)
switch s.captcha.Name() {
case "turnstile":
return fmt.Sprintf(`<script src="https://challenges.cloudflare.com/turnstile/v0/api.js" async defer></script>
<form method="POST" action="/challenge/captcha-verify">
<input type="hidden" name="token" value="%s">
<input type="hidden" name="nonce" value="%s">
<div class="cf-turnstile" data-sitekey="%s" data-callback="csmCaptchaCallback"></div>
<input type="hidden" name="captcha-token" id="captchaToken">
</form>
<script>function csmCaptchaCallback(t){document.getElementById('captchaToken').value=t;document.forms[0].submit();}</script>`,
token, nonce, siteKey)
case "hcaptcha":
return fmt.Sprintf(`<script src="https://js.hcaptcha.com/1/api.js" async defer></script>
<form method="POST" action="/challenge/captcha-verify">
<input type="hidden" name="token" value="%s">
<input type="hidden" name="nonce" value="%s">
<div class="h-captcha" data-sitekey="%s" data-callback="csmCaptchaCallback"></div>
<input type="hidden" name="captcha-token" id="captchaToken">
</form>
<script>function csmCaptchaCallback(t){document.getElementById('captchaToken').value=t;document.forms[0].submit();}</script>`,
token, nonce, siteKey)
default:
return ""
}
}
func (s *Server) makeToken(ip, nonce string) string {
mac := hmac.New(sha256.New, s.secret)
mac.Write([]byte(ip + ":" + nonce))
return hex.EncodeToString(mac.Sum(nil))
}
// bypassByVerifyCookie returns true when the visitor presents a valid
// csm_verified cookie issued by this daemon for this IP. The cookie binds
// to one IP and an expiry, so it cannot be replayed from another network or
// after the allow window; a daemon restart rotates the signing key and
// invalidates every outstanding cookie.
func (s *Server) bypassByVerifyCookie(r *http.Request, ip string) bool {
if s.verifySigner == nil {
return false
}
c, err := r.Cookie("csm_verified")
if err != nil || c.Value == "" {
return false
}
return s.verifySigner.Verify(c.Value, ip) == nil
}
// CleanExpired removes old verification records, prunes the
// admin-failure log, and evicts stale crawler-cache entries. Called
// from the daemon's challengeEscalator ticker every 60 seconds; under
// a sustained scan from many source IPs, this is the only thing
// keeping per-IP map entries from accumulating until restart.
func (s *Server) CleanExpired() {
now := time.Now()
s.verifiedMu.Lock()
verifiedCutoff := now.Add(-4 * time.Hour)
for k, t := range s.verified {
if t.Before(verifiedCutoff) {
delete(s.verified, k)
}
}
s.verifiedMu.Unlock()
// Drop admin-failure entries whose latest failure has aged out of
// the rate-limit window. An IP that hammered the endpoint once
// and stopped will otherwise sit in the map forever.
s.adminFailuresMu.Lock()
failureCutoff := now.Add(-adminFailureWindow)
for ip, times := range s.adminFailures {
kept := times[:0]
for _, t := range times {
if t.After(failureCutoff) {
kept = append(kept, t)
}
}
if len(kept) == 0 {
delete(s.adminFailures, ip)
} else {
s.adminFailures[ip] = kept
}
}
s.adminFailuresMu.Unlock()
if s.crawlers != nil {
s.crawlers.cleanExpired(now)
}
}
// extractIP returns the client IP from the request. X-Forwarded-For is only
// trusted when the direct peer is in the configured trusted_proxies list.
// Uses the rightmost XFF entry (the one appended by the trusted proxy),
// not the leftmost (which the client controls). Without trusted proxies,
// RemoteAddr is always used — this prevents attackers from spoofing their IP
// to mint firewall allow rules for arbitrary addresses.
func (s *Server) extractIP(r *http.Request) string {
peerIP, _, _ := net.SplitHostPort(r.RemoteAddr)
peerIP = canonicalIP(peerIP)
if len(s.trustedProxies) > 0 && s.trustedProxies[peerIP] {
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
parts := strings.Split(xff, ",")
// Use rightmost entry — the one the trusted proxy appended
for i := len(parts) - 1; i >= 0; i-- {
ip := strings.TrimSpace(parts[i])
if net.ParseIP(ip) != nil {
return canonicalIP(ip)
}
}
}
}
return peerIP
}
// canonicalIP normalizes an IP string so downstream exact-match lookups (the
// firewall's blocked-IP index) agree on one form. A dual-stack listener
// reports an IPv4 peer as "::ffff:1.2.3.4"; net.IP.String collapses that to
// the plain IPv4 form the firewall stores. Non-parseable input is returned
// unchanged.
func canonicalIP(ip string) string {
if parsed := net.ParseIP(ip); parsed != nil {
return parsed.String()
}
return ip
}
func generateNonce() string {
b := make([]byte, 16)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
// sanitizeRedirectDest validates that a redirect destination is a safe same-origin
// relative path or absolute URL matching the request host. Rejects cross-origin
// redirects, javascript: URIs, backslash-based bypasses, and other open-redirect payloads.
// Returns a reconstructed URL from parsed components to prevent any raw-string injection.
func sanitizeRedirectDest(dest, requestHost string) string {
if dest == "" {
return "/"
}
// Parse through url.Parse to normalize and detect scheme/host
parsed, err := url.Parse(dest)
if err != nil {
return "/"
}
// Scheme whitelist — applies even for opaque URLs with empty Host.
// Without this, `javascript:alert(1)` produces an opaque URL with
// Host="" and Scheme="javascript", which would slip past the
// host-equality check below and end up reconstructed as
// `"javascript:"`. The only acceptable schemes are the empty
// string (pure-path relatives) and the two HTTP variants.
scheme := strings.ToLower(parsed.Scheme)
if scheme != "" && scheme != "http" && scheme != "https" {
return "/"
}
// Reject anything with a host component that doesn't match the request host.
// This catches protocol-relative (//evil.com), backslash tricks (/\evil.com
// which some browsers normalize to //evil.com), and explicit cross-origin URLs.
if parsed.Host != "" {
destHost := parsed.Hostname()
reqHost := requestHost
if h, _, err := net.SplitHostPort(requestHost); err == nil {
reqHost = h
}
if destHost != reqHost {
return "/"
}
}
// For relative paths: reject anything that doesn't start with a clean /
if parsed.Host == "" && parsed.Scheme == "" {
if !strings.HasPrefix(parsed.Path, "/") {
return "/"
}
// Reject backslash in path (browser normalization attack)
if strings.ContainsRune(parsed.Path, '\\') {
return "/"
}
}
// Reconstruct from parsed components to prevent raw-string injection
safe := &url.URL{
Scheme: parsed.Scheme,
Host: parsed.Host,
Path: parsed.Path,
RawQuery: parsed.RawQuery,
Fragment: parsed.Fragment,
}
return safe.String()
}
// verifyPoW checks that SHA256(nonce + solution) starts with `difficulty` zero nibbles.
func verifyPoW(nonce, solution string, difficulty int) bool {
h := sha256.Sum256([]byte(nonce + solution))
hexHash := hex.EncodeToString(h[:])
for i := 0; i < difficulty; i++ {
if i >= len(hexHash) || hexHash[i] != '0' {
return false
}
}
return true
}
package challenge
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/binary"
"errors"
"strings"
"time"
)
// AdminSessionSigner mints and verifies the signed cookies that let
// authenticated operators bypass the PoW. The signing key is generated
// on construction; rebuilding the signer (i.e., daemon restart)
// invalidates every previously-issued cookie -- the rotation contract.
type AdminSessionSigner struct {
key []byte
ttl time.Duration
}
// ErrSessionExpired is returned by Verify for cookies whose embedded
// expiry has passed.
var ErrSessionExpired = errors.New("session expired")
// ErrSessionBadSignature is returned by Verify for cookies whose HMAC
// does not match. Includes tampered payloads and cookies signed by a
// previous AdminSessionSigner instance (post-rotation).
var ErrSessionBadSignature = errors.New("session signature invalid")
// ErrSessionIPMismatch is returned when the cookie was issued for a
// different IP than the one presenting it. Stops cookie theft from a
// different network.
var ErrSessionIPMismatch = errors.New("session IP mismatch")
// ErrSessionMalformed wraps decoding errors so a corrupt cookie has a
// distinct sentinel from a tampered one.
var ErrSessionMalformed = errors.New("session payload malformed")
// NewAdminSessionSigner generates a fresh 32-byte signing key. The
// caller must keep the returned pointer for the lifetime of the
// challenge server; never construct a second signer for the same
// server, or already-issued cookies will be invalidated mid-session.
func NewAdminSessionSigner(ttl time.Duration) (*AdminSessionSigner, error) {
if ttl <= 0 {
ttl = 4 * time.Hour
}
key := make([]byte, 32)
if _, err := rand.Read(key); err != nil {
return nil, err
}
return &AdminSessionSigner{key: key, ttl: ttl}, nil
}
// TTL exposes the configured cookie lifetime so the server can set the
// matching Max-Age on the Set-Cookie header.
func (s *AdminSessionSigner) TTL() time.Duration { return s.ttl }
// Issue returns a cookie value of the form "<base64(payload)>.<base64 hmac>".
// The payload binds the cookie to a single IP and a single expiry so a
// stolen cookie does not work elsewhere or after the TTL.
func (s *AdminSessionSigner) Issue(ip string) string {
return s.issueAt(ip, time.Now().Add(s.ttl))
}
// issueAt is the test seam for issuing a cookie with an explicit
// expiry. Lets expiry tests construct already-expired cookies without
// reaching into encodeSessionPayload directly.
func (s *AdminSessionSigner) issueAt(ip string, exp time.Time) string {
payload := encodeSessionPayload(ip, exp)
mac := hmac.New(sha256.New, s.key)
mac.Write(payload)
sig := mac.Sum(nil)
return base64.RawURLEncoding.EncodeToString(payload) + "." + base64.RawURLEncoding.EncodeToString(sig)
}
// Verify checks the HMAC, payload format, expiry, and IP binding. Use
// errors.Is to branch on the failure mode.
func (s *AdminSessionSigner) Verify(cookieValue, ip string) error {
dot := strings.LastIndexByte(cookieValue, '.')
if dot <= 0 || dot == len(cookieValue)-1 {
return ErrSessionMalformed
}
payloadEnc := cookieValue[:dot]
sigEnc := cookieValue[dot+1:]
payload, err := base64.RawURLEncoding.DecodeString(payloadEnc)
if err != nil {
return ErrSessionMalformed
}
sig, err := base64.RawURLEncoding.DecodeString(sigEnc)
if err != nil {
return ErrSessionMalformed
}
mac := hmac.New(sha256.New, s.key)
mac.Write(payload)
expected := mac.Sum(nil)
if subtle.ConstantTimeCompare(sig, expected) != 1 {
return ErrSessionBadSignature
}
cookieIP, exp, err := decodeSessionPayload(payload)
if err != nil {
return ErrSessionMalformed
}
if time.Now().After(exp) {
return ErrSessionExpired
}
if cookieIP != ip {
return ErrSessionIPMismatch
}
return nil
}
// CompareAdminSecret returns true when stored and presented secrets
// match in constant time. Empty stored secret always returns false so a
// misconfigured admin_secret cannot accidentally accept any caller.
func CompareAdminSecret(stored, presented string) bool {
if stored == "" {
return false
}
return subtle.ConstantTimeCompare([]byte(stored), []byte(presented)) == 1
}
// encodeSessionPayload formats: 1 byte version | 8 byte unix expiry BE
// | n byte IP. Variable-length tail keeps it simple; the IP read stops
// at end-of-buffer.
func encodeSessionPayload(ip string, exp time.Time) []byte {
out := make([]byte, 0, 9+len(ip))
out = append(out, 1)
var ts [8]byte
// #nosec G115 -- exp is always future-dated (now + positive TTL via NewAdminSessionSigner / Issue), so Unix() is positive and the int64->uint64 cast cannot overflow.
binary.BigEndian.PutUint64(ts[:], uint64(exp.Unix()))
out = append(out, ts[:]...)
out = append(out, []byte(ip)...)
return out
}
func decodeSessionPayload(p []byte) (ip string, exp time.Time, err error) {
if len(p) < 9 || p[0] != 1 {
return "", time.Time{}, ErrSessionMalformed
}
tsRaw := binary.BigEndian.Uint64(p[1:9])
if tsRaw > uint64(1<<62) {
return "", time.Time{}, ErrSessionMalformed
}
return string(p[9:]), time.Unix(int64(tsRaw), 0), nil
}
package checks
import (
"errors"
"os"
"path/filepath"
"sort"
"strings"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/control"
)
// accountSuspended reports whether cPanel has suspended an account. cPanel
// locks the account's database users at suspension, so every query against its
// installs fails with an access error; scanning one produces no coverage, only
// a gap that no operator action can clear. A name that is not a single path
// element is rejected rather than resolved, so it cannot escape the directory.
func accountSuspended(account string) bool {
if account == "" || account == "." || account == ".." || strings.ContainsAny(account, `/\`) {
return false
}
_, err := osFS.Stat(filepath.Join("/var/cpanel/suspended", account))
return err == nil
}
// EnumerateScanAccounts returns the sorted list of cPanel account usernames
// eligible for a server-wide scan. Source of truth: the cPanel user registry
// (/var/cpanel/users) intersected with present home directories (/home/<user>).
// All filesystem access goes through the package osFS hook so it is fakeable.
//
// Fallback: when /var/cpanel/users is absent (non-cPanel platform), the
// function falls back to enumerating /home subdirectories whose names pass
// name validation. This makes the function usable on generic Linux hosts.
//
// A hard FS error reading the registry (not os.ErrNotExist) is propagated as
// an error. An empty registry returns ([]string{}, nil).
func EnumerateScanAccounts(_ *config.Config) ([]string, error) {
registryEntries, err := osFS.ReadDir("/var/cpanel/users")
if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, err
}
if errors.Is(err, os.ErrNotExist) {
// Non-cPanel platform: fall back to /home subdirectories.
return enumerateFromHome()
}
// Build a set of names present under the account roots for the
// intersection step.
homes, _ := listAccountHomes()
homeSet := make(map[string]struct{}, len(homes))
for _, h := range homes {
homeSet[h.Name()] = struct{}{}
}
seen := make(map[string]struct{}, len(registryEntries))
var accounts []string
for _, e := range registryEntries {
name := e.Name()
if !control.ValidScanAccountTarget(name) {
continue
}
if _, inHome := homeSet[name]; !inHome {
continue
}
if _, dup := seen[name]; dup {
continue
}
seen[name] = struct{}{}
accounts = append(accounts, name)
}
sort.Strings(accounts)
if accounts == nil {
accounts = []string{}
}
return accounts, nil
}
// enumerateFromHome lists /home subdirectory names that pass name validation.
// Used when the cPanel registry is absent. Because /home is the sole source of
// truth in this path, a hard read error (anything other than os.ErrNotExist) is
// propagated so callers can distinguish a broken enumeration from an empty host.
func enumerateFromHome() ([]string, error) {
homes, err := listAccountHomes()
if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, err
}
homeEntries := make([]os.DirEntry, 0, len(homes))
for _, h := range homes {
homeEntries = append(homeEntries, h.Entry)
}
seen := make(map[string]struct{}, len(homeEntries))
var accounts []string
for _, e := range homeEntries {
name := e.Name()
if !control.ValidScanAccountTarget(name) {
continue
}
if _, dup := seen[name]; dup {
continue
}
seen[name] = struct{}{}
accounts = append(accounts, name)
}
sort.Strings(accounts)
if accounts == nil {
accounts = []string{}
}
return accounts, nil
}
package checks
import (
"strings"
"sync"
"time"
)
// Domain->owner is read from cPanel's /etc/userdomains ("domain: user").
// Cached for ownerCacheTTL so a busy host does not re-read it per scan.
const ownerCacheTTL = 60 * time.Second
var (
ownerMu sync.Mutex
ownerMap map[string]string
ownerLoadedAt time.Time
)
// resetDomainOwnerCache clears the cache. Test-only seam.
func resetDomainOwnerCache() {
ownerMu.Lock()
ownerMap = nil
ownerLoadedAt = time.Time{}
ownerMu.Unlock()
}
// domainAccountOwner returns the cPanel account that owns domain, or "" when
// the map is unavailable (non-cPanel) or the domain is not present. The
// wildcard "*: nobody" line and any non "domain: user" line are ignored.
func domainAccountOwner(domain string) string {
domain = strings.ToLower(strings.TrimSpace(domain))
if domain == "" {
return ""
}
ownerMu.Lock()
defer ownerMu.Unlock()
if ownerMap == nil || time.Since(ownerLoadedAt) > ownerCacheTTL {
ownerMap = loadDomainOwners()
ownerLoadedAt = time.Now()
}
return ownerMap[domain]
}
func loadDomainOwners() map[string]string {
out := make(map[string]string)
data, err := osFS.ReadFile("/etc/userdomains")
if err != nil {
return out
}
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
idx := strings.LastIndexByte(line, ':')
if idx <= 0 {
continue
}
dom := strings.ToLower(strings.TrimSpace(line[:idx]))
owner := strings.TrimSpace(line[idx+1:])
if dom == "" || dom == "*" || owner == "" || owner == "nobody" {
continue
}
out[dom] = owner
}
return out
}
package checks
import (
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/platform"
)
// accountOwnerForDomain is the production mapping: cPanel's /etc/userdomains
// file, which misses on every other panel. Tests in other packages inject a
// table through SetAccountOwnerLookupForTest.
var accountOwnerForDomain = func(domain string) (string, bool) {
if !platform.Detect().IsCPanel() {
return "", false
}
owner := domainAccountOwner(domain)
return owner, owner != ""
}
// AccountOwnerForDomain resolves the hosting account that owns a mail
// domain, for producers that know a verified local mailbox or domain.
// Correlation never calls it; owners are resolved where the mailbox is known.
func AccountOwnerForDomain(domain string) (string, bool) {
return accountOwnerForDomain(domain)
}
// SetAccountOwnerLookupForTest replaces the domain-to-owner mapping and
// returns a function that restores it. Test-only seam for packages that
// cannot reach this package's filesystem fakes.
func SetAccountOwnerLookupForTest(fn func(domain string) (string, bool)) func() {
prev := accountOwnerForDomain
accountOwnerForDomain = fn
return func() { accountOwnerForDomain = prev }
}
// MailOwner returns the hosting account for a mailbox or bare domain, or ""
// when the mapping is unavailable. It never returns the mailbox or domain
// itself as an owner.
func MailOwner(mailboxOrDomain string) string {
domain := strings.TrimSpace(mailboxOrDomain)
if at := strings.LastIndexByte(domain, '@'); at >= 0 {
domain = domain[at+1:]
}
if domain == "" {
return ""
}
owner, _ := AccountOwnerForDomain(domain)
return owner
}
// hostingAccountForUser is the production mapping from a system user name
// to a hosting account. Tests in other packages inject a table through
// SetHostingAccountLookupForTest.
var hostingAccountForUser = func(name string) string {
if name == "" || name == "root" || name == "unknown" || strings.ContainsAny(name, "/\\") {
return ""
}
home := defaultUIDCache.HomeDir(name)
if !filepath.IsAbs(home) {
return ""
}
if root, account, ok := accountRootOf(filepath.Join(home, "probe")); ok && account == name && filepath.Clean(home) == filepath.Join(root, account) {
return name
}
return ""
}
// HostingAccountForUser returns name when it is a hosting account: its
// passwd home directory sits directly under a configured account root.
// Root, system and service users, unknown names and the lookup sentinel
// resolve to "". The passwd cache is the same one process findings use for
// uid resolution, so tests point both at one fixture file.
func HostingAccountForUser(name string) string {
return hostingAccountForUser(name)
}
// SetHostingAccountLookupForTest replaces the user-to-account mapping and
// returns a function that restores it. Test-only seam for packages that
// cannot reach this package's passwd and account-root fakes.
func SetHostingAccountLookupForTest(fn func(name string) string) func() {
prev := hostingAccountForUser
hostingAccountForUser = fn
return func() { hostingAccountForUser = prev }
}
// ftpAccountOwner resolves an FTP login name: a virtual account is
// user@domain and belongs to the domain's owner; anything else is a system
// account name that must itself be a hosting account.
func ftpAccountOwner(account string) string {
if strings.Contains(account, "@") {
return MailOwner(account)
}
return HostingAccountForUser(account)
}
// accountHomeExists reports whether a configured account root contains a
// directory named user and passwd identifies it as the user's home. A
// leftover directory must not promote a service user to a hosting owner.
func accountHomeExists(user string) bool {
if HostingAccountForUser(user) == "" {
return false
}
for _, root := range accountHomeRoots() {
if info, err := osFS.Stat(filepath.Join(root, user)); err == nil && info.IsDir() {
return true
}
}
return false
}
// installOwner resolves the hosting account that owns a CMS install from
// its configuration path. ok is false outside every account root; callers
// keep their display label but must not stamp a tenant.
func installOwner(configPath string) (string, bool) {
_, account, ok := accountRootOf(configPath)
return account, ok
}
package checks
import (
"os"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
)
// accountHomeRoots answers "where do hosting accounts live" for every
// account-scoped check, remediation and re-check. A seam so tests can point
// the package at a temp tree; production follows platform detection.
var accountHomeRoots = func() []string { return platform.Detect().AccountHomeRoots() }
// accountHome is one account directory under one of the roots.
type accountHome struct {
Root string
Entry os.DirEntry
}
// Path is the account's home directory.
func (h accountHome) Path() string { return filepath.Join(h.Root, h.Entry.Name()) }
// Name is the account name.
func (h accountHome) Name() string { return h.Entry.Name() }
// listAccountHomes provides a best-effort account inventory. Stateful scanners
// use readAccountHomes so a partial inventory cannot retire unseen findings.
func listAccountHomes() ([]accountHome, error) {
homes, err := readAccountHomes()
if len(homes) > 0 {
return homes, nil
}
return homes, err
}
// readAccountHomes skips absent roots, but reports other read failures even
// when another root supplies accounts.
func readAccountHomes() ([]accountHome, error) {
var homes []accountHome
var firstErr error
for _, root := range accountHomeRoots() {
entries, err := osFS.ReadDir(root)
if err != nil {
if !os.IsNotExist(err) && firstErr == nil {
firstErr = err
}
continue
}
for _, e := range entries {
homes = append(homes, accountHome{Root: root, Entry: e})
}
}
return homes, firstErr
}
// AccountHomeRoots returns the directories hosting accounts live under on
// this platform: /home, or /var/www/vhosts on Plesk.
func AccountHomeRoots() []string {
return accountHomeRoots()
}
// accountHomeDir resolves an account's home directory: the first root that
// holds it, or the first root when it exists nowhere (callers that need
// existence check it themselves).
func accountHomeDir(account string) string {
return AccountHomeDirIn(accountHomeRoots(), account)
}
// AccountHomeDirIn is accountHomeDir over the given roots.
func AccountHomeDirIn(roots []string, account string) string {
for _, root := range roots {
candidate := filepath.Join(root, account)
if _, err := osFS.Stat(candidate); err == nil {
return candidate
}
}
if len(roots) == 0 {
return filepath.Join("/home", account)
}
return filepath.Join(roots[0], account)
}
// accountHomeGlob globs "<root>/<pattern>" under every root and returns the
// concatenated matches. pattern typically starts with "*" (any account).
func accountHomeGlob(pattern string) ([]string, error) {
var out []string
var firstErr error
for _, root := range accountHomeRoots() {
matches, err := osFS.Glob(filepath.Join(root, pattern))
if err != nil && firstErr == nil {
firstErr = err
}
out = append(out, matches...)
}
if out == nil {
return nil, firstErr
}
return out, nil
}
// accountHomePatterns returns "<root>/*" for every account root: the glob
// form of "every account home".
func accountHomePatterns() []string {
roots := accountHomeRoots()
out := make([]string, 0, len(roots))
for _, root := range roots {
out = append(out, filepath.Join(root, "*"))
}
return out
}
// accountHomeSubPatterns returns "<root>/*/<sub>" for every account root.
func accountHomeSubPatterns(sub string) []string {
roots := accountHomeRoots()
out := make([]string, 0, len(roots))
for _, root := range roots {
out = append(out, filepath.Join(root, "*", sub))
}
return out
}
// accountRootOf reports the root and account a path belongs to. The path
// must lie strictly inside an account directory: a root or an account home
// itself is not "inside an account".
func accountRootOf(path string) (root, account string, ok bool) {
return accountRootOfAt(path, accountHomeRoots())
}
func accountRootOfAt(path string, roots []string) (root, account string, ok bool) {
clean := filepath.Clean(path)
for _, r := range roots {
r = filepath.Clean(r)
rest, found := strings.CutPrefix(clean, r+string(filepath.Separator))
if !found {
continue
}
account, tail, hasTail := strings.Cut(rest, string(filepath.Separator))
if account == "" || !hasTail || tail == "" {
continue
}
return r, account, true
}
return "", "", false
}
// isAccountRoot reports whether dir is one of the account roots.
func isAccountRoot(dir string) bool {
clean := filepath.Clean(dir)
for _, r := range accountHomeRoots() {
if filepath.Clean(r) == clean {
return true
}
}
return false
}
// underAccountRoot reports whether path is at or below any account root.
func underAccountRoot(path string) bool {
clean := filepath.Clean(path)
for _, r := range accountHomeRoots() {
r = filepath.Clean(r)
if clean == r || strings.HasPrefix(clean, r+string(filepath.Separator)) {
return true
}
}
return false
}
// accountRootPrefixes returns every account root with a trailing separator
// (the prefix form used to test "is this path inside an account") followed
// by extra prefixes verbatim.
func accountRootPrefixes(extra ...string) []string {
roots := accountHomeRoots()
out := make([]string, 0, len(roots)+len(extra))
for _, root := range roots {
out = append(out, filepath.Clean(root)+string(filepath.Separator))
}
return append(out, extra...)
}
// accountNameInTextAt returns the account named by the first
// "<root>/<account>/" reference in free text (a finding message or
// details), or "" when none is present.
func accountNameInTextAt(text string, roots []string) string {
for _, root := range roots {
prefix := filepath.Clean(root) + string(filepath.Separator)
idx := strings.Index(text, prefix)
if idx < 0 {
continue
}
rest := text[idx+len(prefix):]
if slash := strings.IndexByte(rest, '/'); slash > 0 {
return rest[:slash]
}
}
return ""
}
// effectiveFixRoots returns the allowed roots for a remediation: an
// explicit override (tests redirect under t.TempDir()) wins; otherwise the
// account roots plus any extra system directories the action may touch.
func effectiveFixRoots(override []string, extra ...string) []string {
if override != nil {
return override
}
roots := append([]string(nil), accountHomeRoots()...)
if cfg := config.Active(); cfg != nil {
// Failed roots are excluded; doctor reports the resolution errors.
// One tenant's symlink must not disable fixes for other accounts.
configured, _ := platform.ResolveAccountRoots(cfg.AccountRoots)
roots = append(roots, configured...)
}
return append(roots, extra...)
}
// quarantineExtraRoots are the non-account directories a quarantine may
// take a file from: the world-writable temp trees droppers land in.
var quarantineExtraRoots = []string{"/tmp", "/dev/shm", "/var/tmp"}
package checks
import (
"context"
"fmt"
"os"
"os/user"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
const accountScanMaxFilesDefault = 10000
func effectiveAccountScanMaxFiles(cfg *config.Config) int {
if cfg == nil || cfg.Thresholds.AccountScanMaxFiles <= 0 {
return accountScanMaxFilesDefault
}
return cfg.Thresholds.AccountScanMaxFiles
}
// rankPathsByMtimeDesc orders paths most-recent-first and optionally caps
// the result at maxFiles. Lexical glob order plus a downstream check timeout
// would otherwise keep cutting iteration off at the same prefix every cycle,
// hiding indicators on late-alphabet accounts.
//
// Stat failures are tolerated: the path is kept with a zero mtime so it
// sorts to the end, letting the cap chop it first when present.
// Best-effort ranking is the goal -- dropping silently here would
// reintroduce the same hidden-input bug class the helper exists to close.
// Downstream readers handle the missing-file case on their own.
//
// A canceled ctx returns nil; nil ctx is treated as Background.
// maxFiles <= 0 disables the cap; the sort still runs.
func rankPathsByMtimeDesc(ctx context.Context, paths []string, maxFiles int) []string {
if ctx == nil {
ctx = context.Background()
}
if err := ctx.Err(); err != nil {
return nil
}
type entry struct {
path string
mtime time.Time
}
ranked := make([]entry, 0, len(paths))
for _, p := range paths {
if err := ctx.Err(); err != nil {
return nil
}
var mt time.Time
if info, err := osFS.Stat(p); err == nil {
mt = info.ModTime()
}
ranked = append(ranked, entry{path: p, mtime: mt})
}
if err := ctx.Err(); err != nil {
return nil
}
sort.Slice(ranked, func(i, j int) bool {
if ranked[i].mtime.Equal(ranked[j].mtime) {
return ranked[i].path < ranked[j].path
}
return ranked[i].mtime.After(ranked[j].mtime)
})
if err := ctx.Err(); err != nil {
return nil
}
var droppedPaths []string
if maxFiles > 0 && len(ranked) > maxFiles {
dropped := len(ranked) - maxFiles
droppedPaths = make([]string, dropped)
for i, e := range ranked[maxFiles:] {
droppedPaths[i] = e.path
}
ranked = ranked[:maxFiles]
}
if len(droppedPaths) > 0 {
recordAccountScanTruncatedPaths(ctx, droppedPaths, maxFiles)
}
out := make([]string, len(ranked))
for i, e := range ranked {
out[i] = e.path
}
return out
}
// RunAccountScan runs all applicable checks scoped to a single cPanel account.
// Returns findings for that account only. Does NOT trigger auto-response actions.
//
// This is a thin wrapper around RunAccountScanWithOptions using
// DefaultAccountScanOptions so all existing callers retain their current behaviour.
func RunAccountScan(cfg *config.Config, store *state.Store, account string) []alert.Finding {
return RunAccountScanWithOptions(context.Background(), cfg, store, account, DefaultAccountScanOptions(cfg))
}
// RunAccountScanWithOptions is the options-aware entry point for per-account scans.
// Scope is propagated through ctx via ContextWithAccountScope, so parallel
// scans of different accounts no longer block on a single process-wide
// mutex and never bleed scope into each other.
func RunAccountScanWithOptions(ctx context.Context, cfg *config.Config, store *state.Store, account string, opts AccountScanOptions) []alert.Finding {
// Verify account exists
homeDir := accountHomeDir(account)
if _, err := osFS.Stat(homeDir); os.IsNotExist(err) {
return []alert.Finding{{
Severity: alert.Warning,
Check: "account_scan",
Message: fmt.Sprintf("Account '%s' not found (no %s directory)", account, homeDir),
Timestamp: time.Now(),
}}
}
// Account-scoped checks (filesystem + account-specific)
accountChecks := []namedCheck{
{"webshells", CheckWebshells},
{"htaccess", CheckHtaccess},
{"wp_core", CheckWPCore},
{"php_content", CheckPHPContent},
{"phishing", CheckPhishing},
{"filesystem", CheckFilesystem},
{"group_writable_php", CheckGroupWritablePHP},
{"nulled_plugins", CheckNulledPlugins},
{"open_basedir", CheckOpenBasedir},
{"symlink_attacks", CheckSymlinkAttacks},
{"db_content", CheckDatabaseContent},
{"php_config_changes", CheckPHPConfigChanges},
}
// Account-specific checks that need the account name
accountChecks = append(accountChecks,
namedCheck{"ssh_keys_account", makeAccountSSHKeyCheck(account)},
namedCheck{"crontab_account", makeAccountCrontabCheck(account)},
namedCheck{"backdoor_binaries", makeAccountBackdoorCheck(account)},
)
// File-index audit: only included for full scans (ForceFileIndex=true).
// In default scans CheckFileIndex writes live state (fileindex.current,
// fileindex.previous, dircache.json) that the host-wide incremental
// baseline relies on. Adding it unconditionally would corrupt that state
// once per triggered account scan. The audit branch (ForceFileIndex=true)
// is read-only and account-scoped so it is safe to include here.
if opts.ForceFileIndex {
accountChecks = append(accountChecks, namedCheck{"file_index", CheckFileIndex})
}
// Run under the host-wide scan budget - filesystem checks all walk the
// same directory tree, so too many concurrent checks starve each other on
// loaded servers with slow I/O, and a periodic tier may be running too.
scanCtx, truncations := withAccountScanTruncationCollector(ctx)
scanCtx = ContextWithAccountScope(scanCtx, account)
scanCtx = ContextWithScanOptions(scanCtx, opts)
scanCtx = withWPInstallCache(scanCtx)
findings := runAccountChecksBounded(scanCtx, cfg, store, accountChecks)
now := time.Now()
findings = append(findings, truncations.findings(now)...)
for i := range findings {
if findings[i].Timestamp.IsZero() {
findings[i].Timestamp = now
}
}
// Filter findings to only include this account's paths
var filtered []alert.Finding
for _, f := range findings {
if accountScanFindingInScope(f, account) {
filtered = append(filtered, f)
}
}
return stampTenantIDIfEmpty(filtered, account)
}
// runAccountChecksBounded runs checks under the host-wide scan budget.
// A check still waiting for a slot when ctx is cancelled never starts: the
// slot wait used to ignore the context, so an operator's cancel left every
// queued check running to its (immediate) end and reporting a timeout.
func runAccountChecksBounded(ctx context.Context, cfg *config.Config, store *state.Store, checks []namedCheck) []alert.Finding {
scansInFlight.Add(1)
defer scansInFlight.Add(-1)
var mu sync.Mutex
var findings []alert.Finding
var wg sync.WaitGroup
budget := scanBudgetFrom(ctx)
dispatches := checkDispatches.begin(len(checks), budget)
checkDispatches.observe(ctx, dispatches)
for i, nc := range checks {
wg.Add(1)
c := nc
task := dispatches[i]
// Account checks run against user filesystem content (unparsed PHP,
// crafted archives, foreign encodings) so a panic is plausible.
// runAccountScanCheck surfaces it as check_panic, keeping the scan and
// daemon alive.
obs.SafeGo("account-scan-runner", task.wrap(func() {
defer wg.Done()
if !task.admit(ctx) {
task.withdraw(ctx)
return
}
if ctx.Err() != nil {
task.withdraw(ctx)
return
}
results := runAccountScanCheck(withCheckDispatch(ctx, task), c, cfg, store, timeoutFor(c.name))
if len(results) > 0 {
mu.Lock()
findings = append(findings, results...)
mu.Unlock()
}
}))
}
wg.Wait()
return findings
}
// runAccountScanCheck runs one check under a timeout, recovering any panic.
// A check cut short because the scan itself was cancelled reports nothing:
// that is not a timeout, and the warning would be persisted with the partial
// results the cancel keeps.
func runAccountScanCheck(ctx context.Context, c namedCheck, cfg *config.Config, store *state.Store, timeout time.Duration) []alert.Finding {
if ctx.Err() != nil {
checkDispatchFrom(ctx).withdraw(ctx)
return nil
}
cctx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
panicIdentity := alert.Finding{Check: "check_panic", DedupKey: fmt.Sprintf("account-scan:%q:%s", AccountFromContext(ctx), c.name)}
execution := executeCheckAsync(cctx, "account-scan-exec", func() []alert.Finding {
return c.fn(cctx, cfg, store)
})
defer execution.finishCaller()
select {
case outcome := <-execution.done:
execution.received()
if outcome.panicErr != "" {
// Keep stack churn out of identity without merging failures from
// different accounts or the scheduled runner.
return []alert.Finding{{
Severity: alert.High,
Check: "check_panic",
TenantID: AccountFromContext(ctx),
Message: fmt.Sprintf("Account scan check '%s' stopped after an internal panic", c.name),
Details: outcome.panicErr,
DedupKey: panicIdentity.DedupKey,
Timestamp: time.Now(),
}}
}
if store != nil && cctx.Err() == nil {
store.RearmFindings([]string{panicIdentity.Key()})
}
return outcome.findings
case <-cctx.Done():
execution.withdraw(cctx.Err())
if ctx.Err() != nil {
return nil
}
return []alert.Finding{{
Severity: alert.Warning,
Check: "check_timeout",
Message: fmt.Sprintf("Account scan check '%s' timed out", c.name),
Timestamp: time.Now(),
}}
}
}
func accountScanFindingInScope(f alert.Finding, account string) bool {
if account == "" {
return true
}
if f.FilePath != "" {
fileAccount := accountFromHomePath(f.FilePath)
if fileAccount != "" {
return fileAccount == account
}
if containsHomeReference(f.FilePath) {
return false
}
}
hasHomeRef, hasAccountRef := textHomeScope(f.Message, account)
detailsHasHomeRef, detailsHasAccountRef := textHomeScope(f.Details, account)
if hasAccountRef || detailsHasAccountRef {
return true
}
if hasHomeRef || detailsHasHomeRef {
return false
}
return true
}
// accountRootPrefixLen reports whether text starts with an account root
// and how long that prefix is. "/home" keeps cPanel's multi-home tolerance
// (/home2, /home3) so findings from those trees still resolve.
func accountRootPrefixLen(text string) (int, bool) {
for _, root := range accountHomeRoots() {
root = filepath.ToSlash(filepath.Clean(root))
if !strings.HasPrefix(text, root) {
continue
}
i := len(root)
if root == "/home" {
for i < len(text) && text[i] >= '0' && text[i] <= '9' {
i++
}
}
if i == len(text) || text[i] == '/' {
return i, true
}
}
return 0, false
}
func containsHomeReference(path string) bool {
_, ok := accountRootPrefixLen(filepath.ToSlash(filepath.Clean(path)))
return ok
}
func textHomeScope(text, account string) (hasHomeRef, hasAccountRef bool) {
for _, root := range accountHomeRoots() {
root = filepath.ToSlash(filepath.Clean(root))
for i := 0; i < len(text); {
idx := strings.Index(text[i:], root)
if idx < 0 {
break
}
start := i + idx
if homeAccount, ok := homeAccountAt(text[start:]); ok {
hasHomeRef = true
if homeAccount == account {
hasAccountRef = true
}
}
i = start + len(root)
}
}
return hasHomeRef, hasAccountRef
}
func homeAccountAt(text string) (string, bool) {
i, ok := accountRootPrefixLen(text)
if !ok {
return "", false
}
if i == len(text) {
return "", true
}
i++
start := i
for i < len(text) && isHomeAccountByte(text[i]) {
i++
}
if i == start {
return "", true
}
return text[start:i], true
}
func isHomeAccountByte(b byte) bool {
return b >= 'a' && b <= 'z' ||
b >= 'A' && b <= 'Z' ||
b >= '0' && b <= '9' ||
b == '_' || b == '-' || b == '.'
}
// stampTenantIDIfEmpty fills in Finding.TenantID with account when the
// detector emitted the finding without explicit tenant attribution.
// Account-scope detectors otherwise leave TenantID empty and the
// correlator falls back to weaker identities (UID, PID, file hash),
// fragmenting one account's incidents across multiple keys. Findings
// the detector did stamp keep their value; an empty account is a no-op.
func stampTenantIDIfEmpty(findings []alert.Finding, account string) []alert.Finding {
if account == "" {
return findings
}
for i := range findings {
if findings[i].TenantID == "" {
findings[i].TenantID = account
}
}
return findings
}
// GetScanHomeDirs returns the list of home directories to scan.
// When ctx carries an account scope (via ContextWithAccountScope), only
// that account is returned. Otherwise every entry under every account root
// is read. Nil ctx is tolerated for legacy callers and treated as host-wide.
// Callers that need the directory path use scanHomeDirPath on each entry.
func GetScanHomeDirs(ctx context.Context) ([]os.DirEntry, error) {
if account := AccountFromContext(ctx); account != "" {
info, err := osFS.Stat(accountHomeDir(account))
if err != nil {
return nil, err
}
return []os.DirEntry{fakeDirEntry{info}}, nil
}
homes, err := listAccountHomes()
if err != nil {
return nil, err
}
entries := make([]os.DirEntry, 0, len(homes))
for _, h := range homes {
entries = append(entries, rootedDirEntry{DirEntry: h.Entry, root: h.Root})
}
return entries, nil
}
// rootedDirEntry remembers which account root an entry came from.
type rootedDirEntry struct {
os.DirEntry
root string
}
// scanHomeDirPath returns the home directory for an entry returned by
// GetScanHomeDirs (or any account enumeration): the entry's own root when it
// carries one, otherwise the root that holds the account.
func scanHomeDirPath(entry os.DirEntry) string {
if r, ok := entry.(rootedDirEntry); ok {
return filepath.Join(r.root, r.Name())
}
return accountHomeDir(entry.Name())
}
// WebRootPatterns returns the configured web-root globs, including the
// platform default when the operator did not set account_roots.
func WebRootPatterns(cfg *config.Config) []string {
switch {
case cfg != nil && len(cfg.AccountRoots) > 0:
return append([]string(nil), cfg.AccountRoots...)
case platform.Detect().IsCPanel():
return accountHomeSubPatterns("public_html")
default:
return nil
}
}
// ResolveWebRoots returns the list of directory paths CSM should scan for
// web-facing content (wp-config.php, .htaccess, public_html trees, etc.).
//
// Resolution order:
// 1. If cfg.AccountRoots is set, expand each glob and return the result.
// Explicit config always wins.
// 2. On cPanel hosts (detected via platform.Detect), fall back to
// /home/*/public_html for backward compatibility.
// 3. On non-cPanel hosts with no config, return an empty list. Callers
// should treat this as "no scanning" and skip cleanly.
//
// Each returned path is an absolute directory that exists on disk.
func ResolveWebRoots(cfg *config.Config) []string {
patterns := WebRootPatterns(cfg)
var roots []string
seen := make(map[string]struct{})
for _, pattern := range patterns {
matches, err := osFS.Glob(pattern)
if err != nil || len(matches) == 0 {
continue
}
for _, m := range matches {
info, err := osFS.Stat(m)
if err != nil || !info.IsDir() {
continue
}
if _, ok := seen[m]; ok {
continue
}
seen[m] = struct{}{}
roots = append(roots, m)
}
}
return roots
}
// fakeDirEntry wraps os.FileInfo to implement os.DirEntry.
type fakeDirEntry struct {
fi os.FileInfo
}
func (f fakeDirEntry) Name() string { return f.fi.Name() }
func (f fakeDirEntry) IsDir() bool { return f.fi.IsDir() }
func (f fakeDirEntry) Type() os.FileMode { return f.fi.Mode().Type() }
func (f fakeDirEntry) Info() (os.FileInfo, error) { return f.fi, nil }
// makeAccountSSHKeyCheck creates a check for SSH keys of a specific account.
func makeAccountSSHKeyCheck(account string) CheckFunc {
return func(_ context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
keyFile := filepath.Join(accountHomeDir(account), ".ssh", "authorized_keys")
hash, err := hashFileContent(keyFile)
if err != nil {
return nil
}
key := fmt.Sprintf("_ssh_user_keys:%s", keyFile)
prev, exists := store.GetRaw(key)
if exists && prev != hash {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "ssh_keys",
Message: fmt.Sprintf("User authorized_keys modified: %s", keyFile),
})
}
return findings
}
}
// makeAccountCrontabCheck creates a check for a specific account's crontab.
func makeAccountCrontabCheck(account string) CheckFunc {
return func(_ context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
crontabFile := filepath.Join(cronSpoolDir(), account)
data, err := osFS.ReadFile(crontabFile)
if err != nil {
return nil
}
content := string(data)
for _, pattern := range MatchCrontabPatternsDeep(content, cfg) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "suspicious_crontab",
Message: fmt.Sprintf("Suspicious pattern in crontab for %s: %s", account, pattern),
Details: fmt.Sprintf("File: %s\nContent:\n%s", crontabFile, content),
FilePath: crontabFile,
})
}
return findings
}
}
// makeAccountBackdoorCheck creates a check for backdoor binaries in account's .config.
func makeAccountBackdoorCheck(account string) CheckFunc {
return func(_ context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
backdoorNames := map[string]bool{
"defunct": true, "defunct.dat": true, "gs-netcat": true,
"gs-sftp": true, "gs-mount": true, "gsocket": true,
}
patterns := []string{
filepath.Join(accountHomeDir(account), ".config", "htop", "*"),
filepath.Join(accountHomeDir(account), ".config", "*", "*"),
}
for _, pattern := range patterns {
matches, _ := osFS.Glob(pattern)
for _, path := range matches {
if backdoorNames[filepath.Base(path)] {
info, _ := osFS.Stat(path)
var details string
if info != nil {
details = fmt.Sprintf("Size: %d bytes, Mtime: %s", info.Size(), info.ModTime().Format("2006-01-02 15:04:05"))
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "backdoor_binary",
Message: fmt.Sprintf("Backdoor binary found: %s", path),
Details: details,
FilePath: path,
})
}
}
}
return findings
}
}
// LookupUID returns the UID for a system account name, or -1 if not found.
func LookupUID(account string) int {
u, err := user.Lookup(account)
if err != nil {
return -1
}
uid := 0
fmt.Sscanf(u.Uid, "%d", &uid)
return uid
}
// AccountHomePatterns returns the glob for every account home ("<root>/*") on
// this platform. The realtime scanner needs it to recognise an account tree
// without hardcoding /home.
func AccountHomePatterns() []string {
return accountHomePatterns()
}
package checks
import (
"context"
"fmt"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
)
type accountScanTruncationContextKey struct{}
type accountScanTruncationCollector struct {
mu sync.Mutex
// droppedBy is keyed by account so operators can see which tenant hit
// which cap. The empty-string key keeps dropped paths that are not
// attributable to /home/<account>/.
droppedBy map[string]map[int]int
}
func withAccountScanTruncationCollector(ctx context.Context) (context.Context, *accountScanTruncationCollector) {
if ctx == nil {
ctx = context.Background()
}
collector := &accountScanTruncationCollector{droppedBy: map[string]map[int]int{}}
return context.WithValue(ctx, accountScanTruncationContextKey{}, collector), collector
}
func recordAccountScanTruncated(ctx context.Context, dropped, cap int) {
if dropped <= 0 || cap <= 0 || ctx == nil {
return
}
collector, ok := ctx.Value(accountScanTruncationContextKey{}).(*accountScanTruncationCollector)
if !ok || collector == nil {
return
}
collector.record(AccountFromContext(ctx), dropped, cap)
}
func recordAccountScanTruncatedPaths(ctx context.Context, droppedPaths []string, cap int) {
if len(droppedPaths) == 0 || cap <= 0 || ctx == nil {
return
}
collector, ok := ctx.Value(accountScanTruncationContextKey{}).(*accountScanTruncationCollector)
if !ok || collector == nil {
return
}
for account, dropped := range accountScanTruncationAccounts(ctx, droppedPaths) {
collector.record(account, dropped, cap)
}
}
func accountScanTruncationAccounts(ctx context.Context, paths []string) map[string]int {
if account := AccountFromContext(ctx); account != "" {
return map[string]int{account: len(paths)}
}
counts := make(map[string]int)
for _, path := range paths {
counts[accountFromHomePath(path)]++
}
return counts
}
func accountFromHomePath(path string) string {
cleaned := filepath.ToSlash(filepath.Clean(path))
i, ok := accountRootPrefixLen(cleaned)
if !ok || i == len(cleaned) || cleaned[i] != '/' {
return ""
}
account := cleaned[i+1:]
if slash := strings.IndexByte(account, '/'); slash >= 0 {
account = account[:slash]
}
if account == "." || account == ".." {
return ""
}
return account
}
func (c *accountScanTruncationCollector) record(account string, dropped, cap int) {
c.mu.Lock()
defer c.mu.Unlock()
caps, ok := c.droppedBy[account]
if !ok {
caps = map[int]int{}
c.droppedBy[account] = caps
}
caps[cap] += dropped
}
func (c *accountScanTruncationCollector) findings(now time.Time) []alert.Finding {
c.mu.Lock()
defer c.mu.Unlock()
if len(c.droppedBy) == 0 {
return nil
}
accounts := make([]string, 0, len(c.droppedBy))
for a := range c.droppedBy {
accounts = append(accounts, a)
}
sort.Strings(accounts) // stable finding order across runs
var findings []alert.Finding
for _, account := range accounts {
caps := c.droppedBy[account]
capValues := make([]int, 0, len(caps))
for cap := range caps {
capValues = append(capValues, cap)
}
sort.Ints(capValues)
for _, cap := range capValues {
dropped := caps[cap]
scope := "host scan"
tenantID := ""
if account != "" {
scope = fmt.Sprintf("account %s", account)
tenantID = account
}
// The skipped count moves with every file added or removed; the
// condition is this scope hitting this cap.
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "account_scan_truncated",
TenantID: tenantID,
Message: fmt.Sprintf("Account scan truncated for %s: %d file(s) skipped past cap of %d", scope, dropped, cap),
Details: "Raise thresholds.account_scan_max_files if recent detection coverage matters more than scan duration.",
DedupKey: fmt.Sprintf("scope=%q cap=%d", account, cap),
Timestamp: now,
})
}
}
return findings
}
package checks
import (
"context"
"path/filepath"
)
type accountScopeKey struct{}
// ContextWithAccountScope returns a derived context that restricts
// filesystem-based checks to a single cPanel/Linux account. Callers
// pass the resulting context into every check; helpers like
// GetScanHomeDirs read the scope back out. Empty account is a no-op
// (returns ctx unchanged, equivalent to a full host scan).
func ContextWithAccountScope(ctx context.Context, account string) context.Context {
if ctx == nil {
ctx = context.Background()
}
if account == "" {
return ctx
}
return context.WithValue(ctx, accountScopeKey{}, account)
}
// AccountFromContext returns the account scope previously attached by
// ContextWithAccountScope, or "" when no scope is set.
func AccountFromContext(ctx context.Context) string {
if ctx == nil {
return ""
}
v, _ := ctx.Value(accountScopeKey{}).(string)
return v
}
// homeGlob globs elem under every account (or the scoped account) across
// every account root.
func homeGlob(ctx context.Context, elem ...string) ([]string, error) {
account := AccountFromContext(ctx)
if account == "" {
account = "*"
}
parts := append([]string{account}, elem...)
return accountHomeGlob(filepath.Join(parts...))
}
package checks
import (
"context"
"fmt"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// afAlgEvent is the parsed view of a single SYSCALL record tagged with the
// csm_af_alg_socket auditd key. Fields are kept as strings because the
// downstream consumer is a human-readable Finding message.
type afAlgEvent struct {
Timestamp string // e.g. "1761826283.452"
Serial string // e.g. "91234"
UID string
AUID string
PID string // process id from the SYSCALL record; needed for live kill reaction
Comm string
Exe string
}
// AFAlgEvent is the package-public view of afAlgEvent used by callers
// outside internal/checks (the daemon's live audit-log listener emits
// findings derived from this shape).
type AFAlgEvent = afAlgEvent
// ParseAFAlgEventLine is the exported alias of parseAFAlgEvent for the
// daemon's live listener. The unexported form stays internal so the
// rest of this package can refer to the type by its short name.
func ParseAFAlgEventLine(line string) (AFAlgEvent, bool) {
return parseAFAlgEvent(line)
}
// AFAlgOwner resolves the hosting account whose process opened the socket.
// The audit uid is resolved through the shared passwd cache; root, service
// users and unknown uids yield "" so the finding stays unattributed.
func AFAlgOwner(ev AFAlgEvent) string {
uid, err := strconv.ParseUint(ev.UID, 10, 32)
if err != nil {
return ""
}
return HostingAccountForUser(LookupUser(uint32(uid)))
}
// after reports whether e is strictly newer than other. Comparison is
// (Timestamp, Serial) lexicographic with numeric semantics.
func (e afAlgEvent) after(other afAlgEvent) bool {
if e.Timestamp != other.Timestamp {
eFloat, _ := strconv.ParseFloat(e.Timestamp, 64)
otherFloat, _ := strconv.ParseFloat(other.Timestamp, 64)
return eFloat > otherFloat
}
// Avoid the local name `os` here — it shadows the stdlib package
// of the same name and would silently break a future edit that
// adds an `os` import to this file.
eSerial, _ := strconv.Atoi(e.Serial)
otherSerial, _ := strconv.Atoi(other.Serial)
return eSerial > otherSerial
}
// parseAFAlgEvent extracts the relevant fields from a single audit log line.
// It returns (event, true) only when the line is a SYSCALL record carrying
// the csm_af_alg_socket key. Anything else returns ok=false.
func parseAFAlgEvent(line string) (afAlgEvent, bool) {
// Require type=SYSCALL: auditd also writes the rule's key into
// CONFIG_CHANGE records on every add_rule / remove_rule (i.e. every
// CSM restart), and those are not exploit signatures.
if !strings.HasPrefix(line, "type=SYSCALL ") {
return afAlgEvent{}, false
}
if !strings.Contains(line, `key="csm_af_alg_socket"`) {
return afAlgEvent{}, false
}
ts, serial, ok := parseAuditMsgID(line)
if !ok {
return afAlgEvent{}, false
}
ev := afAlgEvent{Timestamp: ts, Serial: serial}
ev.UID = extractAuditField(line, "uid")
ev.AUID = extractAuditField(line, "auid")
ev.PID = extractAuditField(line, "pid")
ev.Comm = extractAuditField(line, "comm")
ev.Exe = extractAuditField(line, "exe")
return ev, true
}
// parseAuditMsgID extracts (timestamp, serial) from `msg=audit(TS:SERIAL):`.
func parseAuditMsgID(line string) (string, string, bool) {
const marker = "msg=audit("
i := strings.Index(line, marker)
if i < 0 {
return "", "", false
}
rest := line[i+len(marker):]
end := strings.Index(rest, ")")
if end < 0 {
return "", "", false
}
inside := rest[:end]
colon := strings.Index(inside, ":")
if colon < 0 {
return "", "", false
}
ts := inside[:colon]
serial := inside[colon+1:]
if _, err := strconv.ParseFloat(ts, 64); err != nil {
return "", "", false
}
if _, err := strconv.Atoi(serial); err != nil {
return "", "", false
}
return ts, serial, true
}
// extractAuditField returns the value of `key=...` from an audit log line.
// Quoted values may contain spaces; bare values are whitespace-delimited.
func extractAuditField(line, key string) string {
prefix := key + "="
idx := 0
for {
i := strings.Index(line[idx:], prefix)
if i < 0 {
return ""
}
i += idx
// Require start-of-line or preceding whitespace so "auid=" doesn't
// match when we asked for "uid=".
if i > 0 && line[i-1] != ' ' && line[i-1] != '\t' {
idx = i + 1
continue
}
rest := line[i+len(prefix):]
if strings.HasPrefix(rest, `"`) {
rest = rest[1:]
end := strings.Index(rest, `"`)
if end < 0 {
return ""
}
return rest[:end]
}
end := strings.IndexAny(rest, " \t")
if end < 0 {
return rest
}
return rest[:end]
}
}
const (
afAlgLogPath = "/var/log/audit/audit.log"
afAlgCursorKey = "_af_alg_last_seen"
)
// CheckAFAlgSocketUsage scans the audit log for csm_af_alg_socket events
// and emits one Critical finding per strictly-newer event. The first run
// alerts on every event found — AF_ALG-from-userland is an exploit signature
// for CVE-2026-31431 ("Copy Fail"), not a baseline metric, so silent seeding
// would hide pre-existing compromise. The cursor in state.Store prevents
// duplicates on subsequent sweeps and survives daemon restarts.
//
// Filtering is delegated to grep so we don't load the whole multi-hundred-MB
// audit log into memory each tick (same precedent as getAuditShadowInfo in
// auth.go). RunAllowNonZero is required because grep returns exit 1 on
// "no match" — the healthy default — and that must not surface as an error.
func CheckAFAlgSocketUsage(_ context.Context, _ *config.Config, st *state.Store) []alert.Finding {
// RunAllowNonZero swallows every non-zero exit (see runCmdAllowNonZeroReal
// in helpers.go) — including exit 1 (no match) and exit 2 (audit log not
// installed). The non-error path is therefore the only one we need to
// reason about here: empty output means "nothing to do".
out, err := cmdExec.RunAllowNonZero("grep", "-a", "csm_af_alg_socket", afAlgLogPath)
if err != nil {
return nil
}
if len(out) == 0 {
return nil
}
cursorRaw, hasCursor := st.GetRaw(afAlgCursorKey)
cursor := decodeCursor(cursorRaw)
var findings []alert.Finding
highest := cursor
highestSet := hasCursor
for _, line := range strings.Split(string(out), "\n") {
ev, ok := parseAFAlgEvent(line)
if !ok {
continue
}
if hasCursor && !ev.after(cursor) {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "af_alg_socket_use",
Message: fmt.Sprintf("AF_ALG socket opened by uid=%s exe=%s", ev.UID, ev.Exe),
TenantID: AFAlgOwner(ev),
Details: fmt.Sprintf(
"Audit event: timestamp=%s serial=%s\nauid=%s uid=%s comm=%q exe=%q\n"+
"AF_ALG is essentially never used by cPanel/PHP workloads. This is\n"+
"the kernel-level exploit signature for CVE-2026-31431 (\"Copy Fail\").\n"+
"Investigate this process immediately and consider unloading algif_aead\n"+
"(modprobe -r algif_aead af_alg) and adding a modprobe.d blacklist.",
ev.Timestamp, ev.Serial, ev.AUID, ev.UID, ev.Comm, ev.Exe,
),
})
if !highestSet || ev.after(highest) {
highest = ev
highestSet = true
}
}
if highestSet {
st.SetRaw(afAlgCursorKey, encodeCursor(highest))
}
return findings
}
func encodeCursor(ev afAlgEvent) string { return ev.Timestamp + ":" + ev.Serial }
func decodeCursor(s string) afAlgEvent {
if s == "" {
return afAlgEvent{}
}
colon := strings.Index(s, ":")
if colon < 0 {
return afAlgEvent{}
}
return afAlgEvent{Timestamp: s[:colon], Serial: s[colon+1:]}
}
package checks
import (
"context"
"errors"
"fmt"
"os"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/systemdrun"
)
// EnforceAction is the discrete outcome of the pure enforcement decision.
// Each value corresponds to one operational step the impure wrapper takes.
type EnforceAction int
const (
EnforceActionNoop EnforceAction = iota
EnforceActionRestoreMarker
EnforceActionUnloadModules
EnforceActionRestoreAndUnload
)
// decideAFAlgEnforcement is the pure, deterministic core of the enforcement
// check. Given the observed state of the marker file and the kernel module
// table, it returns exactly one action.
//
// Inputs:
// - markerPresent: /etc/modprobe.d/csm-copy-fail-mitigation.conf exists.
// - markerContentValid: that file's contents match the canonical CSM-managed
// content (so a hand-edited version still triggers a rewrite).
// - loaded: algif_aead OR af_alg is currently in /proc/modules.
//
// The "marker absent + modules loaded" combination intentionally returns
// Noop. The operator has not opted in to enforcement (no marker), so we will
// not unilaterally unload kernel modules they may legitimately be using —
// the existing hardening audit and auditd tripwire still surface the gap.
func decideAFAlgEnforcement(markerPresent, markerContentValid, loaded bool) EnforceAction {
if !markerPresent {
return EnforceActionNoop
}
switch {
case markerContentValid && !loaded:
return EnforceActionNoop
case markerContentValid && loaded:
return EnforceActionUnloadModules
case !markerContentValid && !loaded:
return EnforceActionRestoreMarker
default: // !markerContentValid && loaded
return EnforceActionRestoreAndUnload
}
}
// afAlgMarkerPath is the canonical location of the CSM-managed mitigation
// marker. Its presence is the signal that operator-driven enforcement is
// active for this host.
const afAlgMarkerPath = "/etc/modprobe.d/csm-copy-fail-mitigation.conf"
// canonicalAFAlgMarker is the byte-exact content the enforcer writes and
// re-asserts on drift. Hand-written variants (`blacklist algif_aead`, etc.)
// still satisfy the hardening audit, but the enforcer rewrites them to this
// canonical form so the file's content can be trivially validated.
const canonicalAFAlgMarker = `# CSM Copy Fail (CVE-2026-31431) mitigation — managed by CSM.
# Restored automatically by the af_alg_enforce critical-tier check.
# Remove this file (and run ` + "`csm harden --copy-fail`" + ` again) if you
# need to re-enable AF_ALG.
install algif_aead /bin/false
install af_alg /bin/false
`
// EnforceResult describes what enforceAFAlgBlocked observed and did, in a
// shape both the CLI subcommand and the periodic Check function can format
// for the operator without re-deriving the same conclusions.
//
// ModuleUnloaded reports the OBSERVED post-call state, not the syscall
// attempt: it is true only when /proc/modules no longer contains the
// targeted modules after `modprobe -r` ran. Use this field to distinguish
// "unload succeeded" from "unload attempted but module is in use".
type EnforceResult struct {
Action EnforceAction
MarkerPresent bool
MarkerValid bool
ModulesLoaded []string // names of currently-loaded targeted modules at start of call
MarkerWritten bool // wrapper wrote/restored the marker file this call
ModuleUnloaded bool // post-call /proc/modules shows targeted modules gone
Notes []string // operator-readable lines (warnings, stuck-module names)
}
func validateMarkerContent(data []byte) bool {
return string(data) == canonicalAFAlgMarker
}
// loadedTargetedModules returns the subset of {algif_aead, af_alg} currently
// present in /proc/modules. Used both before unload (to decide what to do)
// and after unload (to verify it actually took effect).
func loadedTargetedModules() []string {
var loaded []string
for _, mod := range loadModuleList() {
if mod == "algif_aead" || mod == "af_alg" {
loaded = append(loaded, mod)
}
}
return loaded
}
// enforceAFAlgBlocked inspects the marker file and /proc/modules, calls the
// pure decideAFAlgEnforcement, and applies the resulting action via osFS
// and cmdExec. Errors from osFS.WriteFile or unexpected osFS.Stat failures
// are returned; modprobe outcomes are observed via a post-call /proc/modules
// re-read, so an attempted unload is never reported as an observed success.
func enforceAFAlgBlocked() (EnforceResult, error) {
res := EnforceResult{}
// Marker presence + content check. ErrNotExist is the expected "advisory
// mode" path; any other Stat failure (e.g. EACCES on a hardened
// /etc/modprobe.d/) is surfaced as an error rather than silently
// classifying the host as advisory mode.
switch _, err := osFS.Stat(afAlgMarkerPath); {
case err == nil:
res.MarkerPresent = true
if data, readErr := osFS.ReadFile(afAlgMarkerPath); readErr == nil {
res.MarkerValid = validateMarkerContent(data)
}
case errors.Is(err, os.ErrNotExist):
// Advisory mode — operator has not opted in.
default:
return res, fmt.Errorf("stat %s: %w", afAlgMarkerPath, err)
}
res.ModulesLoaded = loadedTargetedModules()
res.Action = decideAFAlgEnforcement(res.MarkerPresent, res.MarkerValid, len(res.ModulesLoaded) > 0)
switch res.Action {
case EnforceActionRestoreMarker, EnforceActionRestoreAndUnload:
if err := osFS.WriteFile(afAlgMarkerPath, []byte(canonicalAFAlgMarker), 0o644); err != nil {
return res, err
}
res.MarkerWritten = true
}
switch res.Action {
case EnforceActionUnloadModules, EnforceActionRestoreAndUnload:
if err := unloadAFAlgModules(); err != nil {
res.Notes = append(res.Notes, fmt.Sprintf("module unload command failed: %v", err))
}
// Observe the kernel state even when the command fails or times out.
stillLoaded := loadedTargetedModules()
if len(stillLoaded) == 0 {
res.ModuleUnloaded = true
} else {
res.Notes = append(res.Notes, fmt.Sprintf(
"modprobe -r attempted but %v still loaded — module is in use; will retry next tick",
stillLoaded,
))
}
}
return res, nil
}
// Module removal is denied by the daemon's seccomp filter and kernel module
// protection. Keep those restrictions and delegate only the fixed opt-in action.
func unloadAFAlgModules() error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
_, err := systemdrun.Run(ctx, cmdExec.LookPath, cmdExec.RunContext, systemdrun.Options{Pipe: true, RuntimeMax: 30 * time.Second}, "modprobe", "-r", "algif_aead", "af_alg")
return err
}
// AFAlgMarkerPath returns the canonical marker file location. Exposed for
// the cmd/csm CLI which prints it to operators; production code in this
// package should reference the unexported constant directly.
func AFAlgMarkerPath() string { return afAlgMarkerPath }
// WriteAFAlgMarker forces the canonical marker content to disk regardless
// of current state. Used by `csm harden --copy-fail` to ensure subsequent
// EnforceAFAlgBlocked() runs see a valid marker even on first install.
func WriteAFAlgMarker() error {
return osFS.WriteFile(afAlgMarkerPath, []byte(canonicalAFAlgMarker), 0o644)
}
// EnforceAFAlgBlocked is the exported alias of enforceAFAlgBlocked for use
// by cmd/csm. The unexported form stays internal to the package so the
// periodic Check (Task 6) can call it without going through the export.
func EnforceAFAlgBlocked() (EnforceResult, error) { return enforceAFAlgBlocked() }
// enforceAFAlgBlockedFn is indirected so the observe-mode gate can be tested
// without a kernel to unload modules from.
var enforceAFAlgBlockedFn = enforceAFAlgBlocked
// CheckAFAlgEnforcement is the periodic critical-tier check that enforces
// the AF_ALG mitigation policy. When the operator has opted in (via
// `csm harden --copy-fail`, which writes the marker file), this check
// reverts any drift on each tick. It is a no-op in advisory mode.
//
// Emits a Warning finding (one per tick that took action) so the operator
// has an alert-pipeline record that system state was modified. Steady-state
// ticks emit no findings. Warning is the lowest severity available in the
// alert.Severity enum (Warning < High < Critical, no Info level).
func CheckAFAlgEnforcement(_ context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if cfg != nil && (cfg.AutoResponse.DisableEnforceAFAlg || cfg.ObserveMode()) {
return nil
}
res, err := enforceAFAlgBlockedFn()
if err != nil {
return []alert.Finding{{
Severity: alert.Warning,
Check: "af_alg_enforcement_corrected",
Message: "AF_ALG enforcement encountered an error",
Details: fmt.Sprintf("error: %v\nresult: %+v", err, res),
}}
}
if res.Action == EnforceActionNoop {
return nil
}
return []alert.Finding{{
Severity: alert.Warning,
Check: "af_alg_enforcement_corrected",
Message: fmt.Sprintf("AF_ALG enforcement re-applied (%s)", actionName(res.Action)),
Details: fmt.Sprintf(
"Action: %s\nMarker present: %v\nMarker valid: %v\nModules loaded: %v\nMarker written: %v\nModule unload succeeded: %v\nNotes: %v",
actionName(res.Action), res.MarkerPresent, res.MarkerValid, res.ModulesLoaded,
res.MarkerWritten, res.ModuleUnloaded, res.Notes,
),
}}
}
func actionName(a EnforceAction) string {
switch a {
case EnforceActionNoop:
return "Noop"
case EnforceActionRestoreMarker:
return "RestoreMarker"
case EnforceActionUnloadModules:
return "UnloadModules"
case EnforceActionRestoreAndUnload:
return "RestoreAndUnload"
}
return fmt.Sprintf("Unknown(%d)", int(a))
}
package checks
import (
"bytes"
"compress/gzip"
"os"
"strings"
)
// copyFailCVE is the CVE identifier used to match KernelCare/kpatch entries.
const copyFailCVE = "CVE-2026-31431"
// configHasBuiltInAFAlgAEAD reports whether the supplied kernel config text
// declares CRYPTO_USER_API_AEAD as built into the kernel image (=y) rather
// than as a loadable module (=m). On =y kernels, the modprobe blacklist
// mitigation is ineffective because there is no module to block from
// loading — the code is always present.
//
// The function tolerates the "is not set" comment form that kconfig uses
// for explicitly-disabled options. An unset CRYPTO_USER_API_AEAD returns
// false (not built-in; modular or absent).
func configHasBuiltInAFAlgAEAD(configText string) bool {
for _, line := range strings.Split(configText, "\n") {
line = strings.TrimSpace(line)
if line == "CONFIG_CRYPTO_USER_API_AEAD=y" {
return true
}
}
return false
}
// configHasModularAFAlgAEAD reports whether CRYPTO_USER_API_AEAD is set
// to =m (loadable module) in the supplied kernel config. Used by the
// "is this host actually exploitable?" policy decision: a kernel built
// without =y AND without =m has no AF_ALG aead interface at all, so
// Copy Fail is not reachable on it.
func configHasModularAFAlgAEAD(configText string) bool {
for _, line := range strings.Split(configText, "\n") {
line = strings.TrimSpace(line)
if line == "CONFIG_CRYPTO_USER_API_AEAD=m" {
return true
}
}
return false
}
// kcareReportsCopyFailPatched reports whether the supplied `kcarectl
// --patch-info` output advertises a patch covering the Copy Fail CVE.
// The kcarectl format emits per-patch records like:
//
// kpatch-name: rhel8/.../CVE-2026-NNNNN-foo.patch
// kpatch-cve: CVE-2026-NNNNN
// kpatch-cve-url: ...
//
// We match on a substring of the literal CVE id rather than the URL or
// filename so a future format change to either does not silently break
// detection.
func kcareReportsCopyFailPatched(out []byte) bool {
return bytes.Contains(out, []byte(copyFailCVE))
}
// kernelHasBuiltInAFAlgAEAD is the impure wrapper: it reads the kernel
// config from /boot/config-$(uname -r), falling back to /proc/config.gz,
// and asks configHasBuiltInAFAlgAEAD whether the AEAD interface is
// statically linked.
//
// Returns (false, nil) when no config file is readable — the caller
// treats "config unknown" the same as "modular" so an inability to read
// the config never silently downgrades a real protection state.
func kernelHasBuiltInAFAlgAEAD() (bool, error) {
cfg, _ := readKernelConfigText()
if cfg == "" {
return false, nil
}
return configHasBuiltInAFAlgAEAD(cfg), nil
}
// readKernelConfigText returns the running kernel's .config text from
// /boot/config-$(uname -r), falling back to /proc/config.gz (gunzipped).
// Returns ("", false) when neither is readable so the caller can decide
// whether to treat "unknown" as conservatively-vulnerable.
func readKernelConfigText() (string, bool) {
if uname, err := readKernelRelease(); err == nil {
if data, err := osFS.ReadFile("/boot/config-" + uname); err == nil {
return string(data), true
}
}
if data, err := osFS.ReadFile("/proc/config.gz"); err == nil {
zr, err := gzip.NewReader(bytes.NewReader(data))
if err != nil {
return "", false
}
defer func() { _ = zr.Close() }()
var buf bytes.Buffer
if _, err := buf.ReadFrom(zr); err != nil {
return "", false
}
return buf.String(), true
}
return "", false
}
// readKernelRelease returns the running kernel release string (the same
// value `uname -r` would print). Used to find the matching config-* file
// under /boot.
func readKernelRelease() (string, error) {
data, err := osFS.ReadFile("/proc/sys/kernel/osrelease")
if err != nil {
return "", err
}
return strings.TrimSpace(string(data)), nil
}
// kcareHasCopyFailPatch runs `kcarectl --patch-info` and reports whether
// KernelCare has applied a livepatch covering Copy Fail. Returns false
// (with a nil error) when kcarectl is absent or fails — KernelCare is
// optional and its absence is not an error condition.
func kcareHasCopyFailPatch() bool {
out, err := cmdExec.RunAllowNonZero("kcarectl", "--patch-info")
if err != nil {
return false
}
if len(out) == 0 {
return false
}
return kcareReportsCopyFailPatched(out)
}
// AFAlgKernelState is the assembled view of how the running kernel
// exposes AF_ALG. Used by the hardening audit, by csm harden, and by
// the live-monitor coordinator to decide whether protection is needed.
type AFAlgKernelState struct {
BuiltIn bool // CONFIG_CRYPTO_USER_API_AEAD=y in the running kernel
Modular bool // CONFIG_CRYPTO_USER_API_AEAD=m (loadable module exists)
ConfigReadable bool // /boot/config-$(uname -r) or /proc/config.gz was parseable
LivepatchActive bool // KernelCare/kpatch has applied a CVE-2026-31431 patch
}
// observeAFAlgKernelState assembles a kernel-state snapshot via the impure
// helpers above. The struct fields document precisely what we know vs
// what we couldn't determine — callers can apply policy without
// re-deriving the same probes.
func observeAFAlgKernelState() AFAlgKernelState {
state := AFAlgKernelState{}
if cfg, ok := readKernelConfigText(); ok {
state.ConfigReadable = true
state.BuiltIn = configHasBuiltInAFAlgAEAD(cfg)
state.Modular = configHasModularAFAlgAEAD(cfg)
}
state.LivepatchActive = kcareHasCopyFailPatch()
return state
}
// IsCopyFailExploitable reports whether this kernel is currently
// vulnerable to Copy Fail (CVE-2026-31431). Used by the daemon's
// live-monitor coordinator to skip starting the listener entirely on
// hosts that don't need protection — saving the inotify watch + tick
// loop for hosts that actually face the threat.
//
// Conservative defaults: when the kernel config is unreadable, treat
// the host as exploitable (better to over-monitor than miss). When a
// KernelCare livepatch is in place, treat as patched regardless of the
// underlying config — the syscall path itself is fixed.
func (s AFAlgKernelState) IsCopyFailExploitable() bool {
if s.LivepatchActive {
return false
}
if s.ConfigReadable && !s.BuiltIn && !s.Modular {
// Kernel was definitively built without the AF_ALG aead interface.
// Nothing to exploit, no listener needed.
return false
}
// Either: confirmed-vulnerable (=y or =m without livepatch), or
// unknown (config unreadable). Both go to "exploitable" so we err
// on the side of monitoring.
return true
}
// String renders the kernel state for inclusion in operator-visible
// messages. The format is stable and short enough to embed in a single
// AuditResult.Message line.
func (s AFAlgKernelState) String() string {
switch {
case s.LivepatchActive:
return "kernel patched by KernelCare (livepatch active for " + copyFailCVE + ")"
case s.BuiltIn:
return "AF_ALG is built into the kernel (CONFIG_CRYPTO_USER_API_AEAD=y) and no livepatch is active"
case s.Modular:
return "AF_ALG aead is a loadable module on this kernel"
case s.ConfigReadable:
return "AF_ALG aead is not present in this kernel build"
default:
return "kernel config unreadable; treating as potentially vulnerable"
}
}
// ObserveAFAlgKernelState is the exported alias for cmd/csm and the
// daemon. Production code inside this package uses the unexported form
// directly.
func ObserveAFAlgKernelState() AFAlgKernelState { return observeAFAlgKernelState() }
// EnsureFile is a sentinel return value: callers can use os.IsNotExist
// to test for the "kernel-config file absent" case explicitly.
var ErrKernelConfigUnreadable = os.ErrNotExist
package checks
import (
"fmt"
"os"
"path/filepath"
"strings"
)
// SeccompDropInBaseName is the file name CSM writes inside each unit's
// /etc/systemd/system/<unit>.d/ override directory. Stable so the
// hardening audit can scan for it and the remove path can clean up
// without guessing.
const SeccompDropInBaseName = "csm-copy-fail-seccomp.conf"
// seccompDropInContent is the byte-exact body of every drop-in CSM
// writes. The marker comment lets a curious sysadmin understand who
// manages the file without consulting external docs.
const seccompDropInContent = `# CSM Copy Fail (CVE-2026-31431) seccomp mitigation - managed by CSM.
# Blocks socket(AF_ALG, ...) for processes spawned by this unit.
# Remove this file to disable.
[Service]
RestrictAddressFamilies=~AF_ALG
`
// afAlgSeccompCandidateUnits is the catalog of systemd units that, on
// shared-hosting servers, regularly spawn untrusted user-level code
// and therefore need the AF_ALG block. Units that do not exist on the
// running host are filtered out at apply time.
//
// The list is intentionally inclusive: a unit that does not exist
// adds zero overhead because we filter via systemctl list-unit-files
// before writing anything.
var afAlgSeccompCandidateUnits = []string{
// Web servers
"lshttpd.service", // LiteSpeed (cPanel default)
"httpd.service", // Apache (RHEL family, cPanel EA4 fallback)
"apache2.service", // Apache (Debian/Ubuntu)
"nginx.service", // Nginx
// PHP-FPM, cPanel EA4
"ea-php72-php-fpm.service",
"ea-php73-php-fpm.service",
"ea-php74-php-fpm.service",
"ea-php80-php-fpm.service",
"ea-php81-php-fpm.service",
"ea-php82-php-fpm.service",
"ea-php83-php-fpm.service",
"ea-php84-php-fpm.service",
"cpanel_php_fpm.service",
// PHP-FPM, generic distro versions
"php-fpm.service",
"php7.4-fpm.service",
"php8.0-fpm.service",
"php8.1-fpm.service",
"php8.2-fpm.service",
"php8.3-fpm.service",
"php8.4-fpm.service",
// Cron and mail (drop privileges to user before running content
// filters or scheduled scripts)
"crond.service",
"cron.service",
"exim.service",
"dovecot.service",
}
// SeccompUnitState describes one unit's mitigation status in operator
// terms. Returned by ScanAFAlgSeccompState so both the CLI and the
// hardening audit can render the same view without re-deriving it.
type SeccompUnitState struct {
Unit string // e.g. "lshttpd.service"
Exists bool // unit is registered with systemd on this host
HasFile bool // CSM-managed drop-in is present on disk
}
// ScanAFAlgSeccompState walks the candidate unit list and returns one
// SeccompUnitState per candidate. Units missing from systemd are still
// reported (Exists=false) so the operator can confirm CSM did not
// silently skip something they expected.
func ScanAFAlgSeccompState() []SeccompUnitState {
var out []SeccompUnitState
for _, u := range afAlgSeccompCandidateUnits {
s := SeccompUnitState{Unit: u}
s.Exists = systemdUnitExists(u)
s.HasFile = seccompDropInPresent(u)
out = append(out, s)
}
return out
}
// SeccompCoverageSummary collapses the per-unit scan into the two
// numbers an operator cares about: how many existing units have the
// CSM drop-in, and how many do not.
type SeccompCoverageSummary struct {
Covered []string // existing units with the CSM drop-in
Uncovered []string // existing units without the drop-in
NotInstalled []string // candidate units not registered with systemd
}
// SummarizeAFAlgSeccompCoverage rolls up ScanAFAlgSeccompState into
// the three-way summary above. Used by the hardening audit and the
// CLI status output.
func SummarizeAFAlgSeccompCoverage() SeccompCoverageSummary {
var s SeccompCoverageSummary
for _, u := range ScanAFAlgSeccompState() {
switch {
case !u.Exists:
s.NotInstalled = append(s.NotInstalled, u.Unit)
case u.HasFile:
s.Covered = append(s.Covered, u.Unit)
default:
s.Uncovered = append(s.Uncovered, u.Unit)
}
}
return s
}
// ApplyAFAlgSeccompDropIns writes the canonical drop-in file for every
// candidate unit that exists on this host AND does not already have
// the file. After all writes, runs systemctl daemon-reload and a
// reload-or-restart per touched unit so the seccomp filter takes
// effect immediately.
//
// Returns the list of units that received a new drop-in this call. An
// empty list with a nil error means everything was already covered
// (idempotent re-run).
func ApplyAFAlgSeccompDropIns() ([]string, error) {
var written []string
for _, u := range afAlgSeccompCandidateUnits {
if !systemdUnitExists(u) {
continue
}
if seccompDropInPresent(u) {
continue
}
if err := writeSeccompDropIn(u); err != nil {
return written, fmt.Errorf("write drop-in for %s: %w", u, err)
}
written = append(written, u)
}
if len(written) == 0 {
return nil, nil
}
if _, err := cmdExec.Run("systemctl", "daemon-reload"); err != nil {
return written, fmt.Errorf("systemctl daemon-reload: %w", err)
}
for _, u := range written {
// try-restart performs a full restart (re-exec the master process)
// only if the unit is currently active. A reload (e.g., SIGUSR2 to
// PHP-FPM) is NOT enough: systemd attaches seccomp filters at
// process spawn time, so the existing master has to be replaced
// for RestrictAddressFamilies to take effect on it and its
// workers. try-restart skips units that were intentionally
// stopped, so we don't surprise-start anything.
if _, err := cmdExec.RunAllowNonZero("systemctl", "try-restart", u); err != nil {
return written, fmt.Errorf("systemctl try-restart %s: %w", u, err)
}
}
return written, nil
}
// RemoveAFAlgSeccompDropIns deletes every CSM-managed seccomp drop-in
// found on disk and runs systemctl daemon-reload + reload-or-restart
// per touched unit so the seccomp filter is dropped from running
// processes. Idempotent: a unit without our drop-in is skipped.
//
// Returns the list of units whose drop-in was removed.
func RemoveAFAlgSeccompDropIns() ([]string, error) {
var removed []string
for _, u := range afAlgSeccompCandidateUnits {
if !seccompDropInPresent(u) {
continue
}
if err := osFS.Remove(seccompDropInPath(u)); err != nil && !os.IsNotExist(err) {
return removed, fmt.Errorf("remove drop-in for %s: %w", u, err)
}
// Best-effort: clean up the now-empty .d directory if we created it.
dir := seccompDropInDir(u)
_ = osFS.Remove(dir) // succeeds only if empty; harmless otherwise
removed = append(removed, u)
}
if len(removed) == 0 {
return nil, nil
}
if _, err := cmdExec.Run("systemctl", "daemon-reload"); err != nil {
return removed, fmt.Errorf("systemctl daemon-reload: %w", err)
}
for _, u := range removed {
if !systemdUnitExists(u) {
continue
}
// Same reasoning as Apply: full re-exec is required so the
// seccomp filter is dropped from the running master.
if _, err := cmdExec.RunAllowNonZero("systemctl", "try-restart", u); err != nil {
return removed, fmt.Errorf("systemctl try-restart %s: %w", u, err)
}
}
return removed, nil
}
// seccompDropInDir returns the override directory for the given unit:
// /etc/systemd/system/<unit>.d
func seccompDropInDir(unit string) string {
return filepath.Join("/etc/systemd/system", unit+".d")
}
// seccompDropInPath returns the full path to the CSM-managed drop-in
// file for the given unit.
func seccompDropInPath(unit string) string {
return filepath.Join(seccompDropInDir(unit), SeccompDropInBaseName)
}
// seccompDropInPresent reports whether the CSM-managed drop-in exists
// for the given unit. Content is not inspected here; the file's
// presence at the canonical path is the policy signal.
func seccompDropInPresent(unit string) bool {
_, err := osFS.Stat(seccompDropInPath(unit))
return err == nil
}
// writeSeccompDropIn creates the override directory and writes the
// canonical drop-in content. Errors propagate to the caller.
func writeSeccompDropIn(unit string) error {
dir := seccompDropInDir(unit)
if err := osFS.MkdirAll(dir, 0o755); err != nil {
return err
}
return osFS.WriteFile(seccompDropInPath(unit), []byte(seccompDropInContent), 0o644)
}
// systemdUnitExists asks systemctl whether the given unit is known to
// the running systemd. Returns false on any error, including missing
// systemctl binary, so a non-systemd host (rare on RHEL/Ubuntu) is
// treated as "no units to mitigate."
func systemdUnitExists(unit string) bool {
out, err := cmdExec.RunAllowNonZero(
"systemctl", "list-unit-files", unit, "--no-legend", "--no-pager",
)
if err != nil {
return false
}
// list-unit-files prints a row per match. An empty stdout means
// the unit name is not registered with this systemd instance.
return len(strings.TrimSpace(string(out))) > 0
}
package checks
import (
"errors"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
)
// A running Apache rejects recursive includes. The audit still bounds descent
// so a changed-on-disk configuration cannot recurse forever before the next
// config test or reload notices the error.
const apacheIncludeMaxDepth = 32
// apacheConfigLine is one line of an assembled Apache configuration, tagged
// with its source so findings can name the file that set a directive.
type apacheConfigLine struct {
File string
Text string
}
var apacheConditionalContainers = map[string]bool{
"if": true,
"else": true,
"elseif": true,
"ifmodule": true,
"ifdefine": true,
"ifversion": true,
"iffile": true,
"ifsection": true,
"ifdirective": true,
}
// assembleApacheConfig reads path and splices Include and IncludeOptional
// targets at the directive position. Apache applies later directives over
// earlier ones, so appending snippets would invert precedence.
func assembleApacheConfig(path string) []apacheConfigLine {
lines, _ := assembleApacheConfigWithStatus(path)
return lines
}
func assembleApacheConfigWithStatus(path string) ([]apacheConfigLine, bool) {
assembler := apacheConfigAssembler{
serverRoot: defaultApacheServerRoot(path),
activePaths: make(map[string]bool),
complete: true,
}
lines := assembler.appendPath(nil, path, false, 0)
if len(assembler.containers) != 0 {
assembler.complete = false
}
return lines, assembler.complete
}
// defaultApacheServerRoot covers the compiled-in layouts used by supported
// distributions when the main file does not set ServerRoot.
func defaultApacheServerRoot(path string) string {
dir := filepath.Dir(path)
if strings.EqualFold(filepath.Base(path), "httpd.conf") &&
strings.EqualFold(filepath.Base(dir), "conf") {
return filepath.Dir(dir)
}
return dir
}
type apacheConfigAssembler struct {
serverRoot string
activePaths map[string]bool
containers []string
complete bool
}
func (a *apacheConfigAssembler) appendPath(dst []apacheConfigLine, path string, optional bool, depth int) []apacheConfigLine {
if depth > apacheIncludeMaxDepth {
a.complete = false
return dst
}
abs, err := filepath.Abs(path)
if err != nil {
abs = filepath.Clean(path)
}
if a.activePaths[abs] {
a.complete = false
return dst
}
a.activePaths[abs] = true
defer delete(a.activePaths, abs)
if info, statErr := osFS.Stat(path); statErr == nil && info.IsDir() {
entries, globErr := osFS.Glob(filepath.Join(path, "*"))
if globErr != nil {
a.complete = false
return dst
}
sort.Strings(entries)
for _, entry := range entries {
dst = a.appendPath(dst, entry, false, depth+1)
}
return dst
}
data, err := osFS.ReadFile(path)
if err != nil {
if !optional || !errors.Is(err, os.ErrNotExist) {
a.complete = false
}
return dst
}
lines := strings.Split(string(data), "\n")
if n := len(lines); n > 0 && lines[n-1] == "" {
lines = lines[:n-1]
}
continued := ""
for _, physical := range lines {
line := continued + physical
if prefix, continues := apacheContinuationPrefix(line); continues {
continued = prefix + " "
continue
}
continued = ""
fields, valid := parseApacheDirectiveFields(stripApacheComment(line))
if !valid {
a.complete = false
dst = append(dst, apacheConfigLine{File: path, Text: line})
continue
}
if len(fields) == 0 {
dst = append(dst, apacheConfigLine{File: path, Text: line})
continue
}
if tag, isContainer, tagValid := parseApacheContainerTag(line); isContainer {
switch {
case !tagValid:
a.complete = false
case tag.closing:
last := len(a.containers) - 1
if last < 0 || !strings.EqualFold(a.containers[last], tag.name) {
a.complete = false
} else {
a.containers = a.containers[:last]
}
default:
a.containers = append(a.containers, tag.name)
}
dst = append(dst, apacheConfigLine{File: path, Text: line})
continue
}
switch {
case strings.EqualFold(fields[0], "ServerRoot"):
if len(a.containers) != 0 || len(fields) != 2 || fields[1] == "" || apacheArgumentHasVariable(fields[1]) || !filepath.IsAbs(fields[1]) {
a.complete = false
} else {
a.serverRoot = filepath.Clean(fields[1])
}
dst = append(dst, apacheConfigLine{File: path, Text: line})
case strings.EqualFold(fields[0], "Include"), strings.EqualFold(fields[0], "IncludeOptional"):
optionalInclude := strings.EqualFold(fields[0], "IncludeOptional")
if len(fields) != 2 || fields[1] == "" || apacheArgumentHasVariable(fields[1]) {
a.complete = false
continue
}
paths, expanded := expandApacheInclude(fields[1], a.serverRoot)
if !expanded || (!optionalInclude && len(paths) == 0) {
a.complete = false
}
for _, included := range paths {
dst = a.appendPath(dst, included, optionalInclude, depth+1)
}
default:
dst = append(dst, apacheConfigLine{File: path, Text: line})
}
}
if continued != "" {
a.complete = false
}
return dst
}
func apacheArgumentHasVariable(arg string) bool {
return strings.Contains(arg, "${")
}
func expandApacheInclude(target, serverRoot string) ([]string, bool) {
if !filepath.IsAbs(target) {
target = filepath.Join(serverRoot, target)
}
if !strings.ContainsAny(target, "*?[") {
return []string{target}, true
}
matches, err := osFS.Glob(target)
if err != nil {
return nil, false
}
sort.Strings(matches)
return matches, true
}
func apacheContinuationPrefix(line string) (string, bool) {
line = strings.TrimRight(line, " \t\r")
backslashes := 0
for i := len(line) - 1; i >= 0 && line[i] == '\\'; i-- {
backslashes++
}
if backslashes%2 == 0 {
return line, false
}
return line[:len(line)-1], true
}
func stripApacheComment(line string) string {
var quote byte
escaped := false
for i := 0; i < len(line); i++ {
c := line[i]
if escaped {
escaped = false
continue
}
if c == '\\' {
escaped = true
continue
}
if quote != 0 {
if c == quote {
quote = 0
}
continue
}
if c == '\'' || c == '"' {
quote = c
continue
}
if c == '#' {
return line[:i]
}
}
return line
}
// parseApacheDirectiveFields handles the quoting and escapes accepted in
// include paths and container arguments.
func parseApacheDirectiveFields(line string) ([]string, bool) {
line = strings.TrimSpace(line)
if line == "" {
return nil, true
}
var fields []string
var field strings.Builder
var quote byte
escaped := false
inField := false
flush := func() {
if !inField {
return
}
fields = append(fields, field.String())
field.Reset()
inField = false
}
for i := 0; i < len(line); i++ {
c := line[i]
if escaped {
field.WriteByte(c)
inField = true
escaped = false
continue
}
if c == '\\' {
escaped = true
inField = true
continue
}
if quote != 0 {
if c == quote {
quote = 0
} else {
field.WriteByte(c)
}
inField = true
continue
}
if c == '\'' || c == '"' {
quote = c
inField = true
continue
}
if c == ' ' || c == '\t' || c == '\r' {
flush()
continue
}
field.WriteByte(c)
inField = true
}
flush()
return fields, quote == 0 && !escaped
}
type apacheParsedContainerTag struct {
name string
label string
closing bool
}
func parseApacheContainerTag(line string) (apacheParsedContainerTag, bool, bool) {
text := strings.TrimSpace(stripApacheComment(line))
if !strings.HasPrefix(text, "<") {
return apacheParsedContainerTag{}, false, true
}
if !strings.HasSuffix(text, ">") {
return apacheParsedContainerTag{}, true, false
}
closing := strings.HasPrefix(text, "</")
prefix := "<"
if closing {
prefix = "</"
}
label := strings.TrimSpace(strings.TrimSuffix(strings.TrimPrefix(text, prefix), ">"))
fields, valid := parseApacheDirectiveFields(label)
if !valid || len(fields) == 0 || (closing && len(fields) != 1) {
return apacheParsedContainerTag{}, true, false
}
return apacheParsedContainerTag{name: fields[0], label: label, closing: closing}, true, true
}
type apacheLineContext struct {
Scope []string
ScopeKey []string
Condition string
LateCondition bool
}
// walkApacheLines calls fn for directive lines and validates the container
// stack. A mismatched close never pops an unrelated scope.
func walkApacheLines(lines []apacheConfigLine, fn func(apacheConfigLine, []string, apacheLineContext)) bool {
type container struct {
name string
label string
key string
scoped bool
conditional bool
lateConditional bool
}
var stack []container
valid := true
serial := 0
context := func() apacheLineContext {
var ctx apacheLineContext
var conditions []string
for _, item := range stack {
if item.scoped {
ctx.Scope = append(ctx.Scope, item.label)
ctx.ScopeKey = append(ctx.ScopeKey, item.key)
}
if item.conditional {
conditions = append(conditions, item.key)
}
if item.lateConditional {
ctx.LateCondition = true
}
}
ctx.Condition = strings.Join(conditions, "\x1e")
return ctx
}
for _, line := range lines {
text := strings.TrimSpace(stripApacheComment(line.Text))
if text == "" {
continue
}
tag, isContainer, tagValid := parseApacheContainerTag(text)
if isContainer {
if !tagValid {
valid = false
continue
}
if tag.closing {
if len(stack) == 0 || !strings.EqualFold(stack[len(stack)-1].name, tag.name) {
valid = false
continue
}
stack = stack[:len(stack)-1]
continue
}
serial++
lowerName := strings.ToLower(tag.name)
conditional := apacheConditionalContainers[lowerName]
key := canonicalApacheContainerKey(tag.label)
if conditional || lowerName == "virtualhost" {
key += "\x00" + strconv.Itoa(serial)
}
stack = append(stack, container{
name: tag.name,
label: normalizeApacheContainerLabel(tag.label),
key: key,
scoped: !conditional,
conditional: conditional,
lateConditional: lowerName == "if" || lowerName == "else" || lowerName == "elseif",
})
continue
}
fields, fieldsValid := parseApacheDirectiveFields(text)
if !fieldsValid {
valid = false
continue
}
if len(fields) > 0 {
fn(line, fields, context())
}
}
return valid && len(stack) == 0
}
func canonicalApacheContainerKey(label string) string {
fields, valid := parseApacheDirectiveFields(label)
if !valid || len(fields) == 0 {
return strings.ToLower(strings.TrimSpace(label))
}
fields[0] = strings.ToLower(fields[0])
return strings.Join(fields, "\x1f")
}
func normalizeApacheContainerLabel(label string) string {
fields, valid := parseApacheDirectiveFields(label)
if !valid {
return strings.TrimSpace(label)
}
return strings.Join(fields, " ")
}
type apacheDirectiveValue struct {
File string
Scope string
Value string
Conditional bool
lateConditional bool
scopeKey string
}
// apacheDirectiveValues returns the last explicit value in each scope and
// keeps conditional branches separate from the unconditional default.
func apacheDirectiveValues(lines []apacheConfigLine, directive string) ([]apacheDirectiveValue, bool) {
const serverScope = "server config"
values := make(map[string]apacheDirectiveValue)
var order []string
conditionalValues := make(map[string]apacheDirectiveValue)
var conditionalOrder []string
directiveValid := true
containersValid := walkApacheLines(lines, func(line apacheConfigLine, fields []string, ctx apacheLineContext) {
if !strings.EqualFold(fields[0], directive) {
return
}
if len(fields) < 2 {
directiveValid = false
return
}
key := serverScope
displayScope := serverScope
if len(ctx.Scope) > 0 {
key = strings.Join(ctx.ScopeKey, "\x1d")
displayScope = strings.Join(ctx.Scope, " > ")
}
value := apacheDirectiveValue{
File: line.File,
Scope: displayScope,
Value: strings.Join(fields[1:], " "),
Conditional: ctx.Condition != "",
lateConditional: ctx.LateCondition,
scopeKey: key,
}
if value.Conditional {
conditionalKey := key + "\x00" + ctx.Condition
if _, seen := conditionalValues[conditionalKey]; !seen {
conditionalOrder = append(conditionalOrder, conditionalKey)
}
conditionalValues[conditionalKey] = value
return
}
if _, seen := values[key]; !seen {
order = append(order, key)
}
values[key] = value
for conditionalKey, conditionalValue := range conditionalValues {
if conditionalValue.scopeKey == key && !conditionalValue.lateConditional {
conditionalValue.File = value.File
conditionalValue.Value = value.Value
conditionalValues[conditionalKey] = conditionalValue
}
}
})
out := make([]apacheDirectiveValue, 0, len(values)+len(conditionalValues))
for _, key := range order {
out = append(out, values[key])
}
for _, key := range conditionalOrder {
out = append(out, conditionalValues[key])
}
return out, containersValid && directiveValid
}
func apacheIndexesScopes(lines []apacheConfigLine) []string {
scopes, _ := apacheIndexesScopesWithStatus(lines)
return scopes
}
func apacheIndexesScopesWithStatus(lines []apacheConfigLine) ([]string, bool) {
const serverScope = "server config"
effective := make(map[string]bool)
conditionalEffective := make(map[string]bool)
conditionalScopes := make(map[string]string)
lateConditional := make(map[string]bool)
displayScopes := make(map[string]string)
seen := make(map[string]bool)
var order []string
optionsValid := true
containersValid := walkApacheLines(lines, func(_ apacheConfigLine, fields []string, ctx apacheLineContext) {
if !strings.EqualFold(fields[0], "Options") {
return
}
if len(fields) < 2 {
optionsValid = false
return
}
key := serverScope
displayScope := serverScope
if len(ctx.Scope) > 0 {
key = strings.Join(ctx.ScopeKey, "\x1d")
displayScope = strings.Join(ctx.Scope, " > ")
}
if !seen[key] {
seen[key] = true
order = append(order, key)
displayScopes[key] = displayScope
}
if ctx.Condition != "" {
conditionalKey := key + "\x00" + ctx.Condition
current, exists := conditionalEffective[conditionalKey]
if !exists {
current = effective[key]
}
conditionalEffective[conditionalKey] = optionsEnableIndexes(fields[1:], current)
conditionalScopes[conditionalKey] = key
lateConditional[conditionalKey] = ctx.LateCondition
return
}
effective[key] = optionsEnableIndexes(fields[1:], effective[key])
for conditionalKey, conditionalValue := range conditionalEffective {
if conditionalScopes[conditionalKey] == key && !lateConditional[conditionalKey] {
conditionalEffective[conditionalKey] = optionsEnableIndexes(fields[1:], conditionalValue)
}
}
})
var out []string
for _, key := range order {
enabled := effective[key]
for conditionalKey, conditionalValue := range conditionalEffective {
if conditionalScopes[conditionalKey] == key && conditionalValue {
enabled = true
break
}
}
if enabled {
out = append(out, displayScopes[key])
}
}
return out, containersValid && optionsValid
}
// optionsEnableIndexes applies one Options directive to the scope's current
// state. Signed tokens merge; bare tokens replace the option set.
func optionsEnableIndexes(tokens []string, current bool) bool {
merge := false
for _, token := range tokens {
if strings.HasPrefix(token, "+") || strings.HasPrefix(token, "-") {
merge = true
break
}
}
if !merge {
for _, token := range tokens {
if strings.EqualFold(token, "Indexes") || strings.EqualFold(token, "All") {
return true
}
}
return false
}
for _, token := range tokens {
name := strings.TrimLeft(token, "+-")
if !strings.EqualFold(name, "Indexes") && !strings.EqualFold(name, "All") {
continue
}
current = strings.HasPrefix(token, "+")
}
return current
}
package checks
import (
"errors"
"fmt"
"os"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
)
// Block sources label who asked for an auto-response block. They feed the
// outcome metric and let evidence rows say which pipeline produced them.
const (
BlockSourceScan = "scan"
BlockSourceChallenge = "challenge"
BlockSourceIncident = "incident"
BlockSourceCentral = "central_intel"
)
// ApplyBlockRequest describes one auto-response IP block. EngineReason is
// handed to the firewall engine (its provenance inference keys on it);
// Reason is the human evidence recorded in the threat DB, the tracker, and
// findings.
type ApplyBlockRequest struct {
// ActionID preserves one admission identity across retries.
ActionID string
// FindingID is the original audit identity, captured before display truncation.
// Empty means this decision has no originating finding.
FindingID string
IP string
EngineReason string
Reason string
TTL time.Duration
Source string
}
// ApplyBlockResult carries the engine outcome plus the auto_block findings
// the caller must route into its alert pipeline (scan batch or the daemon's
// async block recorder) so digests and alerting see every block source.
type ApplyBlockResult struct {
Outcome firewall.BlockOutcome
Findings []alert.Finding
}
// ErrNoIPBlocker is returned when no firewall engine is wired. Callers must
// treat it as "the block did not happen", never as success.
var ErrNoIPBlocker = errors.New("firewall engine not available")
// ApplyBlock is the single chokepoint for auto-response IP blocks issued
// outside the scan loop (challenge escalation, central intel, incident
// spray). It performs the block and the same evidence bookkeeping a scan
// auto-block gets: threat-DB row, blocked-IPs tracker entry, auto_block
// finding with the Cloudflare coverage warning, and permanent-block
// escalation counting. The scan loop shares the inner implementation and
// keeps its own batch semantics (rate limit, pending queue) around it.
//
// Non-scan sources deliberately neither consume nor enforce
// auto_response.max_blocks_per_hour: challenge escalation and central intel
// are already gated upstream, and letting them starve or be starved by the
// scan budget would change containment behavior.
func ApplyBlock(cfg *config.Config, req ApplyBlockRequest) (ApplyBlockResult, error) {
blocker := getIPBlocker()
if blocker == nil {
res := ApplyBlockResult{Outcome: firewall.BlockOutcomeNoop}
observeBlockOutcome(res.Outcome, ErrNoIPBlocker, req.Source)
return res, ErrNoIPBlocker
}
work := autoBlockQueues.acquire()
defer work.finish()
state := work.loadState(cfg.StatePath)
attemptAt := autoBlockNow()
res, err := applyBlockLocked(cfg, blocker, state, req, work.progress, func(err error) { work.directOutcome(req.IP, attemptAt, err) })
if !errors.Is(err, firewall.ErrIPProtected) {
work.observe(err)
}
work.progress()
work.saveState(cfg.StatePath, state)
work.complete()
return res, err
}
// durableActionBlocker supplies atomic admission when the engine owns durable
// action state. The existing scan policy is passed to that transaction.
type durableActionBlocker interface {
DurableActionsEnabled() bool
BlockIPRequest(firewall.ActionRequest, *firewall.ScanAdmission) (firewall.BlockOutcome, error)
}
// applyBlockLocked performs one block attempt plus the evidence bookkeeping
// a live outcome requires. The caller holds blockStateMu and owns loading
// and saving state. It writes no stderr lines for live blocks - callers
// keep their own operational logging - and emits findings instead of
// dispatching them.
func applyBlockLocked(cfg *config.Config, blocker IPBlocker, state *blockState, req ApplyBlockRequest, progress func(), observe func(error)) (ApplyBlockResult, error) {
progress()
var outcome firewall.BlockOutcome
var err error
durableOutcome := false
if durable, ok := blocker.(durableActionBlocker); ok && durable.DurableActionsEnabled() {
durableOutcome = true
var admission *firewall.ScanAdmission
if req.Source == BlockSourceScan {
limit := cfg.AutoResponse.MaxBlocksPerHour
if limit <= 0 {
limit = config.DefaultMaxBlocksPerHour
}
admission = &firewall.ScanAdmission{Window: autoBlockNow().Format("2006-01-02T15"), Limit: limit}
}
outcome, err = durable.BlockIPRequest(firewall.ActionRequest{
ID: req.ActionID, Operation: "block", Target: req.IP, Reason: req.EngineReason,
Source: req.Source, FindingID: req.FindingID, TTL: req.TTL, Actor: "daemon", Automatic: true,
}, admission)
} else {
outcome, err = callBlockIP(blocker, req.IP, req.EngineReason, req.TTL, req.FindingID)
}
observe(err)
observeBlockOutcome(outcome, err, req.Source)
res := ApplyBlockResult{Outcome: outcome}
verifiedAuditPending := durableOutcome && outcome == firewall.BlockOutcomeLive && errors.Is(err, firewall.ErrActionAuditPending)
if err != nil && !verifiedAuditPending {
return res, err
}
switch outcome {
case firewall.BlockOutcomeLive:
// nft was mutated; record the block below.
case firewall.BlockOutcomeDryRun:
// dry-run intercepted: nft was NOT mutated. Do not record a real
// block locally or in the permanent threat DB; emit a Warning
// notice instead so operators see what would have been blocked.
res.Findings = append(res.Findings, alert.Finding{
Severity: alert.Warning,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-BLOCK [dry-run]: %s would be blocked (expires in %s)", req.IP, req.TTL),
Details: fmt.Sprintf("Reason: %s", req.Reason),
Timestamp: time.Now(),
SourceIP: req.IP,
})
return res, nil
case firewall.BlockOutcomeAllowed, firewall.BlockOutcomeAllowlisted, firewall.BlockOutcomeNoop:
// Allowed: the verdict callback deliberately declined. Allowlisted:
// operator allow or verified bot. Noop: already blocked or guard
// rejected. Nothing to record for any of them.
return res, nil
default:
fmt.Fprintf(os.Stderr, "auto-block: unknown block outcome %q for %s, skipping local state\n", outcome, req.IP)
return res, nil
}
// Record in the local threat DB with the same lifetime as the firewall
// block. A permanent record here would re-flag the IP via ip_reputation
// after the temp block lapses and re-block it forever (permablock loop).
progress()
if db := GetThreatDB(); db != nil {
db.AddTemporary(req.IP, req.Reason, req.TTL)
}
state.IPs = append(state.IPs, blockedIP{
IP: req.IP,
FindingID: req.FindingID,
Reason: req.Reason,
BlockedAt: time.Now(),
ExpiresAt: time.Now().Add(req.TTL),
})
details := fmt.Sprintf("Reason: %s", req.Reason)
progress()
if cc, ok := blocker.(cloudflareCoverChecker); ok && cc.CloudflareCovers(req.IP) {
details += " (warning: " + firewall.CloudflareCoverageWarning + ")"
}
res.Findings = append(res.Findings, alert.Finding{
Severity: alert.Critical,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-BLOCK: %s blocked (expires in %s)", req.IP, req.TTL),
Details: details,
Timestamp: time.Now(),
SourceIP: req.IP,
})
// Permanent block escalation: promote after N temp blocks within the
// interval. Every source counts - a challenge-timeout or central-intel
// block is the same repeat-offender evidence as a scan block.
if cfg.AutoResponse.PermBlock {
// Load fills the default and validation rejects anything lower, so
// this only catches a Config assembled in code.
count := cfg.AutoResponse.PermBlockCount
if count < config.MinBlockEscalationCount {
count = config.DefaultPermBlockCount
}
interval := parseExpiryWithDefault(cfg.AutoResponse.PermBlockInterval, config.DefaultPermBlockInterval)
progress()
if checkPermBlockEscalation(cfg.StatePath, req.IP, count, interval) {
permReason := fmt.Sprintf("PERMBLOCK: %d temp blocks within %s", count, interval)
progress()
if promoteToPermanentBlock(blocker, req.IP, permReason, req.FindingID) {
res.Findings = append(res.Findings, alert.Finding{
Severity: alert.Critical,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-PERMBLOCK: %s promoted to permanent block (%d temp blocks)", req.IP, count),
Timestamp: time.Now(),
SourceIP: req.IP,
})
}
}
}
return res, err
}
package checks
import (
"context"
"encoding/json"
"fmt"
"net"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"gopkg.in/yaml.v3"
)
func CheckShadowChanges(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
info, err := osFS.Stat("/etc/shadow")
if err != nil {
return nil
}
mtime := info.ModTime()
mtimeKey := "_shadow_mtime"
hashKey := "_shadow_hash"
// Load current shadow entries (user:hash pairs, no sensitive data stored)
currentEntries := parseShadowUsers()
currentHash := hashBytes([]byte(fmt.Sprintf("%v", currentEntries)))
prevMtimeRaw, mtimeExists := store.GetRaw(mtimeKey)
prevHash, hashExists := store.GetRaw(hashKey)
if mtimeExists {
var lastMtime time.Time
if err := json.Unmarshal([]byte(prevMtimeRaw), &lastMtime); err == nil {
if mtime.After(lastMtime) {
// Shadow file was modified - find what changed
var details string
if hashExists && prevHash != currentHash {
changed := diffShadowChanges(store, currentEntries)
if len(changed) > 0 {
details = fmt.Sprintf("Previous: %s\nCurrent: %s\nAccounts changed: %s",
lastMtime.Format("2006-01-02 15:04:05"),
mtime.Format("2006-01-02 15:04:05"),
strings.Join(changed, ", "))
}
}
if details == "" {
details = fmt.Sprintf("Previous: %s\nCurrent: %s",
lastMtime.Format("2006-01-02 15:04:05"),
mtime.Format("2006-01-02 15:04:05"))
}
// Check if within upcp window
sev := alert.Critical
if cfg.Suppressions.UPCPWindowStart != "" {
now := time.Now()
h, m := now.Hour(), now.Minute()
nowMin := h*60 + m
start := parseTimeMin(cfg.Suppressions.UPCPWindowStart)
end := parseTimeMin(cfg.Suppressions.UPCPWindowEnd)
if nowMin >= start && nowMin <= end {
sev = alert.Warning
}
}
// Check auditd for who made the change
auditInfo := getAuditShadowInfo()
if auditInfo != "" {
details += "\n" + auditInfo
}
// Suppress alerts for password changes made by infra IPs
// (admin-initiated password resets via WHM/xml-api), but only
// when the explaining log event is newer than the shadow file's
// last-seen mtime. A stale infra event from a prior change must
// not mask a fresh, unexplained /etc/shadow edit.
if isInfraShadowChange(cfg, lastMtime) {
// Still update state, but don't alert
goto storeState
}
// Separate root password change (higher severity)
changed := diffShadowChanges(store, currentEntries)
rootChanged := false
userCount := 0
for _, c := range changed {
if c == "root" {
rootChanged = true
} else {
userCount++
}
}
if rootChanged {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "root_password_change",
Message: "Root password changed",
Details: details,
})
}
// Bulk password changes (5+ accounts at once)
if userCount >= 5 {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "bulk_password_change",
Message: fmt.Sprintf("Bulk password change: %d accounts modified", userCount),
Details: details,
})
} else {
findings = append(findings, alert.Finding{
Severity: sev,
Check: "shadow_change",
Message: "/etc/shadow modified",
Details: details,
})
}
}
}
}
storeState:
// Store current state
mtimeData, _ := json.Marshal(mtime)
store.SetRaw(mtimeKey, string(mtimeData))
store.SetRaw(hashKey, currentHash)
// Store per-user hashes for diff next time
for user, hash := range currentEntries {
store.SetRaw("_shadow_user:"+user, hash)
}
return findings
}
// parseShadowUsers reads /etc/shadow and returns a map of user -> password hash.
// Only stores a hash of the hash, not the actual password hash.
func parseShadowUsers() map[string]string {
data, err := osFS.ReadFile("/etc/shadow")
if err != nil {
return nil
}
entries := make(map[string]string)
for _, line := range strings.Split(string(data), "\n") {
parts := strings.SplitN(line, ":", 3)
if len(parts) < 2 || parts[0] == "" {
continue
}
// Store a hash of the password field, not the field itself
entries[parts[0]] = hashBytes([]byte(parts[1]))
}
return entries
}
// diffShadowChanges compares current entries against stored per-user hashes.
func diffShadowChanges(store *state.Store, current map[string]string) []string {
var changed []string
for user, hash := range current {
prev, exists := store.GetRaw("_shadow_user:" + user)
if exists && prev != hash {
changed = append(changed, user)
} else if !exists {
changed = append(changed, user+" (new)")
}
}
return changed
}
// getAuditShadowInfo checks auditd for recent shadow change events.
func getAuditShadowInfo() string {
out, err := runCmd("grep", "csm_shadow_change", "/var/log/audit/audit.log")
if err != nil || len(out) == 0 {
return ""
}
lines := strings.Split(strings.TrimSpace(string(out)), "\n")
if len(lines) == 0 {
return ""
}
// Get the last event
last := lines[len(lines)-1]
// Extract exe= field
exe := ""
for _, part := range strings.Fields(last) {
if strings.HasPrefix(part, "exe=") {
exe = strings.Trim(strings.TrimPrefix(part, "exe="), "\"")
break
}
}
// Decode hex comm if present
comm := ""
for _, part := range strings.Fields(last) {
if strings.HasPrefix(part, "comm=") {
raw := strings.Trim(strings.TrimPrefix(part, "comm="), "\"")
decoded := decodeHexString(raw)
if decoded != "" {
comm = decoded
} else {
comm = raw
}
break
}
}
if exe != "" || comm != "" {
return fmt.Sprintf("Changed by: %s (command: %s)", exe, comm)
}
return ""
}
// decodeHexString tries to decode a hex-encoded string (auditd encodes some comm fields).
func decodeHexString(s string) string {
if len(s)%2 != 0 || len(s) < 4 {
return ""
}
// Check if it looks like hex (all hex chars)
for _, c := range s {
isHex := (c >= '0' && c <= '9') || (c >= 'a' && c <= 'f') || (c >= 'A' && c <= 'F')
if !isHex {
return ""
}
}
var result []byte
for i := 0; i < len(s); i += 2 {
// #nosec G115 -- hexVal returns 0..15; (h<<4)|h fits in a byte.
b := byte(hexVal(s[i])<<4 | hexVal(s[i+1]))
if b == 0 {
break
}
result = append(result, b)
}
if len(result) == 0 {
return ""
}
return string(result)
}
func CheckUID0Accounts(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
data, err := osFS.ReadFile("/etc/passwd")
if err != nil {
return nil
}
for _, line := range strings.Split(string(data), "\n") {
user, unauthorized := classifyUID0Line(line)
if unauthorized {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "uid0_account",
Message: fmt.Sprintf("Unauthorized UID 0 account: %s", user),
Details: line,
})
}
}
return findings
}
func CheckSSHKeys(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
// Check root authorized_keys
rootKeys := "/root/.ssh/authorized_keys"
if hash, err := hashFileContent(rootKeys); err == nil {
key := "_ssh_root_keys_hash"
prev, exists := store.GetRaw(key)
if exists && prev != hash {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "ssh_keys",
Message: "Root authorized_keys modified",
Details: fmt.Sprintf("File: %s", rootKeys),
})
}
store.SetRaw(key, hash)
}
// Check for new authorized_keys in /home. Rank by mtime desc so
// recently-touched accounts are processed first when the check timeout
// cuts iteration short.
homes, _ := accountHomeGlob("*/.ssh/authorized_keys")
for _, keyFile := range rankPathsByMtimeDesc(ctx, homes, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
break
}
hash, err := hashFileContent(keyFile)
if err != nil {
continue
}
key := fmt.Sprintf("_ssh_user_keys:%s", keyFile)
prev, exists := store.GetRaw(key)
if exists && prev != hash {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "ssh_keys",
Message: fmt.Sprintf("User authorized_keys modified: %s", keyFile),
})
}
store.SetRaw(key, hash)
}
return findings
}
func CheckAPITokens(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
// WHM root API tokens
if f, fire := checkWHMRootAPITokens(store); fire {
findings = append(findings, f)
}
// User API tokens - read directly from disk instead of spawning uapi per user.
// Token files are JSON at /home/<user>/.cpanel/api_tokens/<token_name>.
// Rank account dirs by mtime desc so recently touched accounts are processed
// first when the check timeout cuts iteration short.
tokenDirs, _ := accountHomeGlob("*/.cpanel/api_tokens")
for _, tokenDir := range rankPathsByMtimeDesc(ctx, tokenDirs, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
break
}
user := filepath.Base(filepath.Dir(filepath.Dir(tokenDir)))
tokenFiles, _ := osFS.Glob(filepath.Join(tokenDir, "*"))
for _, tokenFile := range tokenFiles {
if ctx.Err() != nil {
return findings
}
tokenName := filepath.Base(tokenFile)
data, err := osFS.ReadFile(tokenFile)
if err != nil {
continue
}
content := string(data)
// Check for full access with no IP whitelist
hasFullAccess := strings.Contains(content, `"has_full_access":1`) ||
strings.Contains(content, `"has_full_access": 1`)
noWhitelist := strings.Contains(content, `"whitelist_ips":null`) ||
strings.Contains(content, `"whitelist_ips": null`) ||
strings.Contains(content, `"whitelist_ips":[]`) ||
strings.Contains(content, `"whitelist_ips": []`) ||
!strings.Contains(content, "whitelist_ips")
if hasFullAccess && noWhitelist {
known := false
for _, t := range cfg.Suppressions.KnownAPITokens {
if tokenName == t {
known = true
break
}
}
if !known {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "api_tokens",
Message: fmt.Sprintf("User %s has full-access API token '%s' with no IP whitelist", user, tokenName),
Details: fmt.Sprintf("File: %s", tokenFile),
})
}
}
}
}
return findings
}
const (
whmAPITokensStateKey = "_whm_api_tokens_state" // #nosec G101 -- bbolt state-store key name, not a credential
whmAPITokensHashKey = "_whm_api_tokens_hash" // #nosec G101 -- bbolt state-store key name, not a credential
)
func checkWHMRootAPITokens(store *state.Store) (alert.Finding, bool) {
out, err := runCmd("whmapi1", "api_token_list", "--output=json")
if err == nil {
if cur, ok := parseWHMTokenSig(out); ok {
return checkStructuredWHMRootAPITokens(store, cur, nil)
}
}
return checkWHMRootAPITokensLegacyHash(store)
}
func checkWHMRootAPITokensLegacyHash(store *state.Store) (alert.Finding, bool) {
out, err := runCmd("whmapi1", "api_token_list")
if err != nil {
return alert.Finding{}, false
}
if cur, ok := parseWHMTokenSigYAML(out); ok {
return checkStructuredWHMRootAPITokens(store, cur, out)
}
return checkWHMRootAPITokensLegacyHashOutput(store, out)
}
func checkWHMRootAPITokensLegacyHashOnly(store *state.Store) (alert.Finding, bool) {
out, err := runCmd("whmapi1", "api_token_list")
if err != nil {
return alert.Finding{}, false
}
return checkWHMRootAPITokensLegacyHashOutput(store, out)
}
func checkStructuredWHMRootAPITokens(store *state.Store, cur tokenSig, legacyOut []byte) (alert.Finding, bool) {
_, hadStructuredState := store.GetRaw(whmAPITokensStateKey)
_, hadLegacyState := store.GetRaw(whmAPITokensHashKey)
finding, fire := diffWHMTokens(store, cur)
switch {
case !hadStructuredState && hadLegacyState:
// Migrating from the legacy hash: the tokens were already vetted under
// the old scheme, so diff against the legacy hash rather than
// re-flagging the whole set as new.
var legacyFinding alert.Finding
var legacyFire bool
if legacyOut != nil {
legacyFinding, legacyFire = checkWHMRootAPITokensLegacyHashOutput(store, legacyOut)
} else {
legacyFinding, legacyFire = checkWHMRootAPITokensLegacyHashOnly(store)
}
if legacyFire {
finding, fire = legacyFinding, true
}
case !hadStructuredState && !hadLegacyState:
// True first run with no prior state of any kind. A root API token
// created by an attacker before CSM was installed would otherwise be
// baselined as "known" and never alert. Surface the pre-existing
// operator tokens for review.
finding, fire = baselineWHMTokens(cur)
}
store.SetRaw(whmAPITokensStateKey, marshalTokenSig(cur))
if legacyOut != nil {
store.SetRaw(whmAPITokensHashKey, hashBytes(legacyOut))
}
return finding, fire
}
func checkWHMRootAPITokensLegacyHashOutput(store *state.Store, out []byte) (alert.Finding, bool) {
hash := hashBytes(out)
prev, exists := store.GetRaw(whmAPITokensHashKey)
store.SetRaw(whmAPITokensHashKey, hash)
if exists && prev != hash {
return alert.Finding{
Severity: alert.Critical,
Check: "api_tokens",
Message: "WHM root API tokens changed",
Details: "Run 'whmapi1 api_token_list' to review",
}, true
}
return alert.Finding{}, false
}
// baselineWHMTokens reports pre-existing root API tokens on the very first
// scan. Routine cluster-managed trust tokens churn on their own and are
// expected, so a baseline made up only of those stays silent; any
// operator/full-access token present at baseline is surfaced for review so a
// token planted before CSM was installed cannot pass as "known".
func baselineWHMTokens(cur tokenSig) (alert.Finding, bool) {
var operator []string
for name, info := range cur {
if isClusterManagedToken(info) || isClusterManagedTokenAddition(name, info) {
continue
}
operator = append(operator, name)
}
if len(operator) == 0 {
return alert.Finding{}, false
}
sort.Strings(operator)
return alert.Finding{
Severity: alert.High,
Check: "api_tokens",
Message: "Pre-existing WHM root API tokens present at baseline",
Details: "Review with 'whmapi1 api_token_list': " + strings.Join(operator, "; "),
}, true
}
// tokenSig maps each WHM root API token name to the security traits compared
// between scans.
type tokenSig map[string]tokenInfo
type tokenInfo struct {
FullAccess bool `json:"full_access"`
ClusterManaged bool `json:"cluster_managed"`
}
// isClusterManagedToken reports whether cPanel owns the token's lifecycle.
// DNS clustering creates, rotates, and deletes these on its own:
// - reverse_trust_<uuid>: trust granted to a remote WHM peer
// - <host>-trust (e.g. ns2-trust): local end of a trust relationship
//
// Their churn is routine and must not page like an attacker-created token.
func isClusterManagedToken(info tokenInfo) bool {
return info.ClusterManaged && !info.FullAccess
}
func isClusterManagedTokenAddition(name string, info tokenInfo) bool {
return strings.HasPrefix(name, "reverse_trust_") && isClusterManagedToken(info)
}
func clusterManagedFromACLs(name string, acls map[string]bool) bool {
if strings.HasPrefix(name, "reverse_trust_") {
return true
}
return strings.HasSuffix(name, "-trust") && acls["clustering"]
}
// parseWHMTokenSig decodes `whmapi1 api_token_list --output=json` into a
// tokenSig. A token whose value is not a JSON object or whose ACLs cannot be
// read is kept with FullAccess=false, so an unparsable entry never hides a
// token's presence. ok=false means the caller should use the legacy hash path.
func parseWHMTokenSig(out []byte) (tokenSig, bool) {
var env struct {
Data struct {
Tokens map[string]json.RawMessage `json:"tokens"`
} `json:"data"`
}
if err := json.Unmarshal(out, &env); err != nil || env.Data.Tokens == nil {
return nil, false
}
sig := make(tokenSig, len(env.Data.Tokens))
for name, raw := range env.Data.Tokens {
var t struct {
ACLs map[string]json.RawMessage `json:"acls"`
}
_ = json.Unmarshal(raw, &t)
acls := decodeWHMTokenACLs(t.ACLs)
sig[name] = tokenInfo{
FullAccess: acls["all"],
ClusterManaged: clusterManagedFromACLs(name, acls),
}
}
return sig, true
}
func parseWHMTokenSigYAML(out []byte) (tokenSig, bool) {
type yamlToken struct {
ACLs map[string]any `yaml:"acls"`
}
type yamlTokenData struct {
Tokens map[string]yamlToken `yaml:"tokens"`
}
var env struct {
Data yamlTokenData `yaml:"data"`
Result struct {
Data yamlTokenData `yaml:"data"`
} `yaml:"result"`
}
if err := yaml.Unmarshal(out, &env); err != nil {
return nil, false
}
tokens := env.Data.Tokens
if tokens == nil {
tokens = env.Result.Data.Tokens
}
if tokens == nil {
return nil, false
}
sig := make(tokenSig, len(tokens))
for name, raw := range tokens {
acls := decodeWHMTokenYAMLACLs(raw.ACLs)
sig[name] = tokenInfo{
FullAccess: acls["all"],
ClusterManaged: clusterManagedFromACLs(name, acls),
}
}
return sig, true
}
func decodeWHMTokenACLs(raw map[string]json.RawMessage) map[string]bool {
acls := make(map[string]bool, len(raw))
for name, value := range raw {
acls[name] = decodeWHMTokenACL(value)
}
return acls
}
func decodeWHMTokenYAMLACLs(raw map[string]any) map[string]bool {
acls := make(map[string]bool, len(raw))
for name, value := range raw {
acls[name] = decodeWHMTokenYAMLACL(value)
}
return acls
}
func decodeWHMTokenACL(raw json.RawMessage) bool {
switch strings.TrimSpace(string(raw)) {
case "1", "true", `"1"`, `"true"`:
return true
case "0", "false", `"0"`, `"false"`, "null", "":
return false
}
var n json.Number
if err := json.Unmarshal(raw, &n); err == nil {
return n.String() == "1"
}
var s string
if err := json.Unmarshal(raw, &s); err == nil {
return s == "1" || strings.EqualFold(s, "true")
}
var b bool
return json.Unmarshal(raw, &b) == nil && b
}
func decodeWHMTokenYAMLACL(raw any) bool {
switch v := raw.(type) {
case bool:
return v
case int:
return v == 1
case int64:
return v == 1
case uint64:
return v == 1
case float64:
return v == 1
case string:
s := strings.TrimSpace(v)
return s == "1" || strings.EqualFold(s, "true")
default:
return false
}
}
// marshalTokenSig serializes a tokenSig deterministically (encoding/json sorts
// map keys), so an unchanged set always produces an identical stored string.
func marshalTokenSig(sig tokenSig) string {
b, _ := json.Marshal(sig)
return string(b)
}
func unmarshalTokenSig(raw string) (tokenSig, bool) {
var sig tokenSig
if err := json.Unmarshal([]byte(raw), &sig); err == nil {
return sig, true
}
var legacy map[string]bool
if err := json.Unmarshal([]byte(raw), &legacy); err != nil {
return nil, false
}
sig = make(tokenSig, len(legacy))
for name, all := range legacy {
sig[name] = tokenInfo{
FullAccess: all,
ClusterManaged: !all && (strings.HasPrefix(name, "reverse_trust_") || strings.HasSuffix(name, "-trust")),
}
}
return sig, true
}
// diffWHMTokens compares the current token set against the previously stored
// one and returns a single finding when something changed. Severity splits on
// intent:
// - Critical: any token added outside generated reverse_trust churn, any
// non-cluster token removed, or ANY token gaining the full-access "all" ACL.
// - Warning: only generated reverse_trust additions/removals or recorded
// cluster trust removals changed. cPanel does this during normal DNS
// clustering and it must not page.
//
// The first scan after the key is introduced just records a baseline.
func diffWHMTokens(store *state.Store, cur tokenSig) (alert.Finding, bool) {
raw, exists := store.GetRaw(whmAPITokensStateKey)
if !exists {
return alert.Finding{}, false
}
prev, ok := unmarshalTokenSig(raw)
if !ok {
return alert.Finding{}, false
}
var critical, cluster []string
for name, info := range cur {
prevInfo, had := prev[name]
if !had {
if isClusterManagedTokenAddition(name, info) {
cluster = append(cluster, "added "+name)
} else {
critical = append(critical, "added "+name)
}
continue
}
if info.FullAccess && !prevInfo.FullAccess {
critical = append(critical, "escalated "+name+" to full access")
}
}
for name, info := range prev {
if _, still := cur[name]; still {
continue
}
if isClusterManagedToken(info) {
cluster = append(cluster, "removed "+name)
} else {
critical = append(critical, "removed "+name)
}
}
switch {
case len(critical) > 0:
sort.Strings(critical)
return alert.Finding{
Severity: alert.Critical,
Check: "api_tokens",
Message: "WHM root API tokens changed",
Details: "Review with 'whmapi1 api_token_list': " + strings.Join(critical, "; "),
}, true
case len(cluster) > 0:
sort.Strings(cluster)
return alert.Finding{
Severity: alert.Warning,
Check: "api_tokens",
Message: "WHM root cluster trust tokens changed",
Details: "cPanel DNS clustering churn (expected): " + strings.Join(cluster, "; "),
}, true
}
return alert.Finding{}, false
}
// shadowMutatingWHMEndpoints lists WHM JSON-API endpoints whose handlers
// rewrite /etc/shadow as a side effect:
// - suspendacct/unsuspendacct: lock/unlock password field, swap login shell
// - passwd/forcepasswordchange: set or expire user password
// - createacct/removeacct/killacct: add or remove the shadow entry entirely
//
// Hits on these endpoints from infra IPs explain shadow mtime changes
// without involving an attacker; hits from non-infra IPs do not.
var shadowMutatingWHMEndpoints = []string{
"/json-api/suspendacct",
"/json-api/unsuspendacct",
"/json-api/passwd",
"/json-api/forcepasswordchange",
"/json-api/createacct",
"/json-api/removeacct",
"/json-api/killacct",
}
// isInfraShadowChange reports whether every recent log signal that could
// explain a /etc/shadow modification was originated by an infra IP. It
// fuses two sources:
//
// 1. session_log PURGE password_change events (WHM/cPanel sets a new
// password, which goes through the session machinery).
// 2. successful api_tokens_log entries for shadow-mutating WHM JSON-API endpoints
// (suspendacct, passwd, createacct, ...). The cPanel session log does
// NOT record these because they are not session events, so the older
// session-only check fired on every internal `suspendacct` call.
//
// Returns true only if at least one such event was seen AND every event in
// both sources came from an infra IP (or loopback / "internal"). Any
// successful external API call or unparseable source short-circuits to false
// so a stolen token or compromised neighbour does not get a free suppression.
// since is the shadow file's previously recorded mtime. Only log events newer
// than since can explain the modification under investigation; stale tail lines
// (a legit infra password change from days ago) must not suppress a fresh,
// unrelated /etc/shadow edit.
func isInfraShadowChange(cfg *config.Config, since time.Time) bool {
sessFound, sessAllInfra := scanSessionLogShadow(cfg, since)
tokFound, tokAllInfra := scanAPITokensLogShadow(cfg, since)
return (sessFound || tokFound) && sessAllInfra && tokAllInfra
}
// parseCPanelLogTime reads the leading "[<timestamp>]" of a cPanel session_log
// or api_tokens_log line. Both the timezone-qualified form
// ("2006-01-02 15:04:05 -0700") and the older bare form ("2006-01-02 15:04:05",
// interpreted in the host's local zone) are accepted. Returns ok=false when no
// bracketed timestamp is present or it does not parse.
func parseCPanelLogTime(line string) (time.Time, bool) {
open := strings.IndexByte(line, '[')
if open < 0 {
return time.Time{}, false
}
end := strings.IndexByte(line[open:], ']')
if end < 0 {
return time.Time{}, false
}
stamp := strings.TrimSpace(line[open+1 : open+end])
for _, layout := range []string{"2006-01-02 15:04:05 -0700", "2006-01-02 15:04:05"} {
if t, err := time.ParseInLocation(layout, stamp, time.Local); err == nil {
return t, true
}
}
return time.Time{}, false
}
func cpanelShadowLogAfterSince(ts, since time.Time) bool {
if since.IsZero() {
return true
}
// cPanel logs only whole seconds, while shadow mtimes can carry
// subsecond precision. Treat the stored mtime's logged second as
// in-window so an infra event at 10:00:00.900 is not rejected just
// because the prior mtime was 10:00:00.500.
return !ts.Before(since.Truncate(time.Second))
}
// scanSessionLogShadow walks the cPanel session log for PURGE password_change
// events and reports whether any were seen and whether every non-loopback,
// non-"internal" source IP belonged to the infra allowlist.
func scanSessionLogShadow(cfg *config.Config, since time.Time) (foundAny, allInfra bool) {
allInfra = true
lines := tailFile("/usr/local/cpanel/logs/session_log", 100)
for i := len(lines) - 1; i >= 0; i-- {
line := lines[i]
if !strings.Contains(line, "PURGE") || !strings.Contains(line, "password_change") {
continue
}
ts, ok := parseCPanelLogTime(line)
if !ok {
// A shadow-mutating line without a usable timestamp cannot prove
// it is stale. Fail toward alerting instead of suppressing.
foundAny = true
allInfra = false
return
}
if !cpanelShadowLogAfterSince(ts, since) {
// Stale line cannot explain this modification.
continue
}
foundAny = true
// Format: [ts] info [xml-api|whostmgr|security] IP PURGE account:token password_change
var ip string
for _, tag := range []string{"[xml-api]", "[whostmgr]", "[security]"} {
if idx := strings.Index(line, tag); idx >= 0 {
rest := strings.TrimSpace(line[idx+len(tag):])
fields := strings.Fields(rest)
if len(fields) > 0 {
ip = fields[0]
}
break
}
}
if ip == "internal" {
continue
}
if !isTrustedShadowSource(ip, cfg) {
allInfra = false
return
}
}
return
}
// scanAPITokensLogShadow walks the WHM api_tokens_log for recent calls to
// JSON-API endpoints that rewrite /etc/shadow. Returns (foundAny, allInfra)
// with the same semantics as scanSessionLogShadow.
func scanAPITokensLogShadow(cfg *config.Config, since time.Time) (foundAny, allInfra bool) {
allInfra = true
lines := tailFile("/usr/local/cpanel/logs/api_tokens_log", 200)
for i := len(lines) - 1; i >= 0; i-- {
line := lines[i]
if !lineHitsShadowEndpoint(line) {
continue
}
if !apiTokensHTTPStatusOK(line) {
continue
}
ts, ok := parseCPanelLogTime(line)
if !ok {
// A shadow-mutating line without a usable timestamp cannot prove
// it is stale. Fail toward alerting instead of suppressing.
foundAny = true
allInfra = false
return
}
if !cpanelShadowLogAfterSince(ts, since) {
// Stale line cannot explain this modification.
continue
}
foundAny = true
ip := extractAPITokensHost(line)
if ip == "internal" {
continue
}
if !isTrustedShadowSource(ip, cfg) {
allInfra = false
return
}
}
return
}
func lineHitsShadowEndpoint(line string) bool {
path := extractAPITokensRequestPath(line)
if path == "" {
return false
}
for _, ep := range shadowMutatingWHMEndpoints {
if path == ep {
return true
}
}
return false
}
func apiTokensHTTPStatusOK(line string) bool {
status := extractAPITokensField(line, "HTTP Status: ['")
return strings.HasPrefix(status, "2")
}
// extractAPITokensHost pulls the source IP out of the api_tokens_log line
// shape used by whostmgrd:
//
// [ts] info [whostmgrd] Host: ['<ip>'] HTTP Status: [...], ...
func extractAPITokensHost(line string) string {
return extractAPITokensField(line, "Host: ['")
}
func extractAPITokensRequestPath(line string) string {
request := extractAPITokensField(line, "Request: ['")
fields := strings.Fields(request)
if len(fields) < 2 {
return ""
}
path := fields[1]
if idx := strings.IndexByte(path, '?'); idx >= 0 {
path = path[:idx]
}
return path
}
func extractAPITokensField(line, marker string) string {
idx := strings.Index(line, marker)
if idx < 0 {
return ""
}
rest := line[idx+len(marker):]
end := strings.Index(rest, "']")
if end < 0 {
return ""
}
return rest[:end]
}
func isTrustedShadowSource(ip string, cfg *config.Config) bool {
parsed := net.ParseIP(ip)
if parsed != nil && parsed.IsLoopback() {
return true
}
return isInfraIP(ip, cfg.InfraIPs)
}
func parseTimeMin(s string) int {
parts := strings.Split(s, ":")
if len(parts) != 2 {
return 0
}
h := 0
m := 0
fmt.Sscanf(parts[0], "%d", &h)
fmt.Sscanf(parts[1], "%d", &m)
return h*60 + m
}
package checks
import (
"crypto/rand"
"encoding/json"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/mailranges"
"github.com/pidginhost/csm/internal/netutil"
"github.com/pidginhost/csm/internal/store"
)
const (
blockStateFile = "blocked_ips.json"
// maxPendingBlocks bounds the retry queue. Under a sustained flood or
// firewall outage the daemon can accumulate more distinct attacker IPs
// than it can block; without a bound the queue grows without limit and
// bloats blocked_ips.json. Overflow is named in stderr and surfaced as
// a warning so the resulting loss is visible to operators.
maxPendingBlocks = 1000
// maxPendingAge drops queued entries whose evidence has gone stale: a
// pending IP survives requeue cycles (rate limit, engine down, block
// errors), and blocking hours after the triggering findings is worse
// than not blocking. Two hours covers one full rate-limit window plus
// slack for the queue to drain.
maxPendingAge = 2 * time.Hour
)
// IPBlocker abstracts the firewall engine for auto-blocking.
// When set, blocks go through nftables firewall engine.
type IPBlocker interface {
BlockIP(ip string, reason string, timeout time.Duration) error
UnblockIP(ip string) error
IsBlocked(ip string) bool
}
// outcomeBlocker is satisfied by engines that report what they actually
// did. When the wired IPBlocker supports it, the auto-block path uses the
// outcome to decide whether to apply local side effects: a dry-run or
// verdict-allowed call must not mutate blocked_ips.json, must not bump the
// hourly counter, and must not emit the operator-facing "AUTO-BLOCK"
// finding (which would falsely claim a real block landed). The plain
// IPBlocker interface stays as a back-compat fallback for tests and any
// legacy implementation.
type outcomeBlocker interface {
BlockIPOutcome(ip, reason string, timeout time.Duration) (firewall.BlockOutcome, error)
}
// liveBlocker is satisfied by engines that can query the live kernel firewall
// state, not just an in-memory cache built from state.json. The tracker
// reconcile loop prefers this because the cache can drift when nft
// auto-expires entries faster than CSM rewrites state.json, or when an
// out-of-band flush dropped entries the cache still claims are live. Falls
// back to IPBlocker.IsBlocked when the live query is unavailable.
type liveBlocker interface {
IsBlockedLive(ip string) (bool, error)
}
// liveBlockedLister is satisfied by engines that can dump their live blocked
// sets once per cycle. The reconcile pass prefers it over liveBlocker: the
// per-IP query dumps an entire family set on every call, so pruning N tracked
// IPs cost N full dumps of the same data.
type liveBlockedLister interface {
LiveBlockedSet() (firewall.LiveBlockedSnapshot, error)
}
type subnetBlocker interface {
BlockSubnet(cidr string, reason string, timeout time.Duration) error
}
// subnetBlockValidator is satisfied by firewall engines that can run their
// subnet safety and capability checks without changing firewall state. Dry-run
// notices use it so they only describe blocks the live path could attempt.
type subnetBlockValidator interface {
ValidateSubnetBlock(cidr string) error
}
// cloudflareCoverChecker is satisfied by engines that can report whether an
// IP falls inside the Cloudflare allow ranges. The input chain accepts
// Cloudflare edges on TCP 80/443 before the blocked drop, so a block of a
// covered IP does not stop its web traffic; findings carry that caveat.
type cloudflareCoverChecker interface {
CloudflareCovers(ip string) bool
}
// allowChecker is satisfied by engines that can report whether an IP is
// firewall-allowed (whitelisted). http_asn_crawl uses it at emit time to drop
// any candidate subnet that contains an observed source IP which is already
// allowed, so a surgical subnet tempban can never include a whitelisted host.
// Optional: blockers that do not implement it leave all candidate CIDRs intact.
type allowChecker interface {
IsAllowed(ip string) bool
}
// permanentPromoter is satisfied by engines that can upgrade an existing
// temporary block to a permanent one. PermBlock escalation runs in the same
// scan cycle as the temp block that triggered it, so the ordinary block path
// (which skips an already-blocked IP and returns BlockOutcomeNoop) would never
// clear the kernel timeout and the "permanent" block would silently expire.
type permanentPromoter interface {
PromoteToPermanentBlock(ip, reason string) error
}
type subnetBlockStatus interface {
IsSubnetBlocked(cidr string) bool
}
// subnetManager is satisfied by firewall engines that expose their blocked
// subnet state and support targeted unblock calls. Used by
// PruneExemptAutoSubnets to enumerate and remove stale subnet blocks whose
// CIDR has become DoS-exempt.
type subnetManager interface {
BlockedSubnets() []firewall.SubnetEntry
UnblockSubnet(cidr string) error
}
// fwBlockerSlot wraps an IPBlocker so atomic.Pointer can store it. The
// extra struct layer is required because atomic.Pointer needs a
// concrete type and interfaces cannot be stored directly.
type fwBlockerSlot struct{ b IPBlocker }
var fwBlockerHolder atomic.Pointer[fwBlockerSlot]
var blockStateMu sync.Mutex
var autoBlockNow = time.Now
// SetIPBlocker installs the firewall engine for auto-blocking. Safe to
// call concurrently with AutoBlockIPs: each call publishes the new
// blocker atomically and any in-flight scan keeps the snapshot it
// already loaded.
func SetIPBlocker(b IPBlocker) {
fwBlockerHolder.Store(&fwBlockerSlot{b: b})
}
// getIPBlocker returns the current blocker via a single atomic load.
// Callers should capture the result into a local variable and reuse it
// for the duration of one operation so a concurrent SetIPBlocker
// cannot split a single scan across two different engines.
func getIPBlocker() IPBlocker {
slot := fwBlockerHolder.Load()
if slot == nil {
return nil
}
return slot.b
}
type blockedIP struct {
FindingID string `json:"finding_id,omitempty"`
IP string `json:"ip"`
Reason string `json:"reason"`
BlockedAt time.Time `json:"blocked_at"`
ExpiresAt time.Time `json:"expires_at"`
}
type pendingIP struct {
ActionID string `json:"action_id,omitempty"`
ActionTTL time.Duration `json:"action_ttl,omitempty"`
FindingID string `json:"finding_id,omitempty"`
IP string `json:"ip"`
Reason string `json:"reason"`
Check string `json:"check,omitempty"`
Severity alert.Severity `json:"severity,omitempty"`
// QueuedAt is when the IP first entered the queue; it survives
// requeue cycles so age accumulates instead of resetting. Stamped on
// the first requeue for eligible entries with no timestamp.
QueuedAt time.Time `json:"queued_at,omitempty"`
queueRecord *autoBlockPendingRecord
queueCandidate *autoBlockCandidate
}
type blockState struct {
IPs []blockedIP `json:"ips"`
Pending []pendingIP `json:"pending,omitempty"` // IPs waiting for another block attempt
CleanupPending []string `json:"cleanup_pending,omitempty"`
BlocksThisHour int `json:"blocks_this_hour"`
HourKey string `json:"hour_key"`
// RateLimitWarnedHour is the HourKey for which the rate-limit warning
// was already emitted. The warning reflects a steady-state condition,
// not a per-IP event, so it fires once per hour window instead of on
// every scan tick -- the per-tick emission flooded the audit log with
// one identical finding every few seconds during a sustained attack.
RateLimitWarnedHour string `json:"rate_limit_warned_hour,omitempty"`
// PendingDropWarnedHour throttles queue-overflow findings independently
// from rate-limit warnings. Engine-down and block-error retries can fill
// the queue without reaching the hourly block limit.
PendingDropWarnedHour string `json:"pending_drop_warned_hour,omitempty"`
}
// blockableCheck reports whether a finding's check may drive a firewall block.
// The policy is carried by the check registry; see response_policy.go.
func blockableCheck(check string, blockCpanelLogins bool) bool {
switch ResponsePolicyFor(check).Block {
case BlockAlways:
return true
case BlockWithCpanelLogins:
return blockCpanelLogins
default:
return false
}
}
func blockableFinding(f alert.Finding, blockCpanelLogins bool) bool {
return blockableCheck(f.Check, blockCpanelLogins) &&
(!ResponsePolicyFor(f.Check).CriticalOnly || f.Severity == alert.Critical)
}
// AutoBlockIPs processes all findings, including repeats, for IP blocking.
func AutoBlockIPs(cfg *config.Config, findings []alert.Finding) []alert.Finding {
return autoBlockIPs(cfg, findings, "")
}
// autoBlockIPs retains the observed source when database response converts one
// finding into several session-IP candidates for the existing block policy.
func autoBlockIPs(cfg *config.Config, findings []alert.Finding, sourceFindingID string) []alert.Finding {
if !cfg.AutoResponse.Enabled || !cfg.AutoResponse.BlockIPs {
return nil
}
work := autoBlockQueues.acquire()
defer work.finish()
// Snapshot the wired firewall engine ONCE per call. A concurrent
// SetIPBlocker (SIGHUP re-wire, test cleanup) can swap the global
// mid-scan; reading the atomic pointer once and reusing the
// returned value keeps every block decision in this batch routed
// to the same engine. The previous unsynchronized read of the
// global also tripped the race detector.
blocker := getIPBlocker()
// One bulk dump of the kernel's blocked sets answers every membership
// question this cycle asks. The per-IP live query dumps the whole set, so
// the reconcile pass below cost one full dump per tracked IP.
var liveBlocked firewall.LiveBlockedSnapshot
var useLiveBlocked bool
if blocker != nil {
work.progress()
liveBlocked, useLiveBlocked = liveBlockedSnapshot(blocker)
}
// exemptLogged deduplicates per-cycle log lines for DoS-exempt CIDR skips
// so each suppressed subnet is logged once per AutoBlockIPs call, not once
// per finding or per IP in the netblock counting sweep.
exemptLogged := make(map[string]struct{})
var actions []alert.Finding
answeredSubnets := make(map[string]bool)
// Load block state
work.progress()
state := work.loadState(cfg.StatePath)
// Prune IPs that the firewall engine no longer has blocked.
// The engine handles expiry natively via nftables timeouts -
// we just sync our state to match. Use the live nftables query
// when the engine supports it so the tracker stays in lock-step
// with the kernel; the in-memory cache (IsBlocked) can lag when
// the kernel expires entries before state.json is rewritten.
var stillBlocked []blockedIP
for _, b := range state.IPs {
work.progress()
if blocker != nil {
if !blockedLiveOrCached(blocker, liveBlocked, useLiveBlocked, b.IP) {
// Engine expired this block - clean up our state
fmt.Fprintf(os.Stderr, "[%s] AUTO-UNBLOCK: %s removed (engine expired)\n", time.Now().Format("2006-01-02 15:04:05"), b.IP)
continue
}
}
stillBlocked = append(stillBlocked, b)
}
state.IPs = stillBlocked
// Prune auto-response subnet blocks that now intersect the DoS-exempt set
// before making new subnet decisions this cycle. Never in dry-run:
// dry-run promises a read-only firewall and pruning is a kernel mutation.
if blocker != nil && isAutoResponseActive(cfg) {
pruneExemptAutoSubnets(cfg, blocker, work.progress)
}
// Check rate limit
currentHour := autoBlockNow().Format("2006-01-02T15")
if state.HourKey != currentHour {
state.HourKey = currentHour
state.BlocksThisHour = 0
}
// Collect IPs to block from findings
ipsToBlock := make(map[string]pendingIP)
// Drain pending queue first (IPs from prior rate-limited or failed
// cycles). Stale entries are dropped by name so the audit trail shows
// exactly which attackers aged out instead of being blocked.
for _, p := range state.Pending {
work.beginPending(p)
ip := normalizeBlockIP(p.IP)
if ip == "" {
fmt.Fprintf(os.Stderr, "auto-block: dropping invalid pending IP %q\n", p.IP)
continue
}
p.IP = ip
// Retry only evidence still eligible under the current policy. Older
// queues lack check identity; a free-text reason cannot establish it.
if !blockableFinding(alert.Finding{Check: p.Check, Severity: p.Severity}, cfg.AutoResponse.BlockCpanelLogins) {
work.completePending(p)
fmt.Fprintf(os.Stderr, "auto-block: dropping ineligible pending %s (check %q)\n", p.IP, p.Check)
continue
}
if !p.QueuedAt.IsZero() && autoBlockNow().Sub(p.QueuedAt) > maxPendingAge {
fmt.Fprintf(os.Stderr, "auto-block: dropping stale pending %s (queued %s)\n",
p.IP, p.QueuedAt.Format(time.RFC3339))
continue
}
if !isAlreadyBlocked(state, p.IP) {
p.queueCandidate = work.candidate(p, ipsToBlock[p.IP].queueCandidate)
ipsToBlock[p.IP] = p
} else {
work.completePending(p)
}
}
state.Pending = nil
// Subnet fast-path: checks that represent a subnet directly.
// Independent of the per-IP rate limit, because a single subnet block
// replaces what would otherwise be hundreds of per-IP blocks.
for _, f := range findings {
work.progress()
if f.Check != "smtp_subnet_spray" && f.Check != "mail_subnet_spray" {
continue
}
cidr := extractCIDRFromFinding(f)
if cidr == "" {
continue
}
if isSubnetAlreadyBlocked(blocker, cidr) {
continue
}
if cidrIntersectsInfra(cfg, cidr) {
continue
}
if shouldSkipAutoSubnet(cfg, cidr, exemptLogged) {
continue
}
if !isAutoResponseActive(cfg) {
if !canDryRunBlockSubnet(blocker, cidr) {
continue
}
// Dry-run: same visibility contract as per-IP blocks - emit a
// Warning notice instead of skipping silently.
actions = append(actions, dryRunSubnetNotice(cidr, "", f.Message))
continue
}
if blocker == nil {
fmt.Fprintf(os.Stderr, "auto-block: firewall engine not available, skipping subnet %s\n", cidr)
continue
}
sb, ok := blocker.(subnetBlocker)
if !ok {
fmt.Fprintf(os.Stderr, "auto-block: firewall engine does not support subnet blocking, skipping %s\n", cidr)
continue
}
reason := fmt.Sprintf("CSM auto-block (subnet): %s", truncate(f.Message, 100))
if !autoFirewallActionApplied(cidr, callBlockSubnet(sb, cidr, reason, parseExpiry(cfg.AutoResponse.BlockExpiry), alert.FindingID(f))) {
continue
}
answeredSubnets[cidr] = true
fmt.Fprintf(os.Stderr, "[%s] AUTO-BLOCK-SUBNET: %s blocked\n", time.Now().Format("2006-01-02 15:04:05"), cidr)
actions = append(actions, alert.Finding{
Severity: alert.Critical,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-BLOCK-SUBNET: %s blocked", cidr),
Details: fmt.Sprintf("Reason: %s", f.Message),
Timestamp: time.Now(),
})
}
for _, f := range findings {
work.progress()
if !blockableFinding(f, cfg.AutoResponse.BlockCpanelLogins) {
continue
}
ip := extractIPFromFinding(f)
if ip == "" {
continue
}
// Never block infra IPs
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
continue
}
// Don't re-block already blocked IPs.
if isAlreadyBlocked(state, ip) || (blocker != nil && blockedLiveOrCached(blocker, liveBlocked, useLiveBlocked, ip)) {
continue
}
// Skip IPs that are already being challenged, but do not let a
// prior challenge suppress a later hard-block-only finding.
if cl := GetChallengeIPList(); cl != nil && cl.Contains(ip) && shouldSkipAutoBlockForChallenge(cfg, f) {
continue
}
// A drained pending entry keeps its QueuedAt when the same IP
// recurs in fresh findings; the check, severity and reason are refreshed.
findingID := sourceFindingID
if findingID == "" {
findingID = alert.FindingID(f)
}
if existing, ok := ipsToBlock[ip]; ok {
if existing.ActionID == "" {
existing.Reason = f.Message
existing.Check = f.Check
existing.Severity = f.Severity
existing.FindingID = findingID
}
ipsToBlock[ip] = existing
} else {
p := pendingIP{IP: ip, Reason: f.Message, Check: f.Check, Severity: f.Severity, FindingID: findingID}
p.queueCandidate = work.candidate(p, nil)
ipsToBlock[ip] = p
}
}
// Block IPs and queue candidates that cannot be attempted or completed.
expiry := parseExpiry(cfg.AutoResponse.BlockExpiry)
maxPerHour := cfg.AutoResponse.MaxBlocksPerHour
if maxPerHour <= 0 {
maxPerHour = config.DefaultMaxBlocksPerHour
}
budgetUnavailable := false
if durable, ok := blocker.(interface {
DurableActionsEnabled() bool
FirewallScanBudget(string) (int, error)
}); ok && durable.DurableActionsEnabled() {
used, budgetErr := durable.FirewallScanBudget(currentHour)
if budgetErr != nil {
budgetUnavailable = true
actions = append(actions, alert.Finding{Severity: alert.Warning, Check: "auto_block", Message: "Firewall action accounting unavailable; scan blocks deferred", Details: budgetErr.Error(), Timestamp: time.Now()})
} else {
state.BlocksThisHour = used
}
}
// http_asn_crawl: surgical subnet tempban for confirmed Critical findings.
// Each CIDR consumes one MaxBlocksPerHour slot. Independent of the per-IP
// list but shares its hourly budget. Skips infra intersections and
// already-blocked subnets; dry-run emits notices instead of blocking and
// consumes no budget.
if sb, ok := blocker.(subnetBlocker); ok {
tempban := parseExpiryWithDefault(cfg.AutoResponse.HTTPASNCrawlTempban, config.DefaultHTTPASNCrawlTempban)
for _, f := range findings {
work.progress()
if f.Check != "http_asn_crawl" || f.Severity != alert.Critical || len(f.CIDRs) == 0 {
continue
}
for _, cidr := range f.CIDRs {
work.progress()
if isSubnetAlreadyBlocked(blocker, cidr) || cidrIntersectsInfra(cfg, cidr) {
continue
}
if shouldSkipAutoSubnet(cfg, cidr, exemptLogged) {
continue
}
if !isAutoResponseActive(cfg) {
if !canDryRunBlockSubnet(blocker, cidr) {
continue
}
actions = append(actions, dryRunSubnetNotice(cidr, " (asn-crawl)", f.Message))
continue
}
if budgetUnavailable || state.BlocksThisHour >= maxPerHour {
break
}
reason := fmt.Sprintf("CSM auto-block (asn-crawl): %s", truncate(f.Message, 100))
var subnetErr error
if durable, ok := blocker.(interface {
DurableActionsEnabled() bool
BlockSubnetRequest(firewall.ActionRequest, *firewall.ScanAdmission) error
}); ok && durable.DurableActionsEnabled() {
subnetErr = durable.BlockSubnetRequest(firewall.ActionRequest{ID: rand.Text(), Operation: "block_subnet", Target: cidr, Reason: reason, TTL: tempban, FindingID: alert.FindingID(f), Actor: "daemon", Source: BlockSourceScan, Automatic: true}, &firewall.ScanAdmission{Window: currentHour, Limit: maxPerHour})
} else {
subnetErr = callBlockSubnet(sb, cidr, reason, tempban, alert.FindingID(f))
}
if !autoFirewallActionApplied(cidr, subnetErr) {
continue
}
state.BlocksThisHour++
actions = append(actions, alert.Finding{
Severity: alert.Critical,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-BLOCK-SUBNET: %s blocked (asn-crawl)", cidr),
Details: fmt.Sprintf("Reason: %s", f.Message),
Timestamp: time.Now(),
})
}
}
}
rateLimited := false
droppedPending := 0
engineUnavailableRequeued := 0
// requeue preserves an IP that could not be blocked this cycle (rate
// limit, engine unavailable, transient block error). QueuedAt is
// stamped on first entry so the drain's age check can retire it; the
// bound keeps a sustained flood from growing the queue without limit,
// and overflow drops are named so they remain identifiable in logs.
requeue := func(p pendingIP) bool {
if p.QueuedAt.IsZero() {
p.QueuedAt = autoBlockNow()
}
if len(state.Pending) < maxPendingBlocks {
state.Pending = append(state.Pending, work.requeueCandidate(p))
return true
}
work.rejectCandidate(p.queueCandidate)
fmt.Fprintf(os.Stderr, "auto-block: pending queue full, dropping %s\n", p.IP)
droppedPending++
return false
}
for ip, cand := range ipsToBlock {
work.progress()
work.startCandidate(cand.queueCandidate)
durableRetry := false
if durable, ok := blocker.(durableActionBlocker); ok {
durableRetry = cand.ActionID != "" && durable.DurableActionsEnabled()
}
// Existing requests need recovery even when their admission used the
// last slot. The durable store still caps any identity not yet admitted.
if budgetUnavailable || (state.BlocksThisHour >= maxPerHour && !durableRetry) {
requeue(cand)
rateLimited = !budgetUnavailable
continue
}
// Block via firewall engine (nftables)
blockReason := fmt.Sprintf("CSM auto-block: %s", truncate(cand.Reason, 100))
if blocker == nil {
work.candidateOutcome(cand.queueCandidate, ErrNoIPBlocker)
observeBlockOutcome(firewall.BlockOutcomeNoop, ErrNoIPBlocker, BlockSourceScan)
if requeue(cand) {
engineUnavailableRequeued++
}
continue
}
if cand.ActionID == "" {
cand.ActionID = rand.Text()
}
requestTTL := expiry
if durable, ok := blocker.(durableActionBlocker); ok && durable.DurableActionsEnabled() {
if cand.ActionTTL == 0 {
cand.ActionTTL = expiry
}
requestTTL = cand.ActionTTL
}
res, err := applyBlockLocked(cfg, blocker, state, ApplyBlockRequest{
ActionID: cand.ActionID,
IP: ip,
EngineReason: blockReason,
Reason: cand.Reason,
TTL: requestTTL,
Source: BlockSourceScan,
FindingID: cand.FindingID,
}, work.progress, func(err error) { work.candidateOutcome(cand.queueCandidate, err) })
verifiedAuditPending := res.Outcome == firewall.BlockOutcomeLive && errors.Is(err, firewall.ErrActionAuditPending)
if verifiedAuditPending {
fmt.Fprintf(os.Stderr, "auto-block: %s verified with audit delivery pending: %v\n", ip, err)
}
if err != nil && !verifiedAuditPending {
if errors.Is(err, firewall.ErrActionFailed) {
// A proven rejection is terminal for this request ID. A later
// attempt gets independent admission under the current policy.
cand.ActionID = ""
cand.ActionTTL = 0
}
// Protected IPs (the server's own interface or infra_ips) are
// intentionally never blocked -- an expected no-op, not a failure.
// The triggering finding still stands, so suspicious activity from a
// protected address is still surfaced. Any other error is treated
// as transient and the IP is requeued when capacity permits; the
// age cap retires it if the failure persists.
if !errors.Is(err, firewall.ErrIPProtected) {
if requeue(cand) {
fmt.Fprintf(os.Stderr, "auto-block: error blocking %s: %v (requeued)\n", ip, err)
} else {
fmt.Fprintf(os.Stderr, "auto-block: error blocking %s: %v (retry dropped)\n", ip, err)
}
} else {
work.finishCandidate(cand.queueCandidate)
}
continue
}
actions = append(actions, res.Findings...)
if res.Outcome != firewall.BlockOutcomeLive {
work.finishCandidate(cand.queueCandidate)
continue
}
if blocker.IsBlocked(ip) {
fmt.Fprintf(os.Stderr, "[%s] AUTO-BLOCK: %s blocked (expires in %s)\n", time.Now().Format("2006-01-02 15:04:05"), ip, requestTTL)
}
state.BlocksThisHour++
work.finishCandidate(cand.queueCandidate)
}
if engineUnavailableRequeued > 0 {
fmt.Fprintf(os.Stderr, "auto-block: firewall engine not available, requeued %d IPs\n", engineUnavailableRequeued)
}
warnRateLimit := rateLimited && state.RateLimitWarnedHour != currentHour
warnPendingDrop := droppedPending > 0 && state.PendingDropWarnedHour != currentHour
if warnRateLimit || warnPendingDrop {
var msg string
if rateLimited {
msg = fmt.Sprintf("Auto-block rate limit reached (%d/hour), %d IPs queued for next cycle", maxPerHour, len(state.Pending))
if droppedPending > 0 {
msg += fmt.Sprintf(", %d dropped (queue full)", droppedPending)
}
} else {
msg = fmt.Sprintf("Auto-block pending queue full, %d IPs queued for retry, %d dropped", len(state.Pending), droppedPending)
}
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_block",
Message: msg,
Timestamp: time.Now(),
})
if rateLimited {
state.RateLimitWarnedHour = currentHour
}
if droppedPending > 0 {
state.PendingDropWarnedHour = currentHour
}
}
// Subnet auto-blocking: detect per-family subnet patterns. Runs in
// dry-run too so escalations that would fire (from live blocks recorded
// before dry-run was enabled) surface as notices.
if cfg.AutoResponse.NetBlock && blocker != nil {
// Load fills the default and validation rejects anything lower, so
// this only catches a Config assembled in code.
threshold := cfg.AutoResponse.NetBlockThreshold
if threshold < config.MinBlockEscalationCount {
threshold = config.DefaultNetBlockThreshold
}
subnetExpiry := parseExpiry(cfg.AutoResponse.BlockExpiry)
// Count addresses blocked inside the window per subnet (IPv4 /24,
// IPv6 /64), not only those blocked right now: a subnet that rotates
// addresses keeps one block live at a time and would never escalate.
now := autoBlockNow()
window := netblockWindow(cfg)
work.progress()
// An unreadable history file is kept for inspection and never
// overwritten. Escalation falls back to the addresses blocked right
// now, which is what it counted before history existed.
history, historyErr := loadNetblockHistory(cfg.StatePath)
if historyErr != nil {
work.observe(historyErr)
fmt.Fprintf(os.Stderr, "autoblock: netblock history unavailable, counting current blocks only: %v\n", historyErr)
history = &netblockHistory{IPs: map[string]time.Time{}, Subnets: map[string]time.Time{}}
}
current := currentlyBlocked(state.IPs, liveBlocked, useLiveBlocked, history, blocker)
historyChanged := recordNetblockHistory(history, state.IPs, current, blocker, now, window)
// Direct mail-spray blocks answer the same prior offenders as a
// threshold-based block, even when the IP blocks have already ended.
for cidr := range answeredSubnets {
history.Subnets[cidr] = now
historyChanged = true
}
subnetCounts := netblockCounts(cfg, history, current, blocker, now, window)
subnetCauses := make(map[string]blockedIP)
subnetBlocked := make(map[string]bool)
for _, b := range state.IPs {
cidr := subnetEscalationCIDR(b.IP)
prior := subnetCauses[cidr]
if cidr != "" && b.FindingID != "" && (prior.FindingID == "" || b.BlockedAt.After(prior.BlockedAt) || (b.BlockedAt.Equal(prior.BlockedAt) && b.FindingID > prior.FindingID)) {
subnetCauses[cidr] = b
}
}
for cidr, count := range subnetCounts {
work.progress()
if count >= threshold && !subnetBlocked[cidr] {
if sb, ok := blocker.(subnetBlocker); ok {
if isSubnetAlreadyBlocked(blocker, cidr) {
continue
}
if cidrIntersectsInfra(cfg, cidr) {
continue
}
if shouldSkipAutoSubnet(cfg, cidr, exemptLogged) {
continue
}
if !isAutoResponseActive(cfg) {
if !canDryRunBlockSubnet(blocker, cidr) {
continue
}
subnetBlocked[cidr] = true
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-NETBLOCK [dry-run]: %s would be blocked (%d IPs from same subnet)", cidr, count),
Timestamp: time.Now(),
})
continue
}
reason := fmt.Sprintf("Auto-netblock: %d IPs from %s within %s", count, cidr, window)
if autoFirewallActionApplied(cidr, callBlockSubnet(sb, cidr, reason, subnetExpiry, subnetCauses[cidr].FindingID)) {
subnetBlocked[cidr] = true
history.Subnets[cidr] = now
historyChanged = true
fmt.Fprintf(os.Stderr, "[%s] AUTO-NETBLOCK: %s blocked (%d IPs from same subnet)\n", time.Now().Format("2006-01-02 15:04:05"), cidr, count)
actions = append(actions, alert.Finding{
Severity: alert.Critical,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-NETBLOCK: %s blocked (%d IPs from same subnet)", cidr, count),
Timestamp: time.Now(),
})
}
}
}
}
if historyChanged && historyErr == nil {
work.progress()
if err := saveNetblockHistory(cfg.StatePath, history); err != nil {
work.observe(err)
fmt.Fprintf(os.Stderr, "autoblock: %v\n", err)
}
}
}
// Save state (expired IPs were already pruned at the top of this function)
work.progress()
work.saveState(cfg.StatePath, state)
work.complete()
return actions
}
// isBlockedLiveOrCached returns the live nftables status when the
// blocker supports it, otherwise falls back to the cached IsBlocked
// view. The reconcile loop relies on this to prune blocked_ips.json
// entries the kernel has already expired even when state.json has not
// caught up yet. Live lookup errors keep the cached answer so transient
// netlink failures do not erase the local tracker.
func isBlockedLiveOrCached(b IPBlocker, ip string) bool {
if lb, ok := b.(liveBlocker); ok {
blocked, err := lb.IsBlockedLive(ip)
if err == nil {
return blocked
}
}
return b.IsBlocked(ip)
}
// liveBlockedSnapshot takes one bulk membership snapshot for the cycle. A
// blocker that cannot produce one leaves callers on the legacy per-IP path.
// A failed bulk dump returns an uncovered snapshot so valid IPs use cached
// status without issuing the same failing dump once per tracked address.
func liveBlockedSnapshot(b IPBlocker) (firewall.LiveBlockedSnapshot, bool) {
lister, ok := b.(liveBlockedLister)
if !ok {
return firewall.LiveBlockedSnapshot{}, false
}
snap, err := lister.LiveBlockedSet()
if err != nil {
// The snapshot still describes whichever families answered, so keep
// it: discarding it would drop a healthy family to cached status
// because the other one failed.
fmt.Fprintf(os.Stderr, "auto-block: live blocked-set dump incomplete, using cached status for uncovered families: %v\n", err)
}
return snap, true
}
// blockedLiveOrCached answers from the cycle's bulk snapshot when it covers
// the IP's address family. An uncovered family uses cached status; retrying a
// live query would repeat the same set dump and could turn an unknown result
// into a false absence. Blockers without snapshot support retain the legacy
// per-IP live lookup.
func blockedLiveOrCached(b IPBlocker, snap firewall.LiveBlockedSnapshot, useSnap bool, ip string) bool {
if useSnap {
if blocked, known := snap.Contains(ip); known {
return blocked
}
return b.IsBlocked(ip)
}
return isBlockedLiveOrCached(b, ip)
}
// callBlockIP dispatches to the outcome-reporting interface when the
// underlying blocker implements it, otherwise falls back to the legacy
// IPBlocker interface and assumes the call landed live (the behaviour
// every IPBlocker had before BlockIPOutcome existed). This keeps tests
// and any third-party implementations of IPBlocker working unchanged.
func callBlockIP(b IPBlocker, ip, reason string, timeout time.Duration, findingID string) (firewall.BlockOutcome, error) {
if ob, ok := b.(interface {
BlockIPOutcomeWithFindingID(string, string, time.Duration, string) (firewall.BlockOutcome, error)
}); ok {
return ob.BlockIPOutcomeWithFindingID(ip, reason, timeout, findingID)
}
if ob, ok := b.(outcomeBlocker); ok {
return ob.BlockIPOutcome(ip, reason, timeout)
}
if err := b.BlockIP(ip, reason, timeout); err != nil {
return firewall.BlockOutcomeNoop, err
}
return firewall.BlockOutcomeLive, nil
}
// shouldSkipAutoBlockForChallenge reports whether an IP carrying this finding
// should be left for the challenge gate instead of hard-blocked. It is the
// exact inverse of responseActionForFinding resolving to a block, so the two
// auto-response paths share one decision.
func shouldSkipAutoBlockForChallenge(cfg *config.Config, f alert.Finding) bool {
return responseActionForFinding(cfg, f) == responseChallenge
}
// promoteToPermanentBlock upgrades an existing temp block to permanent. The
// real engine implements permanentPromoter and clears the kernel timeout in
// place. Legacy blockers that only implement BlockIP have not marked the IP
// blocked in a way that trips skipExisting, so a fresh zero-timeout block on
// them lands live; that fallback preserves pre-existing behaviour for tests
// and third-party implementations.
func promoteToPermanentBlock(b IPBlocker, ip, reason, findingID string) bool {
if pp, ok := b.(interface {
PromoteToPermanentBlockWithFindingID(string, string, string) error
}); ok {
return autoFirewallActionApplied(ip, pp.PromoteToPermanentBlockWithFindingID(ip, reason, findingID))
}
if pp, ok := b.(permanentPromoter); ok {
return autoFirewallActionApplied(ip, pp.PromoteToPermanentBlock(ip, reason))
}
outcome, err := callBlockIP(b, ip, reason, 0, findingID)
return outcome == firewall.BlockOutcomeLive && autoFirewallActionApplied(ip, err)
}
func autoFirewallActionApplied(target string, err error) bool {
if err == nil {
return true
}
// A verified mutation still needs its normal response evidence when the
// separate audit delivery is pending. Keep the degradation visible.
fmt.Fprintf(os.Stderr, "auto-block: firewall action for %s: %v\n", target, err)
return errors.Is(err, firewall.ErrActionAuditPending)
}
func isSubnetAlreadyBlocked(b IPBlocker, cidr string) bool {
sb, ok := b.(subnetBlockStatus)
return ok && sb.IsSubnetBlocked(cidr)
}
// ExtractIPFromFinding extracts an IP address from a finding.
func ExtractIPFromFinding(f alert.Finding) string {
return extractIPFromFinding(f)
}
// ManualBlockIP returns the address an operator may block from a finding: the
// one auto-block would act on, for any check that reports an attacker,
// including the login-failure checks that block only when configured. It is
// empty for every other check, whose messages can quote a victim or customer
// address from a log line.
func ManualBlockIP(f alert.Finding) string {
if !blockableCheck(f.Check, true) {
return ""
}
return extractIPFromFinding(f)
}
func extractIPFromFinding(f alert.Finding) string {
if strings.TrimSpace(f.SourceIP) != "" {
return normalizeBlockIP(f.SourceIP)
}
msg := f.Message
// Fallback for detectors that have not yet adopted the structured SourceIP
// field. Only findings whose Check is auto-block-eligible reach this path,
// and those detectors format their own messages with a CSM-parsed IP at the
// tail. Use LastIndex so the rightmost (CSM-appended) IP wins over any
// log-injected content earlier in the message.
for _, sep := range []string{" from ", ": "} {
if idx := strings.LastIndex(msg, sep); idx >= 0 {
rest := msg[idx+len(sep):]
fields := strings.Fields(rest)
if len(fields) > 0 {
if candidate, ok := netutil.ParseIPToken(fields[0]); ok {
if ip := normalizeBlockIP(candidate); ip != "" {
return ip
}
}
}
}
}
return ""
}
func normalizeBlockIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if host, _, err := net.SplitHostPort(raw); err == nil {
raw = host
}
raw = strings.Trim(raw, "[]")
ip := net.ParseIP(raw)
if ip == nil || ip.IsLoopback() || ip.IsUnspecified() {
return ""
}
return ip.String()
}
func isAlreadyBlocked(state *blockState, ip string) bool {
for _, b := range state.IPs {
if b.IP == ip {
return true
}
}
return false
}
func parseExpiry(s string) time.Duration {
return parseExpiryWithDefault(s, config.DefaultBlockExpiry)
}
func parseExpiryWithDefault(s, fallback string) time.Duration {
d, err := time.ParseDuration(s)
if s != "" && err == nil && d > 0 {
return d
}
d, _ = time.ParseDuration(fallback)
return d
}
func loadBlockState(statePath string) *blockState {
state, err := readBlockState(statePath)
if err != nil {
fmt.Fprintf(os.Stderr, "autoblock: %v; ignoring queued blocks\n", err)
return &blockState{}
}
return state
}
func readBlockState(statePath string) (*blockState, error) {
path := filepath.Join(statePath, blockStateFile)
data, err := osFS.ReadFile(path)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return &blockState{}, nil
}
return nil, fmt.Errorf("reading %s: %w", path, err)
}
state := &blockState{}
if err := json.Unmarshal(data, state); err != nil {
return nil, fmt.Errorf("reading %s: %w", path, err)
}
return state, nil
}
func saveBlockState(statePath string, s *blockState) {
if err := writeBlockState(statePath, s); err != nil {
logBlockStateFailure(statePath, err)
}
}
func writeBlockState(statePath string, s *blockState) error {
return atomicio.AtomicWriteJSON(filepath.Join(statePath, blockStateFile), 0o600, s)
}
func logBlockStateFailure(statePath string, err error) {
fmt.Fprintf(os.Stderr, "autoblock: persist %s failed: %v\n", filepath.Join(statePath, blockStateFile), err)
}
// subnetEscalationCIDR returns the canonical CIDR used by the
// auto-netblock escalation path for the given IP. IPv4 collapses to
// /24 (the historical block size); IPv6 collapses to /64 because most
// providers hand out /64 prefixes to end users -- /128 would let
// attackers rotate addresses inside the same /64 and never escalate,
// while a wider prefix would risk taking down legitimate neighbours.
// Returns "" for unparseable input.
func subnetEscalationCIDR(ip string) string {
parsed := net.ParseIP(ip)
if parsed == nil {
return ""
}
if ip4 := parsed.To4(); ip4 != nil {
return fmt.Sprintf("%d.%d.%d.0/24", ip4[0], ip4[1], ip4[2])
}
ip16 := parsed.To16()
if ip16 == nil {
return ""
}
mask := net.CIDRMask(64, 128)
network := ip16.Mask(mask)
return (&net.IPNet{IP: network, Mask: mask}).String()
}
// --- Permanent block escalation (LF_PERMBLOCK) ---
type permBlockTracker struct {
IPs map[string][]time.Time `json:"ips"` // IP -> list of block timestamps
}
// checkPermBlockEscalation records a new block and returns true if the IP
// has been temp-blocked count times within interval.
func checkPermBlockEscalation(statePath, ip string, count int, interval time.Duration) bool {
tracker := loadPermBlockTracker(statePath)
now := time.Now()
cutoff := now.Add(-interval)
// Add current block timestamp
tracker.IPs[ip] = append(tracker.IPs[ip], now)
// Clean old entries for this IP
var recent []time.Time
for _, t := range tracker.IPs[ip] {
if t.After(cutoff) {
recent = append(recent, t)
}
}
tracker.IPs[ip] = recent
// Clean old IPs entirely (haven't been seen in 2x the interval)
for k, times := range tracker.IPs {
if len(times) == 0 {
delete(tracker.IPs, k)
continue
}
latest := times[len(times)-1]
if now.Sub(latest) > interval*2 {
delete(tracker.IPs, k)
}
}
savePermBlockTracker(statePath, tracker)
return len(recent) >= count
}
func loadPermBlockTracker(statePath string) *permBlockTracker {
tracker := &permBlockTracker{IPs: make(map[string][]time.Time)}
path := filepath.Join(statePath, "permblock_tracker.json")
data, err := osFS.ReadFile(path)
if err == nil {
if uerr := json.Unmarshal(data, tracker); uerr != nil {
fmt.Fprintf(os.Stderr, "autoblock: %s is corrupt, ignoring escalation history: %v\n", path, uerr)
}
if tracker.IPs == nil {
tracker.IPs = make(map[string][]time.Time)
}
}
return tracker
}
func savePermBlockTracker(statePath string, tracker *permBlockTracker) {
path := filepath.Join(statePath, "permblock_tracker.json")
if err := atomicio.AtomicWriteJSON(path, 0o600, tracker); err != nil {
fmt.Fprintf(os.Stderr, "autoblock: persist %s failed: %v\n", path, err)
}
}
// AutoBlockFlushResult reports which phases of a coordinated firewall flush
// completed and whether its best-effort persisted-state snapshot was readable.
type AutoBlockFlushResult struct {
Flushed bool
BlockedCount int
SnapshotErr error
}
// FlushAutoBlockState snapshots the engine's pre-flush IPs, runs an operator
// firewall flush, and clears the auto-block bookkeeping in one critical
// section. Without this the flush was self-reverting: surviving ThreatDB temp
// rows re-flagged every flushed IP through ip_reputation on the next scan and
// re-blocked it, and stale tracker entries suppressed re-block accounting.
// Tracker entries the engine never held are cleaned up too. Pending entries
// are kept - they are queued candidates, not blocks.
//
// The firewall mutation and cleanup stay serialized with AutoBlockIPs so a
// concurrent scan either finishes before the snapshot or starts after cleanup
// with fresh evidence. Result.Flushed is true with a non-nil error when the
// firewall was flushed but bookkeeping cleanup was only partial. SnapshotErr
// is advisory because the tracker-side union still covers tracked auto-blocks.
func FlushAutoBlockState(statePath string, flush func() error) (AutoBlockFlushResult, error) {
work := autoBlockQueues.acquire()
defer work.finish()
var result AutoBlockFlushResult
work.progress()
engineState, snapshotErr := firewall.LoadState(statePath)
result.SnapshotErr = snapshotErr
var ips []string
if snapshotErr == nil {
result.BlockedCount = len(engineState.Blocked)
ips = make([]string, 0, result.BlockedCount)
for _, b := range engineState.Blocked {
ips = append(ips, b.IP)
}
}
work.progress()
if err := flush(); err != nil {
work.observe(err)
work.complete()
return result, fmt.Errorf("flushing blocked IPs: %w", err)
}
result.Flushed = true
work.beginCleanup(statePath, ips, snapshotErr)
var cleanupErr error
work.progress()
state, err := work.readState(statePath)
seenCapacity := len(ips)
if state != nil {
seenCapacity += len(state.IPs) + len(state.CleanupPending)
}
seen := make(map[string]bool, seenCapacity)
cleanupIPs := make([]string, 0, seenCapacity)
addCleanupIP := func(ip string) {
if !seen[ip] {
seen[ip] = true
cleanupIPs = append(cleanupIPs, ip)
work.admitCleanup(ip)
}
}
for _, ip := range ips {
addCleanupIP(ip)
}
if err != nil {
work.observe(err)
cleanupErr = errors.Join(cleanupErr, fmt.Errorf("reading auto-block state: %w", err))
} else {
for _, b := range state.IPs {
addCleanupIP(b.IP)
}
for _, ip := range state.CleanupPending {
addCleanupIP(ip)
}
}
sdb := store.Global()
tdb := GetThreatDB()
failed := make([]string, 0)
for _, ip := range cleanupIPs {
work.startCleanup(ip)
cleanupFailed := false
if sdb != nil {
if _, err := sdb.RemoveTemporaryBlock(ip); err != nil {
cleanupFailed = true
work.cleanupOutcome(ip, true)
work.observe(err)
cleanupErr = errors.Join(cleanupErr, fmt.Errorf("removing auto-block store row for %s: %w", ip, err))
failed = append(failed, ip)
}
}
if tdb != nil {
tdb.RemoveTemporary(ip)
}
work.cleanupOutcome(ip, cleanupFailed)
}
work.progress()
if err := saveNetblockHistory(statePath, &netblockHistory{}); err != nil {
work.observe(err)
cleanupErr = errors.Join(cleanupErr, fmt.Errorf("clearing netblock history: %w", err))
}
if state != nil {
state.IPs = nil
// A failed bbolt cleanup must survive in the tracker after the
// firewall state is empty, or a retry has no way to identify the
// stale row that can recreate the block after restart.
state.CleanupPending = failed
work.progress()
if err := work.writeState(statePath, state); err != nil {
work.observe(err)
logBlockStateFailure(statePath, err)
cleanupErr = errors.Join(cleanupErr, fmt.Errorf("clearing auto-block state: %w", err))
}
}
work.complete()
return result, cleanupErr
}
// dryRunSubnetNotice mirrors the per-IP dry-run notice for the subnet block
// paths so operators evaluating dry-run see subnet decisions instead of
// silence. The message deliberately differs from the live
// "AUTO-BLOCK-SUBNET:" token so alert-filter suppression never treats a
// notice as a real block.
func dryRunSubnetNotice(cidr, kind, reason string) alert.Finding {
return alert.Finding{
Severity: alert.Warning,
Check: "auto_block",
Message: fmt.Sprintf("AUTO-BLOCK-SUBNET [dry-run]: %s would be blocked%s", cidr, kind),
Details: fmt.Sprintf("Reason: %s", reason),
Timestamp: time.Now(),
}
}
// canDryRunBlockSubnet applies the firewall engine's read-only preflight when
// available. A missing or incapable engine cannot make the claimed live block,
// so it must not produce a "would be blocked" notice.
func canDryRunBlockSubnet(blocker IPBlocker, cidr string) bool {
if blocker == nil {
fmt.Fprintf(os.Stderr, "auto-block: firewall engine not available, skipping dry-run subnet %s\n", cidr)
return false
}
if _, ok := blocker.(subnetBlocker); !ok {
fmt.Fprintf(os.Stderr, "auto-block: firewall engine does not support subnet blocking, skipping dry-run subnet %s\n", cidr)
return false
}
if validator, ok := blocker.(subnetBlockValidator); ok {
if err := validator.ValidateSubnetBlock(cidr); err != nil {
fmt.Fprintf(os.Stderr, "auto-block: dry-run subnet %s rejected: %v\n", cidr, err)
return false
}
}
return true
}
// isAutoResponseActive reports whether real blocking should happen now:
// auto-response enabled, IP blocking on, and not in dry-run.
// DryRun defaults to true (safe) when nil — operators must explicitly set
// dry_run: false to enable live nftables blocking.
func isAutoResponseActive(cfg *config.Config) bool {
return cfg.AutoResponse.Enabled && cfg.AutoResponse.BlockIPs && !cfg.AutoResponseDryRunEnabled()
}
// cidrIntersectsInfra reports whether the CIDR contains an operator infra
// IP/range (or loopback), so the subnet tempban never blackholes protected
// addresses. An unparseable CIDR fails safe (treated as intersecting and
// skipped). The firewall engine's dynamic per-IP allowlist is not
// enumerable across a subnet, so infra_ips is the operator's mechanism to
// exempt a specific address from subnet tempban.
func cidrIntersectsInfra(cfg *config.Config, cidr string) bool {
_, ipnet, err := net.ParseCIDR(cidr)
if err != nil {
return true
}
if ipnet.IP.IsLoopback() {
return true
}
candidates := append([]string{"127.0.0.1", "::1"}, cfg.InfraIPs...)
for _, raw := range candidates {
raw = strings.TrimSpace(raw)
if raw == "" {
continue
}
if p := net.ParseIP(raw); p != nil && ipnet.Contains(p) {
return true
}
if _, infraNet, err := net.ParseCIDR(raw); err == nil &&
(ipnet.Contains(infraNet.IP) || infraNet.Contains(ipnet.IP)) {
return true
}
}
return false
}
// dosExemptNets returns the effective DoS-exempt networks (operator ranges
// union-ed with the current mail-provider overlay) split into IPv4 and IPv6
// slices. Delegates to firewall.EffectiveDOSExemptNets so the same logic
// governs both nftables sets and auto-block guards.
func dosExemptNets(cfg *config.Config) (v4, v6 []*net.IPNet) {
var fc *firewall.FirewallConfig
if cfg != nil {
fc = cfg.Firewall
}
return firewall.EffectiveDOSExemptNets(fc, mailranges.ProviderNets())
}
// cidrIntersectsDOSExempt reports whether the CIDR overlaps any DoS-exempt
// network, so the subnet tempban never blocks exempt sources (e.g., operator-
// designated ranges or known mail-provider egress). An unparseable CIDR fails
// safe (treated as intersecting, block skipped).
func cidrIntersectsDOSExempt(cfg *config.Config, cidr string) bool {
_, ipnet, err := net.ParseCIDR(cidr)
if err != nil {
return true // fail-safe: skip block when CIDR is unreadable
}
v4, v6 := dosExemptNets(cfg)
for _, exempt := range append(v4, v6...) { //nolint:gocritic // intentional inline join
if ipnet.Contains(exempt.IP) || exempt.Contains(ipnet.IP) {
return true
}
}
return false
}
// shouldSkipAutoSubnet reports whether the auto-block path must suppress a
// BlockSubnet call for cidr because it intersects a DoS-exempt operator range.
// Logs once per CIDR per cycle (via the logged dedupe map) at stderr info level
// with reason dos_exempt_range. Returns true when the block must be suppressed.
// Manual operator subnet denies bypass this function entirely.
func shouldSkipAutoSubnet(cfg *config.Config, cidr string, logged map[string]struct{}) bool {
if !cidrIntersectsDOSExempt(cfg, cidr) {
return false
}
if _, already := logged[cidr]; !already {
logged[cidr] = struct{}{}
fmt.Fprintf(os.Stderr, "auto-block: skipping subnet %s (dos_exempt_range)\n", cidr)
}
return true
}
// PruneExemptAutoSubnets removes auto-response subnet blocks whose CIDR now
// intersects the DoS-exempt set (operator ranges or mail-provider overlay).
// Only entries with Source == firewall.SourceAutoResponse are touched; manual,
// CLI, web-UI, challenge, whitelist, dyndns, system, and unknown-source blocks
// are left untouched. If b does not implement subnetManager, returns 0.
// UnblockSubnet errors are logged and the entry is not counted as pruned.
func PruneExemptAutoSubnets(cfg *config.Config, b IPBlocker) int {
return pruneExemptAutoSubnets(cfg, b, func() {})
}
func pruneExemptAutoSubnets(cfg *config.Config, b IPBlocker, progress func()) int {
sm, ok := b.(subnetManager)
if !ok {
return 0
}
pruned := 0
progress()
for _, entry := range sm.BlockedSubnets() {
progress()
if entry.Source != firewall.SourceAutoResponse {
continue
}
if !cidrIntersectsDOSExempt(cfg, entry.CIDR) {
continue
}
if err := sm.UnblockSubnet(entry.CIDR); err != nil {
fmt.Fprintf(os.Stderr, "auto-block: prune exempt subnet %s: %v\n", entry.CIDR, err)
continue
}
pruned++
}
return pruned
}
// extractCIDRFromFinding returns the CIDR appearing in the message after
// the canonical " from " separator. Returns "" if the value does not parse
// as a CIDR.
func extractCIDRFromFinding(f alert.Finding) string {
msg := f.Message
idx := strings.LastIndex(msg, " from ")
if idx < 0 {
return ""
}
rest := msg[idx+len(" from "):]
fields := strings.Fields(rest)
if len(fields) == 0 {
return ""
}
// A CIDR never ends in ':' (the prefix length is last), so a trailing
// separator can be dropped here; IP tokens go through netutil instead.
candidate := strings.TrimRight(strings.TrimRight(fields[0], ",;)([]"), ":")
_, ipnet, err := net.ParseCIDR(candidate)
if err != nil {
return ""
}
return ipnet.String()
}
// callBlockSubnet uses causal metadata when the engine supports it, preserving
// the legacy interface for other blocker implementations.
func callBlockSubnet(b subnetBlocker, cidr, reason string, timeout time.Duration, findingID string) error {
if sb, ok := b.(interface {
BlockSubnetWithFindingID(string, string, time.Duration, string) error
}); ok {
return sb.BlockSubnetWithFindingID(cidr, reason, timeout, findingID)
}
return b.BlockSubnet(cidr, reason, timeout)
}
package checks
import (
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
// The state-call gate serializes real cleanup. Its separate metadata mutex
// keeps health available while filesystem and database operations are blocked.
type autoBlockCleanupQueue struct {
path string
records map[string]*autoBlockCleanupRecord
sources map[string]bool
blockSources map[string]map[autoBlockCleanupBlock]bool
depthKnown, historyUnknown, lowerBound bool
readFailed, writeFailed, snapshotFailed bool
loss *queuehealth.Tracker
}
type autoBlockCleanupRecord struct {
at time.Time
inFlight, completed, failed, lossKnown bool
blocks map[autoBlockCleanupBlock]bool
blocksKnown bool
}
type autoBlockCleanupBlock struct {
blockedAt, expiresAt time.Time
}
type autoBlockCleanupCycle struct {
entries map[string]*autoBlockCleanupRecord
settled bool
}
func newAutoBlockCleanupQueue() *autoBlockCleanupQueue {
return &autoBlockCleanupQueue{
records: make(map[string]*autoBlockCleanupRecord),
loss: queuehealth.New(0, time.Minute),
}
}
func (c *autoBlockCleanupQueue) setPath(path string) {
if c.path == path {
return
}
c.path = path
clear(c.records)
c.sources = nil
c.blockSources = nil
c.depthKnown, c.historyUnknown = false, false
c.readFailed, c.writeFailed, c.snapshotFailed = false, false, false
}
func (c *autoBlockCleanupQueue) record(ip string, known bool) *autoBlockCleanupRecord {
r := c.records[ip]
if r == nil {
r = &autoBlockCleanupRecord{at: time.Now(), lossKnown: known}
c.records[ip] = r
}
return r
}
func (w *autoBlockStateWork) beginCleanup(path string, ips []string, snapshotErr error) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
c := q.cleanup
c.setPath(path)
c.snapshotFailed = snapshotErr != nil
c.lowerBound = c.lowerBound || snapshotErr != nil
w.cleanupCycle = &autoBlockCleanupCycle{entries: make(map[string]*autoBlockCleanupRecord)}
for _, ip := range ips {
r := c.record(ip, true)
// A live engine entry is new cleanup demand even if an earlier
// completed cleanup left its retry marker after a failed save.
r.completed, r.lossKnown = false, true
w.cleanupCycle.entries[ip] = r
}
}
func (w *autoBlockStateWork) admitCleanup(ip string) {
w.queue.mu.Lock()
c := w.queue.cleanup
r := c.record(ip, !c.historyUnknown)
blocks := c.blockSources[ip]
if c.depthKnown && r.blocksKnown {
for version := range blocks {
if !r.blocks[version] {
r.completed, r.lossKnown = false, true
break
}
}
}
// Each snapshot owns its immutable source set. Retain only the latest
// cleanup's generations, rather than accumulating every earlier block.
r.blocks = blocks
r.blocksKnown = c.depthKnown
w.cleanupCycle.entries[ip] = r
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) startCleanup(ip string) {
w.queue.mu.Lock()
w.cleanupCycle.entries[ip].inFlight = true
w.at = time.Now()
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) cleanupOutcome(ip string, failed bool) {
w.queue.mu.Lock()
r := w.cleanupCycle.entries[ip]
r.completed = r.completed || !failed
r.failed = failed
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) observeCleanupState(state *blockState, settle bool) {
c := w.queue.cleanup
sources := make(map[string]bool, len(state.IPs)+len(state.CleanupPending))
blocks := make(map[string]map[autoBlockCleanupBlock]bool, len(state.IPs))
for _, b := range state.IPs {
sources[b.IP] = true
if blocks[b.IP] == nil {
blocks[b.IP] = make(map[autoBlockCleanupBlock]bool)
}
blocks[b.IP][autoBlockCleanupBlock{b.BlockedAt.UTC(), b.ExpiresAt.UTC()}] = true
}
for _, ip := range state.CleanupPending {
sources[ip] = true
c.record(ip, !c.historyUnknown)
}
c.sources = sources
c.blockSources = blocks
c.depthKnown = true
for ip, r := range c.records {
if !r.blocksKnown {
// First rediscovery establishes a baseline, not fresh demand.
// Keep it until the next cleanup can distinguish newer blocks.
r.blocks, r.blocksKnown = blocks[ip], true
}
}
w.settleCleanupRecords(settle)
}
func (w *autoBlockStateWork) settleCleanupRecords(settle bool) {
c := w.queue.cleanup
for ip, r := range c.records {
active := w.cleanupCycle != nil && w.cleanupCycle.entries[ip] == r
if !c.sources[ip] && (settle || !active) {
if !r.completed && r.lossKnown {
c.loss.Lose(time.Now(), 1)
}
delete(c.records, ip)
} else if settle {
r.inFlight = false
}
}
}
func (w *autoBlockStateWork) forgetCleanupState() {
c := w.queue.cleanup
c.sources = nil
c.blockSources = nil
c.depthKnown, c.historyUnknown, c.lowerBound = false, true, true
for ip, r := range c.records {
if w.cleanupCycle != nil && w.cleanupCycle.entries[ip] != r {
delete(c.records, ip)
}
}
}
func (w *autoBlockStateWork) finishCleanupLocked() {
c := w.queue.cleanup
switch {
case w.readingState:
c.readFailed = true
w.forgetCleanupState()
case w.retryCycle != nil && w.retryCycle.saving && !w.retryCycle.settled:
c.writeFailed = true
w.forgetCleanupState()
}
if w.cleanupCycle == nil {
return
}
if !w.completed {
for _, r := range w.cleanupCycle.entries {
if r.inFlight && !r.completed {
r.failed = true
}
}
}
if w.cleanupCycle.settled {
return
}
if c.depthKnown {
w.settleCleanupRecords(true)
} else {
// Keep only the last observed batch across uncertainty so a later
// tracker read can recover its retry sources and acknowledgments.
for _, r := range c.records {
r.inFlight = false
}
}
}
func (c *autoBlockCleanupQueue) status(now time.Time, active *autoBlockStateWork) queuehealth.Status {
row := c.loss.Snapshot(now)
row.CapacityUnavailable = true
row.DepthUnavailable = !c.depthKnown
row.DroppedLowerBound = c.lowerBound
row.LagBasis = "deferred_checkpoint"
if active != nil && active.cleanupCycle != nil && len(c.records) > 0 {
row.ProcessingSeconds = max(0, now.Sub(active.at).Seconds())
}
failed := false
for ip, r := range c.records {
if r.inFlight {
row.InFlight++
} else if c.depthKnown || active != nil && active.cleanupCycle != nil && active.cleanupCycle.entries[ip] == r {
row.Depth++
row.LagSeconds = max(row.LagSeconds, now.Sub(r.at).Seconds())
}
failed = failed || r.failed
}
switch {
case c.readFailed || c.writeFailed || c.snapshotFailed:
row.Status, row.Reason = "degraded", "state_io"
case failed:
row.Status, row.Reason = "degraded", "retry_failed"
case row.ProcessingSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "processing_lag"
}
return row
}
package checks
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"time"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
)
const netblockHistoryFile = "netblock_history.json"
// netblockHistoryPruneEvery bounds how often aged entries are rewritten out of
// the history file. Counting ignores them either way, and the file would
// otherwise be rewritten on nearly every cycle as entries cross the window.
const netblockHistoryPruneEvery = time.Hour
// netblockHistory remembers when each address was last seen blocked, so a
// subnet that rotates through addresses one block at a time still reaches
// the netblock threshold. Subnets records when each subnet was last blocked:
// only offenders seen after that count toward the next subnet block.
type netblockHistory struct {
IPs map[string]time.Time `json:"ips"`
Active map[string]bool `json:"active"`
Subnets map[string]time.Time `json:"subnets,omitempty"`
PrunedAt time.Time `json:"pruned_at,omitempty"`
}
func loadNetblockHistory(statePath string) (*netblockHistory, error) {
h := &netblockHistory{}
path := filepath.Join(statePath, netblockHistoryFile)
data, err := osFS.ReadFile(path)
if err != nil && !errors.Is(err, os.ErrNotExist) {
return nil, fmt.Errorf("read %s: %w", path, err)
}
if err == nil {
if err := json.Unmarshal(data, h); err != nil {
return nil, fmt.Errorf("decode %s: %w", path, err)
}
}
if h.IPs == nil {
h.IPs = make(map[string]time.Time)
}
if h.Subnets == nil {
h.Subnets = make(map[string]time.Time)
}
return h, nil
}
func saveNetblockHistory(statePath string, h *netblockHistory) error {
path := filepath.Join(statePath, netblockHistoryFile)
if err := atomicio.AtomicWriteJSON(path, 0o600, h); err != nil {
return fmt.Errorf("persist %s: %w", path, err)
}
return nil
}
// ForgetNetblockHistory serializes the operator's firewall/database mutation
// and history removal with auto-block cycles. clear may be nil for history-only
// cleanup; it must not call another auto-block state operation.
func ForgetNetblockHistory(statePath, ip string, clear func()) error {
work := autoBlockQueues.acquire()
defer work.finish()
work.progress()
if clear != nil {
clear()
}
h, err := loadNetblockHistory(statePath)
if err == nil {
if _, ok := h.IPs[ip]; ok {
delete(h.IPs, ip)
delete(h.Active, ip)
err = saveNetblockHistory(statePath, h)
}
}
work.observe(err)
work.complete()
return err
}
// netblockWindow resolves the counting window. Load fills the default and
// validation rejects anything unparseable, so the fallback only covers a
// Config assembled in code.
func netblockWindow(cfg *config.Config) time.Duration {
return parseExpiryWithDefault(cfg.AutoResponse.NetBlockWindow, config.DefaultNetBlockWindow)
}
// recordNetblockHistory notes every address blocked right now: tracker
// entries with their block times, and addresses only the live kernel set
// knows (operator and permanent blocks), stamped when first seen. It reports
// whether the history changed.
func recordNetblockHistory(h *netblockHistory, tracked []blockedIP, current map[string]bool, blocker IPBlocker, now time.Time, window time.Duration) bool {
changed := false
// Persist membership transitions, not per-cycle timestamps: an operator
// re-block is fresh evidence, but a long-lived block must not churn the file.
if h.Active == nil {
h.Active = make(map[string]bool)
for ip := range h.IPs {
h.Active[ip] = true
}
changed = true
}
if allow, ok := blocker.(allowChecker); ok {
for ip := range current {
if allow.IsAllowed(ip) {
delete(current, ip)
}
}
for ip := range h.IPs {
if allow.IsAllowed(ip) {
delete(h.IPs, ip)
delete(h.Active, ip)
changed = true
}
}
}
note := func(ip string, at time.Time) {
if prev, ok := h.IPs[ip]; !ok || at.After(prev) {
h.IPs[ip] = at
changed = true
}
}
inTracker := make(map[string]bool, len(tracked))
for _, b := range tracked {
inTracker[b.IP] = true
if current[b.IP] {
note(b.IP, b.BlockedAt)
}
}
for ip := range current {
if !inTracker[ip] {
if _, ok := h.IPs[ip]; !ok || !h.Active[ip] {
note(ip, now)
}
}
}
for ip := range h.Active {
if !current[ip] {
delete(h.Active, ip)
changed = true
}
}
for ip := range current {
if !h.Active[ip] {
h.Active[ip] = true
changed = true
}
}
if now.Sub(h.PrunedAt) >= netblockHistoryPruneEvery {
for ip, at := range h.IPs {
if !current[ip] && now.Sub(at) > window {
delete(h.IPs, ip)
}
}
for cidr, at := range h.Subnets {
if now.Sub(at) > window {
delete(h.Subnets, cidr)
}
}
h.PrunedAt = now
changed = true
}
return changed
}
// currentlyBlocked is every address blocked right now, from the tracker and
// from the live kernel set when the engine can list it.
func currentlyBlocked(tracked []blockedIP, live firewall.LiveBlockedSnapshot, useLive bool, h *netblockHistory, blocker IPBlocker) map[string]bool {
current := make(map[string]bool, len(tracked)+len(live.V4)+len(live.V6))
for _, b := range tracked {
current[b.IP] = true
}
if useLive {
for _, set := range []map[string]struct{}{live.V4, live.V6} {
for ip := range set {
current[ip] = true
}
}
}
// A missing family snapshot is unknown, not evidence that an operator's
// permanent block ended. Use the same cached fallback as tracker reconciliation.
for ip := range h.IPs {
if current[ip] {
continue
}
if _, known := live.Contains(ip); useLive && known {
continue
}
if blocker.IsBlocked(ip) {
current[ip] = true
}
}
return current
}
// netblockCounts groups addresses by subnet: every address blocked now, and
// every address whose block ended inside the window. A block that is still in
// place counts however long ago it began. Addresses the firewall now allows
// are no longer evidence.
func netblockCounts(cfg *config.Config, h *netblockHistory, current map[string]bool, blocker IPBlocker, now time.Time, window time.Duration) map[string]int {
allow, _ := blocker.(allowChecker)
counts := make(map[string]int)
for ip, at := range h.IPs {
if !current[ip] && now.Sub(at) > window {
continue
}
cidr := subnetEscalationCIDR(ip)
// Exempt IPs do not contribute toward the netblock threshold so that
// a cluster of blocked addresses inside an operator-declared DoS-exempt
// range cannot inadvertently auto-block that range as a subnet.
if cidr == "" || cidrIntersectsDOSExempt(cfg, cidr) {
continue
}
// Ended blocks that an earlier subnet block already answered must
// not re-block the subnet for the rest of the window. Blocks still
// in place count as before.
if last, ok := h.Subnets[cidr]; ok && !current[ip] && !at.After(last) {
continue
}
if allow != nil && allow.IsAllowed(ip) {
continue
}
counts[cidr]++
}
return counts
}
package checks
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var autoBlockQueues = newAutoBlockQueue()
type autoBlockQueueMonitor struct {
mu sync.Mutex
waiting map[*autoBlockStateWork]struct{}
active *autoBlockStateWork
retries *autoBlockRetryQueue
cleanup *autoBlockCleanupQueue
idleSince time.Time
waitingLoss, activeLoss *queuehealth.Tracker
}
type autoBlockStateWork struct {
queue *autoBlockQueueMonitor
at time.Time
failed, completed bool
readingState bool
retryCycle *autoBlockRetryCycle
cleanupCycle *autoBlockCleanupCycle
}
func newAutoBlockQueue() *autoBlockQueueMonitor {
return &autoBlockQueueMonitor{cleanup: newAutoBlockCleanupQueue(), retries: newAutoBlockRetryQueue(), waiting: make(map[*autoBlockStateWork]struct{}), waitingLoss: queuehealth.New(0, time.Minute), activeLoss: queuehealth.New(1, time.Minute)}
}
func (q *autoBlockQueueMonitor) acquire() *autoBlockStateWork {
w := &autoBlockStateWork{queue: q, at: time.Now()}
q.mu.Lock()
if q.active == nil && len(q.waiting) == 0 {
q.idleSince = w.at
}
q.waiting[w] = struct{}{}
q.mu.Unlock()
blockStateMu.Lock()
q.mu.Lock()
delete(q.waiting, w)
q.active = w
w.at = time.Now()
q.mu.Unlock()
return w
}
func (w *autoBlockStateWork) progress() {
w.queue.mu.Lock()
w.at = time.Now()
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) failLocked() {
if !w.failed {
w.failed = true
w.queue.activeLoss.Lose(time.Now(), 1)
}
}
func (w *autoBlockStateWork) observe(err error) {
if err == nil {
return
}
w.queue.mu.Lock()
w.failLocked()
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) complete() {
w.queue.mu.Lock()
w.completed = true
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) finish() {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
// Release the real slot with its owner, including when the accounting
// below fails. A successor must not publish over a callback whose cleanup
// is still running, and a failure must not strand every later block.
defer blockStateMu.Unlock()
if !w.completed {
w.failLocked()
}
w.finishRetriesLocked()
w.finishCleanupLocked()
q.active = nil
q.idleSince = time.Now()
}
// AutoBlockQueueStatuses reads queue memory without waiting for the state mutex,
// filesystem, firewall or database. A batch is timed by operation progress.
func AutoBlockQueueStatuses(now time.Time) map[string]queuehealth.Status {
q := autoBlockQueues
q.mu.Lock()
defer q.mu.Unlock()
waiting := q.waitingLoss.Snapshot(now)
waiting.CapacityUnavailable = true
active := q.activeLoss.Snapshot(now)
active.LagBasis = "operation_progress"
stalled := q.active == nil && !q.idleSince.IsZero() && now.Sub(q.idleSince) >= time.Minute
if w := q.active; w != nil {
active.InFlight = 1
active.ProcessingSeconds = max(0, now.Sub(w.at).Seconds())
if active.ProcessingSeconds >= time.Minute.Seconds() {
active.Status, active.Reason = "degraded", "processing_lag"
stalled = true
}
}
for w := range q.waiting {
waiting.Depth++
waiting.LagSeconds = max(waiting.LagSeconds, now.Sub(w.at).Seconds())
if stalled {
waiting.Status, waiting.Reason = "degraded", "backlog_lag"
}
}
pending, candidates := q.retries.statuses(now, q.active)
return map[string]queuehealth.Status{"waiting": waiting, "active": active, "pending": pending, "candidates": candidates, "cleanup": q.cleanup.status(now, q.active)}
}
package checks
import (
"errors"
"fmt"
"os"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/queuehealth"
)
var persistAutoBlockState = writeBlockState
// All fields use autoBlockQueueMonitor.mu. Disk identities live only as long
// as their records; an unreadable write outcome never accumulates old versions.
type autoBlockRetryQueue struct {
path string
disk []pendingIP
records []*autoBlockPendingRecord
candidates map[*autoBlockCandidate]struct{}
pendingLoss, candidateLoss *queuehealth.Tracker
depthKnown, historyUnknown, lowerBound bool
readFailed, writeFailed bool
}
type autoBlockPendingRecord struct {
at time.Time
inFlight, completed, lossKnown, failed bool
}
type autoBlockCandidate struct {
at time.Time
origins []*autoBlockPendingRecord
queued *autoBlockPendingRecord
running, completed, finished, failed, lost bool
}
type autoBlockRetryCycle struct {
saving, settled bool
}
func newAutoBlockRetryQueue() *autoBlockRetryQueue {
return &autoBlockRetryQueue{
candidates: make(map[*autoBlockCandidate]struct{}),
pendingLoss: queuehealth.New(maxPendingBlocks, maxPendingAge),
candidateLoss: queuehealth.New(0, time.Minute),
}
}
// InitAutoBlockQueueHealth observes existing retries before daemon consumers
// start, including when automatic blocking is disabled. It never changes state.
func InitAutoBlockQueueHealth(statePath string) error {
work := autoBlockQueues.acquire()
defer work.finish()
_, err := work.readState(statePath)
work.complete()
return err
}
func (w *autoBlockStateWork) readState(path string) (*blockState, error) {
w.progress()
q := w.queue
q.mu.Lock()
w.readingState = true
q.cleanup.setPath(path)
q.mu.Unlock()
state, err := readBlockState(path)
q.mu.Lock()
defer q.mu.Unlock()
w.readingState = false
r := q.retries
if r.path != path {
r.path = path
r.disk, r.records = nil, nil
r.depthKnown, r.historyUnknown = false, false
r.readFailed, r.writeFailed = false, false
}
w.retryCycle = &autoBlockRetryCycle{}
q.cleanup.readFailed = err != nil
r.readFailed = err != nil
if err != nil {
w.failLocked()
r.forgetUncertain()
w.forgetCleanupState()
} else {
r.reconcile(state.Pending, nil, time.Now())
w.observeCleanupState(state, false)
}
return state, err
}
func (w *autoBlockStateWork) loadState(path string) *blockState {
state, err := w.readState(path)
if err != nil {
fmt.Fprintf(os.Stderr, "autoblock: %v; ignoring queued blocks\n", err)
return &blockState{}
}
return state
}
type autoBlockPendingKey struct {
ip, check string
queuedAt time.Time
severity alert.Severity
}
func pendingRecordKey(p pendingIP) autoBlockPendingKey {
return autoBlockPendingKey{ip: p.IP, check: p.Check, queuedAt: p.QueuedAt.UTC(), severity: p.Severity}
}
func (r *autoBlockRetryQueue) reconcile(actual, proposed []pendingIP, now time.Time) {
pool := make(map[autoBlockPendingKey][]*autoBlockPendingRecord, len(r.disk)+len(proposed))
for _, version := range [][]pendingIP{r.disk, proposed} {
for _, p := range version {
if p.queueRecord != nil {
key := pendingRecordKey(p)
pool[key] = append(pool[key], p.queueRecord)
}
}
}
kept := make(map[*autoBlockPendingRecord]bool, len(actual))
records := make([]*autoBlockPendingRecord, 0, len(actual))
for i := range actual {
p := &actual[i]
record := p.queueRecord
if record == nil {
key := pendingRecordKey(*p)
available := pool[key]
for len(available) > 0 {
known := available[0]
available = available[1:]
if !kept[known] {
record = known
break
}
}
pool[key] = available
}
if record == nil {
at := p.QueuedAt
if at.IsZero() {
at = now
}
record = &autoBlockPendingRecord{at: at, lossKnown: !r.historyUnknown}
}
p.queueRecord = record
if !p.QueuedAt.IsZero() {
record.at = p.QueuedAt
}
record.inFlight = false
kept[record] = true
records = append(records, record)
}
// A duplicate coalesced into a retained retry still has an owner. A
// successful block remains completed even when its record survives rollback.
for c := range r.candidates {
survives := kept[c.queued]
for _, origin := range c.origins {
survives = survives || kept[origin]
}
if c.completed || survives {
for _, origin := range c.origins {
if c.completed || !kept[origin] {
origin.completed = true
}
}
} else if len(c.origins) == 0 && !c.lost {
r.candidateLoss.Lose(now, 1)
c.lost = true
}
c.finished = true
}
for _, old := range r.disk {
record := old.queueRecord
if !kept[record] && !record.completed && record.lossKnown {
r.pendingLoss.Lose(now, 1)
}
}
r.records = records
r.disk = append([]pendingIP(nil), actual...)
for i := range r.disk {
r.disk[i].queueCandidate = nil
}
r.depthKnown = true
}
func (r *autoBlockRetryQueue) forgetUncertain() {
r.disk, r.records = nil, nil
r.depthKnown, r.historyUnknown, r.lowerBound = false, true, true
for c := range r.candidates {
c.finished = true
}
}
func (w *autoBlockStateWork) beginPending(p pendingIP) {
w.queue.mu.Lock()
p.queueRecord.inFlight = true
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) completePending(p pendingIP) {
w.queue.mu.Lock()
p.queueRecord.completed = true
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) candidate(p pendingIP, existing *autoBlockCandidate) *autoBlockCandidate {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
c := existing
if c == nil {
at := p.QueuedAt
if at.IsZero() {
at = time.Now()
}
c = &autoBlockCandidate{at: at}
q.retries.candidates[c] = struct{}{}
}
if p.queueRecord != nil {
c.origins = append(c.origins, p.queueRecord)
if p.queueRecord.at.Before(c.at) {
c.at = p.queueRecord.at
}
}
return c
}
func (w *autoBlockStateWork) startCandidate(c *autoBlockCandidate) {
w.queue.mu.Lock()
c.running = true
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) candidateOutcome(c *autoBlockCandidate, err error) {
success := err == nil || errors.Is(err, firewall.ErrIPProtected)
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
c.completed = c.completed || success
c.failed = !success
for _, origin := range c.origins {
origin.completed = origin.completed || success
origin.failed = !success
}
}
// A direct source can complete an eligible retry without consuming it from
// the scan queue. Preserve that acknowledgment before secondary bookkeeping.
func (w *autoBlockStateWork) directOutcome(ip string, attemptAt time.Time, err error) {
if err != nil && !errors.Is(err, firewall.ErrIPProtected) {
return
}
ip = normalizeBlockIP(ip)
if ip == "" {
return
}
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
for _, p := range q.retries.disk {
if normalizeBlockIP(p.IP) == ip && (p.QueuedAt.IsZero() || attemptAt.Sub(p.QueuedAt) <= maxPendingAge) {
p.queueRecord.completed = true
p.queueRecord.failed = false
}
}
}
func (w *autoBlockStateWork) finishCandidate(c *autoBlockCandidate) {
w.queue.mu.Lock()
c.finished = true
w.queue.mu.Unlock()
}
func (w *autoBlockStateWork) rejectCandidate(c *autoBlockCandidate) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
if len(c.origins) == 0 && !c.completed && !c.lost {
q.retries.candidateLoss.Lose(time.Now(), 1)
c.lost = true
}
c.finished = true
}
func (w *autoBlockStateWork) requeueCandidate(p pendingIP) pendingIP {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
c := p.queueCandidate
record := p.queueRecord
if record == nil {
record = &autoBlockPendingRecord{at: p.QueuedAt, lossKnown: true}
q.retries.records = append(q.retries.records, record)
}
record.inFlight = true
record.failed = c.failed || record.failed
p.queueRecord = record
c.queued = record
return p
}
func (w *autoBlockStateWork) writeState(path string, state *blockState) error {
q := w.queue
q.mu.Lock()
w.retryCycle.saving = true
q.mu.Unlock()
err := persistAutoBlockState(path, state)
q.mu.Lock()
q.retries.writeFailed = err != nil
q.cleanup.writeFailed = err != nil
if err != nil {
w.failLocked()
q.retries.depthKnown = false
q.cleanup.depthKnown = false
}
q.mu.Unlock()
// Rename may have succeeded before directory fsync returned an error.
// Only readback can distinguish a retained old record from a new one.
actual := state
var readErr error
if err != nil {
w.progress()
actual, readErr = readBlockState(path)
}
q.mu.Lock()
defer q.mu.Unlock()
if readErr != nil {
q.retries.readFailed = true
q.retries.forgetUncertain()
q.cleanup.readFailed = true
w.forgetCleanupState()
} else {
q.retries.reconcile(actual.Pending, state.Pending, time.Now())
w.observeCleanupState(actual, true)
q.retries.readFailed = false
q.cleanup.readFailed = false
if err == nil {
q.retries.historyUnknown = false
q.cleanup.historyUnknown = false
}
}
w.retryCycle.settled = true
if w.cleanupCycle != nil && readErr == nil {
w.cleanupCycle.settled = true
}
return err
}
func (w *autoBlockStateWork) saveState(path string, state *blockState) {
if err := w.writeState(path, state); err != nil {
logBlockStateFailure(path, err)
}
}
func (w *autoBlockStateWork) finishRetriesLocked() {
if w.readingState {
w.queue.retries.readFailed = true
w.queue.retries.forgetUncertain()
}
if w.retryCycle == nil {
return
}
r := w.queue.retries
if !w.completed {
for c := range r.candidates {
if c.running && !c.completed {
for _, origin := range c.origins {
origin.failed = true
}
}
}
}
if !w.retryCycle.settled {
switch {
case w.retryCycle.saving:
r.writeFailed = true
r.forgetUncertain()
case r.depthKnown:
r.reconcile(r.disk, nil, time.Now())
default:
// The original state was unreadable, but no write started. Fresh work
// accepted since that read cannot have reached the durable file.
for c := range r.candidates {
if !c.completed && !c.lost && len(c.origins) == 0 {
r.candidateLoss.Lose(time.Now(), 1)
c.lost = true
}
}
}
}
clear(r.candidates)
}
func (r *autoBlockRetryQueue) statuses(now time.Time, active *autoBlockStateWork) (queuehealth.Status, queuehealth.Status) {
pending := r.pendingLoss.Snapshot(now)
pending.DroppedLowerBound = r.lowerBound
pending.DepthUnavailable = !r.depthKnown
pending.LagBasis = "queued_or_observed_age"
if !r.depthKnown {
pending.LagBasis = "unavailable"
}
candidates := r.candidateLoss.Snapshot(now)
candidates.CapacityUnavailable = true
candidates.DroppedLowerBound = r.lowerBound
candidates.LagBasis = "operation_progress"
processing := 0.0
if active != nil {
processing = max(0, now.Sub(active.at).Seconds())
}
failed := false
for _, record := range r.records {
if record.inFlight {
pending.InFlight++
pending.ProcessingSeconds = processing
} else if r.depthKnown {
pending.Depth++
pending.LagSeconds = max(pending.LagSeconds, now.Sub(record.at).Seconds())
}
failed = failed || record.failed
}
for c := range r.candidates {
if c.finished {
continue
}
if c.running {
candidates.InFlight++
candidates.ProcessingSeconds = processing
} else {
candidates.Depth++
candidates.LagSeconds = max(candidates.LagSeconds, now.Sub(c.at).Seconds())
}
failed = failed || c.failed
}
switch {
case r.readFailed || r.writeFailed:
pending.Status, pending.Reason = "degraded", "state_io"
case failed:
pending.Status, pending.Reason = "degraded", "retry_failed"
case pending.InFlight > 0 && processing >= time.Minute.Seconds():
pending.Status, pending.Reason = "degraded", "processing_lag"
case pending.LagSeconds > maxPendingAge.Seconds():
pending.Status, pending.Reason = "degraded", "backlog_lag"
}
if candidates.InFlight+candidates.Depth > 0 && processing >= time.Minute.Seconds() {
candidates.Status, candidates.Reason = "degraded", "processing_lag"
}
return pending, candidates
}
package checks
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
"syscall"
"time"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/processhandle"
)
// var (not const) so tests can redirect to t.TempDir().
var quarantineDir = "/opt/csm/quarantine"
// autoQuarantineChecks are the findings the scheduled auto-responder may
// quarantine on its own: the manual move set plus the kill-and-quarantine
// and handler-abuse families, and the realtime signature match, which must
// additionally pass isHighConfidenceRealtimeMatch. Membership is pinned by
// test against the check registry.
var autoQuarantineChecks = map[string]bool{
"webshell": true,
"backdoor_binary": true,
"new_webshell_file": true,
"new_executable_in_config": true,
"obfuscated_php": true,
"suspicious_php_content": true,
"new_php_in_languages": true,
"new_php_in_upgrade": true,
"phishing_page": true,
"phishing_directory": true,
"htaccess_handler_abuse": true,
"signature_match_realtime": true,
}
var signalProcess = processhandle.Signal
var errProcessNotEligible = errors.New("process is no longer eligible for termination")
// AutoKillProcesses kills processes that match critical findings.
// Only targets: fake kernel threads, reverse shells, GSocket processes.
// Never kills root system services or cPanel processes.
func AutoKillProcesses(ctx context.Context, cfg *config.Config, findings []alert.Finding) []alert.Finding {
if !cfg.AutoResponse.Enabled || !cfg.AutoResponse.KillProcesses {
return nil
}
var actions []alert.Finding
for _, f := range findings {
// Only act on specific high-confidence critical checks
switch f.Check {
case "fake_kernel_thread", "suspicious_process", "php_suspicious_execution":
default:
continue
}
if f.Severity != alert.Critical {
continue
}
// Use structured PID field when available, fall back to text extraction
pid := fmt.Sprintf("%d", f.PID)
if f.PID == 0 {
pid = extractPID(f.Details)
if pid == "" {
continue
}
}
pidInt, validPID := parseProcessPID(pid)
if !validPID {
continue
}
pid = strconv.Itoa(pidInt)
var uid, exe string
err := signalProcess(ctx, pidInt, syscall.SIGKILL, func() error {
uid, exe = getProcessUID(pid), getProcessExe(pid)
if uid == "0" || uid == "" || exe == "" || isSafeProcess(exe) || !processStartedBefore(pid, f.Timestamp) {
return errProcessNotEligible
}
return nil
})
if err != nil {
if !errors.Is(err, errProcessNotEligible) && !errors.Is(err, os.ErrProcessDone) && ctx.Err() == nil {
csmlog.Warn("auto-kill: safe process signaling failed", "pid", pidInt, "err", err)
}
recordKillAction(&f, pid, exe, err)
continue
}
recordKillAction(&f, pid, exe, nil)
actions = append(actions, alert.Finding{
Severity: alert.Critical,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-KILL: Process %s killed (was: %s)", pid, f.Check),
Timestamp: time.Now(),
Details: fmt.Sprintf("Original finding: %s\nProcess: %s (UID: %s)", f.Message, exe, uid),
})
}
return actions
}
// recordKillAction writes the unified action record for one termination
// attempt. A refusal is recorded as well as a kill: "the safety rules stopped
// this" is the answer to a question an operator will ask about a process that
// is still running.
func recordKillAction(f *alert.Finding, pid, exe string, err error) {
rec := actionlog.Record{
Op: "respond.kill_process",
Actor: actionlog.DefaultActor(),
Target: "pid " + pid,
ActorDetail: exe,
Reason: "manual process termination",
FindingID: "",
Result: actionlog.Applied,
}
if f != nil {
rec.Reason = f.Check
rec.FindingID = alert.FindingID(*f)
}
switch {
case errors.Is(err, errProcessNotEligible):
rec.Result = actionlog.Refused
rec.Error = "process is not eligible for automatic termination"
case errors.Is(err, os.ErrProcessDone):
rec.Result = actionlog.Refused
rec.Error = "process had already exited"
case err != nil:
rec.Result = actionlog.Failed
rec.Error = err.Error()
}
actionlog.Write(rec)
}
// AutoQuarantineFiles moves malicious files to quarantine directory.
// Preserves original path and metadata in a sidecar .meta file.
// Marks evaluated input findings so alert delivery cannot repeat a response.
func AutoQuarantineFiles(cfg *config.Config, findings []alert.Finding) []alert.Finding {
if cfg == nil || !cfg.AutoResponse.Enabled || !cfg.AutoResponse.QuarantineFiles || cfg.ObserveMode() {
return nil
}
var actions []alert.Finding
seen := make(map[string]bool)
for i, f := range findings {
if f.AutoFileResponseEvaluated || !autoQuarantineChecks[f.Check] || f.Severity != alert.Critical {
continue
}
path := f.FilePath
if path == "" {
path = extractFilePath(f.Message)
}
if path == "" {
continue
}
findings[i].AutoFileResponseEvaluated = true
key := filepath.Clean(path)
if seen[key] {
continue
}
realtime := f.Check == "signature_match_realtime"
if realtime && !isHighConfidenceRealtimeMatch(f, path, nil) {
continue
}
info, err := osFS.Lstat(path)
if err != nil || info.Mode()&os.ModeSymlink != 0 {
continue
}
// Multiple checks can report one file. Do not re-clean a repaired
// target or charge repeated failures for the same batch of evidence.
seen[key] = true
paused := runAutoFileResponse(cfg, path, info, func() error {
// Cleaning is one response attempt. A failed cleaner leaves the file
// and any backup for review; it must not escalate to removing the file.
if !realtime && ShouldCleanInsteadOfQuarantine(path) {
result := cleanInfectedFileIdentified(path, info)
if result.Error != "" {
outcome := "failed"
if result.Refused {
outcome = "refused"
}
actions = append(actions, alert.Finding{Severity: alert.Warning, Check: "auto_response", Message: fmt.Sprintf("AUTO-CLEAN %s for %s; manual review required", outcome, path), Details: result.Error, Timestamp: time.Now()})
// Safety refusals consume capacity without charging a failure.
if result.Refused {
return nil
}
return errors.New(result.Error)
}
if result.Cleaned {
actions = append(actions, alert.Finding{Severity: alert.Critical, Check: "auto_response", Message: fmt.Sprintf("AUTO-CLEAN: %s surgically cleaned", path), Details: fmt.Sprintf("Backup: %s\n%s", result.BackupPath, strings.Join(result.Removals, "\n")), Timestamp: time.Now()})
}
return nil
}
qPath := newQuarantinePath(quarantineDir, path)
meta := quarantineMetadata(path, info, f.Message)
meta.FindingID = alert.FindingID(f)
err := quarantineTarget(path, qPath, info, meta)
warning := ""
if err != nil {
var completed bool
warning, completed = completedQuarantineWarning(err)
if !completed {
return err
}
}
details := fmt.Sprintf("Quarantined to: %s\nOriginal finding: %s", qPath, f.Message)
if warning != "" {
details += "\nWarning: " + warning
}
actions = append(actions, alert.Finding{Severity: alert.Critical, Check: "auto_response", Message: fmt.Sprintf("AUTO-QUARANTINE: %s moved to quarantine", path), Timestamp: time.Now(), Details: details})
return nil
})
if paused != nil {
actions = append(actions, *paused)
}
}
return actions
}
// AutoFixPermissions sets world/group-writable PHP files to 0644.
// Returns the auto-response action findings and the keys of original findings
// that were successfully fixed (so the caller can dismiss them from the UI).
func AutoFixPermissions(cfg *config.Config, findings []alert.Finding) (actions []alert.Finding, fixedKeys []string) {
if !cfg.AutoResponse.Enabled || !cfg.AutoResponse.EnforcePermissions {
return nil, nil
}
for _, f := range findings {
switch f.Check {
case "world_writable_php", "group_writable_php":
default:
continue
}
path := extractFilePath(f.Message)
if path == "" {
continue
}
path, info, err := resolveExistingFixPath(path, effectiveFixRoots(fixPermissionsAllowedRoots))
if err != nil || info.IsDir() {
continue
}
oldMode := info.Mode().Perm()
// #nosec G302 -- same as fixPermissions: restoring canonical web-content
// mode on a user file flagged for dangerous (e.g. world-writable) perms.
if err := os.Chmod(path, 0644); err != nil {
continue
}
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-FIX: %s permissions set to 644 (was %o)", path, oldMode),
Timestamp: time.Now(),
})
fixedKeys = append(fixedKeys, f.Check+":"+f.Message)
}
return actions, fixedKeys
}
// AutoFixWPCron disables WP-Cron and installs a per-user system cron for every
// perf_wp_cron finding. Returns the auto-response action findings and the keys
// of the originals so the caller can dismiss them. Gated behind an explicit
// opt-in because it edits customer wp-config.php and crontabs.
func AutoFixWPCron(cfg *config.Config, findings []alert.Finding) (actions []alert.Finding, fixedKeys []string) {
if !cfg.AutoResponse.Enabled || !cfg.AutoResponse.FixWPCron {
return nil, nil
}
opts := WPCronFixOptions{
IntervalMinutes: cfg.Performance.WPCronFix.IntervalMinutes,
PHPBin: cfg.Performance.WPCronFix.PHPBin,
}
allowedRoots := ResolveWPCronRoots(cfg)
for _, f := range findings {
if f.Check != "perf_wp_cron" {
continue
}
path := extractWPConfigPath(f.Details)
if path == "" {
continue
}
res := FixDisableWPCronInRoots(path, allowedRoots, opts)
if !res.Success {
continue
}
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-FIX: %s", res.Description),
Timestamp: time.Now(),
})
fixedKeys = append(fixedKeys, f.Key())
}
return actions, fixedKeys
}
// extractWPConfigPath pulls the wp-config.php path out of a perf_wp_cron
// finding's Details, formatted as "File: <path> - add define(...)".
func extractWPConfigPath(details string) string {
const prefix = "File: "
idx := strings.Index(details, prefix)
if idx < 0 {
return ""
}
rest := details[idx+len(prefix):]
if j := strings.Index(rest, " - "); j >= 0 {
rest = rest[:j]
}
return strings.TrimSpace(rest)
}
func extractPID(details string) string {
// Look for "PID: 12345" pattern. Stop at the first whitespace, comma,
// or newline so a trailing word ("PID: 42 exe=/bin/ls") doesn't get
// returned as part of the PID string.
idx := strings.Index(details, "PID: ")
if idx < 0 {
return ""
}
rest := details[idx+5:]
for i, c := range rest {
if c == ' ' || c == ',' || c == '\n' || c == '\t' {
return strings.TrimSpace(rest[:i])
}
}
return strings.TrimSpace(rest)
}
func extractFilePath(message string) string {
// Look for /home/... or /tmp/... paths in the message. Order matters:
// longer/more-specific prefixes (/var/tmp/, /dev/shm/) must come BEFORE
// shorter ones (/tmp/) — otherwise "/tmp/" would match inside "/var/tmp/"
// and we'd silently misclassify the path.
for _, prefix := range accountRootPrefixes("/var/tmp/", "/dev/shm/", "/tmp/") {
if idx := strings.Index(message, prefix); idx >= 0 {
rest := message[idx:]
// Path ends at space, comma, or end of string
endIdx := len(rest)
for i, c := range rest {
if c == ' ' || c == ',' || c == '\n' {
endIdx = i
break
}
}
return rest[:endIdx]
}
}
return ""
}
func getProcessUID(pid string) string {
data, err := osFS.ReadFile(filepath.Join("/proc", pid, "status"))
if err != nil {
return ""
}
for _, line := range strings.Split(string(data), "\n") {
if rest, found := strings.CutPrefix(line, "Uid:"); found {
fields := strings.Fields(rest)
if len(fields) != 4 {
return ""
}
var owner uint64
for index, field := range fields {
uid, err := strconv.ParseUint(field, 10, 32)
if err != nil {
return ""
}
// Effective, saved, and filesystem root credentials are also
// privileged even when the real UID still names a tenant.
if uid == 0 {
return "0"
}
if index == 0 {
owner = uid
}
}
return strconv.FormatUint(owner, 10)
}
}
return ""
}
func getProcessExe(pid string) string {
exe, err := osFS.Readlink(filepath.Join("/proc", pid, "exe"))
if err != nil {
return ""
}
return exe
}
func isSafeProcess(exe string) bool {
safePrefixes := []string{
"/usr/local/cpanel/",
"/usr/sbin/",
"/usr/bin/",
"/usr/libexec/",
"/opt/cpanel/",
"/opt/cloudlinux/",
"/opt/imunify360/",
}
for _, prefix := range safePrefixes {
if strings.HasPrefix(exe, prefix) {
return true
}
}
return false
}
// isHighConfidenceRealtimeMatch validates whether a realtime signature match
// is truly malicious and safe to auto-quarantine. Prevents false positives
// on legitimate libraries (PHPMailer, zip) and theme code.
//
// The data parameter should be the file content already read by the caller
// (fanotify fd or scanner) to avoid TOCTOU re-reads. Pass nil to read from path.
//
// Criteria:
// 1. Category must be "dropper" or "webshell"
// 2. File must be >= 512 bytes (entropy unreliable below this)
// 3. Content must show obfuscation indicators:
// Shannon entropy >= 5.5 OR hex density > 20% plus an execution signal.
// This applies to BOTH dropper and webshell categories to avoid
// false-positive quarantine of legitimate plugins that happen to
// match a dropper rule (e.g. curl_exec + eval on distant lines).
func isHighConfidenceRealtimeMatch(f alert.Finding, path string, data []byte) bool {
cat := extractCategory(f.Details)
switch cat {
case "dropper", "webshell":
default:
return false
}
if data == nil {
var err error
data, err = osFS.ReadFile(path)
if err != nil {
return false
}
}
if len(data) < 512 {
return false
}
// Both dropper and webshell categories go through the same entropy/encoding
// checks to avoid false positives. Normal PHP with heavy class constants
// (binary literals, many named constants) lands in the 4.5-5.3 range --
// measured: WPML wpml_zip.php = 5.25, Breakdance google-fonts.php = 4.90.
// The 5.5 floor leaves that headroom. Obfuscated packers land at 5.8+.
// Files that use long hex-string payloads (LEVIATHAN signature) are
// caught by the hex-density arm instead, which stays at 20%.
content := string(data)
// High Shannon entropy is a strong standalone signal: packed/encrypted
// payloads land at 5.8+, while ordinary library code (even with a handful
// of binary constants) stays below 5.5 -- measured WPML wpml_zip.php =
// 5.25, Breakdance google-fonts.php = 4.90.
if shannonEntropy(content) >= 5.5 {
return true
}
// High hex density alone is NOT enough: a ZIP/PDF library's magic-byte
// constant tables ("\x50\x4b\x03\x04" ...) saturate that metric while being
// inert data. Require a structural obfuscated-execution signal too. This
// replaces a hardcoded library-path allowlist (vendor/, node_modules/,
// named plugin slugs) that an attacker could defeat by planting a webshell
// under any "trusted" directory -- the file is now judged by content, so a
// hex-encoded packer still quarantines wherever it hides and a benign
// data-heavy library file is spared on any path.
if hexEncodingDensity(content) > 0.20 {
return hasObfuscatedExecutionSignal(content)
}
return false
}
var (
reVariableFunctionCall = regexp.MustCompile(`\$[A-Za-z_]\w*\s*\(`)
reHexEscapedStringConcat = regexp.MustCompile(`(?i)"(?:\\x[0-9a-f]{2})+"\s*\.\s*"(?:\\x[0-9a-f]{2})+"`)
)
// hasObfuscatedExecutionSignal reports whether content carries a structural
// sign of obfuscated code execution, distinguishing a packed webshell from
// inert binary data such as a ZIP library's magic-byte constants. Any signal
// is sufficient:
// - LEVIATHAN-style control-flow obfuscation (goto spaghetti).
// - Function names built from concatenated hex escapes ("\x65"."\x76"... to
// dodge literal-name detection).
// - A variable bound to a decoder/exec primitive and later invoked.
// - A decoder (base64/gz/rot13/openssl/hex2bin) paired with an executor
// (eval/assert/create_function, a literal dangerous callback, or a
// request-scoped variable-function call).
func hasObfuscatedExecutionSignal(content string) bool {
code := stripPHPCommentsFromCode(content)
codeNoStrings := strings.ToLower(stripPHPStringsFromCode(code))
if countOccurrences(codeNoStrings, "goto ") > 10 {
return true
}
// Function-name obfuscation: many double-quoted hex string literals joined
// by the concatenation operator. Standalone hex constant tables (no concat)
// are inert data and do not match.
if countHexEscapedStringConcats(code) > 10 {
return true
}
if detectVarFuncDangerousAssignment(code) {
return true
}
if !containsDirectPHPFunctionCall(codeNoStrings, []string{
"base64_decode", "gzinflate", "gzuncompress", "gzdecode",
"str_rot13", "openssl_decrypt", "hex2bin", "convert_uudecode",
}) {
return false
}
if containsDirectPHPFunctionCall(codeNoStrings, []string{
"eval", "assert", "create_function",
}) {
return true
}
if hasLiteralCallbackExecutor(code) {
return true
}
return hasRequestScopedVariableFunctionCall(codeNoStrings)
}
func countHexEscapedStringConcats(code string) int {
return len(reHexEscapedStringConcat.FindAllStringIndex(code, -1))
}
func containsDirectPHPFunctionCall(codeNoStrings string, names []string) bool {
for _, name := range names {
if containsStandaloneFunc(codeNoStrings, name+"(") {
return true
}
}
return false
}
func hasRequestScopedVariableFunctionCall(codeNoStrings string) bool {
for _, line := range strings.Split(codeNoStrings, "\n") {
if containsRequestSuperglobal(line) && reVariableFunctionCall.MatchString(line) {
return true
}
}
return false
}
func hasLiteralCallbackExecutor(code string) bool {
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
nameStart := i
if code[i] == '\\' {
if i+1 >= len(code) || !isPHPIdentifierStart(code[i+1]) || !canStartGlobalPHPFunction(code, i) {
continue
}
nameStart = i + 1
} else if !isPHPIdentifierStart(code[i]) || !canStartPHPFunctionName(code, i) {
continue
}
nameEnd := nameStart + 1
for nameEnd < len(code) && isPHPIdentifierPart(code[nameEnd]) {
nameEnd++
}
name := strings.ToLower(code[nameStart:nameEnd])
if _, ok := callbackFirstArgFuncs[name]; !ok {
i = nameEnd - 1
continue
}
openParen := skipPHPWhitespace(code, nameEnd)
if openParen >= len(code) || code[openParen] != '(' {
i = nameEnd - 1
continue
}
firstArg := skipPHPWhitespace(code, openParen+1)
if firstArg >= len(code) || !isPHPQuote(code[firstArg]) {
i = nameEnd - 1
continue
}
callbackName, _, ok := readPHPFunctionString(code, firstArg)
if !ok {
i = nameEnd - 1
continue
}
if _, dangerous := callbackExecNames[callbackName]; dangerous {
return true
}
i = nameEnd - 1
}
return false
}
// hexEncodingDensity returns the fraction of a string's bytes that are part
// of PHP hex escape sequences (\xNN). LEVIATHAN AES-encrypted webshells
// encode their payload as long hex strings - the \x prefix repeats so
// frequently that Shannon entropy drops to ~3.5 (below normal PHP), but
// the hex density reaches 40-60%.
func hexEncodingDensity(s string) float64 {
if len(s) == 0 {
return 0
}
hexBytes := 0
for i := 0; i < len(s)-3; i++ {
if s[i] == '\\' && s[i+1] == 'x' &&
isHexDigit(s[i+2]) && isHexDigit(s[i+3]) {
hexBytes += 4
i += 3 // skip past this sequence
}
}
return float64(hexBytes) / float64(len(s))
}
func isHexDigit(b byte) bool {
return (b >= '0' && b <= '9') || (b >= 'a' && b <= 'f') || (b >= 'A' && b <= 'F')
}
// InlineQuarantineGated applies the operator's quarantine policy before
// InlineQuarantine moves anything. The realtime fanotify path detects malware
// continuously, but moving a file is a customer-impacting auto-response
// action: it must honor the same master switch and quarantine opt-in as the
// batch AutoQuarantineFiles dispatcher, never act on detection alone. An
// operator in monitor mode (auto-response off, or quarantine_files off) gets
// the alert without having files moved out from under them.
func InlineQuarantineGated(cfg *config.Config, f alert.Finding, path string, data []byte) (string, bool) {
path, ok, _ := InlineQuarantineGatedIdentified(cfg, &f, path, data, nil)
return path, ok
}
// InlineQuarantineGatedIdentified applies the auto-response policy gate and
// then quarantines the exact file the caller scanned. See
// InlineQuarantineIdentified for why the identity matters.
// Marks the finding evaluated only once it reaches the shared budget gate.
func InlineQuarantineGatedIdentified(cfg *config.Config, f *alert.Finding, path string, data []byte, scanned os.FileInfo) (string, bool, *alert.Finding) {
if f == nil || cfg == nil || !cfg.AutoResponse.Enabled || !cfg.AutoResponse.QuarantineFiles || cfg.ObserveMode() {
return "", false, nil
}
info, ok := inlineQuarantineInfo(*f, path, data, scanned)
if !ok {
return "", false, nil
}
var qPath string
f.AutoFileResponseEvaluated = true
paused := runAutoFileResponse(cfg, path, info, func() error {
var err error
qPath, err = quarantineInlineTarget(*f, path, info)
return err
})
return qPath, qPath != "", paused
}
// InlineQuarantine moves a file to quarantine immediately if it passes the
// high-confidence validation gates. Called from fanotify's analyzeFile to
// quarantine malware without waiting for the 5-second batch dispatcher.
// The data parameter is the file content already read by the caller (avoids
// TOCTOU re-read). Pass nil to read from path.
// Returns the quarantine path and true if the file was quarantined.
func InlineQuarantine(f alert.Finding, path string, data []byte) (string, bool) {
return InlineQuarantineIdentified(f, path, data, nil)
}
// InlineQuarantineIdentified is InlineQuarantine with the identity of the file
// the caller actually scanned. The realtime scanner reads content from the
// fanotify event descriptor, so passing that descriptor's stat pins the move to
// the object that was examined: a file replaced between detection and
// quarantine fails the identity check instead of being moved in place of the
// malware. A nil identity keeps the older path-based behaviour for callers that
// began from a path in the first place, such as the batch dispatcher.
func InlineQuarantineIdentified(f alert.Finding, path string, data []byte, scanned os.FileInfo) (string, bool) {
info, ok := inlineQuarantineInfo(f, path, data, scanned)
if !ok {
return "", false
}
qPath, err := quarantineInlineTarget(f, path, info)
if err != nil {
csmlog.Warn("inline quarantine refused", "path", path, "err", err)
}
return qPath, err == nil
}
func inlineQuarantineInfo(f alert.Finding, path string, data []byte, scanned os.FileInfo) (os.FileInfo, bool) {
if !isHighConfidenceRealtimeMatch(f, path, data) {
return nil, false
}
info, err := osFS.Lstat(path)
if err != nil || info.Mode()&os.ModeSymlink != 0 {
return nil, false
}
if scanned != nil && (!sameFileIdentity(info, scanned) || !sameContentShape(info, scanned)) {
return nil, false
}
return info, true
}
func quarantineInlineTarget(f alert.Finding, path string, info os.FileInfo) (string, error) {
qPath := newQuarantinePath(quarantineDir, path)
meta := quarantineMetadata(path, info, "Inline quarantine: high-confidence realtime signature match")
meta.FindingID = alert.FindingID(f)
if err := quarantineTarget(path, qPath, info, meta); err != nil {
if warning, completed := completedQuarantineWarning(err); completed {
csmlog.Warn("inline quarantine completed with warning", "warning", warning)
} else {
return "", err
}
}
return qPath, nil
}
// extractCategory parses "Category: <value>" from a finding's Details field.
func extractCategory(details string) string {
for _, line := range strings.Split(details, "\n") {
if strings.HasPrefix(line, "Category: ") {
return strings.TrimPrefix(line, "Category: ")
}
}
return ""
}
// AutoCleanHtaccess runs the hardened .htaccess cleaner against
// every finding emitted by the new detector registry, gated by
// AutoResponse.CleanHtaccess. Skipped when the daemon's auto-response
// pipeline is disabled overall.
//
// Unlike AutoQuarantineFiles, this routes around the
// quarantine/clean fork (.htaccess files are infrastructure -- moving
// them to /opt/csm/quarantine breaks the site). Each invocation
// backs up the original to /opt/csm/quarantine/pre_clean/<ts>_*
// inside CleanHtaccessFile before atomic-replacing.
// Marks evaluated input findings so alert delivery cannot repeat a response.
func AutoCleanHtaccess(cfg *config.Config, findings []alert.Finding) []alert.Finding {
if cfg == nil || !cfg.AutoResponse.Enabled || !cfg.AutoResponse.CleanHtaccess || cfg.ObserveMode() {
return nil
}
var actions []alert.Finding
seen := make(map[string]struct{})
for i, f := range findings {
if f.AutoFileResponseEvaluated || !isHtaccessHardenedFinding(f.Check) {
continue
}
path := f.FilePath
if path == "" {
path = extractFilePath(f.Message)
}
if path == "" {
continue
}
findings[i].AutoFileResponseEvaluated = true
// One Clean per file per autoresponse pass: multiple
// detector findings on the same file converge on a single
// cleaning call (CleanHtaccessFile re-runs every detector).
key := filepath.Clean(path)
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
info, err := osFS.Lstat(path)
if err != nil || info.Mode()&os.ModeSymlink != 0 {
continue
}
paused := runAutoFileResponse(cfg, path, info, func() error {
result := cleanHtaccessFileIdentified(path, info)
if result.Success {
actions = append(actions, alert.Finding{
Severity: alert.Critical,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-CLEAN: %s hardened directives removed", path),
Details: result.Description,
Timestamp: time.Now(),
})
} else if result.Error != "" && !result.Refused {
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-CLEAN failed: %s", path),
Details: result.Error,
Timestamp: time.Now(),
})
}
if result.Error != "" && !result.Refused {
return errors.New(result.Error)
}
return nil
})
if paused != nil {
actions = append(actions, *paused)
}
}
return actions
}
func isHtaccessHardenedFinding(check string) bool {
for _, detector := range htaccessDetectors {
if check == detector.Name {
return true
}
}
return false
}
package checks
import (
"fmt"
"net"
"sync/atomic"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
)
type asnLookupFunc func(ip string) (asn uint, org string)
type asnLookupHolder struct {
fn asnLookupFunc
}
// asnLookup resolves an IP to its autonomous system number and organization
// via the GeoLite2-ASN database. The daemon injects it at startup/reload;
// nil (no ASN database, or unit tests that do not exercise the live path)
// disables bad-ASN classification on the outbound connection scan.
var asnLookup atomic.Pointer[asnLookupHolder]
// SetASNLookup wires the GeoLite2-ASN resolver used by the outbound
// connection scan. Passing nil clears it.
func SetASNLookup(fn func(ip string) (asn uint, org string)) {
if fn == nil {
asnLookup.Store(nil)
return
}
asnLookup.Store(&asnLookupHolder{fn: fn})
}
// CurrentASNLookup returns the wired ASN resolver, or nil when none is set.
// Both the polling connection scan and the live BPF connection evaluator
// use it so bad-ASN classification behaves identically on either path.
func CurrentASNLookup() func(ip string) (asn uint, org string) {
h := asnLookup.Load()
if h == nil {
return nil
}
return h.fn
}
// EvaluateBadASNOutbound classifies one outbound connection's destination by
// autonomous system and returns a finding when the ASN is bad. It is a pure
// function: the caller supplies the destination IP and the ASN/org already
// resolved from the GeoLite2-ASN database, so the classifier has no IO and
// is the third leg of the host-takeover chain correlator.
//
// Classification:
// - blocked_asns always flags (e.g. known bulletproof hosters);
// - when allowed_asns is non-empty, any ASN outside it flags (allowlist
// mode for hosts whose legitimate egress is confined to a few providers).
//
// An ASN of 0 (no AS found for the IP) is skipped: classifying it would flag
// every destination missing from the ASN database. Private, loopback,
// link-local, and unspecified destinations are skipped because ASN lookup is
// meaningless for them.
func EvaluateBadASNOutbound(cfg *config.Config, dstIP net.IP, asn uint, asOrg string) (alert.Finding, bool) {
if cfg == nil || !cfg.Detection.BadASNOutbound.Enabled {
return alert.Finding{}, false
}
if dstIP == nil || dstIP.IsLoopback() || dstIP.IsUnspecified() ||
dstIP.IsPrivate() || dstIP.IsLinkLocalUnicast() || dstIP.IsLinkLocalMulticast() {
return alert.Finding{}, false
}
if asn == 0 {
return alert.Finding{}, false
}
if !asnIsBad(cfg, asn) {
return alert.Finding{}, false
}
org := asOrg
if org == "" {
org = "unknown organization"
}
dst := dstIP.String()
if dstIP.To4() == nil {
dst = "[" + dst + "]"
}
return alert.Finding{
Severity: alert.High,
Check: "bad_asn_outbound",
Message: fmt.Sprintf("Outbound connection to bad ASN %d (%s): %s", asn, org, dst),
Details: fmt.Sprintf("Destination: %s\nASN: %d (%s)\n"+
"Combined with a new uid-0 account or a planted suid binary this escalates to a host takeover.",
dst, asn, org),
SourceIP: dstIP.String(),
}, true
}
// asnIsBad applies the blocklist-then-allowlist policy to a single ASN.
func asnIsBad(cfg *config.Config, asn uint) bool {
for _, b := range cfg.Detection.BadASNOutbound.BlockedASNs {
if b == asn {
return true
}
}
allowed := cfg.Detection.BadASNOutbound.AllowedASNs
if len(allowed) == 0 {
return false
}
for _, a := range allowed {
if a == asn {
return false
}
}
return true
}
package checks
import (
"errors"
"sync"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/metrics"
)
// Operator block sources, alongside the auto-response sources in
// applyblock.go. Operator paths bypass the chokepoint (they force-block and
// audit separately) but report into the same outcome metric.
const (
BlockSourceCLI = "cli"
BlockSourceWebUI = "web_ui"
blockSourceUnknown = "unknown"
blockOutcomeError = "error"
blockOutcomeProtected = "protected"
)
var (
blockOutcomeMetric *metrics.CounterVec
blockOutcomeMetricOnce sync.Once
blockOutcomeLabels = [...]string{
string(firewall.BlockOutcomeLive),
string(firewall.BlockOutcomeDryRun),
string(firewall.BlockOutcomeAllowed),
string(firewall.BlockOutcomeAllowlisted),
string(firewall.BlockOutcomeNoop),
blockOutcomeProtected,
blockOutcomeError,
}
blockSourceLabels = [...]string{
BlockSourceScan,
BlockSourceChallenge,
BlockSourceIncident,
BlockSourceCentral,
BlockSourceCLI,
BlockSourceWebUI,
blockSourceUnknown,
}
)
func init() {
// An aliveness alert needs zero-value series before the first attempt;
// otherwise Prometheus treats a dead block path as missing data.
blockOutcomeCounter()
}
// blockOutcomeCounter registers the outcome metric once and creates every
// closed label combination so zero-attempt paths remain visible to scrapes.
func blockOutcomeCounter() *metrics.CounterVec {
blockOutcomeMetricOnce.Do(func() {
blockOutcomeMetric = metrics.NewCounterVec(
"csm_firewall_block_outcome_total",
"Firewall IP block attempts by outcome (live, dry_run, allowed, allowlisted, noop, protected, error) and source (scan, challenge, incident, central_intel, cli, web_ui, unknown).",
[]string{"outcome", "source"},
)
// CounterVec reserves one child for overflow, so allow the complete
// closed set plus that defensive sentinel and no further growth.
blockOutcomeMetric.SetMaxChildren(len(blockOutcomeLabels)*len(blockSourceLabels) + 1)
for _, outcome := range blockOutcomeLabels {
for _, source := range blockSourceLabels {
blockOutcomeMetric.With(outcome, source)
}
}
metrics.MustRegister("csm_firewall_block_outcome_total", blockOutcomeMetric)
})
return blockOutcomeMetric
}
// observeBlockOutcome counts one block attempt. Protected-IP refusals get
// their own label so expected no-ops do not read as failures on dashboards;
// every other error is outcome=error.
func observeBlockOutcome(outcome firewall.BlockOutcome, err error, source string) {
blockOutcomeCounter().With(blockOutcomeLabel(outcome, err), blockSourceLabel(source)).Inc()
}
func blockOutcomeLabel(outcome firewall.BlockOutcome, err error) string {
if errors.Is(err, firewall.ErrIPProtected) {
return blockOutcomeProtected
}
if err != nil {
return blockOutcomeError
}
switch outcome {
case firewall.BlockOutcomeLive,
firewall.BlockOutcomeDryRun,
firewall.BlockOutcomeAllowed,
firewall.BlockOutcomeAllowlisted,
firewall.BlockOutcomeNoop:
return string(outcome)
default:
// An implementation that returns an undocumented outcome violated the
// engine contract; report it as an error without creating a new label.
return blockOutcomeError
}
}
func blockSourceLabel(source string) string {
switch source {
case BlockSourceScan,
BlockSourceChallenge,
BlockSourceIncident,
BlockSourceCentral,
BlockSourceCLI,
BlockSourceWebUI:
return source
default:
return blockSourceUnknown
}
}
// ObserveOperatorBlock reports an operator-initiated force block (CLI or
// web UI) into the outcome metric. Force blocks bypass the dry-run gate, so
// a nil error means the block landed live.
func ObserveOperatorBlock(err error, source string) {
outcome := firewall.BlockOutcomeLive
observeBlockOutcome(outcome, err, source)
}
package checks
import (
"context"
"fmt"
"net"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// domlogDiscoveryDropped counts per-vhost log paths the discovery helper
// dropped silently (broken symlink, Stat failure). Same family of
// hidden-input bug as the lex-order issue scanDomlogs already fixed:
// without telemetry, operators have no way to notice when discovery
// loses a sizable fraction of the vhosts they expected to scan.
var (
domlogDiscoveryDropped *metrics.CounterVec
domlogDiscoveryDroppedOnce sync.Once
)
func observeDomlogDrop(reason string) {
domlogDiscoveryDroppedOnce.Do(func() {
domlogDiscoveryDropped = metrics.NewCounterVec(
"csm_checks_domlog_discovery_dropped_total",
"Per-vhost access-log paths the WP brute-force domlog discovery helper dropped before scanning. Labels: reason (evalsymlinks_error|stat_error). Steady growth means a chunk of vhosts is being silently skipped each cycle -- usually a broken symlink farm or a permissions regression on the log directory. Stale-mtime drops are intentional filtering, not counted here.",
[]string{"reason"},
)
metrics.MustRegister("csm_checks_domlog_discovery_dropped_total", domlogDiscoveryDropped)
})
domlogDiscoveryDropped.With(reason).Inc()
}
const (
wpLoginThreshold = 20 // attempts per IP across all logs
// xmlrpcThreshold is the fallback used only when emitLegacy is called with a
// nil config (tests). The live default comes from config
// (thresholds.xmlrpc_threshold, default DefaultXMLRPCThreshold); keep this in
// sync with that default.
xmlrpcThreshold = 100
ftpFailThreshold = 10
webmailThreshold = 10
apiFailThreshold = 10
// domlogTailLines is the built-in default for how many trailing lines
// to read from each domlog. Operators can override via
// cfg.Thresholds.DomlogTailLines. 500 covers ~10 minutes of traffic
// on a busy site.
domlogTailLines = 500
// domlogMaxAge skips domlogs not modified recently (inactive sites).
domlogMaxAge = 30 * time.Minute
// domlogMaxFiles caps the number of domlogs scanned per cycle
// to prevent unbounded I/O on servers with thousands of domains.
domlogMaxFiles = 500
)
// CheckWPBruteForce detects brute force attacks against wp-login.php and
// xmlrpc.php by scanning access logs. Always scans per-domain domlogs
// because on LiteSpeed+cPanel, virtual host traffic only appears there.
// The central access log is scanned as a supplement.
//
// Aggregates per-IP counts across ALL domains -- catches attackers who
// distribute requests across many sites to stay under per-site thresholds.
func CheckWPBruteForce(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
window := cfg.Thresholds.BruteForceWindow
if window <= 0 {
window = 5000
}
stats := newDomlogStats()
// 1. Per-domain domlogs -- primary source on LiteSpeed.
// Glob both SSL and non-SSL logs: attackers may use HTTP.
scanned := scanDomlogsStats(ctx, cfg, stats)
// 2. Central access log -- supplement for non-vhost traffic.
// On LiteSpeed this mostly has WHM/server-level requests.
// On Apache it duplicates domlog data; minor double-counting is
// acceptable since thresholds are high enough.
for _, p := range platform.Detect().AccessLogPaths {
lines := tailFile(p, window)
if len(lines) == 0 {
continue
}
for _, line := range lines {
rec, ok := parseAccessLogRecord(line)
if !ok {
continue
}
stats.scan(rec, cfg, currentBotClassifier(cfg))
}
break
}
findings := stats.emit(cfg)
// Replace the generic legacy Details with the actual scanned-file count.
for i := range findings {
if findings[i].Details == "Aggregated across per-vhost access logs" {
findings[i].Details = "Aggregated across " + itoa(scanned) + " per-vhost access logs"
}
}
return findings
}
// knownCentralAccessLogPaths lists central web-server log paths that may
// appear in broad per-vhost glob patterns. CheckWPBruteForce tails these on
// its own pass, so per-vhost discovery filters them to avoid counting the
// same lines twice.
var knownCentralAccessLogPaths = []string{
"/var/log/apache2/access.log",
"/var/log/apache2/access_log",
"/var/log/httpd/access.log",
"/var/log/httpd/access_log",
"/var/log/nginx/access.log",
"/usr/local/lsws/logs/access.log",
}
// discoverFreshDomlogs returns per-vhost access-log paths ready to tail.
// It globs platform.DomlogGlobs, dedupes by resolved-symlink real path,
// excludes the well-known central logs (so they are not double-counted),
// drops files untouched in the last maxAge, ranks survivors
// most-recent-first, and caps the result at maxFiles.
//
// Mtime-desc + cap is the fairness invariant: lexical glob order plus a
// hard cap would systematically hide brute force on late-alphabet
// domains. maxFiles <= 0 falls back to the built-in domlogMaxFiles
// default; maxAge <= 0 falls back to the built-in domlogMaxAge default.
// A canceled ctx returns nil and stops before any further work.
//
// Shared by scanDomlogs and scanDomlogsStats so the discovery semantics
// stay locked together; each caller layers its own per-line aggregator
// on top.
func discoverFreshDomlogs(ctx context.Context, maxFiles int, maxAge time.Duration) []string {
if maxFiles <= 0 {
maxFiles = domlogMaxFiles
}
if maxAge <= 0 {
maxAge = domlogMaxAge
}
if ctx == nil {
ctx = context.Background()
}
if err := ctx.Err(); err != nil {
return nil
}
platformInfo := platform.Detect()
globs := platformInfo.DomlogGlobs
centralLogs := centralAccessLogSet(platformInfo.AccessLogPaths)
var domlogs []string
for _, pattern := range globs {
if err := ctx.Err(); err != nil {
return nil
}
matches, _ := osFS.Glob(pattern)
domlogs = append(domlogs, matches...)
}
type domlogEntry struct {
path string
mtime time.Time
}
fresh := make([]domlogEntry, 0, len(domlogs))
seen := make(map[string]bool)
cutoff := time.Now().Add(-maxAge)
for _, dl := range domlogs {
if err := ctx.Err(); err != nil {
return nil
}
// Resolve symlinks first -- cPanel symlinks SSL and non-SSL
// logs to the same backing file; dedupe on the real path.
real, err := filepath.EvalSymlinks(dl)
if err != nil {
observeDomlogDrop("evalsymlinks_error")
continue
}
if seen[real] || centralLogs[real] {
continue
}
seen[real] = true
// Inactive sites add no signal; filter before the cap so they
// cannot crowd active sites out of the budget.
info, err := osFS.Stat(real)
if err != nil {
observeDomlogDrop("stat_error")
continue
}
if info.ModTime().Before(cutoff) {
continue
}
fresh = append(fresh, domlogEntry{path: real, mtime: info.ModTime()})
}
if err := ctx.Err(); err != nil {
return nil
}
sort.Slice(fresh, func(i, j int) bool {
if fresh[i].mtime.Equal(fresh[j].mtime) {
return fresh[i].path < fresh[j].path
}
return fresh[i].mtime.After(fresh[j].mtime)
})
if len(fresh) > maxFiles {
fresh = fresh[:maxFiles]
}
out := make([]string, len(fresh))
for i, e := range fresh {
out[i] = e.path
}
return out
}
func centralAccessLogSet(configured []string) map[string]bool {
out := make(map[string]bool, len(knownCentralAccessLogPaths)+len(configured))
for _, p := range knownCentralAccessLogPaths {
addCentralAccessLog(out, p)
}
// Multiple AccessLogPaths are fallback candidates; CheckWPBruteForce
// tails only the first one with data, so excluding every candidate here
// would drop later logs that were never scanned centrally.
if len(configured) == 1 {
addCentralAccessLog(out, configured[0])
}
return out
}
func addCentralAccessLog(out map[string]bool, path string) {
if path == "" {
return
}
out[path] = true
if real, err := filepath.EvalSymlinks(path); err == nil {
out[real] = true
}
}
// effectiveDomlogTailLines returns the operator-configured
// thresholds.domlog_tail_lines value or the built-in default when unset.
func effectiveDomlogTailLines(cfg *config.Config) int {
if cfg == nil || cfg.Thresholds.DomlogTailLines <= 0 {
return domlogTailLines
}
return cfg.Thresholds.DomlogTailLines
}
// effectiveDomlogMaxAge returns the operator-configured
// thresholds.domlog_max_age_min as a Duration, or the built-in default
// when unset.
func effectiveDomlogMaxAge(cfg *config.Config) time.Duration {
if cfg == nil || cfg.Thresholds.DomlogMaxAgeMin <= 0 {
return domlogMaxAge
}
return time.Duration(cfg.Thresholds.DomlogMaxAgeMin) * time.Minute
}
// tailDomlogsInto tails each discovered path and feeds every parsed
// access-log record into stats. Returns the number of files actually
// tailed (loop exits early if ctx is cancelled mid-pass).
//
// Single tail-and-aggregate loop shared by every domlog scanner so the
// per-file ctx gate, the parse-or-skip behaviour, and the scanned
// counter cannot drift between callers.
func tailDomlogsInto(ctx context.Context, paths []string, cfg *config.Config, stats *domlogStats, classifier botClassifier, tailLines int) int {
scanned := 0
for _, p := range paths {
if ctx != nil {
if err := ctx.Err(); err != nil {
break
}
}
domain := domainFromDomlogPath(p)
account := domainAccountOwner(domain)
for _, line := range tailFile(p, tailLines) {
rec, ok := parseAccessLogRecord(line)
if !ok {
continue
}
rec.Domain = domain
rec.Account = account
stats.scan(rec, cfg, classifier)
}
scanned++
}
return scanned
}
// domainFromDomlogPath derives the vhost from a per-domain domlog file
// path. Returns "" for paths that do not look like a domain log so the
// central access log and odd filenames do not pollute the per-IP vhost set.
func domainFromDomlogPath(p string) string {
base := filepath.Base(p)
if domain, ok := pleskDomlogDomain(p, base); ok {
return cleanDomlogDomain(domain)
}
if domain, ok := trimDomlogSuffix(base); ok {
return cleanDomlogDomain(domain)
}
return cleanDomlogDomain(base)
}
func pleskDomlogDomain(p, base string) (string, bool) {
switch strings.ToLower(strings.TrimSpace(base)) {
case "access_log", "access_ssl_log", "proxy_access_ssl_log":
return filepath.Base(filepath.Dir(filepath.Dir(p))), true
default:
return "", false
}
}
func trimDomlogSuffix(base string) (string, bool) {
trimmed := strings.TrimSpace(base)
low := strings.ToLower(trimmed)
for _, suffix := range []string{
".access.log",
"-access.log",
"_access.log",
"-access_log",
"_access_log",
"-ssl_log",
"_log",
".log",
} {
if strings.HasSuffix(low, suffix) {
return trimmed[:len(trimmed)-len(suffix)], true
}
}
return trimmed, false
}
func cleanDomlogDomain(domain string) string {
domain = strings.ToLower(strings.TrimSpace(domain))
if domain == "" || len(domain) > 253 || !strings.Contains(domain, ".") {
return ""
}
if strings.HasPrefix(domain, ".") || strings.HasSuffix(domain, ".") ||
strings.Contains(domain, "..") {
return ""
}
if net.ParseIP(domain) != nil {
return ""
}
labels := strings.Split(domain, ".")
for _, label := range labels {
if label == "" || len(label) > 63 ||
strings.HasPrefix(label, "-") || strings.HasSuffix(label, "-") {
return ""
}
for _, c := range label {
if (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') || c == '-' {
continue
}
return ""
}
}
return domain
}
// scanDomlogsStats discovers per-vhost logs honouring the operator's
// thresholds and feeds each parsed record into stats. Production entry
// point used by CheckWPBruteForce. Returns the number of files actually
// tailed.
func scanDomlogsStats(ctx context.Context, cfg *config.Config, stats *domlogStats) int {
if cfg == nil {
cfg = &config.Config{}
}
paths := discoverFreshDomlogs(ctx, cfg.Thresholds.DomlogMaxFiles, effectiveDomlogMaxAge(cfg))
return tailDomlogsInto(ctx, paths, cfg, stats, currentBotClassifier(cfg), effectiveDomlogTailLines(cfg))
}
// scanDomlogs is the legacy infra-IPs-only entry kept for test fixtures
// that drive the brute-force counters directly. Production code calls
// scanDomlogsStats. Both share discoverFreshDomlogs + tailDomlogsInto so
// path selection and per-file ctx semantics cannot diverge.
//
// maxFiles <= 0 falls back to the built-in domlogMaxFiles default.
func scanDomlogs(ctx context.Context, infraIPs []string, maxFiles int, wpLogin, xmlrpc, userEnum map[string]int) int {
cfg := &config.Config{InfraIPs: infraIPs}
stats := newDomlogStats()
stats.wpLogin = wpLogin
stats.xmlrpc = xmlrpc
stats.userEnum = userEnum
paths := discoverFreshDomlogs(ctx, maxFiles, 0)
return tailDomlogsInto(ctx, paths, cfg, stats, nopBotClassifier{}, domlogTailLines)
}
// countBruteForce parses Combined Log Format lines and increments per-IP
// counters via the shared domlogStats aggregator. Kept as a thin shim
// for tests that feed lines directly (no file discovery / tail step).
func countBruteForce(lines []string, infraIPs []string, wpLogin, xmlrpc, userEnum map[string]int) {
cfg := &config.Config{InfraIPs: infraIPs}
stats := newDomlogStats()
stats.wpLogin = wpLogin
stats.xmlrpc = xmlrpc
stats.userEnum = userEnum
for _, line := range lines {
rec, ok := parseAccessLogRecord(line)
if !ok {
continue
}
stats.scan(rec, cfg, nopBotClassifier{})
}
}
// syslogMessagesTailLinesDefault is the built-in fallback for the legacy
// direct CheckFTPLogins path used when no state store is available.
// Operator override: cfg.Thresholds.SyslogMessagesTailLines.
const syslogMessagesTailLinesDefault = 200
// CheckFTPLogins detects pure-ftpd brute force. With a state store it reads
// /var/log/messages forward-only and accumulates per-IP failures over a sliding
// window; without a store it falls back to the legacy per-cycle tail.
func CheckFTPLogins(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if cfg == nil {
cfg = &config.Config{}
}
if store == nil {
return checkFTPLoginsLegacy(cfg)
}
ftpTrackerMu.Lock()
defer ftpTrackerMu.Unlock()
now := time.Now()
tracker := loadFTPFailTracker(store)
lines, next, skipped, err := readNewSyslogLines(ftpSyslogPath, tracker.Follow)
if err != nil {
return nil // leave stored state untouched
}
if skipped > 0 {
observeFTPSkippedBytes(skipped)
}
windowMin := effectiveFTPFailWindowMin(cfg)
tracker.evict(now, windowMin)
cutoff := now.Add(-time.Duration(windowMin) * time.Minute)
var findings []alert.Finding
for _, line := range lines {
if !isPureFTPDLogFields(strings.Fields(line)) {
continue
}
ip := extractIPFromLog(line)
if ignoredFTPClientIP(ip, cfg.InfraIPs) {
continue
}
switch {
case strings.Contains(line, "Authentication failed"), strings.Contains(line, "auth failed"):
// Count the failure when the log recorded it, not when it was
// read: a first run or a large gap catches up on hours or days
// of history, and stamping that with now would turn scattered
// failures into one burst and auto-block the address.
at, ok := syslogLineTime(line, now)
if !ok || at.After(now) {
// A missing or future timestamp cannot safely seed a future
// minute bucket that survives eviction indefinitely. Treat the
// record as current, matching the timestamp-free fallback.
at = now
}
if at.Before(cutoff) {
continue
}
tracker.record(ip, at)
case strings.Contains(line, "is now logged in"):
findings = append(findings, ftpLoginFinding(ip, line, tracker.count(ip)))
}
}
tracker.capIPs(maxTrackedIPs)
for _, off := range tracker.offenders(ftpFailThreshold) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "ftp_bruteforce",
SourceIP: off.IP,
Message: fmt.Sprintf("FTP brute force from %s: %d failed attempts in %dm", off.IP, off.Count, windowMin),
})
}
tracker.Follow = next
tracker.save(store)
return findings
}
// checkFTPLoginsLegacy is the pre-store tail-based detector, used only when no
// state store is available (direct callers / tests). The daemon always supplies
// a store and uses the forward-only store-backed path above.
func checkFTPLoginsLegacy(cfg *config.Config) []alert.Finding {
var findings []alert.Finding
tailLines := syslogMessagesTailLinesDefault
if cfg != nil && cfg.Thresholds.SyslogMessagesTailLines > 0 {
tailLines = cfg.Thresholds.SyslogMessagesTailLines
}
lines := tailFile("/var/log/messages", tailLines)
if len(lines) == 0 {
return nil
}
failedFTP := make(map[string]int)
for _, line := range lines {
fields := strings.Fields(line)
// pure-ftpd logs: "pure-ftpd: ... [WARNING] Authentication failed for user"
if !isPureFTPDLogFields(fields) {
continue
}
ip := extractIPFromLog(line)
if ignoredFTPClientIP(ip, cfg.InfraIPs) {
continue
}
switch {
case strings.Contains(line, "Authentication failed"), strings.Contains(line, "auth failed"):
failedFTP[ip]++
case strings.Contains(line, "is now logged in"):
// failedFTP holds failures seen earlier in this batch; pure-ftpd logs
// chronologically so a brute-then-success appears as fails before the
// "logged in" line, letting the shared builder escalate.
findings = append(findings, ftpLoginFinding(ip, line, failedFTP[ip]))
}
}
for ip, count := range failedFTP {
if count >= ftpFailThreshold {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "ftp_bruteforce",
SourceIP: ip,
Message: fmt.Sprintf("FTP brute force from %s: %d failed attempts", ip, count),
})
}
}
return findings
}
// CheckWebmailLogins parses cPanel access log for webmail logins from non-infra IPs.
func CheckWebmailLogins(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if cfg.Suppressions.SuppressWebmail {
return nil
}
// Webmail ports and this log are cPanel-specific; on other panels the
// path does not exist and the check would silently return nothing.
if !platform.Detect().IsCPanel() {
return nil
}
var findings []alert.Finding
lines := tailFile("/usr/local/cpanel/logs/access_log", 300)
loginAttempts := make(map[string]int)
for _, line := range lines {
// Webmail ports: 2095 (HTTP), 2096 (HTTPS)
if !strings.Contains(line, "2095") && !strings.Contains(line, "2096") {
continue
}
fields := strings.Fields(line)
if len(fields) < 1 {
continue
}
ip := fields[0]
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
continue
}
// Count login attempts per IP
if strings.Contains(line, "POST") && (strings.Contains(line, "login") || strings.Contains(line, "auth")) {
loginAttempts[ip]++
}
}
for ip, count := range loginAttempts {
if count >= webmailThreshold {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "webmail_bruteforce",
Message: fmt.Sprintf("Webmail brute force from %s: %d attempts", ip, count),
})
}
}
return findings
}
// CheckAPIAuthFailures parses cPanel access log for failed API authentication.
func CheckAPIAuthFailures(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
// The cPanel/WHM API access log is cPanel-specific.
if !platform.Detect().IsCPanel() {
return nil
}
var findings []alert.Finding
lines := tailFile("/usr/local/cpanel/logs/access_log", 300)
failedAPI := make(map[string]int)
for _, line := range lines {
// Look for 401/403 responses on API endpoints
if !strings.Contains(line, "\" 401 ") && !strings.Contains(line, "\" 403 ") {
continue
}
// Only API endpoints
if !strings.Contains(line, "json-api") && !strings.Contains(line, "/execute/") &&
!strings.Contains(line, "cpsess") {
continue
}
fields := strings.Fields(line)
if len(fields) < 1 {
continue
}
ip := fields[0]
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
continue
}
failedAPI[ip]++
}
for ip, count := range failedAPI {
if count >= apiFailThreshold {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "api_auth_failure",
Message: fmt.Sprintf("cPanel API auth failures from %s: %d attempts", ip, count),
Details: "Possible API token brute force or unauthorized API access",
})
}
}
return findings
}
// FTPLoginFinding builds the finding for a successful pure-ftpd login line,
// reporting false when the line is not a login or the client address is
// loopback, infrastructure, or unusable. The daemon's realtime log watcher
// calls it so a login seen live and the same line re-read by CheckFTPLogins
// carry one identity; without that the state store sees two findings and the
// operator gets one login reported twice. It has no brute-force history, so
// the escalation to ftp_login_after_bruteforce stays with the scheduled check.
func FTPLoginFinding(line string, cfg *config.Config) (alert.Finding, bool) {
if cfg == nil {
cfg = &config.Config{}
}
if !strings.Contains(line, "is now logged in") {
return alert.Finding{}, false
}
ip := extractIPFromLog(line)
if ignoredFTPClientIP(ip, cfg.InfraIPs) {
return alert.Finding{}, false
}
return ftpLoginFinding(ip, line, 0), true
}
// ftpLoginFinding builds the finding for a successful FTP login from a
// non-infra, non-loopback IP. A login from a source that has already crossed
// the brute-force threshold is a likely cracked credential and pages as
// Critical; any other login is audit-level (Warning), matching the cPanel-login
// detector ("logins are audit trail, not paging-level").
func ftpLoginFinding(ip, line string, recentFails int) alert.Finding {
if recentFails >= ftpFailThreshold {
msg := fmt.Sprintf("FTP login succeeded from brute-force source %s after %d failed attempts", ip, recentFails)
owner := ""
if account := parseFTPLoginAccount(line); account != "" {
msg = fmt.Sprintf("FTP login succeeded for account %s from brute-force source %s after %d failed attempts", account, ip, recentFails)
owner = ftpAccountOwner(account)
}
return alert.Finding{
Severity: alert.Critical,
Check: "ftp_login_after_bruteforce",
DedupKey: loginRecordKey(line),
SourceIP: ip,
Message: msg,
Details: truncate(line, 200),
TenantID: owner,
}
}
return alert.Finding{
Severity: alert.Warning,
Check: "ftp_login",
DedupKey: loginRecordKey(line),
SourceIP: ip,
Message: fmt.Sprintf("FTP login from non-infra IP: %s", ip),
Details: truncate(line, 200),
}
}
// parseFTPLoginAccount extracts the account name from a pure-ftpd
// "<user> is now logged in" line. The username is attacker-influenced (the
// login name appears verbatim in the log), so this only splits on delimiters
// and never interprets the value.
func parseFTPLoginAccount(line string) string {
const marker = " is now logged in"
i := strings.Index(line, marker)
if i < 0 {
return ""
}
pre := strings.TrimRight(line[:i], " ")
if j := strings.LastIndexByte(pre, ' '); j >= 0 {
return pre[j+1:]
}
return pre
}
func ignoredFTPClientIP(ip string, infraIPs []string) bool {
if ip == "" || isInfraIP(ip, infraIPs) {
return true
}
parsed := net.ParseIP(ip)
return parsed == nil || parsed.IsLoopback()
}
// extractIPFromLog tries to extract an IP address from a log line.
func extractIPFromLog(line string) string {
fields := strings.Fields(line)
if isPureFTPDLogFields(fields) {
// pure-ftpd logs the peer as a parenthesised "(user@ip)" token -- or
// "(?@ip)" when the username is unknown -- so the address (IPv4 or
// IPv6) is glued inside the parens, never a standalone field.
for _, f := range fields {
if ip := ipFromParenPeer(f); ip != "" {
return ip
}
}
}
// Fallback: a bare space-delimited IPv4 field (web access logs,
// fail2ban-style "banned <ip>" lines).
for _, f := range fields {
// Simple IP detection: starts with digit, contains dots
if len(f) >= 7 && f[0] >= '0' && f[0] <= '9' && strings.Count(f, ".") == 3 {
// Strip trailing punctuation
f = strings.TrimRight(f, ",:;)([]")
return f
}
}
return ""
}
func isPureFTPDLogFields(fields []string) bool {
if len(fields) == 0 {
return false
}
if isPureFTPDProgramToken(fields[0]) {
return true
}
if len(fields) >= 5 && isSyslogTimestampPrefix(fields) && isPureFTPDProgramToken(fields[4]) {
return true
}
if len(fields) >= 3 {
if _, err := time.Parse(time.RFC3339Nano, fields[0]); err == nil {
return isPureFTPDProgramToken(fields[2])
}
}
return false
}
func isPureFTPDProgramToken(field string) bool {
if field == "pure-ftpd:" {
return true
}
if !strings.HasPrefix(field, "pure-ftpd[") || !strings.HasSuffix(field, "]:") {
return false
}
pid := field[len("pure-ftpd[") : len(field)-len("]:")]
if pid == "" {
return false
}
for _, r := range pid {
if r < '0' || r > '9' {
return false
}
}
return true
}
func isSyslogTimestampPrefix(fields []string) bool {
return isSyslogMonth(fields[0]) && isSyslogDay(fields[1]) && isSyslogClock(fields[2])
}
func isSyslogMonth(s string) bool {
switch s {
case "Jan", "Feb", "Mar", "Apr", "May", "Jun", "Jul", "Aug", "Sep", "Oct", "Nov", "Dec":
return true
default:
return false
}
}
func isSyslogDay(s string) bool {
if len(s) < 1 || len(s) > 2 {
return false
}
for _, r := range s {
if r < '0' || r > '9' {
return false
}
}
return true
}
func isSyslogClock(s string) bool {
if len(s) != len("00:00:00") || s[2] != ':' || s[5] != ':' {
return false
}
for i, r := range s {
if i == 2 || i == 5 {
continue
}
if r < '0' || r > '9' {
return false
}
}
return true
}
// ipFromParenPeer extracts an IP from pure-ftpd's "(user@ip)" / "(?@ip)" peer
// token, supporting both IPv4 and IPv6 addresses. Returns "" when f is not
// such a token or the candidate after the last '@' does not parse as an IP.
func ipFromParenPeer(f string) string {
if len(f) < 4 || f[0] != '(' || f[len(f)-1] != ')' {
return ""
}
inner := f[1 : len(f)-1]
at := strings.LastIndexByte(inner, '@')
if at < 0 || at+1 >= len(inner) {
return ""
}
cand := inner[at+1:]
if net.ParseIP(cand) == nil {
return ""
}
return cand
}
package checks
import (
"fmt"
"os"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
)
var (
challengeRoutedMetric *metrics.CounterVec
challengeRoutedMetricOnce sync.Once
)
// observeChallengeRouted counts one IP routed to the proof-of-work challenge,
// labelled by the source check that flagged it, so operators can graph
// challenge volume per detector (e.g. http_scanner_profile). Registered lazily
// on first use, mirroring the auto-response metric in the runner.
func observeChallengeRouted(check string) {
challengeRoutedMetricOnce.Do(func() {
challengeRoutedMetric = metrics.NewCounterVec(
"csm_challenge_routed_total",
"IPs routed to the proof-of-work challenge, by the source check that flagged them.",
[]string{"check"},
)
metrics.MustRegister("csm_challenge_routed_total", challengeRoutedMetric)
})
challengeRoutedMetric.With(check).Inc()
}
// ChallengeIPList abstracts the challenge IP list for routing.
type ChallengeIPList interface {
Add(ip string, reason string, duration time.Duration)
AddNonEscalating(ip string, reason string, duration time.Duration)
Remove(ip string)
Contains(ip string) bool
}
var challengeIPList ChallengeIPList
// SetChallengeIPList sets the challenge IP list for routing.
func SetChallengeIPList(list ChallengeIPList) {
challengeIPList = list
}
// GetChallengeIPList returns the current challenge IP list (for AutoBlockIPs skip check).
func GetChallengeIPList() ChallengeIPList {
return challengeIPList
}
func isChallengeableCheck(check string) bool {
return ResponsePolicyFor(check).ChallengeFirst
}
// Auto-response actions a challengeable check can resolve to.
const (
responseChallenge = "challenge"
responseBlock = "block"
)
// responseActionForCheck returns the effective auto-response for a check:
// "challenge" to route the IP to the PoW gate, or "block" to hard-block it.
// Challengeable checks default to "challenge" only while challenge routing is
// enabled; otherwise they fall through to "block". An operator-selectable
// override (currently only http_scanner_profile via
// auto_response.http_scanner_action) forces "block". Non-challengeable checks
// always resolve to "block". This is the single source of truth for the
// challenge-vs-block decision, shared by ChallengeRouteIPs and AutoBlockIPs so
// the two cannot diverge.
func responseActionForCheck(cfg *config.Config, check string) string {
if !cfg.Challenge.Enabled || !challengeRoutesCheck(cfg, check) {
return responseBlock
}
return responseChallenge
}
// challengeRoutesCheck is the challenge-vs-block policy for a check with
// challenge routing assumed on.
func challengeRoutesCheck(cfg *config.Config, check string) bool {
if !isChallengeableCheck(check) {
return false
}
return check != "http_scanner_profile" || cfg.AutoResponse.HTTPScannerAction != responseBlock
}
// responseActionForFinding narrows responseActionForCheck for one finding.
// ip_reputation grades its sighting severity by detection vector
// (reputationSightingSeverity: HTTP and cPanel access are High, every
// other vector Critical), so a Critical reputation sighting came from a
// browserless channel (SMTP, IMAP, FTP, SSH) where nothing can ever
// answer the PoW page -- challenge-routing it just leaves the attacker
// unblocked, retrying daily. Those resolve to a hard block.
func responseActionForFinding(cfg *config.Config, f alert.Finding) string {
if !cfg.Challenge.Enabled || !challengeRoutesFinding(cfg, f) {
return responseBlock
}
return responseChallenge
}
// challengeRoutesFinding narrows challengeRoutesCheck for one finding, with
// challenge routing assumed on.
func challengeRoutesFinding(cfg *config.Config, f alert.Finding) bool {
if f.Check == "ip_reputation" && f.Severity == alert.Critical {
return false
}
return challengeRoutesCheck(cfg, f.Check)
}
// isHardBlockCheck reports whether a check must never be routed to the
// challenge: its registry policy says so, or it is a runtime-built name the
// prefix contract covers.
func isHardBlockCheck(check string) bool {
return ResponsePolicyFor(check).NeverChallenge || neverChallengeDynamicName(check)
}
const challengeDuration = 30 * time.Minute
// ChallengeThenBlock runs the two IP-disposition stages in their required
// order -- challenge routing first so an eligible IP is on the challenge list
// before AutoBlockIPs checks membership, then hard-blocking -- and returns both
// action sets. Auto-response call sites use this single helper instead of
// hand-ordering the two calls, so the "challenge before block" invariant cannot
// be silently broken by reordering in one path. Both stages run on the same
// finding set (the full/repeat-offender set); callers append the returned
// actions wherever their pipeline expects them.
func ChallengeThenBlock(cfg *config.Config, findings []alert.Finding) (challengeActions, blockActions []alert.Finding) {
challengeActions = ChallengeRouteIPs(cfg, findings)
blockActions = AutoBlockIPs(cfg, findings)
return challengeActions, blockActions
}
// ChallengeRouteIPs processes findings and routes eligible IPs to the challenge
// list instead of hard-blocking them. Must be called BEFORE AutoBlockIPs so
// that challenged IPs are on the list when AutoBlockIPs checks Contains().
func ChallengeRouteIPs(cfg *config.Config, findings []alert.Finding) []alert.Finding {
if !cfg.Challenge.Enabled || challengeIPList == nil {
return nil
}
var actions []alert.Finding
routed := make(map[string]bool)
for _, f := range findings {
// Challenge timeouts can hard-block too, so gated authentication
// checks must honor the same opt-in as direct firewall responses.
if ResponsePolicyFor(f.Check).Block == BlockWithCpanelLogins && !cfg.AutoResponse.BlockCpanelLogins {
continue
}
if isHardBlockCheck(f.Check) {
continue
}
// Only route checks that are known to contain attacker IPs.
// This is an allowlist: a new IP-bearing check must be given a
// ChallengeFirst Response in the check registry. Defaulting to skip
// prevents version numbers, sizes, and other numeric finding fields
// from being blocked as IPs.
if !isChallengeableCheck(f.Check) {
continue
}
// The scanner-profile response is operator-selectable: "block"
// skips routing here so AutoBlockIPs hard-blocks the IP instead.
if responseActionForFinding(cfg, f) == responseBlock {
continue
}
ip := extractIPFromFinding(f)
if ip == "" || routed[ip] {
continue
}
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
continue
}
if challengeIPList.Contains(ip) {
continue
}
addChallengeIP(f.Check, ip, f.Message, challengeDuration, alert.FindingID(f))
routed[ip] = true
observeChallengeRouted(f.Check)
recordChallengeRouteStat(ip, f.Check, time.Now())
fmt.Fprintf(os.Stderr, "[%s] CHALLENGE: %s routed to challenge (check: %s)\n",
time.Now().Format("2006-01-02 15:04:05"), ip, f.Check)
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "challenge_route",
Message: fmt.Sprintf("CHALLENGE: %s sent to PoW challenge (expires in %s)", ip, challengeDuration),
Details: fmt.Sprintf("Reason: %s", f.Message),
Timestamp: time.Now(),
})
}
return actions
}
func addChallengeIP(check, ip, reason string, duration time.Duration, findingID string) {
if check == "http_claimed_bot_unverified" {
challengeIPList.AddNonEscalating(ip, reason, duration)
return
}
if list, ok := challengeIPList.(interface {
AddWithFindingID(string, string, time.Duration, string)
}); ok {
list.AddWithFindingID(ip, reason, duration, findingID)
return
}
challengeIPList.Add(ip, reason, duration)
}
func removeChallengeIP(ip string) {
if challengeIPList == nil {
return
}
challengeIPList.Remove(ip)
}
package checks
import (
"sync"
"time"
)
// challengeRecentMax bounds the recent-routes ring buffer surfaced to the web
// UI. Small: the panel shows the latest handful, history holds the rest.
const challengeRecentMax = 20
// ChallengeRouteRecord is one IP routed to the challenge, for the web UI's
// recent-activity list.
type ChallengeRouteRecord struct {
IP string `json:"ip"`
Check string `json:"check"`
At time.Time `json:"at"`
}
// ChallengeUIStatsSnapshot is the web-UI view of challenge routing: cumulative
// per-check counts since daemon start plus the most recent routes. It is a
// copy, safe for the caller to read without locking.
type ChallengeUIStatsSnapshot struct {
RoutedByCheck map[string]int `json:"routed_by_check"`
Recent []ChallengeRouteRecord `json:"recent"`
}
var (
challengeStatsMu sync.Mutex
challengeRoutedByCheck = map[string]int{}
challengeRecentRoutes []ChallengeRouteRecord
)
// recordChallengeRouteStat records one route for the web-UI stats. Kept
// separate from the Prometheus counter (observeChallengeRouted) so the UI does
// not depend on scraping /metrics.
func recordChallengeRouteStat(ip, check string, at time.Time) {
challengeStatsMu.Lock()
defer challengeStatsMu.Unlock()
challengeRoutedByCheck[check]++
challengeRecentRoutes = append(challengeRecentRoutes, ChallengeRouteRecord{IP: ip, Check: check, At: at})
if len(challengeRecentRoutes) > challengeRecentMax {
challengeRecentRoutes = challengeRecentRoutes[len(challengeRecentRoutes)-challengeRecentMax:]
}
}
// ChallengeUIStats returns a copy of the current challenge routing stats for
// the web UI. Most recent route is last in Recent.
func ChallengeUIStats() ChallengeUIStatsSnapshot {
challengeStatsMu.Lock()
defer challengeStatsMu.Unlock()
byCheck := make(map[string]int, len(challengeRoutedByCheck))
for k, v := range challengeRoutedByCheck {
byCheck[k] = v
}
recent := make([]ChallengeRouteRecord, len(challengeRecentRoutes))
copy(recent, challengeRecentRoutes)
return ChallengeUIStatsSnapshot{RoutedByCheck: byCheck, Recent: recent}
}
package checks
import (
"context"
"maps"
"os"
"path/filepath"
"sync"
"github.com/pidginhost/csm/internal/state"
)
// incompleteCheckCollector records owners whose coverage this run could not
// complete. Known file gaps can be preserved in the eventual store transaction;
// an unknown range prevents the owner from retiring anything.
type incompleteCheckCollector struct {
mu sync.Mutex
names map[string]struct{}
skipped map[string]bool
}
type incompleteCheckContextKey struct{}
type coverageGapsContextKey struct{}
type coveragePathCollectorContextKey struct{}
type coveragePathCollector struct {
mu sync.Mutex
pathsByOwner map[string]map[string]bool
scopesByOwner map[string]map[string]bool
}
// CoverageGaps holds file gaps and completed database scopes for a scan.
// The caller passes it to the atomic purge-and-merge operation so a concurrent
// update made after the scanner read LatestFindings cannot be retired by a
// stale carry-forward snapshot.
type CoverageGaps struct {
mu sync.Mutex
pathsByCheck map[string]map[string]bool
completedScopes map[string]map[string]bool
incompleteChecks map[string]bool
}
// Paths returns an isolated snapshot of the completed run's path gaps. Each
// inner key is a stable lexical or resolved alias captured during the scan.
func (g *CoverageGaps) Paths() map[string]map[string]bool {
if g == nil {
return nil
}
g.mu.Lock()
defer g.mu.Unlock()
return cloneCoverageGapPaths(g.pathsByCheck)
}
// WithCoverageGaps requests an atomic coverage-aware store operation from the
// caller. The runner publishes only its completed snapshot into the handle.
func WithCoverageGaps(ctx context.Context) (context.Context, *CoverageGaps) {
if ctx == nil {
ctx = context.Background()
}
gaps := &CoverageGaps{}
return context.WithValue(ctx, coverageGapsContextKey{}, gaps), gaps
}
func withIncompleteCheckCollector(ctx context.Context) (context.Context, *incompleteCheckCollector) {
if ctx == nil {
ctx = context.Background()
}
collector := &incompleteCheckCollector{names: make(map[string]struct{})}
return context.WithValue(ctx, incompleteCheckContextKey{}, collector), collector
}
func withCoveragePathCollector(ctx context.Context) (context.Context, *coveragePathCollector) {
collector := &coveragePathCollector{pathsByOwner: make(map[string]map[string]bool)}
return context.WithValue(ctx, coveragePathCollectorContextKey{}, collector), collector
}
func coverageGapsFrom(ctx context.Context) *CoverageGaps {
if ctx == nil {
return nil
}
gaps, _ := ctx.Value(coverageGapsContextKey{}).(*CoverageGaps)
return gaps
}
func (g *CoverageGaps) replace(pathsByCheck map[string]map[string]bool, completedScopes map[string]map[string]bool, incompleteChecks map[string]bool) {
if g == nil {
return
}
g.mu.Lock()
g.pathsByCheck = cloneCoverageGapPaths(pathsByCheck)
g.completedScopes = cloneCoverageGapPaths(completedScopes)
g.incompleteChecks = maps.Clone(incompleteChecks)
g.mu.Unlock()
}
func cloneCoverageGapPaths(pathsByCheck map[string]map[string]bool) map[string]map[string]bool {
if len(pathsByCheck) == 0 {
return nil
}
out := make(map[string]map[string]bool, len(pathsByCheck))
for check, paths := range pathsByCheck {
if len(paths) == 0 {
continue
}
cloned := make(map[string]bool, len(paths))
for path := range paths {
cloned[path] = true
}
out[check] = cloned
}
if len(out) == 0 {
return nil
}
return out
}
// markCheckIncomplete records a coverage gap that cannot be attributed to
// particular files, so the owner keeps every finding it has until it completes.
func markCheckIncomplete(ctx context.Context, name string) {
collector := incompleteCollectorFrom(ctx)
if collector == nil {
return
}
collector.mu.Lock()
collector.names[name] = struct{}{}
collector.mu.Unlock()
}
// markCheckSkipped preserves both discovered state and the last run's coverage
// summary when an internal refresh interval prevents any new audit work.
func markCheckSkipped(ctx context.Context, name string) {
collector := incompleteCollectorFrom(ctx)
if collector == nil {
return
}
collector.mu.Lock()
defer collector.mu.Unlock()
collector.names[name] = struct{}{}
if collector.skipped == nil {
collector.skipped = make(map[string]bool)
}
collector.skipped[name] = true
}
func (c *incompleteCheckCollector) wasSkipped(name string) bool {
c.mu.Lock()
defer c.mu.Unlock()
return c.skipped[name]
}
// A removed path is covered absence; a failed read is not evidence of cleanup.
func markScanReadError(ctx context.Context, owner string, err error) {
if err != nil && !os.IsNotExist(err) {
markCheckIncomplete(ctx, owner)
}
}
// Unlike the best-effort inventory, partial discovery cannot authorize a
// stateful scanner to retire findings from accounts it never enumerated.
func scanHomeDirsWithCoverage(ctx context.Context, owner string) []os.DirEntry {
if AccountFromContext(ctx) != "" {
entries, err := GetScanHomeDirs(ctx)
markScanReadError(ctx, owner, err)
return entries
}
homes, err := readAccountHomes()
markScanReadError(ctx, owner, err)
entries := make([]os.DirEntry, 0, len(homes))
for _, home := range homes {
entries = append(entries, rootedDirEntry{DirEntry: home.Entry, root: home.Root})
}
return entries
}
// recordCoverageGapPaths records the stable aliases captured when a known file
// gap was observed. The store must consume these aliases as identities, without
// resolving them again after a symlink may have changed targets.
func recordCoverageGapPaths(ctx context.Context, owner string, paths []string) {
if ctx == nil || len(paths) == 0 {
return
}
collector, _ := ctx.Value(coveragePathCollectorContextKey{}).(*coveragePathCollector)
if collector == nil {
return
}
collector.mu.Lock()
if collector.pathsByOwner[owner] == nil {
collector.pathsByOwner[owner] = make(map[string]bool)
}
for _, path := range paths {
if path != "" {
collector.pathsByOwner[owner][path] = true
}
}
collector.mu.Unlock()
}
// coveragePathAliases captures every path spelling a Finding.FilePath emitted
// by this walk can use: its absolute lexical form and, while the observed path
// still resolves, its symlink-resolved form. Callers retain the returned set so
// a later symlink retarget cannot change the preservation identity.
func coveragePathAliases(path string) []string {
lexical := coverageLexicalPath(path)
if lexical == "" {
return nil
}
aliases := []string{lexical}
if real, err := filepath.EvalSymlinks(lexical); err == nil {
real = filepath.Clean(real)
if real != lexical {
aliases = append(aliases, real)
}
}
return aliases
}
func coverageLexicalPath(path string) string {
if path == "" {
return ""
}
lexical := filepath.Clean(path)
if absolute, err := filepath.Abs(lexical); err == nil {
lexical = filepath.Clean(absolute)
}
return lexical
}
// stableCoveragePathAliases binds aliases to the file metadata that caused the
// scanner's decision. Re-checking both the lexical and resolved paths prevents
// a symlink retarget during alias construction from preserving a different file
// while retiring the one that actually went unexamined.
func stableCoveragePathAliases(path string, expected os.FileInfo) ([]string, bool) {
aliases := coveragePathAliases(path)
if expected == nil || len(aliases) == 0 {
return aliases, false
}
lexicalInfo, err := osFS.Lstat(aliases[0])
if err != nil || lexicalInfo.Mode()&os.ModeSymlink != 0 || !os.SameFile(expected, lexicalInfo) {
return aliases, false
}
for _, alias := range aliases[1:] {
resolvedInfo, statErr := osFS.Stat(alias)
if statErr != nil || !os.SameFile(expected, resolvedInfo) {
return aliases, false
}
}
lexicalInfo, err = osFS.Lstat(aliases[0])
if err != nil || lexicalInfo.Mode()&os.ModeSymlink != 0 || !os.SameFile(expected, lexicalInfo) {
return aliases, false
}
return aliases, true
}
func incompleteCollectorFrom(ctx context.Context) *incompleteCheckCollector {
if ctx == nil {
return nil
}
collector, _ := ctx.Value(incompleteCheckContextKey{}).(*incompleteCheckCollector)
return collector
}
func checkMarkedIncomplete(ctx context.Context, name string) bool {
collector := incompleteCollectorFrom(ctx)
return collector != nil && collector.contains(name)
}
func (c *coveragePathCollector) gapPaths(owner string) map[string]bool {
c.mu.Lock()
defer c.mu.Unlock()
out := make(map[string]bool, len(c.pathsByOwner[owner]))
for path := range c.pathsByOwner[owner] {
out[path] = true
}
return out
}
func (c *incompleteCheckCollector) contains(name string) bool {
c.mu.Lock()
defer c.mu.Unlock()
_, ok := c.names[name]
return ok
}
// Snapshot includes both file gaps and completed database scopes from the same
// runner result. Callers must pass this snapshot to the atomic store operation.
func (g *CoverageGaps) Snapshot() *state.ScanCoverage {
if g == nil {
return nil
}
g.mu.Lock()
defer g.mu.Unlock()
return &state.ScanCoverage{
PreservePaths: cloneCoverageGapPaths(g.pathsByCheck),
CompletedScopes: cloneCoverageGapPaths(g.completedScopes),
IncompleteChecks: maps.Clone(g.incompleteChecks),
}
}
func recordCompletedCoverageScopes(ctx context.Context, owner string, scopes map[string]bool) {
if ctx == nil {
return
}
collector, _ := ctx.Value(coveragePathCollectorContextKey{}).(*coveragePathCollector)
if collector == nil {
return
}
collector.mu.Lock()
defer collector.mu.Unlock()
if collector.scopesByOwner == nil {
collector.scopesByOwner = make(map[string]map[string]bool)
}
complete := make(map[string]bool)
for scope, covered := range scopes {
if covered && scope != "" {
complete[scope] = true
}
}
collector.scopesByOwner[owner] = complete
}
func (c *coveragePathCollector) completedScopes(owner string) map[string]bool {
c.mu.Lock()
defer c.mu.Unlock()
out := make(map[string]bool, len(c.scopesByOwner[owner]))
for scope := range c.scopesByOwner[owner] {
out[scope] = true
}
return out
}
package checks
import (
"context"
"errors"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
const checkDispatchControlBudget = time.Minute
var checkDispatches = newCheckDispatchMonitor()
type checkDispatchMonitor struct {
mu sync.Mutex
batches map[*checkDispatchBatch]struct{}
losses *queuehealth.Tracker
}
type checkDispatchBatch struct {
monitor *checkDispatchMonitor
tasks map[*checkDispatch]struct{}
budget *scanBudget
progress time.Time
observer *CheckDispatchProgress
}
// Timing, failure and batch membership fields are guarded by the monitor mutex.
type checkDispatch struct {
batch *checkDispatchBatch
slot *scanSlot // set by the runner before starting the execution
started time.Time
deadline time.Time
failed bool
}
func newCheckDispatchMonitor() *checkDispatchMonitor {
return &checkDispatchMonitor{
batches: make(map[*checkDispatchBatch]struct{}),
losses: queuehealth.New(0, time.Minute),
}
}
func (m *checkDispatchMonitor) begin(count int, budget *scanBudget) []*checkDispatch {
now := time.Now()
batch := &checkDispatchBatch{
monitor: m, tasks: make(map[*checkDispatch]struct{}, count),
budget: budget, progress: now,
}
tasks := make([]*checkDispatch, count)
for i := range tasks {
tasks[i] = &checkDispatch{batch: batch}
batch.tasks[tasks[i]] = struct{}{}
}
if count > 0 {
m.mu.Lock()
m.batches[batch] = struct{}{}
m.mu.Unlock()
}
return tasks
}
func (t *checkDispatch) admit(ctx context.Context) bool {
t.slot = t.batch.budget.acquire(ctx)
if t.slot == nil {
return false
}
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
t.started = now
t.deadline = now.Add(checkDispatchControlBudget)
t.batch.progressed(now)
return true
}
func (t *checkDispatch) executing(ctx context.Context) {
if t == nil {
return
}
// A caller without a deadline is unbounded rather than already overdue.
deadline, _ := ctx.Deadline()
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
t.deadline = deadline
t.batch.progressed(time.Now())
}
func (t *checkDispatch) returned() {
if t == nil {
return
}
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
t.deadline = now.Add(checkDispatchControlBudget)
t.batch.progressed(now)
}
func (t *checkDispatch) withdraw(ctx context.Context) {
if t == nil || !errors.Is(ctx.Err(), context.DeadlineExceeded) {
return
}
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
t.failLocked(time.Now())
}
func (t *checkDispatch) failLocked(now time.Time) {
if !t.failed {
t.failed = true
t.batch.monitor.losses.Lose(now, 1)
}
}
func (t *checkDispatch) wrap(fn func()) func() {
return func() {
completed := false
defer func() {
t.slot.release()
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
now := time.Now()
if !completed {
t.failLocked(now)
}
delete(t.batch.tasks, t)
t.batch.progressed(now)
if len(t.batch.tasks) == 0 {
delete(m.batches, t.batch)
if t.batch.observer != nil {
delete(t.batch.observer.batches, t.batch)
}
}
}()
fn()
completed = true
}
}
func (m *checkDispatchMonitor) QueueStatus(now time.Time) queuehealth.Status {
m.mu.Lock()
defer m.mu.Unlock()
status := m.losses.Snapshot(now)
status.CapacityUnavailable = true
status.LagBasis = "consumer_progress"
var dispatchLate, runnerLate bool
for batch := range m.batches {
waiting, running := 0, 0
for task := range batch.tasks {
if task.started.IsZero() {
waiting++
} else {
running++
status.ProcessingSeconds = max(status.ProcessingSeconds, now.Sub(task.started).Seconds())
runnerLate = runnerLate || overdue(now, task.deadline)
}
}
status.Depth += waiting
status.InFlight += running
if waiting > 0 {
lag := now.Sub(batch.progress)
status.LagSeconds = max(status.LagSeconds, lag.Seconds())
// Other batches and withdrawn executions can still own slots.
// Only unused shared capacity indicates stalled dispatch.
dispatchLate = dispatchLate || (batch.budget.hasCapacity() && lag >= checkDispatchControlBudget)
}
}
switch {
case runnerLate:
status.Reason = "processing_lag"
case dispatchLate:
status.Reason = "backlog_lag"
}
if status.Reason != "" {
status.Status = "degraded"
}
return status
}
type checkDispatchContextKey struct{}
func withCheckDispatch(ctx context.Context, task *checkDispatch) context.Context {
return context.WithValue(ctx, checkDispatchContextKey{}, task)
}
func checkDispatchFrom(ctx context.Context) *checkDispatch {
task, _ := ctx.Value(checkDispatchContextKey{}).(*checkDispatch)
return task
}
// CheckDispatchQueueStatus measures pending checks and their runner wrappers.
func CheckDispatchQueueStatus(now time.Time) queuehealth.Status {
return checkDispatches.QueueStatus(now)
}
package checks
import (
"context"
"time"
)
// DispatchProgressSnapshot measures only the check batches owned by one caller.
// LastProgress remains available after the last wrapper exits, so its caller
// can time result handling without borrowing the completed check's budget.
type DispatchProgressSnapshot struct {
Active bool
Overdue bool
LastProgress time.Time
}
// CheckDispatchProgress is memory-only evidence for a scan's orchestration.
// All fields are guarded by the owning dispatch monitor's mutex.
type CheckDispatchProgress struct {
monitor *checkDispatchMonitor
batches map[*checkDispatchBatch]struct{}
last time.Time
}
type dispatchProgressContextKey struct{}
// WithCheckDispatchProgress binds subsequent check batches to this operation.
// A new binding isolates late callbacks from a previously canceled operation.
func WithCheckDispatchProgress(ctx context.Context) (context.Context, *CheckDispatchProgress) {
p := &CheckDispatchProgress{monitor: checkDispatches, batches: make(map[*checkDispatchBatch]struct{}), last: time.Now()}
return context.WithValue(ctx, dispatchProgressContextKey{}, p), p
}
func (m *checkDispatchMonitor) observe(ctx context.Context, tasks []*checkDispatch) {
p, _ := ctx.Value(dispatchProgressContextKey{}).(*CheckDispatchProgress)
if p == nil || len(tasks) == 0 {
return
}
m.mu.Lock()
defer m.mu.Unlock()
batch := tasks[0].batch
batch.observer = p
p.batches[batch] = struct{}{}
p.last = time.Now()
}
func (b *checkDispatchBatch) progressed(now time.Time) {
b.progress = now
if b.observer != nil {
b.observer.last = now
}
}
// Snapshot does not query contexts, the filesystem or the job store.
func (p *CheckDispatchProgress) Snapshot(now time.Time) DispatchProgressSnapshot {
p.monitor.mu.Lock()
defer p.monitor.mu.Unlock()
s := DispatchProgressSnapshot{Active: len(p.batches) != 0, LastProgress: p.last}
for batch := range p.batches {
waiting := 0
for task := range batch.tasks {
if task.started.IsZero() {
waiting++
} else {
s.Overdue = s.Overdue || !now.Before(task.deadline)
}
}
if waiting > 0 && batch.budget.hasCapacity() && now.Sub(batch.progress) >= checkDispatchControlBudget {
s.Overdue = true
}
}
return s
}
package checks
import (
"context"
"fmt"
"runtime/debug"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/obs"
)
type checkExecutionOutcome struct {
findings []alert.Finding
panicErr string
}
func executeCheckAsync(ctx context.Context, component string, fn func() []alert.Finding) *checkExecution {
return checkExecutions.execute(ctx, component, fn)
}
func (e *checkExecution) run(component string, fn func() []alert.Finding) {
defer e.slot.release()
e.monitor.mu.Lock()
e.started = time.Now()
e.monitor.mu.Unlock()
defer e.release()
outcome := checkExecutionOutcome{}
completed := false
defer func() {
if !completed {
e.fail()
}
if recovered := recover(); recovered != nil {
panicValue := fmt.Sprint(recovered)
outcome.panicErr = fmt.Sprintf("%s\n%s", panicValue, debug.Stack())
obs.CaptureMsg(component, "security check panic: "+panicValue)
}
e.done <- outcome
}()
outcome.findings = fn()
completed = true
}
package checks
import (
"context"
"errors"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
)
var checkExecutions = newCheckExecutionMonitor()
type checkExecutionMonitor struct {
mu sync.Mutex
pending map[*checkExecution]struct{}
losses *queuehealth.Tracker
}
func newCheckExecutionMonitor() *checkExecutionMonitor {
return &checkExecutionMonitor{
pending: make(map[*checkExecution]struct{}),
losses: queuehealth.New(0, time.Minute),
}
}
type checkExecution struct {
monitor *checkExecutionMonitor
dispatch *checkDispatch
slot *scanSlot
queued time.Time
started time.Time // guarded by monitor.mu
deadline time.Time
remaining atomic.Int32
failOnce sync.Once
done chan checkExecutionOutcome
// Only the caller changes this; the function owns its separate release.
callerSettled bool
}
func (m *checkExecutionMonitor) begin(deadline time.Time) *checkExecution {
execution := &checkExecution{
monitor: m, queued: time.Now(), deadline: deadline,
done: make(chan checkExecutionOutcome, 1),
}
execution.remaining.Store(2)
m.mu.Lock()
m.pending[execution] = struct{}{}
m.mu.Unlock()
return execution
}
func (m *checkExecutionMonitor) execute(ctx context.Context, component string, fn func() []alert.Finding) *checkExecution {
// Both runners construct a bounded per-check context before dispatch. A
// caller without one is unbounded rather than already overdue.
deadline, _ := ctx.Deadline()
execution := m.begin(deadline)
execution.dispatch = checkDispatchFrom(ctx)
execution.dispatch.executing(ctx)
if execution.dispatch != nil {
execution.slot = execution.dispatch.slot
execution.slot.retain()
}
go execution.run(component, fn)
return execution
}
func (e *checkExecution) fail() {
e.failOnce.Do(func() { e.monitor.losses.Lose(time.Now(), 1) })
}
func (e *checkExecution) received() {
if !e.callerSettled {
e.dispatch.returned()
e.callerSettled = true
}
}
func (e *checkExecution) withdraw(err error) {
e.received()
if errors.Is(err, context.DeadlineExceeded) {
e.fail()
}
}
func (e *checkExecution) finishCaller() {
if !e.callerSettled {
e.fail()
e.dispatch.returned()
}
e.release()
}
func (e *checkExecution) release() {
// A deadline can release the caller while the function still runs. A
// finished function likewise retains its buffered result for the caller.
if e.remaining.Add(-1) == 0 {
e.monitor.mu.Lock()
delete(e.monitor.pending, e)
e.monitor.mu.Unlock()
}
}
func (m *checkExecutionMonitor) QueueStatus(now time.Time) queuehealth.Status {
m.mu.Lock()
defer m.mu.Unlock()
status := m.losses.Snapshot(now)
status.CapacityUnavailable = true
var waitingLate, runningLate bool
for execution := range m.pending {
if execution.started.IsZero() {
status.Depth++
status.LagSeconds = max(status.LagSeconds, now.Sub(execution.queued).Seconds())
waitingLate = waitingLate || overdue(now, execution.deadline)
} else {
status.InFlight++
status.ProcessingSeconds = max(status.ProcessingSeconds, now.Sub(execution.started).Seconds())
runningLate = runningLate || overdue(now, execution.deadline)
}
}
// Each call has its own deadline. A heavy check must not lend its longer
// budget to an overdue short check, or get a fresh budget after dispatch.
switch {
case waitingLate:
status.Reason = "backlog_lag"
case runningLate:
status.Reason = "processing_lag"
}
if status.Reason != "" {
status.Status = "degraded"
}
return status
}
// CheckExecutionQueueStatus reads memory only, including after a runner exits.
func CheckExecutionQueueStatus(now time.Time) queuehealth.Status {
return checkExecutions.QueueStatus(now)
}
// overdue reports whether a bounded deadline has passed. Work with no deadline
// has no bound to miss.
func overdue(now, deadline time.Time) bool {
return !deadline.IsZero() && !now.Before(deadline)
}
package checks
import (
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"io"
"math"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"syscall"
"unicode"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/actionlog"
)
// CleanResult describes the outcome of a cleaning attempt.
type CleanResult struct {
Path string
Cleaned bool
BackupPath string
Removals []string // descriptions of what was removed
Error string
// Refused marks a safety refusal without changing the customer file.
// It consumes attempt capacity but does not indicate a broken action.
Refused bool
}
var (
// Leading class is [\s\x00\x0b] rather than \s: an injection can pad
// the line start with NUL or vertical-tab bytes that Go's \s does not
// cover. Detection still flags the file (analyzePHPContent is not
// anchored); widening the class here lets the surgical cleaner strip
// the injection instead of falling back to whole-file quarantine.
cleanRegexpMaliciousInclude = regexp.MustCompile(
`(?i)^[\s\x00\x0b]*@include\s*\(\s*(?:` +
`['"](?:/tmp/|/dev/shm/|/var/tmp/)` +
`|base64_decode\s*\(` +
`|str_rot13\s*\(` +
`|gzinflate\s*\(` +
`)`)
cleanRegexpVarInclude = regexp.MustCompile(`(?i)^[\s\x00\x0b]*@include\s*\(\s*\$[a-zA-Z_]+\s*\)`)
cleanRegexpCloseOpen = regexp.MustCompile(`(?i)\?>\s*<\?php`)
cleanRegexpInlineEval = regexp.MustCompile(
`(?i)^[\s\x00\x0b]*(?:@?)eval\s*\(\s*(?:base64_decode|gzinflate|gzuncompress|str_rot13)\s*\(`)
cleanRegexpMultiB64 = regexp.MustCompile(`(?i)(?:base64_decode\s*\(\s*){2,}`)
cleanRegexpChainedB64 = regexp.MustCompile(`(?i)\$\w+\s*=\s*base64_decode\s*\(\s*base64_decode`)
cleanRegexpChrChain = regexp.MustCompile(`(?i)\bchr\s*\(\s*\d+\s*\)(?:\s*\.?\s*\bchr\s*\(\s*\d+\s*\)){4,}`)
cleanRegexpPackHex = regexp.MustCompile(`(?i)pack\s*\(\s*["']H\*["']\s*,`)
cleanRegexpHexVar = regexp.MustCompile(`(?:"\x5c\x78[0-9a-fA-F]{2}){3,}|(?:\\x[0-9a-fA-F]{2}){3,}`)
)
// cleanMaxFileSize bounds how large a file surgical cleaning will read
// into memory. The detector that routes a file here (analyzePHPContent)
// only inspects bounded head and tail windows, so an attacker can match a
// signature inside those windows and pad the rest to many gigabytes. Reading
// that whole file with io.ReadAll plus the strings.Split and regex passes
// below would OOM the root daemon. Above this ceiling we refuse; the caller
// decides whether whole-file quarantine is safe for that remediation path.
// Legitimate plugin/theme PHP files are far smaller than this. Var, not
// const, so tests can lower it. 8 MiB.
var cleanMaxFileSize int64 = 8 << 20
// CleanInfectedFile attempts to surgically remove malicious code from a PHP file
// while preserving the legitimate content. Always creates a backup first.
//
// Cleaning strategies (tried in order):
// 1. @include injection - remove @include lines pointing to /tmp, eval, base64, or via variables
// 2. Prepend injection - remove malicious code blocks at start of file (entropy-validated)
// 3. Append injection - remove malicious code after closing ?> or end of PSR-12 file
// 4. Inline eval injection - remove eval(base64_decode(...)) single-line injections
func CleanInfectedFile(path string) CleanResult {
return cleanInfectedFileIdentified(path, nil)
}
func cleanInfectedFileIdentified(path string, expected os.FileInfo) (result CleanResult) {
result = CleanResult{Path: path}
audit := newCleanAction(path)
defer func() { audit.finish(result.Error) }()
target, err := openCleanTarget(path)
if err != nil {
result.Refused = errors.Is(fileResponseSourceError(err), errFileResponseRefused)
result.Error = fmt.Sprintf("cannot read file: %v", err)
return result
}
defer target.Close()
if expected != nil && (!sameFileIdentity(expected, target.Info) || !sameContentShape(expected, target.Info)) {
result.Refused = true
result.Error = "file changed before automatic cleaning"
return result
}
if sz := target.Info.Size(); sz > cleanMaxFileSize {
result.Refused = true
result.Error = fmt.Sprintf("file too large to clean (%d bytes > %d)", sz, cleanMaxFileSize)
return result
}
audit.rec.Result = actionlog.Failed
data, err := io.ReadAll(target.File)
if err != nil {
result.Error = fmt.Sprintf("cannot read file: %v", err)
return result
}
audit.capture(target, data)
// Create backup before any modification
backupDir := filepath.Join(quarantineDir, "pre_clean")
backupPath := newQuarantinePath(backupDir, path)
// Metadata sidecar derived from the same fd we read, so a directory
// race after open cannot change the metadata we record.
meta := quarantineMetadata(path, target.Info, "Pre-clean backup (surgical cleaning)")
if err := storeQuarantineBackup(backupPath, data, meta, 0600); err != nil {
result.Error = fmt.Sprintf("cannot create durable backup: %v", err)
return result
}
result.BackupPath = backupPath
content := string(data)
originalLen := len(content)
var removals []string
// Strategy 1: Remove @include injections (including variable-based)
content, removed := removeIncludeInjections(content)
removals = append(removals, removed...)
// Strategy 2: Remove prepend injections (with entropy validation)
content, removed = removePrependInjection(content)
removals = append(removals, removed...)
// Strategy 3: Remove append injections (handles files with and without closing ?>)
content, removed = removeAppendInjection(content)
removals = append(removals, removed...)
// Strategy 4: Remove inline eval(base64_decode(...)) injections
content, removed = removeInlineEvalInjections(content)
removals = append(removals, removed...)
// Strategy 5: Remove multi-layer base64 decode chains
content, removed = removeMultiLayerBase64(content)
removals = append(removals, removed...)
// Strategy 6: Remove chr()/pack() constructed code
content, removed = removeChrPackInjections(content)
removals = append(removals, removed...)
// Strategy 7: Remove hex-encoded variable injections
content, removed = removeHexVarInjections(content)
removals = append(removals, removed...)
// If nothing was removed, file couldn't be cleaned
if len(removals) == 0 || len(content) == originalLen {
audit.rec.Result = actionlog.Refused
result.Refused = true
result.Error = "no known injection patterns found - file may need manual review"
return result
}
audit.rec.Reason = strings.Join(removals, "; ")
if err := audit.replace(target, []byte(content), backupPath); err != nil {
result.Refused = errors.Is(err, errFileResponseRefused)
result.Error = fmt.Sprintf("cannot write cleaned file: %v", err)
return result
}
result.Cleaned = true
result.Removals = removals
return result
}
type cleanTarget struct {
Path string
DirFD int
Name string
File *os.File
Info os.FileInfo
UID int
GID int
replacementInfo os.FileInfo
installed bool
OwnerKnown bool
}
func (t *cleanTarget) Close() {
if t.File != nil {
_ = t.File.Close()
}
if t.DirFD >= 0 {
_ = unix.Close(t.DirFD)
}
}
func openCleanTarget(path string) (*cleanTarget, error) {
dir, name := filepath.Split(path)
if name == "" || name == "." || name == ".." {
return nil, fmt.Errorf("invalid target path %q", path)
}
if dir == "" {
dir = "."
}
dir = filepath.Clean(dir)
parentInfo, err := os.Lstat(dir)
if err != nil {
return nil, fmt.Errorf("stat parent directory: %w", err)
}
if parentInfo.Mode()&os.ModeSymlink != 0 {
return nil, refuseFileResponse(errors.New("refusing symlinked parent directory"))
}
// Pin the immediate parent. A swap of that directory to a symlink
// after detection must not redirect either the read or the writeback.
dirFD, err := unix.Open(dir, unix.O_RDONLY|unix.O_DIRECTORY|unix.O_NOFOLLOW|unix.O_CLOEXEC, 0)
if err != nil {
return nil, fmt.Errorf("open parent directory: %w", err)
}
closeDir := true
defer func() {
if closeDir {
_ = unix.Close(dirFD)
}
}()
if pinErr := verifyCleanParentStillPinned(dirFD, parentInfo); pinErr != nil {
return nil, pinErr
}
fd, err := unix.Openat(dirFD, name, unix.O_RDONLY|unix.O_NOFOLLOW|unix.O_CLOEXEC|unix.O_NONBLOCK, 0)
if err != nil {
return nil, err
}
// #nosec G115 -- unix.Openat returned a non-negative fd because err is nil.
file := os.NewFile(uintptr(fd), path)
closeFile := true
defer func() {
if closeFile {
_ = file.Close()
}
}()
info, err := file.Stat()
if err != nil {
return nil, fmt.Errorf("stat file: %w", err)
}
if !info.Mode().IsRegular() {
return nil, refuseFileResponse(fmt.Errorf("refusing non-regular file (mode=%v)", info.Mode()))
}
target := &cleanTarget{
Path: path,
DirFD: dirFD,
Name: name,
File: file,
Info: info,
}
if stat, ok := info.Sys().(*syscall.Stat_t); ok {
target.UID = int(stat.Uid)
target.GID = int(stat.Gid)
target.OwnerKnown = true
}
closeDir = false
closeFile = false
return target, nil
}
func verifyCleanParentStillPinned(dirFD int, want os.FileInfo) error {
var got unix.Stat_t
if err := unix.Fstat(dirFD, &got); err != nil {
return fmt.Errorf("stat opened parent directory: %w", err)
}
if !sameUnixStatIdentity(want, got) {
return refuseFileResponse(errors.New("parent directory changed during cleaning"))
}
return nil
}
func sameUnixStatIdentity(info os.FileInfo, stat unix.Stat_t) bool {
if info == nil {
return false
}
want, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return false
}
// #nosec G115 -- Dev is signed on darwin and unsigned on linux; widening
// both sides to uint64 is what makes the comparison portable.
return uint64(want.Dev) == uint64(stat.Dev) && uint64(want.Ino) == uint64(stat.Ino)
}
var closeCleanTemp = (*os.File).Close
var syncCleanParent = unix.Fsync
// writeCleanedFileAtomic stages cleaned content through a hidden sibling
// name under the pinned parent directory and renames it over the original
// only after the path still resolves to the inode we read.
func writeCleanedFileAtomic(target *cleanTarget, content []byte) error {
tmp, tmpName, err := createCleanTempFile(target.DirFD)
if err != nil {
return err
}
removeTmp := true
defer func() {
if removeTmp {
_ = unix.Unlinkat(target.DirFD, tmpName, 0)
}
_ = tmp.Close()
}()
if _, err := tmp.Write(content); err != nil {
return err
}
if target.OwnerKnown {
if err := tmp.Chown(target.UID, target.GID); err != nil {
return err
}
}
if err := tmp.Chmod(cleanReplacementMode(target.Info.Mode())); err != nil {
return err
}
if err := tmp.Sync(); err != nil {
return err
}
replacementInfo, statErr := tmp.Stat()
if statErr != nil {
return statErr
}
if err := closeCleanTemp(tmp); err != nil {
return err
}
if err := verifyCleanTargetUnchanged(target); err != nil {
return err
}
if err := unix.Renameat(target.DirFD, tmpName, target.DirFD, target.Name); err != nil {
return err
}
removeTmp = false
target.installed = true
target.replacementInfo = replacementInfo
if err := syncCleanParent(target.DirFD); err != nil {
return fmt.Errorf("cleaned file installed but directory sync failed; recovery backup retained: %w", err)
}
return nil
}
func createCleanTempFile(dirFD int) (*os.File, string, error) {
for i := 0; i < 100; i++ {
var raw [8]byte
if _, err := rand.Read(raw[:]); err != nil {
return nil, "", fmt.Errorf("random temp name: %w", err)
}
name := ".csm-clean-" + hex.EncodeToString(raw[:])
fd, err := unix.Openat(dirFD, name, unix.O_WRONLY|unix.O_CREAT|unix.O_EXCL|unix.O_NOFOLLOW|unix.O_CLOEXEC, 0o600)
if err == unix.EEXIST {
continue
}
if err != nil {
return nil, "", err
}
// #nosec G115 -- unix.Openat returned a non-negative fd because err is nil.
return os.NewFile(uintptr(fd), name), name, nil
}
return nil, "", fmt.Errorf("could not allocate temp file name")
}
func verifyCleanTargetUnchanged(target *cleanTarget) error {
fd, err := unix.Openat(target.DirFD, target.Name, unix.O_RDONLY|unix.O_NOFOLLOW|unix.O_CLOEXEC|unix.O_NONBLOCK, 0)
if err != nil {
return fmt.Errorf("open target before rename: %w", fileResponseSourceError(err))
}
// #nosec G115 -- unix.Openat returned a non-negative fd because err is nil.
file := os.NewFile(uintptr(fd), target.Path)
defer func() { _ = file.Close() }()
info, err := file.Stat()
if err != nil {
return fmt.Errorf("stat target before rename: %w", err)
}
if !info.Mode().IsRegular() {
return refuseFileResponse(fmt.Errorf("refusing non-regular target before rename (mode=%v)", info.Mode()))
}
if !sameFileIdentity(info, target.Info) || !sameCleanContentShape(info, target.Info) {
return refuseFileResponse(errors.New("file changed during cleaning"))
}
return nil
}
func sameCleanContentShape(a, b os.FileInfo) bool {
if a == nil || b == nil {
return false
}
return a.Size() == b.Size() && a.ModTime().Equal(b.ModTime())
}
func cleanReplacementMode(mode os.FileMode) os.FileMode {
return mode & (os.ModePerm | os.ModeSetuid | os.ModeSetgid | os.ModeSticky)
}
// ShouldCleanInsteadOfQuarantine returns true if the file should be cleaned
// (surgical removal) instead of quarantined (full removal).
// WP core files and plugin files are better cleaned - removing them breaks the site.
// Unknown standalone files (droppers, webshells) should be quarantined.
func ShouldCleanInsteadOfQuarantine(path string) bool {
// WP core files - always clean, never quarantine
if strings.Contains(path, "/wp-includes/") || strings.Contains(path, "/wp-admin/") {
return true
}
// Plugin/theme main files - clean to preserve functionality
if strings.Contains(path, "/wp-content/plugins/") || strings.Contains(path, "/wp-content/themes/") {
// But not if the file itself is the malware (h4x0r.php inside a theme)
name := strings.ToLower(filepath.Base(path))
if isWebshellName(name) {
return false // quarantine this - it's a standalone webshell
}
return true
}
// Everything else - quarantine
return false
}
// removeIncludeInjections removes @include lines that load malicious files.
// Catches:
// - @include("/tmp/...") - literal paths to temp dirs
// - @include(base64_decode("...")) - encoded includes
// - @include($var) where $var is built from obfuscated strings nearby
func removeIncludeInjections(content string) (string, []string) {
var removals []string
lines := strings.Split(content, "\n")
var clean []string
for i, line := range lines {
if cleanRegexpMaliciousInclude.MatchString(line) {
removals = append(removals, fmt.Sprintf("removed @include injection: %s", trimCleanRemovalLine(line)))
continue
}
// Variable-based @include - check surrounding context for obfuscation
if cleanRegexpVarInclude.MatchString(line) {
context := getLineContext(lines, i, 3)
contextLower := strings.ToLower(context)
isObfuscated := strings.Contains(contextLower, "base64_decode") ||
strings.Contains(contextLower, "str_rot13") ||
strings.Contains(contextLower, "chr(") ||
strings.Contains(contextLower, `"\x`) ||
strings.Count(contextLower, ". ") > 5 // heavy string concatenation
if isObfuscated {
removals = append(removals, fmt.Sprintf("removed obfuscated @include: %s", trimCleanRemovalLine(line)))
continue
}
}
clean = append(clean, line)
}
return strings.Join(clean, "\n"), removals
}
func trimCleanRemovalLine(line string) string {
return strings.TrimFunc(line, func(r rune) bool {
return r == '\x00' || unicode.IsSpace(r)
})
}
// removePrependInjection removes malicious PHP code injected before the
// legitimate file content. Uses entropy analysis to verify the prefix
// is actually obfuscated (not legitimate minified code).
func removePrependInjection(content string) (string, []string) {
var removals []string
trimmed := strings.TrimSpace(content)
// PHP open tags are case-insensitive (<?php, <?PHP, <?Php all execute), so
// match without lowercasing the whole file.
if len(trimmed) < 5 || !strings.EqualFold(trimmed[:5], "<?php") {
return content, nil
}
// Find if there's a malicious block at the start followed by ?><?php
loc := cleanRegexpCloseOpen.FindStringIndex(content)
if loc == nil {
return content, nil
}
prefix := content[:loc[0]]
prefixLower := strings.ToLower(prefix)
// Check if the prefix contains malicious patterns
hasMaliciousPatterns := strings.Contains(prefixLower, "eval(") ||
strings.Contains(prefixLower, "base64_decode") ||
strings.Contains(prefixLower, "gzinflate") ||
strings.Contains(prefixLower, "str_rot13") ||
strings.Contains(prefixLower, "@include")
if !hasMaliciousPatterns {
return content, nil
}
// Additional safety: verify the prefix has high entropy (obfuscated code)
// or contains long encoded strings. This prevents false positives on
// legitimate minified PHP that happens to use ?><?php patterns.
entropy := shannonEntropy(prefix)
hasLongStrings := containsLongEncodedString(prefix, 100)
if entropy < 4.5 && !hasLongStrings {
// Low entropy and no long encoded strings - likely legitimate code
return content, nil
}
// Remove everything before the second <?php
cleaned := "<?php" + content[loc[1]:]
removals = append(removals, fmt.Sprintf("removed %d-byte prepend injection (entropy: %.2f)", loc[1], entropy))
return cleaned, removals
}
// removeAppendInjection removes malicious code appended after the end of
// legitimate PHP content. Handles both files with closing ?> and files
// without (PSR-12 style).
func removeAppendInjection(content string) (string, []string) {
var removals []string
// Case 1: File has closing ?> with malicious code after it
lastClose := strings.LastIndex(content, "?>")
if lastClose >= 0 {
after := content[lastClose+2:]
afterTrimmed := strings.TrimSpace(after)
if afterTrimmed != "" {
afterLower := strings.ToLower(afterTrimmed)
isMalicious := strings.Contains(afterLower, "eval(") ||
strings.Contains(afterLower, "base64_decode") ||
strings.Contains(afterLower, "gzinflate") ||
strings.Contains(afterLower, "system(") ||
strings.Contains(afterLower, "exec(") ||
strings.Contains(afterLower, "@include") ||
strings.Contains(afterLower, "<?php") // second PHP block appended
if isMalicious {
cleaned := content[:lastClose+2] + "\n"
removals = append(removals, fmt.Sprintf("removed %d-byte append injection (after ?>)", len(after)))
return cleaned, removals
}
}
}
// Case 2: PSR-12 style file (no closing ?>) - check if there's a malicious
// block appended at the very end, separated by multiple newlines
lines := strings.Split(content, "\n")
if len(lines) < 5 {
return content, nil
}
// Check last 10 lines for injected code block
startCheck := len(lines) - 10
if startCheck < 0 {
startCheck = 0
}
blankLineIdx := -1
for i := startCheck; i < len(lines); i++ {
if strings.TrimSpace(lines[i]) == "" && blankLineIdx < 0 {
blankLineIdx = i
}
}
if blankLineIdx >= 0 {
tailBlock := strings.Join(lines[blankLineIdx:], "\n")
tailLower := strings.ToLower(tailBlock)
if strings.Contains(tailLower, "eval(") && strings.Contains(tailLower, "base64_decode") {
cleaned := strings.Join(lines[:blankLineIdx], "\n") + "\n"
removals = append(removals, fmt.Sprintf("removed %d-byte PSR-12 append injection", len(tailBlock)))
return cleaned, removals
}
}
return content, nil
}
// removeInlineEvalInjections removes single-line eval(base64_decode("..."));
// injections that are inserted as standalone lines in PHP files.
func removeInlineEvalInjections(content string) (string, []string) {
var removals []string
lines := strings.Split(content, "\n")
var clean []string
for _, line := range lines {
trimmedLine := strings.TrimSpace(line)
// Strip /* ... */ block comments and trailing // / # comments
// before matching. Attackers wedge comments between the keyword
// and the open paren ("@eval/*x*/(base64_decode(...))") to slip
// past a strict regex; the cleaner has to see the line the way
// the PHP tokenizer does, not byte-for-byte.
normalized := stripPHPComments(trimmedLine)
if cleanRegexpInlineEval.MatchString(normalized) {
// Length gate stays on the ORIGINAL line so a short legitimate
// eval() does not get sucked into the removal path.
if len(trimmedLine) > 50 {
removals = append(removals, fmt.Sprintf("removed inline eval injection (%d chars)", len(trimmedLine)))
continue
}
}
clean = append(clean, line)
}
return strings.Join(clean, "\n"), removals
}
// stripPHPComments removes PHP comments while leaving quoted strings
// intact, so comment-looking payload data does not change the code the
// cleaner evaluates.
func stripPHPComments(line string) string {
return strings.TrimSpace(stripPHPCommentsFromCode(line))
}
// --- Helper functions ---
// shannonEntropy calculates the Shannon entropy of a string.
// Obfuscated/encoded code typically has entropy > 5.0.
// Normal PHP code typically has entropy 4.0-4.5.
func shannonEntropy(s string) float64 {
if len(s) == 0 {
return 0
}
freq := make(map[byte]float64)
for i := 0; i < len(s); i++ {
freq[s[i]]++
}
length := float64(len(s))
entropy := 0.0
for _, count := range freq {
p := count / length
if p > 0 {
entropy -= p * math.Log2(p)
}
}
return entropy
}
// containsLongEncodedString checks if the text contains a long base64-like
// string (alphanumeric + /+ without spaces).
func containsLongEncodedString(s string, minLength int) bool {
if minLength <= 0 {
return true
}
run := 0
for i := 0; i < len(s); i++ {
if isEncodedStringByte(s[i]) {
run++
if run >= minLength {
return true
}
continue
}
run = 0
}
return false
}
func isEncodedStringByte(b byte) bool {
return b >= 'A' && b <= 'Z' ||
b >= 'a' && b <= 'z' ||
b >= '0' && b <= '9' ||
b == '+' || b == '/' || b == '='
}
// getLineContext returns N lines before and after the given line index.
func getLineContext(lines []string, idx, window int) string {
start := idx - window
if start < 0 {
start = 0
}
end := idx + window + 1
if end > len(lines) {
end = len(lines)
}
return strings.Join(lines[start:end], "\n")
}
// FormatCleanResult returns a human-readable summary of a clean operation.
func FormatCleanResult(r CleanResult) string {
if r.Refused {
return fmt.Sprintf("REFUSED to clean %s: %s", r.Path, r.Error)
}
if r.Error != "" {
return fmt.Sprintf("FAILED to clean %s: %s", r.Path, r.Error)
}
if !r.Cleaned {
return fmt.Sprintf("No changes made to %s", r.Path)
}
var b strings.Builder
fmt.Fprintf(&b, "CLEANED %s\n", r.Path)
fmt.Fprintf(&b, " Backup: %s\n", r.BackupPath)
for _, removal := range r.Removals {
fmt.Fprintf(&b, " - %s\n", removal)
}
return b.String()
}
// --- Strategy 5: Multi-layer base64 decode chains ---
// Catches: eval(base64_decode(base64_decode("...")))
// Catches: $x=base64_decode("...");$y=base64_decode($x);eval($y);
func removeMultiLayerBase64(content string) (string, []string) {
var removals []string
lines := strings.Split(content, "\n")
var clean []string
for _, line := range lines {
trimmed := strings.TrimSpace(line)
if len(trimmed) < 80 {
clean = append(clean, line)
continue
}
if cleanRegexpMultiB64.MatchString(trimmed) || cleanRegexpChainedB64.MatchString(trimmed) {
removals = append(removals, fmt.Sprintf("removed multi-layer base64 chain (%d chars)", len(trimmed)))
continue
}
clean = append(clean, line)
}
return strings.Join(clean, "\n"), removals
}
// --- Strategy 6: chr()/pack() constructed code ---
// Catches: eval(chr(115).chr(121).chr(115)...);
// Catches: $f = pack("H*", "73797374656d"); $f($_POST['cmd']);
func removeChrPackInjections(content string) (string, []string) {
var removals []string
lines := strings.Split(content, "\n")
// Find chr-chain spans across the full content first so multi-line
// chains (chr(115)\n.chr(121)\n...) are recognized as one chain
// instead of slipping past the per-line 5+ count.
dropLine := make(map[int]bool)
offsets := make([]int, len(lines))
off := 0
for i, line := range lines {
offsets[i] = off
off += len(line) + 1 // +1 for the newline strings.Split consumed
}
for _, span := range cleanRegexpChrChain.FindAllStringIndex(content, -1) {
startLine, endLine := chrChainStatementLineRange(lines, offsets, span)
for i := startLine; i <= endLine; i++ {
dropLine[i] = true
}
}
var clean []string
for i, line := range lines {
if dropLine[i] {
removals = append(removals, fmt.Sprintf("removed chr() chain injection (line %d, %d chars)", i+1, len(strings.TrimSpace(line))))
continue
}
trimmed := strings.TrimSpace(line)
if cleanRegexpPackHex.MatchString(trimmed) {
lower := strings.ToLower(trimmed)
// Only remove if combined with execution
if strings.Contains(lower, "eval") || strings.Contains(lower, "$_") ||
strings.Contains(lower, "system") || strings.Contains(lower, "exec") {
removals = append(removals, fmt.Sprintf("removed pack() code construction (line %d)", i+1))
continue
}
}
clean = append(clean, line)
}
return strings.Join(clean, "\n"), removals
}
func chrChainStatementLineRange(lines []string, offsets []int, span []int) (int, int) {
startLine := lineIndexForOffset(offsets, span[0])
endLine := lineIndexForOffset(offsets, span[1]-1)
for startLine > 0 && chrChainPrefixContinues(lines[startLine-1]) {
startLine--
}
for endLine+1 < len(lines) && !strings.Contains(lines[endLine], ";") && chrChainSuffixContinues(lines[endLine+1]) {
endLine++
if strings.Contains(lines[endLine], ";") {
break
}
}
return startLine, endLine
}
func chrChainPrefixContinues(line string) bool {
trimmed := strings.TrimSpace(line)
if trimmed == "" || trimmed == "<?php" {
return false
}
return strings.HasSuffix(trimmed, "=") ||
strings.HasSuffix(trimmed, ".") ||
strings.HasSuffix(trimmed, "(") ||
strings.HasSuffix(trimmed, ",") ||
strings.HasSuffix(trimmed, "[")
}
func chrChainSuffixContinues(line string) bool {
trimmed := strings.TrimSpace(line)
if trimmed == "" {
return true
}
return strings.HasPrefix(trimmed, ".") ||
strings.HasPrefix(trimmed, ")") ||
strings.HasPrefix(trimmed, ",") ||
strings.HasPrefix(trimmed, ";")
}
// lineIndexForOffset returns the index of the line containing the byte
// at offset within the original content.
func lineIndexForOffset(offsets []int, byteOffset int) int {
if len(offsets) == 0 || byteOffset < 0 {
return 0
}
idx := sort.Search(len(offsets), func(i int) bool {
return offsets[i] > byteOffset
}) - 1
if idx < 0 {
return 0
}
if idx >= len(offsets) {
return len(offsets) - 1
}
return idx
}
// --- Strategy 7: Hex-encoded variable injections ---
// Catches: $GLOBALS["\x61\x64\x6d\x69\x6e"] = eval(...)
// Catches: ${"\x47\x4c\x4f\x42\x41\x4c\x53"}[...] = ...
func removeHexVarInjections(content string) (string, []string) {
var removals []string
lines := strings.Split(content, "\n")
var clean []string
for i, line := range lines {
trimmed := strings.TrimSpace(line)
if len(trimmed) < 30 {
clean = append(clean, line)
continue
}
if cleanRegexpHexVar.MatchString(trimmed) {
lower := strings.ToLower(trimmed)
// Only remove if combined with dangerous operations
if strings.Contains(lower, "eval") || strings.Contains(lower, "system(") ||
strings.Contains(lower, "exec(") || strings.Contains(lower, "base64_decode") ||
strings.Contains(lower, "assert(") || strings.Contains(lower, "$_post") ||
strings.Contains(lower, "$_request") || strings.Contains(lower, "$_get") {
removals = append(removals, fmt.Sprintf("removed hex-encoded variable injection (line %d, %d chars)", i+1, len(trimmed)))
continue
}
}
clean = append(clean, line)
}
return strings.Join(clean, "\n"), removals
}
package checks
import (
"errors"
"github.com/pidginhost/csm/internal/actionlog"
)
type cleanAction struct{ rec actionlog.Record }
func newCleanAction(path string) *cleanAction {
return &cleanAction{rec: actionlog.Record{Op: "respond.clean_file", Target: path, Before: actionlog.Metadata(path), Result: actionlog.Refused}}
}
func (a *cleanAction) capture(target *cleanTarget, data []byte) {
a.rec.Target = target.Path
a.rec.Before = actionlog.ContentState(target.Info, data)
}
func (a *cleanAction) replace(target *cleanTarget, content []byte, backupPath string) error {
a.rec.Result = actionlog.Failed
a.rec.RecoveryPath = backupPath
err := writeCleanedFileAtomic(target, content)
// Rename can succeed even when the following directory sync fails. The
// record must still describe the installed bytes and retained recovery copy.
if target.installed {
a.rec.Result = actionlog.Applied
a.rec.After = actionlog.ContentState(target.replacementInfo, content)
} else if errors.Is(err, errFileResponseRefused) {
a.rec.Result = actionlog.Refused
}
return err
}
func (a *cleanAction) finish(message string) {
a.rec.Error = message
if a.rec.After == nil {
a.rec.After = actionlog.Metadata(a.rec.Target)
}
actionlog.Write(a.rec)
}
package checks
import (
"context"
"errors"
"io"
"os"
"golang.org/x/sys/unix"
)
const maxCMSConfigBytes = 1 << 20
var errCMSConfigTooLarge = errors.New("CMS configuration exceeds the read limit")
// Configuration paths belong to tenants. Open without following the final
// symlink or waiting for a FIFO writer, then validate the opened object.
func openCMSConfig(path string) (*os.File, error) {
if fs, production := osFS.(realOS); production {
return fs.openRegularFile(path, unix.O_NOFOLLOW)
}
file, err := osFS.Open(path)
if err != nil {
return nil, err
}
info, err := file.Stat()
if err != nil {
_ = file.Close()
return nil, err
}
if !info.Mode().IsRegular() {
_ = file.Close()
return nil, errNonRegularFile
}
return file, nil
}
func readCMSConfig(ctx context.Context, path string) ([]byte, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
file, err := openCMSConfig(path)
if err != nil {
return nil, err
}
defer file.Close()
before, err := file.Stat()
if err != nil {
return nil, err
}
if before.Size() > maxCMSConfigBytes {
return nil, errCMSConfigTooLarge
}
data, err := io.ReadAll(io.LimitReader(cmsConfigReader{ctx: ctx, file: file}, maxCMSConfigBytes+1))
if err != nil {
return nil, err
}
if err = ctx.Err(); err != nil {
return nil, err
}
if len(data) > maxCMSConfigBytes {
return nil, errCMSConfigTooLarge
}
after, err := file.Stat()
if err != nil {
return nil, err
}
if !sameFileSnapshot(before, after) || int64(len(data)) != before.Size() {
return nil, errFileChanged
}
return data, nil
}
type cmsConfigReader struct {
ctx context.Context
file *os.File
}
func (r cmsConfigReader) Read(p []byte) (int, error) {
if err := r.ctx.Err(); err != nil {
return 0, err
}
return r.file.Read(p)
}
package checks
import (
"crypto/sha256"
"encoding/hex"
"io"
"sync"
)
// CMSHashCache stores SHA256 hashes of verified-clean CMS core files.
// After wp core verify-checksums confirms an installation is clean,
// all its core files are hashed and cached. The real-time scanner
// checks this cache before reporting signature matches - if a file's
// hash is in the cache, it's a known-clean CMS file and signature
// matches on it are false positives.
type CMSHashCache struct {
mu sync.RWMutex
hashes map[string]bool // SHA256 hex -> true
sizes map[int64]bool // file sizes that can possibly match a cached hash
}
var (
globalCache *CMSHashCache
globalCacheOnce sync.Once
)
// GlobalCMSCache returns the singleton cache, creating it on first call.
func GlobalCMSCache() *CMSHashCache {
globalCacheOnce.Do(func() {
globalCache = &CMSHashCache{
hashes: make(map[string]bool),
sizes: make(map[int64]bool),
}
})
return globalCache
}
// Add inserts a file hash and its content size into the cache.
func (c *CMSHashCache) Add(hash string, size int64) {
c.mu.Lock()
c.hashes[hash] = true
c.sizes[size] = true
c.mu.Unlock()
}
// Contains checks if a file hash is in the cache.
func (c *CMSHashCache) Contains(hash string) bool {
c.mu.RLock()
ok := c.hashes[hash]
c.mu.RUnlock()
return ok
}
// Size returns the number of cached hashes.
func (c *CMSHashCache) Size() int {
c.mu.RLock()
n := len(c.hashes)
c.mu.RUnlock()
return n
}
// MayContainSize reports whether a verified file of size bytes was cached.
// It lets realtime scanning reject attacker-sized files before hashing them.
func (c *CMSHashCache) MayContainSize(size int64) bool {
c.mu.RLock()
ok := c.sizes[size]
c.mu.RUnlock()
return ok
}
// Clear removes all cached hashes (used before rebuilding).
func (c *CMSHashCache) Clear() {
c.mu.Lock()
c.hashes = make(map[string]bool)
c.sizes = make(map[int64]bool)
c.mu.Unlock()
}
// HashFile computes the SHA256 hash of a file. Returns empty string on error.
func HashFile(path string) string {
f, err := osFS.Open(path)
if err != nil {
return ""
}
defer func() { _ = f.Close() }()
h := sha256.New()
if _, err := io.Copy(h, f); err != nil {
return ""
}
return hex.EncodeToString(h.Sum(nil))
}
// IsVerifiedCMSFile checks if a file at the given path matches a
// known-clean CMS core file by comparing its SHA256 hash against the cache.
//
// The cache is keyed by SHA256 hash alone (not path+hash) - this is correct:
// - SHA256 preimage resistance makes it computationally infeasible for an
// attacker to craft a file that produces the same hash as a legitimate
// WP core file. Birthday attacks do not apply here because the attacker
// must hit a specific pre-existing hash, not merely find any collision.
// - If file content matches a known WP core file byte-for-byte, it IS that
// file regardless of where it is located on disk. The path is irrelevant
// to whether the content is clean.
func IsVerifiedCMSFile(path string) bool {
cache := GlobalCMSCache()
if cache.Size() == 0 {
return false
}
hash := HashFile(path)
if hash == "" {
return false
}
return cache.Contains(hash)
}
// IsVerifiedCMSHash reports whether a content hash belongs to a verified CMS
// core file. Callers that already hold the content -- the realtime scanner
// hashes the event descriptor -- use this instead of IsVerifiedCMSFile, whose
// re-read resolves the path again and can hash something other than what was
// examined.
func IsVerifiedCMSHash(hash string) bool {
if hash == "" {
return false
}
cache := GlobalCMSCache()
if cache.Size() == 0 {
return false
}
return cache.Contains(hash)
}
// CMSCacheEmpty reports whether any verified core files are cached, so a caller
// can skip hashing entirely when the answer cannot be yes.
func CMSCacheEmpty() bool { return GlobalCMSCache().Size() == 0 }
// CMSCacheMayContainSize reports whether hashing a file of size bytes can
// possibly produce a cached CMS hash.
func CMSCacheMayContainSize(size int64) bool { return GlobalCMSCache().MayContainSize(size) }
package checks
import (
"context"
"encoding/hex"
"fmt"
"net"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/netutil"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// safeRemotePorts and safeUsers are package-level so both legacy polling and
// the BPF live coordinator share one source of truth.
var safeRemotePorts = map[uint16]bool{
53: true, 80: true, 443: true, 25: true, 587: true, 465: true,
993: true, 995: true, 110: true, 143: true,
}
var safeUsers = map[string]bool{
"imunify360-webshield": true,
"named": true,
"mysql": true,
"memcached": true,
"icinga": true,
"dovecot": true,
"mailman": true,
}
var serverLocalPorts = map[uint16]bool{
21: true, 25: true, 26: true, 53: true, 80: true, 110: true,
143: true, 443: true, 465: true, 587: true, 993: true, 995: true,
2082: true, 2083: true, 2086: true, 2087: true, 2095: true, 2096: true,
3306: true, 4190: true,
52223: true, 52224: true, 52227: true, 52228: true,
52229: true, 52230: true, 52231: true, 52232: true,
}
// EvaluateConnection returns a populated alert.Finding and true when the
// connection should be reported, or a zero finding and false when it should
// be ignored. Host-interface lookups are cached. Used by the BPF live backend
// (per-event) and the polling backend (per row of /proc/net/tcp[6]).
func EvaluateConnection(
cfg *config.Config,
uid uint32,
dstIP net.IP,
dstPort uint16,
localPort uint16,
proto string,
user string,
) (alert.Finding, bool) {
if uid == 0 {
return alert.Finding{}, false
}
if dstIP == nil || dstIP.IsLoopback() || dstIP.IsUnspecified() {
return alert.Finding{}, false
}
// Panel front ends proxy to their own backend over the machine's public
// address rather than loopback, so the packet never leaves the host and is
// no more an outbound connection than the loopback case above. Left
// reported, the proxy hop is classed as C2 traffic and drives the host's
// own address to a critical local threat score.
if netutil.IsHostAddress(dstIP.String()) {
return alert.Finding{}, false
}
if serverLocalPorts[localPort] {
return alert.Finding{}, false
}
if safeRemotePorts[dstPort] {
return alert.Finding{}, false
}
if isInfraIP(dstIP.String(), cfg.InfraIPs) {
return alert.Finding{}, false
}
if safeUsers[user] {
return alert.Finding{}, false
}
dst := dstIP.String()
if dstIP.To4() == nil {
dst = "[" + dst + "]"
}
return alert.Finding{
Severity: alert.High,
Check: "user_outbound_connection",
Message: fmt.Sprintf("Non-root user connecting to unusual destination: %s:%d", dst, dstPort),
Details: fmt.Sprintf("UID: %d (%s), Local port: %d, Proto: %s", uid, user, localPort, proto),
}, true
}
// CheckOutboundUserConnections looks for non-root user processes making
// outbound connections to IPs that aren't infra or well-known services.
// Catches compromised accounts phoning home.
func CheckOutboundUserConnections(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
if data, err := osFS.ReadFile("/proc/net/tcp"); err == nil {
findings = append(findings, scanProcNetTCP(cfg, data, false)...)
} else {
// Preserve historical behaviour: if /proc/net/tcp is unreadable,
// return nil rather than continuing to tcp6.
return nil
}
if tcp6Data, err := osFS.ReadFile("/proc/net/tcp6"); err == nil {
findings = append(findings, scanProcNetTCP(cfg, tcp6Data, true)...)
}
return findings
}
// scanProcNetTCP parses one /proc/net/tcp[6] dump and returns findings for
// every ESTABLISHED row that EvaluateConnection flags.
//
// A first pass collects local sockets in LISTEN state so that an ESTABLISHED
// row whose local address and port have a listener is recognised as the
// accept side of an inbound connection (e.g. pure-ftpd PASV data channels,
// user-owned daemons listening on high ports) and not an outbound connect().
// Wildcard listeners match every local address for that port.
func scanProcNetTCP(cfg *config.Config, data []byte, ipv6 bool) []alert.Finding {
var findings []alert.Finding
proto := "tcp"
if ipv6 {
proto = "tcp6"
}
lines := strings.Split(string(data), "\n")
listeners := collectListenSockets(lines, ipv6)
directSMTPEnabled := DirectSMTPEgressBackendEnabled(cfg, "legacy")
var mta platform.MTAIdents
if directSMTPEnabled {
// Resolve MTA identities once per scan; legacy poller has no per-PID context,
// so the EvaluateDirectSMTPEgress UID/user gate carries the load.
mta = platform.LocalMTAIdentities(platform.Detect())
}
for _, line := range lines {
fields := strings.Fields(line)
if len(fields) < 8 || fields[0] == "sl" {
continue
}
// State 01 = ESTABLISHED
if fields[3] != "01" {
continue
}
uidStr := fields[7]
uidU64, err := strconv.ParseUint(uidStr, 10, 32)
if err != nil {
continue
}
uidU32 := uint32(uidU64)
var (
localIP net.IP
dstIP net.IP
dstPort int
localPort int
)
if ipv6 {
localIP, localPort = parseHex6Addr(fields[1])
dstIP, dstPort = parseHex6Addr(fields[2])
} else {
localAddr, parsedLocalPort := parseHexAddr(fields[1])
localIP = net.ParseIP(localAddr)
localPort = parsedLocalPort
remoteIP, remotePort := parseHexAddr(fields[2])
dstIP = net.ParseIP(remoteIP)
dstPort = remotePort
}
if localIP == nil || dstIP == nil || localPort <= 0 || dstPort <= 0 {
continue
}
if listeners.has(localIP, localPort) {
continue
}
// Bad-ASN egress is classified for every UID, root included: a
// post-exploit root process exfiltrating to a bad ASN is the
// host-takeover signal. The live BPF tracker drops root events for
// flood control, so this periodic scan is where root egress is seen.
// The non-root detectors below intentionally skip root.
if lookup := CurrentASNLookup(); lookup != nil && cfg.Detection.BadASNOutbound.Enabled {
asn, org := lookup(dstIP.String())
if f, ok := EvaluateBadASNOutbound(cfg, dstIP, asn, org); ok {
AttributeSocketOwner(&f, uidU32)
f.Timestamp = time.Now()
findings = append(findings, f)
}
}
if uidU64 == 0 {
continue
}
user := LookupUser(uidU32)
// #nosec G115 -- ports parsed from /proc/net/tcp[6] are bounded by uint16.
if f, ok := EvaluateConnection(cfg, uidU32, dstIP, uint16(dstPort), uint16(localPort), proto, user); ok {
f.Timestamp = time.Now()
findings = append(findings, f)
}
if directSMTPEnabled {
// #nosec G115 -- ports parsed from /proc/net/tcp[6] are bounded by uint16.
if f, ok := EvaluateDirectSMTPEgress(cfg, DirectSMTPEgressInput{
UID: uidU32,
User: user,
DstIP: dstIP,
DstPort: uint16(dstPort),
MTA: mta,
}); ok {
f.Timestamp = time.Now()
findings = append(findings, f)
}
}
}
return findings
}
type listenSocket struct {
address string
port int
}
type listenSocketSet struct {
wildcardPorts map[int]bool
sockets map[listenSocket]bool
}
func (s listenSocketSet) has(ip net.IP, port int) bool {
if port <= 0 {
return false
}
if s.wildcardPorts[port] {
return true
}
if ip == nil {
return false
}
return s.sockets[listenSocket{address: normalizeListenIP(ip), port: port}]
}
// collectListenSockets scans /proc/net/tcp[6] rows for state 0A (LISTEN) and
// returns the set of local sockets a process is bound to.
func collectListenSockets(lines []string, ipv6 bool) listenSocketSet {
listeners := listenSocketSet{
wildcardPorts: make(map[int]bool),
sockets: make(map[listenSocket]bool),
}
for _, line := range lines {
fields := strings.Fields(line)
if len(fields) < 8 || fields[0] == "sl" {
continue
}
if fields[3] != "0A" {
continue
}
var (
localIP net.IP
localPort int
)
if ipv6 {
localIP, localPort = parseHex6Addr(fields[1])
} else {
localAddr, parsedLocalPort := parseHexAddr(fields[1])
localIP = net.ParseIP(localAddr)
localPort = parsedLocalPort
}
if localIP == nil || localPort <= 0 {
continue
}
if localIP.IsUnspecified() {
listeners.wildcardPorts[localPort] = true
continue
}
listeners.sockets[listenSocket{address: normalizeListenIP(localIP), port: localPort}] = true
}
return listeners
}
func normalizeListenIP(ip net.IP) string {
if v4 := ip.To4(); v4 != nil {
return net.IP(v4).String()
}
return ip.String()
}
// parseHex6Addr parses an IPv6 address:port from /proc/net/tcp6 format.
// IPv6 addresses are 32 hex chars (128 bits) in little-endian 4-byte groups.
func parseHex6Addr(s string) (net.IP, int) {
parts := strings.SplitN(s, ":", 2)
if len(parts) != 2 {
return nil, 0
}
hexIP := parts[0]
hexPort := parts[1]
if len(hexIP) != 32 {
return nil, 0
}
port, ok := parseProcNetHexPort(hexPort)
if !ok {
return nil, 0
}
// Parse as 4 little-endian 32-bit words
ip := make(net.IP, 16)
for i := 0; i < 4; i++ {
word := hexIP[i*8 : (i+1)*8]
b, _ := hex.DecodeString(word)
if len(b) != 4 {
return nil, 0
}
// Reverse bytes within each 32-bit word (little-endian to big-endian)
ip[i*4+0] = b[3]
ip[i*4+1] = b[2]
ip[i*4+2] = b[1]
ip[i*4+3] = b[0]
}
return ip, port
}
// CheckSSHDConfig monitors sshd_config for dangerous changes.
func CheckSSHDConfig(ctx context.Context, _ *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
parsed := parseSSHDConfig()
if !parsed.Present() {
return nil
}
// Hash every file sshd reads, not only the root: a drop-in under an
// Include can flip PermitRootLogin while the root file stays identical.
hash := parsed.Digest()
current := settingsFromSSHDConfig(parsed)
hashKey := "_sshd_config_hash"
passKey := "_sshd_passwordauthentication"
rootKey := "_sshd_permitrootlogin"
prevHash, exists := store.GetRaw(hashKey)
prevPass, _ := store.GetRaw(passKey)
prevRoot, _ := store.GetRaw(rootKey)
if exists && prevHash != hash {
// Only alert when the effective setting changed into a dangerous value.
// This avoids false positives from commented defaults or Match blocks.
if current.PasswordAuthentication == "yes" && prevPass != "yes" {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "sshd_config_change",
Message: "PasswordAuthentication changed to 'yes' in sshd_config",
Details: "This allows password-based SSH login - high risk if passwords are compromised",
})
}
if current.PermitRootLogin == "yes" && prevRoot != "yes" {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "sshd_config_change",
Message: "PermitRootLogin changed to 'yes' in sshd_config",
})
}
// Generic change alert if no specific dangerous setting found
if len(findings) == 0 {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "sshd_config_change",
Message: "sshd_config modified",
})
}
}
store.SetRaw(hashKey, hash)
store.SetRaw(passKey, current.PasswordAuthentication)
store.SetRaw(rootKey, current.PermitRootLogin)
return findings
}
// CheckNulledPlugins scans WordPress plugin directories for signs of
// nulled/pirated plugins: missing licenses, known crack patterns, GPL
// bypass code, and plugins not found on wordpress.org.
func CheckNulledPlugins(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Known crack/null signatures in PHP files
crackSignatures := []string{
"nulled by", "cracked by", "gpl-club", "gpldl.com",
"developer license", "remove license check",
"license_key_bypass", "activation_bypass",
"@remove_license", "null_license",
}
homeDirs, _ := GetScanHomeDirs(ctx)
for _, homeEntry := range homeDirs {
if !homeEntry.IsDir() {
continue
}
pluginsDir := filepath.Join(scanHomeDirPath(homeEntry), "public_html", "wp-content", "plugins")
plugins, err := osFS.ReadDir(pluginsDir)
if err != nil {
continue
}
for _, plugin := range plugins {
if !plugin.IsDir() {
continue
}
pluginDir := filepath.Join(pluginsDir, plugin.Name())
// Check main plugin PHP file for crack signatures
mainFiles, _ := osFS.Glob(filepath.Join(pluginDir, "*.php"))
for _, mainFile := range mainFiles {
// Only read the first 10KB of each file
data := readFileHead(mainFile, 10*1024)
if data == nil {
continue
}
contentLower := strings.ToLower(string(data))
for _, sig := range crackSignatures {
if strings.Contains(contentLower, sig) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "nulled_plugin",
Message: fmt.Sprintf("Possible nulled plugin: %s/%s", homeEntry.Name(), plugin.Name()),
Details: fmt.Sprintf("File: %s\nSignature: %s", mainFile, sig),
})
break
}
}
}
}
}
return findings
}
// readFileHead reads the first N bytes of a file.
func readFileHead(path string, maxBytes int) []byte {
f, err := osFS.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
buf := make([]byte, maxBytes)
n, _ := f.Read(buf)
if n == 0 {
return nil
}
return buf[:n]
}
package checks
import (
"context"
"github.com/pidginhost/csm/internal/alert"
)
// LatestFindingStore is the subset of the state store the sweep needs.
type LatestFindingStore interface {
LatestFindings() []alert.Finding
// Each mutation takes the snapshot verification actually looked at, so a
// scan or realtime alert that refreshed the same key while the re-check was
// in flight is never overwritten by the older verdict.
DismissFindingIfLatest(expected alert.Finding) bool
DemoteLatestFinding(expected alert.Finding, severity alert.Severity) bool
RestoreLatestFindingSeverity(expected alert.Finding) bool
}
// ContentReverifyDismissal records one finding the sweep cleared, for the
// caller to audit-log (the checks package has no logger of its own).
type ContentReverifyDismissal struct {
Check string
Path string
Detail string
// Demoted and Promoted distinguish the outcomes for the audit log: a
// cleared finding is gone, a demoted one is still listed at a lower
// severity, a promoted one had an earlier demotion reversed.
Demoted bool
Promoted bool
}
// isAutomaticallyDemoted reports whether this finding is sitting at Warning
// because a previous sweep lowered it, rather than because it was raised there.
func isAutomaticallyDemoted(f alert.Finding) bool {
return f.Severity == alert.Warning &&
f.DemotedFrom >= alert.High && f.DemotedFrom <= alert.Critical
}
// ShouldRestoreSeverity reports whether an automatic demotion must be reversed.
// A demotion holds only while the replacement keeps satisfying the inert-content
// gate. Restore on a positive match and on every uncertain or newly-active shape
// alike; otherwise a second edit into a detection gap would leave live malware
// at Warning.
func ShouldRestoreSeverity(f alert.Finding, res VerifyResult) bool {
return isAutomaticallyDemoted(f) && !res.Demote
}
// ShouldDemoteSeverity reports whether a verdict retires a remediated but
// unproven finding from the live queue. It is never a clear: an attacker must
// not retire a finding by editing the file. Demoting an already-Warning finding
// would be churn.
//
// The unattended sweep and the operator's Re-check both ask this, so the two
// cannot disagree about what a verdict means.
func ShouldDemoteSeverity(f alert.Finding, res VerifyResult) bool {
return res.Checked && res.Demote && f.Severity > alert.Warning
}
// autoReverifiable reports whether the sweep may re-check and dismiss a finding
// of this type on its own. Membership is deliberately narrow: only families
// whose verifier re-runs the same test that raised the finding and fails closed
// on any uncertainty. Every other registered verifier stays operator-driven
// through the web UI.
func autoReverifiable(check string) bool {
return IsContentReverifiable(check) || isExposedVerifiable(check)
}
// ReverifyStaleFindings re-checks every auto-reverifiable finding in the store
// against current detection logic and dismisses those that are now confirmed
// stale: for content findings a file that is gone, or identical bytes the
// current classifier no longer flags; for web_exposed_* findings an exposure
// that a complete pinned probe no longer confirms. Dispatch goes through the
// verifier registry, so each family keeps its own safety invariant -- a
// still-present file is cleared only when its bytes are unchanged since
// detection, and an exposure only when a complete probe says the server no
// longer serves it. Returns the dismissed findings for the caller to log.
// Read-only except for dismissing confirmed-stale findings.
func ReverifyStaleFindings(store LatestFindingStore) []ContentReverifyDismissal {
dismissed, _ := ReverifyStaleFindingsContext(context.Background(), store)
return dismissed
}
// ReverifyStaleFindingsContext is the cancellable form used by the daemon so a
// large exposure queue cannot delay shutdown for every remaining probe. The
// bool is false after cancellation so the daemon leaves the sweep version
// uncommitted and retries it on the next start.
// ReverifySweepStats describes what a sweep actually did. A sweep that changed
// nothing is otherwise silent, which makes "ran and found nothing"
// indistinguishable from "never ran" and from "could not check a single
// finding" -- the difference an operator needs when findings are not draining.
type ReverifySweepStats struct {
Considered int
Cleared int
Demoted int
Promoted int
Unchecked int
// TopUncheckedReason is the most common reason a finding could not be
// re-checked at all, which is where a silent sweep usually goes wrong.
TopUncheckedReason string
}
func ReverifyStaleFindingsContext(ctx context.Context, store LatestFindingStore) ([]ContentReverifyDismissal, bool) {
out, _, complete := ReverifyStaleFindingsStats(ctx, store)
return out, complete
}
// ReverifyStaleFindingsStats is ReverifyStaleFindingsContext with a summary of
// everything the sweep looked at, including the findings it could not check.
func ReverifyStaleFindingsStats(ctx context.Context, store LatestFindingStore) ([]ContentReverifyDismissal, ReverifySweepStats, bool) {
if ctx == nil {
ctx = context.Background()
}
var dismissed []ContentReverifyDismissal
var stats ReverifySweepStats
if ctx.Err() != nil {
return nil, stats, false
}
uncheckedReasons := map[string]int{}
var exposureVhosts *exposureVhostIndex
for _, f := range store.LatestFindings() {
if ctx.Err() != nil {
return dismissed, stats, false
}
if !autoReverifiable(f.Check) {
continue
}
stats.Considered++
in := VerifyInput{
Check: f.Check, Message: f.Message, Details: f.Details, Path: f.FilePath,
ContentSHA256: f.ContentSHA256, DetectLogic: f.DetectLogic,
Context: ctx,
}
if isExposedVerifiable(f.Check) {
if exposureVhosts == nil {
loaded := loadExposureVhostIndex()
exposureVhosts = &loaded
}
in.exposureVhosts = exposureVhosts
}
res := VerifyFindingInput(in)
if ctx.Err() != nil {
return dismissed, stats, false
}
switch {
case res.Checked && res.Resolved:
if store.DismissFindingIfLatest(f) {
stats.Cleared++
dismissed = append(dismissed, ContentReverifyDismissal{Check: f.Check, Path: f.FilePath, Detail: res.Detail})
}
case ShouldRestoreSeverity(f, res):
if store.RestoreLatestFindingSeverity(f) {
stats.Promoted++
dismissed = append(dismissed, ContentReverifyDismissal{
Check: f.Check, Path: f.FilePath, Detail: res.Detail, Promoted: true})
}
case ShouldDemoteSeverity(f, res):
if store.DemoteLatestFinding(f, alert.Warning) {
stats.Demoted++
dismissed = append(dismissed, ContentReverifyDismissal{
Check: f.Check, Path: f.FilePath, Detail: res.Detail, Demoted: true})
}
case !res.Checked:
// The verifier could not form an opinion at all -- a scanner that
// was unavailable, a file that changed under it. This is the
// bucket that looks identical to a sweep that never ran.
stats.Unchecked++
uncheckedReasons[res.Detail]++
}
}
top, topN := "", 0
for reason, n := range uncheckedReasons {
if n > topN || (n == topN && reason < top) {
top, topN = reason, n
}
}
stats.TopUncheckedReason = top
return dismissed, stats, true
}
package checks
import (
"crypto/sha256"
"fmt"
"io"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/yara"
)
// ContentLogicVersion identifies the current shape of the PHP content-analysis
// heuristic set (analyzePHPContent and its helpers). BUMP IT in the same commit
// as any change to those heuristics, so findings produced by the previous logic
// are re-verified and cleared by the daemon sweep. See
// docs/superpowers/specs/2026-06-20-stale-content-finding-reverification-design.md.
const ContentLogicVersion = 1
// ContentScannerVersion identifies scanner behavior that is not represented by
// the loaded YAML signature version or YARA rule count. Bump it when shared
// content classification changes so the daemon re-checks stale findings.
const ContentScannerVersion = 4
// JSTaintLogicVersion identifies the current semantics of the JavaScript
// keystroke taint analyzer (internal/jstaint). BUMP IT in the same commit as
// any change to its sources, propagation, sinks, resource limits, parser
// version, or content pre-filter so findings produced by the previous logic
// are re-verified under the new one.
const JSTaintLogicVersion = 2
// PHPTaintLogicVersion identifies the current semantics of the PHP remote-
// source taint analyzer (internal/phptaint). BUMP IT in the same commit as any
// change to its parser, pre-filter, propagation, sinks, resource limits, or
// evidence semantics so existing findings are re-verified by the isolated
// worker under the new logic.
const PHPTaintLogicVersion = 4
// contentReverifiableChecks are content findings whose condition can be
// re-evaluated here by re-running the classifier that produced them on the
// file's current bytes. Unlike presenceVerifiableChecks, a still-present file
// may be resolved -- but ONLY when its bytes are byte-for-byte identical to
// detection time (ContentSHA256 match) and the current classifier no longer
// flags them. A file modified since detection is never auto-cleared (it could
// be a partial clean or an evasion edit), preserving the guarantee behind the
// presence-only design.
var contentReverifiableChecks = []string{
"suspicious_php_content",
"obfuscated_php",
"signature_match_realtime", "yara_match_realtime", "yara_match_scheduled",
"js_keylogger_dataflow",
"php_remote_taint",
}
var contentReverifiableSet = func() map[string]struct{} {
m := make(map[string]struct{}, len(contentReverifiableChecks))
for _, c := range contentReverifiableChecks {
m[c] = struct{}{}
}
return m
}()
// IsContentReverifiable reports whether a check type is re-evaluated by
// re-running the content classifier (vs presence-only).
func IsContentReverifiable(check string) bool {
_, ok := contentReverifiableSet[check]
return ok
}
// ContentDetectionVersion returns a token identifying the full content-detection
// logic in effect: the heuristic and scanner versions, the loaded signature-set
// version, the loaded YARA rule count, and both taint analyzer versions. The
// re-verifier always re-runs the real classifier, so this token only gates the
// daemon sweep and enriches audit detail; its precision is not security-critical.
func ContentDetectionVersion() string {
sigVer := 0
if s := signatures.Global(); s != nil {
sigVer = s.Version()
}
yaraRules := 0
if y := yara.Active(); y != nil {
yaraRules = y.RuleCount()
}
return contentDetectionVersionToken(
ContentLogicVersion,
ContentScannerVersion,
sigVer,
yaraRules,
JSTaintLogicVersion,
PHPTaintLogicVersion,
)
}
// reverifySweepLogicVersion identifies the sweep's own semantics, as opposed to
// the detection logic it re-runs: which outcomes it may apply and which paths
// its verifiers can reach. BUMP IT in the same commit as any such change, so an
// upgraded host sweeps at startup instead of carrying the old behaviour until
// its next deep-scan cycle.
const reverifySweepLogicVersion = 2
// FindingReverifyVersion identifies every verifier family the startup sweep
// runs unattended, plus the sweep's own semantics. Including the exposure
// verifier forces one sweep when that family is first deployed or its
// fail-closed semantics change.
func FindingReverifyVersion() string {
return fmt.Sprintf("%s;exposed=%d;reverify=%d",
ContentDetectionVersion(), exposedReverifyLogicVersion, reverifySweepLogicVersion)
}
// contentDetectionVersionToken renders the version components. Pure so a test
// can pass two different component values and prove the tokens differ.
func contentDetectionVersionToken(phpVer, scanVer, sigVer, yaraRules, jsTaintVer, phpTaintVer int) string {
return fmt.Sprintf(
"php=%d;scan=%d;sig=%d;yara=%d;jstaint=%d;phptaint=%d",
phpVer, scanVer, sigVer, yaraRules, jsTaintVer, phpTaintVer,
)
}
// contentFingerprintMaxBytes caps the file size hashed for a finding's
// fingerprint. Larger files get an empty fingerprint, so the re-verifier treats
// them as un-fingerprinted (never auto-cleared while present).
const contentFingerprintMaxBytes = 16 << 20 // 16 MiB
// FileContentSHA256 returns the hex SHA-256 of the whole file, or "" if the
// file cannot be read, is not a regular file, or exceeds the size cap.
func FileContentSHA256(path string) string {
info, err := osFS.Stat(path)
if err != nil || !info.Mode().IsRegular() || info.Size() > contentFingerprintMaxBytes {
return ""
}
f, err := osFS.Open(path)
if err != nil {
return ""
}
defer f.Close()
h := sha256.New()
if _, err := io.Copy(h, f); err != nil {
return ""
}
return fmt.Sprintf("%x", h.Sum(nil))
}
// StampContentFingerprint records the detection-time content fingerprint on a
// content-reverifiable finding so the Re-check / sweep can later distinguish a
// superseded-heuristic false positive from a file edited after detection.
// A producer that analyzed an already-open snapshot may supply its exact hash;
// retain that fingerprint instead of reopening a path that may now name
// different content. No-op for non-content findings or findings without a path.
func StampContentFingerprint(f *alert.Finding) {
if f == nil || f.FilePath == "" || !IsContentReverifiable(f.Check) {
return
}
if f.ContentSHA256 != "" {
if f.DetectLogic == "" {
f.DetectLogic = ContentDetectionVersion()
}
return
}
f.ContentSHA256 = FileContentSHA256(f.FilePath)
f.DetectLogic = ContentDetectionVersion()
}
package checks
import (
"fmt"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// correlationWindow bounds how far apart two findings may be and still count
// as one cross-account event in the persisted active set. A dispatch batch
// supplies its own grouping and can include carried-forward timestamps, so
// the batch callers do not apply this bound.
//
// Without it the aggregate is a latch rather than an alert: replaying a
// 100-day recording of one production host left coordinated_attack raised for
// 76% of the recording, and a two-day recording of a second host for 99%,
// because the first three accounts that ever carried a critical finding never
// left the set. One hour was chosen against those recordings: it holds the
// aggregate raised for 2.4% of the first recording where six hours leaves
// 19.6% and a day leaves 46.1%, and it still spans a full scan sweep, whose
// findings land together. Re-derive it with scripts/correlation-calibrate.
const correlationWindow = time.Hour
// CorrelationResult is the output of one CorrelateFindings call.
type CorrelationResult struct {
// Derived findings: the coordinated_attack result first, then one
// cross_account_malware result per qualifying check, sorted by check.
// Timestamps are left unset for the caller to stamp.
Derived []alert.Finding
// Unattributed counts qualifying input rows per check that carried no
// account identity. It is a snapshot for this call, not a running total,
// using the same window as the aggregates. Batch correlation has no age
// filter; persisted correlation also counts unstamped legacy rows.
Unattributed map[string]int
// CriticalAccounts is the distinct attributed account count after the
// same eligibility and time filters used to derive coordinated_attack.
CriticalAccounts int
}
// Correlator shares production rules with offline replay. Construct one with
// NewCorrelator to supply recording roots without triggering host discovery.
type Correlator struct {
window time.Duration
accountOf func(alert.Finding) string
}
// NewCorrelator uses only the supplied account roots, with no host lookups.
// Window must be nonnegative; zero reproduces unbounded correlation.
func NewCorrelator(window time.Duration, accountRoots []string) Correlator {
roots := append([]string(nil), accountRoots...)
return Correlator{window: window, accountOf: func(f alert.Finding) string {
return extractAccountFromFindingAt(f, func() []string { return roots })
}}
}
var defaultCorrelator = Correlator{window: correlationWindow, accountOf: extractAccountFromFinding}
// CorrelateFindings raises cross-account findings. Eligibility comes from
// the registry classification and identity from extractAccountFromFinding.
// Callers initialize platform.Detect before correlation so account roots
// come from its cache. It does not log or mutate its input.
func CorrelateFindings(findings []alert.Finding) CorrelationResult {
return defaultCorrelator.Correlate(findings, time.Time{})
}
// CorrelateBatchFindings preserves dispatch grouping even when a scan carries
// forward a prior finding whose original timestamp lies outside the window.
func CorrelateBatchFindings(findings []alert.Finding) CorrelationResult {
batch := defaultCorrelator
batch.window = 0
return batch.Correlate(findings, time.Time{})
}
// Correlate derives aggregates at the supplied observation time. A zero time
// uses the newest non-derived input as its reference. Persisted state supplies
// the merge time so even an empty scan can expire old evidence.
func (c Correlator) Correlate(findings []alert.Finding, at time.Time) CorrelationResult {
res := CorrelationResult{Unattributed: make(map[string]int)}
accounts := make(map[string]bool)
malwareByCheck := make(map[string]map[string]bool)
cutoff := c.cutoff(findings, at)
for _, f := range findings {
class := correlationClassOf(f.Check)
if class != CorrelationSecurityEvent && class != CorrelationMalwareArtifact {
continue
}
if observed := observedAt(f); c.window > 0 && !observed.IsZero() && observed.Before(cutoff) {
continue
}
countsForAttack := f.Severity == alert.Critical
countsForMalware := class == CorrelationMalwareArtifact
if !countsForAttack && !countsForMalware {
continue
}
account := c.accountOf(f)
if account == "" {
res.Unattributed[f.Check]++
continue
}
if countsForAttack {
accounts[account] = true
}
if countsForMalware {
if malwareByCheck[f.Check] == nil {
malwareByCheck[f.Check] = make(map[string]bool)
}
malwareByCheck[f.Check][account] = true
}
}
res.CriticalAccounts = len(accounts)
if len(accounts) >= 3 {
names := sortedKeys(accounts)
res.Derived = append(res.Derived, alert.Finding{
Severity: alert.Critical,
Check: "coordinated_attack",
Message: fmt.Sprintf("Possible coordinated attack: %d accounts have critical security events", len(names)),
Details: fmt.Sprintf("Affected accounts: %s", strings.Join(names, ", ")),
})
}
for _, check := range sortedKeys(malwareByCheck) {
names := sortedKeys(malwareByCheck[check])
if len(names) < 2 {
continue
}
res.Derived = append(res.Derived, alert.Finding{
Severity: alert.Critical,
Check: "cross_account_malware",
Message: fmt.Sprintf("Same malware type (%s) found in %d accounts", check, len(names)),
Details: fmt.Sprintf("Accounts: %s", strings.Join(names, ", ")),
})
}
return res
}
// CorrelationInputOf reports how one finding enters cross-account
// correlation: the hosting account it resolves to, empty when none could be
// determined, and whether its check is an eligible input at all. It applies
// the same registry classification and identity rules CorrelateFindings uses,
// so a caller can explain or calibrate a correlation result without
// re-implementing them. Eligibility here is the check's class only; whether a
// given finding then counts also depends on its severity.
func CorrelationInputOf(f alert.Finding) (account string, eligible bool) {
return defaultCorrelator.InputOf(f)
}
// InputOf uses the same classification and identity rules as Correlate.
// Eligibility is the check class only, before severity and window filtering.
func (c Correlator) InputOf(f alert.Finding) (account string, eligible bool) {
class := correlationClassOf(f.Check)
eligible = class == CorrelationSecurityEvent || class == CorrelationMalwareArtifact
return c.accountOf(f), eligible
}
// observedAt is when a finding's condition started: its first observation,
// falling back to its report time. A scan re-emits every finding it still
// sees with a fresh report time, so judging membership by Timestamp let a
// months-old condition re-enter the window on every cycle.
func observedAt(f alert.Finding) time.Time {
if !f.FirstSeen.IsZero() {
return f.FirstSeen
}
return f.Timestamp
}
// A missing timestamp must not silently discard legacy stored evidence.
func (c Correlator) cutoff(findings []alert.Finding, at time.Time) time.Time {
if c.window == 0 {
return time.Time{}
}
if at.IsZero() {
for _, f := range findings {
// Synthesized findings must not feed back into either aggregate
// membership or attribution-health accounting. The reference is
// the newest report, not the newest first observation: it stands
// for "now" when the caller supplied no observation time.
if !IsDerivedCorrelationCheck(f.Check) && f.Timestamp.After(at) {
at = f.Timestamp
}
}
}
if at.IsZero() {
return time.Time{}
}
return at.Add(-c.window)
}
func uniqueStrings(input []string) []string {
seen := make(map[string]bool)
var result []string
for _, s := range input {
if !seen[s] {
seen[s] = true
result = append(result, s)
}
}
return result
}
func sortedKeys[V any](m map[string]V) []string {
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
sort.Strings(out)
return out
}
// extractAccountFromFinding resolves the hosting account a finding belongs
// to: the producer's TenantID verbatim, then the account home that contains
// an absolute FilePath, then the legacy free-text scan of Message and
// Details. The structured sources win so a path mentioned in free text
// cannot re-attribute a finding whose producer knew its owner.
func extractAccountFromFinding(f alert.Finding) string {
return extractAccountFromFindingAt(f, accountHomeRoots)
}
func extractAccountFromFindingAt(f alert.Finding, accountRoots func() []string) string {
if f.TenantID != "" {
return f.TenantID
}
roots := accountRoots()
if filepath.IsAbs(f.FilePath) {
if _, account, ok := accountRootOfAt(f.FilePath, roots); ok {
return account
}
}
for _, s := range []string{f.Message, f.Details} {
if account := accountNameInTextAt(s, roots); account != "" {
return account
}
}
return ""
}
package checks
import (
"fmt"
"sort"
"sync"
)
// CorrelationClass is a check's role in cross-account correlation. Every
// registry entry sets one; the zero value fails the completeness test.
type CorrelationClass uint8
const (
// CorrelationUnclassified is the zero value and never valid.
CorrelationUnclassified CorrelationClass = iota
// CorrelationIgnored is never an input to correlation; a reason is required.
CorrelationIgnored
// CorrelationSecurityEvent: an attributed Critical counts toward coordinated_attack.
CorrelationSecurityEvent
// CorrelationMalwareArtifact is a SecurityEvent that also raises
// cross_account_malware when the same check appears on two accounts.
CorrelationMalwareArtifact
// CorrelationDerived is an output of correlation and never an input.
CorrelationDerived
)
// Ignore reasons. The value is a short token; the sentence is what the
// generated policy table prints and what a reviewer reads.
const (
reasonPosture = "posture"
reasonAttackerSide = "attacker-side"
reasonInformational = "informational"
reasonSelfHealth = "self-health"
reasonResponse = "response"
reasonHostScope = "host-scope"
reasonAccountAggregate = "account-aggregate"
reasonPerformance = "performance"
)
var correlationReasonSentences = map[string]string{
reasonPosture: "static configuration, hardening or hygiene state; a Critical means a misconfiguration, not an attack on the account",
reasonAttackerSide: "attacker activity or attempted access, not evidence of compromise of the named victim",
reasonInformational: "audit trail or inventory event with no compromise claim",
reasonSelfHealth: "CSM's own health, capacity or coverage state",
reasonResponse: "record of an automatic action already taken; feeding it back would double count",
reasonHostScope: "host-wide condition with no account to attribute; a cross-account count cannot use it even when it is a real compromise",
reasonAccountAggregate: "already summarizes several accounts without a single victim identity",
reasonPerformance: "resource usage",
}
// Attribution gaps documented on eligible checks. A gap never changes
// eligibility: an attributed Critical still counts.
const gapEnvelopeSender = "envelope-sender"
var correlationGapSentences = map[string]string{
gapEnvelopeSender: "sender-domain volume aggregate is unattributed when contributing submissions are unverified or belong to different accounts",
}
// validateCorrelationPolicy returns the first policy violation in entries,
// naming the offending check.
func validateCorrelationPolicy(entries []CheckInfo) error {
for _, c := range entries {
switch c.Correlation {
case CorrelationIgnored:
if _, ok := correlationReasonSentences[c.CorrelationReason]; !ok {
return fmt.Errorf("check %q is ignored without a known reason (%q)", c.Name, c.CorrelationReason)
}
if c.CorrelationGap != "" {
return fmt.Errorf("check %q is ignored but carries gap %q", c.Name, c.CorrelationGap)
}
case CorrelationSecurityEvent, CorrelationMalwareArtifact:
if c.CorrelationReason != "" {
return fmt.Errorf("check %q is eligible but carries reason %q", c.Name, c.CorrelationReason)
}
if c.CorrelationGap != "" {
if _, ok := correlationGapSentences[c.CorrelationGap]; !ok {
return fmt.Errorf("check %q carries unknown gap %q", c.Name, c.CorrelationGap)
}
}
case CorrelationDerived:
if c.CorrelationReason != "" || c.CorrelationGap != "" {
return fmt.Errorf("check %q is derived but carries a reason or gap", c.Name)
}
default:
return fmt.Errorf("check %q has no correlation classification (class %d)", c.Name, c.Correlation)
}
}
return nil
}
// correlationIndex is built once from the registry. It never reads the
// filesystem, the store or the network.
type correlationIndex struct {
classes map[string]CorrelationClass
reasons map[string]string
derived []string
}
var (
correlationOnce sync.Once
correlationTable *correlationIndex
)
func loadCorrelationIndex() *correlationIndex {
correlationOnce.Do(func() {
idx := &correlationIndex{
classes: make(map[string]CorrelationClass, len(checkRegistry)),
reasons: make(map[string]string, len(checkRegistry)),
}
for _, c := range checkRegistry {
idx.classes[c.Name] = c.Correlation
idx.reasons[c.Name] = c.CorrelationReason
if c.Correlation == CorrelationDerived {
idx.derived = append(idx.derived, c.Name)
}
}
sort.Strings(idx.derived)
correlationTable = idx
})
return correlationTable
}
func correlationClassOf(name string) CorrelationClass {
return loadCorrelationIndex().classes[name]
}
// correlationReasonOf returns the ignore reason a check is registered with,
// or "" for an eligible or unknown check.
func correlationReasonOf(name string) string {
return loadCorrelationIndex().reasons[name]
}
// securityEventEligible reports whether an attributed Critical finding of
// this check counts toward coordinated_attack.
func securityEventEligible(name string) bool {
switch correlationClassOf(name) {
case CorrelationSecurityEvent, CorrelationMalwareArtifact:
return true
}
return false
}
// DerivedCorrelationChecks returns the names correlation itself emits,
// sorted. The caller owns the slice.
func DerivedCorrelationChecks() []string {
return append([]string(nil), loadCorrelationIndex().derived...)
}
// IsDerivedCorrelationCheck reports whether name is an output of correlation.
func IsDerivedCorrelationCheck(name string) bool {
return correlationClassOf(name) == CorrelationDerived
}
package checks
import (
"sync"
"time"
csmlog "github.com/pidginhost/csm/internal/log"
)
// AttributionReport is the operator-facing view of correlation attribution.
// Current is what the latest-state active set looks like right now;
// Cumulative is the history since the daemon started. The two answer
// different questions: a producer that lost attribution for weeks and one
// that missed once look identical in a log line, and neither is visible in
// a health endpoint without this.
type AttributionReport struct {
// Current holds, per check, the qualifying rows in the latest-state
// active set that carry no hosting owner and fall inside the correlation
// window at its most recent merge. Unstamped legacy rows also count.
// A later merge clears rows that gain an owner or age out of the window.
Current map[string]int
// Cumulative sums every unattributed row reported since start, by
// check, across active-set merges and per-batch derivations.
Cumulative map[string]int
// ActiveSetUpdates counts active-set merges since start.
ActiveSetUpdates int
// Since is when the first active set was recorded; zero before then.
Since time.Time
}
// unattributedReporter logs once per check name per process when
// cross-account correlation could not attribute a qualifying finding to a
// hosting account, and keeps the counts behind AttributionHealth. Per-call
// counts are never summed into the active-set snapshot; only the cumulative
// history adds them up.
type unattributedReporter struct {
mu sync.Mutex
seen map[string]struct{}
current map[string]int
cumulative map[string]int
activeSetUpdates int
since time.Time
warn func(msg string, args ...any)
}
func newUnattributedReporter(warn func(string, ...any)) *unattributedReporter {
return &unattributedReporter{
seen: make(map[string]struct{}),
current: make(map[string]int),
cumulative: make(map[string]int),
warn: warn,
}
}
// eligibleCounts keeps the entries that describe a real attribution loss:
// positive counts for checks that correlation would otherwise count.
func eligibleCounts(counts map[string]int) map[string]int {
out := make(map[string]int, len(counts))
for check, n := range counts {
if n > 0 && securityEventEligible(check) {
out[check] = n
}
}
return out
}
type unattributedWarning struct {
check string
n int
}
// record publishes all counters together and returns pending warnings so
// callers can release their merge locks before invoking the logger.
func (r *unattributedReporter) record(counts map[string]int, activeSet bool) []unattributedWarning {
filtered := eligibleCounts(counts)
var fresh []unattributedWarning
r.mu.Lock()
defer r.mu.Unlock()
if activeSet {
r.current = filtered
r.activeSetUpdates++
if r.since.IsZero() {
r.since = time.Now()
}
}
for check, n := range filtered {
r.cumulative[check] += n
if _, dup := r.seen[check]; !dup {
r.seen[check] = struct{}{}
fresh = append(fresh, unattributedWarning{check, n})
}
}
return fresh
}
func (r *unattributedReporter) warnCounts(fresh []unattributedWarning) {
for _, f := range fresh {
r.warn("cross-account correlation could not attribute findings to an account", "check", f.check, "rows", f.n)
}
}
// Report records a per-batch derivation: the rows count toward the
// cumulative history and warn once per check, but the active-set snapshot
// is untouched because a batch is not the persisted state.
func (r *unattributedReporter) Report(counts map[string]int) {
r.warnCounts(r.record(counts, false))
}
// RecordActiveSet records the latest-state merge: it replaces the
// active-set snapshot, so a merge whose rows all carry owners clears it,
// and adds to the history like a batch report.
func (r *unattributedReporter) RecordActiveSet(counts map[string]int) {
r.warnCounts(r.record(counts, true))
}
// Health returns a copy of the current state.
func (r *unattributedReporter) Health() AttributionReport {
r.mu.Lock()
defer r.mu.Unlock()
return AttributionReport{
Current: copyCounts(r.current),
Cumulative: copyCounts(r.cumulative),
ActiveSetUpdates: r.activeSetUpdates,
Since: r.since,
}
}
func copyCounts(m map[string]int) map[string]int {
out := make(map[string]int, len(m))
for k, v := range m {
out[k] = v
}
return out
}
var defaultUnattributedReporter = newUnattributedReporter(csmlog.Warn)
// ReportUnattributedCorrelation records unattributed rows from a per-batch
// derivation through the process-wide reporter. Callers invoke it after
// releasing any state store lock; it never re-enters the store.
func ReportUnattributedCorrelation(counts map[string]int) {
defaultUnattributedReporter.Report(counts)
}
// RecordUnattributedActiveSet records the unattributed rows of the
// latest-state active set after a merge. Same locking contract as
// ReportUnattributedCorrelation.
func RecordUnattributedActiveSet(counts map[string]int) {
defaultUnattributedReporter.RecordActiveSet(counts)
}
// AttributionHealth reports the process-wide attribution state for the
// health snapshot and doctor.
func AttributionHealth() AttributionReport {
return defaultUnattributedReporter.Health()
}
// ResetAttributionHealthForTest replaces the process-wide reporter with a
// fresh one. Test-only.
func ResetAttributionHealthForTest() {
defaultUnattributedReporter = newUnattributedReporter(csmlog.Warn)
}
package checks
import (
"context"
"fmt"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
const (
sessionLogPath = "/usr/local/cpanel/logs/session_log"
sessionLogTailLines = 1000
defaultMultiIPThreshold = 3
defaultMultiIPWindowMin = 60
)
// CheckCpanelLogins parses the cPanel session log for suspicious login activity:
// - cPanel (cpaneld) logins from non-infra IPs
// - Same account logged in from multiple distinct IPs (credential compromise indicator)
// - Password change purge events (attacker or auto-response password resets)
func CheckCpanelLogins(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
lines := tailFile(sessionLogPath, sessionLogTailLines)
if len(lines) == 0 {
return nil
}
// Track logins per account for multi-IP correlation
accountIPs := make(map[string]map[string]bool)
var passwordChanges []string
// Determine cutoff - only alert on events within the scan window
cutoff := time.Now().Add(-time.Duration(multiIPWindowMin(cfg)) * time.Minute)
for _, line := range lines {
// Session log format:
// [2026-03-25 07:19:47 +0200] info [cpaneld] 203.0.113.133 NEW user:token address=IP,...
// [2026-03-25 07:38:18 +0200] info [security] internal PURGE user:token password_change
// Parse timestamp
ts := parseSessionTimestamp(line)
if ts.IsZero() || ts.Before(cutoff) {
continue
}
// Detect cPanel logins from non-infra IPs
// Skip API/portal sessions (create_user_session) - only alert on direct form login
if strings.Contains(line, "[cpaneld]") && strings.Contains(line, " NEW ") {
if cfg.Suppressions.SuppressCpanelLogin {
// Still track IPs for multi-IP correlation even when suppressed
} else if strings.Contains(line, "method=create_user_session") ||
strings.Contains(line, "method=create_session") ||
strings.Contains(line, "create_user_session") {
continue
}
ip, account := parseCpanelLogin(line)
if ip == "" || account == "" {
continue
}
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" || ip == "internal" {
continue
}
// Always track for multi-IP correlation (even when login alerts suppressed)
if accountIPs[account] == nil {
accountIPs[account] = make(map[string]bool)
}
accountIPs[account][ip] = true
// WARNING severity - logins are audit trail, not paging-level.
// Multi-IP correlation stays CRITICAL via its own check below.
if !cfg.Suppressions.SuppressCpanelLogin {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "cpanel_login",
Message: fmt.Sprintf("cPanel direct login from non-infra IP: %s (account: %s)", ip, account),
Details: truncateString(line, 300),
})
}
}
// Detect password change purge events
if strings.Contains(line, "PURGE") && strings.Contains(line, "password_change") {
account := parsePurgeAccount(line)
if account != "" {
passwordChanges = append(passwordChanges, account)
}
}
}
// Multi-IP correlation: same account from 3+ distinct non-infra IPs
threshold := multiIPThreshold(cfg)
for account, ips := range accountIPs {
if len(ips) >= threshold {
ipList := make([]string, 0, len(ips))
for ip := range ips {
ipList = append(ipList, ip)
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "cpanel_multi_ip_login",
Message: fmt.Sprintf("Account '%s' logged in from %d distinct IPs (credential compromise likely)", account, len(ips)),
Details: fmt.Sprintf("IPs: %s\nThreshold: %d IPs within %d minutes", strings.Join(ipList, ", "), threshold, multiIPWindowMin(cfg)),
TenantID: HostingAccountForUser(account),
})
}
}
// Password change events - deduplicate by account
seen := make(map[string]bool)
for _, account := range passwordChanges {
if seen[account] {
continue
}
seen[account] = true
// Check if triggered by security module (Imunify auto-response) vs user action
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "cpanel_password_purge",
Message: fmt.Sprintf("cPanel sessions purged via password change for account: %s", account),
Details: "This may indicate an automated security response or attacker-initiated password change",
})
}
return findings
}
// CheckCpanelFileManager parses the cPanel access log for file management
// operations from non-infra IPs (file uploads, edits via cPanel File Manager).
func CheckCpanelFileManager(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
lines := tailFile("/usr/local/cpanel/logs/access_log", 300)
// Only match actual write actions - not read-only calls like get_homedir.
// Skip 401/403 responses - the server rejected the request, no write occurred.
// Match against request URI only, not the full line (referer can contain "upload").
filemanWriteActions := []string{
"fileman/save_file_content",
"fileman/upload_files",
"fileman/save_file",
"fileman/paste",
"fileman/rename",
"fileman/delete",
}
for _, line := range lines {
// Only check cPanel (port 2083) entries
if !strings.Contains(line, "2083") {
continue
}
// Skip rejected requests - no write occurred
if strings.Contains(line, "\" 401 ") || strings.Contains(line, "\" 403 ") {
continue
}
fields := strings.Fields(line)
if len(fields) < 1 {
continue
}
ip := fields[0]
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
continue
}
// Extract request URI (between first pair of quotes) to avoid
// matching "upload" in referer URLs like upload-ajax.html
requestURI := extractRequestURIChecks(line)
// Common log format: the authenticated cPanel user is the third
// field, "-" when the request was unauthenticated.
owner := ""
if len(fields) >= 3 && fields[2] != "-" {
owner = HostingAccountForUser(fields[2])
}
for _, action := range filemanWriteActions {
if strings.Contains(strings.ToLower(requestURI), strings.ToLower(action)) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "cpanel_file_upload",
Message: fmt.Sprintf("cPanel File Manager write operation from non-infra IP: %s", ip),
Details: truncateString(line, 300),
TenantID: owner,
})
break
}
}
}
return findings
}
// parseSessionTimestamp extracts the timestamp from a session log line.
// Format: [2026-03-25 07:19:47 +0200]
func parseSessionTimestamp(line string) time.Time {
start := strings.Index(line, "[")
end := strings.Index(line, "]")
if start < 0 || end < 0 || end <= start+1 {
return time.Time{}
}
tsStr := line[start+1 : end]
// Try common cPanel session log formats
for _, layout := range []string{
"2006-01-02 15:04:05 -0700",
"2006-01-02 15:04:05 +0000",
} {
if t, err := time.Parse(layout, tsStr); err == nil {
return t
}
}
return time.Time{}
}
// parseCpanelLogin extracts IP and account from a session NEW line.
// Format: [timestamp] info [cpaneld] 203.0.113.133 NEW user:token address=IP,...
func parseCpanelLogin(line string) (ip, account string) {
// Find IP after [cpaneld]
idx := strings.Index(line, "[cpaneld]")
if idx < 0 {
return "", ""
}
rest := strings.TrimSpace(line[idx+len("[cpaneld]"):])
fields := strings.Fields(rest)
if len(fields) < 3 {
return "", ""
}
ip = fields[0]
// Find account from "NEW user:token" or "NEW user:token address=..."
for i, f := range fields {
if f == "NEW" && i+1 < len(fields) {
userToken := fields[i+1]
parts := strings.SplitN(userToken, ":", 2)
if len(parts) >= 1 {
account = parts[0]
}
break
}
}
return ip, account
}
// parsePurgeAccount extracts the account name from a PURGE password_change line.
// Format: [timestamp] info [security] internal PURGE user:token password_change
func parsePurgeAccount(line string) string {
idx := strings.Index(line, "PURGE")
if idx < 0 {
return ""
}
rest := strings.TrimSpace(line[idx+len("PURGE"):])
fields := strings.Fields(rest)
if len(fields) < 1 {
return ""
}
parts := strings.SplitN(fields[0], ":", 2)
if len(parts) >= 1 {
return parts[0]
}
return ""
}
func multiIPThreshold(cfg *config.Config) int {
if cfg.Thresholds.MultiIPLoginThreshold > 0 {
return cfg.Thresholds.MultiIPLoginThreshold
}
return defaultMultiIPThreshold
}
func multiIPWindowMin(cfg *config.Config) int {
if cfg.Thresholds.MultiIPLoginWindowMin > 0 {
return cfg.Thresholds.MultiIPLoginWindowMin
}
return defaultMultiIPWindowMin
}
// extractRequestURIChecks extracts the request line from an access log entry.
// Format: ... "METHOD /path HTTP/1.1" ... → returns "METHOD /path HTTP/1.1"
func extractRequestURIChecks(line string) string {
start := strings.Index(line, "\"")
if start < 0 {
return ""
}
end := strings.Index(line[start+1:], "\"")
if end < 0 {
return ""
}
return line[start+1 : start+1+end]
}
package checks
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// credentialReuseMinAccounts is the threshold for emitting the
// password-reuse finding: the same WordPress admin password hash present
// on this many or more distinct accounts. Two is the common shared-hosting
// pattern (an agency reusing one admin password across client sites) where
// a single credential disclosure compromises every site at once.
const credentialReuseMinAccounts = 2
// CheckCredentialReuse flags WordPress administrator accounts that share
// an identical password hash across two or more distinct hosting accounts.
// Password hashes are salted on modern WordPress installs, so this is an
// exact at-rest hash reuse signal, not a weak-password detector.
//
// Privacy: the raw password hash is never stored, logged, or emitted. Only
// a truncated one-way fingerprint is used to group identical hashes, and
// findings report the affected accounts and a count -- not the hash.
func CheckCredentialReuse(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
wpConfigs := credentialReuseWPConfigs(ctx)
// fingerprint -> set of distinct accounts carrying that admin hash.
byFingerprint := map[string]map[string]struct{}{}
for _, wpConfig := range wpConfigs {
if ctx.Err() != nil {
return nil
}
account := wpConfigUser(filepath.Dir(wpConfig))
if account == "" {
continue
}
creds, complete := parseWPConfigChecked(wpConfig)
if !complete {
markCheckIncomplete(ctx, "credential_reuse")
continue
}
if creds.dbName == "" {
markCheckIncomplete(ctx, "credential_reuse")
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
markCheckIncomplete(ctx, "credential_reuse")
continue
}
creds.tablePrefix = prefix
fingerprints, err := adminPasswordFingerprintsForSite(creds, prefix)
if err != nil {
markCheckIncomplete(ctx, "credential_reuse")
continue
}
for _, fp := range fingerprints {
if fp == "" {
continue
}
if byFingerprint[fp] == nil {
byFingerprint[fp] = map[string]struct{}{}
}
byFingerprint[fp][account] = struct{}{}
}
}
return buildCredentialReuseFindings(byFingerprint, credentialReuseMinAccounts)
}
// adminPasswordFingerprintsForSite returns fingerprints for the admin
// user_pass hashes currently stored on the WordPress site. Uses root MySQL
// because wp-config passwords drift on cPanel hosts (same rationale as
// adminEmailsForSite).
func adminPasswordFingerprintsForSite(creds wpDBCreds, prefix string) ([]string, error) {
query := fmt.Sprintf(
"SELECT DISTINCT u.user_pass FROM `%susers` u "+
"JOIN `%susermeta` um ON u.ID = um.user_id "+
"WHERE um.meta_key = '%scapabilities' AND um.meta_value LIKE '%%administrator%%'",
prefix, prefix, prefix,
)
rows, err := runMySQLQueryRootWithError(creds.dbName, query)
if err != nil {
return nil, err
}
var out []string
for _, row := range rows {
fp := credentialHashFingerprint(strings.TrimSpace(row))
if fp != "" {
out = append(out, fp)
}
}
return out, nil
}
// credentialHashFingerprint maps a raw password hash to a short,
// non-reversible grouping key. Two identical hashes map to the same
// fingerprint without returning the raw hash. Empty input yields "".
func credentialHashFingerprint(rawHash string) string {
if rawHash == "" {
return ""
}
sum := sha256.Sum256([]byte(rawHash))
return "fp:" + hex.EncodeToString(sum[:])[:16]
}
// buildCredentialReuseFindings emits one Warning per fingerprint shared by
// at least minAccounts distinct accounts. The finding never includes the
// hash or fingerprint-as-secret -- only the affected account list and a
// count, so an operator can rotate the shared credential.
func buildCredentialReuseFindings(byFingerprint map[string]map[string]struct{}, minAccounts int) []alert.Finding {
if minAccounts < 2 {
minAccounts = 2
}
type credentialReuseGroup struct {
accounts []string
}
groups := make([]credentialReuseGroup, 0, len(byFingerprint))
for _, accountSet := range byFingerprint {
if len(accountSet) < minAccounts {
continue
}
accounts := make([]string, 0, len(accountSet))
for a := range accountSet {
accounts = append(accounts, a)
}
sort.Strings(accounts)
groups = append(groups, credentialReuseGroup{accounts: accounts})
}
sort.Slice(groups, func(i, j int) bool {
return strings.Join(groups[i].accounts, "\x00") < strings.Join(groups[j].accounts, "\x00")
})
var out []alert.Finding
for _, group := range groups {
accounts := group.accounts
out = append(out, alert.Finding{
Severity: alert.Warning,
Check: "credential_reuse",
Message: fmt.Sprintf("Identical WordPress admin password hash reused across %d accounts: %s",
len(accounts), strings.Join(accounts, ", ")),
Details: fmt.Sprintf("Accounts sharing one admin password hash: %s\n"+
"Rotate the shared credential: a single disclosure compromises every listed site.",
strings.Join(accounts, ", ")),
Timestamp: time.Now(),
})
}
return out
}
// credentialReuseWPConfigs lists the WordPress installs this check fingerprints.
// A hash reused between a primary site and an addon install is still reuse.
func credentialReuseWPConfigs(ctx context.Context) []string {
installs := wpInstalls(ctx, "credential_reuse")
out := make([]string, 0, len(installs))
for _, in := range installs {
out = append(out, in.ConfigPath)
}
return out
}
package checks
import "github.com/pidginhost/csm/internal/platform"
// cronSpoolDir returns the directory holding per-user crontabs for this
// platform (cronie: /var/spool/cron; Debian cron: /var/spool/cron/crontabs).
// Var so tests can pin a layout without touching the host.
var cronSpoolDir = func() string { return platform.Detect().CronSpoolDir() }
// webServerUsers returns the accounts the web server runs as on this
// platform; a seam so tests can stand in for platform detection.
var webServerUsers = func() []string { return platform.Detect().WebServerUsers() }
package checks
import (
"fmt"
"os"
"path/filepath"
"github.com/pidginhost/csm/internal/quarantinefs"
)
// fixCrontabAllowedRoots limits suspicious_crontab remediation to the cron
// spool. Declared as a var so tests can redirect it under t.TempDir()
// without touching the real /var/spool/cron.
var fixCrontabAllowedRoots = []string{"/var/spool/cron"}
// fixSuspiciousCrontab copies a user crontab matching known-bad persistence
// markers into quarantine, writes a restore-ready metadata sidecar, and then
// truncates the live file to empty. Truncation (not deletion) keeps cron(8)
// from re-reading stale content and preserves the caller's ability to
// inspect file perms while the malware is gone.
func fixSuspiciousCrontab(path string) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
path, info, err := resolveExistingFixPath(path, fixCrontabAllowedRoots)
if err != nil {
return RemediationResult{Error: err.Error()}
}
data, err := osFS.ReadFile(path)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot read: %v", err)}
}
user := filepath.Base(path)
qPath := newQuarantinePath(quarantineDir, "crontab_"+user)
meta := quarantineMetadata(path, info, "suspicious_crontab remediation")
if err := storeQuarantineBackup(qPath, data, meta, 0600); err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot create durable crontab backup: %v", err)}
}
// Truncate live crontab. 0600 is the mode cron(8) enforces for user
// spool files; any other mode makes cron skip the file with a warning.
// #nosec G306 -- cron(8) rejects world-readable user crontabs, so 0600
// is the only safe mode for /var/spool/cron/<user>.
if err := os.WriteFile(path, []byte{}, 0600); err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot truncate crontab: %v", err)}
}
if err := quarantinefs.SyncFilePath(path); err != nil {
return RemediationResult{Error: fmt.Sprintf("crontab truncated but not synced; recovery copy retained at %s: %v", qPath, err)}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("quarantined crontab %s -> %s and truncated", path, qPath),
Description: fmt.Sprintf("Truncated %d-byte crontab; copy saved to quarantine", len(data)),
}
}
package checks
import (
"context"
"encoding/base64"
"fmt"
"path/filepath"
"regexp"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/state"
)
// crontabBase64Truncated counts base64 candidates that hit
// crontabBase64BlobMaxBytes. Sustained growth means an attacker is
// padding outer blobs to push the real payload past the decode window;
// raise the cap or split the scanner.
var (
crontabBase64Truncated *metrics.Counter
crontabBase64TruncatedOnce sync.Once
)
func observeCrontabBase64Truncation() {
crontabBase64TruncatedOnce.Do(func() {
crontabBase64Truncated = metrics.NewCounter(
"csm_checks_crontab_base64_truncated_total",
"Crontab base64 candidates that exceeded the per-blob decode cap before decoded-content pattern matching ran. Sustained growth means the scanner inspected only the leading decoded window of large encoded cron content.",
)
metrics.MustRegister("csm_checks_crontab_base64_truncated_total", crontabBase64Truncated)
})
crontabBase64Truncated.Inc()
}
func CheckCrontabs(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
// Rank account crontabs by mtime desc so recently-touched users
// process first when the check timeout cuts iteration short. Keep
// root outside the account cap; it is a system baseline, not an
// account-scoped path.
if ctx.Err() != nil {
return findings
}
crontabs, _ := osFS.Glob(filepath.Join(cronSpoolDir(), "*"))
var rootCrontabs []string
accountCrontabs := make([]string, 0, len(crontabs))
for _, path := range crontabs {
if filepath.Base(path) == "root" {
rootCrontabs = append(rootCrontabs, path)
continue
}
accountCrontabs = append(accountCrontabs, path)
}
rankedRootCrontabs := rankPathsByMtimeDesc(ctx, rootCrontabs, 0)
if ctx.Err() != nil {
return findings
}
for _, path := range rankedRootCrontabs {
if ctx.Err() != nil {
return findings
}
hash, err := hashFileContent(path)
if err != nil {
continue
}
if ctx.Err() != nil {
return findings
}
key := "_crontab_root_hash"
prev, exists := store.GetRaw(key)
if exists && prev != hash {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "crontab_change",
Message: "Root crontab modified",
Details: "Review with: crontab -l",
})
}
store.SetRaw(key, hash)
}
rankedCrontabs := rankPathsByMtimeDesc(ctx, accountCrontabs, accountScanMaxFiles(ctx, cfg))
if ctx.Err() != nil {
return findings
}
for _, path := range rankedCrontabs {
if ctx.Err() != nil {
return findings
}
user := filepath.Base(path)
data, err := osFS.ReadFile(path)
if err != nil {
continue
}
if ctx.Err() != nil {
return findings
}
// The spool basename is a hosting owner only when an account home
// of that name exists; service users keep the finding unattributed.
owner := ""
if accountHomeExists(user) {
owner = user
}
content := string(data)
for _, pattern := range MatchCrontabPatternsDeep(content, cfg) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "suspicious_crontab",
Message: fmt.Sprintf("Suspicious pattern in crontab for user %s: %s", user, pattern),
Details: fmt.Sprintf("File: %s\nContent:\n%s", path, truncate(content, 500)),
FilePath: path,
TenantID: owner,
})
}
}
// Check /etc/cron.d for new files. This is a system directory, so
// account_scan_max_files must not hide older cron.d baselines.
if ctx.Err() != nil {
return findings
}
cronDFiles, globErr := osFS.Glob("/etc/cron.d/*")
rankedCronDFiles := rankPathsByMtimeDesc(ctx, cronDFiles, 0)
if ctx.Err() != nil {
return findings
}
// Before the first complete pass every file is install backlog and
// only gets baselined; afterwards a file with no stored hash appeared
// since the last run and is reported, not silently absorbed.
_, cronDBaselined := store.GetRaw(cronDBaselineKey)
cronDBaselineComplete := globErr == nil
for _, path := range rankedCronDFiles {
if ctx.Err() != nil {
return findings
}
data, err := osFS.ReadFile(path)
if err != nil {
cronDBaselineComplete = false
continue
}
if ctx.Err() != nil {
return findings
}
hash := hashBytes(data)
key := fmt.Sprintf("_crond:%s", filepath.Base(path))
prev, exists := store.GetRaw(key)
// The same file reaches the realtime write detector, which rescores
// a vendor-driven change instead of paging. Scoring the scheduled
// diff on its own left an upgrade or a panel maintenance run
// reporting High through whichever detector saw it first.
switch {
case exists && prev != hash:
findings = append(findings, rescoreSensitive(alert.Finding{
Severity: alert.High,
Check: "crond_change",
Message: fmt.Sprintf("Cron.d file modified: %s", path),
}, "cron", data, 0, time.Now()))
case !exists && cronDBaselined:
findings = append(findings, rescoreSensitive(alert.Finding{
Severity: alert.High,
Check: "crond_change",
Message: fmt.Sprintf("Cron.d file added: %s", path),
Details: fmt.Sprintf("File: %s\nContent: %s", path, alert.RedactCommandLine(truncate(strings.TrimSpace(string(data)), cronDExcerptLen))),
}, "cron", data, 0, time.Now()))
}
store.SetRaw(key, hash)
}
if cronDBaselineComplete && ctx.Err() == nil {
store.SetRaw(cronDBaselineKey, "1")
}
return findings
}
// cronDBaselineKey marks that /etc/cron.d was fully enumerated once; from
// then on an unknown file is new rather than backlog.
const cronDBaselineKey = "_crond:_baseline_complete"
// cronDExcerptLen bounds the content quoted for a new cron.d file.
const cronDExcerptLen = 300
func truncate(s string, maxLen int) string {
if len(s) <= maxLen {
return s
}
return s[:maxLen] + "..."
}
// crontabSuspiciousPatterns is the shared allowlist of case-insensitive
// substrings that mark a crontab line as likely malicious. Single source of
// truth for CheckCrontabs (system scan) and makeAccountCrontabCheck
// (per-account scan) so the two cannot drift apart.
var crontabSuspiciousPatterns = []string{
"defunct-kernel",
"SEED PRNG",
"base64_decode",
"base64 -d|bash",
"base64 -d | bash",
"base64 --decode|bash",
"eval(",
"/dev/tcp/",
"gsocket",
"gs-netcat",
"reverse",
"bash -i",
"/bin/sh -i",
"nc -e",
"ncat -e",
"python -c",
"perl -e",
}
// matchCrontabPatterns returns patterns from crontabSuspiciousPatterns that
// appear as case-insensitive substrings of content, preserving list order.
func matchCrontabPatterns(content string) []string {
lower := strings.ToLower(content)
var matched []string
for _, pattern := range crontabSuspiciousPatterns {
if strings.Contains(lower, strings.ToLower(pattern)) {
matched = append(matched, pattern)
}
}
return matched
}
// crontabBase64BlobMaxBytesDefault is the built-in fallback cap for a
// single base64 candidate before decoding. 16384 encoded bytes
// (~12 KiB decoded) comfortably fits any realistic gsocket /
// `base64 -d|bash` payload while bounding work on adversarial input.
// Operator override: cfg.Thresholds.CrontabBase64BlobMaxBytes.
//
// Must stay a multiple of 4 -- standard base64 needs aligned input or
// DecodeString errors and the candidate is silently skipped. The
// validator rejects non-aligned operator values.
const crontabBase64BlobMaxBytesDefault = 16384
// effectiveCrontabBase64BlobMaxBytes returns the operator-configured cap
// or the built-in default when unset. The validator enforces multiple-of-4
// alignment so this returns a safe value without further checks.
func effectiveCrontabBase64BlobMaxBytes(cfg *config.Config) int {
if cfg == nil || cfg.Thresholds.CrontabBase64BlobMaxBytes <= 0 {
return crontabBase64BlobMaxBytesDefault
}
return cfg.Thresholds.CrontabBase64BlobMaxBytes
}
// crontabBase64BlobMaxCount caps the number of base64 candidates examined
// per crontab. A realistic gsocket cron entry has one outer blob; we
// allow enough headroom for a handful without doing unbounded work.
const crontabBase64BlobMaxCount = 16
// crontabBase64BlobRE matches contiguous standard-alphabet base64 of
// length >= 40 (with optional padding). The 40-char floor avoids matching
// short config IDs and noise like Wordfence cookie names.
var crontabBase64BlobRE = regexp.MustCompile(`[A-Za-z0-9+/]{40,}={0,2}`)
// MatchCrontabPatternsDeep is matchCrontabPatterns plus a single base64
// decode pass: it pulls out base64 candidates from content and re-runs
// pattern matching on the decoded bytes. Catches attackers who wrap the
// `base64 -d|bash` pipe chain in an outer base64 layer so the literal
// markers never appear in the cron file as written. Single decode depth;
// no recursion. cfg nil uses the built-in defaults; pass the live
// operator config to honour `thresholds.crontab_base64_blob_max_bytes`.
func MatchCrontabPatternsDeep(content string, cfg *config.Config) []string {
maxBytes := effectiveCrontabBase64BlobMaxBytes(cfg)
matched := matchCrontabPatterns(content)
seen := make(map[string]bool, len(matched))
for _, m := range matched {
seen[m] = true
}
candidates := crontabBase64BlobRE.FindAllString(content, crontabBase64BlobMaxCount)
for _, blob := range candidates {
if len(blob) > maxBytes {
observeCrontabBase64Truncation()
blob = blob[:maxBytes]
}
decoded, err := base64.StdEncoding.DecodeString(blob)
if err != nil {
continue
}
for _, m := range matchCrontabPatterns(string(decoded)) {
if !seen[m] {
matched = append(matched, m)
seen[m] = true
}
}
}
return matched
}
package checks
import (
"context"
"crypto/sha256"
"fmt"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// adminEmailRetention bounds how far back an admin observation stays
// relevant for overlap detection. A contractor email seen six months
// ago on one account and never since is not actionable signal -- the
// access likely lapsed. The window is a deliberate balance against the
// alternative of evicting on every scan, which would lose overlaps when
// scans run asynchronously across customer accounts.
const adminEmailRetention = 90 * 24 * time.Hour
// adminEmailDefaultMinAccounts is the default threshold for emitting
// the cross-account overlap finding. Matches the most common
// compromise pattern on shared hosting: a contractor administering two
// or more customer cPanels.
const adminEmailDefaultMinAccounts = 2
// CheckAdminEmailOverlap records every WordPress administrator email
// encountered during an account scan into a server-wide bbolt bucket,
// then emits a Warning finding for each email whose owner list now
// spans the configured minimum number of distinct accounts. The
// detection surface is shared-hosting credential leakage: a single
// compromised contractor account is one credential disclosure away
// from administrator access on every site they touch.
//
// The check is silent when the bbolt store is unavailable (early
// daemon startup, test harness without state injection) -- it can't
// observe overlap without persistence between scans, and falling
// silent is better than a misleading partial result.
func CheckAdminEmailOverlap(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
db := store.Global()
if db == nil {
return nil
}
now := time.Now()
wpConfigs := adminOverlapWPConfigs(ctx)
for _, wpConfig := range wpConfigs {
if ctx.Err() != nil {
return nil
}
account := wpConfigUser(filepath.Dir(wpConfig))
creds, complete := parseWPConfigChecked(wpConfig)
if !complete {
markCheckIncomplete(ctx, "admin_overlap")
continue
}
if creds.dbName == "" {
markCheckIncomplete(ctx, "admin_overlap")
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
markCheckIncomplete(ctx, "admin_overlap")
continue
}
creds.tablePrefix = prefix
emails, err := adminEmailsForSite(creds, prefix)
if err != nil {
markCheckIncomplete(ctx, "admin_overlap")
continue
}
for _, email := range emails {
if err := db.RecordAdminEmail(email, account, creds.dbName, now); err != nil {
markCheckIncomplete(ctx, "admin_overlap")
}
}
}
min := adminEmailDefaultMinAccounts
if cfg != nil && cfg.Detection.AdminOverlapMinAccounts > 0 {
min = cfg.Detection.AdminOverlapMinAccounts
}
overlaps, err := db.OverlappingAdminEmails(min, adminEmailRetention)
if err != nil {
markCheckIncomplete(ctx, "admin_overlap")
return nil
}
if len(overlaps) == 0 {
return nil
}
overlaps = filterTrustedAdminOverlaps(overlaps, cfg)
return buildAdminOverlapFindings(overlaps)
}
// adminEmailsForSite returns the lowercase admin emails currently
// configured on the WordPress site. Uses the existing root-MySQL
// helper so it works on cPanel hosts where wp-config passwords drift.
func adminEmailsForSite(creds wpDBCreds, prefix string) ([]string, error) {
query := fmt.Sprintf(
"SELECT DISTINCT LOWER(u.user_email) FROM `%susers` u "+
"JOIN `%susermeta` um ON u.ID = um.user_id "+
"WHERE um.meta_key = '%scapabilities' AND um.meta_value LIKE '%%administrator%%'",
prefix, prefix, prefix,
)
rows, err := runMySQLQueryRootWithError(creds.dbName, query)
if err != nil {
return nil, err
}
var out []string
for _, row := range rows {
row = strings.TrimSpace(row)
if row != "" {
out = append(out, row)
}
}
return out, nil
}
// buildAdminOverlapFindings collapses each overlap entry into a single Warning
// finding. The sorted, de-duplicated account set feeds both operator-facing
// text and the explicit identity, so input order and multiple schemas owned by
// one account cannot change the finding key.
func buildAdminOverlapFindings(overlaps map[string][]store.AdminEmailEntry) []alert.Finding {
emails := make([]string, 0, len(overlaps))
for email := range overlaps {
emails = append(emails, email)
}
sort.Strings(emails)
out := make([]alert.Finding, 0, len(emails))
for _, email := range emails {
owners := overlaps[email]
accountSet := make(map[string]struct{}, len(owners))
for _, o := range owners {
accountSet[o.Account] = struct{}{}
}
accounts := make([]string, 0, len(accountSet))
for a := range accountSet {
accounts = append(accounts, a)
}
sort.Strings(accounts)
details := strings.Builder{}
fmt.Fprintf(&details, "Email: %s\nAccounts: %s\n", email, strings.Join(accounts, ", "))
for _, o := range owners {
fmt.Fprintf(&details, "- %s (schema %s, last seen %s)\n", o.Account, o.Schema, o.LastSeen.Format(time.RFC3339))
}
out = append(out, alert.Finding{
Severity: alert.Warning,
Check: "admin_cross_account_overlap",
// The overlap itself is the identity: this email on this set of
// accounts. Details carry each account's last-seen time, and
// Finding.Key() hashes Details, so without an explicit key every
// scan minted a new finding for an unchanged overlap.
DedupKey: adminOverlapDedupKey(email, accounts),
Message: fmt.Sprintf("Admin email %s appears on %d accounts: %s", email, len(accounts), strings.Join(accounts, ", ")),
Details: details.String(),
Timestamp: time.Now(),
})
}
return out
}
// adminOverlapDedupKey identifies one overlap by its substance: the shared
// email and the set of accounts carrying it. An email that spreads to another
// account is a new situation and gets its own key; the same overlap re-observed
// on the next scan keeps this one. accounts is already sorted and de-duplicated
// by the caller.
func adminOverlapDedupKey(email string, accounts []string) string {
identity := strings.Join(append([]string{email}, accounts...), "\x00")
digest := sha256.Sum256([]byte(identity))
return fmt.Sprintf("admin-overlap:%x", digest[:12])
}
func filterTrustedAdminOverlaps(overlaps map[string][]store.AdminEmailEntry, cfg *config.Config) map[string][]store.AdminEmailEntry {
if cfg == nil || (len(cfg.Detection.AdminOverlapTrustedEmails) == 0 && len(cfg.Detection.AdminOverlapTrustedDomains) == 0) {
return overlaps
}
out := make(map[string][]store.AdminEmailEntry, len(overlaps))
for email, owners := range overlaps {
if trustedAdminOverlapEmail(email, cfg) {
continue
}
out[email] = owners
}
return out
}
func trustedAdminOverlapEmail(email string, cfg *config.Config) bool {
email = strings.ToLower(strings.TrimSpace(email))
if email == "" {
return false
}
for _, trusted := range cfg.Detection.AdminOverlapTrustedEmails {
if email == strings.ToLower(strings.TrimSpace(trusted)) {
return true
}
}
domain := adminEmailDomain(email)
if domain == "" {
return false
}
for _, trusted := range cfg.Detection.AdminOverlapTrustedDomains {
if domain == strings.ToLower(strings.TrimSpace(trusted)) {
return true
}
}
return false
}
func adminEmailDomain(email string) string {
at := strings.LastIndexByte(email, '@')
if at < 0 || at == len(email)-1 {
return ""
}
return email[at+1:]
}
// adminOverlapWPConfigs lists the WordPress installs this check compares.
// Overlap between a primary site and a subdomain install is the shape this
// check exists to catch, so both must be discovered.
func adminOverlapWPConfigs(ctx context.Context) []string {
installs := wpInstalls(ctx, "admin_overlap")
out := make([]string, 0, len(installs))
for _, in := range installs {
out = append(out, in.ConfigPath)
}
return out
}
package checks
import (
"context"
"fmt"
"net"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/mysqlclient"
)
// AutoRespondDBMalware processes database injection findings and takes
// automated action: blocks attacker IPs extracted from WordPress session
// tokens, revokes compromised user sessions, and cleans confirmed
// malicious content from wp_options or stored database objects.
//
// Only acts on high-confidence findings:
// - db_options_injection with confirmed malicious external script URLs
// - db_siteurl_hijack (siteurl/home pointing to malicious content)
// - db_malicious_trigger/event/procedure/function with structured metadata
//
// Does NOT act on:
// - db_spam_injection (spam posts — needs manual review)
// - db_post_injection (script in posts — too many FPs from page builders)
// - db_options_injection without confirmed malicious URLs
func AutoRespondDBMalware(cfg *config.Config, findings []alert.Finding) []alert.Finding {
return AutoRespondDBMalwareWithPolicy(cfg, findings, nil)
}
// AutoRespondDBMalwareWithPolicy keeps session IP enforcement independent of
// permission to edit a database or revoke sessions. A nil policy permits both;
// callers with suppressions supply a per-finding remediation decision.
func AutoRespondDBMalwareWithPolicy(cfg *config.Config, findings []alert.Finding, canRemediate func(alert.Finding) bool) []alert.Finding {
if !cfg.AutoResponse.Enabled || !cfg.AutoResponse.CleanDatabase {
return nil
}
var actions []alert.Finding
for _, f := range findings {
remediate := canRemediate == nil || canRemediate(f)
switch f.Check {
case "db_options_injection":
acts := handleMaliciousOption(cfg, f, remediate)
actions = append(actions, acts...)
case "db_siteurl_hijack":
acts := handleSiteurlHijack(cfg, f, remediate)
actions = append(actions, acts...)
case "db_malicious_trigger", "db_malicious_event",
"db_malicious_procedure", "db_malicious_function":
if !remediate {
continue
}
acts := handleMaliciousDBObject(f)
actions = append(actions, acts...)
}
}
return actions
}
// dbDropObjectFn is the seam through which handleMaliciousDBObject performs
// the backup-then-DROP. Overridden in tests so the routing and action
// emission can be exercised without a live MySQL server.
var dbDropObjectFn = DBDropObject
// handleMaliciousDBObject auto-cleans a confirmed malicious stored database
// object (trigger/event/procedure/function). Detection always fires; the
// DROP only runs when the operator has enabled auto_response.clean_database
// (checked by the caller). The object kind comes from the check name
// (db_malicious_<kind>); account/schema/name come from the finding details.
// DBDropObject records a SHOW CREATE backup in bbolt before dropping, so the
// action is reversible.
func handleMaliciousDBObject(f alert.Finding) []alert.Finding {
kind := maliciousDBObjectKind(f.Check)
if kind == "" {
return nil
}
account, schema, detailKind, name := parseDBObjectFindingDetails(f.Details)
if account == "" || schema == "" || detailKind == "" || name == "" {
return nil
}
if detailKind != kind {
return nil
}
res := dbDropObjectFn(account, schema, kind, name, false)
if !res.Success {
return []alert.Finding{{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-DB-CLEAN failed to drop %s %s.%s: %s", kind, schema, name, res.Message),
Timestamp: time.Now(),
}}
}
return []alert.Finding{{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-DB-CLEAN: Dropped malicious %s %s.%s (backup retained for restore)", kind, schema, name),
Timestamp: time.Now(),
}}
}
func maliciousDBObjectKind(check string) string {
const prefix = "db_malicious_"
if !strings.HasPrefix(check, prefix) {
return ""
}
kind := strings.TrimPrefix(check, prefix)
if !IsDBObjectKind(kind) {
return ""
}
return kind
}
// parseDBObjectFindingDetails extracts the structured header fields a
// db_malicious_<kind> finding carries in its Details block. The SQL body is
// attacker-controlled and may contain lines that look like metadata, so parsing
// stops at Body and keeps the first value for each header key.
func parseDBObjectFindingDetails(details string) (account, schema, kind, name string) {
for _, line := range strings.Split(details, "\n") {
line = strings.TrimSpace(line)
if strings.HasPrefix(line, "Body:") {
break
}
switch {
case strings.HasPrefix(line, "Account: "):
if account == "" {
account = strings.TrimSpace(strings.TrimPrefix(line, "Account: "))
}
case strings.HasPrefix(line, "Schema: "):
if schema == "" {
schema = strings.TrimSpace(strings.TrimPrefix(line, "Schema: "))
}
case strings.HasPrefix(line, "Kind: "):
if kind == "" {
kind = strings.TrimSpace(strings.TrimPrefix(line, "Kind: "))
}
case strings.HasPrefix(line, "Name: "):
if name == "" {
name = strings.TrimSpace(strings.TrimPrefix(line, "Name: "))
}
}
}
return
}
// handleMaliciousOption checks if a db_options_injection finding contains
// a confirmed malicious external script URL, and if so:
// 1. Extracts attacker IPs from WP sessions and emits block findings
// 2. Revokes sessions for users with non-infra, non-private IPs only
// 3. Backs up and cleans the malicious content from the option
func handleMaliciousOption(cfg *config.Config, f alert.Finding, remediate bool) []alert.Finding {
var actions []alert.Finding
dbName, optionName := parseDBFindingDetails(f.Details)
if dbName == "" || optionName == "" {
return nil
}
// Validate option name — must be a plausible WP option name.
if !isValidOptionName(optionName) {
return nil
}
// Never act on CSM backup options — they preserve original malicious
// content for recovery. Acting on them causes cascading backup loops.
if strings.HasPrefix(optionName, "csm_backup_") {
return nil
}
creds := findCredsForDB(dbName)
if creds.dbName == "" {
return nil
}
prefix := creds.tablePrefix
if prefix == "" {
prefix = "wp_"
}
// Re-read the FULL option value from the database — the finding's
// Details field only has a truncated 200-char preview.
fullValue := readOptionValue(creds, prefix, optionName)
if fullValue == "" {
return nil
}
// Only act on options with confirmed malicious external script URLs.
maliciousURL := extractMaliciousScriptURL(fullValue)
if maliciousURL == "" {
return nil
}
// 1. Extract and block attacker IPs from active WP sessions through the
// real auto-block path (dry-run, rate limits, and allowlists all apply).
suspiciousIPs := extractSuspiciousSessionIPs(creds, prefix, cfg.InfraIPs)
actions = append(actions, blockSessionAttackerIPs(cfg, suspiciousIPs,
fmt.Sprintf("active WP session on compromised site, DB: %s", dbName), alert.FindingID(f))...)
if !remediate {
return actions
}
// 2. Revoke sessions only for users with suspicious IPs.
// This preserves the site admin's session if they're on an infra IP.
revoked := revokeCompromisedSessions(creds, prefix, cfg.InfraIPs)
if revoked > 0 {
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-DB-CLEAN: Revoked %d compromised WordPress sessions (DB: %s)", revoked, dbName),
Timestamp: time.Now(),
})
}
// 3. Back up the original value, then clean the malicious content.
cleaned := backupAndCleanOption(creds, prefix, optionName, fullValue, maliciousURL)
if cleaned {
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-DB-CLEAN: Removed malicious script from wp_options '%s' (DB: %s, URL: %s)", optionName, dbName, maliciousURL),
Timestamp: time.Now(),
})
}
return actions
}
// handleSiteurlHijack handles siteurl/home hijacking by revoking sessions
// and blocking attacker IPs. Does NOT modify siteurl/home values.
func handleSiteurlHijack(cfg *config.Config, f alert.Finding, remediate bool) []alert.Finding {
var actions []alert.Finding
dbName, _ := parseDBFindingDetails(f.Details)
if dbName == "" {
return nil
}
creds := findCredsForDB(dbName)
if creds.dbName == "" {
return nil
}
prefix := creds.tablePrefix
if prefix == "" {
prefix = "wp_"
}
suspiciousIPs := extractSuspiciousSessionIPs(creds, prefix, cfg.InfraIPs)
actions = append(actions, blockSessionAttackerIPs(cfg, suspiciousIPs,
fmt.Sprintf("active session on hijacked site, DB: %s", dbName), alert.FindingID(f))...)
if !remediate {
return actions
}
revoked := revokeCompromisedSessions(creds, prefix, cfg.InfraIPs)
if revoked > 0 {
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-DB-CLEAN: Revoked %d sessions on hijacked site (DB: %s)", revoked, dbName),
Timestamp: time.Now(),
})
}
return actions
}
// blockSessionAttackerIPs routes attacker IPs recovered from active WordPress
// sessions through the standard auto-block path so each one lands as a real
// firewall block subject to dry-run, rate limiting, allowlists, and the
// expiring threat record. It returns the genuine AUTO-BLOCK / dry-run findings
// AutoBlockIPs emits.
//
// The synthetic findings carry the local_threat_score check -- an existing
// always-block signal meaning "this IP is a confirmed local threat" -- plus a
// structured SourceIP, so AutoBlockIPs blocks exactly that address. Emitting a
// fabricated "auto_block: AUTO-BLOCK <ip>" finding here instead -- as the code
// once did -- never blocked anything, yet alert.FilterBlockedAlerts trusted it
// as proof-of-block and suppressed the IP's reputation alert, so the address was
// neither blocked nor surfaced.
func blockSessionAttackerIPs(cfg *config.Config, ips []string, siteContext, findingID string) []alert.Finding {
if len(ips) == 0 {
return nil
}
findings := make([]alert.Finding, 0, len(ips))
for _, ip := range ips {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "local_threat_score",
Message: fmt.Sprintf("attacker session IP %s (%s)", ip, siteContext),
SourceIP: ip,
Timestamp: time.Now(),
})
}
return autoBlockIPs(cfg, findings, findingID)
}
// --- URL analysis ---
// scriptSrcRe matches <script src="..."> or <script src=...> patterns.
// Accepts https://, http://, and protocol-relative // URLs, since real
// attackers use all three forms to load external payloads.
var scriptSrcRe = regexp.MustCompile(`(?i)<script[^>]+src\s*=\s*["']?((?:https?:)?//[^"'\s>]+)`)
// escapedSlashRe matches a forward slash behind one or more backslashes, the
// form every JSON-encoded and re-serialised option value takes.
var escapedSlashRe = regexp.MustCompile(`\\+/`)
// unescapeStoredSlashes restores the slashes of a URL stored inside a JSON or
// serialised option value. WordPress writes json_encode output straight into
// wp_options, so a stored payload reads "https:\/\/host\/payload.js"; without
// this the script matcher sees no "//" after the scheme and extracts nothing.
func unescapeStoredSlashes(value string) string {
if !strings.Contains(value, `\/`) {
return value
}
return escapedSlashRe.ReplaceAllString(value, "/")
}
// knownSafeDomains are legitimate services that embed scripts in wp_options.
var knownSafeDomains = []string{
"googletagmanager.com",
"google-analytics.com",
"googleapis.com",
"gstatic.com",
"google.com",
"facebook.net",
"facebook.com",
"fbcdn.net",
"connect.facebook.net",
"chimpstatic.com",
"mailchimp.com",
"hotjar.com",
"clarity.ms",
"cloudflare.com",
"cdnjs.cloudflare.com",
"jquery.com",
"jsdelivr.net",
"unpkg.com",
"wp.com",
"wordpress.com",
"gravatar.com",
"tawk.to",
"crisp.chat",
"tidio.co",
"intercom.io",
"zendesk.com",
"hubspot.com",
"hubspot.net",
"hs-scripts.com",
"hs-analytics.net",
"hsforms.com",
"mautic.net",
"pinterest.com",
"twitter.com",
"linkedin.com",
"addthis.com",
"sharethis.com",
"recaptcha.net",
"stripe.com",
"paypal.com",
"brevo-mail.com",
}
// extractMaliciousScriptURL finds a <script src="..."> URL in the content
// that is classified as an attacker script by isAttackerScriptURL.
//
// The classification uses an attack-indicator model (see url_reputation.go):
// a URL flags only when it shows attacker-characteristic markers (raw IP
// host, abused TLD, plaintext HTTP, known-bad exfil host, or no valid
// TLD). The previous allowlist-only model produced HIGH-severity findings
// for legitimate third-party widgets (OneTrust, Issuu, regional video
// embeds, regional tax-form widgets) whose domains were not on the
// allowlist; the attack-indicator model eliminates those false positives
// while still catching the injection patterns attackers actually use.
//
// knownSafeDomains is retained as a fast-path optimisation and operator-
// pre-approved list — see isAttackerScriptURL for the composition order.
func extractMaliciousScriptURL(content string) string {
// WordPress stores json_encode output verbatim, so a stored loader reads
// "https:\/\/host\/payload.js" and the src grammar never matches it.
matches := scriptSrcRe.FindAllStringSubmatch(unescapeStoredSlashes(content), -1)
for _, match := range matches {
if len(match) < 2 {
continue
}
url := match[1]
if isAttackerScriptURL(url) {
return url
}
}
return ""
}
// isSafeScriptDomain checks if a script URL is from a known safe domain.
// Handles https://host, http://host, //host (protocol-relative), and
// host-with-port forms.
func isSafeScriptDomain(url string) bool {
urlLower := strings.ToLower(url)
urlLower = strings.TrimPrefix(urlLower, "https://")
urlLower = strings.TrimPrefix(urlLower, "http://")
urlLower = strings.TrimPrefix(urlLower, "//")
host := urlLower
if idx := strings.IndexByte(host, '/'); idx >= 0 {
host = host[:idx]
}
if idx := strings.IndexByte(host, ':'); idx >= 0 {
host = host[:idx]
}
for _, safe := range knownSafeDomains {
if host == safe || strings.HasSuffix(host, "."+safe) {
return true
}
}
return false
}
// --- Validation ---
// validOptionNameRe allows alphanumeric, underscores, hyphens, colons, and dots.
// Rejects anything that could be SQL injection.
var validOptionNameRe = regexp.MustCompile(`^[a-zA-Z0-9_\-:.]+$`)
// isValidOptionName validates that an option name is safe for SQL interpolation.
func isValidOptionName(name string) bool {
return len(name) > 0 && len(name) <= 191 && validOptionNameRe.MatchString(name)
}
// --- DB helpers ---
// parseDBFindingDetails extracts the database name and option name from
// a finding's Details field.
func parseDBFindingDetails(details string) (dbName, optionName string) {
for _, line := range strings.Split(details, "\n") {
line = strings.TrimSpace(line)
if strings.HasPrefix(line, "Database: ") {
dbName = strings.TrimPrefix(line, "Database: ")
}
if strings.HasPrefix(line, "Option: ") {
optionName = strings.TrimPrefix(line, "Option: ")
}
}
return
}
// findCredsForDB finds wp-config.php credentials that match a database name.
// Skips wp-configs whose $table_prefix fails the safety check -- those
// values come straight from a cPanel-user-writable file and end up in
// root-credentialled SQL via handleMaliciousOption / handleSiteurlHijack.
func findCredsForDB(dbName string) wpDBCreds {
wpConfigs := wpInstallConfigPaths(wpInstalls(context.Background(), "db_content"))
for _, path := range wpConfigs {
creds := parseWPConfig(path)
if creds.dbName != dbName {
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
continue
}
creds.tablePrefix = prefix
return creds
}
return wpDBCreds{}
}
// readOptionValue reads the full value of a wp_option from the database.
//
// mysqlclient returns mysql batch-mode output, where control bytes are rendered
// as escape sequences (a real newline becomes the two bytes "\n"). The value is
// unescaped back to its true bytes before returning so callers that write it
// back (the backup copy and the cleaned value) persist the original content
// rather than the escaped text, keeping PHP-serialized length prefixes valid.
//
// This path intentionally bypasses runMySQLQuery: that legacy scan helper trims
// each returned row before handing it to callers, but wp_options values may
// contain significant leading/trailing whitespace that must survive byte-for-
// byte when CSM writes the backup and cleaned option value.
func readOptionValue(creds wpDBCreds, prefix, optionName string) string {
if !isValidOptionName(optionName) {
return ""
}
query := fmt.Sprintf(
"SELECT option_value FROM %soptions WHERE option_name='%s' LIMIT 1",
prefix, escapeSQLString(optionName))
lines, err := mysqlclient.PerAccountQuery(context.Background(), mysqlclient.Creds{
User: creds.dbUser,
Password: creds.dbPass,
Host: creds.dbHost,
DBName: creds.dbName,
}, query)
if err != nil {
return ""
}
if len(lines) == 0 {
return ""
}
return mysqlclient.BatchUnescape(lines[0])
}
// extractSuspiciousSessionIPs reads WP session tokens and returns IPs that
// are NOT infra IPs, not private, and not loopback.
func extractSuspiciousSessionIPs(creds wpDBCreds, prefix string, infraIPs []string) []string {
query := fmt.Sprintf(
"SELECT meta_value FROM %susermeta WHERE meta_key='session_tokens' AND meta_value != ''",
prefix)
lines := runMySQLQuery(creds, query)
seen := make(map[string]bool)
var ips []string
ipRe := regexp.MustCompile(`"ip";s:\d+:"([^"]+)"`)
for _, line := range lines {
matches := ipRe.FindAllStringSubmatch(line, -1)
for _, m := range matches {
if len(m) < 2 {
continue
}
ip := m[1]
parsed := net.ParseIP(ip)
if parsed == nil || parsed.IsLoopback() || parsed.IsPrivate() {
continue
}
if isInfraIP(ip, infraIPs) {
continue
}
if !seen[ip] {
seen[ip] = true
ips = append(ips, ip)
}
}
}
return ips
}
// revokeCompromisedSessions clears session_tokens only for WP users whose
// sessions contain non-infra, non-private IPs. Returns count of users revoked.
func revokeCompromisedSessions(creds wpDBCreds, prefix string, infraIPs []string) int {
// Get user IDs with active sessions.
query := fmt.Sprintf(
"SELECT user_id, meta_value FROM %susermeta WHERE meta_key='session_tokens' AND meta_value != ''",
prefix)
lines := runMySQLQuery(creds, query)
ipRe := regexp.MustCompile(`"ip";s:\d+:"([^"]+)"`)
revoked := 0
for _, line := range lines {
parts := strings.SplitN(line, "\t", 2)
if len(parts) != 2 {
continue
}
userID := strings.TrimSpace(parts[0])
sessionData := parts[1]
// Check if this user has any suspicious (non-infra, non-private) IPs.
hasSuspicious := false
matches := ipRe.FindAllStringSubmatch(sessionData, -1)
for _, m := range matches {
if len(m) < 2 {
continue
}
ip := m[1]
parsed := net.ParseIP(ip)
if parsed == nil || parsed.IsLoopback() || parsed.IsPrivate() {
continue
}
if isInfraIP(ip, infraIPs) {
continue
}
hasSuspicious = true
break
}
if hasSuspicious {
revokeQuery := fmt.Sprintf(
"UPDATE %susermeta SET meta_value='' WHERE user_id=%s AND meta_key='session_tokens'",
prefix, escapeSQLString(userID))
runMySQLQuery(creds, revokeQuery)
revoked++
}
}
return revoked
}
// backupAndCleanOption saves the original value to a backup option, then
// removes malicious script injections from the option value.
func backupAndCleanOption(creds wpDBCreds, prefix, optionName, originalValue, maliciousURL string) bool {
if maliciousURL == "" {
return false
}
cleaned := removeMaliciousScripts(originalValue)
// Never claim a clean (nor write back) unless the confirmed attacker
// script is actually gone. This keeps removal locked to detection: if
// a script form is flagged but the remover cannot strip it, report the
// finding but never persist a value that still carries a live payload.
// Plain text references to the same URL are inert option data and must
// not block a valid script cleanup.
if optionInjectionRemains(optionName, cleaned) {
return false
}
if cleaned == originalValue {
return false
}
// Save original value as a backup option (csm_backup_<name>_<timestamp>).
backupName := fmt.Sprintf("csm_backup_%s_%d", optionName, time.Now().Unix())
if len(backupName) > 191 {
backupName = backupName[:191]
}
backupQuery := fmt.Sprintf(
"INSERT INTO %soptions (option_name, option_value, autoload) VALUES ('%s', '%s', 'no')",
prefix, escapeSQLString(backupName), escapeSQLString(originalValue))
runMySQLQuery(creds, backupQuery)
// Write the cleaned value.
updateQuery := fmt.Sprintf(
"UPDATE %soptions SET option_value='%s' WHERE option_name='%s'",
prefix, escapeSQLString(cleaned), escapeSQLString(optionName))
runMySQLQuery(creds, updateQuery)
return true
}
// A notice sink makes executable markup malicious regardless of URL
// reputation. Removing one known attacker URL must not permit a partial write
// while an ordinary HTTPS loader or inline script survives in the same row.
func optionInjectionRemains(option, value string) bool {
if extractMaliciousScriptURL(value) != "" {
return true
}
_, sink := pluginNoticeSinkOptions[strings.ToLower(strings.TrimSpace(option))]
return sink && executableMarkupRe.MatchString(unescapeStoredSlashes(value))
}
// --- Script removal ---
// maliciousScriptRe matches the style-break injection pattern:
// </style><script src=...></script><style>
var maliciousScriptRe = regexp.MustCompile(
`(?i)</style>\s*<script[^>]*src\s*=\s*[^>]+>\s*</script>\s*<style>`)
// simpleScriptRe matches standalone <script src="..."></script> tags.
// The src grammar mirrors scriptSrcRe (https://, http://, and
// protocol-relative //) so removal stays paired with detection: a URL
// form the detector flags must be one the remover can strip.
var simpleScriptRe = regexp.MustCompile(
`(?i)<script[^>]*src\s*=\s*["']?(?:https?:)?//[^"'\s>]+["']?[^>]*>\s*</script>`)
// removeMaliciousScripts strips malicious <script> injections from content,
// preserving scripts that are not classified as attacker scripts.
//
// Uses the same isAttackerScriptURL predicate as extractMaliciousScriptURL
// so detection and removal stay semantically paired. If the detector
// would not flag a given URL as malicious, the remover must not strip
// it — otherwise an operator running DBCleanOption on an option that
// contains a real injection alongside a legitimate third-party embed
// (OneTrust, Issuu, regional widget) would silently lose the legitimate
// embed along with the attacker's script.
func removeMaliciousScripts(content string) string {
// First pass: remove style-break patterns only when the embedded URL has
// attacker indicators. The wrapper is suspicious, but the cleaner must
// not remove a legitimate embed just because it appears next to a real
// attacker script in the same option.
content = maliciousScriptRe.ReplaceAllStringFunc(content, func(match string) string {
urls := scriptSrcRe.FindStringSubmatch(match)
if len(urls) >= 2 && isAttackerScriptURL(urls[1]) {
return ""
}
return match
})
// Second pass: remove standalone script tags only when the URL
// shows attacker indicators.
content = simpleScriptRe.ReplaceAllStringFunc(content, func(match string) string {
urls := scriptSrcRe.FindStringSubmatch(match)
if len(urls) >= 2 && isAttackerScriptURL(urls[1]) {
return ""
}
return match
})
return strings.TrimSpace(content)
}
// escapeSQLString escapes special characters for MySQL string interpolation.
func escapeSQLString(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, `'`, `\'`)
s = strings.ReplaceAll(s, "\x00", `\0`)
s = strings.ReplaceAll(s, "\n", `\n`)
s = strings.ReplaceAll(s, "\r", `\r`)
s = strings.ReplaceAll(s, "\x1a", `\Z`)
return s
}
package checks
import (
"context"
"fmt"
"path/filepath"
"regexp"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/mysqlclient"
)
// spamPostIDRe guards the ID list interpolated into root DELETE statements.
var spamPostIDRe = regexp.MustCompile(`^[0-9]+$`)
// DBCleanResult describes the outcome of a database cleanup operation.
type DBCleanResult struct {
Account string
Database string
Action string // "clean-option", "revoke-user", "delete-spam"
Success bool
Message string
Details []string // individual actions taken
BackupNames []string // names of backup options created
}
// DBCleanOption removes malicious script injections from a wp_option value.
// Creates a backup option before modifying. Returns a result describing
// what was done. If preview is true, reports what would be done without
// modifying the database.
func DBCleanOption(account, optionName string, preview bool) DBCleanResult {
result := DBCleanResult{
Account: account,
Action: "clean-option",
}
if !isValidOptionName(optionName) {
result.Message = fmt.Sprintf("Invalid option name: %q", optionName)
return result
}
creds, prefix := findCredsForAccount(account)
if creds.dbName == "" {
result.Message = fmt.Sprintf("No WordPress database found for account %q", account)
return result
}
result.Database = creds.dbName
// Read current value.
value := readOptionValue(creds, prefix, optionName)
if value == "" {
result.Message = fmt.Sprintf("Option %q not found or empty in %s", optionName, creds.dbName)
return result
}
// Check for malicious content.
maliciousURL := extractMaliciousScriptURL(value)
if maliciousURL == "" {
result.Message = fmt.Sprintf("No malicious external script found in %q", optionName)
result.Details = append(result.Details, "Option exists but contains no confirmed malicious URLs")
return result
}
cleaned := removeMaliciousScripts(value)
if cleaned == value {
result.Message = "Content unchanged after cleaning"
return result
}
if optionInjectionRemains(optionName, cleaned) {
result.Message = "Failed to remove all malicious scripts"
return result
}
result.Details = append(result.Details, fmt.Sprintf("Malicious URL: %s", maliciousURL))
result.Details = append(result.Details, fmt.Sprintf("Original length: %d, Cleaned length: %d", len(value), len(cleaned)))
if preview {
result.Message = fmt.Sprintf("PREVIEW: Would clean malicious script from %q", optionName)
result.Success = true
return result
}
// Backup and clean.
if backupAndCleanOption(creds, prefix, optionName, value, maliciousURL) {
backupName := fmt.Sprintf("csm_backup_%s_%d", optionName, time.Now().Unix())
if len(backupName) > 191 {
backupName = backupName[:191]
}
result.BackupNames = append(result.BackupNames, backupName)
result.Details = append(result.Details, fmt.Sprintf("Backup saved as: %s", backupName))
result.Message = fmt.Sprintf("Cleaned malicious script from %q", optionName)
result.Success = true
} else {
result.Message = "Failed to clean option"
}
return result
}
// DBRevokeUser revokes WordPress sessions for a specific user and optionally
// demotes them to subscriber role. If preview is true, reports what would be
// done without modifying the database.
func DBRevokeUser(account string, userID int, demote, preview bool) DBCleanResult {
result := DBCleanResult{
Account: account,
Action: "revoke-user",
}
creds, prefix := findCredsForAccount(account)
if creds.dbName == "" {
result.Message = fmt.Sprintf("No WordPress database found for account %q", account)
return result
}
result.Database = creds.dbName
// Verify user exists.
query := fmt.Sprintf(
"SELECT user_login, user_email FROM %susers WHERE ID=%d LIMIT 1",
prefix, userID)
lines := runMySQLQueryRoot(creds.dbName, query)
if len(lines) == 0 {
result.Message = fmt.Sprintf("User ID %d not found in %s", userID, creds.dbName)
return result
}
parts := strings.SplitN(lines[0], "\t", 2)
login := parts[0]
email := ""
if len(parts) > 1 {
email = parts[1]
}
result.Details = append(result.Details, fmt.Sprintf("User: %s (email: %s)", login, email))
// Check current sessions.
sessQuery := fmt.Sprintf(
"SELECT LEFT(meta_value, 200) FROM %susermeta WHERE user_id=%d AND meta_key='session_tokens'",
prefix, userID)
sessLines := runMySQLQueryRoot(creds.dbName, sessQuery)
sessionCount := 0
if len(sessLines) > 0 && sessLines[0] != "" {
sessionCount = strings.Count(sessLines[0], `"expiration"`)
}
result.Details = append(result.Details, fmt.Sprintf("Active sessions: %d", sessionCount))
if preview {
msg := fmt.Sprintf("PREVIEW: Would revoke %d sessions for user %s (ID %d)", sessionCount, login, userID)
if demote {
msg += " and demote to subscriber"
}
result.Message = msg
result.Success = true
return result
}
// Revoke sessions.
revokeQuery := fmt.Sprintf(
"UPDATE %susermeta SET meta_value='' WHERE user_id=%d AND meta_key='session_tokens'",
prefix, userID)
runMySQLQueryRoot(creds.dbName, revokeQuery)
result.Details = append(result.Details, "Sessions revoked")
// Demote to subscriber.
if demote {
// Read current capabilities to find the meta_key (varies by prefix).
capQuery := fmt.Sprintf(
"SELECT meta_key FROM %susermeta WHERE user_id=%d AND meta_key LIKE '%%capabilities'",
prefix, userID)
capLines := runMySQLQueryRoot(creds.dbName, capQuery)
if len(capLines) > 0 {
capKey := capLines[0]
demoteQuery := fmt.Sprintf(
"UPDATE %susermeta SET meta_value='a:1:{s:10:\"subscriber\";b:1;}' WHERE user_id=%d AND meta_key='%s'",
prefix, userID, escapeSQLString(capKey))
runMySQLQueryRoot(creds.dbName, demoteQuery)
result.Details = append(result.Details, "Demoted to subscriber role")
}
}
result.Message = fmt.Sprintf("Revoked sessions for user %s (ID %d)", login, userID)
result.Success = true
return result
}
// DBDeleteSpam deletes published posts matching spam patterns from a WordPress
// database. Only deletes posts of type 'post' with status 'publish' to avoid
// touching pages, attachments, or plugin data. If preview is true, reports
// counts without deleting.
//
// Two restrictions keep a keyword match from destroying real content: the
// SQL LIKE result is re-tested Go-side on a word boundary, and only the
// patterns marked deletable are considered. Injections that live inside an
// otherwise legitimate page are out of scope here -- deleting the page
// would destroy the customer's own content, so those need the injected
// markup stripped instead.
func DBDeleteSpam(account string, preview bool) DBCleanResult {
result := DBCleanResult{
Account: account,
Action: "delete-spam",
}
creds, prefix := findCredsForAccount(account)
if creds.dbName == "" {
result.Message = fmt.Sprintf("No WordPress database found for account %q", account)
return result
}
result.Database = creds.dbName
// SQL LIKE narrows candidates; the Go word-boundary check decides. LIKE
// alone matches "cialis" inside "specialist" and "pharma" inside a
// conference name, so a substring-only hit must never reach a DELETE.
// dbSpamPatterns is shared with the scanner so detection and deletion
// can never drift apart again.
spamIDs := make(map[string]struct{})
var order []string
var keywordCounts []string
for _, sp := range dbSpamPatterns {
if !sp.deletable {
continue
}
query := "SELECT ID, CONCAT_WS(' ', post_title, post_content) FROM " + prefix +
"posts WHERE post_type='post' AND post_status='publish' AND (post_content LIKE '" +
sp.likeFragment + "' OR post_title LIKE '" + sp.likeFragment + "')"
matched := 0
rows, err := runMySQLQueryRootWithError(creds.dbName, query)
if err != nil {
result.Message = fmt.Sprintf("Failed to query spam posts for %q: %v", sp.keyword, err)
return result
}
for _, line := range rows {
id, text, ok := splitSpamCandidate(line)
if !ok || len(spamKeywordMatchIndexes(sp, text)) == 0 {
continue
}
matched++
if _, seen := spamIDs[id]; !seen {
spamIDs[id] = struct{}{}
order = append(order, id)
}
}
if matched > 0 {
keywordCounts = append(keywordCounts,
fmt.Sprintf("%s=%s", sp.keyword, formatSpamPostCount(matched)))
}
}
if len(order) == 0 {
result.Message = "No spam posts found"
result.Success = true
return result
}
result.Details = append(result.Details,
"Matched "+formatCountedNoun(len(order), "unique post", "unique posts"),
"Keyword match counts overlap when a post contains several keywords: "+
strings.Join(keywordCounts, ", "))
if preview {
result.Message = "PREVIEW: Would delete " + formatSpamPostCount(len(order))
result.Success = true
return result
}
// Delete spam posts (and their revisions/meta) in batches.
deleted := 0
for i := 0; i < len(order); i += 100 {
end := i + 100
if end > len(order) {
end = len(order)
}
idList := strings.Join(order[i:end], ",")
affected, err := deleteSpamBatch(creds.dbName, prefix, idList)
if err != nil {
result.Message = fmt.Sprintf("Spam cleanup failed after deleting %s: %v",
formatSpamPostCount(deleted), err)
result.Details = append(result.Details,
"Cleanup stopped at the first failed database statement; rerun it after fixing the database error.")
return result
}
deleted += affected
}
result.Message = fmt.Sprintf("Deleted %s and %s metadata",
formatSpamPostCount(deleted), spamPostPossessive(deleted))
if notDeleted := len(order) - deleted; notDeleted > 0 {
verb := "were"
if notDeleted == 1 {
verb = "was"
}
result.Message += fmt.Sprintf("; %s %s not deleted",
formatCountedNoun(notDeleted, "candidate", "candidates"), verb)
}
result.Success = true
return result
}
func deleteSpamBatch(dbName, prefix, idList string) (int, error) {
statements := []struct {
action string
query string
}{
{
action: "delete post metadata",
query: fmt.Sprintf("DELETE FROM %spostmeta WHERE post_id IN (%s)", prefix, idList),
},
{
action: "delete post revisions",
query: fmt.Sprintf(
"DELETE FROM %sposts WHERE post_parent IN (%s) AND post_type='revision'",
prefix, idList),
},
}
for _, statement := range statements {
if _, err := runMySQLExecRootAffected(dbName, statement.query); err != nil {
return 0, fmt.Errorf("%s: %w", statement.action, err)
}
}
affected, err := runMySQLExecRootAffected(dbName, fmt.Sprintf(
"DELETE FROM %sposts WHERE ID IN (%s) AND post_type='post' AND post_status='publish'",
prefix, idList))
if err != nil {
return 0, fmt.Errorf("delete posts: %w", err)
}
return int(affected), nil
}
func formatSpamPostCount(count int) string {
return formatCountedNoun(count, "spam post", "spam posts")
}
func formatCountedNoun(count int, singular, plural string) string {
noun := plural
if count == 1 {
noun = singular
}
return fmt.Sprintf("%d %s", count, noun)
}
func spamPostPossessive(count int) string {
if count == 1 {
return "its"
}
return "their"
}
// FormatDBCleanResult formats a DBCleanResult for terminal output.
func FormatDBCleanResult(r DBCleanResult) string {
var sb strings.Builder
status := "FAILED"
if r.Success {
status = "OK"
}
fmt.Fprintf(&sb, "[%s] %s — %s\n", status, r.Action, r.Message)
if r.Database != "" {
fmt.Fprintf(&sb, " Database: %s\n", r.Database)
}
for _, d := range r.Details {
fmt.Fprintf(&sb, " %s\n", d)
}
return sb.String()
}
// --- helpers ---
// splitSpamCandidate splits an "ID\ttext" row from the spam candidate
// query. mysqlclient keeps each database row in one string and batch-escapes
// embedded controls; the bounded split also keeps literal tabs in test or
// alternate input from becoming extra columns. A nonnumeric ID is dropped
// rather than interpolated into a DELETE ... IN (...) list.
func splitSpamCandidate(line string) (id, text string, ok bool) {
parts := strings.SplitN(line, "\t", 2)
if len(parts) != 2 {
return "", "", false
}
id = strings.TrimSpace(parts[0])
if !spamPostIDRe.MatchString(id) {
return "", "", false
}
return id, mysqlclient.BatchUnescape(parts[1]), true
}
// findCredsForAccount finds WP database credentials for a cPanel account.
// Returns root-authenticated credentials that use /root/.my.cnf instead of
// wp-config.php passwords (which are often stale on cPanel servers).
func findCredsForAccount(account string) (wpDBCreds, string) {
patterns := wpInstallConfigPaths(wpInstallsForAccount(context.Background(), "db_content", account))
if len(patterns) == 0 {
return wpDBCreds{}, ""
}
sort.Strings(patterns)
primary := filepath.Join(accountHomeDir(account), "public_html", "wp-config.php")
if canonical, err := canonicalWPInstallPath(primary); err == nil {
primary = canonical
}
for i, path := range patterns {
if path == primary {
patterns[0], patterns[i] = patterns[i], patterns[0]
break
}
}
for _, path := range patterns {
creds := parseWPConfig(path)
if creds.dbName == "" {
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
continue
}
creds.tablePrefix = prefix
// Use root auth — CSM runs as root with /root/.my.cnf.
// wp-config.php passwords are unreliable (cPanel password
// rotations don't always update the file).
creds.dbUser = ""
creds.dbPass = ""
creds.dbHost = "localhost"
return creds, prefix
}
return wpDBCreds{}, ""
}
// resolveTablePrefix returns the safe `wp_options`-style prefix for the
// parsed wp-config. Empty input defaults to "wp_". Anything outside
// [A-Za-z0-9_]+ returns ("", false) so callers refuse to interpolate
// attacker-controlled $table_prefix into root-credentialled MySQL.
func resolveTablePrefix(creds wpDBCreds) (string, bool) {
prefix := creds.tablePrefix
if prefix == "" {
prefix = "wp_"
}
if !validTablePrefix.MatchString(prefix) {
return "", false
}
return prefix, true
}
// runMySQLExecRoot runs a non-SELECT MySQL statement under root
// credentials and returns the underlying exec error verbatim. Used
// by the persistence-mechanism cleaner where success is signalled by
// a zero exit + empty stdout, which runMySQLQueryRoot misclassifies
// as "no output, treat as failure".
func runMySQLExecRoot(dbName, stmt string) error {
_, err := runMySQLExecRootAffected(dbName, stmt)
return err
}
func runMySQLExecRootAffected(dbName, stmt string) (int64, error) {
return mysqlclient.RootExecSchema(context.Background(), dbName, stmt)
}
// runMySQLQueryRoot runs a MySQL query using root credentials from
// /root/.my.cnf (no explicit user/password args).
func runMySQLQueryRoot(dbName, query string) []string {
rows, _ := runMySQLQueryRootWithError(dbName, query)
return rows
}
func runMySQLQueryRootWithError(dbName, query string) ([]string, error) {
rows, err := mysqlclient.RootQuerySchema(context.Background(), dbName, query)
if err != nil {
return nil, err
}
if len(rows) == 0 {
return nil, nil
}
var lines []string
for _, line := range rows {
line = strings.TrimSpace(line)
if line != "" {
lines = append(lines, line)
}
}
return lines, nil
}
package checks
import (
"context"
"fmt"
"path/filepath"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// DB persistence-mechanism scanner.
//
// Vanilla CMS installs (WordPress, Joomla, Drupal, Magento, OpenCart)
// ship zero triggers / events / stored procedures / stored functions.
// Any presence is operator-review territory at minimum, and a body
// matching known-malware patterns is critical -- attacker persistence
// often re-injects on the next INSERT after a file-level cleanup, so
// detection here closes a real gap.
//
// The scanner reuses the existing parseWPConfig + runMySQLQuery
// infrastructure from dbscan.go. When the multi-CMS adapter layer
// lands later, this file's helpers become reusable across CMSes by
// swapping the credential discovery, not the queries.
//
// Per spec: detection only, no auto-drop. Operators drop manually
// via `csm db-clean drop-object`.
// dbPersistenceMalwarePatterns supplements dbMalwarePatterns from
// dbscan.go with MySQL-specific persistence-attack signals. Lowercase
// for case-insensitive matching against SQL bodies.
var dbPersistenceMalwarePatterns = []string{
"sys_exec",
"lib_mysqludf_sys",
"into outfile",
"into dumpfile",
"load_file(",
"load data infile",
}
// magicTokenRegex extracts high-entropy activation tokens from trigger
// bodies that gate privileged actions on `display_name LIKE
// '%<token>%'`. Common display-name filters such as "%administrator%"
// must not escalate a merely unexpected trigger into Critical.
var magicTokenRegex = regexp.MustCompile(`(?i)display_name\s+like\s+['"]%([A-Za-z0-9_-]{10,32})%['"]`)
// validTablePrefix matches the character class WordPress accepts for
// $table_prefix (alphanumerics and underscore). Untrusted prefixes
// from a malformed wp-config.php fail this check, which keeps
// scanMagicTokenUsers from concatenating attacker-controlled data into
// its SQL literal.
var validTablePrefix = regexp.MustCompile(`^[A-Za-z0-9_]+$`)
// dbPersistenceMalwareRegexes catches multi-token shapes that no
// substring set can match cleanly: role-escalation writes and
// password-hash exfiltration reads. Pre-compiled at package init time
// -- a regex parse error here is a build-time bug, not a runtime one.
//
// Patterns intentionally case-insensitive ((?i) prefix) and tolerant of
// whitespace / line breaks across MySQL trigger bodies. The role-write
// pattern requires the literal string "administrator" inside the
// serialized capabilities payload -- promotion to subscriber/customer
// is the legitimate WP-signup shape and must not match.
var dbPersistenceMalwareRegexes = []*regexp.Regexp{
// Role escalation: UPDATE on *_usermeta writing administrator caps.
// The (?s) flag lets `.` match newlines so multi-line trigger
// bodies with the UPDATE split across lines still hit.
regexp.MustCompile(`(?is)update\s+` + "`?" + `\w*usermeta` + "`?" + `\s+set\s+meta_value\s*=.*?(?:s:13:["\x60]administrator["\x60]|["\x60]administrator["\x60])`),
// Password-hash exfil read: SELECT user_pass FROM <users-like>
// table. Real WP code goes through wp_check_password() in PHP, never
// raw SELECT user_pass from SQL.
regexp.MustCompile(`(?i)select\s+user_pass\s+from\s+` + "`?" + `\w*users`),
}
// dbObjectKind names the four MySQL object types this scanner
// inspects. Used in finding categories and CLI subcommands.
type dbObjectKind string
const (
dbObjectTrigger dbObjectKind = "trigger"
dbObjectEvent dbObjectKind = "event"
dbObjectProcedure dbObjectKind = "procedure"
dbObjectFunction dbObjectKind = "function"
)
// dbObjectAllKinds lists all valid kinds for the CLI's type validator.
var dbObjectAllKinds = []dbObjectKind{
dbObjectTrigger, dbObjectEvent, dbObjectProcedure, dbObjectFunction,
}
// dbObjectFinding describes one detector hit before it is converted
// to an alert.Finding -- carrying the structured fields the CLI
// drop-object subcommand needs to look up the same row.
type dbObjectFinding struct {
Account string
Schema string
Kind dbObjectKind
Name string
Body string
IsMalw bool // true: malware pattern hit; false: unexpected presence
}
// CheckDatabaseObjects scans every WordPress installation's database
// for triggers, events, procedures, and functions. Critical findings
// fire when the body matches a known-malware pattern; Warning
// findings fire when an object exists at all (vanilla CMSes ship
// none). The Detection.DBObjectScanning kill-switch silences both
// emit paths without disabling the manual drop-object CLI.
func CheckDatabaseObjects(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !dbObjectScanningEnabled(cfg) {
return nil
}
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
wpConfigs := dbObjectWPConfigs(ctx)
if len(wpConfigs) == 0 {
return nil
}
allowlist := dbObjectAllowlistMap(cfg)
// Rank by mtime desc so recently touched WP installs are processed
// first when the check timeout cuts iteration short.
for _, wpConfig := range rankPathsByMtimeDesc(ctx, wpConfigs, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
return findings
}
account := wpConfigUser(filepath.Dir(wpConfig))
creds, complete := parseWPConfigChecked(wpConfig)
if !complete {
markCheckIncomplete(ctx, "db_objects")
continue
}
if creds.dbName == "" || creds.dbUser == "" {
markCheckIncomplete(ctx, "db_objects")
continue
}
var installFindings []alert.Finding
hits, err := dbObjectScanner(account, creds)
if err != nil {
markCheckIncomplete(ctx, "db_objects")
}
for _, h := range hits {
if !h.IsMalw && allowlist[allowlistKey(h)] {
continue
}
installFindings = append(installFindings, h.toFinding())
}
// Retro-scan: when a trigger gates a privileged action on a
// secret token in display_name, find users whose display_name
// still carries the token. Zero matches is itself useful for
// the incident report ("no evidence the backdoor fired").
for _, h := range hits {
if h.Kind != dbObjectTrigger {
continue
}
tokens := magicTokensOf(h.Body)
if len(tokens) == 0 {
continue
}
tokenFindings, err := magicTokenScanner(account, creds.dbName, creds.tablePrefix, tokens)
if err != nil {
markCheckIncomplete(ctx, "db_objects")
}
installFindings = append(installFindings, tokenFindings...)
}
// The display label may be a lookup sentinel; only a resolved
// account root owner is stamped, per install, before merging.
owner, ok := installOwner(wpConfig)
if ok {
installFindings = stampTenantIDIfEmpty(installFindings, owner)
}
findings = append(findings, installFindings...)
}
return findings
}
// Per-install scan boundaries. Tests replace them with inert scanners to
// prove ownership stamping for every finding name the check owns.
var (
dbObjectScanner = scanDBObjects
magicTokenScanner = scanMagicTokenUsers
magicTokensOf = extractMagicTokens
)
// scanDBObjects runs the three INFORMATION_SCHEMA queries and
// classifies every row. Pure function over the cmdExec injector --
// tests provide canned MySQL CLI output and assert on the structured
// findings without touching a real database.
//
// Connections use root credentials via /root/.my.cnf. WP-config
// passwords are unreliable on cPanel hosts (password rotations
// don't update the file), so a WP-creds path here would silently
// miss persistence objects on the very platform we care most about.
// The existing db-clean code (db_clean.go: findCredsForAccount)
// hits the same constraint and reaches the same conclusion.
func scanDBObjects(account string, creds wpDBCreds) ([]dbObjectFinding, error) {
if creds.dbName == "" {
return nil, nil
}
schema := creds.dbName
schemaLit := mysqlSchemaLiteral(schema)
var hits []dbObjectFinding
// TRIGGERS
rows, err := runMySQLQueryRootWithError(schema, fmt.Sprintf(
`SELECT TRIGGER_NAME, ACTION_STATEMENT FROM INFORMATION_SCHEMA.TRIGGERS WHERE TRIGGER_SCHEMA = %s`,
schemaLit))
if err != nil {
return hits, err
}
for _, row := range rows {
name, body := splitTabRow(row)
if name == "" {
continue
}
hits = append(hits, classifyDBObject(account, schema, dbObjectTrigger, name, body))
}
// EVENTS
rows, err = runMySQLQueryRootWithError(schema, fmt.Sprintf(
`SELECT EVENT_NAME, EVENT_DEFINITION FROM INFORMATION_SCHEMA.EVENTS WHERE EVENT_SCHEMA = %s`,
schemaLit))
if err != nil {
return hits, err
}
for _, row := range rows {
name, body := splitTabRow(row)
if name == "" {
continue
}
hits = append(hits, classifyDBObject(account, schema, dbObjectEvent, name, body))
}
// ROUTINES (procedures + functions)
rows, err = runMySQLQueryRootWithError(schema, fmt.Sprintf(
`SELECT ROUTINE_NAME, ROUTINE_TYPE, ROUTINE_DEFINITION FROM INFORMATION_SCHEMA.ROUTINES WHERE ROUTINE_SCHEMA = %s`,
schemaLit))
if err != nil {
return hits, err
}
for _, row := range rows {
name, rtype, body := splitTabRow3(row)
if name == "" {
continue
}
kind := dbObjectProcedure
if strings.EqualFold(rtype, "FUNCTION") {
kind = dbObjectFunction
}
hits = append(hits, classifyDBObject(account, schema, kind, name, body))
}
return hits, nil
}
// classifyDBObject decides whether a row matches the malware
// patterns (Critical) or merely exists (Warning).
func classifyDBObject(account, schema string, kind dbObjectKind, name, body string) dbObjectFinding {
return dbObjectFinding{
Account: account,
Schema: schema,
Kind: kind,
Name: name,
Body: body,
IsMalw: bodyHasMalwarePattern(body),
}
}
// bodyHasMalwarePattern returns true when the SQL body matches any of
// the three classifier tiers:
//
// 1. dbMalwarePatterns / dbPersistenceMalwarePatterns: substring tokens
// for OS-exec UDFs and file-IO sinks (sys_exec, INTO OUTFILE, etc.).
// 2. extractMagicTokens: high-entropy display_name activation gates.
// 3. dbPersistenceMalwareRegexes: multi-token shapes for role-escalation
// writes and password-hash exfiltration reads.
//
// Substring matching stays case-insensitive via ToLower; the regex tier
// keeps its own `(?i)` flags so its semantics travel with the pattern.
func bodyHasMalwarePattern(body string) bool {
body = normalizeDBPatternBody(body)
lower := strings.ToLower(body)
for _, p := range dbMalwarePatterns {
if strings.Contains(lower, strings.ToLower(p.pattern)) {
return true
}
}
for _, p := range dbPersistenceMalwarePatterns {
if strings.Contains(lower, p) {
return true
}
}
if len(extractMagicTokens(body)) > 0 {
return true
}
for _, re := range dbPersistenceMalwareRegexes {
if re.MatchString(body) {
return true
}
}
return false
}
func normalizeDBPatternBody(body string) string {
repl := strings.NewReplacer(`\n`, "\n", `\r`, "\r", `\t`, "\t")
return repl.Replace(body)
}
// toFinding renders the structured hit into the alert.Finding shape
// the rest of the pipeline expects. Finding category encodes both
// kind and severity tier so operators can suppress per attack type.
func (h dbObjectFinding) toFinding() alert.Finding {
check := fmt.Sprintf("db_unexpected_%s", h.Kind)
severity := alert.Warning
intro := "Unexpected"
if h.IsMalw {
check = fmt.Sprintf("db_malicious_%s", h.Kind)
severity = alert.Critical
intro = "Malicious"
}
excerpt := h.Body
if len(excerpt) > 240 {
excerpt = excerpt[:240] + "..."
}
return alert.Finding{
Severity: severity,
Check: check,
Message: fmt.Sprintf("%s %s %s in %s.%s", intro, h.Kind, h.Name, h.Account, h.Schema),
Details: fmt.Sprintf("Account: %s\nSchema: %s\nKind: %s\nName: %s\nBody: %s", h.Account, h.Schema, h.Kind, h.Name, excerpt),
Timestamp: time.Now(),
}
}
// allowlistKey shapes the suppression key per spec:
// `<account>:<schema>:<type>:<name>`. Used for the Warning tier
// only -- Critical malware-pattern hits always fire.
func allowlistKey(h dbObjectFinding) string {
return fmt.Sprintf("%s:%s:%s:%s", h.Account, h.Schema, h.Kind, h.Name)
}
func dbObjectAllowlistMap(cfg *config.Config) map[string]bool {
out := map[string]bool{}
if cfg == nil {
return out
}
for _, e := range cfg.Detection.DBObjectAllowlist {
out[strings.TrimSpace(e)] = true
}
return out
}
// splitTabRow returns the first two tab-separated fields from a
// MySQL `-B -N` row. Empty strings if the row has fewer than two
// fields.
func splitTabRow(row string) (string, string) {
parts := strings.SplitN(row, "\t", 2)
if len(parts) < 2 {
return "", ""
}
return parts[0], parts[1]
}
// splitTabRow3 returns the first three tab-separated fields. Used
// for ROUTINES which carries (name, type, body).
func splitTabRow3(row string) (string, string, string) {
parts := strings.SplitN(row, "\t", 3)
if len(parts) < 3 {
return "", "", ""
}
return parts[0], parts[1], parts[2]
}
// mysqlSchemaLiteral wraps the schema name as a single-quoted
// string literal with backslash escaping. The DB name comes from
// wp-config.php (operator-controlled) so the risk is low, but the
// string-literal route is consistent with how the existing dbscan
// queries handle string args.
func mysqlSchemaLiteral(name string) string {
escaped := strings.ReplaceAll(name, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `'`, `\'`)
return "'" + escaped + "'"
}
// dbObjectScanningEnabled resolves the tri-state cfg flag: nil and
// missing-config both mean default-on; an explicit *false means off.
func dbObjectScanningEnabled(cfg *config.Config) bool {
if cfg == nil {
return true
}
if cfg.Detection.DBObjectScanning == nil {
return true
}
return *cfg.Detection.DBObjectScanning
}
// IsDBObjectKind reports whether s is one of the four valid kinds.
// Used by the CLI subcommand to validate user input before opening
// a connection.
func IsDBObjectKind(s string) bool {
for _, k := range dbObjectAllKinds {
if string(k) == s {
return true
}
}
return false
}
// extractMagicTokens returns the secret activation tokens referenced in
// a trigger body's `display_name LIKE '%<token>%'` clauses. The body
// classifier uses the same helper, and the user retro scan reuses the
// returned tokens to search the *_users table for matches. Tokens are
// deduplicated to keep query count bounded when a trigger references
// the same token across multiple branches.
//
// Returns nil for benign bodies so callers can skip MySQL entirely.
func extractMagicTokens(body string) []string {
body = normalizeDBPatternBody(body)
matches := magicTokenRegex.FindAllStringSubmatch(body, -1)
if len(matches) == 0 {
return nil
}
seen := make(map[string]struct{}, len(matches))
var out []string
for _, m := range matches {
if len(m) < 2 {
continue
}
tok := m[1]
if !validMagicToken(tok) {
continue
}
if _, ok := seen[tok]; ok {
continue
}
seen[tok] = struct{}{}
out = append(out, tok)
}
return out
}
func validMagicToken(tok string) bool {
if len(tok) < 10 || len(tok) > 32 {
return false
}
hasUpper, hasLower, hasDigit := false, false, false
for _, r := range tok {
switch {
case r >= 'A' && r <= 'Z':
hasUpper = true
case r >= 'a' && r <= 'z':
hasLower = true
case r >= '0' && r <= '9':
hasDigit = true
case r == '_' || r == '-':
default:
return false
}
}
return hasUpper && hasLower && hasDigit
}
// scanMagicTokenUsers searches the WordPress users table for accounts
// whose display_name carries a backdoor activation token. A match is
// forensic evidence that the trigger fired against that user -- they
// may still be administrator, or the attacker may have demoted them
// after promotion. Either way the user requires manual review and is
// surfaced as Critical.
//
// The function is conservative about query construction. Tokens are
// guaranteed to be high-entropy [A-Za-z0-9_-]{10,32} strings by
// extractMagicTokens, and the table prefix is validated against
// [A-Za-z0-9_]+ before concatenation. Anything outside those character
// classes causes the scan to skip the query entirely rather than emit a
// half-built SQL statement against an untrusted prefix.
func scanMagicTokenUsers(account, schema, tablePrefix string, tokens []string) ([]alert.Finding, error) {
if len(tokens) == 0 || tablePrefix == "" || !validTablePrefix.MatchString(tablePrefix) {
return nil, nil
}
var findings []alert.Finding
for _, tok := range tokens {
if !validMagicToken(tok) {
continue
}
query := fmt.Sprintf(
"SELECT ID, user_login, user_email, display_name FROM `%susers` WHERE display_name LIKE '%%%s%%'",
tablePrefix, tok,
)
rows, err := runMySQLQueryRootWithError(schema, query)
if err != nil {
return findings, err
}
for _, row := range rows {
parts := strings.SplitN(row, "\t", 4)
if len(parts) < 4 {
continue
}
userID, userLogin, userEmail, displayName := parts[0], parts[1], parts[2], parts[3]
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "db_magic_token_user",
Message: fmt.Sprintf("User %s (ID %s) carries backdoor activation token in %s.%susers", userLogin, userID, account, tablePrefix),
Details: fmt.Sprintf("Account: %s\nSchema: %s\nTable prefix: %s\nToken: %s\nUser ID: %s\nUser login: %s\nUser email: %s\nDisplay name: %s", account, schema, tablePrefix, tok, userID, userLogin, userEmail, displayName),
Timestamp: time.Now(),
})
}
}
return findings, nil
}
// dbObjectWPConfigs lists the WordPress installs this check scans. Discovery is
// shared (wpinstalls.go): a subdomain or nested install carries the same
// injected objects as a primary one.
func dbObjectWPConfigs(ctx context.Context) []string {
installs := wpInstalls(ctx, "db_objects")
out := make([]string, 0, len(installs))
for _, in := range installs {
out = append(out, in.ConfigPath)
}
return out
}
package checks
import (
"context"
"errors"
"fmt"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/mysqlclient"
"github.com/pidginhost/csm/internal/store"
)
// showCreateStatement extracts the CREATE statement from one batch-mode SHOW
// CREATE row. The client returns a row as tab-joined columns with newlines
// escaped, so the statement is one column, not the row: TRIGGER, PROCEDURE and
// FUNCTION put it third after the name and sql_mode, EVENT inserts time_zone
// before it. Replaying the whole row would fail with a syntax error every
// time, after the object is already gone, so anything else is rejected.
func showCreateStatement(kind, row string) (string, error) {
cols := strings.Split(row, "\t")
idx := 2
if strings.EqualFold(kind, "event") {
idx = 3
}
if len(cols) <= idx {
return "", fmt.Errorf("expected at least %d columns, got %d", idx+1, len(cols))
}
stmt := mysqlclient.BatchUnescape(cols[idx])
fields := strings.Fields(stmt)
if len(fields) < 2 || !strings.EqualFold(fields[0], "CREATE") {
return "", errors.New("statement column does not hold a CREATE statement")
}
objectIdx := 1
if len(fields) > objectIdx+1 && strings.EqualFold(fields[objectIdx], "OR") && strings.EqualFold(fields[objectIdx+1], "REPLACE") {
objectIdx += 2
}
if len(fields) > objectIdx && strings.HasPrefix(strings.ToUpper(fields[objectIdx]), "DEFINER=") {
objectIdx++
}
if len(fields) <= objectIdx || !strings.EqualFold(fields[objectIdx], kind) {
return "", fmt.Errorf("CREATE statement is not for a %s", kind)
}
return stmt, nil
}
// reAccountName matches the cPanel-username shape we accept from
// operator CLI input. Constrained on purpose: anything outside this
// charset would either fail later validation (QuoteIdent on schema)
// or escape /home via the path interpolation in findAccountSchemas
// when an unwary glob expanded a `*`.
var reAccountName = regexp.MustCompile(`^[a-zA-Z][a-zA-Z0-9_-]{0,31}$`)
// errInvalidAccountName flags an account string that fails the
// allowed-charset check. Surfaced through the CLI so the operator
// sees a clear error before any filesystem or SQL lookup.
var errInvalidAccountName = errors.New("invalid account name (want [a-zA-Z][a-zA-Z0-9_-]{0,31})")
// DBDropObject drops a single trigger / event / stored procedure /
// stored function from the operator-supplied account+schema, after:
//
// 1. Validating the kind ("trigger" | "event" | "procedure" | "function").
// 2. Validating that <schema> is one of the databases this account
// hosts. The account is taken from /home/<account>/* wp-config.php
// files; an attacker who can pass an arbitrary <schema> here gets
// no further than DROP'ping their own database.
// 3. QuoteIdent on both <schema> and <name>, so identifier strings
// never participate in SQL string concatenation.
// 4. SHOW CREATE the object and persist the result to the
// db_object_backups bbolt bucket as the backup -- replaying the
// CREATE SQL restores the object byte-for-byte.
// 5. DROP the object.
//
// preview=true short-circuits before step 4: the function reports
// what it would do (kind, schema, name, captured CREATE SQL) without
// touching the database.
//
// Per spec: detection is always-on; drop is operator-driven.
func DBDropObject(account, schema, kind, name string, preview bool) DBCleanResult {
result := DBCleanResult{
Account: account,
Action: "drop-object",
}
if !reAccountName.MatchString(account) {
result.Message = fmt.Sprintf("%v: %q", errInvalidAccountName, account)
return result
}
if !IsDBObjectKind(kind) {
result.Message = fmt.Sprintf("Invalid object kind %q (want trigger|event|procedure|function)", kind)
return result
}
quotedSchema, err := QuoteIdent(schema)
if err != nil {
result.Message = fmt.Sprintf("Invalid schema name: %v", err)
return result
}
quotedName, err := QuoteIdent(name)
if err != nil {
result.Message = fmt.Sprintf("Invalid object name: %v", err)
return result
}
knownSchemas := findAccountSchemas(account)
if !containsString(knownSchemas, schema) {
result.Message = fmt.Sprintf("Schema %q is not one of the databases discovered for account %q (known: %v)",
schema, account, knownSchemas)
return result
}
result.Database = schema
// SHOW CREATE captures the backup. Different MySQL grammars per
// kind: TRIGGER and EVENT use the schema-qualified name in
// `<schema>.<name>` form; PROCEDURE and FUNCTION accept the same
// shape under modern MySQL. Use the unified form for consistency.
showCreateSQL := fmt.Sprintf("SHOW CREATE %s %s.%s",
strings.ToUpper(kind), quotedSchema, quotedName)
createOutput := runMySQLQueryRoot(schema, showCreateSQL)
if len(createOutput) == 0 {
result.Message = fmt.Sprintf("SHOW CREATE returned no rows for %s %s.%s -- object missing or permission denied",
kind, schema, name)
return result
}
createSQL, err := showCreateStatement(kind, createOutput[0])
if err != nil {
result.Message = fmt.Sprintf("SHOW CREATE for %s %s.%s did not yield a restorable statement (refusing to drop): %v",
kind, schema, name, err)
return result
}
if preview {
result.Message = fmt.Sprintf("PREVIEW: would drop %s %s.%s", kind, schema, name)
result.Details = []string{
fmt.Sprintf("Captured CREATE SQL (%d bytes)", len(createSQL)),
"No backup written and no DROP executed in preview mode.",
}
result.Success = true
return result
}
// Persist backup BEFORE the drop so a SQL failure on DROP still
// leaves the operator with a record of what existed.
sdb := store.Global()
if sdb == nil {
result.Message = "bbolt store not available; refusing to drop without a recorded backup"
return result
}
if err := sdb.PutDBObjectBackup(store.DBObjectBackup{
Account: account,
Schema: schema,
Kind: kind,
Name: name,
CreateSQL: createSQL,
DroppedAt: time.Now().UTC(),
DroppedBy: "csm-cli",
}); err != nil {
result.Message = fmt.Sprintf("recording backup failed (refusing to drop): %v", err)
return result
}
dropSQL := fmt.Sprintf("DROP %s IF EXISTS %s.%s",
strings.ToUpper(kind), quotedSchema, quotedName)
// runMySQLExecRoot reports the mysql client's exec error
// directly. The previous use of runMySQLQueryRoot misread a
// zero-exit + empty-stdout (the success signature for DROP) as
// failure.
if err := runMySQLExecRoot(schema, dropSQL); err != nil {
result.Message = fmt.Sprintf("DROP %s %s.%s failed: %v", kind, schema, name, err)
return result
}
result.Details = []string{
fmt.Sprintf("Dropped %s %s.%s", kind, schema, name),
fmt.Sprintf("Backup recorded in bbolt: %d bytes", len(createSQL)),
}
result.Message = fmt.Sprintf("Dropped %s %s.%s (backup retained)", kind, schema, name)
result.Success = true
return result
}
// findAccountSchemas returns every distinct database name discovered
// across the account's wp-config.php files. Multiple WordPress
// installations under the same account commonly reuse one database
// but can use several; the CLI relies on this list to validate
// operator input before opening any connection.
func findAccountSchemas(account string) []string {
patterns := wpInstallConfigPaths(wpInstallsForAccount(context.Background(), "db_objects", account))
seen := map[string]struct{}{}
var out []string
for _, path := range patterns {
// parseWPConfig handles missing files silently, so the bare
// non-glob first-entry path is harmless when the account has
// no public_html/wp-config.php.
creds := parseWPConfig(path)
if creds.dbName == "" {
continue
}
if _, ok := seen[creds.dbName]; ok {
continue
}
seen[creds.dbName] = struct{}{}
out = append(out, creds.dbName)
}
return out
}
// containsString reports whether haystack contains needle. Local
// because the package's other helper of the same name lives in a
// _test.go file (waf_test.go) and is not visible to production
// builds.
func containsString(haystack []string, needle string) bool {
for _, h := range haystack {
if h == needle {
return true
}
}
return false
}
// RestoreDBObjectBackup re-executes the captured CREATE SQL for a
// previously-dropped MySQL trigger / event / procedure / function.
// Looks up the row in the db_object_backups bbolt bucket by exact
// key; the caller (typically the web UI's cleanup-history page)
// supplies the key it got from the listing endpoint.
//
// Per spec the operation is operator-driven: there is no auto-
// restore. The webui handler enforces the same.
func RestoreDBObjectBackup(backupKey string) DBCleanResult {
result := DBCleanResult{Action: "restore-object"}
sdb := store.Global()
if sdb == nil {
result.Message = "bbolt store not available"
return result
}
rec, ok, err := sdb.GetDBObjectBackupByKey(backupKey)
if err != nil {
result.Message = fmt.Sprintf("looking up backup: %v", err)
return result
}
if !ok {
result.Message = "backup not found (may have been pruned)"
return result
}
result.Account = rec.Account
result.Database = rec.Schema
if rec.CreateSQL == "" {
result.Message = "backup record has no CREATE SQL"
return result
}
if err := runMySQLExecRoot(rec.Schema, rec.CreateSQL); err != nil {
result.Message = fmt.Sprintf("re-executing CREATE failed: %v", err)
return result
}
result.Details = []string{
fmt.Sprintf("Restored %s %s.%s", rec.Kind, rec.Schema, rec.Name),
fmt.Sprintf("Original drop: %s by %s", rec.DroppedAt.Format(time.RFC3339), rec.DroppedBy),
}
if err := sdb.MarkDBObjectBackupRestored(backupKey, time.Now().UTC()); err != nil {
result.Details = append(result.Details, fmt.Sprintf("Restore state was not recorded: %v", err))
}
result.Message = fmt.Sprintf("Restored %s %s.%s from backup", rec.Kind, rec.Schema, rec.Name)
result.Success = true
return result
}
package checks
import (
"fmt"
"net/url"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/store"
)
// externalScriptHosts returns the hosts of every <script src> in content
// that is neither a known-safe service nor structurally malicious. The
// structural classifier already reports the latter as Critical; the former
// are pre-approved. What remains is an unremarkable external host, which is
// exactly the shape a careful injection takes and which no other signal
// covers.
func externalScriptHosts(content string) []string {
var hosts []string
seen := map[string]bool{}
// A payload stored as JSON carries escaped slashes, which the src
// grammar does not match; normalise before extracting.
content = unescapeStoredSlashes(content)
for _, match := range scriptSrcRe.FindAllStringSubmatch(content, -1) {
if len(match) < 2 {
continue
}
raw := match[1]
if isSafeScriptDomain(raw) {
continue
}
if bad, _ := scriptSrcMaliciousReason(raw); bad {
continue
}
normalised := raw
if strings.HasPrefix(normalised, "//") {
normalised = "https:" + normalised
}
u, err := url.Parse(normalised)
if err != nil || u == nil {
continue
}
host := strings.ToLower(u.Hostname())
if host == "" || seen[host] {
continue
}
seen[host] = true
hosts = append(hosts, host)
}
return hosts
}
// newExternalScriptFindings reports, as a Warning, each external script host
// in an option value that firstSeen has not recorded before. Severity stays
// below the auto-response threshold: this is a "look at this" signal, not
// proof of injection.
func newExternalScriptFindings(user string, creds wpDBCreds, prefix, option, value string, firstSeen func(option, host string) bool) []alert.Finding {
var findings []alert.Finding
for _, host := range externalScriptHosts(value) {
if !firstSeen(option, host) {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "db_options_new_external_script",
Message: fmt.Sprintf("New external script host in wp_options '%s' (account: %s): %s", option, user, host),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Option: %s", option),
fmt.Sprintf("Script host: %s", host),
fmt.Sprintf("Content preview: %s", truncateDB(value, 200)),
"First appearance of this host since the site's baseline; verify it is a service the site owner added."),
DedupKey: dbContentDedupKey(user, creds, prefix,
fmt.Sprintf("Option: %s", option),
fmt.Sprintf("Script host: %s", host),
fmt.Sprintf("Content preview: %s", truncateDB(value, 200)),
"First appearance of this host since the site's baseline; verify it is a service the site owner added."),
})
}
return findings
}
// externalScriptSiteKey identifies one WordPress install's options table in
// the first-seen store.
func externalScriptSiteKey(dbName, prefix string) string {
return dbName + "|" + prefix
}
// storeFirstSeen adapts the bbolt store to the firstSeen callback; without a
// store nothing is ever new, so the scan stays silent rather than reporting
// every host on every cycle.
func storeFirstSeen(site string) func(option, host string) bool {
db := store.Global()
if db == nil {
return func(string, string) bool { return false }
}
return func(option, host string) bool {
isNew, err := db.MarkExternalScriptSeen(site, option, host, time.Now())
return err == nil && isNew
}
}
package checks
import (
"encoding/hex"
"fmt"
"regexp"
"sort"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// pluginNoticeSinkOptions are plugin status options whose contents WordPress
// renders into the admin dashboard as a notice. A stored cross-site scripting
// bug in the plugin turns one of these rows into a loader that runs with the
// privileges of whichever administrator opens the dashboard next.
//
// LiteSpeed Cache before 5.7.0.1 (CVE-2023-40000) lets an unauthenticated
// request write its CDN setup status, and the observed campaign puts a
// <script src> into cdn_setup_err, then into the rendered message list. Both
// rows hold plugin state, never site content, so markup that executes is
// proof of injection on its own.
var pluginNoticeSinkOptions = map[string]string{
"litespeed.cdn_setup._summary": "LiteSpeed Cache CDN setup status",
"litespeed.admin_display.messages": "LiteSpeed Cache admin notices",
"litespeed.admin_display.msg_pin": "LiteSpeed Cache pinned admin notice",
}
// executableMarkupRe matches markup that runs code when a notice is rendered.
// Notice sinks legitimately carry layout markup -- LiteSpeed writes its own
// errors as a styled div -- so only the executing constructs count.
var executableMarkupRe = regexp.MustCompile(`(?i)<script[\s>]|<iframe[\s>]|javascript:|\bon(?:error|load|click|mouseover)\s*=|eval\s*\(\s*atob`)
const maxPluginNoticeBytes = 65536
// Notice lists can grow beyond a preview-sized read. Bound each value in bytes
// and carry its stored length so omitted content never looks like a clean scan.
// Hex preserves whitespace and stored escapes through the batch-row transport.
func checkWPPluginNotices(user string, creds wpDBCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"SELECT option_name, OCTET_LENGTH(option_value), "+
"CONCAT('x', HEX(LEFT(CAST(option_value AS BINARY), %d))) FROM %soptions "+
"WHERE option_name IN (%s) LIMIT %d",
maxPluginNoticeBytes, prefix, pluginNoticeSinkNameList(), len(pluginNoticeSinkOptions))
var findings []alert.Finding
for _, line := range runMySQLQuery(creds, query) {
option, value, complete := parsePluginNoticeRow(line)
if !complete {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if finding := pluginNoticeInjectionFinding(user, creds, prefix, option, value); finding != nil {
findings = append(findings, *finding)
}
}
return findings
}
func parsePluginNoticeRow(line string) (option, value string, complete bool) {
parts := strings.SplitN(line, "\t", 3)
if len(parts) != 3 {
return "", "", false
}
size, err := strconv.ParseInt(parts[1], 10, 64)
if err != nil || size < 0 || size > maxPluginNoticeBytes {
return "", "", false
}
encoded := parts[2]
if !strings.HasPrefix(encoded, "x") || int64(len(encoded)-1) != 2*size {
return "", "", false
}
decoded, err := hex.DecodeString(encoded[1:])
if err != nil {
return "", "", false
}
return parts[0], string(decoded), true
}
// pluginNoticeInjectionFinding reports executable markup stored in a plugin
// status option that WordPress renders as an admin notice. The option's
// identity carries the verdict: neither host reputation nor a first-seen
// baseline applies, so a payload that predates CSM's baseline and a loader on
// an ordinary HTTPS host are both reported.
func pluginNoticeInjectionFinding(user string, creds wpDBCreds, prefix, option, value string) *alert.Finding {
name := strings.ToLower(strings.TrimSpace(option))
sink, ok := pluginNoticeSinkOptions[name]
if !ok {
return nil
}
decoded := unescapeStoredSlashes(value)
marker := executableMarkupRe.FindString(decoded)
if marker == "" {
return nil
}
detail := []string{
fmt.Sprintf("Option: %s (%s)", option, sink),
fmt.Sprintf("Executable markup: %s", strings.TrimSpace(marker)),
}
if hosts := externalScriptHosts(value); len(hosts) > 0 {
detail = append(detail, fmt.Sprintf("Script host: %s", strings.Join(hosts, ", ")))
}
detail = append(detail,
fmt.Sprintf("Content preview: %s", truncateDB(decoded, 200)),
"This option holds plugin state, not site content. WordPress prints it in the dashboard, so the code runs for the next administrator who opens it.")
return &alert.Finding{
Severity: alert.Critical,
Check: "db_options_plugin_notice_injection",
Message: fmt.Sprintf("Executable markup stored in plugin notice option '%s' (account: %s)", option, user),
Details: dbContentFindingDetails(creds, prefix, detail...),
DedupKey: dbContentDedupKey(user, creds, prefix, detail...),
}
}
// pluginNoticeSinkNameList renders the sink names for a SQL IN clause.
func pluginNoticeSinkNameList() string {
names := make([]string, 0, len(pluginNoticeSinkOptions))
for name := range pluginNoticeSinkOptions {
names = append(names, "'"+name+"'")
}
sort.Strings(names)
return strings.Join(names, ", ")
}
package checks
import (
"bufio"
"context"
"crypto/sha256"
"encoding/binary"
"fmt"
"io"
"net/netip"
"net/url"
"path/filepath"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/mysqlclient"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// Malicious patterns in WordPress database content.
//
// requiresExternalScript: when true, a matching row is only reported if
// its content also contains a <script src=...> pointing at a domain NOT
// on the known-safe list. This filters out the legitimate analytics and
// widget embeds that site owners place in page content (Google Tag
// Manager, Google merchant badge, HubSpot, Mailchimp, etc.) without
// weakening detection of attacker-injected external loaders.
var dbMalwarePatterns = []struct {
pattern string
severity alert.Severity
desc string
requiresExternalScript bool
}{
// The script-tag entry catches BOTH inline <script> blocks and
// <script src=...> loaders as a fast LIKE pre-filter; the Go post-
// filter (hasMaliciousExternalScript) verifies the presence of a
// non-safe-domain external src before raising a finding. Inline
// obfuscation without an external src is caught by the subsequent
// code-pattern entries below.
{"<script", alert.High, "injected <script> tag with non-safe external src", true},
{"eval(", alert.High, "eval() in database content", false},
{"base64_decode", alert.High, "base64_decode in database content", false},
{"document.write(", alert.High, "document.write injection", false},
{"String.fromCharCode", alert.High, "JavaScript obfuscation (fromCharCode)", false},
{".workers.dev", alert.Critical, "Cloudflare Workers exfiltration URL", false},
{"gist.githubusercontent.com", alert.Critical, "GitHub Gist payload URL", false},
{"pastebin.com/raw", alert.Critical, "Pastebin payload URL", false},
}
// nonDocRootDirs are common account-data directories that are never candidate
// document roots during the home-directory walk. "www" is deliberately absent:
// where it is cPanel's alias for public_html the discovery walk collapses the
// symlink (canonicalWPInstallPath), and where it is a real directory it is a
// document root serving a real site.
var nonDocRootDirs = map[string]bool{
"mail": true, "etc": true, "logs": true, "ssl": true, "tmp": true,
"public_ftp": true, "cache": true, ".cagefs": true,
"access-logs": true, "access_logs": true, "backups": true,
"cgi-bin": true, "perl5": true, "spamassassin": true, "var": true,
}
// servedState records whether the panel currently serves a document root.
// A dormant install is not harmless -- it holds a live database and becomes
// public again the moment the domain is re-pointed -- but it is not being
// served to anyone today, and triage that cannot tell the two apart orders its
// queue wrongly in both directions.
type servedState int
const (
// servedUnknown is the honest answer when the panel's domain map could not
// be read. It is not "not served".
servedUnknown servedState = iota
servedByPanel
notServed
)
// wpConfigPaths returns direct wp-config.php files at account document roots,
// each with whether the panel currently serves that root.
func wpConfigPaths(ctx context.Context) ([]string, map[string]servedState) {
paths, served, _ := wpConfigPathsWithDomains(ctx)
return paths, served
}
// wpConfigPathsWithDomains projects the shared install seam into the shapes
// CheckDatabaseContent works in. Discovery itself lives in wpinstalls.go, so
// this check, the object and overlap scanners, the core verifier and every
// fixer see the same installs.
func wpConfigPathsWithDomains(ctx context.Context) ([]string, map[string]servedState, map[string][]string) {
installs, panelDomains := wpInstallsWithDomains(ctx, "db_content")
paths := make([]string, 0, len(installs))
served := make(map[string]servedState, len(installs))
for _, in := range installs {
paths = append(paths, in.ConfigPath)
served[in.ConfigPath] = in.Served
}
return paths, served, panelDomains
}
// wpConfigOwners maps each discovered wp-config.php to the hosting account
// discovery attributed it to; unattributable installs are absent so their
// findings stay unstamped.
func wpConfigOwners(installs []wpInstall) map[string]string {
owners := make(map[string]string, len(installs))
for _, in := range installs {
if in.Account != "" {
owners[in.ConfigPath] = in.Account
}
}
return owners
}
// The shared vhost parser omits wildcard names because they cannot be used as
// an HTTP Host for exposure probes. They still declare a served document root
// and tenant ownership, so the database scan parses those rows separately.
func parseWildcardUserdataDomainRootsChecked(content string) ([]vhost, bool) {
var out []vhost
complete := true
for _, line := range strings.Split(content, "\n") {
line = strings.TrimSpace(line)
if !strings.HasPrefix(line, "*.") {
continue
}
parsed, lineComplete := parseUserdataDomainRootsChecked(strings.TrimPrefix(line, "*.") + "\n")
if !lineComplete || len(parsed) != 1 {
complete = false
continue
}
parsed[0].domain = "*." + parsed[0].domain
out = append(out, parsed[0])
}
return out, complete
}
func docrootBelongsToCPanelUser(root, user string) bool {
parts := strings.Split(filepath.Clean(root), string(filepath.Separator))
for i, part := range parts {
if !isCPanelHomeBase(part) || i+2 >= len(parts) || parts[i+1] != user {
continue
}
return true
}
return false
}
func isCPanelHomeBase(name string) bool {
if name == "home" {
return true
}
if !strings.HasPrefix(name, "home") || len(name) == len("home") {
return false
}
for _, r := range name[len("home"):] {
if r < '0' || r > '9' {
return false
}
}
return true
}
const maxWPSecondaryBlogs = 100
// dbSpamSampleLimit bounds the rows pulled back per spam pattern. When a
// pattern fills it, the reported count is a floor rather than a total.
const dbSpamSampleLimit = 200
func spamCountLabel(n int, truncated bool) string {
if truncated {
return fmt.Sprintf("at least %d", n)
}
return strconv.Itoa(n)
}
// dbScanCoverage counts why discovered installs could not be inspected and
// keeps one example path per reason. Multisite limits have their own detailed
// findings and are not counted again here.
//
// The owner remains incomplete when any install fails. Independently
// completed database scopes can still retire their own previous findings.
type dbScanCoverage struct {
discovered int
discoveryIncomplete bool
counts map[string]int
examples map[string]string
queryFailures map[string]int
queryFailureOverflow int
}
func (c *dbScanCoverage) record(reason, configPath string) {
if c == nil {
return
}
if c.counts == nil {
c.counts = make(map[string]int, 5)
c.examples = make(map[string]string, 5)
}
c.counts[reason]++
if c.examples[reason] == "" {
c.examples[reason] = configPath
}
}
func (c *dbScanCoverage) skipped() int {
if c == nil {
return 0
}
var n int
for _, v := range c.counts {
n += v
}
return n
}
// summary renders the reason breakdown, or the empty string when nothing was
// attributed, including when discovery stops before reaching any install.
func (c *dbScanCoverage) summary() string {
if c.skipped() == 0 {
return ""
}
reasons := make([]string, 0, len(c.counts))
for reason := range c.counts {
reasons = append(reasons, reason)
}
sort.Strings(reasons)
var b strings.Builder
fmt.Fprintf(&b, "%d of %d discovered installs could not be fully inspected.\n", c.skipped(), c.discovered)
for _, reason := range reasons {
// Account-controlled names must not forge reason lines or terminal
// commands. Bound the escaped display so expansion cannot grow it.
example := strconv.QuoteToASCII(c.examples[reason])
example = truncateDB(example[1:len(example)-1], 200)
fmt.Fprintf(&b, "%s=%d (example: %s)\n", reason, c.counts[reason], example)
}
b.WriteString(c.queryFailureSummary())
if c.discoveryIncomplete {
b.WriteString("Document-root discovery was incomplete; additional installs may be missing.\n")
}
return b.String()
}
// CheckDatabaseContent scans WordPress databases for injected malware,
// spam content, siteurl hijacking, and rogue admin accounts.
func CheckDatabaseContent(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
coverage := &dbScanCoverage{}
installs, panelDomains := wpInstallsWithDomains(ctx, "db_content")
coverage.discoveryIncomplete = checkMarkedIncomplete(ctx, "db_content")
if len(installs) == 0 {
return appendDatabaseScanIncompleteFinding(ctx, nil, coverage)
}
// cPanel locks an account's database users when it suspends the account,
// so every query against its installs fails with an access error. Such an
// install is dropped before coverage is counted: leaving it in made the
// scan permanently short of full coverage, with a warning no operator
// action could clear while the account stayed suspended.
wpConfigs := make([]string, 0, len(installs))
servedRoots := make(map[string]servedState, len(installs))
scannable := make([]wpInstall, 0, len(installs))
for _, in := range installs {
if accountSuspended(wpConfigUser(filepath.Dir(in.ConfigPath))) {
continue
}
scannable = append(scannable, in)
wpConfigs = append(wpConfigs, in.ConfigPath)
servedRoots[in.ConfigPath] = in.Served
}
installs = scannable
if len(installs) == 0 {
return appendDatabaseScanIncompleteFinding(ctx, nil, coverage)
}
coverage.discovered = len(installs)
owners := wpConfigOwners(installs)
domainOwnership := newPanelDomainOwnership(panelDomains)
// Cache the coverage outcome as well as the scan: aliases of an unreadable
// database are affected installs too, but must not repeat its queries.
seenDatabases := make(map[string]string, len(wpConfigs))
completedScopes := make(map[string]bool)
for _, wpConfig := range wpConfigs {
if ctx.Err() != nil {
return findings
}
user := wpConfigUser(filepath.Dir(wpConfig))
creds, complete := parseWPConfigChecked(wpConfig)
if !complete {
coverage.record("unreadable_config", wpConfig)
markCheckIncomplete(ctx, "db_content")
continue
}
if creds.dbName == "" || creds.dbUser == "" {
// A missing login does not erase an otherwise known scope. Its
// healthy alias must not retire findings this install could not
// examine, regardless of which config discovery returned first.
if prefix, ok := resolveTablePrefix(creds); creds.dbName != "" && ok {
completedScopes[dbContentDedupKey(user, creds, prefix)] = false
}
coverage.record("missing_credentials", wpConfig)
markCheckIncomplete(ctx, "db_content")
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
coverage.record("unresolved_table_prefix", wpConfig)
markCheckIncomplete(ctx, "db_content")
continue
}
creds.tablePrefix = prefix
creds.docrootServed = servedRoots[wpConfig]
creds.panelDomains = domainOwnership
databaseKey := strings.Join([]string{
user, creds.dbHost, creds.dbName, creds.dbUser, creds.dbPass, prefix,
strconv.FormatBool(creds.multisite),
}, "\x00")
if reason, duplicate := seenDatabases[databaseKey]; duplicate {
if reason != "" {
coverage.record(reason, wpConfig)
}
continue
}
// Isolate content-read gaps so an earlier install's failure cannot
// mask this one's. The scanner keeps the outer context for multisite
// limits, which already emit a separate account-specific finding.
queryCtx, contentIncomplete := withIncompleteCheckCollector(ctx)
creds.queryCtx = queryCtx
creds.queryState = new(dbQueryState)
// Stamp each install's own slice before merging: the host-wide
// summary appended below must never inherit an owner.
installFindings := capPhantomAuthorFindings(wpInstallScanner(ctx, user, creds, prefix), maxPhantomAuthorsReported)
scope := dbContentDedupKey(user, creds, prefix)
multisiteLimited := false
for i := range installFindings {
installFindings[i].CoverageScope = scope
if installFindings[i].Check == "db_content_scan_incomplete" {
multisiteLimited = true
}
}
findings = append(findings, stampTenantIDIfEmpty(installFindings, owners[wpConfig])...)
var reason string
if creds.queryState.failed {
reason = "query_failed"
} else if contentIncomplete.contains("db_content") {
reason = "incomplete_content"
}
scopeComplete := reason == "" && !multisiteLimited
if prior, seen := completedScopes[scope]; seen {
scopeComplete = scopeComplete && prior
}
completedScopes[scope] = scopeComplete
coverage.recordQueryFailures(creds.queryState)
seenDatabases[databaseKey] = reason
if reason != "" {
coverage.record(reason, wpConfig)
markCheckIncomplete(ctx, "db_content")
}
}
recordCompletedCoverageScopes(ctx, "db_content", completedScopes)
return appendDatabaseScanIncompleteFinding(ctx, findings, coverage)
}
// wpInstallScanner is the per-install scan boundary. Tests replace it with
// an inert scanner to prove ownership stamping for every finding name
// without driving each SQL scanner.
var wpInstallScanner = scanWPInstall
// scanWPInstall runs every content, user and multisite scanner for one
// discovered install and returns the unstamped findings.
func scanWPInstall(ctx context.Context, user string, creds wpDBCreds, prefix string) []alert.Finding {
var installFindings []alert.Finding
// Always scan the main-site (or single-site) tables. In
// multisite, blog ID 1 keeps the unprefixed names; in a
// single-site install these are the only tables.
installFindings = append(installFindings, scanWPBlog(user, creds, prefix, prefix)...)
// wp_users / wp_usermeta are network-wide in multisite, so
// the user-table scan runs once regardless of the layout.
installFindings = append(installFindings, checkWPUsers(user, creds.withQueryStage("users"), prefix)...)
// Multisite: enumerate active secondary blog IDs and scan
// each one's wp_<N>_options / wp_<N>_posts. Spam, archived,
// and deleted blogs are excluded -- their content is
// already operator-suppressed at the WP level, and most
// hosts have stale ones we'd otherwise alert on
// indefinitely.
if creds.multisite {
installFindings = append(installFindings, scanMultisiteSecondaryBlogs(ctx, user, creds, prefix)...)
}
return installFindings
}
// scanWPBlog runs checks whose tables belong to one blog. usersPrefix stays
// separate because multisite blogs share the network-wide users table.
func scanWPBlog(user string, creds wpDBCreds, sitePrefix, usersPrefix string) []alert.Finding {
var findings []alert.Finding
findings = append(findings, checkWPOptions(user, creds.withQueryStage("options"), sitePrefix)...)
findings = append(findings, checkWPPosts(user, creds.withQueryStage("posts"), sitePrefix)...)
findings = append(findings, checkWPStoredCode(user, creds.withQueryStage("stored_code"), sitePrefix)...)
findings = append(findings, checkWPSpamTaxonomy(user, creds.withQueryStage("taxonomy"), sitePrefix)...)
findings = append(findings, checkWPHiddenLinks(user, creds.withQueryStage("hidden_links"), sitePrefix)...)
findings = append(findings, checkWPCloakConfig(user, creds.withQueryStage("cloak_config"), sitePrefix)...)
// Rate change rather than vocabulary: the next kit will use different
// words, but it will still publish a flood onto a long-quiet site.
findings = append(findings, checkWPPostVolumeBurst(user, creds.withQueryStage("post_burst"), sitePrefix)...)
findings = append(findings,
checkWPPhantomAuthors(user, creds.withQueryStage("phantom_authors"), sitePrefix, usersPrefix, maxPhantomAuthorsReported)...)
return findings
}
func wpConfigUser(path string) string {
parts := strings.Split(filepath.Clean(path), string(filepath.Separator))
for i, part := range parts {
if isCPanelHomeBase(part) && i+1 < len(parts) {
return parts[i+1]
}
}
return extractUser(path)
}
// dbContentHostCoverageDedupKey is the identity of the host-wide coverage
// summary. Its Details carry per-reason counts and examples that shift between
// cycles while the same degradation persists. The multisite limit warning keeps
// its own per-install identity.
const dbContentHostCoverageDedupKey = "install_coverage"
func appendDatabaseScanIncompleteFinding(ctx context.Context, findings []alert.Finding, coverage *dbScanCoverage) []alert.Finding {
if !checkMarkedIncomplete(ctx, "db_content") {
return findings
}
// A multisite-limit warning covers only its own network. Suppress the
// generic fallback only when no other coverage gaps need reporting.
if coverage.skipped() == 0 && !coverage.discoveryIncomplete {
for _, finding := range findings {
if finding.Check == "db_content_scan_incomplete" {
return findings
}
}
}
return append(findings, alert.Finding{
Severity: alert.Warning,
Check: "db_content_scan_incomplete",
Message: "WordPress database scan could not inspect every discovered install",
Details: databaseScanIncompleteDetails(coverage),
DedupKey: dbContentHostCoverageDedupKey,
})
}
// databaseScanIncompleteDetails names what was skipped and why when the scan
// got far enough to attribute a cause, and falls back to the generic sentence
// when no install-specific cause was recorded.
func databaseScanIncompleteDetails(coverage *dbScanCoverage) string {
const retained = "Findings without complete database coverage are retained."
if summary := coverage.summary(); summary != "" {
return summary + retained
}
return "A document-root record, wp-config.php file, or database query could not be read safely. " + retained
}
// scanMultisiteSecondaryBlogs queries wp_blogs for active blog IDs other than 1
// and runs the per-blog scans against each. The user-table scan does not
// iterate because WP shares wp_users / wp_usermeta across the entire network
// by default. A site-specific user table only exists on configurations that
// override that, which we ignore here for v1. The phantom-author scan does
// iterate each posts table, but joins it to that shared users table.
//
// blog_id=1 is excluded because its tables are unprefixed and were
// already scanned by the caller.
func scanMultisiteSecondaryBlogs(ctx context.Context, user string, creds wpDBCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"SELECT blog_id FROM %sblogs WHERE archived = 0 AND deleted = 0 AND spam = 0 AND blog_id != 1 ORDER BY blog_id LIMIT %d",
prefix, maxWPSecondaryBlogs+1,
)
rows := runMySQLQuery(creds.withQueryStage("multisite_discovery"), query)
var findings []alert.Finding
truncated := len(rows) > maxWPSecondaryBlogs
if truncated {
rows = rows[:maxWPSecondaryBlogs]
}
for _, row := range rows {
if ctx.Err() != nil {
return findings
}
blogID := strings.TrimSpace(row)
if blogID == "" || blogID == "1" {
continue
}
// Guard against any garbage in the row -- only digits.
if !isAllDigits(blogID) {
markCheckIncomplete(creds.queryCtx, "db_content")
markCheckIncomplete(ctx, "db_content")
continue
}
sitePrefix := fmt.Sprintf("%s%s_", prefix, blogID)
findings = append(findings, scanWPBlog(user, creds, sitePrefix, prefix)...)
}
if truncated {
markCheckIncomplete(ctx, "db_content")
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "db_content_scan_incomplete",
Message: fmt.Sprintf("WordPress multisite database scan reached its %d-site safety limit (account: %s)", maxWPSecondaryBlogs, user),
Details: dbContentFindingDetails(creds, prefix,
"The network has more active secondary sites than one scheduled scan can safely inspect."),
DedupKey: dbContentDedupKey(user, creds, prefix,
"The network has more active secondary sites than one scheduled scan can safely inspect."),
})
}
return findings
}
func isAllDigits(s string) bool {
if s == "" {
return false
}
for _, r := range s {
if r < '0' || r > '9' {
return false
}
}
return true
}
type wpDBCreds struct {
dbName string
dbUser string
dbPass string
dbHost string
tablePrefix string
// docrootServed records whether the panel serves this install's document
// root, so a finding says whether it is reachable today.
docrootServed servedState
// panelDomains is the complete panel domain ownership map. The foreign-host
// check needs every account, not just this one, so a more-specific domain
// delegated to another tenant wins over this account's parent domain.
panelDomains *panelDomainOwnership
// queryCtx ties scheduled database work to the runner's deadline. Command
// paths leave it nil and retain the per-query timeout below.
queryCtx context.Context
// queryOwner identifies the CMS check whose coverage depends on a query.
// WordPress callers use the default owner when this is empty.
queryOwner string
// queryState is shared by the sequential queries for one install.
// Coverage failures and an unusable connection are tracked separately.
queryState *dbQueryState
queryStage string
// multisite is set when wp-config.php declares
// `define('MULTISITE', true)`. In multisite, the main blog
// (ID 1) keeps the unprefixed table names and secondary blogs
// live under `wp_<N>_options` / `wp_<N>_posts`. CheckDatabaseContent
// scans both layouts when this is set; a single-site install
// (multisite=false) skips the wp_blogs lookup and per-site
// iteration entirely.
multisite bool
}
// parseWPConfig extracts database credentials from wp-config.php.
func parseWPConfig(path string) wpDBCreds {
creds, complete := parseWPConfigChecked(path)
if !complete {
return wpDBCreds{}
}
return creds
}
// parseWPConfigChecked bounds account-controlled input so a special or very
// large wp-config.php cannot strand the scheduled database scan.
func parseWPConfigChecked(path string) (wpDBCreds, bool) {
f, err := openCMSConfig(path)
if err != nil {
return wpDBCreds{}, false
}
defer func() { _ = f.Close() }()
var creds wpDBCreds
limited := &io.LimitedReader{R: f, N: maxCMSConfigBytes + 1}
scanner := bufio.NewScanner(limited)
scanner.Buffer(make([]byte, 64*1024), maxCMSConfigBytes+1)
for scanner.Scan() {
line := scanner.Text()
// Match: define( 'DB_NAME', 'value' );
if val := extractDefine(line, "DB_NAME"); val != "" {
creds.dbName = val
}
if val := extractDefine(line, "DB_USER"); val != "" {
creds.dbUser = val
}
if val := extractDefine(line, "DB_PASSWORD"); val != "" {
creds.dbPass = val
}
if val := extractDefine(line, "DB_HOST"); val != "" {
creds.dbHost = val
}
// Match: $table_prefix = 'wp_';
if strings.Contains(line, "$table_prefix") {
if val := extractPHPString(line); val != "" {
creds.tablePrefix = val
}
}
// Match: define( 'MULTISITE', true );
if extractDefineBool(line, "MULTISITE") {
creds.multisite = true
}
}
if creds.dbHost == "" {
creds.dbHost = "localhost"
}
return creds, scanner.Err() == nil && limited.N > 0
}
// extractDefine extracts the value from: define( 'KEY', 'value' );
func extractDefine(line, key string) string {
if !strings.Contains(line, key) {
return ""
}
// Skip comments
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "//") || strings.HasPrefix(trimmed, "#") || strings.HasPrefix(trimmed, "/*") {
return ""
}
// After the literal key, step past the first comma so
// extractPHPString picks up the VALUE's opening quote rather than
// the KEY's trailing closing quote. Without this, on input
// define( 'DB_NAME', 'wordpress_db' );
// extractPHPString would see `', 'wordpress_db' );` and return
// `, ` — the substring between the closing quote of 'DB_NAME' and
// the opening quote of 'wordpress_db'. Every real WordPress
// install's wp-config.php triggered this, which silently broke
// the entire WP database scan check.
rest := line[strings.Index(line, key)+len(key):]
if commaIdx := strings.Index(rest, ","); commaIdx >= 0 {
rest = rest[commaIdx+1:]
}
return extractPHPString(rest)
}
// extractDefineBool returns true when line is a non-comment
// define('<key>', true) -- i.e., a bare boolean value rather than a
// quoted string. Used for `MULTISITE` and any future bool defines
// CSM cares about. Whitespace is permissive, case-insensitive on
// the literal `true`, trailing PHP/inline comments tolerated.
//
// The key must appear inside its enclosing PHP quotes (single or
// double). This avoids a substring-match collision: a WordPress
// wp-config.php commonly carries `define('WP_ALLOW_MULTISITE',
// true)` to enable the admin network creator on single-site
// installs; matching MULTISITE as a bare substring would falsely
// detect those as multisite hosts.
//
// Operators using anything other than the canonical `true` literal
// (e.g., `!false`, `1`, `defined('FOO')`) won't get multisite
// scanning. That's preferable to running an arbitrary PHP expression
// evaluator over wp-config.php.
func extractDefineBool(line, key string) bool {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "//") || strings.HasPrefix(trimmed, "#") || strings.HasPrefix(trimmed, "/*") {
return false
}
if !strings.Contains(trimmed, "define") {
return false
}
// Find the key's quoted form and seek past the closing quote.
// Two acceptable openings: 'KEY' and "KEY". The key never
// appears unquoted inside a define() literal in valid PHP.
var keyEnd int
switch {
case strings.Contains(trimmed, "'"+key+"'"):
keyEnd = strings.Index(trimmed, "'"+key+"'") + len(key) + 2
case strings.Contains(trimmed, `"`+key+`"`):
keyEnd = strings.Index(trimmed, `"`+key+`"`) + len(key) + 2
default:
return false
}
rest := trimmed[keyEnd:]
commaIdx := strings.Index(rest, ",")
if commaIdx < 0 {
return false
}
value := rest[commaIdx+1:]
// The value runs until the closing paren; everything after it
// is the statement terminator and any trailing comment.
if closeIdx := strings.Index(value, ")"); closeIdx >= 0 {
value = value[:closeIdx]
}
return strings.EqualFold(strings.TrimSpace(value), "true")
}
// extractPHPString extracts the first quoted string value from a line.
func extractPHPString(s string) string {
// Find opening quote
for _, quote := range []byte{'\'', '"'} {
start := strings.IndexByte(s, quote)
if start < 0 {
continue
}
rest := s[start+1:]
end := strings.IndexByte(rest, quote)
if end < 0 {
continue
}
return rest[:end]
}
return ""
}
// runMySQLQuery executes a MySQL query via the in-process database/sql
// driver and returns each row tab-joined, matching the legacy
// `mysql -N -B -e <query>` output shape so existing tab-split callers
// keep working unchanged. Returns nil on any open / query / scan
// error (the legacy implementation swallowed errors the same way).
// Var so tests can serve canned rows without a live database.
var runMySQLQuery = func(creds wpDBCreds, query string) []string {
if creds.queryState != nil && creds.queryState.halted {
return nil
}
parent := creds.queryCtx
if parent == nil {
parent = context.Background()
}
ctx, cancel := context.WithTimeout(parent, 2*time.Minute)
defer cancel()
rows, err := mysqlclient.PerAccountQuery(ctx, mysqlclient.Creds{
User: creds.dbUser,
Password: creds.dbPass,
Host: creds.dbHost,
DBName: creds.dbName,
}, query)
if err != nil {
creds.queryState.record(creds.queryStage, err)
markCheckIncomplete(creds.queryCtx, creds.queryCheck())
return nil
}
out := make([]string, 0, len(rows))
for _, line := range rows {
line = strings.TrimSpace(line)
if line != "" {
out = append(out, line)
}
}
if len(out) == 0 {
return nil
}
return out
}
func (c wpDBCreds) queryCheck() string {
if c.queryOwner != "" {
return c.queryOwner
}
return "db_content"
}
// siteURLPoisonReason reports why a siteurl/home value cannot be a real site
// address, and whether it is one at all. WordPress concatenates this value to
// build every asset URL it emits, so an attacker who rewrites it makes the
// address it names load on every page without touching a single file.
//
// The test is the URL's shape, not its host. Hosting a site under a domain
// that is not served locally is ordinary -- sites move, staging lives
// elsewhere, and a theme demo keeps its vendor address -- so a host check
// would report dozens of healthy installs. Shape does not have that problem:
// a site address is an origin plus an optional subdirectory. A backslash, a
// query string, a fragment, a script for a path, or a non-web scheme cannot
// appear in one.
func siteURLPoisonReason(value string) (string, bool) {
// MySQL query rows preserve batch-mode escaping. Parse the stored bytes,
// not the escaped transport form, or control characters can look like an
// ordinary path and evade the shape checks below.
value = strings.TrimSpace(mysqlclient.BatchUnescape(value))
if value == "" {
return "", false
}
if strings.ContainsRune(value, '\\') {
return "address carries a backslash", true
}
u, err := url.Parse(value)
if err != nil {
return "value is not a URL", true
}
if scheme := strings.ToLower(u.Scheme); scheme != "http" && scheme != "https" {
return "scheme is not http or https", true
}
if u.Hostname() == "" {
return "no host", true
}
if port := u.Port(); port != "" {
n, err := strconv.Atoi(port)
if err != nil || n < 1 || n > 65535 {
return "port is outside the valid range", true
}
}
if u.RawQuery != "" || strings.Contains(value, "?") {
return "address carries a query string", true
}
if u.Fragment != "" || strings.Contains(value, "#") {
return "address carries a fragment", true
}
if isScriptPath(u.Path) {
return "address resolves to a script", true
}
return "", false
}
// isScriptPath reports whether a URL path's final segment names a client- or
// server-side script. A site address always ends at a directory.
func isScriptPath(path string) bool {
segment := path
if i := strings.LastIndex(segment, "/"); i >= 0 {
segment = segment[i+1:]
}
segment = strings.ToLower(segment)
if isExecutablePHPName(segment) {
return true
}
switch filepath.Ext(segment) {
case ".js", ".mjs", ".cjs",
".asp", ".aspx", ".ashx", ".asmx",
".jsp", ".jspx", ".cfm",
".cgi", ".pl", ".py", ".rb":
return true
default:
return false
}
}
// checkWPOptions checks for siteurl/home hijacking and injected JavaScript.
func checkWPOptions(user string, creds wpDBCreds, prefix string) []alert.Finding {
var findings []alert.Finding
foreignByOption := make(map[string]*alert.Finding, 2)
// Check siteurl and home for hijacking
query := fmt.Sprintf(
"SELECT option_name, option_value FROM %soptions WHERE option_name IN ('siteurl', 'home', 'admin_email') LIMIT 10",
prefix)
lines := runMySQLQuery(creds, query)
for _, line := range lines {
parts := strings.SplitN(line, "\t", 2)
if len(parts) != 2 {
continue
}
// WordPress's default option_name collation is case-insensitive, so a
// differently-cased row can satisfy get_option("siteurl") and this SQL
// query. Keep the Go-side security check consistent with that lookup.
optName := strings.ToLower(strings.TrimSpace(parts[0]))
optValue := strings.ToLower(parts[1])
if optName == "siteurl" || optName == "home" {
if strings.Contains(optValue, "eval(") || strings.Contains(optValue, "<script") {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "db_siteurl_hijack",
Message: fmt.Sprintf("WordPress %s contains malicious code (account: %s)", optName, user),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("%s = %s", optName, truncateDB(parts[1], 200))),
DedupKey: dbContentDedupKey(user, creds, prefix,
fmt.Sprintf("%s = %s", optName, truncateDB(parts[1], 200))),
})
} else if reason, bad := siteURLPoisonReason(parts[1]); bad {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "db_siteurl_invalid",
Message: fmt.Sprintf("WordPress %s is not a site address (account: %s): %s", optName, user, reason),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("%s = %s\nWordPress builds every asset URL from this value, so the address it names is loaded on every page.",
optName, truncateDB(parts[1], 200))),
DedupKey: dbContentDedupKey(user, creds, prefix,
"reason="+reason,
fmt.Sprintf("%s = %s\nWordPress builds every asset URL from this value, so the address it names is loaded on every page.",
optName, truncateDB(parts[1], 200))),
})
} else if foreign := foreignSiteURLFinding(user, creds, prefix, optName, parts[1]); foreign != nil {
// siteurl and home commonly hold the same address. Emit one stable
// condition per blog, preferring siteurl regardless of row order.
if current := foreignByOption[optName]; current == nil || foreign.Details < current.Details {
foreignByOption[optName] = foreign
}
}
}
}
for _, option := range []string{"siteurl", "home"} {
if foreign := foreignByOption[option]; foreign != nil {
findings = append(findings, *foreign)
break
}
}
// Path 1: External script URLs in any option — only flag non-safe domains.
query = fmt.Sprintf(
"SELECT option_name, option_value FROM %soptions WHERE option_value LIKE '%%<script%%src%%' LIMIT 20",
prefix)
lines = runMySQLQuery(creds, query)
firstSeen := storeFirstSeen(externalScriptSiteKey(creds.dbName, prefix))
for _, line := range lines {
parts := strings.SplitN(line, "\t", 2)
if len(parts) != 2 {
continue
}
optName := parts[0]
optValue := parts[1]
// Skip CSM backup options — they preserve the original malicious
// content for recovery and should not be re-detected/re-cleaned.
if strings.HasPrefix(optName, "csm_backup_") {
continue
}
maliciousURL := extractMaliciousScriptURL(optValue)
if maliciousURL == "" {
// No attacker marker. A loader on an unremarkable HTTPS host is
// still reported once, the first time it appears after the
// site's baseline.
findings = append(findings, newExternalScriptFindings(user, creds, prefix, optName, optValue, firstSeen)...)
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "db_options_injection",
Message: fmt.Sprintf("Malicious script injection in wp_options '%s' (account: %s)", optName, user),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Option: %s", optName),
fmt.Sprintf("Malicious URL: %s", maliciousURL),
fmt.Sprintf("Content preview: %s", truncateDB(optValue, 200))),
DedupKey: dbContentDedupKey(user, creds, prefix,
fmt.Sprintf("Option: %s", optName),
fmt.Sprintf("Malicious URL: %s", maliciousURL),
fmt.Sprintf("Content preview: %s", truncateDB(optValue, 200))),
})
}
queryComplete := creds.queryState == nil || !creds.queryState.failed
if queryComplete {
if sdb := store.Global(); sdb != nil {
_ = sdb.FinishExternalScriptBaseline(externalScriptSiteKey(creds.dbName, prefix), time.Now())
}
}
// Path 1b: Plugin status options that WordPress renders as admin
// notices. These are queried by name because the generic script lookup
// above caps its result set and requires a src attribute, while an
// injection here may be inline. The option's identity is the verdict,
// so neither host reputation nor the first-seen baseline applies.
findings = append(findings, checkWPPluginNotices(user, creds, prefix)...)
// Path 2: Inline script/code injection in core WP options that should
// NEVER contain JavaScript (siteurl, home, blogname, blogdescription).
coreOpts := "siteurl', 'home', 'blogname', 'blogdescription', 'admin_email"
codePatterns := "<script"
query = fmt.Sprintf(
"SELECT option_name, LEFT(option_value, 500) FROM %soptions WHERE option_name IN ('%s') AND option_value LIKE '%%%s%%'",
prefix, coreOpts, codePatterns)
lines = runMySQLQuery(creds, query)
for _, line := range lines {
parts := strings.SplitN(line, "\t", 2)
if len(parts) != 2 {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "db_options_injection",
Message: fmt.Sprintf("Malicious content in core wp_option '%s' (account: %s)", parts[0], user),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Option: %s", parts[0]),
fmt.Sprintf("Content preview: %s", truncateDB(parts[1], 200))),
DedupKey: dbContentDedupKey(user, creds, prefix,
fmt.Sprintf("Option: %s", parts[0]),
fmt.Sprintf("Content preview: %s", truncateDB(parts[1], 200))),
})
}
return findings
}
// checkWPPosts checks post content for injected scripts and malware.
//
// Two classes of false positive are suppressed compared to a naive LIKE-
// based scan:
//
// - post_types used for plugin-managed storage (form submissions,
// revisions, templates, minified bundles) are excluded via the
// shared nonScannablePostTypes denylist. See dbscan_filters.go for
// the rationale and the full list.
//
// - Patterns that match too broadly at the SQL layer (the bare
// <script substring, and bare-word spam keywords like "cialis")
// are post-filtered in Go against word-boundary regexes and the
// known-safe-domain list. Legitimate analytics embeds and
// substring coincidences ("specialist" containing "cialis") no
// longer produce findings.
//
// The denylist is defense-in-depth: custom post_types created by a
// theme or plugin remain in scope, so attackers cannot evade by
// inventing a new post_type value.
func checkWPPosts(user string, creds wpDBCreds, prefix string) []alert.Finding {
var findings []alert.Finding
postTypeExcl := nonScannablePostTypesSQLList()
// Keep each pattern's independent LIMIT, but send the bounded selects as
// one UNION. On a host with hundreds of installs this avoids one database
// connection and round trip per signature without letting a noisy pattern
// consume every candidate slot for the others.
malwareSelects := make([]string, 0, len(dbMalwarePatterns))
for i, mp := range dbMalwarePatterns {
pattern := mysqlEscapeForLike(mp.pattern)
selectedContent := "'_content_not_required'"
if mp.requiresExternalScript {
selectedContent = "CONCAT_WS(CHAR(10), post_content, post_content_filtered)"
}
malwareSelects = append(malwareSelects, fmt.Sprintf(
"(SELECT %d AS pattern_index, ID, %s FROM %sposts WHERE post_status='publish' AND post_type NOT IN (%s) AND (post_content LIKE '%%%s%%' OR post_content_filtered LIKE '%%%s%%') LIMIT 20)",
i, selectedContent, prefix, postTypeExcl, pattern, pattern))
}
malwareRows := runMySQLQuery(creds, strings.Join(malwareSelects, " UNION ALL "))
confirmedByPattern := make([][]string, len(dbMalwarePatterns))
seenByPattern := make([]map[string]struct{}, len(dbMalwarePatterns))
for _, row := range malwareRows {
parts := strings.SplitN(row, "\t", 3)
if len(parts) != 3 {
continue
}
patternIndex, err := strconv.Atoi(parts[0])
if err != nil || patternIndex < 0 || patternIndex >= len(dbMalwarePatterns) {
continue
}
mp := dbMalwarePatterns[patternIndex]
content := mysqlclient.BatchUnescape(parts[2])
if mp.requiresExternalScript && !hasMaliciousExternalScriptInPost(content) {
continue
}
if seenByPattern[patternIndex] == nil {
seenByPattern[patternIndex] = make(map[string]struct{})
}
postID := strings.TrimSpace(parts[1])
if postID == "" {
continue
}
if _, duplicate := seenByPattern[patternIndex][postID]; duplicate {
continue
}
seenByPattern[patternIndex][postID] = struct{}{}
if len(confirmedByPattern[patternIndex]) < 5 {
confirmedByPattern[patternIndex] = append(confirmedByPattern[patternIndex], postID)
}
}
for i, confirmedIDs := range confirmedByPattern {
if len(confirmedIDs) == 0 {
continue
}
mp := dbMalwarePatterns[i]
findings = append(findings, alert.Finding{
Severity: mp.severity,
Check: "db_post_injection",
Message: fmt.Sprintf("WordPress posts contain %s (account: %s, %d posts)", mp.desc, user, len(confirmedIDs)),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Affected post IDs: %s", strings.Join(confirmedIDs, ", ")),
fmt.Sprintf("Pattern: %s", mp.pattern)),
DedupKey: dbContentDedupKey(user, creds, prefix,
fmt.Sprintf("Affected post IDs: %s", strings.Join(confirmedIDs, ", ")),
fmt.Sprintf("Pattern: %s", mp.pattern)),
})
}
// Spam keyword scan. Three-layer filter:
//
// 1. SQL LIKE as a fast server-side pre-filter (reduces rows).
// 2. Word-boundary regex in countCloakedSpamMatches (rejects
// substring false positives like "specialist" / "cialis").
// 3. SEO-context requirement in contentHasSpamContext: a keyword
// hit only counts when accompanied by CSS cloaking, an
// injection fingerprint, or an external anchor whose URL
// path contains the keyword. Bare prose mentions (industry
// verticals, advisor bios, product catalogs listing a
// pharmaceutical supply chain) do not fire.
//
// The context requirement catches the real attack pattern — hidden
// off-screen div with external commercial link — while leaving
// legitimate content silent. See spam_context.go for the full
// signal catalog.
spamSelects := make([]string, 0, len(dbSpamPatterns))
for i, sp := range dbSpamPatterns {
spamSelects = append(spamSelects, fmt.Sprintf(
"(SELECT %d AS pattern_index, ID, post_content FROM %sposts WHERE post_status='publish' AND post_type NOT IN (%s) AND post_content LIKE '%s' LIMIT %d)",
i, prefix, postTypeExcl, mysqlEscapeForLike(sp.likeFragment), dbSpamSampleLimit))
}
spamRows := runMySQLQuery(creds, strings.Join(spamSelects, " UNION ALL "))
spamContents := make([][]string, len(dbSpamPatterns))
spamSampled := make([]int, len(dbSpamPatterns))
for _, row := range spamRows {
parts := strings.SplitN(row, "\t", 3)
patternIndex, err := strconv.Atoi(parts[0])
if err != nil || patternIndex < 0 || patternIndex >= len(dbSpamPatterns) {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
spamSampled[patternIndex]++
if len(parts) != 3 {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
spamContents[patternIndex] = append(spamContents[patternIndex], mysqlclient.BatchUnescape(parts[2]))
}
for i, sp := range dbSpamPatterns {
n := countCloakedSpamMatches(sp, spamContents[i])
if n == 0 {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "db_spam_injection",
// The per-pattern LIMIT bounds the sample, so a full sample means
// the real figure may be larger. Reporting it as exact understates the
// scale, and scale is what decides whether an operator looks.
Message: fmt.Sprintf("WordPress posts contain cloaked spam keyword '%s' (%s posts, account: %s)",
sp.keyword, spamCountLabel(n, spamSampled[i] >= dbSpamSampleLimit), user),
Details: dbContentFindingDetails(creds, prefix),
// The pattern is the identity; the count is not. Spam grows between
// scans, and that is the same finding, not a new one.
DedupKey: dbContentDedupKey(user, creds, prefix, "keyword="+sp.keyword),
})
}
return findings
}
// dbContentDedupKey pins a database-content finding's identity to the database
// it was found in, the account using it, and what was found there. Callers
// include stable distinctions from Message as well as Details, excluding
// observation-only changes such as site age and document-root served state.
//
// The document-root note is deliberately excluded. It reports what the panel's
// domain map said during this scan, not anything the scan found in the
// database, and that map read fails transiently -- when it does the note
// disappears, the default Message+Details identity changes with it, and the
// store keeps a second copy of a finding that never changed.
func dbContentDedupKey(user string, creds wpDBCreds, prefix string, lines ...string) string {
identity := make([]byte, 0, 128)
appendField := func(value string) {
identity = binary.BigEndian.AppendUint64(identity, uint64(len(value)))
identity = append(identity, value...)
}
appendField(user)
appendField(creds.dbHost)
appendField(creds.dbName)
appendField(prefix)
for _, line := range lines {
appendField(line)
}
digest := sha256.Sum256(identity)
return fmt.Sprintf("db-content:%x", digest[:12])
}
// dbContentFindingDetails renders a database-content finding's details. The
// document-root note it adds is scan-time context, not part of what was found,
// so every caller supplies a DedupKey that excludes this note.
func dbContentFindingDetails(creds wpDBCreds, prefix string, lines ...string) string {
out := []string{
fmt.Sprintf("Database: %s", creds.dbName),
fmt.Sprintf("Table prefix: %s", prefix),
}
if note := docrootServedNote(creds.docrootServed); note != "" {
out = append(out, note)
}
out = append(out, lines...)
return strings.Join(out, "\n")
}
// wpInstallAdminGrace is how far after the first user registration an
// admin account still counts as created by the site install itself.
// Softaculous and similar installers register the initial admin (and any
// setup-wizard co-admins) within moments of creating the users table, so
// those accounts are the install, not a takeover. Anything later is a
// change to an existing site and stays alert-worthy.
const wpInstallAdminGrace = 15 * time.Minute
const wpRegisteredLayout = "2006-01-02 15:04:05"
func parseWPRegistered(value string) (time.Time, error) {
return time.Parse(wpRegisteredLayout, strings.TrimSpace(value))
}
// wpInstallEraAdmin reports whether an admin registration timestamp falls
// within the install grace of the site's first user registration. Fails
// open: unparseable timestamps never suppress.
func wpInstallEraAdmin(registered, firstRegistered string) bool {
reg, err := parseWPRegistered(registered)
if err != nil {
return false
}
first, err := parseWPRegistered(firstRegistered)
if err != nil {
return false
}
diff := reg.Sub(first)
return diff >= 0 && diff <= wpInstallAdminGrace
}
// checkWPUsers checks for rogue admin accounts created recently.
func checkWPUsers(user string, creds wpDBCreds, prefix string) []alert.Finding {
var findings []alert.Finding
// Find admin users created in the last 7 days. Missing registration
// timestamps stay in scope so the suppression fails open.
// MIN(user_registered) rides along as the install marker, but any NULL
// invalidates that marker. EXISTS keeps duplicate capability metadata
// from consuming the bounded result.
query := fmt.Sprintf(
"SELECT u.ID, u.user_login, u.user_email, u.user_registered, "+
"(SELECT CASE WHEN COUNT(*) <> COUNT(user_registered) "+
"THEN NULL ELSE MIN(user_registered) END FROM %susers) FROM %susers u "+
"WHERE EXISTS (SELECT 1 FROM %susermeta m "+
"WHERE m.user_id = u.ID AND m.meta_key = '%scapabilities' "+
"AND m.meta_value LIKE '%%administrator%%') "+
"AND (u.user_registered >= DATE_SUB(NOW(), INTERVAL 7 DAY) "+
"OR u.user_registered IS NULL "+
"OR CAST(u.user_registered AS CHAR) = '0000-00-00 00:00:00') "+
"ORDER BY (u.user_registered IS NULL OR "+
"CAST(u.user_registered AS CHAR) = '0000-00-00 00:00:00') DESC, "+
"u.user_registered DESC "+
"LIMIT 10",
prefix, prefix, prefix, prefix)
lines := runMySQLQuery(creds, query)
for _, line := range lines {
parts := strings.SplitN(line, "\t", 5)
if len(parts) < 3 {
continue
}
if wpInstallEraAdmin(safeGet(parts, 3), safeGet(parts, 4)) {
continue
}
registered := safeGet(parts, 3)
message := fmt.Sprintf("New WordPress admin account created in last 7 days: %s (account: %s)", parts[1], user)
if _, err := parseWPRegistered(registered); err != nil {
message = fmt.Sprintf("WordPress admin account has a missing or invalid registration timestamp: %s (account: %s)", parts[1], user)
}
severity := alert.Critical
details := fmt.Sprintf("Database: %s\nTable prefix: %s\nUser ID: %s\nLogin: %s\nEmail: %s\nRegistered: %s",
creds.dbName, prefix, parts[0], parts[1], parts[2], registered)
// A stable session pattern can support the legitimate developer or agency
// case, but session metadata is not authoritative. Downgrade -- never
// suppress -- so a forged session record cannot hide the account.
loginIPs := wpAdminLoginIPs(creds, prefix, parts[0])
distinct := uniqueStrings(loginIPs)
switch {
case len(loginIPs) == 0:
details += "\nSession-token IP evidence: none recorded."
case len(loginIPs) >= wpEstablishedLoginSessions && len(distinct) == 1:
severity = alert.Warning
details += fmt.Sprintf("\nSession-token IP evidence: %d stored sessions from a single IP (%s) -- consistent with a stable operator login pattern, but not proof the account is legitimate. Verify with the account owner before acting.",
len(loginIPs), distinct[0])
default:
details += fmt.Sprintf("\nSession-token IP evidence: %d stored sessions from %d IP(s): %s",
len(loginIPs), len(distinct), strings.Join(firstN(distinct, 5), ", "))
}
findings = append(findings, alert.Finding{
Severity: severity,
Check: "db_rogue_admin",
Message: message,
Details: details,
})
}
// Check for admin users with suspicious email patterns
query = fmt.Sprintf(
"SELECT u.user_login, u.user_email FROM %susers u "+
"INNER JOIN %susermeta m ON u.ID = m.user_id "+
"WHERE m.meta_key = '%scapabilities' AND m.meta_value LIKE '%%administrator%%' "+
"LIMIT 50",
prefix, prefix, prefix)
lines = runMySQLQuery(creds, query)
for _, line := range lines {
parts := strings.SplitN(line, "\t", 2)
if len(parts) != 2 {
continue
}
email := strings.ToLower(parts[1])
// Flag suspicious admin emails (disposable/temporary email domains)
suspiciousDomains := []string{
"tempmail", "guerrillamail", "mailinator", "throwaway",
"yopmail", "sharklasers", "trashmail", "maildrop",
}
for _, sd := range suspiciousDomains {
if strings.Contains(email, sd) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "db_suspicious_admin_email",
Message: fmt.Sprintf("WordPress admin '%s' has disposable email (account: %s)", parts[0], user),
Details: fmt.Sprintf("Database: %s\nTable prefix: %s\nEmail: %s", creds.dbName, prefix, email),
})
break
}
}
}
return findings
}
func safeGet(parts []string, idx int) string {
if idx < len(parts) {
return parts[idx]
}
return ""
}
func truncateDB(s string, maxLen int) string {
if len(s) <= maxLen {
return s
}
return s[:maxLen] + "..."
}
// CleanDatabaseSpam removes known spam/malware patterns from WordPress database content.
// Targets wp_posts and wp_options tables. Returns findings for each cleaned row.
func CleanDatabaseSpam(account string) []alert.Finding {
var findings []alert.Finding
wpConfigs := spamCleanWPConfigs(account)
for _, wpConfig := range wpConfigs {
creds := parseWPConfig(wpConfig)
if creds.dbName == "" {
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
continue
}
creds.tablePrefix = prefix
// Clean spam from wp_posts
spamPatterns := []struct {
pattern string
desc string
}{
{"<script>", "injected script tag"},
{"eval(", "eval() in post content"},
{"base64_decode(", "base64_decode in post content"},
{"document.write(", "document.write injection"},
}
for _, sp := range spamPatterns {
// Count affected rows first
countQuery := fmt.Sprintf(
"SELECT COUNT(*) FROM %sposts WHERE post_content LIKE '%%%s%%'",
prefix, sp.pattern)
countLines := runMySQLQuery(creds, countQuery)
if len(countLines) == 0 || countLines[0] == "0" {
continue
}
// Clean: remove the malicious pattern from post_content
cleanQuery := fmt.Sprintf(
"UPDATE %sposts SET post_content = REPLACE(post_content, '%s', '') WHERE post_content LIKE '%%%s%%'",
prefix, sp.pattern, sp.pattern)
runMySQLQuery(creds, cleanQuery)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "db_spam_cleaned",
Message: fmt.Sprintf("Cleaned %s from %s posts in %s (account: %s)", sp.desc, countLines[0], creds.dbName, account),
Timestamp: time.Now(),
})
}
// Scan for spam keywords in wp_posts. Uses the same word-boundary
// regex + post_type denylist + SEO-context requirement as
// checkWPPosts so an operator-initiated cleanup surfaces the
// same set of findings the periodic scan does.
postTypeExcl := nonScannablePostTypesSQLList()
for _, sp := range dbSpamPatterns {
query := fmt.Sprintf(
"SELECT ID, post_content FROM %sposts WHERE post_status='publish' AND post_type NOT IN (%s) AND post_content LIKE '%s' LIMIT 200",
prefix, postTypeExcl, sp.likeFragment)
lines := runMySQLQuery(creds, query)
if len(lines) == 0 {
continue
}
contents := make([]string, 0, len(lines))
for _, line := range lines {
parts := strings.SplitN(line, "\t", 2)
if len(parts) < 2 {
continue
}
contents = append(contents, parts[1])
}
n := countCloakedSpamMatches(sp, contents)
if n == 0 {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "db_spam_found",
Message: fmt.Sprintf("Found spam keyword '%s' in %d published posts in %s (account: %s) - manual review recommended", sp.keyword, n, creds.dbName, account),
})
}
}
return findings
}
// wpEstablishedLoginSessions is the number of stored login sessions from a
// single stable IP that is strong enough to downgrade the finding for review.
const wpEstablishedLoginSessions = 5
// parseSessionTokenIPs extracts structurally valid login source IPs from a
// WordPress session_tokens usermeta blob in session order.
func parseSessionTokenIPs(serialized string) []string {
var ips []string
for i := 0; i < len(serialized); {
key, next, ok := parsePHPSerializedStringAt(serialized, i)
if !ok {
if serialized[i] == 's' && i+1 < len(serialized) && serialized[i+1] == ':' {
return nil
}
i++
continue
}
if key != "ip" {
i = next
continue
}
value, afterValue, ok := parsePHPSerializedStringAt(serialized, next)
if !ok {
return nil
}
addr, err := netip.ParseAddr(value)
if err == nil {
ips = append(ips, addr.Unmap().String())
}
i = afterValue
}
return ips
}
// parsePHPSerializedStringAt reads one PHP s:<length>:"<value>"; token. Using
// the declared byte length prevents a forged token-looking fragment inside a
// session's attacker-controlled user-agent string from being treated as a key.
func parsePHPSerializedStringAt(serialized string, start int) (string, int, bool) {
if start < 0 || start+2 > len(serialized) ||
serialized[start] != 's' || serialized[start+1] != ':' {
return "", start, false
}
lengthStart := start + 2
lengthEnd := lengthStart
for lengthEnd < len(serialized) &&
serialized[lengthEnd] >= '0' && serialized[lengthEnd] <= '9' {
lengthEnd++
}
if lengthEnd == lengthStart || lengthEnd+2 > len(serialized) ||
serialized[lengthEnd] != ':' || serialized[lengthEnd+1] != '"' {
return "", start, false
}
valueLen, err := strconv.Atoi(serialized[lengthStart:lengthEnd])
if err != nil {
return "", start, false
}
valueStart := lengthEnd + 2
if valueLen > len(serialized)-valueStart {
return "", start, false
}
valueEnd := valueStart + valueLen
if valueEnd+2 > len(serialized) ||
serialized[valueEnd] != '"' || serialized[valueEnd+1] != ';' {
return "", start, false
}
return serialized[valueStart:valueEnd], valueEnd + 2, true
}
// wpAdminLoginIPs returns source IPs encoded in an admin's stored login
// sessions. The caller treats this forgeable metadata only as downgrade
// evidence and never uses it to suppress the finding.
func wpAdminLoginIPs(creds wpDBCreds, prefix, userID string) []string {
if !wpNumericID(userID) {
return nil
}
q := fmt.Sprintf(
"SELECT meta_value FROM %susermeta WHERE user_id = %s AND meta_key = 'session_tokens' LIMIT 1",
prefix, userID)
serialized := mysqlclient.BatchUnescape(strings.Join(runMySQLQuery(creds, q), ""))
return parseSessionTokenIPs(serialized)
}
// wpNumericID guards the user id (a prior query row value) before it is
// interpolated into the session lookup.
func wpNumericID(s string) bool {
if s == "" {
return false
}
for _, r := range s {
if r < '0' || r > '9' {
return false
}
}
return true
}
// firstN caps a slice for display in finding details.
func firstN(in []string, n int) []string {
if len(in) <= n {
return in
}
return in[:n]
}
// docrootServedNote states whether this install is reachable today. Both
// answers change how a finding should be queued: a dormant install is not
// serving anyone right now, and a served one is. Silence when the panel's map
// could not be read -- claiming either would be a guess.
func docrootServedNote(state servedState) string {
switch state {
case servedByPanel:
return "Document root: served by the panel, so this is live now."
case notServed:
return "Document root: not currently served. The database is still live " +
"and the content publishes again the moment a domain is pointed here."
default:
return ""
}
}
// spamCleanWPConfigs lists the installs the spam cleaner acts on. Shared
// discovery: spam left in a nested or panel-mapped install is the same spam.
func spamCleanWPConfigs(account string) []string {
return wpInstallConfigPaths(wpInstallsForAccount(context.Background(), "db_content", account))
}
package checks
import (
"encoding/base64"
"encoding/hex"
"fmt"
"regexp"
"sort"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// Cloaking infrastructure kept in the options table.
//
// Two shapes, both from doorway kits found on one host:
//
// - Configuration stored under a digest of the site's own hostname, so the
// row is unfindable without knowing which site you are looking at, with a
// base64 layer over the serialized array so its contents do not appear in
// any search of the table.
// - One rewrite rule per doorway cluster, routing a numbered sitemap
// straight into a matching numbered feed, to hand crawlers the generated
// pages without them appearing in the site's real sitemap.
//
// Neither is code execution. Both are the scaffolding a doorway network needs,
// and both survive the deletion of every spam post.
const (
// maxCloakOptionBytes bounds one option value. Configuration blobs are
// normally small, but rewrite_rules can exceed this on a large site. The
// query also returns the original byte length so a bounded read can never
// be mistaken for a complete value.
maxCloakOptionBytes = 256 * 1024
// maxCloakOptionRows bounds the digest-named candidates read.
maxCloakOptionRows = 50
// maxCloakSamplesShown bounds what the finding names.
maxCloakSamplesShown = 10
)
// digestOptionName matches an option named by a 32-character hex digest and
// nothing else. WordPress core and plugins name options after what they hold.
var digestOptionName = regexp.MustCompile(`^[0-9a-f]{32}$`)
// numberedSitemapRoute and numberedSitemapFeed are the two halves of the
// doorway routing. Reporting needs both with the same number: real sitemap
// plugins add rewrite rules too, but none of them route sitemap<N>.xml into a
// feed named xmlsitemap<N>.
var (
numberedSitemapRoute = regexp.MustCompile(`(?i)(?:^|\^|/)sitemap([0-9]+)\\?\.xml(?:\$|$)`)
numberedSitemapFeed = regexp.MustCompile(`(?i)(?:^|[?&])feed=xmlsitemap([0-9]+)(?:&|$)`)
)
// hostnameKeyedOption reports whether an option is cloak configuration keyed
// by a digest, returning the decoded size. Both halves are required: a plugin
// may hash a cache key, and base64 alone is ordinary.
func hostnameKeyedOption(name, value string) (int, bool) {
if !digestOptionName.MatchString(strings.ToLower(name)) {
return 0, false
}
decoded, ok := decodeBase64Payload(value)
if !ok || !isPHPSerializedArray(decoded) {
return 0, false
}
return len(decoded), true
}
// decodeBase64Payload decodes a stored base64 value. Kits wrap the stored text
// at arbitrary widths, so whitespace is removed before decoding, and padding is
// not always present.
func decodeBase64Payload(value string) ([]byte, bool) {
if value == "" || len(value) > maxCloakOptionBytes {
return nil, false
}
compactLen := 0
for i := 0; i < len(value); i++ {
c := value[i]
switch c {
case ' ', '\t', '\n', '\r', '\f', '\v':
continue
}
if (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') ||
(c >= '0' && c <= '9') || c == '+' || c == '/' || c == '=' {
compactLen++
continue
}
// Reject base64url, NUL, non-ASCII, and other invalid bytes
// before allocating a decoder-sized output buffer.
return nil, false
}
if compactLen == 0 {
return nil, false
}
compact := make([]byte, 0, compactLen)
for i := 0; i < len(value); i++ {
switch value[i] {
case ' ', '\t', '\n', '\r', '\f', '\v':
continue
default:
compact = append(compact, value[i])
}
}
encoding := base64.StdEncoding
switch len(compact) % 4 {
case 0:
// Padded standard base64 and unpadded complete quanta both land
// here and are accepted by StdEncoding.
case 2, 3:
encoding = base64.RawStdEncoding
default:
return nil, false
}
decoded := make([]byte, encoding.DecodedLen(len(compact)))
n, err := encoding.Decode(decoded, compact)
if err != nil {
// Decode may have written a valid prefix. Never inspect it after an
// error or trailing garbage could fabricate the serialized marker.
return nil, false
}
return decoded[:n], true
}
type phpSerializedFrame struct {
remaining int
closing byte
hasKeys bool
}
// isPHPSerializedArray validates one complete PHP serialization without
// materializing it. Checking the full grammar prevents a decoded prefix such
// as "a:1:{" (including one followed by NUL or garbage) from becoming a
// finding. The explicit frame stack also keeps hostile nesting off Go's call
// stack.
func isPHPSerializedArray(data []byte) bool {
if len(data) == 0 || data[0] != 'a' {
return false
}
frames := []phpSerializedFrame{{remaining: 1}}
pos := 0
for len(frames) > 0 {
frameIndex := len(frames) - 1
frame := frames[frameIndex]
if frame.remaining == 0 {
if frame.closing == 0 {
return len(frames) == 1 && pos == len(data)
}
if pos >= len(data) || data[pos] != frame.closing {
return false
}
pos++
frames = frames[:frameIndex]
continue
}
if pos >= len(data) {
return false
}
expectingKey := frame.hasKeys && frame.remaining%2 == 0
if expectingKey &&
data[pos] != 'i' && data[pos] != 's' {
return false
}
frames[frameIndex].remaining--
next, child, hasChild, ok := consumePHPSerializedValue(data, pos)
if !ok {
return false
}
pos = next
if hasChild {
frames = append(frames, child)
}
}
return false
}
func consumePHPSerializedValue(data []byte, pos int) (int, phpSerializedFrame, bool, bool) {
var noChild phpSerializedFrame
if pos >= len(data) {
return pos, noChild, false, false
}
switch data[pos] {
case 'N':
if pos+2 <= len(data) && string(data[pos:pos+2]) == "N;" {
return pos + 2, noChild, false, true
}
case 'b':
if pos+4 <= len(data) && data[pos+1] == ':' &&
(data[pos+2] == '0' || data[pos+2] == '1') && data[pos+3] == ';' {
return pos + 4, noChild, false, true
}
case 'i':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
next, ok := consumePHPSerializedInteger(data, pos+2)
return next, noChild, false, ok
case 'd':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
end := pos + 2
for end < len(data) && data[end] != ';' {
end++
}
if end == len(data) || end == pos+2 || end-(pos+2) > 64 {
break
}
number := string(data[pos+2 : end])
if number == "INF" || number == "-INF" || number == "NAN" {
return end + 1, noChild, false, true
}
if _, err := strconv.ParseFloat(number, 64); err == nil {
return end + 1, noChild, false, true
}
case 's', 'E':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
next, ok := consumePHPSerializedBytes(data, pos+2, ';')
return next, noChild, false, ok
case 'R', 'r':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
ref, next, ok := parsePHPSerializedUint(data, pos+2, ';')
return next, noChild, false, ok && ref > 0
case 'a':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
count, next, ok := parsePHPSerializedUint(data, pos+2, ':')
if !ok || count > len(data) || count > int(^uint(0)>>1)/2 ||
next >= len(data) || data[next] != '{' {
break
}
return next + 1, phpSerializedFrame{
remaining: count * 2,
closing: '}',
hasKeys: true,
}, true, true
case 'O':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
next, ok := consumePHPSerializedBytes(data, pos+2, ':')
if !ok {
break
}
count, next, ok := parsePHPSerializedUint(data, next, ':')
if !ok || count > len(data) || count > int(^uint(0)>>1)/2 ||
next >= len(data) || data[next] != '{' {
break
}
return next + 1, phpSerializedFrame{
remaining: count * 2,
closing: '}',
// __serialize() may return integer as well as string keys.
hasKeys: true,
}, true, true
case 'C':
if pos+2 > len(data) || data[pos+1] != ':' {
break
}
next, ok := consumePHPSerializedBytes(data, pos+2, ':')
if !ok {
break
}
payloadLen, next, ok := parsePHPSerializedUint(data, next, ':')
if !ok || next >= len(data) || data[next] != '{' ||
payloadLen > len(data)-(next+1) {
break
}
next += 1 + payloadLen
if next < len(data) && data[next] == '}' {
return next + 1, noChild, false, true
}
}
return pos, noChild, false, false
}
func consumePHPSerializedInteger(data []byte, pos int) (int, bool) {
start := pos
if pos < len(data) && data[pos] == '-' {
pos++
}
digitStart := pos
for pos < len(data) && data[pos] >= '0' && data[pos] <= '9' {
pos++
}
if pos == digitStart || pos >= len(data) || data[pos] != ';' || pos-start > 20 {
return start, false
}
if _, err := strconv.ParseInt(string(data[start:pos]), 10, 64); err != nil {
return start, false
}
return pos + 1, true
}
// consumePHPSerializedBytes consumes <length>:"<raw bytes>"<terminator>.
func consumePHPSerializedBytes(data []byte, pos int, terminator byte) (int, bool) {
length, next, ok := parsePHPSerializedUint(data, pos, ':')
if !ok || next >= len(data) || data[next] != '"' || length > len(data)-(next+1) {
return pos, false
}
next += 1 + length
if next+1 >= len(data) || data[next] != '"' || data[next+1] != terminator {
return pos, false
}
return next + 2, true
}
func parsePHPSerializedUint(data []byte, pos int, delimiter byte) (int, int, bool) {
start := pos
n := 0
maxInt := int(^uint(0) >> 1)
for pos < len(data) && data[pos] >= '0' && data[pos] <= '9' {
digit := int(data[pos] - '0')
if n > (maxInt-digit)/10 {
return 0, start, false
}
n = n*10 + digit
pos++
}
if pos == start || pos >= len(data) || data[pos] != delimiter {
return 0, start, false
}
return n, pos + 1, true
}
// doorwaySitemapRoutes returns the canonical cluster numbers routed from a
// numbered sitemap key into the matching numbered feed value. Requiring both
// halves in the same serialized rewrite-rule pair avoids correlating unrelated
// rules that merely occur in the same option.
func doorwaySitemapRoutes(rewriteRules string) []string {
routes, _ := doorwaySitemapRoutesChecked(rewriteRules)
return routes
}
func doorwaySitemapRoutesChecked(rewriteRules string) ([]string, bool) {
data := []byte(rewriteRules)
if len(data) < 6 || data[0] != 'a' || data[1] != ':' {
return nil, false
}
count, pos, ok := parsePHPSerializedUint(data, 2, ':')
if !ok || count > len(data) || pos >= len(data) || data[pos] != '{' {
return nil, false
}
pos++
routes := make(map[string]struct{})
for i := 0; i < count; i++ {
key, next, ok := parsePHPSerializedStringAt(rewriteRules, pos)
if !ok {
return nil, false
}
value, nextValue, ok := parsePHPSerializedStringAt(rewriteRules, next)
if !ok {
return nil, false
}
pos = nextValue
keyNumbers := cloakClusterNumbers(numberedSitemapRoute, key)
valueNumbers := cloakClusterNumbers(numberedSitemapFeed, value)
for number := range keyNumbers {
if _, paired := valueNumbers[number]; paired {
routes[number] = struct{}{}
}
}
}
if pos >= len(data) || data[pos] != '}' || pos+1 != len(data) {
return nil, false
}
out := make([]string, 0, len(routes))
for number := range routes {
out = append(out, number)
}
sort.Slice(out, func(i, j int) bool {
if len(out[i]) != len(out[j]) {
return len(out[i]) < len(out[j])
}
return out[i] < out[j]
})
return out, true
}
func cloakClusterNumbers(pattern *regexp.Regexp, value string) map[string]struct{} {
numbers := make(map[string]struct{})
for _, match := range pattern.FindAllStringSubmatch(value, -1) {
if len(match) < 2 || len(match[1]) > 10 {
continue
}
number := strings.TrimLeft(match[1], "0")
if number == "" {
number = "0"
}
numbers[number] = struct{}{}
}
return numbers
}
type cloakOptionRow struct {
kind string
name string
value []byte
}
func parseCloakOptionRow(line string) (cloakOptionRow, bool, bool) {
var row cloakOptionRow
parts := strings.SplitN(strings.TrimRight(line, "\r\n"), "\t", 4)
if len(parts) != 4 {
return row, false, false
}
valueBytes, err := strconv.ParseInt(strings.TrimSpace(parts[2]), 10, 64)
if err != nil || valueBytes < 0 {
return row, false, false
}
encoded := strings.TrimSpace(parts[3])
if !strings.HasPrefix(encoded, "x") {
return row, false, false
}
value, err := hex.DecodeString(encoded[1:])
if err != nil {
return row, false, false
}
expected := valueBytes
if expected > maxCloakOptionBytes {
expected = maxCloakOptionBytes
}
if int64(len(value)) != expected {
return row, false, false
}
return cloakOptionRow{
kind: strings.TrimSpace(parts[0]),
name: parts[1],
value: value,
}, true, valueBytes <= maxCloakOptionBytes
}
// checkWPCloakConfig reports doorway scaffolding kept in the options table.
func checkWPCloakConfig(user string, creds wpDBCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"(SELECT 'opt' AS kind, option_name, OCTET_LENGTH(option_value), "+
"CONCAT('x', HEX(LEFT(CAST(option_value AS BINARY), %d))) FROM %soptions "+
"WHERE autoload IN ('yes', 'on', 'auto-on', 'auto') "+
"AND OCTET_LENGTH(option_name) = 32 AND UNHEX(option_name) IS NOT NULL "+
"ORDER BY option_name LIMIT %d) UNION ALL "+
"(SELECT 'rules', option_name, OCTET_LENGTH(option_value), "+
"CONCAT('x', HEX(LEFT(CAST(option_value AS BINARY), %d))) FROM %soptions "+
"WHERE option_name = 'rewrite_rules' LIMIT 1)",
maxCloakOptionBytes, prefix, maxCloakOptionRows+1, maxCloakOptionBytes, prefix)
var keyed []string
routeSet := make(map[string]struct{})
candidates := 0
candidatesTruncated := false
for _, line := range runMySQLQuery(creds, query) {
row, ok, complete := parseCloakOptionRow(line)
if !ok {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if !complete {
markCheckIncomplete(creds.queryCtx, "db_content")
}
switch row.kind {
case "opt":
candidates++
if candidates > maxCloakOptionRows {
candidatesTruncated = true
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if complete {
if size, matched := hostnameKeyedOption(row.name, string(row.value)); matched {
keyed = append(keyed, fmt.Sprintf("%s (%d bytes decoded)", row.name, size))
}
}
case "rules":
if complete {
routes, parsed := doorwaySitemapRoutesChecked(string(row.value))
if !parsed {
continue
}
for _, route := range routes {
routeSet[route] = struct{}{}
}
}
default:
markCheckIncomplete(creds.queryCtx, "db_content")
}
}
sort.Strings(keyed)
routes := make([]string, 0, len(routeSet))
for route := range routeSet {
routes = append(routes, route)
}
sort.Slice(routes, func(i, j int) bool {
if len(routes[i]) != len(routes[j]) {
return len(routes[i]) < len(routes[j])
}
return routes[i] < routes[j]
})
var findings []alert.Finding
if len(keyed) > 0 {
count := strconv.Itoa(len(keyed))
if candidatesTruncated {
count = "at least " + count
}
noun, verb, hold := "options", "are", "hold"
if len(keyed) == 1 {
noun, verb, hold = "option", "is", "holds"
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "db_hostname_keyed_option",
Message: fmt.Sprintf("%s autoloaded WordPress %s %s named by digest and %s encoded data (account: %s)",
count, noun, verb, hold, user),
Details: dbContentFindingDetails(creds, prefix,
"An option named after a digest cannot be found without already knowing "+
"the key, and the base64 layer keeps its contents out of any search of "+
"the table. Cloak kits key that digest to the site's own hostname so one "+
"payload serves many sites. The row is autoloaded, so it is read on every request.",
cloakSample("Options", keyed)),
DedupKey: dbContentDedupKey(user, creds, prefix,
"An option named after a digest cannot be found without already knowing "+
"the key, and the base64 layer keeps its contents out of any search of "+
"the table. Cloak kits key that digest to the site's own hostname so one "+
"payload serves many sites. The row is autoloaded, so it is read on every request.",
cloakSample("Options", keyed)),
})
}
if len(routes) > 0 {
noun, verb := "routes", "feed"
if len(routes) == 1 {
noun, verb = "route", "feeds"
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "db_doorway_sitemap_routes",
Message: fmt.Sprintf("%d numbered sitemap %s %s generated pages to crawlers (account: %s)",
len(routes), noun, verb, user),
Details: dbContentFindingDetails(creds, prefix,
"Each rule routes sitemap<N>.xml straight into a matching feed, one per "+
"doorway cluster, so crawlers are handed the generated pages without them "+
"appearing in the site's real sitemap. Sitemap plugins add rewrite rules "+
"too, but none of them pair a numbered sitemap with a feed of the same number.",
cloakSample("Clusters", routes)),
DedupKey: dbContentDedupKey(user, creds, prefix,
"Each rule routes sitemap<N>.xml straight into a matching feed, one per "+
"doorway cluster, so crawlers are handed the generated pages without them "+
"appearing in the site's real sitemap. Sitemap plugins add rewrite rules "+
"too, but none of them pair a numbered sitemap with a feed of the same number.",
cloakSample("Clusters", routes)),
})
}
return findings
}
func cloakSample(label string, values []string) string {
shown := values
if len(shown) > maxCloakSamplesShown {
shown = shown[:maxCloakSamplesShown]
label += fmt.Sprintf(" (showing %d of %d)", len(shown), len(values))
}
return label + ": " + strings.Join(shown, ", ")
}
package checks
import (
"context"
"fmt"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/state"
)
// cmsScanRowLimit bounds every non-WordPress CMS content, settings and
// admin query. The WordPress scanner caps each of its selects; these
// adapters used to pull whole tables through the mysql client.
const cmsScanRowLimit = 200
// One extra row distinguishes a complete result at the cap from a truncated
// query. Failed or partial queries must not establish an administrator baseline.
func runCMSQuery(creds wpDBCreds, query string) ([]string, bool) {
if creds.queryState == nil {
creds.queryState = new(dbQueryState)
}
markIncomplete := func() {
creds.queryState.failed = true
creds.queryState.halted = true
markCheckIncomplete(creds.queryCtx, creds.queryCheck())
}
if creds.queryCtx != nil && creds.queryCtx.Err() != nil {
markIncomplete()
return nil, false
}
rows := runMySQLQuery(creds, fmt.Sprintf("%s LIMIT %d", query, cmsScanRowLimit+1))
if creds.queryCtx != nil && creds.queryCtx.Err() != nil {
markIncomplete()
return nil, false
}
if len(rows) > cmsScanRowLimit {
rows = rows[:cmsScanRowLimit]
markIncomplete()
}
return rows, !creds.queryState.failed
}
// cmsDiscover globs every pattern under every account root and returns the
// unique matches. Installs live under public_html and under addon-domain
// document roots (<home>/<domain>/...), so callers pass both shapes.
func cmsDiscover(ctx context.Context, owner string, patterns ...string) []string {
var out []string
for _, p := range patterns {
for _, root := range accountHomeRoots() {
if ctx.Err() != nil {
markCheckIncomplete(ctx, owner)
return uniqueStrings(out)
}
// One root's matches cannot cover another root's discovery error.
matches, err := osFS.Glob(filepath.Join(root, p))
if err != nil {
markCheckIncomplete(ctx, owner)
}
out = append(out, matches...)
}
}
return uniqueStrings(out)
}
func rankCMSConfigs(ctx context.Context, owner string, paths []string, maxFiles int) []string {
ranked := rankPathsByMtimeDesc(ctx, paths, maxFiles)
if len(ranked) < len(paths) || ctx.Err() != nil {
markCheckIncomplete(ctx, owner)
}
return ranked
}
// cmsAdminFindings reports CMS administrator rows. With a store, the first
// complete pass records every admin id for the install and stays quiet;
// from then on only an id not seen before is reported, once, as a High
// finding. Without a store (ad-hoc runs, tests) it keeps the historical
// per-row visibility Warning. describe renders the message tail and the
// details for one row's tab-separated fields; fields[0] is the id.
func cmsAdminFindings(store *state.Store, cms, check, account string, creds wpDBCreds, rows []string, complete bool, describe func(fields []string) (message, details string)) []alert.Finding {
for _, row := range rows {
id, _, _ := strings.Cut(row, "\t")
if !isAllDigits(id) {
complete = false
markCheckIncomplete(creds.queryCtx, creds.queryCheck())
}
}
var findings []alert.Finding
// Account-wide keys cannot establish which installation supplied an id.
// Each database and prefix gets a fresh baseline after the key migration.
siteKey := dbContentDedupKey(account, creds, creds.tablePrefix, "cms-admin="+cms)
baselineKey := "_cmsadmin_baseline:v2:" + siteKey
baselined := false
if store != nil {
_, baselined = store.GetRaw(baselineKey)
}
seenIDs := make(map[string]bool, len(rows))
for _, row := range rows {
fields := strings.Split(row, "\t")
if !isAllDigits(fields[0]) || seenIDs[fields[0]] {
continue
}
seenIDs[fields[0]] = true
message, details := describe(fields)
details = dbContentFindingDetails(creds, creds.tablePrefix,
"Database host: "+creds.dbHost, details)
dedupKey := dbContentDedupKey(account, creds, creds.tablePrefix, "cms-admin="+cms, "id="+fields[0])
if store == nil {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: check,
Message: message,
Details: details,
DedupKey: dedupKey,
})
continue
}
key := "_cmsadmin:v2:" + dedupKey
if _, seen := store.GetRaw(key); seen {
continue
}
if complete {
store.SetRaw(key, "seen")
}
if !baselined {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: check,
Message: "New " + message,
Details: details + "\nThis administrator was not present when the install was baselined.",
DedupKey: dedupKey,
})
}
if store != nil && !baselined && complete {
store.SetRaw(baselineKey, "1")
}
return findings
}
package checks
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// Drupal database content scanner.
//
// v1 covers Drupal 8 and later (the unified-table layout: config /
// node_field_data / users_field_data). Drupal 7 uses a different
// schema (variable / node / users) and reached EOL in January 2025;
// scanning it lands as a follow-up if any operator reports D7 sites
// still in production.
//
// Discovery: glob /home/*/public_html/sites/default/settings.php
// and confirm core/lib/Drupal.php exists in the site root --
// that file is the canonical D8+ marker (D7 has bootstrap.inc /
// modules/ but no core/ directory).
//
// Credentials: parsed by regex over the canonical $databases
// array literal. Drupal allows both array() and short [] syntax;
// the regex accepts either.
//
// Scanned tables (all unprefixed -- D8+ does not use a per-site
// table prefix in standard installs):
//
// config name + data; data is a
// serialized PHP array carrying
// site name, slogan, theme, etc.
// Common hijack target.
// node_revision__body entity_id + body_value; the
// actual article body. Scanned for
// any pattern in dbMalwarePatterns
// with the same external-script
// post-filter the WP and Joomla
// scanners use.
// users_field_data user table; joined with
// user__roles on entity_id = uid. Rows where
// roles_target_id = 'administrator'
// surface as drupal_admin_injection.
//
// Three new finding categories with CMS-explicit names:
// drupal_settings_injection, drupal_content_injection,
// drupal_admin_injection.
// drupalAdminRoleID is the canonical role identifier for Drupal
// site administrators in vanilla D8+. Operators on hardened
// installs may have renumbered or renamed; v1 narrows to this.
const drupalAdminRoleID = "administrator"
// drupalSettingsRe pulls credentials out of the $databases array
// literal. Each field is matched independently rather than
// trying to parse the array structure -- attackers occasionally
// reorder keys, and a key-only regex ignores layout differences.
var (
drupalDBNameRe = regexp.MustCompile(`'database'\s*=>\s*['"]([^'"]+)['"]`)
drupalDBUserRe = regexp.MustCompile(`'username'\s*=>\s*['"]([^'"]+)['"]`)
drupalDBPassRe = regexp.MustCompile(`'password'\s*=>\s*['"]([^'"]+)['"]`)
drupalDBHostRe = regexp.MustCompile(`'host'\s*=>\s*['"]([^'"]+)['"]`)
)
// drupalCreds carries the parsed connection details. Mirrors the
// jConfigCreds shape so existing helpers (runMySQLQuery,
// asWPDBCreds) work uniformly.
type drupalCreds struct {
// ctx ties every query for this install to the runner's deadline.
ctx context.Context
dbName string
dbUser string
dbPass string
dbHost string
path string
queryState *dbQueryState
}
func (c drupalCreds) asWPDBCreds() wpDBCreds {
return wpDBCreds{
dbName: c.dbName,
dbUser: c.dbUser,
dbPass: c.dbPass,
dbHost: c.dbHost,
queryCtx: c.ctx,
queryOwner: "db_content_drupal",
queryState: c.queryState,
}
}
// CheckDrupalContent discovers Drupal 8+ sites and scans the three
// canonical attacker-touched tables. Mirrors CheckJoomlaContent
// without sharing code -- the credential layout and table set are
// distinct enough that a generic dispatcher would be more
// abstraction than a 4-CMS pipeline calls for.
// scanDrupalInstall scans one discovered install and stamps its findings
// with the owner resolved from the settings path. The display label stays
// as before; an install outside every account root is not stamped.
func scanDrupalInstall(ctx context.Context, path string, store *state.Store) []alert.Finding {
// public_html is three dirs up from sites/default/settings.php.
publicHTML := filepath.Dir(filepath.Dir(filepath.Dir(path)))
matched, err := looksLikeDrupal8Plus(publicHTML)
if err != nil {
markCheckIncomplete(ctx, "db_content_drupal")
return nil
}
if !matched {
return nil
}
// /home/<account> is one level above public_html.
account := extractUser(filepath.Dir(publicHTML))
creds, err := parseDrupalSettings(ctx, path)
if err != nil || creds.dbName == "" || creds.dbUser == "" {
markCheckIncomplete(ctx, "db_content_drupal")
return nil
}
creds.ctx = ctx
creds.queryState = new(dbQueryState)
var findings []alert.Finding
findings = append(findings, scanDrupalConfig(account, creds)...)
findings = append(findings, scanDrupalContent(account, creds)...)
findings = append(findings, scanDrupalAdmins(store, account, creds)...)
if owner, ok := installOwner(path); ok {
findings = stampTenantIDIfEmpty(findings, owner)
}
return findings
}
func CheckDrupalContent(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
settings := cmsDiscover(ctx, "db_content_drupal", "*/public_html/sites/default/settings.php", "*/*/sites/default/settings.php")
if len(settings) == 0 {
return nil
}
// Rank by mtime desc so recently touched Drupal sites are processed
// first when the check timeout cuts iteration short.
for _, path := range rankCMSConfigs(ctx, "db_content_drupal", settings, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
return findings
}
findings = append(findings, scanDrupalInstall(ctx, path, store)...)
}
return findings
}
// The marker distinguishes D8+ from D7 without reading its contents. A
// symlink or special file cannot establish the installation's version.
func looksLikeDrupal8Plus(publicHTML string) (bool, error) {
marker := filepath.Join(publicHTML, "core", "lib", "Drupal.php")
info, err := osFS.Lstat(marker)
if errors.Is(err, os.ErrNotExist) {
return false, nil
}
if err != nil {
return false, err
}
if !info.Mode().IsRegular() {
return false, errNonRegularFile
}
return true, nil
}
// parseDrupalSettings reads settings.php and returns the database
// credentials from the $databases['default']['default'] entry. If
// settings.php uses split-DB or per-environment overrides, only
// the first 'default' connection is reported -- the rest are
// followed by the same regex on subsequent calls.
func parseDrupalSettings(ctx context.Context, path string) (drupalCreds, error) {
creds := drupalCreds{path: path}
data, err := readCMSConfig(ctx, path)
if err != nil {
return creds, err
}
body := string(data)
if m := drupalDBNameRe.FindStringSubmatch(body); m != nil {
creds.dbName = m[1]
}
if m := drupalDBUserRe.FindStringSubmatch(body); m != nil {
creds.dbUser = m[1]
}
if m := drupalDBPassRe.FindStringSubmatch(body); m != nil {
creds.dbPass = m[1]
}
if m := drupalDBHostRe.FindStringSubmatch(body); m != nil {
creds.dbHost = m[1]
}
if creds.dbHost == "" {
creds.dbHost = "localhost"
}
return creds, nil
}
// scanDrupalConfig pulls rows from the config table whose data
// blob matches any malware pattern, then refines via
// classifyMalwareRow (strict / config-storage variant).
func scanDrupalConfig(account string, creds drupalCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT name, data FROM config WHERE %s",
paramsLikeClause("data"))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
name, body := splitTabRow(row)
if name == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, false)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "drupal_settings_injection",
Message: fmt.Sprintf("Drupal config injection on %s: %s (%s)", account, name, desc),
Details: fmt.Sprintf("Account: %s\nConfig name: %s\nMatch: %s", account, name, desc),
})
}
return findings
}
// scanDrupalContent walks node_revision__body for malware-pattern
// matches in article bodies. The looser post-filter
// (hasMaliciousExternalScriptInPost) applies because article
// content is author-written and may carry pre-TLS-era embeds.
func scanDrupalContent(account string, creds drupalCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT entity_id, body_value FROM node_revision__body WHERE %s",
paramsLikeClause("body_value"))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
entityID, body := splitTabRow(row)
if entityID == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, true)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "drupal_content_injection",
Message: fmt.Sprintf("Drupal article content injection on %s: node %s (%s)", account, entityID, desc),
Details: fmt.Sprintf("Account: %s\nNode entity_id: %s\nMatch: %s", account, entityID, desc),
})
}
return findings
}
// scanDrupalAdmins joins users_field_data with user__roles to
// surface every account in the administrator role. Same Warning
// severity / per-row emission as the Joomla equivalent --
// legitimate site admin shows up here too, so this is operator
// review territory rather than auto-actionable.
//
// users_field_data is multilingual in D8+: a single uid can appear
// once per language code on translated sites. The default_langcode
// = 1 filter keeps each admin to exactly one finding regardless of
// how many translations the site has.
func scanDrupalAdmins(store *state.Store, account string, creds drupalCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT u.uid, u.name, u.mail FROM users_field_data u JOIN user__roles r ON u.uid = r.entity_id WHERE r.roles_target_id = '%s' AND u.default_langcode = 1",
drupalAdminRoleID)
rows, complete := runCMSQuery(creds.asWPDBCreds(), query)
return cmsAdminFindings(store, "drupal", "drupal_admin_injection", account, creds.asWPDBCreds(), rows, complete, func(fields []string) (string, string) {
return fmt.Sprintf("Drupal administrator account on %s: %s", account, fields[0]),
fmt.Sprintf("Account: %s\nRow: %s\nReview: confirm this is the legitimate site administrator.", account, strings.Join(fields, "\t"))
})
}
package checks
import (
"regexp"
"strings"
"unicode"
"unicode/utf8"
)
// This file contains pure-function helpers used by the database content
// scanner (checkWPPosts in dbscan.go). Keeping them pure (no MySQL, no
// filesystem) makes them deterministically testable and independently
// reusable.
//
// Two classes of false positive were historically observed on real
// production traffic:
//
// 1. db_post_injection fired on every post containing a script tag,
// including site-owner-added analytics and widget embeds (Google Tag
// Manager, Google merchant rating badge, etc.).
//
// 2. db_spam_injection used substring LIKE matching, so "specialist"
// triggered on "cialis", "pharmaceutical" triggered on "pharma",
// "casino resort" triggered on "casino", etc. It also scanned all
// post_types including Contact Form 7 / WPForms / Jetpack stored
// submissions, which routinely contain spambot form fills the site
// owner never displays.
//
// The helpers below encode the decisions needed to eliminate those FPs
// without opening detection holes: word-boundary keyword matching,
// post_type filtering against a denylist (not an allowlist, so attackers
// cannot hide a post by renaming post_type to one we didn't anticipate),
// and safe-domain filtering for external script-tag sources.
// nonScannablePostTypes are post_type values that legitimately store
// non-site-content data (form submissions, revisions, templates, feeds).
// These are excluded from malware and spam scans because their content
// is operator-invisible storage, not material rendered to site visitors.
//
// This is a DENYLIST, not an allowlist. A custom post_type created by a
// theme or plugin (for example WooCommerce `product`, events, portfolios)
// is still scanned. Adding a new value here is safe; an attacker cannot
// hide a post by choosing a new post_type, because we default to
// scanning anything not on this list.
var nonScannablePostTypes = []string{
// WordPress internals / templates / navigation
"revision",
"customize_changeset",
"oembed_cache",
"nav_menu_item",
"wp_template",
"wp_template_part",
"wp_global_styles",
"wp_navigation",
// Minification plugins (store compiled bundles that legitimately
// contain JavaScript and obfuscated character sequences).
"wphb_minify_group",
// Form builders (store plugin configuration and visitor submissions.
// Contact-form spam landing here is noise, not site compromise.)
"wpforms",
"wpforms_entries",
"wpforms-log",
"wpcf7_contact_form",
"flamingo_inbound",
"flamingo_outbound",
"cf7_message",
"feedback",
"jetpack_feedback",
}
// isScannablePostType returns true if the given post_type should be
// included in malware and spam scans. The decision mirrors the SQL
// `post_type NOT IN (...)` clause used in checkWPPosts, so Go-side
// callers and test assertions stay consistent with the live SQL.
func isScannablePostType(postType string) bool {
for _, t := range nonScannablePostTypes {
if t == postType {
return false
}
}
return true
}
// nonScannablePostTypesSQLList returns the denylist as a comma-separated
// SQL literal list (for example "'revision','wp_template',...") suitable
// for use inside a `post_type NOT IN (...)` clause. The values are
// hardcoded and contain only [a-z_-] characters, so SQL injection is not
// a risk here; nonetheless the function escapes defensively.
func nonScannablePostTypesSQLList() string {
parts := make([]string, 0, len(nonScannablePostTypes))
for _, t := range nonScannablePostTypes {
// Hardcoded list contains only [a-z_-], but defensive escape in
// case future maintainers add a value with a quote or backslash.
escaped := strings.ReplaceAll(t, `\`, `\\`)
escaped = strings.ReplaceAll(escaped, `'`, `\'`)
parts = append(parts, "'"+escaped+"'")
}
return strings.Join(parts, ",")
}
// dbSpamPattern is a single spam-keyword detector. The LIKE fragment
// is a MySQL server-side pre-filter that quickly narrows the set of
// candidate posts; the Go regex finds case-insensitive occurrences and
// spamKeywordMatchIndexes applies strict word boundaries so that
// "specialist" does not match "cialis", "pharmacy" does not match
// "pharma", and so on.
//
// Patterns that end with a non-word character (dash) already have an
// implicit right boundary from that character and only need a left
// word boundary. Patterns ending in a word character need boundaries
// on both sides, even when they contain an internal dash.
type dbSpamPattern struct {
keyword string // human-readable keyword used in finding messages
regex *regexp.Regexp // applied Go-side to candidate rows
likeFragment string // SQL LIKE fragment, always bracketed with '%'
// deletable gates the DESTRUCTIVE path only. Scanning uses every
// pattern; db-clean --delete-spam uses this subset. "pharma" and
// "betting" are genuine spam keywords that are also ordinary words in
// legitimate publishing ("Health and Pharma Summit", coverage of a
// licensed bookmaker), and a word boundary cannot tell the two apart.
// Flagging those for review is useful; deleting on them is not.
deletable bool
}
// dbSpamPatterns enumerates the keywords we flag as SEO/pharma/gambling
// spam in WordPress post content. Each entry pairs a fast SQL LIKE with
// a case-insensitive Go regex; spamKeywordMatchIndexes applies the strict
// Unicode word-boundary check.
//
// The regexes are case-insensitive to catch CIALIS / Cialis / cialis.
// The LIKE fragments are lowercase because MySQL LIKE is case-insensitive
// under the default _ci collation used by cPanel MariaDB.
var dbSpamPatterns = []dbSpamPattern{
newDBSpamPattern("viagra", "%viagra%", true),
newDBSpamPattern("cialis", "%cialis%", true),
newDBSpamPattern("pharma", "%pharma%", false),
newDBSpamPattern("betting", "%betting%", false),
// Dashed variants: the trailing dash is itself a non-word char and
// serves as the right boundary. Only a left word-boundary is needed.
// The dash makes these URL-slug shaped, which prose does not produce,
// so they carry enough signal to delete on.
newDBSpamPattern("casino-", "%casino-%", true),
newDBSpamPattern("buy-cheap-", "%buy-cheap-%", true),
// These slug forms end in a letter, so the matcher requires both
// boundaries and rejects longer words such as "free-downloadable".
newDBSpamPattern("free-download", "%free-download%", true),
newDBSpamPattern("crack-serial", "%crack-serial%", true),
}
func newDBSpamPattern(keyword, likeFragment string, deletable bool) dbSpamPattern {
return dbSpamPattern{
keyword: keyword,
regex: regexp.MustCompile(`(?i)` + regexp.QuoteMeta(keyword)),
likeFragment: likeFragment,
deletable: deletable,
}
}
// spamKeywordMatchIndexes returns only whole-word keyword matches. Go's \b
// is ASCII-only, so checking adjacent runes explicitly prevents a keyword
// surrounded by non-ASCII letters from being treated as a standalone word.
// A keyword ending in punctuation, such as "casino-", needs no right check:
// the punctuation is already the delimiter that makes the pattern specific.
func spamKeywordMatchIndexes(pattern dbSpamPattern, content string) [][]int {
matches := pattern.regex.FindAllStringIndex(content, -1)
confirmed := matches[:0]
for _, match := range matches {
if hasSpamWordBoundaries(content, match[0], match[1]) {
confirmed = append(confirmed, match)
}
}
return confirmed
}
func hasSpamWordBoundaries(content string, start, end int) bool {
first, _ := utf8.DecodeRuneInString(content[start:end])
if start > 0 && isSpamWordRune(first) {
previous, _ := utf8.DecodeLastRuneInString(content[:start])
if isSpamWordRune(previous) {
return false
}
}
last, _ := utf8.DecodeLastRuneInString(content[start:end])
if end < len(content) && isSpamWordRune(last) {
next, _ := utf8.DecodeRuneInString(content[end:])
if isSpamWordRune(next) {
return false
}
}
return true
}
func isSpamWordRune(r rune) bool {
return r == '_' || r == '\u200c' || r == '\u200d' ||
r == unicode.ReplacementChar || unicode.IsLetter(r) ||
unicode.IsNumber(r) || unicode.IsMark(r) || unicode.Is(unicode.Pc, r)
}
// countSpamMatches returns the number of candidate rows with a bounded
// keyword match. The caller is responsible for passing only rows that were
// already narrowed by the pattern.likeFragment SQL pre-filter.
func countSpamMatches(pattern dbSpamPattern, contents []string) int {
n := 0
for _, c := range contents {
if len(spamKeywordMatchIndexes(pattern, c)) > 0 {
n++
}
}
return n
}
// hasMaliciousExternalScript reports whether the content contains a
// script-tag with a src attribute pointing at a domain NOT on the known-
// safe list (see knownSafeDomains in db_autoresponse.go).
//
// This predicate uses the STRICT classifier (isAttackerScriptURL): it
// flags raw-IP hosts, abused TLDs, known-bad exfil hosts, empty hosts,
// AND plaintext HTTP external scripts. It is the right predicate for
// wp_options and similar configuration storage where fresh-today
// content is expected; see hasMaliciousExternalScriptInPost for the
// looser post_content predicate.
//
// Inline script blocks without a src attribute are not classified by
// this function; those are covered by the separate code-pattern entries
// in dbMalwarePatterns which catch common inline obfuscation techniques.
//
// Rationale: a bare script-tag match was the primary source of false
// positives on real traffic. Legitimate analytics embeds (Google Tag
// Manager, Google Analytics, Google merchant rating badge, Mailchimp,
// HubSpot, etc.) install both an external loader tag AND an inline
// initialization block. Flagging the inline block alone produced many
// HIGH severity noise findings on customer sites. Requiring a
// non-safe-domain external src reduces this to zero FPs in practice
// while still catching attackers who inject a tag pointing at an
// untrusted domain.
func hasMaliciousExternalScript(content string) bool {
return extractMaliciousScriptURL(content) != ""
}
// hasMaliciousExternalScriptInPost is the post_content variant of
// hasMaliciousExternalScript. It applies the same regex-based script-tag
// extraction but classifies each src URL with isAttackerScriptURLInPost,
// which drops the plaintext-HTTP indicator.
//
// Rationale: post_content carries author-written text that can contain
// pre-TLS-era embeds (e.g. a 2013-era video embed on a Romanian video
// site). Those embeds use http:// and, under the strict classifier,
// would flag on scheme alone even though the post has not been modified
// in a decade. Attackers in 2026 almost never land on plaintext-HTTP
// mainstream-TLD URLs — they use raw IPs, abused TLDs, or cheap exfil
// hosts — so dropping the HTTP signal for this context eliminates the
// legacy-embed false positives without giving up meaningful detection.
func hasMaliciousExternalScriptInPost(content string) bool {
// Page builders keep their markup as JSON inside the post row, which
// escapes every slash of an injected loader URL.
matches := scriptSrcRe.FindAllStringSubmatch(unescapeStoredSlashes(content), -1)
for _, match := range matches {
if len(match) < 2 {
continue
}
if isAttackerScriptURLInPost(match[1]) {
return true
}
}
return false
}
package checks
import (
"fmt"
"math"
"net"
"net/url"
"sort"
"strconv"
"strings"
"unicode/utf8"
"github.com/tdewolff/parse/v2"
cssparser "github.com/tdewolff/parse/v2/css"
"golang.org/x/net/html"
"golang.org/x/net/idna"
"golang.org/x/net/publicsuffix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/mysqlclient"
)
// Hidden-container link injection stored in the database.
//
// The kit wraps outbound links in a container the reader never sees -- an
// absolutely positioned block pushed thousands of pixels off the canvas, a
// text-indent far outside the viewport, or a display:none wrapper -- so a
// visitor sees an ordinary page while crawlers follow the links. CSM already
// ships file-side rules for this shape, but these injections live in
// post_content and option rows, where no file scanner ever looks.
//
// A hidden container on its own is ordinary: themes hide screen-reader text,
// mobile menus and collapsed panels. What has no benign reading is a hidden
// container wrapping links to somebody else's domain.
const (
// maxHiddenLinkRows bounds the rows pulled per table.
maxHiddenLinkRows = 200
// maxHiddenLinkValueBytes bounds one row's markup. Injected blocks sit at
// the top or bottom of the content, so half is read from each end rather
// than letting padding at the front hide a trailing block.
maxHiddenLinkValueBytes = 128 * 1024
maxHiddenLinkSampleBytes = maxHiddenLinkValueBytes / 2
// maxHiddenLinkSiteURLBytes prevents a poisoned site address from making
// the otherwise bounded candidate query return an attacker-sized value.
maxHiddenLinkSiteURLBytes = 4 * 1024
// maxHiddenLinkNodes bounds one row's parsed markup. The walk is iterative
// so nesting cannot exhaust the stack, but a hostile row must not be able
// to spend unbounded time either.
maxHiddenLinkNodes = 200000
// maxHiddenLinkHostsShown bounds the hosts named in the finding.
maxHiddenLinkHostsShown = 12
// maxHiddenLinkRowsShown bounds the rows named in the finding.
maxHiddenLinkRowsShown = 10
// offScreenPixels is how far outside the canvas an offset must sit before
// it is cloaking rather than layout. Real layouts nudge elements by a few
// pixels; the kits observed used four- and five-digit offsets.
offScreenPixels = 1000
// Font-relative offsets reach the same distance with smaller numbers. The
// threshold remains high enough to exclude ordinary indentation.
offScreenFontUnits = 100
)
// hiddenLinkCandidateCondition selects rows whose sampled markup may hide a
// link. The Go parser joins the samples when they cover the whole value, so
// SQL must examine that value together too. Larger values use separate end
// windows. Character windows cover every sampled byte, with bounded work.
//
// commented adds matching for declarations split by CSS comments. That can
// exhaust the server's regex work limit on dense content. The selection
// without it stays well inside that limit for a sampled window, so callers
// retry without it after a regex timeout.
func hiddenLinkCandidateCondition(column string, commented bool) string {
head := fmt.Sprintf("LEFT(%s, CASE WHEN OCTET_LENGTH(%s) <= %d THEN %d ELSE %d END)",
column, column, maxHiddenLinkValueBytes, maxHiddenLinkValueBytes, maxHiddenLinkSampleBytes)
return fmt.Sprintf("(%s OR (OCTET_LENGTH(%s) > %d AND %s))",
hiddenLinkWindowCondition(head, commented), column, maxHiddenLinkValueBytes,
hiddenLinkWindowCondition(fmt.Sprintf("RIGHT(%s, %d)", column, maxHiddenLinkSampleBytes), commented))
}
// hiddenLinkPlainDeclarations are the declaration prefixes the parser can
// hide or move off-screen, after whitespace and calc( are removed. Longer
// property names such as margin-left end in one of these.
var hiddenLinkPlainDeclarations = []string{
"display:none", "visibility:hidden", "visibility:collapse",
"opacity:0", "opacity:+0", "opacity:.0", "opacity:+.0", "opacity:-",
"text-indent:-", "left:-", "top:-", "right:-", "bottom:-",
}
// hiddenLinkRegexPattern is matched against the reversed, normalized window.
// Encodings there are backslashes, and each plain declaration has become
// style= followed by a backslash. From each backslash the pattern searches
// back to style=, stopping at the preceding tag delimiter or backslash, so no
// stretch of text is traversed twice and encoded page text outside a style
// attribute is not selected. A single alternative keeps the cost of each
// backslash low enough for dense content.
func hiddenLinkRegexPattern() string {
return `\\[^>\\]*=elyts`
}
// hiddenLinkCommentPattern reads the unmodified ASCII-converted window.
// Removing whitespace or calc( inside a comment can manufacture a closing
// delimiter and lose a declaration that the CSS parser recognizes as hidden.
func hiddenLinkCommentPattern() string {
// Reverse the whole comment grammar: a comment can contain /*, so its
// reversed body can contain */. A comment fallback must keep the
// declaration's tokens adjacent, or a visible declaration could join
// unrelated page text and exhaust the candidate limit.
// ASCII conversion maps Unicode whitespace to ?, without joining comment
// body characters. Keep gaps outside the comment repetition unambiguous.
space := `[[:space:]?]*`
gap := space + `(/[*]+([^*]*[^*/][*]+)*[^*]*[*]/` + space + `)*`
return strings.Join([]string{
"enon" + gap + ":" + gap + "yalpsid",
"(neddih|espalloc)" + gap + ":" + gap + "ytilibisiv",
// Positive values with negative exponents can underflow to zero.
"(-|0[.]?[+]?|-e[0-9_.]+[+]?)" + gap + ":" + gap + "yticapo",
"-" + gap + "([(]" + gap + "clac" + gap + ")?:" + gap + "(tnedni-txet|tfel|pot|thgir|mottob)",
}, "|")
}
// hiddenLinkWindowCondition selects one sample window. Literal replacements
// do the bulk of the work because they carry no regex budget. The chain is
// written in both branches of one CASE, so the server evaluates it once.
func hiddenLinkWindowCondition(window string, commented bool) string {
ascii := func(code int) string { return fmt.Sprintf("CHAR(%d USING ascii)", code) }
lower := "CAST(LOWER(" + window + ") AS BINARY)"
// Conversion turns every other character into a question mark. Removing
// those and ASCII whitespace joins each declaration's tokens; the value
// parser trims Unicode whitespace, so every parsed form stays adjacent.
normalized := "CONVERT(LOWER(" + window + ") USING ascii)"
for _, code := range []int{'\t', '\n', '\v', '\f', '\r', ' ', '?'} {
normalized = fmt.Sprintf("REPLACE(%s, %s, '')", normalized, ascii(code))
}
normalized = fmt.Sprintf("REPLACE(%s, 'calc(', '')", normalized)
// A plain declaration becomes a match the regex finds without searching.
matched := "CONCAT('style=', " + ascii(92) + ")"
for _, declaration := range hiddenLinkPlainDeclarations {
normalized = fmt.Sprintf("REPLACE(%s, '%s', %s)", normalized, declaration, matched)
}
// Numeric entities only need their first digit, as in the old prefix match.
normalized = fmt.Sprintf("REPLACE(REPLACE(%s, '&#x', '&#'), '&colon', %s)", normalized, ascii(92))
for _, digit := range "0123456789abcdef" {
normalized = fmt.Sprintf("REPLACE(%s, '&#%c', %s)", normalized, digit, ascii(92))
}
// Only rows carrying an encoding or comment need a regex. Hex literals
// and CHAR keep backslashes independent of SQL escape modes.
encoded := fmt.Sprintf("LOCATE('&#', %[1]s) > 0 OR LOCATE('&colon', %[1]s) > 0 OR LOCATE(CHAR(92), %[1]s) > 0", lower)
pattern := fmt.Sprintf("CONVERT(0x%x USING ascii)", hiddenLinkRegexPattern())
selection := fmt.Sprintf("(CASE WHEN %s THEN REVERSE(%s) REGEXP %s ELSE LOCATE(%s, %s) > 0 END)",
encoded, normalized, pattern, ascii(92), normalized)
if commented {
selection = fmt.Sprintf("(CASE WHEN %s THEN 1 WHEN LOCATE('/*', %s) > 0 "+
"THEN REVERSE(CONVERT(LOWER(%s) USING ascii)) REGEXP CONVERT(0x%x USING ascii) ELSE 0 END)",
selection, lower, window, hiddenLinkCommentPattern())
}
return fmt.Sprintf("(LOCATE('style', %s) > 0 AND %s)", lower, selection)
}
// hiddenLinkHit is what one row's markup revealed.
type hiddenLinkHit struct {
// offScreen records the strong signal: a container moved outside the
// canvas. Nothing legitimate positions readable content there.
offScreen bool
// hosts are the distinct off-site hosts linked from inside hidden
// containers, sorted.
hosts []string
// domains are the distinct registrable domains behind hosts. Severity is
// based on these so subdomains of one target do not look like a link farm.
domains []string
// multiDomain records that one hidden container, rather than merely one
// database row, links to at least two registrable domains.
multiDomain bool
// spammy records gambling or pharmacy vocabulary in the hidden anchors.
spammy bool
}
// hiddenOffsiteLinks reports the off-site hosts linked from inside a container
// the page hides from readers. siteHost is the site's own address; links back
// to it are ordinary navigation whatever their styling.
//
// This tokenizes rather than building a tree. x/net/html refuses markup nested
// deeper than 512 elements, and returning nothing for those rows would hand
// every kit a one-line evasion. The tokenizer has no depth limit, and tracking
// open elements on an explicit heap stack keeps attacker-controlled nesting off
// the goroutine stack, where an overflow is fatal and unrecoverable.
func hiddenOffsiteLinks(markup, siteHost string) hiddenLinkHit {
return hiddenOffsiteLinksForSites(markup, []string{siteHost})
}
func hiddenOffsiteLinksForSites(markup string, siteHosts []string) hiddenLinkHit {
siteDomains := make(map[string]bool, len(siteHosts))
for _, host := range siteHosts {
if domain := registrableDomain(host); domain != "" {
siteDomains[domain] = true
}
}
if len(siteDomains) == 0 {
return hiddenLinkHit{}
}
type openElement struct {
name string
hidden bool
hiddenGroup int
visibilityHidden bool
visibilityGroup int
offScreen bool
}
var stack []openElement
inherited := func() openElement {
if len(stack) == 0 {
return openElement{}
}
return stack[len(stack)-1]
}
var hit hiddenLinkHit
seenHosts := make(map[string]bool)
seenDomains := make(map[string]bool)
nextHiddenGroup := 1
groupDomains := make(map[int]map[string]bool)
newHiddenGroup := func() int {
group := nextHiddenGroup
nextHiddenGroup++
return group
}
recordGroupDomain := func(group int, domain string) {
if group == 0 {
return
}
domains := groupDomains[group]
if domains == nil {
domains = make(map[string]bool)
groupDomains[group] = domains
}
domains[domain] = true
if len(domains) >= 2 {
hit.multiDomain = true
}
}
// anchorHost is the host of the hidden anchor currently open, so its link
// text can be graded when the anchor closes.
anchorHost, anchorText := "", strings.Builder{}
anchorDepth := -1
closeAnchor := func() {
if anchorHost == "" {
return
}
if termNameSpamVocabulary.MatchString(anchorText.String()) || termNameSpamVocabulary.MatchString(anchorHost) {
hit.spammy = true
}
anchorHost, anchorText = "", strings.Builder{}
anchorDepth = -1
}
popOpenElement := func(name string) bool {
for i := len(stack) - 1; i >= 0; i-- {
if stack[i].name != name {
continue
}
if anchorDepth >= i {
closeAnchor()
}
stack = stack[:i]
return true
}
return false
}
z := html.NewTokenizer(strings.NewReader(markup))
for tokens := 0; tokens < maxHiddenLinkNodes; tokens++ {
switch z.Next() {
case html.ErrorToken:
closeAnchor()
sort.Strings(hit.hosts)
sort.Strings(hit.domains)
return hit
case html.TextToken:
if anchorHost != "" && anchorText.Len() < 512 &&
(len(stack) == 0 || !rawTextHTMLElements[stack[len(stack)-1].name]) {
text := z.Text()
remaining := 512 - anchorText.Len()
if len(text) > remaining {
text = text[:remaining]
}
anchorText.Write(text)
}
case html.StartTagToken, html.SelfClosingTagToken:
name, style, href := tokenAttrs(z)
if name == "a" {
popOpenElement("a")
closeAnchor()
}
state := inherited()
if style != "" {
styleState := parseHiddenCSSState(style)
if styleState.hidden {
if !state.hidden {
state.hiddenGroup = newHiddenGroup()
}
state.hidden = true
}
if styleState.visibilitySet {
if styleState.visibilityHidden {
if !state.visibilityHidden {
state.visibilityGroup = newHiddenGroup()
}
} else {
state.visibilityGroup = 0
}
state.visibilityHidden = styleState.visibilityHidden
}
if styleState.offScreen {
state.offScreen = true
}
}
if name == "a" && (state.hidden || state.visibilityHidden) && href != "" {
if host, ok := absoluteLinkHost(href); ok {
domain := registrableDomain(host)
if domain != "" && !siteDomains[domain] {
if !seenHosts[host] {
seenHosts[host] = true
hit.hosts = append(hit.hosts, host)
}
if !seenDomains[domain] {
seenDomains[domain] = true
hit.domains = append(hit.domains, domain)
}
recordGroupDomain(state.hiddenGroup, domain)
recordGroupDomain(state.visibilityGroup, domain)
if state.offScreen {
hit.offScreen = true
}
anchorHost = host
anchorDepth = len(stack)
}
}
}
if !voidHTMLElements[name] {
state.name = name
stack = append(stack, state)
}
case html.EndTagToken:
name, _, _ := tokenAttrs(z)
// Unclosed tags are ordinary in real content, so pop back to the
// nearest matching element rather than assuming balance.
popOpenElement(name)
if name == "a" {
closeAnchor()
}
}
}
closeAnchor()
sort.Strings(hit.hosts)
sort.Strings(hit.domains)
return hit
}
// voidHTMLElements never have an end tag, so they must not be pushed onto the
// open-element stack.
var voidHTMLElements = map[string]bool{
"area": true, "base": true, "br": true, "col": true, "embed": true,
"hr": true, "img": true, "input": true, "link": true, "meta": true,
"param": true, "source": true, "track": true, "wbr": true,
}
var rawTextHTMLElements = map[string]bool{
"script": true, "style": true, "textarea": true, "title": true,
}
// tokenAttrs returns the current token's lowercased tag name plus the two
// attributes this check reads.
func tokenAttrs(z *html.Tokenizer) (name, style, href string) {
raw, hasAttr := z.TagName()
name = strings.ToLower(string(raw))
for hasAttr {
var key, val []byte
key, val, hasAttr = z.TagAttr()
switch strings.ToLower(string(key)) {
case "style":
style = string(val)
case "href":
href = strings.TrimSpace(string(val))
}
}
return name, style, href
}
// absoluteLinkHost returns the host of a link that leaves the current page. A
// relative link cannot leave the site and is not one.
func absoluteLinkHost(href string) (string, bool) {
parsed, err := url.Parse(href)
if err != nil || parsed.Host == "" {
return "", false
}
switch strings.ToLower(parsed.Scheme) {
case "http", "https", "":
return normalizeHost(parsed.Hostname()), true
default:
return "", false
}
}
// registrableDomain canonicalizes a host to the public suffix plus one. IP and
// single-label hosts remain their own identity so sites served on either can
// still distinguish their own links from external ones.
func registrableDomain(host string) string {
host = normalizeHost(host)
if host == "" {
return ""
}
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
domain, err := publicsuffix.EffectiveTLDPlusOne(host)
if err != nil {
return host
}
return domain
}
func normalizeHost(host string) string {
host = strings.TrimSuffix(strings.ToLower(strings.TrimSpace(host)), ".")
if host == "" {
return ""
}
if ip := net.ParseIP(host); ip != nil {
return ip.String()
}
ascii, err := idna.Lookup.ToASCII(host)
if err != nil {
return ""
}
return strings.ToLower(ascii)
}
type cssDeclaration struct {
value string
important bool
}
type hiddenCSSState struct {
hidden bool
visibilitySet bool
visibilityHidden bool
offScreen bool
}
func parseHiddenCSSState(declarations string) hiddenCSSState {
effective := make(map[string]cssDeclaration)
parser := cssparser.NewParser(parse.NewInputString(declarations), true)
for {
grammar, _, propertyBytes := parser.Next()
if grammar == cssparser.ErrorGrammar {
break
}
if grammar != cssparser.DeclarationGrammar {
continue
}
property := strings.ToLower(cssUnescape(string(propertyBytes)))
var rawValue strings.Builder
for _, token := range parser.Values() {
rawValue.Write(token.Data)
}
value, important := cssDeclarationValue(strings.ToLower(cssUnescape(rawValue.String())))
previous, exists := effective[property]
if exists && previous.important && !important {
continue
}
effective[property] = cssDeclaration{value: value, important: important}
}
var state hiddenCSSState
if declaration, ok := effective["display"]; ok && declaration.value == "none" {
state.hidden = true
}
if declaration, ok := effective["visibility"]; ok {
switch declaration.value {
case "hidden", "collapse":
state.visibilitySet = true
state.visibilityHidden = true
case "visible", "initial":
state.visibilitySet = true
}
}
if declaration, ok := effective["opacity"]; ok && cssOpacityIsHidden(declaration.value) {
state.hidden = true
}
for _, property := range []string{
"text-indent", "left", "top", "right", "bottom", "margin-left", "margin-top",
} {
if declaration, ok := effective[property]; ok && cssLengthIsOffScreen(declaration.value) {
state.hidden = true
state.offScreen = true
break
}
}
return state
}
func cssUnescape(value string) string {
if !strings.ContainsRune(value, '\\') {
return value
}
var out strings.Builder
out.Grow(len(value))
for i := 0; i < len(value); i++ {
if value[i] != '\\' {
out.WriteByte(value[i])
continue
}
i++
if i >= len(value) {
break
}
if isCSSHex(value[i]) {
codePoint := uint32(0)
digits := 0
for i < len(value) && digits < 6 && isCSSHex(value[i]) {
codePoint = codePoint*16 + uint32(cssHexValue(value[i]))
i++
digits++
}
if codePoint == 0 || codePoint > utf8.MaxRune || 0xD800 <= codePoint && codePoint <= 0xDFFF {
out.WriteRune(utf8.RuneError)
} else {
out.WriteRune(rune(codePoint))
}
if i < len(value) && isCSSWhitespace(value[i]) {
if value[i] == '\r' && i+1 < len(value) && value[i+1] == '\n' {
i++
}
} else {
i--
}
continue
}
if value[i] == '\r' && i+1 < len(value) && value[i+1] == '\n' {
i++
continue
}
if value[i] == '\n' || value[i] == '\r' || value[i] == '\f' {
continue
}
r, size := utf8.DecodeRuneInString(value[i:])
out.WriteRune(r)
i += size - 1
}
return out.String()
}
func isCSSHex(value byte) bool {
return value >= '0' && value <= '9' || value >= 'a' && value <= 'f' || value >= 'A' && value <= 'F'
}
func cssHexValue(value byte) byte {
switch {
case value >= '0' && value <= '9':
return value - '0'
case value >= 'a' && value <= 'f':
return value - 'a' + 10
default:
return value - 'A' + 10
}
}
func isCSSWhitespace(value byte) bool {
return value == ' ' || value == '\t' || value == '\n' || value == '\r' || value == '\f'
}
func cssOpacityIsHidden(value string) bool {
value = strings.TrimSpace(value)
value = strings.TrimSpace(strings.TrimSuffix(value, "%"))
number, err := strconv.ParseFloat(value, 64)
return err == nil && number <= 0 || math.IsInf(number, -1)
}
func cssDeclarationValue(value string) (string, bool) {
value = strings.TrimSpace(value)
marker := strings.LastIndexByte(value, '!')
if marker < 0 || strings.TrimSpace(value[marker+1:]) != "important" {
return value, false
}
return strings.TrimSpace(value[:marker]), true
}
func cssLengthIsOffScreen(value string) bool {
value = strings.TrimSpace(value)
if strings.HasPrefix(value, "calc(") && strings.HasSuffix(value, ")") {
value = strings.TrimSpace(value[len("calc(") : len(value)-1])
}
unit := ""
for _, candidate := range []string{"vmin", "vmax", "rem", "px", "em", "vw", "vh", "%"} {
if strings.HasSuffix(value, candidate) {
unit = candidate
value = strings.TrimSpace(strings.TrimSuffix(value, candidate))
break
}
}
number, err := strconv.ParseFloat(value, 64)
if err != nil {
return math.IsInf(number, -1)
}
threshold := float64(offScreenPixels)
if unit == "em" || unit == "rem" {
threshold = offScreenFontUnits
}
return number <= -threshold
}
// hiddenLinkRow is one database row that carried a hidden link block.
type hiddenLinkRow struct {
label string
hit hiddenLinkHit
}
// checkWPHiddenLinks reports link blocks the page hides from its readers.
func checkWPHiddenLinks(user string, creds wpDBCreds, prefix string) []alert.Finding {
siteHosts, optionRows := hiddenLinkOptionRows(creds, prefix)
if len(siteHosts) == 0 {
// Every absolute link would look external. Report nothing rather than
// flood, and let the incomplete marker say why.
markCheckIncomplete(creds.queryCtx, "db_content")
return nil
}
rows := make([]hiddenLinkRow, 0, len(optionRows))
for _, row := range optionRows {
if hit := hiddenOffsiteLinkSamples(row, siteHosts); len(hit.hosts) > 0 {
rows = append(rows, hiddenLinkRow{label: "option " + strconv.Quote(row.label), hit: hit})
}
}
for _, row := range hiddenLinkPostRows(creds, prefix) {
if hit := hiddenOffsiteLinkSamples(row, siteHosts); len(hit.hosts) > 0 {
rows = append(rows, hiddenLinkRow{label: "post " + strconv.Quote(row.label), hit: hit})
}
}
return buildHiddenLinkFindings(user, creds, prefix, rows)
}
type hiddenLinkSource struct {
label string
markup string
tailMarkup string
valueBytes int
}
func hiddenOffsiteLinkSamples(source hiddenLinkSource, siteHosts []string) hiddenLinkHit {
if source.valueBytes > 0 && source.valueBytes <= maxHiddenLinkValueBytes && source.tailMarkup != "" {
overlap := len(source.markup) + len(source.tailMarkup) - source.valueBytes
if overlap >= 0 && overlap <= len(source.markup) && overlap <= len(source.tailMarkup) &&
source.markup[len(source.markup)-overlap:] == source.tailMarkup[:overlap] {
source.markup += source.tailMarkup[overlap:]
source.tailMarkup = ""
}
}
// Treat the two samples as separate fragments. Concatenating them could
// carry an unclosed hidden container across the omitted middle and turn a
// visible trailing link into a false positive.
hit := hiddenOffsiteLinksForSites(source.markup, siteHosts)
if source.tailMarkup == "" || source.tailMarkup == source.markup {
return hit
}
return mergeHiddenLinkHits(hit, hiddenOffsiteLinksForSites(source.tailMarkup, siteHosts))
}
func mergeHiddenLinkHits(left, right hiddenLinkHit) hiddenLinkHit {
merged := hiddenLinkHit{
offScreen: left.offScreen || right.offScreen,
multiDomain: left.multiDomain || right.multiDomain,
spammy: left.spammy || right.spammy,
}
for _, values := range [][]string{left.hosts, right.hosts} {
for _, value := range values {
if len(merged.hosts) == 0 || merged.hosts[len(merged.hosts)-1] != value {
merged.hosts = append(merged.hosts, value)
}
}
}
for _, values := range [][]string{left.domains, right.domains} {
for _, value := range values {
if len(merged.domains) == 0 || merged.domains[len(merged.domains)-1] != value {
merged.domains = append(merged.domains, value)
}
}
}
sort.Strings(merged.hosts)
merged.hosts = compactSortedStrings(merged.hosts)
sort.Strings(merged.domains)
merged.domains = compactSortedStrings(merged.domains)
return merged
}
func compactSortedStrings(values []string) []string {
if len(values) < 2 {
return values
}
out := values[:1]
for _, value := range values[1:] {
if value != out[len(out)-1] {
out = append(out, value)
}
}
return out
}
// runHiddenLinkCandidateQuery runs the statement built around a candidate
// condition. A regex timeout stops only commented-style matching: the failure
// stays recorded, so coverage remains incomplete, and the retry keeps plain
// and encoded styles covered.
func runHiddenLinkCandidateQuery(creds wpDBCreds, query func(commented bool) string) []string {
timeouts := creds.queryState.regexTimeoutCount()
rows := runMySQLQuery(creds, query(true))
if creds.queryState.regexTimeoutCount() == timeouts {
return rows
}
return runMySQLQuery(creds, query(false))
}
// hiddenLinkOptionRows reads the site address and the option rows worth
// parsing in one round trip.
func hiddenLinkOptionRows(creds wpDBCreds, prefix string) ([]string, []hiddenLinkSource) {
query := func(commented bool) string {
return fmt.Sprintf(
"(SELECT 'site' AS kind, option_name, LEFT(CAST(option_value AS BINARY), %d), "+
"'', OCTET_LENGTH(option_value), 'site' FROM %soptions "+
"WHERE option_name IN ('siteurl', 'home') LIMIT 4) UNION ALL "+
"(SELECT 'opt', option_name, LEFT(CAST(option_value AS BINARY), %d), "+
"RIGHT(CAST(option_value AS BINARY), %d), OCTET_LENGTH(option_value), 'opt' FROM %soptions "+
"WHERE %s LIMIT %d)",
maxHiddenLinkSiteURLBytes, prefix, maxHiddenLinkSampleBytes, maxHiddenLinkSampleBytes, prefix,
hiddenLinkCandidateCondition("option_value", commented), maxHiddenLinkRows+1)
}
var siteHosts []string
seenSiteHosts := make(map[string]bool)
var out []hiddenLinkSource
rows := runHiddenLinkCandidateQuery(creds, query)
optionRowsSeen := 0
truncated := false
for _, line := range rows {
// Column separators are literal tabs while tabs and newlines inside a
// value remain batch escapes. Split that transport form before decoding
// individual columns or an embedded tab becomes indistinguishable from
// a separator.
parts := strings.SplitN(strings.TrimRight(line, "\r\n"), "\t", 6)
if len(parts) != 6 {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
kind := strings.TrimSpace(parts[0])
if kind != "site" {
optionRowsSeen++
if optionRowsSeen > maxHiddenLinkRows {
if !truncated {
markCheckIncomplete(creds.queryCtx, "db_content")
truncated = true
}
continue
}
}
name := strings.TrimSpace(mysqlclient.BatchUnescape(parts[1]))
encodedValue := parts[2]
value := mysqlclient.BatchUnescape(encodedValue)
valueBytes, err := strconv.Atoi(strings.TrimSpace(parts[4]))
if err != nil || valueBytes < 0 {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if kind == "site" {
if valueBytes > maxHiddenLinkSiteURLBytes || len(value) != valueBytes {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if reason, _ := siteURLPoisonReason(encodedValue); reason != "" {
continue
}
if parsed, err := url.Parse(strings.TrimSpace(value)); err == nil {
host := normalizeHost(parsed.Hostname())
if registrableDomain(host) != "" && !seenSiteHosts[host] {
seenSiteHosts[host] = true
siteHosts = append(siteHosts, host)
}
}
continue
}
tailMarkup := mysqlclient.BatchUnescape(parts[3])
expectedSampleBytes := valueBytes
if expectedSampleBytes > maxHiddenLinkSampleBytes {
expectedSampleBytes = maxHiddenLinkSampleBytes
}
if len(value) != expectedSampleBytes || len(tailMarkup) != expectedSampleBytes {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if valueBytes > maxHiddenLinkValueBytes {
markCheckIncomplete(creds.queryCtx, "db_content")
}
out = append(out, hiddenLinkSource{
label: name, markup: value, tailMarkup: tailMarkup, valueBytes: valueBytes,
})
}
return siteHosts, out
}
func hiddenLinkPostRows(creds wpDBCreds, prefix string) []hiddenLinkSource {
query := func(commented bool) string {
return fmt.Sprintf(
"SELECT ID, LEFT(CAST(post_content AS BINARY), %d), "+
"RIGHT(CAST(post_content AS BINARY), %d), OCTET_LENGTH(post_content), 'post' "+
"FROM %sposts WHERE post_status = 'publish' "+
"AND post_type NOT IN (%s) AND %s LIMIT %d",
maxHiddenLinkSampleBytes, maxHiddenLinkSampleBytes, prefix,
nonScannablePostTypesSQLList(), hiddenLinkCandidateCondition("post_content", commented), maxHiddenLinkRows+1)
}
rows := runHiddenLinkCandidateQuery(creds, query)
if len(rows) > maxHiddenLinkRows {
markCheckIncomplete(creds.queryCtx, "db_content")
rows = rows[:maxHiddenLinkRows]
}
out := make([]hiddenLinkSource, 0, len(rows))
for _, line := range rows {
parts := strings.SplitN(strings.TrimRight(line, "\r\n"), "\t", 5)
if len(parts) != 5 {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
markup := mysqlclient.BatchUnescape(parts[1])
tailMarkup := mysqlclient.BatchUnescape(parts[2])
valueBytes, err := strconv.Atoi(strings.TrimSpace(parts[3]))
if err != nil || valueBytes < 0 {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
expectedSampleBytes := valueBytes
if expectedSampleBytes > maxHiddenLinkSampleBytes {
expectedSampleBytes = maxHiddenLinkSampleBytes
}
if len(markup) != expectedSampleBytes || len(tailMarkup) != expectedSampleBytes {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
if valueBytes > maxHiddenLinkValueBytes {
markCheckIncomplete(creds.queryCtx, "db_content")
}
out = append(out, hiddenLinkSource{
label: strings.TrimSpace(mysqlclient.BatchUnescape(parts[0])),
markup: markup,
tailMarkup: tailMarkup,
valueBytes: valueBytes,
})
}
return out
}
// buildHiddenLinkFindings grades the rows. An off-canvas container has no
// benign reading and is reported on its own. display:none does have one --
// themes hide panels that link out -- so it is reported only once the block
// points at several unrelated domains or carries spam vocabulary.
func buildHiddenLinkFindings(user string, creds wpDBCreds, prefix string, rows []hiddenLinkRow) []alert.Finding {
var reported []hiddenLinkRow
offScreen := false
hosts := make(map[string]bool)
domains := make(map[string]bool)
for _, row := range rows {
if !row.hit.offScreen && !row.hit.multiDomain && !row.hit.spammy {
continue
}
reported = append(reported, row)
if row.hit.offScreen {
offScreen = true
}
for _, host := range row.hit.hosts {
hosts[host] = true
}
for _, domain := range row.hit.domains {
domains[domain] = true
}
}
if len(reported) == 0 {
return nil
}
named := make([]string, 0, len(hosts))
for host := range hosts {
named = append(named, host)
}
sort.Strings(named)
shownHosts := named
if len(shownHosts) > maxHiddenLinkHostsShown {
shownHosts = shownHosts[:maxHiddenLinkHostsShown]
}
labels := make([]string, 0, len(reported))
for _, row := range reported {
if len(labels) >= maxHiddenLinkRowsShown {
break
}
labels = append(labels, row.label)
}
details := []string{
"A container the page hides from readers wraps links to other domains. " +
"Crawlers still follow them, which is the point: the site's ranking " +
"is lent to the linked domains without a visitor ever seeing it.",
hiddenLinkSample("Linked hosts", shownHosts, len(named)),
hiddenLinkSample("Rows", labels, len(reported)),
}
// An off-canvas container is cloaking by construction; a merely hidden one
// needed corroboration to be reported at all, which is a weaker claim.
severity := alert.Warning
if offScreen {
severity = alert.High
}
// Concealment strength and the complete destination set identify this
// site's finding. Row growth and display samples must not re-alert, while
// stronger concealment must survive a prior baseline or dismissal.
identity := append([]string{fmt.Sprintf("off-screen=%t", offScreen)}, named...)
return []alert.Finding{{
Severity: severity,
Check: "db_hidden_link_injection",
Message: fmt.Sprintf("%d WordPress rows hide outbound links to %d hosts across %d domains (account: %s)",
len(reported), len(named), len(domains), user),
Details: dbContentFindingDetails(creds, prefix, details...),
DedupKey: dbContentDedupKey(user, creds, prefix, identity...),
}}
}
func hiddenLinkSample(label string, shown []string, total int) string {
if total > len(shown) {
label += fmt.Sprintf(" (showing %d of %d)", len(shown), total)
}
return label + ": " + strings.Join(shown, ", ")
}
package checks
import (
"context"
"fmt"
"path/filepath"
"regexp"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// Joomla database content scanner.
//
// Discovery: glob /home/*/public_html/configuration.php and
// verify the file contains `class JConfig` -- the canonical marker
// for a Joomla site, distinguishing it from PHP files that happen
// to share the configuration.php filename. Credentials are read
// via regex over public-property assignments (`public $host = ...;`)
// rather than PHP eval; the parser ignores anything outside the
// JConfig class body.
//
// Scanned tables (all prefixed; the prefix is operator-controlled
// via configuration.php's $dbprefix and defaults to `jos_` /
// `<random>_` on fresh installs):
//
// <prefix>extensions params blob -- live_site, sitename,
// offline_message are common hijack targets
// <prefix>content article body for malware patterns
// <prefix>users user table; joined with
// <prefix>user_usergroup_map to find rogue Super Users
// (group_id = 8 in vanilla Joomla)
//
// Three new finding categories. CMS-explicit names so operators
// running mixed-CMS hosts can suppress per-CMS:
//
// joomla_extensions_injection (Critical) -- malware pattern in
// an extension's params
// joomla_content_injection (Critical) -- malware pattern in
// an article body
// joomla_admin_injection (Critical) -- rogue Super User
// account
// jConfigCredsPattern parses a `public $foo = 'value';` line
// (single OR double quotes). Anchored to the start of the line so
// arbitrary text inside string literals further along can't be
// misread as a credential.
var jConfigCredsPattern = regexp.MustCompile(`^\s*public\s+\$(\w+)\s*=\s*['"]([^'"]*)['"]\s*;`)
// joomlaSuperUserGroupID is the canonical group id for "Super Users"
// in vanilla Joomla 3+. Operators on hardened installs may have
// renumbered; the spec narrows to 8 for v1.
const joomlaSuperUserGroupID = 8
// jConfigCreds carries the credentials extracted from a Joomla
// configuration.php. Mirrors the wpDBCreds shape so the existing
// runMySQLQuery / mysqlSchemaLiteral helpers can be reused, but
// kept distinct because Joomla configuration.php and WordPress
// wp-config.php are not interchangeable.
type jConfigCreds struct {
// ctx ties every query for this install to the runner's deadline.
ctx context.Context
dbName string
dbUser string
dbPass string
dbHost string
dbPrefix string
path string
queryState *dbQueryState
}
// asWPDBCreds returns the equivalent wpDBCreds for runMySQLQuery
// reuse. The `multisite` field is irrelevant here; the existing
// mysql client wrapper does not look at it.
func (c jConfigCreds) asWPDBCreds() wpDBCreds {
return wpDBCreds{
dbName: c.dbName,
dbUser: c.dbUser,
dbPass: c.dbPass,
dbHost: c.dbHost,
tablePrefix: c.dbPrefix,
queryCtx: c.ctx,
queryOwner: "db_content_joomla",
queryState: c.queryState,
}
}
// CheckJoomlaContent scans every Joomla installation under
// /home/*/public_html for malware-pattern matches in the three
// canonical attacker-touched tables. Mirrors the structure of
// CheckDatabaseContent without sharing code -- the credentials and
// table layout differ enough that a generic dispatcher is more
// abstraction than this point in the codebase needs.
func CheckJoomlaContent(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
configs := cmsDiscover(ctx, "db_content_joomla", "*/public_html/configuration.php", "*/*/configuration.php")
if len(configs) == 0 {
return nil
}
// Rank by mtime desc so recently touched Joomla installs are processed
// first when the check timeout cuts iteration short.
for _, path := range rankCMSConfigs(ctx, "db_content_joomla", configs, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
return findings
}
findings = append(findings, scanJoomlaInstall(ctx, path, store)...)
}
return findings
}
// scanJoomlaInstall scans one discovered install and stamps its findings
// with the owner resolved from the configuration path. The display label
// stays as before; an install outside every account root is not stamped.
func scanJoomlaInstall(ctx context.Context, path string, store *state.Store) []alert.Finding {
matched, err := looksLikeJoomlaConfig(ctx, path)
if err != nil {
markCheckIncomplete(ctx, "db_content_joomla")
return nil
}
if !matched {
return nil
}
account := extractUser(filepath.Dir(path))
creds, err := parseJConfig(ctx, path)
if err != nil || creds.dbName == "" || creds.dbUser == "" {
markCheckIncomplete(ctx, "db_content_joomla")
return nil
}
creds.ctx = ctx
creds.queryState = new(dbQueryState)
prefix := creds.dbPrefix
if prefix == "" {
prefix = "jos_"
}
var findings []alert.Finding
findings = append(findings, scanJoomlaExtensions(account, creds, prefix)...)
findings = append(findings, scanJoomlaContent(account, creds, prefix)...)
findings = append(findings, scanJoomlaSuperUsers(store, account, creds, prefix)...)
if owner, ok := installOwner(path); ok {
findings = stampTenantIDIfEmpty(findings, owner)
}
return findings
}
// A defaced configuration can still identify a Joomla installation. Its
// credentials are parsed without executing any PHP.
func looksLikeJoomlaConfig(ctx context.Context, path string) (bool, error) {
data, err := readCMSConfig(ctx, path)
if err != nil {
return false, err
}
return strings.Contains(strings.ToLower(string(data)), "class jconfig"), nil
}
// parseJConfig reads configuration.php and pulls credentials out of
// the public-property assignments. Lines outside the class body
// (PHP comments, namespaced statements, etc.) are tolerated
// silently because the regex is line-anchored to "public $foo = ...".
func parseJConfig(ctx context.Context, path string) (jConfigCreds, error) {
creds := jConfigCreds{path: path}
data, err := readCMSConfig(ctx, path)
if err != nil {
return creds, err
}
for _, line := range strings.Split(string(data), "\n") {
m := jConfigCredsPattern.FindStringSubmatch(line)
if m == nil {
continue
}
switch strings.ToLower(m[1]) {
case "host":
creds.dbHost = m[2]
case "user":
creds.dbUser = m[2]
case "password":
creds.dbPass = m[2]
case "db":
creds.dbName = m[2]
case "dbprefix":
creds.dbPrefix = m[2]
}
}
if creds.dbHost == "" {
creds.dbHost = "localhost"
}
return creds, nil
}
// scanJoomlaExtensions queries the extensions table for params
// blobs that match the malware-pattern pre-filter, then applies a
// Go-side post-filter to drop rows whose only match was <script>
// LIKE noise (legitimate analytics embeds in extension params).
//
// Two-phase classifier:
//
// 1. SQL pre-filter via LIKE keeps the result set bounded -- a
// vanilla Joomla #__extensions table has ~50 rows, but a
// plugin-heavy install can exceed 200, and we don't want to
// pull every params blob into the daemon.
//
// 2. Go classifyMalwareRow re-checks each pattern individually
// against the full body, applying the same requiresExternalScript
// filter the WP scanner uses for wp_options. Strict predicate
// here (hasMaliciousExternalScript) because params is config
// storage.
func scanJoomlaExtensions(account string, creds jConfigCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"SELECT name, params FROM %sextensions WHERE %s",
prefix, paramsLikeClause("params"))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
name, body := splitTabRow(row)
if name == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, false)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "joomla_extensions_injection",
Message: fmt.Sprintf("Joomla extension params injection on %s: %s (%s)", account, name, desc),
Details: fmt.Sprintf("Account: %s\nExtension: %s\nMatch: %s", account, name, desc),
})
}
return findings
}
// scanJoomlaContent queries article bodies (introtext) for malware
// patterns. Same two-phase classifier as scanJoomlaExtensions but
// uses the looser post-filter (hasMaliciousExternalScriptInPost)
// because articles are author-written and may carry pre-TLS-era
// embeds the strict predicate would flag on scheme alone.
//
// fulltext_ is not scanned in v1: it's almost never populated on
// modern Joomla installs (the read-more split is a layout choice
// most templates don't bother with), and adding it doubles the
// query cost for marginal coverage. Follow-up if operators see
// missed detections.
func scanJoomlaContent(account string, creds jConfigCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"SELECT id, title, introtext FROM %scontent WHERE %s",
prefix, paramsLikeClause("introtext"))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
fields := strings.SplitN(row, "\t", 3)
if len(fields) < 3 {
continue
}
id, title, body := fields[0], fields[1], fields[2]
sev, desc, ok := classifyMalwareRow(body, true)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "joomla_content_injection",
Message: fmt.Sprintf("Joomla article content injection on %s: id=%s title=%q (%s)", account, id, title, desc),
Details: fmt.Sprintf("Account: %s\nArticle ID: %s\nTitle: %s\nMatch: %s", account, id, title, desc),
})
}
return findings
}
// classifyMalwareRow walks dbMalwarePatterns against body and
// returns the strongest pattern match that survives the
// requiresExternalScript filter. Returns ok=false when nothing
// genuine matched -- the caller skips that row entirely.
//
// inPostContext switches between the strict
// hasMaliciousExternalScript (for config-storage rows like Joomla
// extension params or Drupal config) and the looser
// hasMaliciousExternalScriptInPost (for author-written article
// content). Mirrors how the WP scanner picks its predicate per
// table. Shared by the Joomla and Drupal scanners; the function
// lives in dbscan_joomla.go for historical reasons (added with
// the Joomla scanner; renamed when Drupal also needed it).
func classifyMalwareRow(body string, inPostContext bool) (alert.Severity, string, bool) {
if body == "" {
return 0, "", false
}
lower := strings.ToLower(body)
var bestSev alert.Severity
var bestDesc string
matched := false
for _, p := range dbMalwarePatterns {
if !strings.Contains(lower, strings.ToLower(p.pattern)) {
continue
}
if p.requiresExternalScript {
ok := hasMaliciousExternalScript(body)
if inPostContext {
ok = hasMaliciousExternalScriptInPost(body)
}
if !ok {
continue
}
}
if !matched || p.severity > bestSev {
bestSev = p.severity
bestDesc = p.desc
}
matched = true
}
return bestSev, bestDesc, matched
}
// scanJoomlaSuperUsers detects rogue accounts in the Super Users
// group (group_id = 8 by default). The two-table join is necessary
// because Joomla stores group membership separately from the user
// row; a single rogue admin shows up only when the join fires.
func scanJoomlaSuperUsers(store *state.Store, account string, creds jConfigCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"SELECT u.id, u.username, u.email FROM %susers u JOIN %suser_usergroup_map m ON u.id = m.user_id WHERE m.group_id = %d",
prefix, prefix, joomlaSuperUserGroupID)
rows, complete := runCMSQuery(creds.asWPDBCreds(), query)
// The legitimate site admin is in this set too: the store baseline
// keeps known Super Users quiet and reports only a newcomer.
dbCreds := creds.asWPDBCreds()
dbCreds.tablePrefix = prefix
return cmsAdminFindings(store, "joomla", "joomla_admin_injection", account, dbCreds, rows, complete, func(fields []string) (string, string) {
return fmt.Sprintf("Joomla Super User account on %s: %s", account, fields[0]),
fmt.Sprintf("Account: %s\nRow: %s\nReview: confirm this is the legitimate site administrator.", account, strings.Join(fields, "\t"))
})
}
// paramsLikeClause builds an OR'd LIKE clause over the supplied
// columns and the existing dbMalwarePatterns list. The patterns
// are escaped for MySQL string literal syntax (single quotes
// doubled, backslashes doubled).
//
// We use plain LIKE rather than full-text search so the query
// runs against any MySQL configuration -- some shared-hosting
// instances have ngram FTS off, and we don't want to depend on
// it for a security scan.
func paramsLikeClause(columns ...string) string {
if len(columns) == 0 {
return "1=0"
}
var clauses []string
for _, col := range columns {
for _, p := range dbMalwarePatterns {
lit := mysqlEscapeForLike(p.pattern)
clauses = append(clauses, fmt.Sprintf("%s LIKE '%%%s%%'", col, lit))
}
}
return strings.Join(clauses, " OR ")
}
// mysqlEscapeForLike escapes a literal for use inside a single-quoted
// MySQL LIKE pattern. Only `'` and `\` need escaping; LIKE's `%` and
// `_` are intentionally left alone because the malware-pattern list
// uses literal substrings and never SQL wildcards.
func mysqlEscapeForLike(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, `'`, `\'`)
return s
}
package checks
import (
"context"
"encoding/xml"
"fmt"
"path/filepath"
"regexp"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// Magento database content scanner.
//
// Single file covering both major versions: M1 (Magento 1.x, EOL'd
// June 2020 but still found on legacy hosts) and M2 (Magento 2.x).
// The two versions share table names and content shape but disagree
// on configuration file format -- M1 stores credentials in
// app/etc/local.xml as XML, M2 stores them in app/etc/env.php as
// PHP arrays.
//
// Discovery is mutually exclusive: a host either has app/etc/env.php
// (M2) or app/etc/local.xml (M1), never both. We probe in that order
// because M2 is the actively maintained version and we want fresh
// installs to be picked up first.
//
// Scanned tables (identical between M1 and M2):
//
// core_config_data (path, value) -- settings.
// web/unsecure/base_url is the
// canonical hijack target;
// attackers redirect the storefront
// by overwriting it.
// catalog_product_entity_text product description text;
// spam-injection vector for SEO
// cms_block + cms_page CMS-managed content; same
// pattern as the product text scan
// admin_user backend administrator accounts.
// One Warning per row -- legitimate
// admin shows up too.
//
// Three new finding categories with CMS-explicit names:
// magento_settings_injection, magento_content_injection,
// magento_admin_injection.
// magentoCreds carries the parsed connection details. Mirrors
// jConfigCreds / drupalCreds; the version field tells the scanner
// which discovery path produced the creds (useful for messages).
type magentoCreds struct {
// ctx ties every query for this install to the runner's deadline.
ctx context.Context
dbName string
dbUser string
dbPass string
dbHost string
dbPrefix string
version string // "M1" | "M2"
path string
queryState *dbQueryState
}
func (c magentoCreds) asWPDBCreds() wpDBCreds {
return wpDBCreds{
dbName: c.dbName,
dbUser: c.dbUser,
dbPass: c.dbPass,
dbHost: c.dbHost,
tablePrefix: c.dbPrefix,
queryCtx: c.ctx,
queryOwner: "db_content_magento",
queryState: c.queryState,
}
}
// magentoM1XMLRoot is the minimum struct surface encoding/xml needs
// to extract the connection block out of a Magento 1.x local.xml.
// CDATA wrapping is transparent to the decoder; both the bare and
// CDATA-wrapped forms produce the same string value.
type magentoM1XMLRoot struct {
XMLName xml.Name `xml:"config"`
Connection magentoM1XMLConnBlock `xml:"global>resources>default_setup>connection"`
Resources magentoM1XMLResources `xml:"global>resources>db"`
}
type magentoM1XMLConnBlock struct {
Host string `xml:"host"`
Username string `xml:"username"`
Password string `xml:"password"`
DBName string `xml:"dbname"`
}
type magentoM1XMLResources struct {
TablePrefix string `xml:"table_prefix"`
}
// M2 env.php is a PHP file returning an array; we extract by regex
// rather than wiring a PHP parser. The patterns match the canonical
// nested-array layout that vendor/magento installers produce; hand-
// rolled env.php files with reordered keys still parse because
// each pattern matches independently.
var (
magentoM2HostRe = regexp.MustCompile(`['"]host['"]\s*=>\s*['"]([^'"]+)['"]`)
magentoM2UserRe = regexp.MustCompile(`['"]username['"]\s*=>\s*['"]([^'"]+)['"]`)
magentoM2PassRe = regexp.MustCompile(`['"]password['"]\s*=>\s*['"]([^'"]+)['"]`)
magentoM2DBRe = regexp.MustCompile(`['"]dbname['"]\s*=>\s*['"]([^'"]+)['"]`)
magentoM2PrefixRe = regexp.MustCompile(`['"]table_prefix['"]\s*=>\s*['"]([^'"]*)['"]`)
)
// CheckMagentoContent discovers Magento installs (M1 + M2) and
// scans the four canonical tables. Mirrors CheckJoomlaContent and
// CheckDrupalContent without sharing code -- the version-branching
// is local to this scanner.
//
// Accounts that produced creds via the M2 (env.php) path are
// tracked in seenAccounts so the M1 fallback doesn't re-scan a
// host that's already been processed -- including the common case
// where M2 found zero malware findings (a clean install). Without
// this, a half-migrated host with both env.php and stale local.xml
// would scan the database twice with different credential sets.
// scanMagentoInstall scans one discovered install (either configuration
// layout) and stamps its findings with the owner resolved from the
// configuration path. The display label stays as before; an install
// outside every account root is not stamped.
func scanMagentoInstall(ctx context.Context, path, account string, creds magentoCreds, store *state.Store) []alert.Finding {
creds.ctx = ctx
creds.queryState = new(dbQueryState)
findings := scanMagentoAll(store, account, creds)
if owner, ok := installOwner(path); ok {
findings = stampTenantIDIfEmpty(findings, owner)
}
return findings
}
func CheckMagentoContent(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
seenAccounts := map[string]bool{}
// M2 discovery first (active version). Rank by mtime desc so recently
// touched installs are processed first when the check timeout cuts
// iteration short.
m2Files := cmsDiscover(ctx, "db_content_magento", "*/public_html/app/etc/env.php", "*/*/app/etc/env.php")
for _, path := range rankCMSConfigs(ctx, "db_content_magento", m2Files, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
return findings
}
account := magentoAccountFromPath(path)
creds, err := parseMagentoM2(ctx, path)
if err != nil || creds.dbName == "" || creds.dbUser == "" {
markCheckIncomplete(ctx, "db_content_magento")
continue
}
seenAccounts[account] = true
findings = append(findings, scanMagentoInstall(ctx, path, account, creds, store)...)
}
// M1 fallback for hosts where env.php is absent or unparseable.
m1Files := cmsDiscover(ctx, "db_content_magento", "*/public_html/app/etc/local.xml", "*/*/app/etc/local.xml")
for _, path := range rankCMSConfigs(ctx, "db_content_magento", m1Files, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
return findings
}
account := magentoAccountFromPath(path)
if seenAccounts[account] {
continue
}
creds, err := parseMagentoM1(ctx, path)
if err != nil || creds.dbName == "" || creds.dbUser == "" {
markCheckIncomplete(ctx, "db_content_magento")
continue
}
findings = append(findings, scanMagentoInstall(ctx, path, account, creds, store)...)
}
return findings
}
// magentoAccountFromPath strips the conventional cPanel prefix
// (/home/<account>/public_html/app/etc/...) down to the account
// component.
func magentoAccountFromPath(path string) string {
// /home/<account>/public_html/app/etc/<file> -- four Dirs up.
cur := path
for i := 0; i < 4; i++ {
cur = filepath.Dir(cur)
}
return extractUser(cur)
}
// parseMagentoM1 reads local.xml and extracts the connection block.
// A read or XML error keeps the installation out of the completed scan.
func parseMagentoM1(ctx context.Context, path string) (magentoCreds, error) {
creds := magentoCreds{path: path, version: "M1"}
data, err := readCMSConfig(ctx, path)
if err != nil {
return creds, err
}
var root magentoM1XMLRoot
if err := xml.Unmarshal(data, &root); err != nil {
return creds, err
}
creds.dbHost = strings.TrimSpace(root.Connection.Host)
creds.dbUser = strings.TrimSpace(root.Connection.Username)
creds.dbPass = strings.TrimSpace(root.Connection.Password)
creds.dbName = strings.TrimSpace(root.Connection.DBName)
creds.dbPrefix = strings.TrimSpace(root.Resources.TablePrefix)
if creds.dbHost == "" {
creds.dbHost = "localhost"
}
return creds, nil
}
// parseMagentoM2 reads env.php and pulls credentials out via the
// field-level regexes. Unlike Drupal we have a stable nested-array
// layout to match against (the one Magento Setup writes), but to
// stay robust against operator-edited env.php files we match each
// key independently.
func parseMagentoM2(ctx context.Context, path string) (magentoCreds, error) {
creds := magentoCreds{path: path, version: "M2"}
data, err := readCMSConfig(ctx, path)
if err != nil {
return creds, err
}
body := string(data)
if m := magentoM2HostRe.FindStringSubmatch(body); m != nil {
creds.dbHost = m[1]
}
if m := magentoM2UserRe.FindStringSubmatch(body); m != nil {
creds.dbUser = m[1]
}
if m := magentoM2PassRe.FindStringSubmatch(body); m != nil {
creds.dbPass = m[1]
}
if m := magentoM2DBRe.FindStringSubmatch(body); m != nil {
creds.dbName = m[1]
}
if m := magentoM2PrefixRe.FindStringSubmatch(body); m != nil {
creds.dbPrefix = m[1]
}
if creds.dbHost == "" {
creds.dbHost = "localhost"
}
return creds, nil
}
// scanMagentoAll runs the four scan paths against one Magento
// install. Helper exists so M1 and M2 dispatch through the same
// post-creds code path.
func scanMagentoAll(store *state.Store, account string, creds magentoCreds) []alert.Finding {
var findings []alert.Finding
findings = append(findings, scanMagentoSettings(account, creds)...)
findings = append(findings, scanMagentoContent(account, creds, "catalog_product_entity_text", "value")...)
findings = append(findings, scanMagentoContent(account, creds, "cms_block", "content")...)
findings = append(findings, scanMagentoContent(account, creds, "cms_page", "content")...)
findings = append(findings, scanMagentoAdmins(store, account, creds)...)
return findings
}
// scanMagentoSettings looks for malware patterns in core_config_data
// values. The path column carries dotted-namespace identifiers
// (web/unsecure/base_url, design/header/welcome, etc.) so we keep
// it in the finding details for triage.
func scanMagentoSettings(account string, creds magentoCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT path, value FROM %score_config_data WHERE %s",
creds.dbPrefix, paramsLikeClause("value"))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
cfgPath, body := splitTabRow(row)
if cfgPath == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, false)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "magento_settings_injection",
Message: fmt.Sprintf("Magento %s settings injection on %s: %s (%s)", creds.version, account, cfgPath, desc),
Details: fmt.Sprintf("Account: %s\nConfig path: %s\nMatch: %s", account, cfgPath, desc),
})
}
return findings
}
// scanMagentoContent walks one CMS table (catalog_product_entity_text,
// cms_block, cms_page) for malware patterns. The looser
// hasMaliciousExternalScriptInPost predicate applies because the
// tables carry author-written content.
func scanMagentoContent(account string, creds magentoCreds, table, valueCol string) []alert.Finding {
idCol := "row_id"
switch table {
case "catalog_product_entity_text":
idCol = "entity_id"
case "cms_block":
idCol = "block_id"
case "cms_page":
idCol = "page_id"
}
query := fmt.Sprintf(
"SELECT %s, %s FROM %s%s WHERE %s",
idCol, valueCol, creds.dbPrefix, table, paramsLikeClause(valueCol))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
id, body := splitTabRow(row)
if id == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, true)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "magento_content_injection",
Message: fmt.Sprintf("Magento %s content injection on %s: %s id=%s (%s)", creds.version, account, table, id, desc),
Details: fmt.Sprintf("Account: %s\nTable: %s\nRow id: %s\nMatch: %s", account, table, id, desc),
})
}
return findings
}
// scanMagentoAdmins enumerates the admin_user table. Rows include
// the legitimate site admin -- one Warning per row, operator
// review territory.
func scanMagentoAdmins(store *state.Store, account string, creds magentoCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT user_id, username, email FROM %sadmin_user",
creds.dbPrefix)
rows, complete := runCMSQuery(creds.asWPDBCreds(), query)
return cmsAdminFindings(store, "magento", "magento_admin_injection", account, creds.asWPDBCreds(), rows, complete, func(fields []string) (string, string) {
return fmt.Sprintf("Magento %s admin account on %s: user_id=%s", creds.version, account, fields[0]),
fmt.Sprintf("Account: %s\nRow: %s\nReview: confirm this is the legitimate site administrator.", account, strings.Join(fields, "\t"))
})
}
package checks
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// OpenCart database content scanner.
//
// Discovery: glob /home/*/public_html/config.php and confirm both
// it AND /home/*/public_html/admin/config.php contain
// `define('DB_DRIVER'`. The admin-side config.php pair is the
// canonical OpenCart marker -- plain PHP sites carry a config.php
// in the document root that's nothing to do with OpenCart.
//
// Credentials use PHP define() constants with OC-specific names:
//
// DB_HOSTNAME DB_USERNAME DB_PASSWORD DB_DATABASE DB_PREFIX
//
// Reuses the existing extractDefine helper from dbscan.go (the WP
// scanner already understands this shape). DB_PREFIX defaults to
// "oc_" on vanilla installs.
//
// Scanned tables (all prefixed):
//
// <prefix>setting k/v pairs; values are JSON
// blobs. config_url / config_ssl
// are the canonical hijack
// targets for storefront redirect.
// <prefix>product_description product description text
// <prefix>information_description CMS-managed information pages
// <prefix>user admin/staff accounts.
// Customer accounts live in the
// oc_customer table, not here;
// every oc_user row is admin-shaped.
//
// Three new finding categories:
// opencart_settings_injection, opencart_content_injection,
// opencart_admin_injection.
type opencartCreds struct {
// ctx ties every query for this install to the runner's deadline.
ctx context.Context
dbName string
dbUser string
dbPass string
dbHost string
dbPrefix string
path string
queryState *dbQueryState
}
func (c opencartCreds) asWPDBCreds() wpDBCreds {
return wpDBCreds{
dbName: c.dbName,
dbUser: c.dbUser,
dbPass: c.dbPass,
dbHost: c.dbHost,
tablePrefix: c.dbPrefix,
queryCtx: c.ctx,
queryOwner: "db_content_opencart",
queryState: c.queryState,
}
}
// CheckOpenCartContent discovers OpenCart installs and scans the
// four canonical attacker-touched tables. Mirrors the other CMS
// scanners; the discovery and credentials parsing are the only
// OC-specific bits.
// scanOpenCartInstall scans one discovered install and stamps its findings
// with the owner resolved from the configuration path. The display label
// stays as before; an install outside every account root is not stamped.
func scanOpenCartInstall(ctx context.Context, path string, store *state.Store) []alert.Finding {
matched, err := looksLikeOpenCart(ctx, path)
if err != nil {
markCheckIncomplete(ctx, "db_content_opencart")
return nil
}
if !matched {
return nil
}
account := extractUser(filepath.Dir(path))
creds, err := parseOpenCartConfig(ctx, path)
if err != nil || creds.dbName == "" || creds.dbUser == "" {
markCheckIncomplete(ctx, "db_content_opencart")
return nil
}
creds.ctx = ctx
creds.queryState = new(dbQueryState)
prefix := creds.dbPrefix
if prefix == "" {
prefix = "oc_"
}
creds.dbPrefix = prefix
var findings []alert.Finding
findings = append(findings, scanOpenCartSettings(account, creds)...)
findings = append(findings, scanOpenCartContentTable(account, creds, "product_description", "description")...)
findings = append(findings, scanOpenCartContentTable(account, creds, "information_description", "description")...)
findings = append(findings, scanOpenCartAdmins(store, account, creds)...)
if owner, ok := installOwner(path); ok {
findings = stampTenantIDIfEmpty(findings, owner)
}
return findings
}
func CheckOpenCartContent(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
configs := cmsDiscover(ctx, "db_content_opencart", "*/public_html/config.php", "*/*/config.php")
if len(configs) == 0 {
return nil
}
// Rank by mtime desc so recently touched OpenCart installs are processed
// first when the check timeout cuts iteration short.
for _, path := range rankCMSConfigs(ctx, "db_content_opencart", configs, accountScanMaxFiles(ctx, cfg)) {
if ctx.Err() != nil {
return findings
}
findings = append(findings, scanOpenCartInstall(ctx, path, store)...)
}
return findings
}
// looksLikeOpenCart confirms both config.php files exist and both
// reference DB_DRIVER. The admin-side file is what distinguishes
// OpenCart from arbitrary PHP sites that happen to ship a
// config.php at the document root.
func looksLikeOpenCart(ctx context.Context, rootConfig string) (bool, error) {
matched, err := configContainsDBDriver(ctx, rootConfig)
if err != nil || !matched {
return false, err
}
adminConfig := filepath.Join(filepath.Dir(rootConfig), "admin", "config.php")
matched, err = configContainsDBDriver(ctx, adminConfig)
if errors.Is(err, os.ErrNotExist) {
return false, nil
}
return matched, err
}
func configContainsDBDriver(ctx context.Context, path string) (bool, error) {
data, err := readCMSConfig(ctx, path)
if err != nil {
return false, err
}
return strings.Contains(string(data), "DB_DRIVER"), nil
}
// parseOpenCartConfig extracts the DB_* defines from a config.php.
// Reuses the WP scanner's extractDefine helper -- the OC defines
// have the same `define('KEY', 'value')` shape WP uses, and the
// helper already strips comments and walks past the key's closing
// quote correctly.
func parseOpenCartConfig(ctx context.Context, path string) (opencartCreds, error) {
creds := opencartCreds{path: path}
data, err := readCMSConfig(ctx, path)
if err != nil {
return creds, err
}
for _, line := range strings.Split(string(data), "\n") {
if v := extractDefine(line, "DB_HOSTNAME"); v != "" {
creds.dbHost = v
}
if v := extractDefine(line, "DB_USERNAME"); v != "" {
creds.dbUser = v
}
if v := extractDefine(line, "DB_PASSWORD"); v != "" {
creds.dbPass = v
}
if v := extractDefine(line, "DB_DATABASE"); v != "" {
creds.dbName = v
}
if v := extractDefine(line, "DB_PREFIX"); v != "" {
creds.dbPrefix = v
}
}
if creds.dbHost == "" {
creds.dbHost = "localhost"
}
return creds, nil
}
// scanOpenCartSettings walks oc_setting k/v rows. The value column
// is a JSON-serialized blob; same external-script post-filter as
// the other CMS settings scanners (strict variant -- this is
// config storage, not author-written content).
func scanOpenCartSettings(account string, creds opencartCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT `key`, value FROM %ssetting WHERE %s",
creds.dbPrefix, paramsLikeClause("value"))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
key, body := splitTabRow(row)
if key == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, false)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "opencart_settings_injection",
Message: fmt.Sprintf("OpenCart settings injection on %s: %s (%s)", account, key, desc),
Details: fmt.Sprintf("Account: %s\nSetting key: %s\nMatch: %s", account, key, desc),
})
}
return findings
}
// scanOpenCartContentTable walks one of the description tables
// (product_description, information_description). Both have an id
// column and a description column; the id-column name varies but
// the schema is consistent enough that we accept it as a parameter.
//
// Looser post-filter (hasMaliciousExternalScriptInPost) because
// these tables carry author-written content.
//
// Both description tables carry one row per language per product
// or page. Without filtering, a multilingual storefront emits N
// findings per malware-injected row (one per installed language).
// language_id = 1 is English / the vanilla default; non-English
// monolingual sites and genuine multilingual coverage need a
// follow-up that reads config_language_id from oc_setting first.
func scanOpenCartContentTable(account string, creds opencartCreds, table, valueCol string) []alert.Finding {
idCol := "product_id"
if table == "information_description" {
idCol = "information_id"
}
query := fmt.Sprintf(
"SELECT %s, %s FROM %s%s WHERE language_id = 1 AND %s",
idCol, valueCol, creds.dbPrefix, table, paramsLikeClause(valueCol))
rows, _ := runCMSQuery(creds.asWPDBCreds(), query)
var findings []alert.Finding
for _, row := range rows {
id, body := splitTabRow(row)
if id == "" {
continue
}
sev, desc, ok := classifyMalwareRow(body, true)
if !ok {
continue
}
findings = append(findings, alert.Finding{
Severity: sev,
Check: "opencart_content_injection",
Message: fmt.Sprintf("OpenCart content injection on %s: %s id=%s (%s)", account, table, id, desc),
Details: fmt.Sprintf("Account: %s\nTable: %s\nRow id: %s\nMatch: %s", account, table, id, desc),
})
}
return findings
}
// scanOpenCartAdmins enumerates the oc_user table (admins/staff,
// not customers -- customers live in oc_customer). Same Warning
// per row as the other CMS adapters.
func scanOpenCartAdmins(store *state.Store, account string, creds opencartCreds) []alert.Finding {
query := fmt.Sprintf(
"SELECT user_id, username, email FROM %suser",
creds.dbPrefix)
rows, complete := runCMSQuery(creds.asWPDBCreds(), query)
return cmsAdminFindings(store, "opencart", "opencart_admin_injection", account, creds.asWPDBCreds(), rows, complete, func(fields []string) (string, string) {
return fmt.Sprintf("OpenCart admin account on %s: user_id=%s", account, fields[0]),
fmt.Sprintf("Account: %s\nRow: %s\nReview: confirm this is the legitimate site administrator.", account, strings.Join(fields, "\t"))
})
}
package checks
import (
"fmt"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// phantomAuthorFarmSize is the published-post count at which a phantom author
// is treated as a content farm rather than a likely orphan. A direct SQL user
// deletion can leave a small number of posts behind; a large group warrants a
// Critical alert.
const phantomAuthorFarmSize = 100
// maxPhantomAuthorsReported bounds the finding count across one installation,
// including all blogs in a multisite network.
const maxPhantomAuthorsReported = 25
// checkWPPhantomAuthors finds published posts whose post_author has no row in
// the network users table. postsPrefix and usersPrefix are kept separate so a
// multisite blog can join its own posts table to the shared users table.
// Callers derive both from resolveTablePrefix, appending only a numeric blog ID
// to postsPrefix, before either value reaches SQL construction.
//
// Doorway kits attribute their pages to invented author IDs and then install a
// filter that removes those IDs from admin queries, patching the post counts to
// match. The dashboard therefore shows a clean site while the posts are served
// to visitors and listed in crawler-facing sitemaps. The orphaned author ID is
// what the cloak cannot hide: it is a property of the data, not of the code
// doing the hiding, so it survives every renaming and re-obfuscation of the
// filter itself.
func checkWPPhantomAuthors(user string, creds wpDBCreds, postsPrefix, usersPrefix string, limit int) []alert.Finding {
var findings []alert.Finding
if limit <= 0 {
return findings
}
query := fmt.Sprintf(
"SELECT p.post_author, COUNT(*) AS c FROM %sposts p "+
"LEFT JOIN %susers u ON u.ID = p.post_author "+
"WHERE u.ID IS NULL AND p.post_author <> 0 "+
"AND p.post_type = 'post' AND p.post_status = 'publish' "+
"GROUP BY p.post_author ORDER BY c DESC LIMIT %d",
postsPrefix, usersPrefix, limit)
for _, line := range runMySQLQuery(creds, query) {
parts := strings.SplitN(strings.TrimSpace(line), "\t", 2)
if len(parts) != 2 {
continue
}
authorID, err := strconv.ParseUint(strings.TrimSpace(parts[0]), 10, 64)
if err != nil || authorID == 0 {
continue
}
count, err := strconv.ParseUint(strings.TrimSpace(parts[1]), 10, 64)
if err != nil || count == 0 {
continue
}
severity := alert.Warning
if count >= phantomAuthorFarmSize {
severity = alert.Critical
}
postNoun := "posts"
if count == 1 {
postNoun = "post"
}
findings = append(findings, alert.Finding{
Severity: severity,
Check: "db_phantom_post_author",
Message: fmt.Sprintf("%d published %s are attributed to a non-existent user (account: %s, author ID %d)",
count, postNoun, user, authorID),
Details: dbContentFindingDetails(creds, postsPrefix,
fmt.Sprintf("post_author = %d has no row in %susers.\n"+
"A small number can remain after an administrator deletes a user directly in SQL. "+
"Large groups can indicate hidden or injected content.",
authorID, usersPrefix)),
// Crossing into a content farm must alert even if the orphan group
// was baselined or dismissed. Counts within either tier stay stable.
DedupKey: dbContentDedupKey(user, creds, postsPrefix,
"farm="+strconv.FormatBool(count >= phantomAuthorFarmSize),
fmt.Sprintf("post_author = %d has no row in %susers.\n"+
"A small number can remain after an administrator deletes a user directly in SQL. "+
"Large groups can indicate hidden or injected content.",
authorID, usersPrefix)),
})
if len(findings) >= limit {
break
}
}
return findings
}
// capPhantomAuthorFindings bounds one installation's alert volume after every
// multisite posts table has been checked. Critical farms take precedence over
// likely orphan warnings, so a noisy primary blog cannot hide a compromised
// secondary blog by consuming the cap first.
func capPhantomAuthorFindings(findings []alert.Finding, limit int) []alert.Finding {
selected := make([]bool, len(findings))
remaining := limit
for i, finding := range findings {
if remaining == 0 {
break
}
if finding.Check == "db_phantom_post_author" && finding.Severity == alert.Critical {
selected[i] = true
remaining--
}
}
for i, finding := range findings {
if remaining == 0 {
break
}
if finding.Check == "db_phantom_post_author" && !selected[i] {
selected[i] = true
remaining--
}
}
out := make([]alert.Finding, 0, len(findings))
for i, finding := range findings {
if finding.Check != "db_phantom_post_author" || selected[i] {
out = append(out, finding)
}
}
return out
}
package checks
import (
"fmt"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
const (
// postBurstWindowDays is the recent period compared against everything the
// site published before it.
postBurstWindowDays = 365
// postBurstMinSiteAgeDays keeps a young site quiet. A launch legitimately
// produces its whole archive at once, and there is no history to judge it
// against.
postBurstMinSiteAgeDays = 400
// postBurstMinRecent avoids reporting ordinary editorial activity on a very
// quiet site, where any uptick is a large multiple of almost nothing.
postBurstMinRecent = 50
// postBurstRatio is how many times the site's entire prior output the recent
// window must exceed. A steady publisher never reaches it; a doorway kit
// clears it by an order of magnitude.
postBurstRatio = 5
)
// checkWPPostVolumeBurst reports a site that suddenly publishes far more than
// it ever did before.
//
// This is the keyword-free half of spam detection. A live compromise published
// 508 posts across seven languages in under a year on a site that had managed
// 17 in the preceding seven; a gambling word list matched fewer than half of
// them, while the change in publishing rate separated spam from real content
// exactly. It says nothing about what the posts contain, which is the point --
// the next kit will use a different vocabulary.
func checkWPPostVolumeBurst(user string, creds wpDBCreds, prefix string) []alert.Finding {
query := fmt.Sprintf(
"SELECT DATEDIFF(NOW(), MIN(post_date)) AS site_age_days, "+
"SUM(post_date < DATE_SUB(NOW(), INTERVAL %d DAY)) AS prior_posts, "+
"SUM(post_date >= DATE_SUB(NOW(), INTERVAL %d DAY)) AS recent_posts "+
"FROM %sposts WHERE post_type = 'post' AND post_status = 'publish'",
postBurstWindowDays, postBurstWindowDays, prefix)
rows := runMySQLQuery(creds, query)
if len(rows) == 0 {
return nil
}
parts := strings.SplitN(strings.TrimSpace(rows[0]), "\t", 3)
if len(parts) != 3 {
return nil
}
ageDays, err1 := strconv.Atoi(strings.TrimSpace(parts[0]))
prior, err2 := strconv.Atoi(strings.TrimSpace(parts[1]))
recent, err3 := strconv.Atoi(strings.TrimSpace(parts[2]))
if err1 != nil || err2 != nil || err3 != nil {
return nil
}
if ageDays < postBurstMinSiteAgeDays || recent < postBurstMinRecent {
return nil
}
// prior == 0 on an established site means every post it has was published
// in the recent window, which is the strongest form of this signal rather
// than a division-by-zero edge case.
if prior > 0 && recent < prior*postBurstRatio {
return nil
}
return []alert.Finding{{
Severity: alert.High,
Check: "db_post_volume_burst",
Message: fmt.Sprintf("WordPress published %d posts in the last year against %d in the %d years before (account: %s)",
recent, prior, ageDays/365, user),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Site has been publishing for %d days. A sudden flood on a long-quiet site is how "+
"doorway spam arrives, and it is visible without knowing what language or vocabulary the "+
"spam uses.\nReview the recent posts before acting: a genuine content migration looks the same.",
ageDays)),
// One burst per site. The details count days and posts, all of which
// move on their own between scans without the burst being a new one.
DedupKey: dbContentDedupKey(user, creds, prefix),
}}
}
package checks
import (
"context"
"database/sql/driver"
"errors"
"fmt"
"io"
"net"
"sort"
"strings"
"github.com/go-sql-driver/mysql"
)
const maxDatabaseQueryDiagnostics = 16
// mysqlRegexTimeout is the server error for a regular expression that
// exhausted its work limit.
const mysqlRegexTimeout = 3699
// dbQueryState separates incomplete coverage from an unusable connection.
// Both prevent a clean baseline, but a statement-local failure must not stop
// independent detectors from reading other tables in the same installation.
type dbQueryState struct {
failed bool
halted bool
regexTimeouts int
failures map[string]int
}
func (c wpDBCreds) withQueryStage(stage string) wpDBCreds {
c.queryStage = stage
return c
}
func (s *dbQueryState) record(stage string, err error) {
if s == nil {
return
}
class, code, halt := databaseQueryErrorClass(err)
s.failed = true
s.halted = s.halted || halt
if code == mysqlRegexTimeout {
s.regexTimeouts++
}
if stage == "" {
stage = "query"
}
if s.failures == nil {
s.failures = make(map[string]int)
}
// Only a code-defined stage, class and numeric error code leave the query
// boundary. Server messages and SQL can contain account data or values.
key := fmt.Sprintf("stage=%s class=%s code=%d", stage, class, code)
s.failures[key]++
}
// regexTimeoutCount lets a caller tell whether its statement stopped at the
// regex work limit. Paths without query state never retry.
func (s *dbQueryState) regexTimeoutCount() int {
if s == nil {
return 0
}
return s.regexTimeouts
}
func databaseQueryErrorClass(err error) (class string, code uint16, halt bool) {
var sqlErr *mysql.MySQLError
if errors.As(err, &sqlErr) {
code = sqlErr.Number
switch code {
case 1054, 1146:
return "schema", code, false
case 1064:
return "syntax", code, false
case 1142, 1143:
return "permission", code, false
case 1139, 1267, 1271:
return "expression", code, false
case mysqlRegexTimeout:
// ICU stops this expression when its work budget is exhausted;
// the connection and independent statements remain usable.
return "timeout", code, false
case 1044, 1045:
return "authentication", code, true
case 1049:
return "database_missing", code, true
case 1040, 1203, 1226:
return "resource", code, true
}
}
// Unknown failures keep the existing stop behavior. Continuing is safe
// only when the server identified a failure confined to this statement.
if errors.Is(err, context.DeadlineExceeded) {
return "timeout", code, true
}
if errors.Is(err, context.Canceled) {
return "canceled", code, true
}
var netErr net.Error
if errors.As(err, &netErr) {
if netErr.Timeout() {
return "timeout", code, true
}
return "connection", code, true
}
if errors.Is(err, driver.ErrBadConn) || errors.Is(err, mysql.ErrInvalidConn) ||
errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return "connection", code, true
}
return "unknown", code, true
}
func (c *dbScanCoverage) recordQueryFailures(state *dbQueryState) {
if state == nil || len(state.failures) == 0 {
return
}
if c.queryFailures == nil {
c.queryFailures = make(map[string]int)
}
keys := make([]string, 0, len(state.failures))
for key := range state.failures {
keys = append(keys, key)
}
sort.Strings(keys)
for _, key := range keys {
count := state.failures[key]
if c.queryFailures[key] == 0 && len(c.queryFailures) >= maxDatabaseQueryDiagnostics {
c.queryFailureOverflow += count
continue
}
c.queryFailures[key] += count
}
}
func (c *dbScanCoverage) queryFailureSummary() string {
keys := make([]string, 0, len(c.queryFailures))
for key := range c.queryFailures {
keys = append(keys, key)
}
sort.Strings(keys)
var b strings.Builder
for _, key := range keys {
fmt.Fprintf(&b, "Query failures: %s queries=%d\n", key, c.queryFailures[key])
}
if c.queryFailureOverflow > 0 {
fmt.Fprintf(&b, "Other query failures: queries=%d\n", c.queryFailureOverflow)
}
return b.String()
}
package checks
import (
"crypto/sha256"
"fmt"
"net"
"net/url"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/mysqlclient"
"golang.org/x/net/publicsuffix"
)
// A WordPress address pointing off the account.
//
// siteURLPoisonReason deliberately tests the shape of siteurl and not its
// host, because hosting a site under a domain that is not served locally is
// ordinary: sites move, staging lives elsewhere, a theme demo keeps its vendor
// address. Testing the host alone would report all of them.
//
// What is not ordinary is a document root the panel is serving right now,
// under domains the account owns, whose WordPress address points at an
// unrelated domain. A site that moved is no longer served here; a site that is
// served here should address itself by a name the account holds. That
// combination is either a hijack or a migration left half-finished, and both
// want an operator.
//
// It matters because the shape check only caught the live case by luck: the
// poisoned value carried a query string. A hijack to a plausible-looking
// address passes every shape test there is.
// foreignSiteURLFinding reports a served install addressed at a domain the
// account does not own, or nil when the value is fine, unreadable, or the
// question cannot be answered.
func foreignSiteURLFinding(user string, creds wpDBCreds, prefix, option, value string) *alert.Finding {
// Only a served root rules out the migration explanation.
if creds.docrootServed != servedByPanel || !creds.panelDomains.hasAccount(user) {
return nil
}
// A value whose shape is wrong is siteURLPoisonReason's finding, not this
// one; reporting it twice says nothing new.
if _, bad := siteURLPoisonReason(value); bad {
return nil
}
parsed, err := url.Parse(strings.TrimSpace(mysqlclient.BatchUnescape(value)))
if err != nil {
return nil
}
host := normalizeHost(parsed.Hostname())
if host == "" {
return nil
}
if panelHostOwnedByAccount(creds.panelDomains, user, host) {
return nil
}
return &alert.Finding{
Severity: alert.High,
Check: "db_siteurl_foreign_host",
DedupKey: foreignSiteURLDedupKey(user, creds, prefix),
Message: fmt.Sprintf("WordPress %s addresses a domain this account does not own (account: %s): %s",
option, user, host),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("%s = %s", option, truncateDB(value, 200)),
"WordPress builds every asset URL from this value, so the address it "+
"names is loaded on every page of a site the panel is serving right now.",
"This is reported only for a served document root. A site that moved "+
"away is no longer served here, which is why an unowned address on a "+
"dormant root is left alone."),
}
}
func foreignSiteURLDedupKey(user string, creds wpDBCreds, prefix string) string {
identity := strings.Join([]string{user, creds.dbHost, creds.dbName, prefix}, "\x00")
digest := sha256.Sum256([]byte(identity))
return fmt.Sprintf("wp-site:%x", digest[:12])
}
// panelDomainOwnership indexes the complete panel map in both directions.
type panelDomainOwnership struct {
owners map[string]string
wildcards map[string]string
accounts map[string]struct{}
}
func newPanelDomainOwnership(panelDomains map[string][]string) *panelDomainOwnership {
if len(panelDomains) == 0 {
return nil
}
ownership := &panelDomainOwnership{
owners: make(map[string]string),
wildcards: make(map[string]string),
accounts: make(map[string]struct{}, len(panelDomains)),
}
for account, domains := range panelDomains {
for _, rawDomain := range domains {
wildcard := strings.HasPrefix(rawDomain, "*.")
domain := normalizeHost(strings.TrimPrefix(rawDomain, "*."))
if domain == "" {
return nil
}
owners := ownership.owners
if wildcard {
owners = ownership.wildcards
}
if owner, exists := owners[domain]; exists && owner != account {
return nil
}
owners[domain] = account
ownership.accounts[account] = struct{}{}
}
}
if len(ownership.owners) == 0 && len(ownership.wildcards) == 0 {
return nil
}
return ownership
}
func (ownership *panelDomainOwnership) hasAccount(account string) bool {
if ownership == nil {
return false
}
_, ok := ownership.accounts[account]
return ok
}
// panelHostOwnedByAccount resolves ownership using the most-specific panel
// domain that covers host. This preserves delegated subdomains: when alice
// owns example.com but bob owns shop.example.com, bob owns both
// shop.example.com and www.shop.example.com.
func panelHostOwnedByAccount(ownership *panelDomainOwnership, account, host string) bool {
host = normalizeHost(host)
if ownership == nil || host == "" {
return false
}
if net.ParseIP(host) != nil {
owner, exists := ownership.owners[host]
return exists && owner == account
}
for domain, exact := host, true; ; exact = false {
// A wildcard delegation is more specific than an exact mapping of its
// parent, but it never owns the parent name itself. Like an exact
// ancestor, its base must be an ownable domain: a malformed wildcard
// such as *.com must not claim every registrable name below it.
if !exact {
if owner, exists := ownership.wildcards[domain]; exists {
if panelDomainOwnsDescendants(domain) {
return owner == account
}
}
}
if owner, exists := ownership.owners[domain]; exists {
// Exact names remain authoritative. An ancestor must itself be an
// ownable domain; a malformed "com" row cannot claim google.com.
if exact {
return owner == account
}
if panelDomainOwnsDescendants(domain) {
return owner == account
}
}
dot := strings.IndexByte(domain, '.')
if dot < 0 {
break
}
domain = domain[dot+1:]
}
return false
}
func panelDomainOwnsDescendants(domain string) bool {
if net.ParseIP(domain) != nil {
return false
}
_, err := publicsuffix.EffectiveTLDPlusOne(domain)
return err == nil
}
package checks
import (
"fmt"
"regexp"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/mysqlclient"
)
const (
maxTaxonomyRowsScanned = 500
maxTaxonomySamplesPerKind = 10
)
const (
termNameURLPattern = `^[[:space:]]*(https?://|www[.])[^[:space:]]+`
termNameSpamPattern = `(^|[^a-z])(casino|casinos|kasino|gambling|betting|bookmaker|mostbet|viagra|cialis|pharmacy|pokies|slots|bahis)([^a-z]|$)`
)
// termNameIsURL matches a term whose name is a URL. Categories and tags are
// human labels; a URL in that field is put there by a link-farm kit, never by
// an editor.
var termNameIsURL = regexp.MustCompile(`(?i)` + termNameURLPattern)
// termNameSpamVocabulary covers gambling and pharmacy vocabulary in term
// names. It is bounded on word edges so "specialist" does not match "cialis".
var termNameSpamVocabulary = regexp.MustCompile(`(?i)` + termNameSpamPattern)
type spamTaxonomyRow struct {
taxonomy string
name string
}
func parseSpamTaxonomyRow(line string) (spamTaxonomyRow, bool) {
parts := strings.SplitN(strings.TrimRight(line, "\r\n"), "\t", 4)
if len(parts) != 4 {
return spamTaxonomyRow{}, false
}
name := strings.TrimSpace(mysqlclient.BatchUnescape(parts[3]))
if name == "" {
return spamTaxonomyRow{}, false
}
return spamTaxonomyRow{
taxonomy: strings.TrimSpace(mysqlclient.BatchUnescape(parts[1])),
name: name,
}, true
}
// checkWPSpamTaxonomy reports attacker-created taxonomy terms.
//
// Removing spam posts leaves their taxonomy behind, and a category archive is a
// public page: a spam taxonomy is a doorway network even with no posts attached.
// This was missed during a live cleanup -- the homepage still served gambling
// links after every spam post had been deleted.
func checkWPSpamTaxonomy(user string, creds wpDBCreds, prefix string) []alert.Finding {
urlCandidate := fmt.Sprintf("LOWER(t.name) REGEXP '%s'", termNameURLPattern)
spamCandidate := fmt.Sprintf("LOWER(t.name) REGEXP '%s'", termNameSpamPattern)
query := fmt.Sprintf(
"SELECT t.term_id, tt.taxonomy, tt.count, t.name FROM %sterms t "+
"JOIN %sterm_taxonomy tt ON tt.term_id = t.term_id "+
"WHERE %s OR %s ORDER BY CASE WHEN %s THEN 0 ELSE 1 END, t.term_id LIMIT %d",
prefix, prefix, urlCandidate, spamCandidate, urlCandidate, maxTaxonomyRowsScanned+1)
var urlNamed, keywordNamed []string
var urlCount, keywordCount int
rows := runMySQLQuery(creds, query)
truncated := len(rows) > maxTaxonomyRowsScanned
if truncated {
markCheckIncomplete(creds.queryCtx, "db_content")
rows = rows[:maxTaxonomyRowsScanned]
}
for _, line := range rows {
row, ok := parseSpamTaxonomyRow(line)
if !ok {
markCheckIncomplete(creds.queryCtx, "db_content")
continue
}
label := fmt.Sprintf("%q (%q)", row.name, row.taxonomy)
switch {
case termNameIsURL.MatchString(row.name):
urlCount++
if len(urlNamed) < maxTaxonomySamplesPerKind {
urlNamed = append(urlNamed, label)
}
case termNameSpamVocabulary.MatchString(row.name):
keywordCount++
if len(keywordNamed) < maxTaxonomySamplesPerKind {
keywordNamed = append(keywordNamed, label)
}
}
}
total := urlCount + keywordCount
if total == 0 {
return nil
}
details := []string{
"Taxonomy terms can feed public archives, navigation and other rendered " +
"content. Deleting spam posts does not remove them.",
}
if urlCount > 0 {
details = append(details, taxonomySampleDetails("Named after a URL", urlNamed, urlCount))
}
if keywordCount > 0 {
details = append(details, taxonomySampleDetails("Spam vocabulary", keywordNamed, keywordCount))
}
// A URL as a term name has no benign reading; vocabulary alone can be a
// legitimate site writing about an industry.
severity := alert.High
if urlCount == 0 {
severity = alert.Warning
}
return []alert.Finding{{
Severity: severity,
Check: "db_spam_taxonomy",
Message: fmt.Sprintf("%s spam taxonomy terms found in WordPress (account: %s)",
spamCountLabel(total, truncated), user),
Details: dbContentFindingDetails(creds, prefix, details...),
DedupKey: dbContentDedupKey(user, creds, prefix, details...),
}}
}
func taxonomySampleDetails(label string, samples []string, total int) string {
if total > len(samples) {
label += fmt.Sprintf(" (showing %d of %d)", len(samples), total)
}
return label + ": " + strings.Join(samples, ", ")
}
package checks
import (
"fmt"
"regexp"
"sort"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// Cloak logic stored in the database.
//
// A cloak has to decide what to serve on every request, which means the page
// must not be cached. Stored code that disables caching and, in the same
// snippet, looks for a search-engine crawler is serving one page to the index
// and another to visitors.
//
// Neither half is evidence on its own, which is why both are required: caching
// plugins invoke cache helpers from their own files as a matter of course, and
// reading the user agent is ordinary. It is a stored snippet doing both that
// has no innocent reading -- the snippet is not the caching plugin, and it has
// no reason to care whether the visitor is Googlebot.
// crawlerUserAgent matches the crawlers a doorway kit cares about. The list is
// the set worth cloaking for: the engines that index and rank, plus the SEO
// crawlers kits hide from to stay out of backlink reports.
var crawlerUserAgent = regexp.MustCompile(
`(?i)\b(googlebot|bingbot|msnbot|yandex(?:bot)?|baiduspider|duckduckbot|slurp|` +
`applebot|sogou|exabot|facebot|ia_archiver|ahrefsbot|semrushbot|mj12bot|dotbot|petalbot|` +
`oai-searchbot|claude-searchbot)\b`)
const (
// The detector only reports known signals, but explicit caps keep finding
// text and derived-literal work bounded if those vocabularies grow.
maxStoredCloakMatches = 64
maxStoredCloakDetailNames = 8
maxStoredCloakDerivedBytes = maxStoredCodeBytes
maxStoredScalarParens = 32
)
// storedCloakComponents returns the cache-defeat and crawler-detection markers
// found in one stored snippet.
func storedCloakComponents(code []byte) (cacheDefeat, crawler []string) {
php := stripPHPCommentsFromCode(string(code))
// Heredoc and nowdoc bodies are string data. Keep them available to the
// crawler-literal scan, but do not parse helper calls or $_SERVER accesses
// written inside them as executable PHP.
executablePHP := blankStoredPHPHeredocs(php)
cacheDefeat = storedCacheDefeatSignals(executablePHP)
if len(cacheDefeat) == 0 {
return nil, nil
}
// A crawler name in documentation or output is not a visitor test. Require
// the snippet to inspect the HTTP user agent, including a name assembled
// from constant string operations.
if !storedHasUserAgentInspection(executablePHP) {
return cacheDefeat, nil
}
derived := storedCloakDerivedStrings(executablePHP)
searchable := php + "\n" + derived
crawler = appendCapturedNames(nil, make(map[string]bool), crawlerUserAgent,
[]byte(searchable), strings.ToLower)
return limitStoredCloakNames(cacheDefeat), limitStoredCloakNames(crawler)
}
// appendCapturedNames collects the first capture group of each match, so the
// finding names the constant or crawler rather than the surrounding syntax.
func appendCapturedNames(out []string, seen map[string]bool, re *regexp.Regexp, code []byte, canonical func(string) string) []string {
for _, m := range re.FindAllSubmatch(code, maxStoredCloakMatches) {
if len(m) < 2 {
continue
}
name := canonical(strings.TrimSpace(string(m[1])))
if name == "" || seen[name] {
continue
}
seen[name] = true
out = append(out, name)
}
sort.Strings(out)
return out
}
func limitStoredCloakNames(names []string) []string {
if len(names) > maxStoredCloakDetailNames {
return names[:maxStoredCloakDetailNames]
}
return names
}
func blankStoredPHPHeredocs(code string) string {
var out strings.Builder
out.Grow(len(code))
for i := 0; i < len(code); i++ {
// A quoted string can contain text that looks like an unterminated
// heredoc opener. Skip it whole so that text cannot blank executable
// code after the closing quote.
if isPHPQuote(code[i]) {
end := skipPHPString(code, i)
out.WriteString(code[i : end+1])
i = end
continue
}
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
end := phpHeredocEnd(code, bodyStart, label)
blankInlineHTML(&out, code[i:end])
i = end - 1
continue
}
out.WriteByte(code[i])
}
return out.String()
}
type storedPHPBool uint8
const (
storedPHPBoolUnknown storedPHPBool = iota
storedPHPBoolFalse
storedPHPBoolTrue
)
type storedPHPBoolAssignment struct {
pos int
value storedPHPBool
}
// storedCacheDefeatSignals recognises operations that stop the current request
// from being cached. Parsing real calls and assignments avoids treating a
// comment, string example, false DONOTCACHE value, or true WP_CACHE value as
// executable cache control.
func storedCacheDefeatSignals(code string) []string {
seen := make(map[string]bool)
var signals []string
var assignments map[string][]storedPHPBoolAssignment
boolValue := func(expr string, before int) storedPHPBool {
value := storedPHPBoolExpression(expr, nil, before)
if value != storedPHPBoolUnknown {
return value
}
if _, ok := singlePHPVariableExpr(trimStoredPHPParens(strings.TrimSpace(expr))); !ok {
return storedPHPBoolUnknown
}
if assignments == nil {
assignments = storedPHPBoolAssignments(code)
}
return storedPHPBoolExpression(expr, assignments, before)
}
callNames := map[string]struct{}{
"define": {},
"header": {},
"nocache_headers": {},
}
for searchFrom := 0; searchFrom < len(code); {
callStart, openParen, closeParen, ok := nextStandalonePHPCall(code, searchFrom, callNames)
if !ok {
break
}
searchFrom = nextSearchOffset(closeParen, len(code))
if closeParen >= len(code) || storedPHPCallIsNonGlobal(code, callStart) {
continue
}
name := storedPHPCallName(code[callStart:openParen])
args := phpCallArguments(code, openParen+1, closeParen)
switch name {
case "define":
if len(args) < 2 {
continue
}
constantName, ok := storedConstantStringExpression(args[0])
if !ok {
continue
}
constantName = strings.ToUpper(strings.TrimSpace(constantName))
value := boolValue(args[1], callStart)
switch constantName {
case "DONOTCACHEPAGE":
if value == storedPHPBoolTrue {
signals = appendStoredCloakName(signals, seen, constantName)
}
case "WP_CACHE":
if value == storedPHPBoolFalse {
signals = appendStoredCloakName(signals, seen, constantName)
}
}
case "header":
if len(args) == 0 {
continue
}
headerValue, ok := storedConstantStringExpression(args[0])
if ok && storedLiteSpeedNoCacheHeader(headerValue) {
signals = appendStoredCloakName(signals, seen, "X-LiteSpeed-Cache-Control")
}
case "nocache_headers":
if len(args) == 0 {
signals = appendStoredCloakName(signals, seen, "nocache_headers")
}
}
}
sort.Strings(signals)
return limitStoredCloakNames(signals)
}
func appendStoredCloakName(out []string, seen map[string]bool, name string) []string {
if name == "" || seen[name] {
return out
}
seen[name] = true
return append(out, name)
}
func storedPHPCallName(callPrefix string) string {
return strings.ToLower(strings.TrimPrefix(strings.TrimSpace(callPrefix), `\`))
}
func storedPHPCallIsNonGlobal(code string, callStart int) bool {
i := callStart - 1
for i >= 0 && isPHPSpace(code[i]) {
i--
}
// PHP permits whitespace around member, static, and namespace operators.
// Those calls do not invoke the WordPress/PHP global cache helpers.
if i >= 0 && (code[i] == '>' || code[i] == ':' || code[i] == '\\') {
return true
}
if i >= 0 && code[i] == '&' {
i--
for i >= 0 && isPHPSpace(code[i]) {
i--
}
}
end := i + 1
for i >= 0 && isPHPIdentifierPart(code[i]) {
i--
}
return strings.EqualFold(code[i+1:end], "function")
}
func storedLiteSpeedNoCacheHeader(value string) bool {
name, directive, ok := strings.Cut(value, ":")
if !ok || !strings.EqualFold(strings.TrimSpace(name), "X-LiteSpeed-Cache-Control") {
return false
}
for _, item := range strings.Split(directive, ",") {
if strings.EqualFold(strings.TrimSpace(item), "no-cache") {
return true
}
}
return false
}
func storedPHPBoolAssignments(code string) map[string][]storedPHPBoolAssignment {
assignments := make(map[string][]storedPHPBoolAssignment)
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
if code[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(code, i)
if !ok {
continue
}
operator := skipPHPWhitespace(code, next)
opLen, direct, _, ok := phpAssignmentOperator(code, operator)
if !ok {
i = next - 1
continue
}
exprStart := skipPHPWhitespace(code, operator+opLen)
exprEnd := phpExpressionEnd(code, exprStart)
value := storedPHPBoolUnknown
if direct {
value = storedPHPBoolExpression(code[exprStart:exprEnd], assignments, exprStart)
}
assignments[variable] = append(assignments[variable], storedPHPBoolAssignment{
pos: exprEnd, value: value,
})
i = exprEnd - 1
}
return assignments
}
func storedPHPBoolExpression(expr string, assignments map[string][]storedPHPBoolAssignment, before int) storedPHPBool {
expr = trimStoredPHPParens(strings.TrimSpace(expr))
switch strings.ToLower(expr) {
case "true":
return storedPHPBoolTrue
case "false", "null":
return storedPHPBoolFalse
}
if value, ok := storedConstantStringExpression(expr); ok {
if value == "" || value == "0" {
return storedPHPBoolFalse
}
return storedPHPBoolTrue
}
if number, err := strconv.ParseFloat(expr, 64); err == nil {
if number == 0 {
return storedPHPBoolFalse
}
return storedPHPBoolTrue
}
variable, ok := singlePHPVariableExpr(expr)
if !ok {
return storedPHPBoolUnknown
}
values := assignments[variable]
index := sort.Search(len(values), func(i int) bool {
return values[i].pos > before
})
if index > 0 {
return values[index-1].value
}
return storedPHPBoolUnknown
}
func trimStoredPHPParens(expr string) string {
// Scalar expressions are attacker-controlled. Cap redundant-wrapper work so
// deeply nested input cannot turn repeated matching scans quadratic.
for depth := 0; depth < maxStoredScalarParens && len(expr) >= 2 && expr[0] == '(' && matchingParen(expr, 0) == len(expr)-1; depth++ {
expr = strings.TrimSpace(expr[1 : len(expr)-1])
}
return expr
}
// storedCloakDerivedStrings evaluates only bounded, constant string operations
// commonly used to hide crawler literals. General PHP evaluation is outside
// this detector; literal XOR remains covered by the backdoor signature.
func storedCloakDerivedStrings(code string) string {
var derived strings.Builder
for i := 0; i < len(code) && derived.Len() < maxStoredCloakDerivedBytes; i++ {
if !isPHPQuote(code[i]) {
continue
}
closeQuote, closed := storedPHPStringClose(code, i)
if !closed {
break
}
dot := skipPHPWhitespace(code, closeQuote+1)
if dot >= len(code) || code[dot] != '.' {
i = closeQuote
continue
}
nextOperand := skipPHPWhitespace(code, dot+1)
if nextOperand >= len(code) || !isPHPQuote(code[nextOperand]) {
i = closeQuote
continue
}
value, end, parts, ok := storedConstantStringAt(code, i)
if !ok {
continue
}
if parts > 1 {
appendStoredDerivedString(&derived, value)
}
i = end - 1
}
rot13 := map[string]struct{}{"str_rot13": {}}
for searchFrom := 0; searchFrom < len(code) && derived.Len() < maxStoredCloakDerivedBytes; {
callStart, openParen, closeParen, ok := nextStandalonePHPCall(code, searchFrom, rot13)
if !ok {
break
}
searchFrom = nextSearchOffset(closeParen, len(code))
if closeParen >= len(code) || storedPHPCallIsNonGlobal(code, callStart) {
continue
}
args := phpCallArguments(code, openParen+1, closeParen)
if len(args) != 1 {
continue
}
value, ok := storedConstantStringExpression(args[0])
if ok {
appendStoredDerivedString(&derived, storedROT13(value))
}
}
return derived.String()
}
func appendStoredDerivedString(out *strings.Builder, value string) {
remaining := maxStoredCloakDerivedBytes - out.Len()
if remaining <= 1 {
return
}
out.WriteByte('\n')
remaining--
if len(value) > remaining {
value = value[:remaining]
}
out.WriteString(value)
}
func storedConstantStringExpression(expr string) (string, bool) {
expr = strings.TrimSpace(expr)
value, end, _, ok := storedConstantStringAt(expr, 0)
return value, ok && skipPHPWhitespace(expr, end) == len(expr)
}
func storedConstantStringAt(code string, start int) (string, int, int, bool) {
if start >= len(code) || !isPHPQuote(code[start]) {
return "", start, 0, false
}
var value strings.Builder
parts := 0
pos := start
for {
if pos >= len(code) || !isPHPQuote(code[pos]) {
return "", start, 0, false
}
closeQuote, closed := storedPHPStringClose(code, pos)
if !closed {
return "", start, 0, false
}
if storedPHPStringInterpolates(code, pos, closeQuote) {
return "", start, 0, false
}
value.WriteString(phpStringLiteralValue(code, pos, closeQuote))
parts++
next := skipPHPWhitespace(code, closeQuote+1)
if next >= len(code) || code[next] != '.' {
return value.String(), closeQuote + 1, parts, true
}
pos = skipPHPWhitespace(code, next+1)
if pos >= len(code) || !isPHPQuote(code[pos]) {
return value.String(), closeQuote + 1, parts, true
}
}
}
func storedPHPStringInterpolates(code string, start, end int) bool {
if code[start] != '"' {
return false
}
for i := start + 1; i < end; i++ {
if code[i] == '\\' && i+1 < end {
i++
continue
}
if code[i] == '$' && i+1 < end && (code[i+1] == '{' || isPHPIdentifierStart(code[i+1])) {
return true
}
}
return false
}
func storedPHPStringClose(code string, start int) (int, bool) {
quote := code[start]
for i := start + 1; i < len(code); i++ {
if code[i] == '\\' && i+1 < len(code) {
i++
continue
}
if code[i] == quote {
return i, true
}
}
return len(code), false
}
// storedHasUserAgentInspection requires an executable $_SERVER access instead
// of accepting HTTP_USER_AGENT in prose or another string literal. Constant
// string concatenation is enough to cover the common lightly-obfuscated key
// without following arbitrary attacker-controlled variables. Parsing just the
// supported scalar key avoids repeatedly searching nested, hostile brackets.
func storedHasUserAgentInspection(code string) bool {
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
if code[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(code, i)
if !ok || !strings.EqualFold(variable, "_SERVER") {
continue
}
openBracket := skipPHPWhitespace(code, next)
if openBracket >= len(code) || code[openBracket] != '[' {
continue
}
keyStart := skipPHPWhitespace(code, openBracket+1)
key, keyEnd, _, ok := storedConstantStringAt(code, keyStart)
if !ok {
continue
}
closeBracket := skipPHPWhitespace(code, keyEnd)
if closeBracket < len(code) && code[closeBracket] == ']' &&
strings.EqualFold(key, "HTTP_USER_AGENT") {
return true
}
}
return false
}
func storedROT13(value string) string {
return strings.Map(func(r rune) rune {
switch {
case r >= 'a' && r <= 'z':
return 'a' + (r-'a'+13)%26
case r >= 'A' && r <= 'Z':
return 'A' + (r-'A'+13)%26
default:
return r
}
}, value)
}
// storedCloakFinding reports a stored snippet that both defeats caching and
// looks for a crawler, or nil when only one half is present.
func storedCloakFinding(user string, creds wpDBCreds, prefix string, row storedCodeRow) *alert.Finding {
cacheDefeat, crawler := storedCloakComponents(row.code)
return storedCloakFindingWithComponents(user, creds, prefix, row, cacheDefeat, crawler)
}
func storedCloakFindingWithComponents(user string, creds wpDBCreds, prefix string, row storedCodeRow, cacheDefeat, crawler []string) *alert.Finding {
if len(cacheDefeat) == 0 || len(crawler) == 0 {
return nil
}
// Only a published snippet runs. A draft still documents the intent.
severity := alert.Warning
if row.status == "publish" {
severity = alert.High
}
return &alert.Finding{
Severity: severity,
Check: "db_stored_cloak_logic",
Message: fmt.Sprintf("Stored PHP snippet %s (%s) serves crawlers differently from visitors (account: %s)",
row.id, row.status, user),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Snippet %s is stored in %sposts, so no filesystem scan reads it.", row.id, prefix),
"It disables caching for the request and, in the same snippet, tests the "+
"visitor against a search or SEO crawler. Cloaks need both: the decision "+
"is per request, so the page must not be served from cache. Cache helpers "+
"and user-agent checks are ordinary separately, but not together here.",
"Cache defeat: "+strings.Join(cacheDefeat, ", "),
"Crawlers named: "+strings.Join(crawler, ", ")),
DedupKey: dbContentDedupKey(user, creds, prefix,
"status="+row.status,
fmt.Sprintf("Snippet %s is stored in %sposts, so no filesystem scan reads it.", row.id, prefix),
"It disables caching for the request and, in the same snippet, tests the "+
"visitor against a search or SEO crawler. Cloaks need both: the decision "+
"is per request, so the page must not be served from cache. Cache helpers "+
"and user-agent checks are ordinary separately, but not together here.",
"Cache defeat: "+strings.Join(cacheDefeat, ", "),
"Crawlers named: "+strings.Join(crawler, ", ")),
}
}
// storedCloakNote adds the cloak components to a snippet that already matched a
// signature, rather than raising a second finding about the same row.
func storedCloakNote(cacheDefeat, crawler []string) string {
if len(cacheDefeat) == 0 || len(crawler) == 0 {
return ""
}
return fmt.Sprintf("\nIt also cloaks: caching is disabled for the request (%s) "+
"while the visitor is tested against %s, so what a crawler is served is "+
"not what a visitor sees.",
strings.Join(cacheDefeat, ", "), strings.Join(crawler, ", "))
}
package checks
import (
"encoding/hex"
"fmt"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
const (
// A snippet is code a human wrote into a form; the ones that matter are far
// smaller than this. The bound keeps one pathological row from dominating a
// scan, and the hex encoding below doubles whatever is transferred.
maxStoredCodeBytes = 65536
maxStoredCodeRows = 200
)
type storedCodeRow struct {
id string
status string
contentSize int64
code []byte
}
// parseStoredCodeRow decodes one tab-separated query row. ok remains true for
// a decoded prefix when the query or transport truncates it; complete tells the
// caller not to present that partial row as a complete scan.
func parseStoredCodeRow(line string) (row storedCodeRow, ok, complete bool) {
parts := strings.SplitN(strings.TrimSpace(line), "\t", 4)
if len(parts) != 4 {
return row, false, false
}
contentSize, err := strconv.ParseInt(strings.TrimSpace(parts[2]), 10, 64)
if err != nil || contentSize < 0 {
return row, false, false
}
encodedCode := strings.TrimSpace(parts[3])
code, decodeErr := hex.DecodeString(encodedCode)
if len(code) == 0 {
return row, false, false
}
return storedCodeRow{
id: strings.TrimSpace(parts[0]),
status: strings.TrimSpace(parts[1]),
contentSize: contentSize,
code: code,
}, true, decodeErr == nil && int64(len(code)) == contentSize
}
// checkWPStoredCode scans WPCode PHP that lives in the database rather than in
// a file.
//
// WPCode executes stored code by design, which makes the posts table an
// executable surface that no filesystem scan covers. On a live
// compromise a 17KB obfuscated backdoor ran on every request from a WPCode row
// while a full file sweep of the same account -- core checksums, eval chains,
// upload shells -- came back clean.
//
// Content is hex-encoded by the query because snippets contain newlines and
// tabs, which would otherwise break row parsing.
func checkWPStoredCode(user string, creds wpDBCreds, prefix string) []alert.Finding {
scanner := contentSignatureScanner()
if scanner == nil {
return nil
}
// WPCode / Insert Headers and Footers. Code Snippets keeps its code in its
// own table and is the next surface worth adding here.
query := fmt.Sprintf(
"SELECT p.ID, p.post_status, OCTET_LENGTH(p.post_content), "+
"HEX(LEFT(CAST(p.post_content AS BINARY), %d)) FROM %sposts p "+
"WHERE p.post_type = 'wpcode' AND p.post_content <> '' "+
"AND EXISTS (SELECT 1 FROM %sterm_relationships tr "+
"JOIN %sterm_taxonomy tt ON tt.term_taxonomy_id = tr.term_taxonomy_id "+
"JOIN %sterms t ON t.term_id = tt.term_id "+
"WHERE tr.object_id = p.ID AND tt.taxonomy = 'wpcode_type' "+
"AND t.slug IN ('php', 'universal')) "+
"ORDER BY CASE p.post_status WHEN 'publish' THEN 0 WHEN 'draft' THEN 1 "+
"WHEN 'trash' THEN 2 ELSE 3 END, p.ID LIMIT %d",
maxStoredCodeBytes, prefix, prefix, prefix, prefix, maxStoredCodeRows+1)
var findings []alert.Finding
rows := runMySQLQuery(creds, query)
if len(rows) > maxStoredCodeRows {
markCheckIncomplete(creds.queryCtx, "db_content")
rows = rows[:maxStoredCodeRows]
}
for _, line := range rows {
row, ok, complete := parseStoredCodeRow(line)
if !complete {
markCheckIncomplete(creds.queryCtx, "db_content")
}
if !ok {
continue
}
// MySQL HEX() always emits valid hex, so a decode error means the row
// was truncated in transport. Scan whatever decoded rather than
// dropping the row: a partially recovered payload still identifies a
// backdoor, and silence here would read as "clean".
cacheDefeat, crawler := storedCloakComponents(row.code)
hits := scanner.ScanContentWithSize(row.code, ".php", row.contentSize)
if len(hits) == 0 {
// No signature matched, but stored code that both defeats caching
// and looks for a crawler is a cloak on its own terms.
if cloak := storedCloakFindingWithComponents(user, creds, prefix, row, cacheDefeat, crawler); cloak != nil {
findings = append(findings, *cloak)
}
continue
}
// Only a published snippet runs. A draft is one click from running; a
// trashed one is evidence of what was run before.
severity := alert.Warning
switch row.status {
case "publish":
severity = alert.Critical
case "draft":
severity = alert.High
}
names := make([]string, 0, len(hits))
for _, h := range hits {
names = append(names, h.RuleName)
}
findings = append(findings, alert.Finding{
Severity: severity,
Check: "db_stored_code_execution",
Message: fmt.Sprintf("Stored PHP snippet %s (%s) matches %s (account: %s)",
row.id, row.status, strings.Join(names, ", "), user),
Details: dbContentFindingDetails(creds, prefix,
fmt.Sprintf("Snippet %s is stored in %sposts for WPCode, "+
"so it is not visible to any filesystem scan.\nMatched: %s%s",
row.id, prefix, strings.Join(names, ", "),
storedCloakNote(cacheDefeat, crawler))),
DedupKey: dbContentDedupKey(user, creds, prefix,
"status="+row.status,
fmt.Sprintf("Snippet %s is stored in %sposts for WPCode, "+
"so it is not visible to any filesystem scan.\nMatched: %s%s",
row.id, prefix, strings.Join(names, ", "),
storedCloakNote(cacheDefeat, crawler))),
})
}
return findings
}
package checks
import (
"fmt"
"net"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/processctx"
)
// DirectSMTPEgressInput is the input to the evaluator. The caller
// (BPF connection consumer or legacy poller) builds it from the live
// event and passes the platform-resolved MTA allowlist as MTA.
//
// Process is optional; when present the resulting finding includes the
// full process-ancestry tree. UID/User/PID/Comm/Exe are the live event
// fields used in finding details and account attribution.
type DirectSMTPEgressInput struct {
UID uint32
User string
PID uint32
Comm string
Exe string
DstIP net.IP
DstPort uint16
MTA platform.MTAIdents
Process *processctx.ProcessContext
// Domain is an optional rDNS-resolved name for DstIP. When set, it
// is included in the finding details. Populating it is the caller's
// responsibility (off-path enrichment lands in Task 6).
Domain string
}
// EvaluateDirectSMTPEgress returns a populated finding when the input
// represents a non-MTA local process opening an outbound SMTP connection.
// Owner resolution uses the shared passwd cache. Detector-disabled config
// returns (zero, false) without inspecting the input.
func EvaluateDirectSMTPEgress(cfg *config.Config, in DirectSMTPEgressInput) (alert.Finding, bool) {
if cfg == nil || !cfg.Detection.DirectSMTPEgress.Enabled || directSMTPEgressBackend(cfg) == "none" {
return alert.Finding{}, false
}
if in.UID == 0 {
return alert.Finding{}, false
}
if in.DstIP == nil || in.DstIP.IsLoopback() || in.DstIP.IsUnspecified() {
return alert.Finding{}, false
}
if !portInList(in.DstPort, cfg.Detection.DirectSMTPEgress.Ports) {
return alert.Finding{}, false
}
if isInfraIP(in.DstIP.String(), cfg.InfraIPs) {
return alert.Finding{}, false
}
if in.MTA.IsMTAUser(in.User) {
return alert.Finding{}, false
}
dst := in.DstIP.String()
if in.DstIP.To4() == nil {
dst = "[" + dst + "]"
}
details := fmt.Sprintf("UID: %d (%s), Process: %s, PID: %d, Destination: %s:%d",
in.UID, in.User, in.Comm, in.PID, dst, in.DstPort)
if in.Domain != "" {
details += ", Domain: " + in.Domain
}
return alert.Finding{
Severity: alert.High,
Check: "direct_smtp_egress",
Message: fmt.Sprintf("Non-MTA process opened outbound SMTP connection to %s:%d", dst, in.DstPort),
Details: details,
TenantID: directSMTPTenant(in),
Process: in.Process,
}, true
}
func DirectSMTPEgressBackendEnabled(cfg *config.Config, backend string) bool {
if cfg == nil || !cfg.Detection.DirectSMTPEgress.Enabled {
return false
}
choice := directSMTPEgressBackend(cfg)
switch choice {
case "auto":
return true
case "bpf", "legacy":
return backend == choice
default:
return false
}
}
func directSMTPEgressBackend(cfg *config.Config) string {
backend := strings.ToLower(strings.TrimSpace(cfg.Detection.DirectSMTPEgress.Backend))
if backend == "" {
return "auto"
}
return backend
}
func directSMTPTenant(in DirectSMTPEgressInput) string {
if in.Process != nil && in.Process.Account != "" {
return HostingAccountForUser(in.Process.Account)
}
return HostingAccountForUser(in.User)
}
func portInList(p uint16, list []int) bool {
for _, q := range list {
if q <= 0 || q > 65535 {
continue
}
if q == int(p) {
return true
}
}
return false
}
package checks
import (
"sync/atomic"
"github.com/pidginhost/csm/internal/metrics"
)
var directSMTPEgressFindingsTotal atomic.Uint64
// RegisterDirectSMTPEgressMetrics binds the per-finding counter to reg.
// Production callers should pass metrics.Default(); tests pass
// metrics.NewRegistry() to keep registration isolated.
func RegisterDirectSMTPEgressMetrics(reg *metrics.Registry) {
reg.RegisterCounterFunc(
"csm_direct_smtp_egress_findings_total",
"Direct SMTP egress findings emitted by the connection consumer.",
func() float64 { return float64(directSMTPEgressFindingsTotal.Load()) },
)
}
// BumpDirectSMTPEgressFindings increments the per-finding counter.
// Called by the connection consumer when EvaluateDirectSMTPEgress
// returns a finding.
func BumpDirectSMTPEgressFindings() {
directSMTPEgressFindingsTotal.Add(1)
}
// resetDirectSMTPEgressMetricsForTest is a test seam.
func resetDirectSMTPEgressMetricsForTest() {
directSMTPEgressFindingsTotal.Store(0)
}
package checks
import (
"bufio"
"context"
"fmt"
"net"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// DNS server processes that legitimately connect to many resolvers
// (e.g. BIND doing recursive resolution on a cPanel server).
var dnsServerUsers = map[string]bool{
"named": true, // BIND
"unbound": true, // Unbound
"pdns": true, // PowerDNS
"systemd-resolve": true, // systemd-resolved forwarding to its upstreams
"dnsmasq": true, // dnsmasq forwarding to its upstreams
}
// resolvedUpstreamsPath lists the real upstreams when /etc/resolv.conf only
// names the systemd-resolved stub.
const resolvedUpstreamsPath = "/run/systemd/resolve/resolv.conf"
const (
dnsmasqConfigPath = "/etc/dnsmasq.conf"
dnsmasqConfigGlob = "/etc/dnsmasq.d/*.conf"
)
// CheckDNSConnections looks for established connections to port 53 on
// DNS servers that are NOT in /etc/resolv.conf. This catches DNS
// tunneling, GSocket relay discovery, and malware using hardcoded resolvers.
// Connections owned by known DNS server processes (e.g. named) are skipped.
func CheckDNSConnections(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Parse configured resolvers
resolvers := parseResolvers()
if len(resolvers) == 0 {
return nil
}
// Also allow infra IPs and localhost
allowed := make(map[string]bool)
allowed["127.0.0.1"] = true
allowed["0.0.0.0"] = true
for _, r := range resolvers {
allowed[r] = true
}
// Build a set of UIDs belonging to DNS server processes (named, unbound,
// etc.) so we can skip their connections without reading /etc/passwd
// on every loop iteration.
dnsServerUIDs := resolveDNSServerUIDs()
// Parse /proc/net/tcp for connections to port 53
data, err := osFS.ReadFile("/proc/net/tcp")
if err != nil {
return nil
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) < 8 || fields[0] == "sl" {
continue
}
// State 01 = ESTABLISHED
if fields[3] != "01" {
continue
}
remoteIP, remotePort := parseHexAddr(fields[2])
if remotePort != 53 {
continue
}
if allowed[remoteIP] {
continue
}
// Check if it's an infra IP
if isInfraIP(remoteIP, cfg.InfraIPs) {
continue
}
// Skip connections owned by DNS server processes (e.g. named
// doing recursive resolution talks to many different servers)
if dnsServerUIDs[fields[7]] {
continue
}
_, localPort := parseHexAddr(fields[1])
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "dns_connection",
Message: fmt.Sprintf("DNS connection to non-configured resolver: %s", remoteIP),
Details: fmt.Sprintf("Local port: %d, Remote: %s:53\nConfigured resolvers: %s", localPort, remoteIP, strings.Join(resolvers, ", ")),
})
}
return findings
}
// resolveDNSServerUIDs returns a set of UIDs that belong to known DNS
// server users (named, unbound, pdns) by reading /etc/passwd once.
func resolveDNSServerUIDs() map[string]bool {
uids := make(map[string]bool)
data, err := osFS.ReadFile("/etc/passwd")
if err != nil {
return uids
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Split(line, ":")
if len(fields) >= 3 && dnsServerUsers[fields[0]] {
uids[fields[2]] = true
}
}
return uids
}
// parseResolvers returns the configured nameservers. When /etc/resolv.conf
// points at a loopback stub (systemd-resolved's 127.0.0.53, a local
// dnsmasq) the upstreams behind it are configured resolvers too: clients
// never reach them directly, but the stub does, and nothing else on the
// host is expected to.
func parseResolvers() []string {
resolvers := parseResolverFile("/etc/resolv.conf")
loopback := false
for _, r := range resolvers {
if ip := net.ParseIP(r); ip != nil && ip.IsLoopback() {
loopback = true
break
}
}
if loopback {
resolvers = append(resolvers, parseResolverFile(resolvedUpstreamsPath)...)
resolvers = append(resolvers, parseDNSMasqServerFile(dnsmasqConfigPath)...)
if paths, err := osFS.Glob(dnsmasqConfigGlob); err == nil {
for _, path := range paths {
resolvers = append(resolvers, parseDNSMasqServerFile(path)...)
}
}
}
seen := make(map[string]bool, len(resolvers))
unique := resolvers[:0]
for _, resolver := range resolvers {
if !seen[resolver] {
seen[resolver] = true
unique = append(unique, resolver)
}
}
return unique
}
func parseResolverFile(path string) []string {
f, err := osFS.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
var resolvers []string
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if strings.HasPrefix(line, "nameserver ") {
ip := strings.TrimSpace(strings.TrimPrefix(line, "nameserver"))
if ip != "" {
resolvers = append(resolvers, ip)
}
}
}
return resolvers
}
func parseDNSMasqServerFile(path string) []string {
f, err := osFS.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
var resolvers []string
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
key, value, ok := strings.Cut(line, "=")
if !ok || strings.TrimSpace(key) != "server" {
continue
}
if fields := strings.Fields(value); len(fields) > 0 {
if ip := dnsmasqServerIP(fields[0]); ip != "" {
resolvers = append(resolvers, ip)
}
}
}
return resolvers
}
func dnsmasqServerIP(value string) string {
value = strings.TrimSpace(value)
if strings.HasPrefix(value, "/") {
if slash := strings.LastIndexByte(value, '/'); slash >= 0 {
value = value[slash+1:]
}
}
if at := strings.IndexByte(value, '@'); at >= 0 {
value = value[:at]
}
if strings.HasPrefix(value, "[") {
if close := strings.IndexByte(value, ']'); close > 1 {
value = value[1:close]
}
} else if hash := strings.LastIndexByte(value, '#'); hash > 0 {
value = value[:hash]
}
if ip := net.ParseIP(value); ip != nil {
return ip.String()
}
return ""
}
package checks
import (
"context"
"encoding/json"
"fmt"
"net/netip"
"path/filepath"
"sort"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// CheckDNSZoneChanges monitors named zone files for tampering.
//
// A raw file-hash watch over /var/named is too coarse: every cPanel serial
// bump, AutoSSL DCV TXT record, DKIM rotation, and customer Zone Editor edit
// rewrites the file, so hashing the whole zone alerts on routine activity and
// buries real hijacks. This check instead compares only the security-relevant
// records and weighs the change against cPanel provenance:
//
// - The "security fingerprint" covers delegation (NS), mail (MX), and apex/
// wildcard address records -- the records an attacker rewrites to take over
// a domain. Serial, TXT/DKIM/SPF/DCV, and ordinary subdomain A records are
// ignored, so legitimate churn stays quiet.
// - cPanel stamps each zone it writes with an "(update_time):" header. A
// security change with no advance of that stamp means the file was edited
// out of band (direct file write, or a non-cPanel path) -- the signature of
// a hijack -- and is reported High. A security change that did go through
// cPanel is trusted more: an NS/MX move still surfaces as a Warning (could
// be a compromised account), while an apex/wildcard address repoint by the
// authenticated owner is routine and stays quiet.
//
// There is deliberately no bulk-suppression gate: a mass out-of-band NS rewrite
// across every hosted domain is exactly the incident operators must see, and
// the per-record/provenance model already keeps benign mass operations quiet.
func CheckDNSZoneChanges(_ context.Context, _ *config.Config, store *state.Store) []alert.Finding {
// cPanel stores zone files in /var/named/
zoneDir := "/var/named"
zones, err := osFS.ReadDir(zoneDir)
if err != nil {
return nil
}
var findings []alert.Finding
for _, zone := range zones {
if zone.IsDir() {
continue
}
name := zone.Name()
if !strings.HasSuffix(name, ".db") {
continue
}
fullPath := filepath.Join(zoneDir, name)
data, err := osFS.ReadFile(fullPath)
if err != nil {
continue
}
key := "_dns_zone:" + name
rawPrev, exists := store.GetRaw(key)
var prev dnsZoneState
prevOK := false
if exists {
prev, prevOK = decodeDNSZoneState(rawPrev)
}
secHash, delegHash := parseZoneSecurity(data, zoneOrigin(name))
cur := dnsZoneState{
File: hashBytes(data),
Sec: secHash,
Deleg: delegHash,
Prov: zoneUpdateTime(data),
}
cur.Panel = cur.Prov > 0
if exists && prevOK && !prev.Panel {
cur.Panel = false
cur.Prov = 0
}
if !cur.Panel {
cur.Prov = 0
} else if exists && prevOK && prev.Panel && cur.Prov < prev.Prov {
cur.Prov = prev.Prov
}
store.SetRaw(key, cur.encode())
if !exists {
continue // first sight: baseline only
}
if !prevOK {
continue // legacy bare-hash state: re-baseline silently, no alert
}
if prev.File == cur.File {
continue // file unchanged
}
if prev.Sec == cur.Sec {
continue // only non-security content changed (serial, TXT, subdomain A)
}
// A security-relevant record changed. Weigh it against cPanel provenance.
if !prev.Panel || !cur.Panel || cur.Prov <= prev.Prov {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "dns_zone_change",
Message: fmt.Sprintf("DNS records changed outside cPanel: %s", name),
Details: fmt.Sprintf("File: %s\nDelegation, mail, or apex address records changed with no matching cPanel zone edit. A direct edit to a zone file that bypasses cPanel is the signature of DNS hijacking.", fullPath),
})
continue
}
if prev.Deleg == cur.Deleg {
continue // cPanel-applied apex/wildcard address repoint: routine owner action
}
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "dns_zone_change",
Message: fmt.Sprintf("DNS delegation or mail records changed: %s", name),
Details: fmt.Sprintf("File: %s\nNS or MX records were changed through cPanel. Confirm this was an authorized account action and not a compromised login redirecting the domain or its mail.", fullPath),
})
}
return findings
}
// dnsZoneState is the per-zone watch state persisted under "_dns_zone:<name>".
// File is the whole-file hash (cheap "did anything change" gate); Sec is the
// hash of the security-relevant record set (NS, MX, apex/wildcard A/AAAA);
// Deleg is the hash of the delegation/mail subset (NS, MX). Panel records
// whether the zone was already known as cPanel-managed; manually managed zones
// cannot become trusted just by adding a cPanel-looking file header later. Prov
// is the cPanel "(update_time):" stamp, 0 when Panel is false.
type dnsZoneState struct {
File string `json:"f"`
Sec string `json:"s"`
Deleg string `json:"d"`
Prov int64 `json:"p"`
Panel bool `json:"c,omitempty"`
}
func (z dnsZoneState) encode() string {
b, _ := json.Marshal(z)
return string(b)
}
// decodeDNSZoneState parses persisted state. It returns ok=false for the legacy
// bare-hash format (a 64-char hex string written by earlier versions) so the
// caller re-baselines instead of treating the upgrade as a zone change.
func decodeDNSZoneState(s string) (dnsZoneState, bool) {
if !strings.HasPrefix(s, "{") {
return dnsZoneState{}, false
}
var raw struct {
File string `json:"f"`
Sec string `json:"s"`
Deleg string `json:"d"`
Prov int64 `json:"p"`
Panel *bool `json:"c"`
}
if err := json.Unmarshal([]byte(s), &raw); err != nil || raw.File == "" {
return dnsZoneState{}, false
}
z := dnsZoneState{
File: raw.File,
Sec: raw.Sec,
Deleg: raw.Deleg,
Prov: raw.Prov,
Panel: raw.Prov > 0,
}
if raw.Panel != nil {
z.Panel = *raw.Panel
}
if !z.Panel {
z.Prov = 0
}
return z, true
}
// zoneOrigin derives the zone apex (FQDN, trailing dot) from a zone file name,
// e.g. "example.com.db" -> "example.com.".
func zoneOrigin(filename string) string {
return canonicalZoneName(strings.TrimSuffix(filename, ".db"))
}
// zoneUpdateTime extracts the epoch from cPanel's "(update_time):<n>" zone
// header. Returns 0 when absent (manually managed or non-cPanel zone).
func zoneUpdateTime(data []byte) int64 {
const marker = "(update_time):"
for _, raw := range strings.Split(string(data), "\n") {
line := strings.TrimSpace(raw)
if line == "" {
continue
}
if !strings.HasPrefix(line, ";") {
return 0
}
if !strings.HasPrefix(line, "; cPanel ") ||
!strings.Contains(line, "Cpanel::ZoneFile::VERSION:") ||
!strings.Contains(line, marker) {
continue
}
rest := line[strings.Index(line, marker)+len(marker):]
j := 0
for j < len(rest) && rest[j] >= '0' && rest[j] <= '9' {
j++
}
if j == 0 {
return 0
}
n, err := strconv.ParseInt(rest[:j], 10, 64)
if err != nil {
return 0
}
return n
}
return 0
}
// parseZoneSecurity parses a BIND zone file and returns two hashes: the
// security fingerprint (NS, MX, apex/wildcard A and AAAA) and the delegation
// subset (NS, MX only). Records are canonicalized and sorted so reordering or
// whitespace changes do not register as a difference. SOA, ordinary subdomain
// addresses, and all TXT/DKIM/SPF/DCV records are excluded by design.
func parseZoneSecurity(data []byte, origin string) (secHash, delegHash string) {
var secRecs, delegRecs []string
lastOwnerFQDN := ""
zoneOrigin := canonicalZoneName(origin)
curOrigin := zoneOrigin
parenDepth := 0
pendingRaw := ""
pendingLine := ""
processRecord := func(raw, line string) {
fields := zoneRecordFields(line)
if len(fields) == 0 {
return
}
// A leading blank means "same owner as the previous record".
fqdn := lastOwnerFQDN
rest := fields
if r := []rune(raw); len(r) > 0 && r[0] != ' ' && r[0] != '\t' {
fqdn = toFQDN(fields[0], curOrigin)
lastOwnerFQDN = fqdn
rest = fields[1:]
}
typ, rdata := zoneRecordType(rest)
if typ == "" || fqdn == "" {
return
}
malformedParenSuffix := zoneMalformedParenSuffix(line)
switch typ {
case "NS":
if len(rdata) >= 1 {
rec := "NS " + fqdn + " " + toFQDN(rdata[0], curOrigin) + malformedParenSuffix
delegRecs = append(delegRecs, rec)
secRecs = append(secRecs, rec)
}
case "MX":
if len(rdata) >= 2 {
rec := "MX " + fqdn + " " + canonicalMXPreference(rdata[0]) + " " + toFQDN(rdata[1], curOrigin) + malformedParenSuffix
delegRecs = append(delegRecs, rec)
secRecs = append(secRecs, rec)
}
case "A", "AAAA":
if len(rdata) >= 1 && isApexOrWildcard(fqdn, zoneOrigin) {
rec := typ + " " + fqdn + " " + canonicalIPLiteral(typ, rdata[0]) + malformedParenSuffix
secRecs = append(secRecs, rec)
}
}
}
lines := strings.Split(string(data), "\n")
pendingCloses := false
for i, raw := range lines {
line := stripZoneComment(raw)
trimmed := strings.TrimSpace(line)
if parenDepth > 0 {
if !pendingCloses && continuationStartsZoneEntry(raw, line) {
processRecord(pendingRaw, pendingLine)
pendingRaw = ""
pendingLine = ""
parenDepth = 0
} else {
if trimmed != "" {
pendingLine += " " + trimmed
}
parenDepth += zoneParenDelta(line)
if parenDepth > 0 {
continue
}
parenDepth = 0
processRecord(pendingRaw, pendingLine)
pendingRaw = ""
pendingLine = ""
pendingCloses = false
continue
}
}
if trimmed == "" {
continue
}
if strings.HasPrefix(trimmed, "$") {
f := strings.Fields(trimmed)
if len(f) == 0 {
continue
}
if len(f) >= 2 && strings.EqualFold(f[0], "$ORIGIN") {
curOrigin = toFQDN(f[1], curOrigin)
continue
}
if rec, ok := zoneDirectiveFingerprint(f); ok {
delegRecs = append(delegRecs, rec)
secRecs = append(secRecs, rec)
}
continue
}
parenDepth = zoneParenDelta(line)
if parenDepth > 0 {
pendingRaw = raw
pendingLine = trimmed
pendingCloses = zoneContinuationCloses(parenDepth, lines[i+1:])
continue
}
if parenDepth < 0 {
parenDepth = 0
}
processRecord(raw, line)
}
if pendingLine != "" {
processRecord(pendingRaw, pendingLine)
}
return hashSortedRecords(secRecs), hashSortedRecords(delegRecs)
}
func zoneContinuationCloses(depth int, lines []string) bool {
for _, raw := range lines {
depth += zoneParenDelta(stripZoneComment(raw))
if depth <= 0 {
return true
}
}
return false
}
// zoneRecordType skips leading TTL and class tokens and returns the record type
// (upper-cased) plus its rdata tokens. Returns "" when no type is present.
func zoneRecordType(tokens []string) (string, []string) {
i := 0
for i < len(tokens) && (isZoneTTL(tokens[i]) || isZoneClass(tokens[i])) {
i++
}
if i >= len(tokens) {
return "", nil
}
return strings.ToUpper(tokens[i]), tokens[i+1:]
}
func isZoneClass(t string) bool {
switch strings.ToUpper(t) {
case "IN", "CH", "HS", "CS":
return true
}
return false
}
// isZoneTTL reports whether a token is a BIND TTL. BIND accepts plain seconds
// and compact unit sequences such as 1h30m.
func isZoneTTL(t string) bool {
if t == "" {
return false
}
i := 0
for i < len(t) {
start := i
for i < len(t) && t[i] >= '0' && t[i] <= '9' {
i++
}
if i == start {
return false
}
if i == len(t) {
return true
}
switch t[i] {
case 's', 'S', 'm', 'M', 'h', 'H', 'd', 'D', 'w', 'W':
i++
default:
return false
}
}
return true
}
func zoneRecordFields(line string) []string {
var b strings.Builder
b.Grow(len(line))
inQuote := false
escaped := false
for i := 0; i < len(line); i++ {
if escaped {
escaped = false
b.WriteByte(line[i])
continue
}
switch line[i] {
case '\\':
if inQuote {
escaped = true
}
b.WriteByte(line[i])
case '"':
inQuote = !inQuote
b.WriteByte(line[i])
case '(', ')':
if inQuote {
b.WriteByte(line[i])
} else {
b.WriteByte(' ')
}
default:
b.WriteByte(line[i])
}
}
return strings.Fields(b.String())
}
func zoneDirectiveFingerprint(fields []string) (string, bool) {
switch strings.ToUpper(fields[0]) {
case "$INCLUDE", "$GENERATE":
return "DIRECTIVE " + strings.ToUpper(fields[0]) + " " + strings.Join(fields[1:], " "), true
}
return "", false
}
func continuationStartsZoneEntry(raw, line string) bool {
if raw == "" || raw[0] == ' ' || raw[0] == '\t' {
return false
}
trimmed := strings.TrimSpace(line)
if trimmed == "" {
return false
}
if strings.HasPrefix(trimmed, "$") {
return true
}
fields := zoneRecordFields(line)
if len(fields) < 2 {
return false
}
typ, _ := zoneRecordType(fields[1:])
return isZoneRecordType(typ)
}
func isZoneRecordType(typ string) bool {
switch typ {
case "A", "AAAA", "CAA", "CNAME", "DNAME", "DNSKEY", "DS", "HTTPS",
"LOC", "MX", "NAPTR", "NS", "PTR", "SOA", "SPF", "SRV", "SSHFP",
"SVCB", "TLSA", "TXT":
return true
}
if !strings.HasPrefix(typ, "TYPE") {
return false
}
_, err := strconv.ParseUint(strings.TrimPrefix(typ, "TYPE"), 10, 16)
return err == nil
}
func canonicalMXPreference(s string) string {
if !isDecimalUint16(s) {
return s
}
n, err := strconv.ParseUint(s, 10, 16)
if err != nil {
return s
}
return strconv.FormatUint(n, 10)
}
func isDecimalUint16(s string) bool {
if s == "" {
return false
}
for i := 0; i < len(s); i++ {
if s[i] < '0' || s[i] > '9' {
return false
}
}
_, err := strconv.ParseUint(s, 10, 16)
return err == nil
}
func canonicalIPLiteral(typ, s string) string {
if strings.Contains(s, "%") {
return s
}
addr, err := netip.ParseAddr(s)
if err != nil {
return s
}
switch typ {
case "A":
if !addr.Is4() {
return s
}
case "AAAA":
if !addr.Is6() {
return s
}
default:
return s
}
return addr.String()
}
func isApexOrWildcard(fqdn, origin string) bool {
return fqdn == origin || strings.HasPrefix(fqdn, "*.")
}
func toFQDN(name, origin string) string {
name = strings.TrimSpace(name)
origin = canonicalZoneName(origin)
if name == "" || name == "@" {
return origin
}
if strings.HasSuffix(name, ".") {
return canonicalZoneName(name)
}
if origin == "." {
return canonicalZoneName(name)
}
return canonicalZoneName(name + "." + origin)
}
func canonicalZoneName(s string) string {
s = strings.TrimSpace(s)
if s == "." {
return "."
}
if strings.HasSuffix(s, ".") {
s = strings.TrimSuffix(s, ".")
if s == "" {
return "."
}
return lowerASCII(s) + "."
}
if s == "" {
return "."
}
return lowerASCII(s) + "."
}
func lowerASCII(s string) string {
var b strings.Builder
changed := false
for i := 0; i < len(s); i++ {
c := s[i]
if c >= 'A' && c <= 'Z' {
if !changed {
b.Grow(len(s))
b.WriteString(s[:i])
changed = true
}
c += 'a' - 'A'
}
if changed {
b.WriteByte(c)
}
}
if !changed {
return s
}
return b.String()
}
func stripZoneComment(line string) string {
inQuote := false
escaped := false
for i := 0; i < len(line); i++ {
if escaped {
escaped = false
continue
}
switch line[i] {
case '\\':
if inQuote {
escaped = true
}
case '"':
inQuote = !inQuote
case ';':
if !inQuote {
return line[:i]
}
}
}
return line
}
func zoneParenDelta(line string) int {
inQuote := false
escaped := false
delta := 0
for i := 0; i < len(line); i++ {
if escaped {
escaped = false
continue
}
switch line[i] {
case '\\':
if inQuote {
escaped = true
}
case '"':
inQuote = !inQuote
case '(':
if !inQuote {
delta++
}
case ')':
if !inQuote {
delta--
}
}
}
return delta
}
func zoneMalformedParenSuffix(line string) string {
if !zoneParensMalformed(line) {
return ""
}
return " MALFORMED_PARENS " + strings.TrimSpace(line)
}
func zoneParensMalformed(line string) bool {
inQuote := false
escaped := false
depth := 0
for i := 0; i < len(line); i++ {
if escaped {
escaped = false
continue
}
switch line[i] {
case '\\':
if inQuote {
escaped = true
}
case '"':
inQuote = !inQuote
case '(':
if !inQuote {
if zoneParenAttachedToToken(line, i) {
return true
}
depth++
}
case ')':
if !inQuote {
if zoneParenAttachedToToken(line, i) {
return true
}
depth--
if depth < 0 {
return true
}
}
}
}
return depth != 0
}
func zoneParenAttachedToToken(line string, i int) bool {
return (i > 0 && !isZoneSpace(line[i-1])) || (i+1 < len(line) && !isZoneSpace(line[i+1]))
}
func isZoneSpace(b byte) bool {
switch b {
case ' ', '\t', '\r', '\n':
return true
}
return false
}
func hashSortedRecords(recs []string) string {
sort.Strings(recs)
n := 0
for _, rec := range recs {
if n == 0 || rec != recs[n-1] {
recs[n] = rec
n++
}
}
return hashBytes([]byte(strings.Join(recs[:n], "\n")))
}
// CheckSSLCertIssuance monitors AutoSSL logs for new certificate issuance.
// Attackers may issue certificates for phishing domains using compromised accounts.
func CheckSSLCertIssuance(ctx context.Context, _ *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
// Check AutoSSL log
logPath := "/var/cpanel/logs/autossl"
entries, err := osFS.ReadDir(logPath)
if err != nil {
return nil
}
// Track the count of log files as a simple change indicator
currentCount := len(entries)
key := "_ssl_autossl_count"
prev, exists := store.GetRaw(key)
store.SetRaw(key, fmt.Sprintf("%d", currentCount))
if !exists {
return nil
}
prevCount := 0
fmt.Sscanf(prev, "%d", &prevCount)
if currentCount > prevCount {
// New AutoSSL activity - check the latest log
var latestLog string
var latestTime int64
for _, entry := range entries {
if entry.IsDir() {
continue
}
info, err := entry.Info()
if err != nil {
continue
}
if info.ModTime().Unix() > latestTime {
latestTime = info.ModTime().Unix()
latestLog = filepath.Join(logPath, entry.Name())
}
}
if latestLog != "" {
// Read the tail of the latest log for certificate issuance
lines := tailFile(latestLog, 50)
for _, line := range lines {
lineLower := strings.ToLower(line)
if strings.Contains(lineLower, "installed") || strings.Contains(lineLower, "issued") {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "ssl_cert_issued",
Message: "New SSL certificate issued via AutoSSL",
Details: truncate(line, 300),
})
}
}
}
}
return findings
}
package checks
import (
"context"
"errors"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var emailHashes = newEmailHashPool(3)
var emailMailboxAudits = newScanBatchMonitor()
type emailHashPool struct {
slots chan struct{}
waiting *queuehealth.Tracker
hashes *queuehealth.Tracker
}
func newEmailHashPool(capacity int) *emailHashPool {
return &emailHashPool{
slots: make(chan struct{}, capacity),
waiting: queuehealth.New(0, checkTimeout),
hashes: queuehealth.NewSharedCapacity(capacity, checkTimeout),
}
}
type emailHashWork struct {
pool *emailHashPool
ticket queuehealth.Ticket
remaining atomic.Int32
failOnce sync.Once
}
func (p *emailHashPool) acquire(ctx context.Context) (_ *emailHashWork, err error) {
waiting := p.waiting.Begin(time.Now())
defer func() {
if errors.Is(err, context.DeadlineExceeded) {
waiting.Reject(time.Now())
} else {
// Explicit cancellation withdraws demand without losing work.
waiting.Finish(time.Now())
}
}()
select {
case p.slots <- struct{}{}:
case <-ctx.Done():
return nil, ctx.Err()
}
if ctxErr := ctx.Err(); ctxErr != nil {
<-p.slots
return nil, ctxErr
}
work := &emailHashWork{pool: p, ticket: p.hashes.Begin(time.Now())}
work.remaining.Store(2)
return work, nil
}
func (w *emailHashWork) fail() {
w.failOnce.Do(func() { w.pool.hashes.Lose(time.Now(), 1) })
}
func (w *emailHashWork) release() {
// KDFs cannot be interrupted. The worker and caller share the slot so
// cancellation cannot hide a running KDF or admit unbounded late results.
if w.remaining.Add(-1) == 0 {
w.ticket.Finish(time.Now())
<-w.pool.slots
}
}
type emailHashResult struct {
match bool
err error
}
func (p *emailHashPool) execute(work *emailHashWork, match func(string) (bool, error), candidate string, done chan<- emailHashResult) {
work.ticket.Start(time.Now())
completed := false
defer func() {
if !completed {
work.fail()
}
work.release()
}()
matched, err := match(candidate)
if err != nil {
work.fail()
// Decoder/KDF errors may embed secret input.
err = errEmailPasswordVerify
}
done <- emailHashResult{match: matched, err: err}
completed = true
}
func (p *emailHashPool) matches(ctx context.Context, match func(string) (bool, error), candidate string) (_ bool, err error) {
if ctxErr := ctx.Err(); ctxErr != nil {
return false, ctxErr
}
if len(candidate) > maxEmailCandidateBytes || strings.ContainsRune(candidate, 0) {
return false, errEmailCandidate
}
work, err := p.acquire(ctx)
if err != nil {
return false, err
}
defer func() {
if errors.Is(err, context.DeadlineExceeded) {
work.fail()
}
work.release()
}()
done := make(chan emailHashResult, 1)
go p.execute(work, match, candidate, done)
select {
case got := <-done:
if ctxErr := ctx.Err(); ctxErr != nil {
return false, ctxErr
}
return got.match, got.err
case <-ctx.Done():
return false, ctx.Err()
}
}
func (p *emailHashPool) QueueStatuses(now time.Time) map[string]queuehealth.Status {
waiting := p.waiting.Snapshot(now)
// Scan invocations share the hash slots; their waiting callers have no
// fixed global cap. Do not label that backlog with the KDF concurrency.
waiting.CapacityUnavailable = true
return map[string]queuehealth.Status{
"waiting": waiting,
"hashes": p.hashes.Snapshot(now),
}
}
// EmailPasswordQueueStatuses reads admission and KDF ownership from memory.
func EmailPasswordQueueStatuses(now time.Time) map[string]queuehealth.Status {
rows := emailHashes.QueueStatuses(now)
rows["mailboxes"] = emailMailboxAudits.snapshot(now)
return rows
}
package checks
import (
"context"
"crypto/md5" // #nosec G501 -- Read-only verification of existing Dovecot MD5 hashes, never password creation.
"crypto/sha1" // #nosec G505 -- Read-only verification of existing Dovecot SHA1 hashes, never password creation.
"crypto/sha256"
"crypto/sha512"
"crypto/subtle"
"encoding/base64"
"encoding/hex"
"errors"
"hash"
"strconv"
"strings"
"github.com/go-crypt/crypt/algorithm"
"github.com/go-crypt/crypt/algorithm/argon2"
"github.com/go-crypt/crypt/algorithm/bcrypt"
"github.com/go-crypt/crypt/algorithm/md5crypt"
"github.com/go-crypt/crypt/algorithm/shacrypt"
)
const (
maxEmailHashBytes = 4096
maxEmailCandidateBytes = 256
)
var (
errEmailHashUnsupported = errors.New("unsupported password hash scheme")
errEmailHashInvalid = errors.New("malformed password hash")
errEmailHashCost = errors.New("password hash exceeds audit cost limits")
errEmailPasswordVerify = errors.New("password verification failed")
errEmailCandidate = errors.New("password candidate exceeds audit limits")
)
type emailPasswordVerifier struct {
match func(string) (bool, error)
}
func (v *emailPasswordVerifier) matches(ctx context.Context, candidate string) (bool, error) {
return emailHashes.matches(ctx, v.match, candidate)
}
func (v *emailPasswordVerifier) firstMatch(ctx context.Context, candidates []string) (string, error) {
for _, candidate := range candidates {
matched, err := v.matches(ctx, candidate)
if err != nil {
return "", err
}
if matched {
return candidate, nil
}
}
return "", ctx.Err()
}
func parseEmailPasswordHash(stored string) (*emailPasswordVerifier, error) {
if len(stored) == 0 || len(stored) > maxEmailHashBytes || strings.ContainsRune(stored, 0) {
return nil, errEmailHashInvalid
}
scheme, encoded := "CRYPT", stored
if strings.HasPrefix(stored, "{") {
end := strings.IndexByte(stored, '}')
if end < 2 {
return nil, errEmailHashInvalid
}
scheme, encoded = strings.ToUpper(stored[1:end]), stored[end+1:]
}
base, encoding, _ := strings.Cut(scheme, ".")
if encoding != "" && (base == "PLAIN" || strings.HasSuffix(base, "CRYPT") || strings.HasPrefix(base, "ARGON2")) {
decoded, err := decodeEmailDigest(encoding, encoded)
if err != nil {
return nil, err
}
encoded, encoding = string(decoded), ""
if strings.ContainsRune(encoded, 0) {
return nil, errEmailHashInvalid
}
}
switch base {
case "PLAIN":
return &emailPasswordVerifier{match: func(candidate string) (bool, error) {
return subtle.ConstantTimeCompare([]byte(encoded), []byte(candidate)) == 1, nil
}}, nil
case "PLAIN-MD5", "LDAP-MD5", "SMD5", "SHA", "SHA1", "SSHA", "SHA256", "SSHA256", "SHA512", "SSHA512":
return parseEmailDigest(base, encoding, encoded)
case "CRYPT", "SHA512-CRYPT", "SHA256-CRYPT", "MD5-CRYPT", "BLF-CRYPT", "ARGON2I", "ARGON2ID":
if encoding != "" {
return nil, errEmailHashUnsupported
}
default:
return nil, errEmailHashUnsupported
}
prefixes := map[string]string{
"SHA512-CRYPT": "$6$", "SHA256-CRYPT": "$5$", "MD5-CRYPT": "$1$",
"BLF-CRYPT": "$2", "ARGON2I": "$argon2i$", "ARGON2ID": "$argon2id$",
}
if prefix := prefixes[base]; prefix != "" && !strings.HasPrefix(encoded, prefix) {
return nil, errEmailHashInvalid
}
decode, err := boundedEmailCryptDecoder(encoded)
if err != nil {
return nil, err
}
digest, err := decode(encoded)
if err != nil {
return nil, errEmailHashInvalid
}
return &emailPasswordVerifier{match: digest.MatchAdvanced}, nil
}
// Validate costs before a library sees the digest. Strict fields also prevent
// duplicate parameters or library defaults from changing the checked cost.
func boundedEmailCryptDecoder(encoded string) (func(string) (algorithm.Digest, error), error) {
p := strings.Split(encoded, "$")
if len(p) < 2 || p[0] != "" {
return nil, errEmailHashUnsupported
}
switch p[1] {
case "1", "5", "6":
keyLen, saltMax := 22, 8
if p[1] != "1" {
saltMax, keyLen = 16, 43
if p[1] == "6" {
keyLen = 86
}
if len(p) == 5 {
rounds, ok := strings.CutPrefix(p[2], "rounds=")
if !ok {
return nil, errEmailHashInvalid
}
if _, err := emailHashNumber(rounds, 1000, 1000000); err != nil {
return nil, err
}
p = append(p[:2:2], p[3:]...)
}
}
if len(p) != 4 || len(p[2]) > saltMax || !emailCryptChars(p[2]) || len(p[3]) != keyLen || !emailCryptChars(p[3]) {
return nil, errEmailHashInvalid
}
if p[1] == "1" {
return md5crypt.Decode, nil
}
return shacrypt.Decode, nil
case "2a", "2b", "2y":
if len(p) != 4 || len(p[2]) != 2 || len(p[3]) != 53 || !emailCryptChars(p[3]) {
return nil, errEmailHashInvalid
}
if _, err := emailHashNumber(p[2], 4, 14); err != nil {
return nil, err
}
return bcrypt.Decode, nil
case "argon2i", "argon2id":
if err := validateEmailArgon2(p); err != nil {
return nil, err
}
return argon2.Decode, nil
default:
return nil, errEmailHashUnsupported
}
}
func emailHashNumber(s string, min, max uint64) (uint64, error) {
if s == "" || strings.IndexFunc(s, func(r rune) bool { return r < '0' || r > '9' }) >= 0 {
return 0, errEmailHashInvalid
}
n, err := strconv.ParseUint(s, 10, 32)
if err != nil || n > max {
return 0, errEmailHashCost
}
if n < min {
return 0, errEmailHashInvalid
}
return n, nil
}
func emailCryptChars(s string) bool {
return s != "" && strings.IndexFunc(s, func(r rune) bool {
valid := r >= 'a' && r <= 'z' || r >= 'A' && r <= 'Z' || r >= '0' && r <= '9' || r == '.' || r == '/'
return !valid
}) < 0
}
func validateEmailArgon2(p []string) error {
if len(p) != 6 || p[2] != "v=19" {
return errEmailHashInvalid
}
params := strings.Split(p[3], ",")
if len(params) != 3 {
return errEmailHashInvalid
}
values := make(map[string]uint64, 3)
limits := map[string]uint64{"m": 65536, "t": 4, "p": 4}
for _, param := range params {
key, value, ok := strings.Cut(param, "=")
if !ok || limits[key] == 0 || values[key] != 0 {
return errEmailHashInvalid
}
n, err := emailHashNumber(value, 1, limits[key])
if err != nil {
return err
}
values[key] = n
}
if values["m"] < 8*values["p"] {
return errEmailHashInvalid
}
for i, field := range p[4:] {
decoded, err := base64.RawStdEncoding.Strict().DecodeString(field)
minLen := 8
if i == 1 {
minLen = 16
}
if err != nil || len(decoded) < minLen || len(decoded) > 64 {
return errEmailHashInvalid
}
}
return nil
}
func parseEmailDigest(scheme, encoding, encoded string) (*emailPasswordVerifier, error) {
var newHash func() hash.Hash
switch scheme {
case "PLAIN-MD5", "LDAP-MD5", "SMD5":
newHash = md5.New // #nosec G401 -- Verify legacy Dovecot hashes; never create stored credentials.
case "SHA", "SHA1", "SSHA":
newHash = sha1.New // #nosec G401 -- Verify legacy Dovecot hashes; never create stored credentials.
case "SHA256", "SSHA256":
newHash = sha256.New
case "SHA512", "SSHA512":
newHash = sha512.New
}
autoEncoding := encoding == ""
if autoEncoding {
encoding = "BASE64"
if scheme == "PLAIN-MD5" {
encoding = "HEX"
}
}
decoded, err := decodeEmailDigest(encoding, encoded)
size := newHash().Size()
salted := strings.HasPrefix(scheme, "SSHA") || scheme == "SMD5"
if autoEncoding && !salted && (err != nil || len(decoded) != size) {
other := "HEX"
if encoding == "HEX" {
other = "BASE64"
}
decoded, err = decodeEmailDigest(other, encoded)
}
if err != nil || len(decoded) < size || !salted && len(decoded) != size || salted && (len(decoded) == size || len(decoded) > size+64) {
return nil, errEmailHashInvalid
}
return &emailPasswordVerifier{match: func(candidate string) (bool, error) {
h := newHash()
_, _ = h.Write([]byte(candidate))
_, _ = h.Write(decoded[size:])
return subtle.ConstantTimeCompare(h.Sum(nil), decoded[:size]) == 1, nil
}}, nil
}
func decodeEmailDigest(encoding, encoded string) ([]byte, error) {
var decoded []byte
var err error
switch encoding {
case "HEX":
decoded, err = hex.DecodeString(encoded)
case "B64", "BASE64":
decoded, err = base64.StdEncoding.Strict().DecodeString(encoded)
default:
return nil, errEmailHashUnsupported
}
if err != nil {
return nil, errEmailHashInvalid
}
return decoded, nil
}
package checks
import (
"bufio"
"context"
"crypto/sha1" // #nosec G505 -- SHA1 is the digest format required by the Have I Been Pwned range API (https://haveibeenpwned.com/API/v3#PwnedPasswords). We send the first 5 chars of the digest and compare remaining chars against the returned list — HIBP does not offer a stronger-hash endpoint.
"crypto/sha256"
"errors"
"fmt"
"io"
"net/http"
"strconv"
"strings"
"sync"
"time"
"unicode"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// hibpClient is used for HIBP API requests.
var hibpClient = &http.Client{Timeout: 10 * time.Second}
// hibpEndpoint is the base URL queried for HIBP password-range lookups.
// Declared as a var (not const) so tests can swap in an httptest server.
// Production callers must not modify this.
var hibpEndpoint = "https://api.pwnedpasswords.com/range/"
// currentYear returns the current year. Called each audit cycle so
// long-running daemons don't use a stale year after Jan 1.
func currentYear() int { return time.Now().Year() }
// weakPasswordCache caches the bundled wordlist (loaded once).
var (
weakPasswordOnce sync.Once
weakPasswords []string
)
// parseShadowLine parses a Dovecot shadow line "mailbox:{scheme}hash".
// Returns empty strings if the line is malformed.
func parseShadowLine(line string) (mailbox, hash string) {
idx := strings.IndexByte(line, ':')
if idx <= 0 || idx >= len(line)-1 {
return "", ""
}
return line[:idx], line[idx+1:]
}
// isLockedHash returns true if the hash indicates a locked/disabled account.
func isLockedHash(hash string) bool {
if hash == "" {
return true
}
return hash[0] == '!' || hash[0] == '*'
}
// generateCandidates creates password candidates from username and domain.
// All candidates are >= 6 characters. No duplicates.
func generateCandidates(username, domain string) []string {
seen := make(map[string]bool)
var candidates []string
add := func(s string) {
if len(s) >= 6 && !seen[s] {
seen[s] = true
candidates = append(candidates, s)
}
}
domainLabel := domain
if idx := strings.IndexByte(domain, '.'); idx > 0 {
domainLabel = domain[:idx]
}
bases := []string{username, domainLabel}
for _, base := range bases {
add(base)
// Capitalize first letter variant
upper := capitalizeFirst(base)
add(upper)
// Year variants: current year +/- 2
year := currentYear()
for y := year - 2; y <= year+2; y++ {
ys := strconv.Itoa(y)
add(base + ys)
add(upper + ys)
}
// Two-digit suffix variants: 00-99
for n := 0; n <= 99; n++ {
suffix := fmt.Sprintf("%02d", n)
add(base + suffix)
add(upper + suffix)
}
}
return candidates
}
// capitalizeFirst returns the string with its first rune upper-cased.
func capitalizeFirst(s string) string {
if len(s) == 0 {
return s
}
runes := []rune(s)
runes[0] = unicode.ToUpper(runes[0])
return string(runes)
}
// hashFingerprint returns a SHA256 hex fingerprint of a password hash
// (used for change detection -- re-audit only when hash changes).
func hashFingerprint(hash string) string {
h := sha256.Sum256([]byte(hash))
return fmt.Sprintf("%x", h[:])
}
// parseHIBPCount searches a HIBP range response body for a hash suffix
// and returns the breach count. Returns 0 if not found.
func parseHIBPCount(body, suffix string) int {
upperSuffix := strings.ToUpper(suffix)
for _, line := range strings.Split(body, "\n") {
line = strings.TrimSpace(line)
parts := strings.SplitN(line, ":", 2)
if len(parts) != 2 {
continue
}
if strings.ToUpper(strings.TrimSpace(parts[0])) == upperSuffix {
count, err := strconv.Atoi(strings.TrimSpace(parts[1]))
if err != nil {
return 0
}
return count
}
}
return 0
}
// checkHIBP queries the HIBP Pwned Passwords API for a plaintext password.
// Returns the breach count (0 if not found or on error).
func checkHIBP(plaintext string) int {
return checkHIBPWithContext(context.Background(), plaintext)
}
func checkHIBPWithContext(ctx context.Context, plaintext string) int {
if ctx == nil {
ctx = context.Background()
}
if ctx.Err() != nil {
return 0
}
// #nosec G401 -- SHA1 is mandated by the HIBP Pwned Passwords range API; see import comment.
h := sha1.Sum([]byte(plaintext))
hex := fmt.Sprintf("%X", h[:])
prefix := hex[:5]
suffix := hex[5:]
req, err := http.NewRequestWithContext(ctx, http.MethodGet, hibpEndpoint+prefix, nil)
if err != nil {
return 0
}
resp, err := hibpClient.Do(req)
if err != nil {
return 0
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return 0
}
body, err := io.ReadAll(resp.Body)
if err != nil {
return 0
}
return parseHIBPCount(string(body), suffix)
}
type shadowFile struct {
path string
account string
domain string
}
type mailboxEntry struct {
account string
domain string
mailbox string
hash string
}
// discoverShadowFiles finds all Dovecot shadow files under /home/*/etc/*/shadow.
// Results are ranked by mtime desc so recently touched mailbox password files
// are inspected first when the check timeout cuts work short. maxFiles caps
// iteration; 0 disables the cap.
func discoverShadowFiles(ctx context.Context, maxFiles int) []shadowFile {
matches, _ := accountHomeGlob("*/etc/*/shadow")
ranked := rankPathsByMtimeDesc(ctx, matches, maxFiles)
results := make([]shadowFile, 0, len(ranked))
for _, m := range ranked {
parts := strings.Split(m, "/")
// /home/{account}/etc/{domain}/shadow
if len(parts) >= 5 {
results = append(results, shadowFile{
path: m,
account: parts[2],
domain: parts[4],
})
}
}
return results
}
// readShadowFile reads all mailbox entries from a Dovecot shadow file.
func readShadowFile(sf shadowFile) []mailboxEntry {
f, err := osFS.Open(sf.path)
if err != nil {
return nil
}
defer f.Close()
var entries []mailboxEntry
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
mailbox, hash := parseShadowLine(line)
if mailbox == "" || hash == "" {
continue
}
if isLockedHash(hash) {
continue
}
entries = append(entries, mailboxEntry{
account: sf.account,
domain: sf.domain,
mailbox: mailbox,
hash: hash,
})
}
return entries
}
// loadWeakPasswords reads the bundled wordlist once and caches it.
func loadWeakPasswords() []string {
weakPasswordOnce.Do(func() {
f, err := osFS.Open("/opt/csm/configs/weak_passwords.txt")
if err != nil {
return
}
defer f.Close()
scanner := bufio.NewScanner(f)
for scanner.Scan() {
word := strings.TrimSpace(scanner.Text())
if len(word) < 6 || strings.HasPrefix(word, "#") {
continue
}
weakPasswords = append(weakPasswords, word)
}
})
return weakPasswords
}
// CheckEmailPasswords audits Dovecot email account passwords for weak/predictable
// patterns. Uses internal throttle: skips if last refresh was less than
// password_check_interval_min ago.
func CheckEmailPasswords(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
db := store.Global()
if db == nil {
return nil
}
// Internal throttle. Unlike CheckOutdatedPlugins, which throttles only its
// cache refresh and still evaluates that cache every cycle, this check has
// nothing to report from until it runs again. Returning nothing silently
// would read as "ran, found nothing" and let the runner retire the weak
// passwords the last run found, so the skipped scope is declared.
if !ForceAll {
lastRefresh := db.GetEmailPWLastRefresh()
interval := time.Duration(cfg.EmailProtection.PasswordCheckIntervalMin) * time.Minute
if time.Since(lastRefresh) < interval {
markCheckSkipped(ctx, "email_weak_password")
return nil
}
}
shadowFiles := discoverShadowFiles(ctx, accountScanMaxFiles(ctx, cfg))
if len(shadowFiles) == 0 {
return nil
}
// Collect all mailbox entries
var allEntries []mailboxEntry
for _, sf := range shadowFiles {
if ctx.Err() != nil {
return nil
}
allEntries = append(allEntries, readShadowFile(sf)...)
}
if ctx.Err() != nil {
return nil
}
if len(allEntries) == 0 {
_ = db.SetEmailPWLastRefresh(time.Now())
return nil
}
var mu sync.Mutex
var findings []alert.Finding
var incomplete int
var unauditable int
var wg sync.WaitGroup
sem := make(chan struct{}, 5)
batch := emailMailboxAudits.begin(len(allEntries), cap(sem))
defer batch.abandon(ctx)
mailboxes:
for i, entry := range allEntries {
select {
case sem <- struct{}{}:
case <-ctx.Done():
break mailboxes
}
work := batch.tasks[i]
work.admit()
wg.Go(func() {
defer func() { <-sem }()
work.run(ctx, checkTimeout, func() {
if err := ctx.Err(); err != nil {
work.withdraw(err)
return
}
fullMailbox := entry.mailbox + "@" + entry.domain
storeKey := fmt.Sprintf("email:pwaudit:%s:%s", entry.account, fullMailbox)
// Older versions also cached failed verifications as clean.
fp := "v2:" + hashFingerprint(entry.hash)
if db.GetMetaString(storeKey) == fp {
return
}
finding, err := auditEmailPassword(ctx, entry)
work.progress()
if err != nil && !errors.Is(err, context.Canceled) &&
!errors.Is(err, errEmailHashUnsupported) && !errors.Is(err, errEmailHashInvalid) &&
!errors.Is(err, errEmailHashCost) && !errors.Is(err, errEmailCandidate) {
work.fail()
}
mu.Lock()
if err != nil {
if emailHashPermanentlyUnauditable(err) {
unauditable++
} else {
incomplete++
}
} else if finding != nil {
findings = append(findings, *finding)
}
mu.Unlock()
if err != nil {
return
}
if err := ctx.Err(); err != nil {
work.withdraw(err)
return
}
if err := db.SetMetaString(storeKey, fp); err != nil {
work.fail()
}
})
})
}
batch.abandon(ctx)
wg.Wait()
if incomplete > 0 || unauditable > 0 || ctx.Err() != nil {
markCheckIncomplete(ctx, "email_weak_password")
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "email_password_audit_incomplete",
Message: "Email password audit did not complete",
Details: fmt.Sprintf("Mailboxes to retry: %d. Mailboxes whose stored hash this build cannot audit at all: %d. See the email password audit documentation for supported formats and limits.", incomplete, unauditable),
// The mailbox counts in Details are not the condition.
DedupKey: "unfinished_verification",
})
}
// A hash this build cannot audit never becomes auditable by rerunning, so
// it must not hold the stamp back: that turned the interval into "every
// scan" and re-verified the whole mailbox set each cycle. A cycle skipped
// by the stamp now declares its skipped scope, so the findings survive it.
if incomplete == 0 && ctx.Err() == nil {
_ = db.SetEmailPWLastRefresh(time.Now())
}
return findings
}
// emailHashPermanentlyUnauditable reports whether a verification failure is a
// property of the stored hash rather than a condition that may clear.
func emailHashPermanentlyUnauditable(err error) bool {
// The parser returns these sentinels directly. A wrapper can describe a
// transient verification failure, so unwrapping is not proof that only
// the stored hash prevented this audit from completing.
return err == errEmailHashUnsupported || err == errEmailHashInvalid || err == errEmailHashCost
}
func auditEmailPassword(ctx context.Context, entry mailboxEntry) (*alert.Finding, error) {
verifier, err := parseEmailPasswordHash(entry.hash)
if err != nil {
return nil, err
}
matched, err := verifier.firstMatch(ctx, generateCandidates(entry.mailbox, entry.domain))
if err != nil {
return nil, err
}
matchType := "heuristic"
if matched == "" {
matched, err = verifier.firstMatch(ctx, loadWeakPasswords())
if err != nil {
return nil, err
}
matchType = "wordlist"
}
if matched == "" {
return nil, nil
}
fullMailbox := entry.mailbox + "@" + entry.domain
details := fmt.Sprintf("Account: %s\nMailbox: %s\nMatch type: %s", entry.account, fullMailbox, matchType)
if breachCount := checkHIBPWithContext(ctx, matched); breachCount > 0 {
details += fmt.Sprintf("\nHIBP: password found in %d data breaches", breachCount)
}
return &alert.Finding{
Severity: alert.Critical,
Check: "email_weak_password",
Message: fmt.Sprintf("Weak email password for %s (account: %s)", fullMailbox, entry.account),
Details: details,
Domain: entry.domain,
Mailbox: fullMailbox,
}, nil
}
package checks
import (
"bufio"
"bytes"
"context"
"fmt"
"io"
"mime"
"net/textproto"
"path/filepath"
"regexp"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailspool"
emime "github.com/pidginhost/csm/internal/mime"
"github.com/pidginhost/csm/internal/state"
)
const emailBodySampleSize = 65536 // Analyze the first 64KB of email body
// Suspicious mailer headers that indicate mass mailing scripts. PHPMailer is
// deliberately absent: it is the default WordPress transport and appears on
// essentially all legitimate WordPress mail, so it carries no signal. The
// "phpmail" substring is excluded for the same reason (it matches PHPMailer).
var suspiciousMailers = []string{
"swiftmailer", "mass mailer", "bulk mailer",
"leaf phpmailer", "mail.php",
}
// Known safe mailers that should not be flagged.
var safeMailers = []string{
"wordpress", "woocommerce", "roundcube", "squirrelmail",
"thunderbird", "outlook", "apple mail", "cpanel",
"postfix", "exim", "dovecot",
}
// Phishing URL patterns in email body.
var emailPhishPatterns = []string{
".workers.dev",
"//bit.ly/", "//tinyurl.com/", "//is.gd/", "//rb.gy/",
"//t.co/",
"/redir?url=", "/redirect?url=", "link?url=",
"effi.redir",
}
// Phishing language in email body.
var emailPhishLanguage = []string{
"verify your account",
"confirm your identity",
"unusual activity",
"your account will be",
"suspended unless",
"click here to verify",
"update your payment",
"confirm your email address",
"security alert",
"unauthorized access",
}
// Brand impersonation in email body (when sender doesn't match).
var emailBrandNames = []string{
"paypal", "microsoft", "apple", "google", "amazon",
"netflix", "facebook", "instagram", "bank of",
"wells fargo", "chase", "citibank",
}
// CheckOutboundEmailContent samples outbound email content from Exim spool
// for phishing URLs, credential harvesting language, suspicious mailers,
// and Reply-To mismatches.
func CheckOutboundEmailContent(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Read recent outbound messages from exim_mainlog
lines := tailFile("/var/log/exim_mainlog", 200)
if len(lines) == 0 {
return nil
}
// Extract message IDs for outbound emails
msgIDRegex := regexp.MustCompile(`^\S+\s+(\S+)\s+<=\s+(\S+)`)
scanned := make(map[string]bool) // avoid scanning same message twice
for _, line := range lines {
matches := msgIDRegex.FindStringSubmatch(line)
if len(matches) < 3 {
continue
}
msgID := matches[1]
sender := matches[2]
if sender == "<>" || sender == "" {
continue // bounce message
}
if scanned[msgID] {
continue
}
scanned[msgID] = true
// Scan the message
result := scanEximMessage(msgID, sender, cfg)
if result != nil {
findings = append(findings, *result)
}
}
return findings
}
// scanEximMessage reads an Exim spool message and checks for suspicious content.
func scanEximMessage(msgID, sender string, cfg *config.Config) *alert.Finding {
// Exim spool paths
// Headers: /var/spool/exim/input/{msgID}-H
// Body: /var/spool/exim/input/{msgID}-D
spoolDirs := []string{
"/var/spool/exim/input",
"/var/spool/exim4/input",
}
var headerPath, bodyPath string
for _, dir := range spoolDirs {
h := filepath.Join(dir, msgID+"-H")
b := filepath.Join(dir, msgID+"-D")
if _, err := osFS.Stat(h); err == nil {
headerPath = h
bodyPath = b
break
}
}
if headerPath == "" {
return nil // message already delivered/removed from spool
}
var indicators []string
// Read and parse headers via the emailspool Exim -H parser. We go through
// osFS.ReadFile + bytes.NewReader rather than emailspool.ParseHeaders(path)
// so check tests can inject mock spool contents through the existing
// osFS seam.
headerData, err := osFS.ReadFile(headerPath)
if err != nil {
return nil
}
parsed, err := emailspool.ParseHeadersReader(bytes.NewReader(headerData))
if err != nil {
// Malformed or truncated -H file: nothing to check, skip silently as
// the previous loose parse would have done.
return nil
}
// Check 1: Reply-To mismatch
if parsed.From != "" && parsed.ReplyTo != "" {
fromDomain := emailspool.ExtractDomain(parsed.From)
replyDomain := emailspool.ExtractDomain(parsed.ReplyTo)
if fromDomain != "" && replyDomain != "" && fromDomain != replyDomain {
indicators = append(indicators, fmt.Sprintf("Reply-To mismatch: From=%s, Reply-To=%s", fromDomain, replyDomain))
}
}
// Check 2: Suspicious X-Mailer (fall back to User-Agent, which the
// emailspool parser surfaces alongside the other RFC 5322 fields).
mailer := parsed.XMailer
if mailer == "" {
mailer = parsed.UserAgent
}
if mailer != "" {
mailerLower := strings.ToLower(mailer)
isSafe := false
for _, safe := range safeMailers {
if strings.Contains(mailerLower, safe) {
isSafe = true
break
}
}
if !isSafe {
for _, suspicious := range suspiciousMailers {
if strings.Contains(mailerLower, suspicious) {
indicators = append(indicators, fmt.Sprintf("suspicious mailer: %s", strings.TrimSpace(mailer)))
break
}
}
}
}
// Check 3: Spoofed display name (brand name in From: but sender is not that brand)
if parsed.From != "" {
fromLower := strings.ToLower(parsed.From)
senderDomain := emailspool.ExtractDomain(sender)
for _, brand := range emailBrandNames {
if strings.Contains(fromLower, brand) && !strings.Contains(strings.ToLower(senderDomain), brand) {
indicators = append(indicators, fmt.Sprintf("spoofed brand in From: '%s' (actual sender: %s)", strings.TrimSpace(parsed.From), sender))
break
}
}
}
// Read and analyze a bounded body sample.
bodyData, _ := osFS.ReadFile(bodyPath)
if len(bodyData) > emailBodySampleSize {
bodyData = bodyData[:emailBodySampleSize]
}
if len(bodyData) > 0 {
bodyLower := strings.ToLower(string(bodyData))
indicators = append(indicators, bodyContentIndicators(bodyLower)...)
// A base64 Content-Transfer-Encoding is standard MIME framing, not an
// indicator by itself. Decode the body and run the same content checks
// on the plaintext so a payload hidden behind base64 is caught -- the
// raw base64 blob would otherwise sail past every text pattern.
// Only the exact Exim marker is framing; arbitrary body text ending
// in -D must remain available to the transfer decoder.
if marker, body, ok := bytes.Cut(bodyData, []byte("\n")); ok && string(bytes.TrimSuffix(marker, []byte("\r"))) == msgID+"-D" {
bodyData = body
}
mimeHeaders := emime.ParseSpoolMIMEHeaders(headerData)
if decoded := decodeBase64Body(bodyData, mimeHeaders.Get("Content-Type"), mimeHeaders.Get("Content-Transfer-Encoding")); decoded != "" {
indicators = append(indicators, bodyContentIndicators(strings.ToLower(decoded))...)
}
}
indicators = uniqueStrings(indicators)
// Require at least two independent indicators before alerting. A lone weak
// signal (Reply-To mismatch, one brand word, one suspicious header) fires
// on ordinary legitimate mail, so a single indicator is not enough.
if len(indicators) < 2 {
return nil
}
severity := alert.High
if len(indicators) >= 3 {
severity = alert.Critical
}
_, senderDomain := alert.SplitEmail(sender)
return &alert.Finding{
Severity: severity,
Check: "email_phishing_content",
Message: fmt.Sprintf("Suspicious outbound email from %s (message: %s)", sender, msgID),
Details: fmt.Sprintf("Indicators:\n- %s", strings.Join(indicators, "\n- ")),
Domain: senderDomain,
Mailbox: sender,
}
}
// bodyContentIndicators runs the phishing-URL and credential-harvesting content
// checks over an already-lowercased body text. Shared by the raw-body pass and
// the decoded-base64 pass so both apply identical logic.
func bodyContentIndicators(bodyLower string) []string {
var indicators []string
for _, pattern := range emailPhishPatterns {
if strings.Contains(bodyLower, pattern) {
indicators = append(indicators, fmt.Sprintf("phishing URL pattern: %s", pattern))
break
}
}
harvestCount := 0
for _, phrase := range emailPhishLanguage {
if strings.Contains(bodyLower, phrase) {
harvestCount++
}
}
if harvestCount >= 2 {
indicators = append(indicators, fmt.Sprintf("credential harvesting language (%d phrases)", harvestCount))
}
return indicators
}
// decodeBase64Body walks actual MIME parts instead of searching body text for
// header names. Boundaries are case-sensitive and only delimit their enclosing
// multipart; other dash lines remain part of the lenient base64 input.
func decodeBase64Body(raw []byte, contentType, transferEncoding string) string {
var decodedParts []string
var walk func([]byte, string, string, int)
walk = func(body []byte, ct, cte string, depth int) {
// Match attachment extraction's nesting bound for attacker-written MIME.
if depth >= 16 {
return
}
visit := func(part []byte) {
reader := bufio.NewReader(bytes.NewReader(part))
headers, err := textproto.NewReader(reader).ReadMIMEHeader()
if err != nil {
// A damaged part must not hide its siblings.
return
}
data, _ := io.ReadAll(reader)
walk(data, headers.Get("Content-Type"), headers.Get("Content-Transfer-Encoding"), depth+1)
}
encoding := strings.ToLower(strings.TrimSpace(withoutMIMEComments(cte)))
mediaType, params, _ := mime.ParseMediaType(withoutMIMEComments(ct))
if strings.HasPrefix(mediaType, "multipart/") {
boundary := params["boundary"]
if boundary == "" {
return
}
visitBase64MIMEParts(body, boundary, visit)
return
}
if (mediaType == "message/rfc822" || mediaType == "message/global") && encoding != "base64" {
visit(body)
return
}
if encoding == "base64" {
for _, variant := range emime.DecodeTransferVariants("base64", body) {
if len(variant) > 0 {
decodedParts = append(decodedParts, string(variant))
}
}
}
}
walk(raw, contentType, transferEncoding, 0)
return strings.Join(decodedParts, "\n")
}
// Frame parts before parsing their headers so malformed headers cannot stop
// iteration at the next sibling. A clipped final part is still worth scanning.
func visitBase64MIMEParts(body []byte, boundary string, visit func([]byte)) {
delimiter := "--" + boundary
start := -1
for offset := 0; offset < len(body); {
end := bytes.IndexByte(body[offset:], '\n')
next := len(body)
if end >= 0 {
next = offset + end + 1
}
line := strings.TrimRight(string(body[offset:next]), "\r\n \t")
if line == delimiter || line == delimiter+"--" {
if start >= 0 {
visit(body[start:offset])
}
if line == delimiter+"--" {
return
}
start = next
}
offset = next
}
if start >= 0 {
visit(body[start:])
}
}
// MIME comments are allowed around tokens, but parentheses inside quoted
// parameters (notably boundary values) are literal. Keep a space where a
// comment was removed so separate tokens cannot be joined accidentally.
func withoutMIMEComments(value string) string {
var out strings.Builder
depth := 0
quoted := false
for i := 0; i < len(value); i++ {
b := value[i]
if depth > 0 {
switch b {
case '\\':
if i+1 < len(value) {
i++
}
case '(':
depth++
case ')':
depth--
}
continue
}
if quoted {
out.WriteByte(b)
if b == '\\' && i+1 < len(value) {
i++
out.WriteByte(value[i])
} else if b == '"' {
quoted = false
}
continue
}
switch b {
case '(':
depth = 1
out.WriteByte(' ')
case '"':
quoted = true
out.WriteByte(b)
default:
out.WriteByte(b)
}
}
return out.String()
}
package checks
import (
"context"
"fmt"
"path/filepath"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// uidStringToUser converts a string-form uid (as read from /proc/<pid>/status)
// to a username via the cached LookupUser. Falls back to the raw uid string
// when it is not parseable as a uint32.
func uidStringToUser(uid string) string {
u64, err := strconv.ParseUint(uid, 10, 32)
if err != nil {
return uid
}
return LookupUser(uint32(u64))
}
// CheckDatabaseDumps detects mysqldump/pg_dump processes running under
// non-root users - potential data exfiltration.
func CheckDatabaseDumps(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
dumpTools := []string{"mysqldump", "pg_dump", "mongodump"}
procs, _ := osFS.Glob("/proc/[0-9]*/cmdline")
for _, cmdPath := range procs {
pid := filepath.Base(filepath.Dir(cmdPath))
// Read UID
statusData, _ := osFS.ReadFile(filepath.Join("/proc", pid, "status"))
var uid string
for _, line := range strings.Split(string(statusData), "\n") {
if strings.HasPrefix(line, "Uid:\t") {
fields := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(fields) > 0 {
uid = fields[0]
}
}
}
// Skip root - root may run legitimate backups
if uid == "0" || uid == "" {
continue
}
cmdline, err := osFS.ReadFile(cmdPath)
if err != nil {
continue
}
cmdStr := strings.ReplaceAll(string(cmdline), "\x00", " ")
safeCmdStr := redactProcCommandLine(cmdline)
for _, tool := range dumpTools {
if strings.Contains(cmdStr, tool) {
user := uidStringToUser(uid)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "database_dump",
Message: fmt.Sprintf("Database dump by non-root user: %s (%s)", user, tool),
Details: fmt.Sprintf("PID: %s, UID: %s, cmdline: %s", pid, uid, safeCmdStr),
})
break
}
}
}
return findings
}
// CheckOutboundPasteSites detects connections to known paste/exfiltration sites.
func CheckOutboundPasteSites(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Check running processes for connections to paste sites
pasteSites := []string{
"pastebin.com", "hastebin.com", "ghostbin.co",
"paste.ee", "dpaste.org", "gist.githubusercontent.com",
"raw.githubusercontent.com", "transfer.sh",
"file.io", "0x0.st", "ix.io",
}
procs, _ := osFS.Glob("/proc/[0-9]*/cmdline")
for _, cmdPath := range procs {
pid := filepath.Base(filepath.Dir(cmdPath))
statusData, _ := osFS.ReadFile(filepath.Join("/proc", pid, "status"))
var uid string
for _, line := range strings.Split(string(statusData), "\n") {
if strings.HasPrefix(line, "Uid:\t") {
fields := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(fields) > 0 {
uid = fields[0]
}
}
}
if uid == "0" {
continue
}
cmdline, err := osFS.ReadFile(cmdPath)
if err != nil {
continue
}
cmdStr := strings.ToLower(strings.ReplaceAll(string(cmdline), "\x00", " "))
safeCmdStr := redactProcCommandLine(cmdline)
// Check if process is connecting to paste sites
for _, site := range pasteSites {
if strings.Contains(cmdStr, site) {
user := uidStringToUser(uid)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "exfiltration_paste_site",
Message: fmt.Sprintf("Process connecting to paste/exfiltration site: %s (user: %s)", site, user),
Details: fmt.Sprintf("PID: %s, cmdline: %s", pid, safeCmdStr),
TenantID: HostingAccountForUser(user),
})
break
}
}
}
return findings
}
package checks
import (
"bytes"
"compress/gzip"
"context"
"crypto/tls"
"errors"
"fmt"
"io"
"io/fs"
"net"
"net/http"
"net/url"
"path/filepath"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// userdataDomainsPath is cPanel's authoritative domain->docroot map. It covers
// addon and subdomain docroots that /home/*/public_html misses.
const userdataDomainsPath = "/etc/userdatadomains"
const cpanelInstallPath = "/usr/local/cpanel"
// Bounds so the deep scan stays cheap on a host with hundreds of docroots.
const (
exposureMaxFilesPerRoot = 4000
exposureMaxDirsPerRoot = 4000
exposureProbeConnTimeout = 4 * time.Second
exposureProbeTotalTimeout = 6 * time.Second
)
// walkSkipDirs are directory names never worth descending for a leaked dump or
// backup and expensive to traverse.
var walkSkipDirs = map[string]bool{
"node_modules": true, ".git": true, ".svn": true,
}
// CheckExposedFiles scans every cPanel docroot for sensitive files (database
// dumps, backup archives, config/source backups, phpinfo) that the web server
// actually serves, and reports each confirmed exposure. It reads only response
// headers except for a bounded phpinfo body used to reject empty stubs.
func CheckExposedFiles(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
content, err := osFS.ReadFile(userdataDomainsPath)
if err != nil {
// Absence is an expected no-op only on a non-cPanel host. Other failures
// make the run incomplete: retaining prior findings is safer than
// clearing them because the map was temporarily unreadable.
if vhostMapFailureIsIncomplete(err) {
markCheckIncomplete(ctx, "exposed_files")
}
return nil
}
vhosts, complete := parseUserdataDomainsChecked(string(content))
if !complete {
markCheckIncomplete(ctx, "exposed_files")
}
return scanVhostsForExposure(ctx, dedupVhostsByDocroot(vhosts), cfg)
}
func vhostMapFailureIsIncomplete(readErr error) bool {
if !errors.Is(readErr, fs.ErrNotExist) {
return true
}
_, statErr := osFS.Stat(cpanelInstallPath)
// A missing map is an expected no-op only when cPanel itself is absent.
// If the install path exists (or cannot be checked), preserve prior
// findings until the authoritative map is readable again.
return statErr == nil || !errors.Is(statErr, fs.ErrNotExist)
}
// dedupVhostsByDocroot keeps one vhost per docroot (multiple domains -- www,
// parked -- can share a docroot; one probe host is enough). Prefer a vhost with
// a usable serving IP, then a real main/addon/subdomain over a parked alias,
// since parked domains commonly redirect and cannot confirm an exposure.
func dedupVhostsByDocroot(vhosts []vhost) []vhost {
chosen := make(map[string]int, len(vhosts))
out := vhosts[:0:0]
for _, vh := range vhosts {
key := filepath.Clean(vh.docroot)
if i, ok := chosen[key]; ok {
if betterProbeVhost(vh, out[i]) {
out[i] = vh
}
continue
}
chosen[key] = len(out)
out = append(out, vh)
}
return out
}
func betterProbeVhost(candidate, current vhost) bool {
candidateUsable := probeHost(candidate) != ""
currentUsable := probeHost(current) != ""
if candidateUsable != currentUsable {
return candidateUsable
}
return preferredProbeVhost(candidate) && !preferredProbeVhost(current)
}
func preferredProbeVhost(vh vhost) bool {
switch strings.ToLower(strings.TrimSpace(vh.typ)) {
case "main", "addon", "sub":
return true
default:
return false
}
}
// scanVhostsForExposure walks each vhost's docroot, classifies candidate files,
// confirms the server serves them, and emits a finding per confirmed exposure.
func scanVhostsForExposure(ctx context.Context, vhosts []vhost, cfg *config.Config) []alert.Finding {
depth := exposureScanDepth(cfg)
var findings []alert.Finding
for _, vh := range vhosts {
if ctx.Err() != nil {
return findings
}
host := probeHost(vh)
if host == "" {
// A loopback fallback is not a meaningful reachability check:
// LiteSpeed rejects loopback-originated requests even when the same
// vhost serves the file on its configured address. Skip this vhost
// and preserve findings from the last complete scan instead.
markCheckIncomplete(ctx, "exposed_files")
continue
}
paths, complete := walkExposureCandidates(ctx, vh.docroot, depth)
if !complete {
markCheckIncomplete(ctx, "exposed_files")
}
for _, path := range paths {
if ctx.Err() != nil {
return findings
}
class := classifyExposedPath(path)
if class == classNone {
// An archive named after the site it holds carries no backup
// token, so the name tells us nothing. Its entry list does.
holdsSite, complete := archiveSiteBackupStatus(ctx, path)
if !complete {
markCheckIncomplete(ctx, "exposed_files")
}
if holdsSite {
class = classBackupArchive
}
}
if class == classNone {
continue
}
rel := relURLPath(vh.docroot, path)
if rel == "" {
continue
}
class = demoteSampleSQL(class, rel)
pr := webProber.probe(ctx, vh.domain, host, rel)
confirmed := confirmExposure(class, pr)
if !confirmed && (!pr.reachable || pr.partial) {
// A transport failure cannot prove that a previous exposure was
// fixed. Keep prior findings until a later probe gets a response.
markCheckIncomplete(ctx, "exposed_files")
}
if !confirmed {
continue
}
if class == classPHPInfo {
// Headers cannot distinguish a real dump from a stub whose
// phpinfo() call is commented out or guarded: both answer 200
// text/html. Only a body carrying actual phpinfo output is an
// information disclosure.
scheme, exposed, complete := confirmPHPInfoBody(ctx, pr.scheme, vh.domain, host, rel)
if !complete {
markCheckIncomplete(ctx, "exposed_files")
}
if !exposed {
continue
}
pr.scheme = scheme
}
findings = append(findings, buildExposedFinding(vh, path, rel, class, pr))
}
}
return findings
}
func exposureScanDepth(cfg *config.Config) int {
depth := config.DefaultExposedFileScanDepth
if cfg != nil && cfg.Thresholds.ExposedFileScanDepth > 0 {
depth = cfg.Thresholds.ExposedFileScanDepth
}
if depth > config.MaxExposedFileScanDepth {
return config.MaxExposedFileScanDepth
}
return depth
}
// walkExposureCandidates returns files under docroot down to maxDepth directory
// levels, capped at exposureMaxFilesPerRoot to bound I/O on large trees. It
// visits shallower directories first so a large cache subtree cannot consume
// the cap before root-level leaks are considered.
func walkExposureCandidates(ctx context.Context, docroot string, maxDepth int) ([]string, bool) {
return walkExposureCandidatesLimit(ctx, docroot, maxDepth, exposureMaxFilesPerRoot)
}
func walkExposureCandidatesLimit(ctx context.Context, docroot string, maxDepth, maxFiles int) ([]string, bool) {
if maxFiles <= 0 {
return nil, false
}
type pendingDir struct {
path string
depth int
}
queue := []pendingDir{{path: docroot}}
queuedDirs := 1
var out []string
complete := true
for len(queue) > 0 && len(out) < maxFiles {
if ctx.Err() != nil {
return out, false
}
dir := queue[0]
queue = queue[1:]
entries, err := osFS.ReadDir(dir.path)
if err != nil {
complete = false
continue
}
for _, e := range entries {
if ctx.Err() != nil {
return out, false
}
if len(out) >= maxFiles {
complete = false
break
}
if e.IsDir() {
// A checked-out repository is never descended (thousands of
// objects), but its marker file is a candidate in its own
// right: a web-served .git/ or .svn/ hands out the site's
// source and, routinely, its credentials.
for _, marker := range repoMetadataMarkers(e.Name()) {
markerPath := filepath.Join(dir.path, e.Name(), marker)
if info, statErr := osFS.Stat(markerPath); statErr == nil && info.Mode().IsRegular() {
out = append(out, markerPath)
}
}
if dir.depth < maxDepth && !walkSkipDirs[e.Name()] {
if queuedDirs >= exposureMaxDirsPerRoot {
complete = false
continue
}
queue = append(queue, pendingDir{
path: filepath.Join(dir.path, e.Name()),
depth: dir.depth + 1,
})
queuedDirs++
}
continue
}
out = append(out, filepath.Join(dir.path, e.Name()))
}
}
if len(queue) > 0 {
complete = false
}
return out, complete
}
func relURLPath(docroot, path string) string {
r, err := filepath.Rel(docroot, path)
if err != nil || r == ".." || strings.HasPrefix(r, ".."+string(filepath.Separator)) {
return ""
}
return "/" + filepath.ToSlash(r)
}
func buildExposedFinding(vh vhost, path, rel string, class exposedClass, pr probeResult) alert.Finding {
size := int64(-1)
if fi, err := osFS.Stat(path); err == nil {
size = fi.Size()
}
scheme := pr.scheme
if scheme == "" {
scheme = "https"
}
exposureURL := (&url.URL{Scheme: scheme, Host: vh.domain, Path: rel}).String()
return alert.Finding{
Severity: class.severity(),
Check: class.findingName(),
Message: fmt.Sprintf("Web-exposed %s reachable at %s", exposureLabel(class), exposureURL),
Details: fmt.Sprintf("File: %s (%d bytes), served as %q. Remove it from the web root or deny HTTP access.",
path, size, pr.contentType),
FilePath: path,
Domain: vh.domain,
TenantID: vh.user,
Timestamp: time.Now(),
}
}
func exposureLabel(class exposedClass) string {
switch class {
case classRepoMetadata:
return "version-control repository (source and history)"
case classConfigLeak:
return "configuration/credentials file"
case classDBDump:
return "database dump"
case classBackupArchive:
return "site backup archive"
case classSourceBackup:
return "source-code backup"
case classPHPInfo:
return "phpinfo diagnostic"
case classSampleSQL:
return "framework/sample SQL file"
default:
return "sensitive file"
}
}
// ---------------------------------------------------------------------------
// Reachability probe seam
// ---------------------------------------------------------------------------
// webProbe abstracts a headers-only reachability check against the local web
// server. Tests inject a fake; production pins the connection to the vhost's
// serving IP while using its domain as SNI/Host, and reads only the status line
// and Content-Type.
type webProbe interface {
probe(ctx context.Context, domain, host, urlPath string) probeResult
probeComplete(ctx context.Context, domain, host, urlPath string) probeResult
}
var webProber webProbe = realWebProbe{}
// SetWebProbe replaces the reachability prober. Test-only seam.
func SetWebProbe(p webProbe) { webProber = p }
type realWebProbe struct{}
func (realWebProbe) probe(ctx context.Context, domain, host, urlPath string) probeResult {
return probeLocalSchemes(ctx, domain, host, urlPath, true, doLocalProbe)
}
// probeComplete always attempts both origin protocols. Detection may stop once
// one protocol proves a raw exposure, but verification needs every protocol to
// answer before a negative result can safely clear an earlier finding.
func (realWebProbe) probeComplete(ctx context.Context, domain, host, urlPath string) probeResult {
return probeLocalSchemes(ctx, domain, host, urlPath, false, doLocalProbe)
}
type localProbeFunc func(context.Context, string, string, string, string) (probeResult, bool)
func probeLocalSchemes(
ctx context.Context,
domain, host, urlPath string,
stopOnRawExposure bool,
probeOne localProbeFunc,
) probeResult {
results := make([]probeResult, 0, 2)
partial := false
for _, scheme := range []string{"https", "http"} {
if pr, ok := probeOne(ctx, scheme, domain, host, urlPath); ok {
results = append(results, pr)
// A successful non-HTML response is sufficient for every raw
// leak class; avoid an unnecessary second request.
if stopOnRawExposure && successfulProbe(pr) && !isHTMLContentType(pr.contentType) {
return pr
}
} else {
partial = true
}
}
pr := bestProbeResult(results)
pr.partial = partial
return pr
}
func successfulProbe(pr probeResult) bool {
return pr.status == http.StatusOK || pr.status == http.StatusPartialContent
}
// bestProbeResult prefers a successful response over an HTTP error and a raw
// response over HTML. This lets an HTTP-only exposure win when HTTPS redirects,
// blocks the path, or executes a different handler.
func bestProbeResult(results []probeResult) probeResult {
var firstReachable probeResult
var firstSuccess probeResult
for _, pr := range results {
if !pr.reachable {
continue
}
if !firstReachable.reachable {
firstReachable = pr
}
if !successfulProbe(pr) {
continue
}
if !isHTMLContentType(pr.contentType) {
return pr
}
if !firstSuccess.reachable {
firstSuccess = pr
}
}
if firstSuccess.reachable {
return firstSuccess
}
return firstReachable
}
// doLocalProbe issues one HEAD (falling back to a ranged GET when HEAD is
// disallowed) to the vhost's own serving IP, presenting domain as the vhost. It
// never reads the response body. The dial is pinned to the host's configured
// serving address (never DNS-resolved) so the request reaches this origin, not
// a CDN, and stays on-box. Loopback is not used: LiteSpeed answers 403 to
// loopback-originated requests even for files it serves on the public IP.
func doLocalProbe(ctx context.Context, scheme, domain, host, urlPath string) (probeResult, bool) {
host = normalizeServingIP(host)
if host == "" {
return probeResult{}, false
}
probeCtx, cancel := context.WithTimeout(ctx, exposureProbeTotalTimeout)
defer cancel()
client, closeIdle := exposureProbeHTTPClient(domain, host)
defer closeIdle()
return doLocalProbeWithClient(probeCtx, client, scheme, domain, urlPath)
}
func doLocalProbeWithClient(ctx context.Context, client *http.Client, scheme, domain, urlPath string) (probeResult, bool) {
u := url.URL{Scheme: scheme, Host: domain, Path: urlPath}
resp, err := doProbeRequest(ctx, client, http.MethodHead, u.String(), "")
if err != nil {
return probeResult{}, false
}
if resp.StatusCode == http.StatusMethodNotAllowed || resp.StatusCode == http.StatusNotImplemented {
_ = resp.Body.Close()
resp, err = doProbeRequest(ctx, client, http.MethodGet, u.String(), "bytes=0-0")
if err != nil {
return probeResult{}, false
}
}
pr := probeResult{scheme: scheme, status: resp.StatusCode, contentType: resp.Header.Get("Content-Type"), reachable: true}
_ = resp.Body.Close()
return pr, true
}
// exposureProbeHTTPClient builds the pinned-dial client shared by the
// headers-only probe and the phpinfo body confirmation: connections dial the
// vhost's configured serving address (never DNS), SNI/Host carry the domain,
// and redirects are never followed. The returned func releases idle
// connections; a fresh transport per probe is needed because SNI
// (ServerName) is per-domain.
func exposureProbeHTTPClient(domain, host string) (*http.Client, func()) {
transport := &http.Transport{
MaxResponseHeaderBytes: 64 << 10,
DialContext: func(ctx context.Context, network, addr string) (net.Conn, error) {
_, port, err := net.SplitHostPort(addr)
if err != nil {
return nil, err
}
d := &net.Dialer{Timeout: exposureProbeConnTimeout}
return d.DialContext(ctx, network, net.JoinHostPort(host, port))
},
// #nosec G402 -- reachability probe to this host's own serving IP (never
// an external endpoint); no trust decision rides on the certificate, and
// SNI must match the requested domain to select the right vhost.
TLSClientConfig: &tls.Config{InsecureSkipVerify: true, ServerName: domain},
}
client := &http.Client{
Transport: transport,
Timeout: exposureProbeTotalTimeout,
CheckRedirect: func(*http.Request, []*http.Request) error {
return http.ErrUseLastResponse
},
}
return client, transport.CloseIdleConnections
}
// phpinfoBodyReadMax bounds how much of the phpinfo.php response body the
// confirmation stage reads. Real dumps put "PHP Version" in the first couple
// of KB; 64 KiB leaves generous room for custom headers while keeping the
// probe cheap.
const phpinfoBodyReadMax = 64 << 10
// phpinfoMinBodyBytes is the minimum confirmed-dump size. Even a minimal PHP
// build renders tens of KB of phpinfo output; the observed false-positive
// stubs answer with empty or sub-100-byte bodies.
const phpinfoMinBodyBytes = 4096
// phpinfoBodyFetcher retrieves up to phpinfoBodyReadMax bytes of the response
// body for a phpinfo candidate. The bool reports whether the result can clear
// an earlier finding: non-success HTTP statuses are complete negatives, while
// transport, body-read, and inconclusive partial-response failures are not.
// This separate seam keeps every other exposure class on its headers-only
// contract.
type phpinfoBodyFetcher func(ctx context.Context, scheme, domain, host, urlPath string) ([]byte, bool)
var fetchPHPInfoBody phpinfoBodyFetcher = realFetchPHPInfoBody
func realFetchPHPInfoBody(ctx context.Context, scheme, domain, host, urlPath string) ([]byte, bool) {
host = normalizeServingIP(host)
if host == "" {
return nil, false
}
fetchCtx, cancel := context.WithTimeout(ctx, exposureProbeTotalTimeout)
defer cancel()
client, closeIdle := exposureProbeHTTPClient(domain, host)
defer closeIdle()
u := url.URL{Scheme: scheme, Host: domain, Path: urlPath}
return fetchPHPInfoBodyWithClient(fetchCtx, client, u.String())
}
func fetchPHPInfoBodyWithClient(ctx context.Context, client *http.Client, u string) ([]byte, bool) {
resp, err := doProbeRequest(ctx, client, http.MethodGet, u, "")
if err != nil {
return nil, false
}
defer func() { _ = resp.Body.Close() }()
partial := resp.StatusCode == http.StatusPartialContent
if resp.StatusCode != http.StatusOK && !partial {
return nil, true
}
reader := io.Reader(resp.Body)
switch strings.ToLower(strings.TrimSpace(resp.Header.Get("Content-Encoding"))) {
case "":
case "gzip":
compressed, gzErr := gzip.NewReader(resp.Body)
if gzErr != nil {
return nil, false
}
defer func() { _ = compressed.Close() }()
reader = compressed
default:
return nil, false
}
body, err := io.ReadAll(io.LimitReader(reader, phpinfoBodyReadMax))
if err != nil {
return nil, false
}
if partial && !isRealPHPInfoBody(body) {
return nil, false
}
return body, true
}
// confirmPHPInfoBody checks both origin protocols because HEAD and GET can be
// routed differently, and HTTPS and HTTP can execute different handlers. A
// dump observed on either protocol is conclusive; if neither exposes a dump,
// every GET must complete before an earlier finding can be cleared.
func confirmPHPInfoBody(ctx context.Context, preferredScheme, domain, host, urlPath string) (string, bool, bool) {
schemes := []string{"https", "http"}
if preferredScheme == "http" {
schemes[0], schemes[1] = schemes[1], schemes[0]
}
complete := true
for _, scheme := range schemes {
body, ok := fetchPHPInfoBody(ctx, scheme, domain, host, urlPath)
if !ok {
complete = false
continue
}
if isRealPHPInfoBody(body) {
return scheme, true, true
}
}
return "", false, complete
}
// isRealPHPInfoBody reports whether a response body is genuine phpinfo()
// output: the version banner every phpinfo variant emits (HTML "PHP Version
// x.y" heading or CLI "PHP Version => x.y") plus a dump-sized body. UTF-16
// variants cover output handlers that transcode the otherwise ASCII banner.
func isRealPHPInfoBody(body []byte) bool {
if len(body) < phpinfoMinBodyBytes {
return false
}
marker := []byte("PHP Version")
if bytes.Contains(body, marker) {
return true
}
le := make([]byte, len(marker)*2)
be := make([]byte, len(marker)*2)
for i, b := range marker {
le[i*2] = b
be[i*2+1] = b
}
return bytes.Contains(body, le) || bytes.Contains(body, be)
}
func doProbeRequest(ctx context.Context, client *http.Client, method, u, rangeHdr string) (*http.Response, error) {
req, err := http.NewRequestWithContext(ctx, method, u, nil)
if err != nil {
return nil, err
}
if rangeHdr != "" {
req.Header.Set("Range", rangeHdr)
}
req.Header.Set("User-Agent", "csm-exposure-probe")
return client.Do(req)
}
// exposedClass labels a docroot file by the kind of exposure it represents
// when the web server serves it as a raw download. The zero value classNone
// means "not a sensitive file" and is never emitted.
type exposedClass int
const (
classNone exposedClass = iota
// classRepoMetadata is a web-served version-control directory (.git,
// .svn): the site's source, its history and, routinely, credentials.
classRepoMetadata
classConfigLeak
classDBDump
classBackupArchive
classSourceBackup
classPHPInfo
// classSampleSQL is a plain SQL file with a sample-specific name under a
// framework/vendor/downloaded-project directory. It is still served, but is
// a Warning rather than a Critical database dump.
classSampleSQL
)
// severity maps a class to its finding severity. Credential- and
// database-bearing exposures are Critical; a leaked source backup is High;
// a phpinfo dump is Warning (information disclosure only).
func (c exposedClass) severity() alert.Severity {
switch c {
case classConfigLeak, classDBDump, classBackupArchive, classRepoMetadata:
return alert.Critical
case classSourceBackup:
return alert.High
default:
return alert.Warning
}
}
// findingName is the stable registry key / audit-log check name per class.
func (c exposedClass) findingName() string {
switch c {
case classRepoMetadata:
return "web_exposed_repo_metadata"
case classConfigLeak:
return "web_exposed_config_leak"
case classDBDump:
return "web_exposed_db_dump"
case classBackupArchive:
return "web_exposed_backup_archive"
case classSourceBackup:
return "web_exposed_source_backup"
case classPHPInfo:
return "web_exposed_phpinfo"
case classSampleSQL:
return "web_exposed_sample_sql"
default:
return ""
}
}
// vhost is one cPanel virtual host: the domain the web server answers for and
// the docroot it serves. Parsed from /etc/userdatadomains, which -- unlike the
// /home/*/public_html glob -- covers addon and subdomain docroots too.
type vhost struct {
domain string
user string
typ string
mainDomain string
docroot string
// ip is the vhost's serving address (from the ip:443/ip:80 columns). The
// reachability probe dials this rather than 127.0.0.1: LiteSpeed returns
// 403 to loopback-originated requests even for files it serves (HTTP 200)
// to a real request on the public IP, so a loopback probe confirms nothing.
ip string
// phpVersion is the vhost's MultiPHP selection ("ea-php83", "alt-php81")
// from the PHP-version column, empty when the row omits it. Callers that run
// code against a docroot need this: the system default interpreter is not
// necessarily the one the site is pinned to.
phpVersion string
}
// parseUserdataDomains parses the /etc/userdatadomains map. Each line is
// "<domain>: <user>==<reseller>==<type>==<maindomain>==<docroot>==...".
// Wildcard, comment, and short/malformed lines are skipped.
func parseUserdataDomains(content string) []vhost {
vhosts, _ := parseUserdataDomainsChecked(content)
return vhosts
}
func parseUserdataDomainsChecked(content string) ([]vhost, bool) {
return parseUserdataDomainsForUse(content, true)
}
// parseUserdataDomainRootsChecked validates the fields needed by filesystem
// scanners. Unlike exposure probing, a local docroot walk does not need a
// serving IP, so a row without one is still complete for this use.
func parseUserdataDomainRootsChecked(content string) ([]vhost, bool) {
return parseUserdataDomainsForUse(content, false)
}
func parseUserdataDomainsForUse(content string, requireServingIP bool) ([]vhost, bool) {
var out []vhost
complete := true
for _, line := range strings.Split(content, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, "*:") || strings.HasPrefix(line, "*.") {
continue
}
colon := strings.Index(line, ": ")
if colon <= 0 {
complete = false
continue
}
domain := strings.ToLower(strings.TrimSpace(line[:colon]))
fields := strings.Split(line[colon+2:], "==")
if !validProbeDomain(domain) || len(fields) < 5 {
complete = false
continue
}
docroot := filepath.Clean(strings.TrimSpace(fields[4]))
user := strings.TrimSpace(fields[0])
if !filepath.IsAbs(docroot) || docroot == string(filepath.Separator) ||
!validAccountName.MatchString(user) {
complete = false
continue
}
servingIP := parseServingIP(fields)
if requireServingIP && servingIP == "" {
// Without a literal serving address the vhost cannot be probed
// reliably. Keep the row so callers can preserve partial-scan state,
// but do not treat this map as a complete scan input.
complete = false
}
out = append(out, vhost{
domain: domain,
user: user,
typ: strings.TrimSpace(fields[2]),
mainDomain: strings.ToLower(strings.TrimSpace(fields[3])),
docroot: docroot,
ip: servingIP,
phpVersion: parseVhostPHPVersion(fields),
})
}
return out, complete
}
// parseServingIP extracts the vhost's serving address from the ip:443 column
// (preferred) or the ip:80 column. Only literal unicast IP addresses are
// accepted so a malformed map can never turn the pinned origin probe into a
// DNS lookup or a connection to a non-serving address.
func parseServingIP(fields []string) string {
bindings := []struct {
field int
port string
}{
{field: 6, port: "443"},
{field: 5, port: "80"},
}
for _, binding := range bindings {
if binding.field >= len(fields) {
continue
}
hp := strings.TrimSpace(fields[binding.field])
if hp == "" {
continue
}
if host, port, err := net.SplitHostPort(hp); err == nil {
if port == binding.port {
if ip := normalizeServingIP(host); ip != "" {
return ip
}
}
continue
}
// cPanel may render IPv6 bindings without brackets. SplitHostPort
// rejects those, so split at the last colon and validate both parts.
if idx := strings.LastIndexByte(hp, ':'); idx > 0 {
if strings.TrimSpace(hp[idx+1:]) == binding.port {
if ip := normalizeServingIP(hp[:idx]); ip != "" {
return ip
}
}
}
}
return ""
}
// normalizeServingIP returns a canonical literal unicast address. In
// particular, it rejects hostnames, unspecified/listener addresses, loopback,
// multicast, and link-local addresses, none of which identify the vhost's
// externally serving origin.
func normalizeServingIP(host string) string {
ip := net.ParseIP(strings.TrimSpace(host))
if ip == nil || !ip.IsGlobalUnicast() {
return ""
}
return ip.String()
}
// probeHost is the literal serving address the reachability probe dials. An
// empty result means the vhost must be skipped and the scan marked incomplete;
// loopback is deliberately not a fallback because it produces false 403s on
// LiteSpeed.
func probeHost(vh vhost) string {
return normalizeServingIP(vh.ip)
}
// validProbeDomain rejects URL authority syntax. Without this guard a corrupt
// map entry containing a port could make the privileged daemon probe an
// arbitrary service on the serving IP instead of the vhost's port 80/443.
func validProbeDomain(domain string) bool {
if domain == "" || len(domain) > 253 || strings.ContainsAny(domain, ":/\\?#@[]%") {
return false
}
for _, r := range domain {
if r <= ' ' || r == 0x7f {
return false
}
}
return true
}
// probeResult is the local reachability check outcome for one candidate file.
// It carries the requested scheme and only headers-level response facts
// (status, Content-Type); the body is never read, so a leaked file's secret
// contents never enter CSM.
type probeResult struct {
scheme string
status int
contentType string
reachable bool
// partial is true when one protocol could not be reached. Unless the other
// protocol confirms an exposure, that uncertainty must prevent purging a
// finding observed during an earlier complete scan.
partial bool
}
// confirmExposure reports whether a classified candidate is a confirmed,
// downloadable exposure. It fails closed: anything the server blocks
// (403/404), redirects (3xx), or that comes back as executed HTML on a
// non-executing class is not a finding.
func confirmExposure(class exposedClass, pr probeResult) bool {
if !pr.reachable {
return false
}
if pr.status != 200 && pr.status != 206 {
return false
}
if class == classPHPInfo {
return true
}
// Non-executing sensitive files: a raw (non-HTML) body means the server
// served the file itself. An HTML body means it executed or returned an
// error/challenge page -- not a confirmed source leak.
return !isHTMLContentType(pr.contentType)
}
func isHTMLContentType(ct string) bool {
ct = strings.ToLower(strings.TrimSpace(ct))
return strings.HasPrefix(ct, "text/html") || strings.HasPrefix(ct, "application/xhtml")
}
// phpExecExts are the extensions a cPanel PHP handler executes rather than
// serving as text. A backup whose FINAL extension is one of these is run by
// the interpreter, emits no source, and is therefore not a leak.
var phpExecExts = map[string]bool{
"php": true, "php3": true, "php4": true, "php5": true, "php7": true,
"php8": true, "phtml": true, "pht": true, "phar": true,
}
// backupSuffixSet are terminal dot-segments that mark a renamed backup a web
// server serves as text (not executed). "bak-<timestamp>" style segments are
// matched by prefix in isBackupSuffix.
var backupSuffixSet = map[string]bool{
"old": true, "bak": true, "save": true, "orig": true, "broken": true,
"swp": true, "swo": true, "tmp": true, "copy": true, "backup": true,
}
// dbDumpSuffixes are raw database-dump endings. Each may optionally be wrapped
// in one of dbDumpArchiveSuffixes.
var dbDumpSuffixes = []string{
".sql", ".dump", ".mysql",
}
var dbDumpArchiveSuffixes = []string{
"", ".gz", ".zip", ".bz2", ".xz", ".zst", ".7z", ".rar",
".tar", ".tar.gz", ".tar.bz2", ".tar.xz", ".tar.zst", ".tgz", ".tbz",
}
// archiveExts are archive endings that, combined with a backup token, denote a
// full-site backup a visitor could download whole.
var archiveExts = []string{
".tar.gz", ".tgz", ".tar.bz2", ".tbz", ".tar.xz", ".tar.zst", ".tar",
".zip", ".rar", ".7z", ".gz", ".bz2", ".xz", ".zst",
}
// backupTokens are name fragments that mark an archive as a site/db backup
// rather than a legitimately-offered download. Kept strong to avoid flagging
// ordinary user zips.
var backupTokens = []string{
"public_html", "wp-content",
}
// These generic words are backup signals only as complete filename tokens;
// substring matching would turn names such as dumpster.zip or immigration.zip
// into Critical false positives.
var delimitedBackupTokens = []string{
"backup", "dump", "full",
}
// sampleSQLDirMarkers are directory-name segments that provide supporting
// context for a sample-specific SQL file name. Directory context alone is not
// enough to demote a dump because these paths are customer-controlled.
var sampleSQLDirMarkers = map[string]bool{
"example": true, "examples": true,
"sample": true, "samples": true,
"demo": true, "demos": true,
"doc": true, "docs": true,
"fixture": true, "fixtures": true, "testdata": true,
"vendor": true, "node_modules": true, "bower_components": true,
}
var sampleSQLNameTokens = map[string]bool{
"demo": true, "example": true, "fixture": true, "install": true,
"migration": true, "sample": true, "schema": true, "seed": true,
"setup": true, "structure": true, "update": true, "upgrade": true,
}
// isSampleSQLPath reports whether rel (a leading-slash URL path) names a plain
// SQL file with a sample/schema-specific name under framework scaffolding.
// Both the file name and directory context must support that classification.
// Archived, renamed, and customer-named dumps remain Critical even under a
// generic docs/vendor path.
func isSampleSQLPath(rel string) bool {
segs := strings.Split(rel, "/")
if len(segs) < 2 {
return false
}
name := strings.ToLower(segs[len(segs)-1])
if !strings.HasSuffix(name, ".sql") {
return false
}
stem := strings.TrimSuffix(name, ".sql")
sampleName := false
for _, token := range strings.FieldsFunc(stem, func(r rune) bool {
return r == '-' || r == '_' || r == '.' || r == ' '
}) {
if sampleSQLNameTokens[token] {
sampleName = true
break
}
}
archiveProject := false
for _, seg := range segs[:len(segs)-1] { // dir segments only, skip file name
s := strings.ToLower(seg)
if s == "" {
continue
}
if sampleName && sampleSQLDirMarkers[s] {
return true
}
// Directories unpacked from a GitHub archive keep a "-master"/"-main"
// suffix, a strong signal the tree is a downloaded sample project.
if (strings.HasSuffix(s, "-master") && len(s) > len("-master")) ||
(strings.HasSuffix(s, "-main") && len(s) > len("-main")) {
archiveProject = true
}
}
return archiveProject && (sampleName || stem == "database")
}
// demoteSampleSQL lowers a database-dump candidate to classSampleSQL when its
// path has specific sample-file and framework-scaffolding signals. Only
// classDBDump is eligible; credential, archive, and source classes are never
// demoted.
func demoteSampleSQL(class exposedClass, rel string) exposedClass {
if class == classDBDump && isSampleSQLPath(rel) {
return classSampleSQL
}
return class
}
// classifyExposedFile classifies a base file name. It is deliberately
// conservative: the benign long tail (samples, examples, live scripts, and
// backups that still execute) returns classNone so the detector does not
// drown operators in false positives.
func classifyExposedFile(name string) exposedClass {
lower := strings.ToLower(strings.TrimSpace(name))
if lower == "" {
return classNone
}
// Diagnostic by its unambiguous conventional name. A generic info.php can
// be any application endpoint, and headers alone cannot prove it calls
// phpinfo(), so classifying that name would create false positives.
if lower == "phpinfo.php" {
return classPHPInfo
}
// Benign long tail excluded before any leak matching.
if isBenignExposedName(lower) {
return classNone
}
// Live scripts are executed by the PHP handler rather than served as
// source. Keep the phpinfo diagnostic exceptions above, but reject every
// other candidate whose final extension still executes -- including names
// such as .env.php that would otherwise match the dotenv family.
if hasPHPExecExtension(lower) {
return classNone
}
// Dotenv family (secrets), including editor backups such as .env~.
if isDotenvName(lower) {
return classConfigLeak
}
if hasDBDumpSuffix(lower) {
return classDBDump
}
if isBackupArchive(lower) {
return classBackupArchive
}
// A sensitive file can itself have been renamed with one or more backup
// suffixes (for example wp-config.php.bak.old).
if stripped, ok := stripBackupSuffixes(lower); ok {
switch {
case isDotenvName(stripped):
return classConfigLeak
case hasDBDumpSuffix(stripped):
return classDBDump
case isBackupArchive(stripped):
return classBackupArchive
case looksLikePHPSource(stripped):
if looksLikeConfig(stripped) {
return classConfigLeak
}
return classSourceBackup
}
}
return classNone
}
func isDotenvName(lower string) bool {
return lower == ".env" || strings.HasPrefix(lower, ".env.")
}
// isBenignExposedName matches shipped samples and templates that carry no
// secret and must never be flagged.
func isBenignExposedName(lower string) bool {
if lower == "wp-config-sample.php" {
return true
}
if stripped, ok := stripBackupSuffixes(lower); ok && stripped == "wp-config-sample.php" {
return true
}
for _, template := range []string{".env.example", ".env.sample", ".env.dist", ".env.default"} {
if lower == template || strings.HasPrefix(lower, template+".") || strings.HasPrefix(lower, template+"~") {
return true
}
}
return strings.HasSuffix(lower, ".dist") || strings.HasSuffix(lower, ".default")
}
func hasDBDumpSuffix(lower string) bool {
for _, base := range dbDumpSuffixes {
for _, archive := range dbDumpArchiveSuffixes {
if strings.HasSuffix(lower, base+archive) {
return true
}
}
}
return false
}
func isBackupArchive(lower string) bool {
if strings.HasSuffix(lower, ".wpress") {
return true
}
hasArchiveExt := false
for _, e := range archiveExts {
if strings.HasSuffix(lower, e) {
hasArchiveExt = true
break
}
}
if !hasArchiveExt {
return false
}
for _, tok := range backupTokens {
if strings.Contains(lower, tok) {
return true
}
}
for _, tok := range delimitedBackupTokens {
if hasDelimitedFilenameToken(lower, tok) {
return true
}
}
return false
}
func hasDelimitedFilenameToken(name, token string) bool {
for start := 0; start <= len(name)-len(token); {
i := strings.Index(name[start:], token)
if i < 0 {
return false
}
i += start
beforeOK := i == 0 || !isASCIIAlphaNumeric(name[i-1])
after := i + len(token)
afterOK := after == len(name) || !isASCIIAlphaNumeric(name[after])
if beforeOK && afterOK {
return true
}
start = i + 1
}
return false
}
func isASCIIAlphaNumeric(b byte) bool {
return b >= 'a' && b <= 'z' || b >= '0' && b <= '9'
}
// stripBackupSuffix removes a single trailing backup marker (a "~" or a
// terminal dot-segment such as ".old" / ".bak-20260515") and reports whether
// one was found.
func stripBackupSuffix(lower string) (string, bool) {
if strings.HasSuffix(lower, "~") {
return strings.TrimSuffix(lower, "~"), true
}
idx := strings.LastIndexByte(lower, '.')
if idx <= 0 {
return lower, false
}
seg := lower[idx+1:]
if isBackupSuffix(seg) {
return lower[:idx], true
}
return lower, false
}
func stripBackupSuffixes(lower string) (string, bool) {
stripped := lower
found := false
for {
next, ok := stripBackupSuffix(stripped)
if !ok {
return stripped, found
}
stripped = next
found = true
}
}
func isBackupSuffix(seg string) bool {
if backupSuffixSet[seg] || isASCIIDigits(seg) {
return true
}
// "bak-20260515-124446", "bak_1", "save1" style timestamped variants.
for _, p := range []string{"bak-", "bak_", "old-", "old_", "save-", "save_", "backup-", "backup_", "orig-", "orig_"} {
if strings.HasPrefix(seg, p) {
return true
}
}
for _, p := range []string{"bak", "old", "save", "backup", "orig"} {
if strings.HasPrefix(seg, p) && isASCIIDigits(strings.TrimPrefix(seg, p)) {
return true
}
}
return false
}
func isASCIIDigits(s string) bool {
if s == "" {
return false
}
for i := 0; i < len(s); i++ {
if s[i] < '0' || s[i] > '9' {
return false
}
}
return true
}
// looksLikePHPSource reports whether the name (with its backup suffix already
// removed) is a PHP source file the server would serve as text.
func looksLikePHPSource(stripped string) bool {
return hasPHPExecExtension(stripped)
}
func hasPHPExecExtension(name string) bool {
idx := strings.LastIndexByte(name, '.')
if idx < 0 {
return false
}
ext := name[idx+1:]
return phpExecExts[ext]
}
// looksLikeConfig reports whether a source backup carries configuration /
// credentials, warranting Critical instead of High.
func looksLikeConfig(stripped string) bool {
for _, tok := range []string{"config", "wp-config", "settings", "database", "db-config", "credentials", "secret"} {
if strings.Contains(stripped, tok) {
return true
}
}
return false
}
package checks
import (
"archive/zip"
"bufio"
"context"
"encoding/binary"
"errors"
"io"
"io/fs"
"os"
"path"
"strings"
"syscall"
)
const (
// The central directory is the only part of a zip this check reads. A byte
// bound covers entry count, names, comments, and extra fields together while
// allowing substantially more than the old 4096-entry cutoff.
archiveDirectoryScanByteLimit = 16 << 20
zipDirectoryHeaderLen = 46
zipDirectoryEndLen = 22
zipDirectory64EndLen = 56
zipDirectory64LocLen = 20
zipMaxCommentLen = 1<<16 - 1
// A site backup names its document root at or just below the archive
// root. Past that a configuration file needs corroboration.
shallowSiteConfigDepth = 2
zipDirectoryHeaderSignature = 0x02014b50
zipDirectoryEndSignature = 0x06054b50
zipDirectory64EndSignature = 0x06064b50
zipDirectory64LocSignature = 0x07064b50
)
var (
errArchiveFormat = errors.New("invalid zip central directory")
errArchiveScanLimit = errors.New("zip central directory exceeds scan limit")
)
// docrootDirNames are hosting conventions for the web root. An archive whose
// entries sit under one is a copy of a served directory tree.
var docrootDirNames = map[string]bool{
"wwwroot": true,
"public_html": true,
"htdocs": true,
"httpdocs": true,
}
type zipDirectory struct {
offset int64
size int64
records uint64
}
// archiveHoldsSiteBackup reports whether a web-reachable archive contains a
// copy of a site, judged by its entry list rather than its file name.
func archiveHoldsSiteBackup(p string) bool {
holds, _ := archiveSiteBackupStatus(context.Background(), p)
return holds
}
// archiveSiteBackupStatus also reports whether the result is complete. A
// resource-limited or interrupted inspection must retain an earlier finding.
// Malformed zip data is a complete negative result: a file merely ending in
// .zip must not hold the whole exposed-files check incomplete forever.
func archiveSiteBackupStatus(ctx context.Context, p string) (holds, complete bool) {
if !strings.HasSuffix(strings.ToLower(p), ".zip") {
return false, true
}
info, err := osFS.Lstat(p)
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return false, true
}
return false, false
}
if !info.Mode().IsRegular() && info.Mode()&os.ModeSymlink == 0 {
return false, true
}
var f *os.File
if _, productionFS := osFS.(realOS); productionFS {
// The account controls this path. A nonblocking open plus the descriptor
// type check allows web-served regular symlinks without letting a FIFO
// swapped in after the directory walk strand the scan.
// #nosec G304 -- read-only document-root candidate; flags reject unsafe types.
f, err = os.OpenFile(p, os.O_RDONLY|syscall.O_NONBLOCK, 0)
} else {
f, err = osFS.Open(p)
}
if err != nil {
if errors.Is(err, fs.ErrNotExist) || errors.Is(err, syscall.ELOOP) {
return false, true
}
return false, false
}
defer func() { _ = f.Close() }()
openedInfo, err := f.Stat()
if err != nil {
return false, false
}
if !openedInfo.Mode().IsRegular() {
return false, true
}
dir, err := readZipDirectory(f, openedInfo.Size())
if err != nil {
if errors.Is(err, errArchiveFormat) {
return false, true
}
return false, false
}
holds, err = scanZipDirectory(ctx, f, dir)
if err != nil {
if errors.Is(err, errArchiveFormat) {
return false, true
}
return holds, false
}
return holds, true
}
func readZipDirectory(f *os.File, size int64) (zipDirectory, error) {
if size < zipDirectoryEndLen {
return zipDirectory{}, errArchiveFormat
}
tailLen := int64(zipDirectoryEndLen + zipMaxCommentLen)
if tailLen > size {
tailLen = size
}
tail := make([]byte, int(tailLen))
if err := readAtFull(f, tail, size-tailLen); err != nil {
return zipDirectory{}, err
}
endIndex := findZipDirectoryEnd(tail)
if endIndex < 0 {
return zipDirectory{}, errArchiveFormat
}
endOffset := size - tailLen + int64(endIndex)
end := tail[endIndex : endIndex+zipDirectoryEndLen]
if binary.LittleEndian.Uint16(end[4:6]) != 0 ||
binary.LittleEndian.Uint16(end[6:8]) != 0 {
return zipDirectory{}, errArchiveFormat
}
recordsThisDisk := uint64(binary.LittleEndian.Uint16(end[8:10]))
records := uint64(binary.LittleEndian.Uint16(end[10:12]))
directorySize := uint64(binary.LittleEndian.Uint32(end[12:16]))
directoryOffset := uint64(binary.LittleEndian.Uint32(end[16:20]))
if recordsThisDisk != records {
return zipDirectory{}, errArchiveFormat
}
directoryEndOffset := endOffset
if records == 0xffff || directorySize == 0xffffffff || directoryOffset == 0xffffffff {
var err error
directoryEndOffset, records, directorySize, directoryOffset, err = readZip64Directory(f, endOffset)
if err != nil {
return zipDirectory{}, err
}
}
if directorySize > archiveDirectoryScanByteLimit {
return zipDirectory{}, errArchiveScanLimit
}
directorySize64, sizeOK := archiveOffset(directorySize)
directoryOffset64, offsetOK := archiveOffset(directoryOffset)
if !sizeOK || !offsetOK || directoryEndOffset < 0 ||
directorySize64 > directoryEndOffset || directoryOffset64 > directoryEndOffset {
return zipDirectory{}, errArchiveFormat
}
if records > directorySize/zipDirectoryHeaderLen {
return zipDirectory{}, errArchiveFormat
}
return zipDirectory{
offset: directoryEndOffset - directorySize64,
size: directorySize64,
records: records,
}, nil
}
func findZipDirectoryEnd(tail []byte) int {
for i := len(tail) - zipDirectoryEndLen; i >= 0; i-- {
if binary.LittleEndian.Uint32(tail[i:i+4]) != zipDirectoryEndSignature {
continue
}
commentLen := int(binary.LittleEndian.Uint16(tail[i+20 : i+22]))
if i+zipDirectoryEndLen+commentLen == len(tail) {
return i
}
}
return -1
}
func readZip64Directory(f *os.File, endOffset int64) (end int64, records, size, offset uint64, err error) {
locatorOffset := endOffset - zipDirectory64LocLen
if locatorOffset < 0 {
return 0, 0, 0, 0, errArchiveFormat
}
var locator [zipDirectory64LocLen]byte
if err := readAtFull(f, locator[:], locatorOffset); err != nil {
return 0, 0, 0, 0, err
}
if binary.LittleEndian.Uint32(locator[0:4]) != zipDirectory64LocSignature ||
binary.LittleEndian.Uint32(locator[4:8]) != 0 ||
binary.LittleEndian.Uint32(locator[16:20]) != 1 {
return 0, 0, 0, 0, errArchiveFormat
}
zip64Offset := binary.LittleEndian.Uint64(locator[8:16])
zip64Offset64, ok := archiveOffset(zip64Offset)
if !ok || zip64Offset64 > locatorOffset {
return 0, 0, 0, 0, errArchiveFormat
}
var record [zipDirectory64EndLen]byte
if err := readAtFull(f, record[:], zip64Offset64); err != nil {
return 0, 0, 0, 0, err
}
if binary.LittleEndian.Uint32(record[0:4]) != zipDirectory64EndSignature ||
binary.LittleEndian.Uint64(record[4:12]) < zipDirectory64EndLen-12 ||
binary.LittleEndian.Uint32(record[16:20]) != 0 ||
binary.LittleEndian.Uint32(record[20:24]) != 0 {
return 0, 0, 0, 0, errArchiveFormat
}
recordsThisDisk := binary.LittleEndian.Uint64(record[24:32])
records = binary.LittleEndian.Uint64(record[32:40])
if recordsThisDisk != records {
return 0, 0, 0, 0, errArchiveFormat
}
return zip64Offset64, records,
binary.LittleEndian.Uint64(record[40:48]),
binary.LittleEndian.Uint64(record[48:56]), nil
}
func archiveOffset(value uint64) (int64, bool) {
if value > 1<<63-1 {
return 0, false
}
return int64(value), true // #nosec G115 -- the MaxInt64 bound is checked above.
}
func readAtFull(f *os.File, dst []byte, offset int64) error {
n, err := f.ReadAt(dst, offset)
if n != len(dst) {
if errors.Is(err, io.EOF) {
return errArchiveFormat
}
if err != nil {
return err
}
return errArchiveFormat
}
if err != nil && !errors.Is(err, io.EOF) {
return err
}
return nil
}
func scanZipDirectory(ctx context.Context, f *os.File, dir zipDirectory) (bool, error) {
reader := bufio.NewReader(io.NewSectionReader(f, dir.offset, dir.size))
var header [zipDirectoryHeaderLen]byte
var records uint64
holdsSite := false
// A wp-config.php nested past the shallow bound only counts alongside a
// WordPress runtime file, so both are accumulated across the whole
// directory rather than decided per entry.
deepConfig := false
wpRuntime := false
// A nested Joomla configuration.php counts only when its own directory
// also holds a Joomla entry point, so both are keyed by directory.
joomlaConfigDirs := map[string]bool{}
joomlaRuntimeDirs := map[string]bool{}
for {
if records%256 == 0 {
if err := ctx.Err(); err != nil {
return holdsSite, err
}
}
signature, err := reader.Peek(4)
if errors.Is(err, io.EOF) {
break
}
if err != nil {
return holdsSite, err
}
if binary.LittleEndian.Uint32(signature) != zipDirectoryHeaderSignature {
break
}
if _, err := io.ReadFull(reader, header[:]); err != nil {
return holdsSite, archiveDirectoryReadError(err)
}
nameLen := int(binary.LittleEndian.Uint16(header[28:30]))
extraLen := int64(binary.LittleEndian.Uint16(header[30:32]))
commentLen := int64(binary.LittleEndian.Uint16(header[32:34]))
name := make([]byte, nameLen)
if _, err := io.ReadFull(reader, name); err != nil {
return holdsSite, archiveDirectoryReadError(err)
}
if _, err := io.CopyN(io.Discard, reader, extraLen+commentLen); err != nil {
return holdsSite, archiveDirectoryReadError(err)
}
if entry := string(name); archiveEntrySignalsSiteBackup(entry) {
holdsSite = true
} else {
if archiveEntryIsDeepWPConfig(entry) {
deepConfig = true
}
if archiveEntryIsWPRuntime(entry) {
wpRuntime = true
}
// Names alone cannot distinguish PHP files from directories or
// symlinks. Interpret central-directory attributes without opening
// or decompressing any entry payload.
entryHeader := zip.FileHeader{
Name: strings.ReplaceAll(entry, `\`, "/"),
CreatorVersion: binary.LittleEndian.Uint16(header[4:6]),
ExternalAttrs: binary.LittleEndian.Uint32(header[38:42]),
}
if entryHeader.Mode().IsRegular() {
if d, ok := archiveNestedJoomlaConfigDir(entry); ok {
joomlaConfigDirs[d] = true
holdsSite = holdsSite || joomlaRuntimeDirs[d]
}
if d, ok := archiveJoomlaRuntimeDir(entry); ok {
joomlaRuntimeDirs[d] = true
holdsSite = holdsSite || joomlaConfigDirs[d]
}
}
}
records++
}
if records != dir.records {
return false, errArchiveFormat
}
return holdsSite || (deepConfig && wpRuntime), nil
}
func archiveDirectoryReadError(err error) error {
if errors.Is(err, io.EOF) || errors.Is(err, io.ErrUnexpectedEOF) {
return errArchiveFormat
}
return err
}
// archiveEntryPath normalises a raw zip entry name and rejects the shapes that
// must never be read as a marker: absolute, drive-qualified, and traversal
// names. It returns nil for anything unusable.
func archiveEntryPath(rawName string) []string {
return archiveEntryPathPreservingCase(strings.ToLower(strings.TrimSpace(rawName)))
}
// Directory identity must retain case and whitespace when pairing entries:
// archives from Linux can contain distinct directories differing only by either.
func archiveEntryPathPreservingCase(rawName string) []string {
name := strings.ReplaceAll(rawName, `\`, "/")
if name == "" || strings.HasPrefix(name, "/") ||
(len(name) >= 3 && ((name[0] >= 'a' && name[0] <= 'z') || (name[0] >= 'A' && name[0] <= 'Z')) && name[1] == ':' && name[2] == '/') {
return nil
}
name = path.Clean(name)
if name == "." || name == ".." || strings.HasPrefix(name, "../") {
return nil
}
parts := strings.Split(strings.Trim(name, "/"), "/")
if len(parts) == 0 || parts[0] == "" {
return nil
}
return parts
}
func archiveEntrySignalsSiteBackup(rawName string) bool {
parts := archiveEntryPath(rawName)
if len(parts) == 0 {
return false
}
for _, part := range parts[:len(parts)-1] {
if docrootDirNames[part] {
return true
}
}
if len(parts) == 1 && hasDBDumpSuffix(parts[0]) {
return true
}
switch parts[len(parts)-1] {
case "wp-config.php":
// A configuration file this close to the archive root is the archive's
// own subject. Deeper ones are ambiguous -- plugins ship fixtures at
// arbitrary depth -- so those are paired with a runtime marker instead,
// via archiveEntryIsDeepWPConfig.
return len(parts) <= shallowSiteConfigDepth
case "configuration.php":
return len(parts) == 1
case "settings.php":
return hasArchivePathSuffix(parts, "sites", "default", "settings.php")
default:
return false
}
}
// archiveEntryIsDeepWPConfig reports a wp-config.php nested past the depth the
// shallow rule accepts on its own.
func archiveEntryIsDeepWPConfig(rawName string) bool {
parts := archiveEntryPath(rawName)
return len(parts) > shallowSiteConfigDepth && parts[len(parts)-1] == "wp-config.php"
}
// archiveEntryIsWPRuntime reports a file only a WordPress installation carries.
// A plugin or theme bundle shipping a configuration fixture has none of these,
// which is what separates a nested backup from a nested fixture.
func archiveEntryIsWPRuntime(rawName string) bool {
parts := archiveEntryPath(rawName)
if len(parts) == 0 {
return false
}
switch parts[len(parts)-1] {
case "wp-load.php", "wp-settings.php", "wp-blog-header.php":
return true
}
for _, part := range parts[:len(parts)-1] {
if part == "wp-includes" {
return true
}
}
return false
}
// archiveNestedJoomlaConfigDir returns the directory of a configuration.php
// below the archive root. A root-level one already counts on its own.
func archiveNestedJoomlaConfigDir(rawName string) (string, bool) {
parts := archiveEntryPathPreservingCase(rawName)
if len(parts) < 2 || !strings.EqualFold(parts[len(parts)-1], "configuration.php") {
return "", false
}
return strings.Join(parts[:len(parts)-1], "/"), true
}
// archiveJoomlaRuntimeDir returns the site directory of a Joomla core entry
// point. Every Joomla release carries these and no extension package does.
func archiveJoomlaRuntimeDir(rawName string) (string, bool) {
parts := archiveEntryPathPreservingCase(rawName)
if len(parts) < 2 {
return "", false
}
parent, file := strings.ToLower(parts[len(parts)-2]), strings.ToLower(parts[len(parts)-1])
if (parent == "includes" && file == "defines.php") ||
(parent == "administrator" && file == "index.php") {
return strings.Join(parts[:len(parts)-2], "/"), true
}
return "", false
}
func hasArchivePathSuffix(parts []string, suffix ...string) bool {
if len(parts) < len(suffix) {
return false
}
start := len(parts) - len(suffix)
for i := range suffix {
if parts[start+i] != suffix[i] {
return false
}
}
return true
}
package checks
import (
"context"
"fmt"
"net/url"
"strings"
)
// Exposure verification shares the detection-time confirmation rule and only
// clears a finding after a complete probe pinned to the current local vhost.
// exposedVerifiableChecks is the web_exposed_* family, keyed the same way the
// findings are.
var exposedVerifiableChecks = []string{
"web_exposed_repo_metadata",
"web_exposed_config_leak",
"web_exposed_db_dump",
"web_exposed_backup_archive",
"web_exposed_source_backup",
"web_exposed_phpinfo",
"web_exposed_sample_sql",
}
// exposedReverifyLogicVersion is part of the daemon sweep token. Bump it when
// unattended exposure verification semantics change so existing findings are
// revisited once after upgrade.
const exposedReverifyLogicVersion = 1
// isExposedVerifiable reports whether check belongs to the web_exposed_* family,
// whose findings the automatic sweep may re-probe and dismiss.
func isExposedVerifiable(check string) bool {
for _, c := range exposedVerifiableChecks {
if c == check {
return true
}
}
return false
}
// exposedClassForCheck is the inverse of exposedClass.findingName. The class
// decides what counts as a confirmed exposure, so verification has to recover
// it or it would judge a phpinfo page by the raw-leak rule.
func exposedClassForCheck(check string) (exposedClass, bool) {
for _, c := range []exposedClass{
classRepoMetadata, classConfigLeak, classDBDump, classBackupArchive,
classSourceBackup, classPHPInfo, classSampleSQL,
} {
if c.findingName() == check {
return c, true
}
}
return classNone, false
}
// exposureURLFromMessage pulls the probed URL back out of the finding message
// ("Web-exposed <label> reachable at <url>").
func exposureURLFromMessage(message string) string {
const marker = " reachable at "
i := strings.Index(message, marker)
if i < 0 {
return ""
}
return strings.TrimSpace(message[i+len(marker):])
}
type exposureVhostIndex struct {
servingIPs map[string]string
complete bool
}
func loadExposureVhostIndex() exposureVhostIndex {
index := exposureVhostIndex{servingIPs: map[string]string{}}
content, err := osFS.ReadFile(userdataDomainsPath)
if err != nil {
return index
}
vhosts, complete := parseUserdataDomainsChecked(string(content))
index.complete = complete
for _, vh := range vhosts {
host := probeHost(vh)
if previous, exists := index.servingIPs[vh.domain]; exists {
if previous != host {
index.complete = false
}
continue
}
index.servingIPs[vh.domain] = host
}
return index
}
func (index exposureVhostIndex) servingIPForDomain(domain string) string {
domain = strings.ToLower(strings.TrimSpace(domain))
return index.servingIPs[domain]
}
// verifyExposedFile re-probes a web_exposed_* finding and resolves it only when
// a complete pinned probe says the server no longer serves the exposure.
func verifyExposedFile(in VerifyInput) VerifyResult {
vhosts := in.exposureVhosts
if vhosts == nil {
loaded := loadExposureVhostIndex()
vhosts = &loaded
}
return verifyExposedFileWithVhosts(in, vhosts)
}
func verifyExposedFileWithVhosts(in VerifyInput, vhosts *exposureVhostIndex) VerifyResult {
ctx := in.Context
if ctx == nil {
ctx = context.Background()
}
class, ok := exposedClassForCheck(in.Check)
if !ok {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("unknown exposure class for '%s'", in.Check)}
}
raw := exposureURLFromMessage(in.Message)
if raw == "" {
return VerifyResult{Checked: false, Detail: "could not extract the exposure URL from the finding"}
}
u, err := url.Parse(raw)
if err != nil || u.Host == "" || u.Path == "" {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("unusable exposure URL %q", raw)}
}
domain := u.Hostname()
if !validProbeDomain(domain) {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("unusable exposure domain %q", domain)}
}
if !vhosts.complete {
return VerifyResult{Checked: false,
Detail: "local vhost routing is incomplete; re-check cannot select a trusted serving address"}
}
// Pin to this host's own serving address. Resolving the domain through
// public DNS would ask whoever owns it now: a domain that has migrated away
// answers from its new provider, and that answer -- 200 with the new site's
// HTML, or a clean 404 -- says nothing about what this server exposes.
host := vhosts.servingIPForDomain(domain)
if host == "" {
return VerifyResult{Checked: false,
Detail: fmt.Sprintf("%s is no longer served by this host; re-check cannot reach the original vhost", domain)}
}
pr := webProber.probeComplete(ctx, domain, host, u.Path)
if !pr.reachable {
return VerifyResult{Checked: false,
Detail: fmt.Sprintf("could not reach %s to re-check; leaving the finding open", domain)}
}
if pr.partial {
return VerifyResult{Checked: false,
Detail: "only one protocol answered; leaving the finding open until a complete probe"}
}
if confirmExposure(class, pr) {
if class == classPHPInfo {
_, exposed, complete := confirmPHPInfoBody(ctx, pr.scheme, domain, host, u.Path)
switch {
case exposed:
return VerifyResult{Checked: true, Resolved: false,
Detail: "still downloadable: confirmed phpinfo output"}
case !complete:
return VerifyResult{Checked: false,
Detail: "could not complete phpinfo body confirmation; leaving the finding open"}
}
return VerifyResult{Checked: true, Resolved: true,
Detail: "no longer served as confirmed phpinfo output"}
}
return VerifyResult{Checked: true, Resolved: false,
Detail: fmt.Sprintf("still downloadable: HTTP %d %s", pr.status, pr.contentType)}
}
return VerifyResult{Checked: true, Resolved: true,
Detail: fmt.Sprintf("no longer served as an exposure (HTTP %d %s)", pr.status, pr.contentType)}
}
package checks
import (
"path/filepath"
"strings"
)
// repoMetadataMarkers lists the files whose presence proves that a directory
// named like a version-control store is one, so the walker can surface the
// exposure without descending into thousands of objects.
func repoMetadataMarkers(dirName string) []string {
switch strings.ToLower(dirName) {
case ".git":
return []string{"HEAD"}
case ".svn":
return []string{"wc.db", "entries"}
default:
return nil
}
}
// isRepoMetadataDir reports whether dir is a version-control store.
func isRepoMetadataDir(dir string) bool {
return len(repoMetadataMarkers(filepath.Base(dir))) > 0
}
// classifyExposedPath classifies a candidate by its full path: a marker file
// inside a repository directory is the repository's exposure, whatever the
// file is called; every other candidate is classified by name.
func classifyExposedPath(path string) exposedClass {
dir := filepath.Dir(path)
base := filepath.Base(path)
for _, marker := range repoMetadataMarkers(filepath.Base(dir)) {
if strings.EqualFold(base, marker) {
return classRepoMetadata
}
}
return classifyExposedFile(base)
}
package checks
import (
"bytes"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io"
"io/fs"
"os"
"path/filepath"
"strings"
"syscall"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/quarantinefs"
)
// Virtual patching for web-exposed files.
//
// A confirmed web_exposed_* finding means a visitor can download a sensitive
// file. Rather than delete the customer's file, CSM can write an .htaccess
// "Require all denied" rule that blocks HTTP access while leaving the file on
// disk (so the application, which reads it from the filesystem, keeps working).
//
// Every write records rollback metadata under the quarantine pre_clean dir so
// the existing /api/v1/quarantine-restore path can restore an existing file or
// remove one CSM created. Gating (off/manual/auto + dry_run) lives in callers.
// chownFunc is overridable in tests. It receives the already-open temporary
// file so ownership cannot be redirected through a path swap.
var chownFunc = func(file *os.File, uid, gid int) error {
return file.Chown(uid, gid)
}
// virtualPatchBeforeCommitForTest simulates a customer or deploy process
// changing .htaccess between the initial read and the atomic commit.
var virtualPatchBeforeCommitForTest func(string, string)
var syncVirtualPatchDirectory = quarantinefs.SyncDir
const maxVirtualPatchHtaccessSize = 4 << 20
const (
QuarantineRestoreReplaceIfUnchanged = "replace_if_unchanged"
QuarantineRestoreRemoveIfUnchanged = "remove_if_unchanged"
)
const (
vpBeginFile = "# BEGIN CSM exposed-file virtual-patch"
vpEndFile = "# END CSM exposed-file virtual-patch"
vpBeginDir = "# BEGIN CSM exposed-file virtual-patch: deny directory"
vpEndDir = "# END CSM exposed-file virtual-patch: deny directory"
// The parent block lives one directory up, outside anything a backup
// plugin owns, so it survives the plugin rewriting its own .htaccess.
vpBeginParent = "# BEGIN CSM exposed-file virtual-patch: deny archive extension"
vpEndParent = "# END CSM exposed-file virtual-patch: deny archive extension"
)
const virtualPatchAlreadyApplied = "already virtual-patched"
var errVirtualPatchAlreadyApplied = errors.New(virtualPatchAlreadyApplied)
// virtualPatchRevertedPrefix marks the finding raised when CSM has to write a
// deny block it already wrote once. Backup plugins own the .htaccess inside
// their own directory and rewrite it on every run, so the deny disappears and
// the archives are downloadable again until the next scan.
const virtualPatchRevertedPrefix = "VIRTUAL-PATCH re-applied after revert:"
var ErrVirtualPatchRestoreConflict = errors.New("virtual-patch restore conflicts with current .htaccess")
var errVirtualPatchRollbackIncomplete = errors.New("virtual-patch rollback incomplete")
type htaccessState struct {
content []byte
info os.FileInfo
existed bool
uid int
gid int
mode os.FileMode
}
type virtualPatchBackup struct {
itemPath string
metaPath string
}
// vpExposedChecks are the web_exposed_* finding names eligible for virtual
// patching -- every web-reachable file class the detector reports.
var vpExposedChecks = map[string]bool{
"web_exposed_repo_metadata": true,
"web_exposed_config_leak": true,
"web_exposed_db_dump": true,
"web_exposed_sample_sql": true,
"web_exposed_backup_archive": true,
"web_exposed_source_backup": true,
"web_exposed_phpinfo": true,
}
func isVirtualPatchableExposedCheck(check string) bool { return vpExposedChecks[check] }
// VirtualPatchExposedFile writes an .htaccess "Require all denied" rule that
// blocks HTTP download of a confirmed web-exposed file without modifying the
// file itself. For an archive inside a known backup-plugin directory the whole
// directory is denied. Returns Success=false with an "already" error when all
// applicable rules are already present (idempotent no-op).
func VirtualPatchExposedFile(filePath string) RemediationResult {
resolved, targetInfo, err := resolveExistingFixPath(filePath, effectiveFixRoots(fixHtaccessAllowedRoots))
if err != nil {
return RemediationResult{Error: err.Error()}
}
if !targetInfo.Mode().IsRegular() {
return RemediationResult{Error: "virtual-patch target is not a regular file"}
}
name := filepath.Base(resolved)
dir := filepath.Dir(resolved)
// A repository directory is denied whole: denying only the marker file
// would leave objects, refs and config downloadable.
dirDeny := isKnownBackupPluginDir(dir) || isRepoMetadataDir(dir)
if !dirDeny {
if nameErr := validDenyName(name); nameErr != nil {
return RemediationResult{Error: nameErr.Error()}
}
}
block := buildDenyBlock(name, dirDeny)
reverted, err := applyHtaccessDeny(dir, block)
primaryChanged := err == nil
if err != nil && !errors.Is(err, errVirtualPatchAlreadyApplied) {
return RemediationResult{Error: err.Error()}
}
target := name
if dirDeny {
target = filepath.Base(dir) + "/ (whole directory)"
}
parentChanged := false
parentNote := ""
if dirDeny && isKnownBackupPluginDir(dir) {
parentChanged, parentNote = denyArchiveExtensionInParent(dir)
}
if !primaryChanged && !parentChanged && parentNote == "" {
return RemediationResult{Error: virtualPatchAlreadyApplied}
}
verb := "Wrote"
if !primaryChanged {
verb = "Kept existing"
}
description := fmt.Sprintf("%s Require all denied for %s in %s", verb, target, filepath.Join(dir, ".htaccess"))
if parentNote != "" {
description += "; " + parentNote
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("denied HTTP access to %s", target),
Description: description,
Reverted: reverted,
}
}
// applyHtaccessDeny appends block to the .htaccess in dir unless the complete
// block is already present, keeping a restorable pre-patch copy. It reports
// whether an earlier CSM block had been removed or damaged in that file.
func applyHtaccessDeny(dir string, block []byte) (bool, error) {
htaccess := filepath.Join(dir, ".htaccess")
if _, err := sanitizeFixPath(htaccess, effectiveFixRoots(fixHtaccessAllowedRoots)); err != nil {
return false, err
}
state, err := readHtaccessState(htaccess, dir)
if err != nil {
return false, err
}
if state.existed && containsHtaccessBlock(state.content, block) {
return false, errVirtualPatchAlreadyApplied
}
reverted := priorPrePatchBackupExists(htaccess, block)
newContent := patchedHtaccessContent(state.content, state.existed, block)
if len(newContent) > maxVirtualPatchHtaccessSize {
return false, fmt.Errorf("refusing patched .htaccess larger than %d bytes", maxVirtualPatchHtaccessSize)
}
backup, created, err := backupHtaccessBeforePatch(htaccess, state, newContent)
if err != nil {
return false, err
}
keepBackup := !created
defer func() {
if !keepBackup {
backup.remove()
}
}()
tmp, tempState, err := writeVirtualPatchTemp(dir, newContent, state)
if err != nil {
return false, err
}
if virtualPatchBeforeCommitForTest != nil {
virtualPatchBeforeCommitForTest(htaccess, tmp)
}
if err := commitVirtualPatchTemp(tmp, htaccess, state, tempState); err != nil {
removeVirtualPatchTemp(tmp, tempState)
if errors.Is(err, errVirtualPatchRollbackIncomplete) {
keepBackup = true
}
return false, err
}
keepBackup = true
if err := syncVirtualPatchDirectory(dir); err != nil {
return false, fmt.Errorf("virtual-patch installed but directory sync failed; backup retained at %s: %w", backup.itemPath, err)
}
return reverted, nil
}
// backupPluginArchiveExt maps a backup-plugin directory to the archive
// extension that can safely be denied from the parent directory. Plugins
// writing .zip or .gz are absent on purpose: wp-content serves those
// legitimately, so a parent-level deny would break working downloads. The
// plugin-directory deny still covers them.
var backupPluginArchiveExt = map[string]string{
"ai1wm-backups": ".wpress",
}
// denyArchiveExtensionInParent writes the durable half of the protection into
// the parent directory, which the plugin does not own. It never fails the
// caller: the plugin-directory deny is the immediate protection and must not
// be given up because the parent could not be written.
func denyArchiveExtensionInParent(dir string) (bool, string) {
ext, ok := backupPluginArchiveExt[strings.ToLower(filepath.Base(dir))]
if !ok {
return false, ""
}
parent := filepath.Dir(dir)
reverted, err := applyHtaccessDeny(parent, buildParentExtensionDenyBlock(ext))
switch {
case err == nil:
if reverted {
return true, fmt.Sprintf("re-applied the durable %s deny under %s", ext, parent)
}
return true, fmt.Sprintf("also denied %s under %s", ext, parent)
case errors.Is(err, errVirtualPatchAlreadyApplied):
return false, ""
default:
return false, fmt.Sprintf("durable %s deny in %s could not be written: %v", ext, filepath.Join(parent, ".htaccess"), err)
}
}
func buildParentExtensionDenyBlock(ext string) []byte {
return []byte(fmt.Sprintf("%s %s\n<FilesMatch \"\\%s$\">\nRequire all denied\n</FilesMatch>\n%s %s\n",
vpBeginParent, ext, ext, vpEndParent, ext))
}
func patchedHtaccessContent(content []byte, existed bool, block []byte) []byte {
if !existed {
return append([]byte(nil), block...)
}
base := ensureTrailingNewline(append([]byte(nil), content...))
return append(base, block...)
}
// priorPrePatchBackupExists reports whether CSM previously wrote this exact
// block to this .htaccess. Reconstructing the recorded post-patch hash avoids
// treating a first patch for another file in the same directory as a revert.
func priorPrePatchBackupExists(htaccess string, block []byte) bool {
found := false
eachPrePatchBackup(func(meta QuarantineMeta, name string, read func(string) ([]byte, error)) bool {
if meta.OriginalPath != htaccess {
return false
}
archived, err := read(strings.TrimSuffix(name, ".meta"))
if err != nil {
return false
}
var existed bool
switch meta.RestoreAction {
case QuarantineRestoreReplaceIfUnchanged:
existed = true
case QuarantineRestoreRemoveIfUnchanged:
if len(archived) != 0 {
return false
}
default:
return false
}
if meta.ExpectedCurrentSHA256 == virtualPatchSHA256(patchedHtaccessContent(archived, existed, block)) {
found = true
return true
}
return false
})
return found
}
// validDenyName rejects file names that could break out of the quoted
// <Files "..."> argument and inject arbitrary .htaccess directives.
func validDenyName(name string) error {
if strings.TrimSpace(name) == "" {
return fmt.Errorf("empty file name")
}
if strings.ContainsAny(name, "\"\n\r<>\\*?[]${}") || strings.HasPrefix(strings.TrimLeft(name, " \t"), "#") {
return fmt.Errorf("file name %q contains characters unsafe for an .htaccess directive", name)
}
for _, r := range name {
if r < 0x20 || r == 0x7f {
return fmt.Errorf("file name %q contains control characters unsafe for an .htaccess directive", name)
}
}
return nil
}
func buildDenyBlock(name string, dirDeny bool) []byte {
if dirDeny {
return []byte(fmt.Sprintf("%s\n<FilesMatch \"^\">\nRequire all denied\n</FilesMatch>\n%s\n", vpBeginDir, vpEndDir))
}
return []byte(fmt.Sprintf("%s %s\n<Files \"%s\">\nRequire all denied\n</Files>\n%s %s\n",
vpBeginFile, name, name, vpEndFile, name))
}
func isKnownBackupPluginDir(dir string) bool {
base := strings.ToLower(filepath.Base(dir))
switch base {
case "ai1wm-backups", "wpvividbackups", "updraft":
return strings.EqualFold(filepath.Base(filepath.Dir(dir)), "wp-content")
default:
return false
}
}
func containsHtaccessBlock(content, block []byte) bool {
normalized := bytes.ReplaceAll(content, []byte("\r\n"), []byte("\n"))
return bytes.Contains(normalized, block)
}
func ensureTrailingNewline(b []byte) []byte {
if len(b) > 0 && b[len(b)-1] != '\n' {
return append(b, '\n')
}
return b
}
func readHtaccessState(htaccess, dir string) (htaccessState, error) {
info, err := os.Lstat(htaccess)
if err != nil {
if !os.IsNotExist(err) {
return htaccessState{}, fmt.Errorf("inspecting .htaccess: %v", err)
}
dirInfo, statErr := os.Stat(dir)
if statErr != nil {
return htaccessState{}, fmt.Errorf("inspecting target directory: %v", statErr)
}
uid, gid, ownerErr := ownerFromInfo(dirInfo)
if ownerErr != nil {
return htaccessState{}, ownerErr
}
return htaccessState{uid: uid, gid: gid, mode: 0644}, nil
}
if info.Mode()&os.ModeSymlink != 0 {
return htaccessState{}, fmt.Errorf("refusing symlinked .htaccess: %s", htaccess)
}
if !info.Mode().IsRegular() {
return htaccessState{}, fmt.Errorf("refusing non-regular .htaccess: %s", htaccess)
}
if info.Size() > maxVirtualPatchHtaccessSize {
return htaccessState{}, fmt.Errorf("refusing .htaccess larger than %d bytes", maxVirtualPatchHtaccessSize)
}
// #nosec G304 -- htaccess is under the resolved and validated target dir;
// O_NOFOLLOW rejects a path swap to a symlink.
file, err := os.OpenFile(htaccess, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return htaccessState{}, fmt.Errorf("opening .htaccess: %v", err)
}
defer file.Close()
openedInfo, err := file.Stat()
if err != nil {
return htaccessState{}, fmt.Errorf("stating .htaccess: %v", err)
}
if !os.SameFile(info, openedInfo) {
return htaccessState{}, fmt.Errorf(".htaccess changed while preparing virtual-patch")
}
content, err := io.ReadAll(io.LimitReader(file, maxVirtualPatchHtaccessSize+1))
if err != nil {
return htaccessState{}, fmt.Errorf("reading .htaccess: %v", err)
}
if len(content) > maxVirtualPatchHtaccessSize {
return htaccessState{}, fmt.Errorf("refusing .htaccess larger than %d bytes", maxVirtualPatchHtaccessSize)
}
uid, gid, err := ownerFromInfo(openedInfo)
if err != nil {
return htaccessState{}, err
}
return htaccessState{
content: content,
info: openedInfo,
existed: true,
uid: uid,
gid: gid,
mode: openedInfo.Mode().Perm(),
}, nil
}
func ownerFromInfo(info os.FileInfo) (int, int, error) {
st, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return 0, 0, fmt.Errorf("cannot determine filesystem owner")
}
return int(st.Uid), int(st.Gid), nil
}
func writeVirtualPatchTemp(dir string, content []byte, state htaccessState) (string, htaccessState, error) {
tmp, err := os.CreateTemp(dir, ".htaccess.csm-vpatch-*")
if err != nil {
return "", htaccessState{}, fmt.Errorf("creating temporary .htaccess: %v", err)
}
tmpPath := tmp.Name()
remove := true
defer func() {
_ = tmp.Close()
if remove {
_ = os.Remove(tmpPath)
}
}()
if _, writeErr := tmp.Write(content); writeErr != nil {
return "", htaccessState{}, fmt.Errorf("writing temporary .htaccess: %v", writeErr)
}
if syncErr := tmp.Sync(); syncErr != nil {
return "", htaccessState{}, fmt.Errorf("syncing temporary .htaccess: %v", syncErr)
}
if chownErr := chownFunc(tmp, state.uid, state.gid); chownErr != nil {
return "", htaccessState{}, fmt.Errorf("setting owner on temporary .htaccess: %v", chownErr)
}
if chmodErr := tmp.Chmod(state.mode); chmodErr != nil {
return "", htaccessState{}, fmt.Errorf("setting mode on temporary .htaccess: %v", chmodErr)
}
if syncErr := tmp.Sync(); syncErr != nil {
return "", htaccessState{}, fmt.Errorf("syncing temporary .htaccess metadata: %v", syncErr)
}
tmpInfo, err := tmp.Stat()
if err != nil {
return "", htaccessState{}, fmt.Errorf("stating temporary .htaccess: %v", err)
}
if err := tmp.Close(); err != nil {
return "", htaccessState{}, fmt.Errorf("closing temporary .htaccess: %v", err)
}
remove = false
return tmpPath, htaccessState{
content: append([]byte(nil), content...),
info: tmpInfo,
existed: true,
uid: state.uid,
gid: state.gid,
mode: state.mode,
}, nil
}
func htaccessStateMatchesPath(path string, state htaccessState) error {
info, err := os.Lstat(path)
if err != nil {
return err
}
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return fmt.Errorf(".htaccess path is no longer a regular file")
}
if state.info == nil || !os.SameFile(state.info, info) {
return fmt.Errorf(".htaccess inode changed")
}
// #nosec G304 -- path is either the validated .htaccess or its randomized
// same-directory staging name; O_NOFOLLOW and the repeated inode check keep
// a path swap from redirecting the read.
file, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return err
}
defer file.Close()
openedInfo, err := file.Stat()
if err != nil {
return err
}
if !os.SameFile(info, openedInfo) || !os.SameFile(state.info, openedInfo) {
return fmt.Errorf(".htaccess inode changed")
}
uid, gid, err := ownerFromInfo(openedInfo)
if err != nil {
return err
}
if uid != state.uid || gid != state.gid || openedInfo.Mode().Perm() != state.mode.Perm() {
return fmt.Errorf(".htaccess ownership or mode changed")
}
content, err := io.ReadAll(io.LimitReader(file, maxVirtualPatchHtaccessSize+1))
if err != nil {
return err
}
if len(content) > maxVirtualPatchHtaccessSize {
return fmt.Errorf(".htaccess content exceeds validation limit")
}
if !bytes.Equal(content, state.content) {
return fmt.Errorf(".htaccess content changed")
}
return nil
}
func removeVirtualPatchTemp(path string, state htaccessState) {
if htaccessStateMatchesPath(path, state) == nil {
_ = os.Remove(path)
}
}
// eachPrePatchBackup calls fn for every archived pre-patch backup, stopping
// when fn returns true. Reads go through an os.Root scoped to the backup
// directory, so a swapped entry cannot redirect them outside it.
func eachPrePatchBackup(fn func(meta QuarantineMeta, name string, read func(string) ([]byte, error)) bool) {
root, err := os.OpenRoot(htaccessBackupDirRoot)
if err != nil {
return
}
defer func() { _ = root.Close() }()
entries, err := fs.ReadDir(root.FS(), ".")
if err != nil {
return
}
read := func(name string) ([]byte, error) {
pathInfo, lstatErr := root.Lstat(name)
if lstatErr != nil {
return nil, lstatErr
}
if !pathInfo.Mode().IsRegular() {
return nil, fmt.Errorf("pre-patch backup entry is not a regular file")
}
file, openErr := root.OpenFile(name, os.O_RDONLY|syscall.O_NONBLOCK, 0)
if openErr != nil {
return nil, openErr
}
defer func() { _ = file.Close() }()
info, statErr := file.Stat()
if statErr != nil {
return nil, statErr
}
if !info.Mode().IsRegular() || !os.SameFile(pathInfo, info) {
return nil, fmt.Errorf("pre-patch backup entry changed while opening")
}
data, readErr := io.ReadAll(io.LimitReader(file, maxVirtualPatchHtaccessSize+1))
if readErr != nil {
return nil, readErr
}
if len(data) > maxVirtualPatchHtaccessSize {
return nil, fmt.Errorf("pre-patch backup entry exceeds %d bytes", maxVirtualPatchHtaccessSize)
}
currentInfo, lstatErr := root.Lstat(name)
if lstatErr != nil || !currentInfo.Mode().IsRegular() || !os.SameFile(pathInfo, currentInfo) {
return nil, fmt.Errorf("pre-patch backup entry changed while reading")
}
return data, nil
}
for _, entry := range entries {
if entry.IsDir() || !strings.HasSuffix(entry.Name(), ".meta") {
continue
}
metaData, readErr := read(entry.Name())
if readErr != nil {
continue
}
var meta QuarantineMeta
if json.Unmarshal(metaData, &meta) != nil {
continue
}
if fn(meta, entry.Name(), read) {
return
}
}
}
// findExistingPrePatchBackup returns an archived pre-patch backup for this
// .htaccess whose stored content and expected post-patch hash both match what
// this patch would produce. Both must match, so a backup recorded for a
// different patch is never reused as this one's rollback point.
func findExistingPrePatchBackup(htaccess string, state htaccessState, patched []byte) (virtualPatchBackup, bool) {
wantPatched := virtualPatchSHA256(patched)
// Rolling back a file CSM created means deleting it; rolling back one it
// appended to means restoring the original bytes. An empty .htaccess and
// no .htaccess hold the same content, so without this the reused backup
// could delete a file the customer owns.
wantAction := QuarantineRestoreReplaceIfUnchanged
if !state.existed {
wantAction = QuarantineRestoreRemoveIfUnchanged
}
var found virtualPatchBackup
ok := false
eachPrePatchBackup(func(meta QuarantineMeta, name string, read func(string) ([]byte, error)) bool {
if meta.OriginalPath != htaccess ||
meta.ExpectedCurrentSHA256 != wantPatched ||
meta.RestoreAction != wantAction ||
meta.Owner != state.uid ||
meta.Group != state.gid ||
meta.Mode != state.mode.String() ||
(state.existed && !meta.OriginalModTime.Equal(state.info.ModTime())) ||
meta.Size != int64(len(state.content)) {
return false
}
itemName := strings.TrimSuffix(name, ".meta")
archived, readErr := read(itemName)
if readErr != nil || !bytes.Equal(archived, state.content) {
return false
}
found = virtualPatchBackup{
itemPath: filepath.Join(htaccessBackupDirRoot, itemName),
metaPath: filepath.Join(htaccessBackupDirRoot, name),
}
ok = true
return true
})
return found, ok
}
// backupHtaccessBeforePatch records both the pre-patch content and the exact
// expected post-patch hash. Restore can then replace or remove .htaccess only
// while it still matches the version CSM wrote, preserving later user edits.
// The bool reports whether this call created the backup, so a failed patch
// never removes an archived copy shared with an earlier successful patch.
func backupHtaccessBeforePatch(htaccess string, state htaccessState, patched []byte) (virtualPatchBackup, bool, error) {
if err := quarantinefs.EnsureDir(htaccessBackupDirRoot, 0750); err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("creating backup dir: %v", err)
}
// Reuse only an identical recovery state. Equal bytes with different
// attributes need a new backup so restore keeps the captured metadata.
if existing, found := findExistingPrePatchBackup(htaccess, state, patched); found {
for _, path := range []string{existing.itemPath, existing.metaPath} {
if err := quarantinefs.SyncFilePath(path); err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("syncing existing backup: %w", err)
}
}
if err := quarantinefs.SyncDir(htaccessBackupDirRoot); err != nil {
return virtualPatchBackup{}, false, err
}
return existing, false, nil
}
stamp := time.Now().UTC().Format("20060102T150405Z")
pathSum := sha256.Sum256([]byte(htaccess))
backupFile, err := os.CreateTemp(htaccessBackupDirRoot, fmt.Sprintf("%s_vpatch_%x_", stamp, pathSum[:6]))
if err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("creating backup: %v", err)
}
backup := virtualPatchBackup{itemPath: backupFile.Name(), metaPath: backupFile.Name() + ".meta"}
keep := false
defer func() {
_ = backupFile.Close()
if !keep {
backup.remove()
}
}()
if state.existed {
if _, writeErr := backupFile.Write(state.content); writeErr != nil {
return virtualPatchBackup{}, false, fmt.Errorf("writing backup: %v", writeErr)
}
}
if chmodErr := backupFile.Chmod(0640); chmodErr != nil {
return virtualPatchBackup{}, false, fmt.Errorf("setting backup mode: %v", chmodErr)
}
if syncErr := backupFile.Sync(); syncErr != nil {
return virtualPatchBackup{}, false, fmt.Errorf("syncing backup: %v", syncErr)
}
if closeErr := backupFile.Close(); closeErr != nil {
return virtualPatchBackup{}, false, fmt.Errorf("closing backup: %v", closeErr)
}
restoreAction := QuarantineRestoreReplaceIfUnchanged
if !state.existed {
restoreAction = QuarantineRestoreRemoveIfUnchanged
}
meta := QuarantineMeta{
OriginalPath: htaccess,
Owner: state.uid,
Group: state.gid,
Mode: state.mode.String(),
Size: int64(len(state.content)),
QuarantineAt: time.Now().UTC(),
Reason: "exposed-file virtual-patch: pre-patch .htaccess backup",
RestoreAction: restoreAction,
ExpectedCurrentSHA256: virtualPatchSHA256(patched),
}
if state.existed {
meta.OriginalModTime = state.info.ModTime().UTC()
}
metaJSON, err := json.Marshal(meta)
if err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("encoding backup meta: %v", err)
}
metaFile, err := os.OpenFile(backup.metaPath, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0600)
if err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("creating backup meta: %v", err)
}
if _, err := metaFile.Write(metaJSON); err != nil {
_ = metaFile.Close()
return virtualPatchBackup{}, false, fmt.Errorf("writing backup meta: %v", err)
}
if err := metaFile.Sync(); err != nil {
_ = metaFile.Close()
return virtualPatchBackup{}, false, fmt.Errorf("syncing backup meta: %v", err)
}
if err := metaFile.Close(); err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("closing backup meta: %v", err)
}
if err := quarantinefs.SyncDir(htaccessBackupDirRoot); err != nil {
return virtualPatchBackup{}, false, fmt.Errorf("syncing backup directory: %w", err)
}
keep = true
return backup, true, nil
}
func (backup virtualPatchBackup) remove() {
if backup.metaPath != "" {
_ = os.Remove(backup.metaPath)
}
if backup.itemPath != "" {
_ = os.Remove(backup.itemPath)
}
}
func virtualPatchSHA256(content []byte) string {
sum := sha256.Sum256(content)
return fmt.Sprintf("sha256:%x", sum[:])
}
func parseVirtualPatchMode(value string) (os.FileMode, error) {
if len(value) != 10 || value[0] != '-' {
return 0, fmt.Errorf("invalid virtual-patch mode %q", value)
}
perms := value[1:]
wantChars := "rwxrwxrwx"
bits := []os.FileMode{0400, 0200, 0100, 0040, 0020, 0010, 0004, 0002, 0001}
var mode os.FileMode
for i, char := range perms {
if char == '-' {
continue
}
if char != rune(wantChars[i]) {
return 0, fmt.Errorf("invalid virtual-patch mode %q", value)
}
mode |= bits[i]
}
return mode, nil
}
// VirtualPatchExposedFindings applies (apply=true) or previews (apply=false) a
// deny rule for each virtual-patchable web_exposed_* finding, deduplicated by
// path. It returns one auto_response action finding per file.
func VirtualPatchExposedFindings(_ *config.Config, findings []alert.Finding, apply bool) []alert.Finding {
var actions []alert.Finding
seen := make(map[string]struct{})
for _, f := range findings {
if !isVirtualPatchableExposedCheck(f.Check) {
continue
}
path := f.FilePath
if path == "" {
path = extractFilePath(f.Message)
}
if path == "" {
continue
}
if _, ok := seen[path]; ok {
continue
}
seen[path] = struct{}{}
if !apply {
details := "Enable auto_response.virtual_patch_exposed_files=auto with dry_run:false, or run `csm virtual-patch --apply`, to write the deny rule."
if f.Check == "web_exposed_sample_sql" {
details = "Warning-only sample SQL is never enforced automatically; run `csm virtual-patch --apply` to write the deny rule."
}
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("VIRTUAL-PATCH (preview): would deny HTTP access to %s", path),
Details: details,
Timestamp: time.Now(),
})
continue
}
res := VirtualPatchExposedFile(path)
switch {
case res.Success && res.Reverted:
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("%s %s", virtualPatchRevertedPrefix, path),
Details: fmt.Sprintf("%s. CSM had already denied this path; something rewrote the .htaccess and the file was downloadable until this scan. "+
"Backup plugins own the .htaccess in their own directory and restore it on every run, so the deny belongs somewhere the plugin does not overwrite.",
res.Description),
Timestamp: time.Now(),
})
case res.Success:
actions = append(actions, alert.Finding{
Severity: alert.Critical,
Check: "auto_response",
Message: fmt.Sprintf("VIRTUAL-PATCH: denied HTTP access to %s", path),
Details: res.Description,
Timestamp: time.Now(),
})
case res.Error != "" && res.Error != virtualPatchAlreadyApplied:
actions = append(actions, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("VIRTUAL-PATCH failed: %s", path),
Details: res.Error,
Timestamp: time.Now(),
})
}
}
return actions
}
// AutoVirtualPatchExposedFiles is the scan-time entry point. It acts only when
// auto_response is enabled and the mode is "auto"; the write is gated by the
// shared auto_response dry_run flag (dry_run reports the intended denials).
func AutoVirtualPatchExposedFiles(cfg *config.Config, findings []alert.Finding) []alert.Finding {
if cfg == nil || !cfg.AutoResponse.Enabled || cfg.VirtualPatchMode() != config.VirtualPatchAuto {
return nil
}
autoEligible := make([]alert.Finding, 0, len(findings))
for _, finding := range findings {
if finding.Check != "web_exposed_sample_sql" {
autoEligible = append(autoEligible, finding)
}
}
return VirtualPatchExposedFindings(cfg, autoEligible, !cfg.AutoResponseDryRunEnabled())
}
//go:build linux
package checks
import (
"errors"
"fmt"
"os"
"golang.org/x/sys/unix"
)
func commitVirtualPatchTemp(tmp, htaccess string, state, tempState htaccessState) error {
if !state.existed {
if err := unix.Renameat2(unix.AT_FDCWD, tmp, unix.AT_FDCWD, htaccess, unix.RENAME_NOREPLACE); err != nil {
if errors.Is(err, unix.EEXIST) {
return fmt.Errorf(".htaccess changed while preparing virtual-patch")
}
return fmt.Errorf("committing new .htaccess: %v", err)
}
if err := htaccessStateMatchesPath(htaccess, tempState); err != nil {
if rollbackErr := unix.Renameat2(unix.AT_FDCWD, htaccess, unix.AT_FDCWD, tmp, unix.RENAME_NOREPLACE); rollbackErr != nil {
return fmt.Errorf("%w: temporary .htaccess changed and rollback failed: %v", errVirtualPatchRollbackIncomplete, rollbackErr)
}
return fmt.Errorf("temporary .htaccess changed before commit: %v", err)
}
return nil
}
if err := unix.Renameat2(unix.AT_FDCWD, tmp, unix.AT_FDCWD, htaccess, unix.RENAME_EXCHANGE); err != nil {
return fmt.Errorf("atomically exchanging .htaccess: %v", err)
}
oldErr := htaccessStateMatchesPath(tmp, state)
newErr := htaccessStateMatchesPath(htaccess, tempState)
if oldErr != nil || newErr != nil {
if rollbackErr := unix.Renameat2(unix.AT_FDCWD, tmp, unix.AT_FDCWD, htaccess, unix.RENAME_EXCHANGE); rollbackErr != nil {
return fmt.Errorf("%w: .htaccess changed while preparing virtual-patch and rollback failed: %v", errVirtualPatchRollbackIncomplete, rollbackErr)
}
if oldErr != nil {
return fmt.Errorf(".htaccess changed while preparing virtual-patch: %v", oldErr)
}
return fmt.Errorf("temporary .htaccess changed before commit: %v", newErr)
}
if err := os.Remove(tmp); err != nil {
if rollbackErr := unix.Renameat2(unix.AT_FDCWD, tmp, unix.AT_FDCWD, htaccess, unix.RENAME_EXCHANGE); rollbackErr != nil {
return fmt.Errorf("%w: removing replaced .htaccess failed (%v) and rollback failed: %v", errVirtualPatchRollbackIncomplete, err, rollbackErr)
}
_ = os.Remove(tmp)
return fmt.Errorf("removing replaced .htaccess: %v", err)
}
return nil
}
package checks
import (
"bytes"
"fmt"
"io"
"os"
"path/filepath"
"syscall"
"github.com/pidginhost/csm/internal/safepath"
)
// RestoreVirtualPatchBackup reverts only the captured patch, using the pinned
// destination directory so tenant renames cannot redirect reads or changes.
func RestoreVirtualPatchBackup(backupPath string, target *safepath.Target, meta QuarantineMeta) error {
if target.Name != ".htaccess" {
return fmt.Errorf("virtual-patch restore applies only to .htaccess")
}
if meta.RestoreAction != QuarantineRestoreReplaceIfUnchanged &&
meta.RestoreAction != QuarantineRestoreRemoveIfUnchanged {
return fmt.Errorf("unsupported virtual-patch restore action %q", meta.RestoreAction)
}
mode, err := parseVirtualPatchMode(meta.Mode)
if err != nil {
return err
}
state, err := readRestoreHtaccess(target.Parent, target.Name)
if err != nil {
return fmt.Errorf("%w: %v", ErrVirtualPatchRestoreConflict, err)
}
if virtualPatchSHA256(state.content) != meta.ExpectedCurrentSHA256 ||
state.uid != meta.Owner || state.gid != meta.Group || state.mode.Perm() != mode.Perm() {
return fmt.Errorf("%w: live file was modified after enforcement", ErrVirtualPatchRestoreConflict)
}
var content []byte
if meta.RestoreAction == QuarantineRestoreReplaceIfUnchanged {
// Quarantine is daemon-owned; the final backup name is still opened
// without following a symlink and validated before reading.
backup, openErr := os.OpenFile(backupPath, os.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_NONBLOCK, 0) // #nosec G304 -- caller supplies a daemon-owned quarantine entry; no-follow and regular-file checks protect the read
if openErr != nil {
return fmt.Errorf("opening virtual-patch backup: %w", openErr)
}
content, err = readRestoreContent(backup)
closeErr := backup.Close()
if err != nil {
return err
}
if closeErr != nil {
return closeErr
}
}
if opErr := target.Check(); opErr != nil {
return fmt.Errorf("%w: %v", ErrVirtualPatchRestoreConflict, opErr)
}
stage, stageName, err := target.Parent.CreatePrivateTemp()
if err != nil {
return err
}
defer func() { _ = stage.Close() }()
keep := false
defer func() {
if !keep {
_ = target.Parent.RemoveDir(stageName)
}
}()
temp, err := stage.CreateTemp()
if err != nil {
return err
}
name := filepath.Base(temp.Name())
defer func() {
_ = temp.Close()
if !keep {
_ = stage.Remove(name)
}
}()
if _, opErr := temp.Write(content); opErr != nil {
return opErr
}
if opErr := temp.Chmod(mode); opErr != nil {
return opErr
}
if opErr := temp.Chown(meta.Owner, meta.Group); opErr != nil {
return opErr
}
if !meta.OriginalModTime.IsZero() {
if opErr := safepath.SetModTime(temp, meta.OriginalModTime); opErr != nil {
return opErr
}
}
if opErr := temp.Sync(); opErr != nil {
return opErr
}
prepared, err := temp.Stat()
if err != nil {
return err
}
if opErr := temp.Close(); opErr != nil {
return opErr
}
if opErr := target.Check(); opErr != nil {
return fmt.Errorf("%w: %v", ErrVirtualPatchRestoreConflict, opErr)
}
remove := meta.RestoreAction == QuarantineRestoreRemoveIfUnchanged
if remove {
if opErr := stage.Remove(name); opErr != nil {
return opErr
}
// Isolate the old file without leaving a placeholder whose later
// unlink could delete a concurrent replacement at the live name.
err = target.Parent.RenameTo(target.Name, stage, name)
} else {
err = stage.ExchangeTo(name, target.Parent, target.Name)
}
if err != nil {
return fmt.Errorf("%w: %v", ErrVirtualPatchRestoreConflict, err)
}
if virtualPatchRestoreAfterMoveForTest != nil {
virtualPatchRestoreAfterMoveForTest()
}
conflict := func(cause error) error {
keep = true
return fmt.Errorf("%w: %v; recovery files retained in %s", ErrVirtualPatchRestoreConflict, cause, stageName)
}
rollback := func(cause error) error {
keep = true
if remove {
if opErr := stage.RenameTo(name, target.Parent, target.Name); opErr != nil {
return conflict(fmt.Errorf("%v; rollback failed: %w", cause, opErr))
}
keep = false
return fmt.Errorf("%w: %v", ErrVirtualPatchRestoreConflict, cause)
}
// Capture the live name before deciding what to put back. A second
// exchange after a check could overwrite another intervening edit.
captured, opErr := stage.CreateTemp()
if opErr != nil {
return conflict(opErr)
}
captureName := filepath.Base(captured.Name())
if opErr := captured.Close(); opErr != nil {
return conflict(opErr)
}
if opErr := stage.Remove(captureName); opErr != nil {
return conflict(opErr)
}
captureErr := target.Parent.RenameTo(target.Name, stage, captureName)
if captureErr != nil && !os.IsNotExist(captureErr) {
return conflict(fmt.Errorf("%v; rollback failed: %w", cause, captureErr))
}
restoreName := name
if captureErr == nil {
current, readErr := readRestoreHtaccess(stage, captureName)
if readErr != nil || !os.SameFile(prepared, current.info) ||
!bytes.Equal(content, current.content) || current.uid != meta.Owner ||
current.gid != meta.Group || current.mode != mode.Perm() {
restoreName = captureName
}
}
if opErr := stage.RenameTo(restoreName, target.Parent, target.Name); opErr != nil {
return conflict(fmt.Errorf("%v; rollback failed: %w", cause, opErr))
}
if captureErr != nil {
keep = false
return fmt.Errorf("%w: %v", ErrVirtualPatchRestoreConflict, cause)
}
// Keep every captured inode on conflict, including one that looked
// unchanged: a writer may still hold an open descriptor to it.
return conflict(cause)
}
oldState, err := readRestoreHtaccess(stage, name)
if err != nil {
return rollback(err)
}
if !os.SameFile(state.info, oldState.info) || !bytes.Equal(state.content, oldState.content) ||
state.uid != oldState.uid || state.gid != oldState.gid || state.mode != oldState.mode {
return rollback(fmt.Errorf("live file changed during restore"))
}
if remove {
if _, statErr := target.Parent.Stat(target.Name); !os.IsNotExist(statErr) {
return conflict(fmt.Errorf("live file was recreated during restore"))
}
} else {
current, readErr := readRestoreHtaccess(target.Parent, target.Name)
if readErr != nil {
return rollback(readErr)
}
if !os.SameFile(prepared, current.info) || !bytes.Equal(content, current.content) ||
current.uid != meta.Owner || current.gid != meta.Group || current.mode != mode.Perm() {
return conflict(fmt.Errorf("prepared restore changed"))
}
}
if opErr := target.Check(); opErr != nil {
return rollback(opErr)
}
if opErr := target.Parent.Sync(); opErr != nil {
keep = true
return fmt.Errorf("restore applied but destination sync failed; recovery files retained in %s: %w", stageName, opErr)
}
if opErr := stage.Remove(name); opErr != nil {
keep = true
return fmt.Errorf("restore applied but replaced file could not be removed from %s: %w", stageName, opErr)
}
return nil
}
var virtualPatchRestoreAfterMoveForTest func()
func readRestoreContent(file *os.File) ([]byte, error) {
info, err := file.Stat()
if err != nil {
return nil, err
}
if !info.Mode().IsRegular() || info.Size() > maxVirtualPatchHtaccessSize {
return nil, fmt.Errorf("restore input must be a bounded regular file")
}
content, err := io.ReadAll(io.LimitReader(file, maxVirtualPatchHtaccessSize+1))
if err != nil {
return nil, err
}
if len(content) > maxVirtualPatchHtaccessSize {
return nil, fmt.Errorf("restore input exceeds size limit")
}
return content, nil
}
func readRestoreHtaccess(dir *safepath.Dir, name string) (htaccessState, error) {
file, err := dir.OpenFile(name, os.O_RDONLY, 0)
if err != nil {
return htaccessState{}, err
}
defer file.Close()
content, err := readRestoreContent(file)
if err != nil {
return htaccessState{}, err
}
info, err := file.Stat()
if err != nil {
return htaccessState{}, err
}
uid, gid, err := ownerFromInfo(info)
if err != nil {
return htaccessState{}, err
}
return htaccessState{content: content, info: info, existed: true, uid: uid, gid: gid, mode: info.Mode().Perm()}, nil
}
package checks
import (
"encoding/json"
"errors"
"io"
"os"
"path/filepath"
"sync"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
)
const fileResponseStateName = "file-response.json"
// An attempt is reserved durably as failed before touching a customer file.
// Interrupted operations therefore keep both their budget and failure charge.
type fileResponseAttempt struct {
At time.Time `json:"at"`
Account string `json:"account"`
Failed bool `json:"failed"`
}
type fileResponseState struct {
Version int `json:"version"`
Attempts []fileResponseAttempt `json:"attempts"`
}
var fileResponseNow = time.Now
var writeFileResponseState = atomicio.AtomicWriteJSON
// A refusal leaves the customer file untouched and consumes attempt capacity,
// but does not indicate a failed response mechanism.
var errFileResponseRefused = errors.New("file response refused")
// fileResponseRefusal classifies an error as a refusal for errors.Is while
// keeping the refusing check's message. Manual remediation shares these
// checks, so the breaker's accounting must not show up in operator text.
type fileResponseRefusal struct{ err error }
func (r fileResponseRefusal) Error() string { return r.err.Error() }
func (r fileResponseRefusal) Unwrap() error { return r.err }
func (fileResponseRefusal) Is(target error) bool { return target == errFileResponseRefused }
func refuseFileResponse(err error) error { return fileResponseRefusal{err: err} }
func fileResponseSourceError(err error) error {
// A socket replacement cannot be opened as a file and returns ENXIO
// before the descriptor-based regular-file check can refuse it.
if errors.Is(err, os.ErrNotExist) || errors.Is(err, unix.ELOOP) || errors.Is(err, unix.ENOTDIR) || errors.Is(err, unix.ENXIO) {
return refuseFileResponse(err)
}
return err
}
// runAutoFileResponse is shared by the automatic PHP/access-file cleaners and
// both quarantine entry points. Manual remediation does not enter this gate.
// A nonblocking process-shared lock keeps concurrency from spending the same
// slot twice, without making a detector wait behind filesystem remediation.
func runAutoFileResponse(cfg *config.Config, path string, info os.FileInfo, apply func() error) *alert.Finding {
now := fileResponseNow()
notice := func(reason, account, detail string) *alert.Finding {
return fileResponseNotice(cfg.StatePath, now, reason, account, detail)
}
if !info.Mode().IsRegular() {
return notice("file_type", "", "Automatic file response refused a directory or special file. Review the detection and use manual remediation for this target.")
}
hostLimit, accountLimit, failureLimit := fileResponseLimits(cfg)
if hostLimit < 1 || accountLimit < 1 || failureLimit < 1 || hostLimit > config.MaxFileResponseLimit || accountLimit > config.MaxFileResponseLimit || failureLimit > config.MaxFileResponseLimit {
return notice("state", "", "Automatic file response is paused because its safety limits are invalid. Correct the configuration; detection continues.")
}
if cfg.StatePath == "" || !filepath.IsAbs(cfg.StatePath) {
return notice("state", "", "Automatic file response is paused because its safety state directory is unavailable. Detection continues.")
}
// The operator-owned state directory is provisioned at daemon startup.
// Never create a new directory here and silently reset a missing state root.
lockPath := filepath.Join(cfg.StatePath, "file-response.lock")
// #nosec G304 -- fixed filename inside the operator-configured state directory.
lock, err := os.OpenFile(lockPath, os.O_CREATE|os.O_RDWR|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0600)
if err != nil {
return notice("state", "", "Automatic file response is paused because its safety state cannot be locked. Detection continues.")
}
defer lock.Close()
lockInfo, err := lock.Stat()
if err != nil || !lockInfo.Mode().IsRegular() {
return notice("state", "", "Automatic file response is paused because its safety lock is not a regular file. Detection continues.")
}
// #nosec G115 -- an open POSIX file descriptor fits in int.
if err = unix.Flock(int(lock.Fd()), unix.LOCK_EX|unix.LOCK_NB); err != nil {
return notice("busy", "", "Automatic file response refused an attempt while its safety state was busy. Detection continues; review the original finding before retrying manually.")
}
// Close releases the flock. The lock file stays in place so all processes
// keep locking the same inode, including callers that already opened it.
statePath := filepath.Join(cfg.StatePath, fileResponseStateName)
state, err := readFileResponseState(statePath, now)
if err != nil {
csmlog.Warn("automatic file response state unreadable", "err", err)
return notice("state", "", "Automatic file response is paused because its safety state cannot be read. Repair the state storage; detection continues.")
}
_, account, _ := accountRootOf(path)
// Unknown paths share one budget. Finding text/TenantID must not select a
// fresh account bucket, and UID alone conflates root-owned tenant files.
used, failures := 0, 0
for _, attempt := range state.Attempts {
if attempt.Account == account {
used++
}
if attempt.Failed {
failures++
}
}
if failures >= failureLimit {
return notice("failures", "", "Automatic file response is paused after repeated action failures in the rolling hour. Detection continues; inspect the action log and recovery copies before manual remediation.")
}
if len(state.Attempts) >= hostLimit {
return notice("host_limit", "", "Automatic file response reached its host limit for the rolling hour. Detection continues; review outstanding findings for manual remediation.")
}
if used >= accountLimit {
return notice("account_limit", account, "Automatic file response reached its account limit for the rolling hour. Other accounts remain eligible; detection continues.")
}
state.Attempts = append(state.Attempts, fileResponseAttempt{At: now, Account: account, Failed: true})
if err = writeFileResponseState(statePath, 0600, state); err != nil {
csmlog.Warn("automatic file response reservation failed", "err", err)
return notice("state", "", "Automatic file response is paused because its safety reservation could not be saved. Detection continues.")
}
// The caller's snapshot predates budget persistence. Recheck it after that
// I/O; the descriptor-based quarantine/cleaner checks it again when opening.
current, err := os.Lstat(path)
err = fileResponseSourceError(err)
if err == nil && (!sameFileIdentity(info, current) || !sameContentShape(info, current)) {
err = refuseFileResponse(errors.New("file changed before automatic response"))
}
if err == nil {
err = apply()
}
if err != nil && !errors.Is(err, errFileResponseRefused) {
csmlog.Warn("automatic file response failed", "path", path, "err", err)
if failures+1 >= failureLimit {
return notice("failures", "", "Automatic file response is paused after repeated action failures in the rolling hour. Detection continues; inspect the action log and recovery copies before manual remediation.")
}
return nil
}
if err != nil {
csmlog.Warn("automatic file response refused", "path", path, "err", err)
}
state.Attempts[len(state.Attempts)-1].Failed = false
if err = writeFileResponseState(statePath, 0600, state); err != nil {
// The durable reservation remains charged even when outcome persistence
// fails. Never retry the action or refund its slot based on this error.
csmlog.Warn("automatic file response outcome save failed", "err", err)
return notice("state", "", "An automatic file response completed but its safety outcome could not be saved. Its reservation remains charged; inspect storage and recovery evidence.")
}
return nil
}
func fileResponseLimits(cfg *config.Config) (host, account, failures int) {
host, account, failures = cfg.AutoResponse.MaxFileActionsPerHour, cfg.AutoResponse.MaxFileActionsPerAccountPerHour, cfg.AutoResponse.MaxFileActionFailuresPerHour
if host == 0 {
host = config.DefaultMaxFileActionsPerHour
}
if account == 0 {
account = config.DefaultMaxFileActionsPerAccountPerHour
}
if failures == 0 {
failures = config.DefaultMaxFileActionFailuresPerHour
}
return
}
func readFileResponseState(path string, now time.Time) (*fileResponseState, error) {
// #nosec G304 -- fixed filename inside the operator-owned state directory.
file, err := os.OpenFile(path, os.O_RDONLY|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if errors.Is(err, os.ErrNotExist) {
return &fileResponseState{Version: 1}, nil
}
if err != nil {
return nil, err
}
defer file.Close()
info, err := file.Stat()
if err != nil {
return nil, err
}
if !info.Mode().IsRegular() {
return nil, errors.New("safety state is not a regular file")
}
// Bound corrupted or externally replaced state before decoding.
const maxSize = 4 << 20
data, err := io.ReadAll(io.LimitReader(file, maxSize+1))
if err != nil {
return nil, err
}
if len(data) > maxSize {
return nil, errors.New("safety state is too large")
}
var stored struct {
Version int `json:"version"`
Attempts json.RawMessage `json:"attempts"`
}
if err := json.Unmarshal(data, &stored); err != nil {
return nil, err
}
if stored.Version != 1 {
return nil, errors.New("unknown safety state format")
}
if len(stored.Attempts) == 0 {
return nil, errors.New("safety state is missing reservations")
}
// Account and Failed have meaningful zero values. Missing or null fields
// must not silently turn a charged reservation into an unknown account
// or a successful action and reopen capacity after state corruption.
var attempts []struct {
At time.Time `json:"at"`
Account *string `json:"account"`
Failed *bool `json:"failed"`
}
if err := json.Unmarshal(stored.Attempts, &attempts); err != nil {
return nil, err
}
if attempts == nil {
return nil, errors.New("safety state has null reservations")
}
state := &fileResponseState{Version: stored.Version}
for _, attempt := range attempts {
if attempt.At.IsZero() {
return nil, errors.New("safety state has an undated reservation")
}
if attempt.Account == nil || attempt.Failed == nil {
return nil, errors.New("safety state has an incomplete reservation")
}
// Future entries remain charged when the clock moves backwards.
if attempt.At.After(now.Add(-time.Hour)) {
state.Attempts = append(state.Attempts, fileResponseAttempt{At: attempt.At, Account: *attempt.Account, Failed: *attempt.Failed})
}
}
return state, nil
}
var fileResponseNotices = struct {
sync.Mutex
last map[string]time.Time
}{last: make(map[string]time.Time)}
// Pause notices must not become a second alert flood during a detector or
// storage failure. Dedup is local (restart reports the pause again), bounded,
// and independent of the safety state so storage failures are reportable.
func fileResponseNotice(statePath string, now time.Time, reason, account, detail string) *alert.Finding {
// An account-wide detector fault can exhaust many tenant budgets. Emit
// one host notice per cause, so those warnings cannot spend the alert
// budget that other non-critical detections need.
key := statePath + "\x00" + reason
fileResponseNotices.Lock()
defer fileResponseNotices.Unlock()
if last, ok := fileResponseNotices.last[key]; ok && now.Sub(last) < time.Hour {
return nil
}
for k, last := range fileResponseNotices.last {
if now.Sub(last) >= time.Hour {
delete(fileResponseNotices.last, k)
}
}
if len(fileResponseNotices.last) >= 1024 {
// Do not let arbitrary state paths grow daemon memory without a bound.
// Keep the existing notices suppressed until their window expires.
return nil
}
fileResponseNotices.last[key] = now
message := "Automatic file response paused"
if account != "" {
message = "Automatic file response paused for accounts at their limit"
}
return &alert.Finding{Check: "auto_response_paused", Severity: alert.Warning, Message: message, Details: detail, DedupKey: reason, Timestamp: now}
}
package checks
import (
"bufio"
"context"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync/atomic"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// fileIndexScanCount schedules a full walk at startup and every sixth scan.
// Failed or interrupted walks reset it so retries cannot trust cached mtimes
// for directories that have become unreadable.
var fileIndexScanCount int32
// fileIndexShrinkSkips counts consecutive scans whose current index shrank far
// enough to trip the mass-deletion guard. It self-heals the guard: after
// fileIndexShrinkPromoteThreshold consecutive shrinking scans the smaller index
// is promoted as the new baseline, so new-file detection resumes instead of
// forever diffing against a stale, permanently-larger previous index.
// The live CheckFileIndex path is serialized by fileIndexLiveScanGate, so this
// counter advances once per completed stateful index build rather than racing
// between a periodic deep scan and an operator-triggered tier run.
var fileIndexShrinkSkips int32
// fileIndexLiveScanGate serializes the stateful CheckFileIndex path. The
// scanner reads and writes fileindex.current, fileindex.previous, dircache.json,
// and the shrink streak as one logical baseline update; overlapping live runs
// can otherwise double-count one shrink window or compare against a half-written
// peer scan. ForceFileIndex audit scans return before this gate because they do
// not touch live baseline state.
var fileIndexLiveScanGate = make(chan struct{}, 1)
// fileIndexShrinkPromoteThreshold is how many consecutive shrinking scans must
// occur before the smaller index becomes the baseline. Three deep cycles is
// long enough to rule out a one-off read glitch or a transient mass rename,
// while still healing the wedge automatically within a few scan intervals so no
// operator has to delete fileindex.previous by hand.
const fileIndexShrinkPromoteThreshold = 3
// evaluateFileIndexShrink classifies the current-vs-previous index sizes and
// advances the consecutive-shrink counter. isShrink reports whether this scan
// tripped the mass-deletion guard (an empty current against a populated
// previous, or a non-empty current below half the previous). promote reports
// whether the shrink has now persisted for fileIndexShrinkPromoteThreshold
// consecutive scans, in which case the caller must adopt the smaller index as
// the baseline. Empty current indexes are never promoted: a persistent read
// failure has the same shape, and adopting it would make recovery alert on
// every indexed file as if the whole host were new.
// A non-shrink scan resets the streak so only *consecutive* shrinks promote.
func evaluateFileIndexShrink(prevLen, curLen int) (isShrink, promote bool) {
if prevLen > 10 && curLen == 0 {
atomic.StoreInt32(&fileIndexShrinkSkips, 0)
return true, false
}
isShrink = prevLen > 0 && curLen > 0 && curLen*2 < prevLen
if !isShrink {
atomic.StoreInt32(&fileIndexShrinkSkips, 0)
return false, false
}
if atomic.AddInt32(&fileIndexShrinkSkips, 1) >= fileIndexShrinkPromoteThreshold {
atomic.StoreInt32(&fileIndexShrinkSkips, 0)
return true, true
}
return true, false
}
// suspiciousExtensions are file extensions worth reading in a web root. Being
// listed here only routes a file into content analysis; the verdict is still
// the content scanner's. ".phps" earns its place despite a stock handler
// rendering it as source rather than executing it: staging a dropper under that
// extension leaves it unread until a rename makes it live.
var suspiciousExtensions = map[string]bool{
".phtml": true, ".pht": true, ".php5": true, ".phps": true,
".haxor": true, ".cgix": true,
}
// dirMtimeCache maps directory paths to their last-known mtime (unix seconds).
// Directories with unchanged mtime are skipped during scanning.
type dirMtimeCache map[string]int64
func loadDirCache(stateDir string) (dirMtimeCache, error) {
cache := make(dirMtimeCache)
data, err := osFS.ReadFile(filepath.Join(stateDir, "dircache.json"))
if os.IsNotExist(err) {
return cache, nil
}
if err != nil {
return cache, err
}
err = json.Unmarshal(data, &cache)
return cache, err
}
func saveDirCache(stateDir string, cache dirMtimeCache) error {
data, _ := json.Marshal(cache)
tmpPath := filepath.Join(stateDir, "dircache.json.tmp")
if err := os.WriteFile(tmpPath, data, 0600); err != nil {
return err
}
return os.Rename(tmpPath, filepath.Join(stateDir, "dircache.json"))
}
// dirChanged returns true if the directory mtime has changed since last scan.
// Updates the cache with the new mtime. If forceFullScan is true, always
// returns true to force a ReadDir regardless of mtime (catches writes that
// bypass parent mtime updates, e.g. hard links or mount tricks).
func dirChanged(dir string, cache dirMtimeCache, forceFullScan bool) bool {
info, err := osFS.Stat(dir)
if err != nil {
return true // can't stat, scan it to be safe
}
mtime := info.ModTime().Unix()
prev, exists := cache[dir]
cache[dir] = mtime
if forceFullScan {
return true
}
if !exists {
return true // first time seeing this dir
}
return mtime != prev
}
type subtreeChangeTracker struct {
cache dirMtimeCache
changedBelow map[string]struct{}
built bool
}
func newSubtreeChangeTracker(cache dirMtimeCache) *subtreeChangeTracker {
snapshot := make(dirMtimeCache, len(cache))
for dir, mtime := range cache {
snapshot[dir] = mtime
}
return &subtreeChangeTracker{cache: snapshot}
}
func (t *subtreeChangeTracker) hasChangedDir(ctx context.Context, dir string) bool {
if t == nil {
return false
}
if !t.built {
t.build(ctx)
}
_, ok := t.changedBelow[dir]
return ok
}
func (t *subtreeChangeTracker) build(ctx context.Context) {
t.built = true
t.changedBelow = make(map[string]struct{})
for cached, mtime := range t.cache {
if ctx != nil && ctx.Err() != nil {
return
}
info, err := osFS.Stat(cached)
if err == nil && info.ModTime().Unix() == mtime {
continue
}
for parent := filepath.Dir(cached); parent != "." && parent != string(filepath.Separator); parent = filepath.Dir(parent) {
if _, ok := t.cache[parent]; ok {
t.changedBelow[parent] = struct{}{}
}
}
}
}
// CheckFileIndex builds an index of suspicious files using pure Go directory
// reads, diffs against the previous index, and alerts on new files.
// Uses directory mtime caching: unchanged dirs carry forward previous entries
// without calling ReadDir, while changed dirs are re-scanned.
//
// When ctx carries AccountScanOptions with ForceFileIndex=true the function
// runs in audit mode: it enumerates only the in-scope account, bypasses the
// directory mtime cache entirely, and writes none of the live state files
// (fileindex.current, fileindex.previous, dircache.json). The normal incremental
// baseline is left byte-for-byte intact. All indexed files are treated as new
// so the caller receives findings for the full current state of the account.
func CheckFileIndex(ctx context.Context, cfg *config.Config, st *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
// Audit mode: enumerate only the in-scope account, bypass all state I/O.
// CRITICAL: none of fileindex.current, fileindex.previous, or dircache.json
// may be written -- corrupting them would silently break future incremental
// scans by shifting their baseline.
if scanForceFileIndex(ctx) {
// Fresh empty cache so every directory is treated as changed.
auditCache := make(dirMtimeCache)
// Empty prevByDir so every indexed file is diff'd as new.
currentEntries := buildFileIndex(ctx, auditCache, nil, true)
if ctx.Err() != nil {
return nil
}
// All entries are "new" -- diff against empty previous set.
return checkFileIndexAnalyzeNewFiles(ctx, cfg, currentEntries)
}
return fileIndexQueues.run(ctx, func(work *fileIndexWork) []alert.Finding {
return checkFileIndexLive(ctx, cfg, st, work)
})
}
func checkFileIndexLive(ctx context.Context, cfg *config.Config, st *state.Store, work *fileIndexWork) []alert.Finding {
ctx = context.WithValue(ctx, fileIndexWorkKey{}, work)
scanNum := atomic.AddInt32(&fileIndexScanCount, 1)
forceFullScan := scanNum == 1 || scanNum%6 == 0
defer func() {
work.local()
if ctx.Err() != nil {
atomic.StoreInt32(&fileIndexScanCount, 0)
}
}()
indexDir := cfg.StatePath
currentPath := filepath.Join(indexDir, "fileindex.current")
previousPath := filepath.Join(indexDir, "fileindex.previous")
// Load caches
dirCache, err := loadDirCache(indexDir)
work.observe(err)
// Build a set of previous entries grouped by directory ancestry, so
// unchanged dirs can carry forward their whole subtree without ReadDir.
previousEntries, err := loadIndex(previousPath)
work.observe(err)
prevByDir := groupEntriesByUploadDir(previousEntries)
// Track completeness locally as well as in the runner: direct calls must
// not replace the baseline with a partial walk either.
indexCtx, completion := withIncompleteCheckCollector(ctx)
work.execution()
currentEntries := buildFileIndex(indexCtx, dirCache, prevByDir, forceFullScan)
incomplete := completion.contains("file_index")
if incomplete {
work.fail()
atomic.StoreInt32(&fileIndexScanCount, 0)
markCheckIncomplete(ctx, "file_index")
}
// A cancelled scan produced a partial index. Do not write or promote it:
// the partial set or its mtimes would make the next scan compare against
// stale cache state instead of the last complete baseline.
if err := ctx.Err(); err != nil {
work.withdraw(err)
return nil
}
prevSet := make(map[string]bool, len(previousEntries))
for _, e := range previousEntries {
prevSet[e] = true
}
// The baseline only tracks new paths. Active findings also need a current
// verdict before this scan can retire them, even on a cached walk. Two of
// the owned names are shared with the content scan, so a re-verified path
// may carry a finding this check never raised.
activeChecks := make(map[string]map[string]bool)
if st != nil {
for _, f := range st.LatestFindings() {
for _, name := range runnerFindingNames["file_index"] {
if f.Check == name && f.FilePath != "" {
if activeChecks[f.FilePath] == nil {
activeChecks[f.FilePath] = make(map[string]bool)
}
activeChecks[f.FilePath][f.Check] = true
break
}
}
}
}
newSet := make(map[string]bool)
var newFiles, activeFiles, scanFiles []string
for _, e := range currentEntries {
if !prevSet[e] {
newSet[e] = true
newFiles = append(newFiles, e)
}
if activeChecks[e] != nil {
activeFiles = append(activeFiles, e)
}
if !prevSet[e] || activeChecks[e] != nil {
scanFiles = append(scanFiles, e)
}
}
// Every finding this check raises asserts the file is new. A path is only
// re-verified because it already carries one, so the re-verification may
// renew that finding but must never raise a different one about a file
// that has been in the baseline all along -- which would also keep itself
// alive, by putting the path back into this set on the next cycle.
keepReverified := func(findings []alert.Finding) []alert.Finding {
out := findings[:0]
for _, f := range findings {
if !newSet[f.FilePath] && activeChecks[f.FilePath] != nil && !activeChecks[f.FilePath][f.Check] {
continue
}
out = append(out, f)
}
return out
}
if incomplete {
// Publishing partial entries or mtimes would let the next cached
// walk retire findings for directories we still have not read.
return keepReverified(checkFileIndexAnalyzeNewFiles(ctx, cfg, scanFiles))
}
work.local()
// Write current index (atomic)
work.observe(writeIndex(currentPath, currentEntries))
// First run - save baseline
if _, err := osFS.Stat(previousPath); os.IsNotExist(err) {
work.observe(copyFile(currentPath, previousPath))
work.observe(saveDirCache(indexDir, dirCache))
work.execution()
return keepReverified(checkFileIndexAnalyzeNewFiles(ctx, cfg, activeFiles))
}
isShrink, promote := evaluateFileIndexShrink(len(previousEntries), len(currentEntries))
if isShrink && !promote {
// The preserved baseline still contains removed paths. Its entries
// cannot be carried forward using mtimes from the smaller current
// walk, or the next cached scan resurrects those paths and resets
// the shrink streak. Rewalk until the smaller baseline is adopted.
dirCache = make(dirMtimeCache)
}
work.observe(saveDirCache(indexDir, dirCache))
// A large shrink (mass deletion, a WP install removed, or a transient read
// failure) must not instantly flush the baseline: promoting an empty or
// much-smaller index would make every surviving file look "new" next cycle
// and flood alerts. Still, entries that are present in the shrunken current
// index but absent from the old baseline are genuinely new and must be
// classified now. Non-promoting shrink cycles merge those paths into the old
// baseline so they do not alert repeatedly while the deletion guard is still
// preserving the removed paths against recovery floods.
if isShrink {
work.execution()
findings := keepReverified(checkFileIndexAnalyzeNewFiles(ctx, cfg, scanFiles))
work.local()
if promote {
fmt.Fprintf(os.Stderr, "file_index: shrink persisted %d scans; adopting smaller index (%d entries, was %d) as new baseline\n",
fileIndexShrinkPromoteThreshold, len(currentEntries), len(previousEntries))
work.observe(copyFile(currentPath, previousPath))
} else {
if len(newFiles) > 0 {
work.observe(writeIndex(previousPath, mergeIndexEntries(previousEntries, newFiles)))
}
fmt.Fprintf(os.Stderr, "file_index: current index (%d) shrank vs previous (%d); preserving prior baseline this cycle\n",
len(currentEntries), len(previousEntries))
}
return findings
}
work.execution()
findings := keepReverified(checkFileIndexAnalyzeNewFiles(ctx, cfg, scanFiles))
work.local()
work.observe(copyFile(currentPath, previousPath))
return findings
}
func mergeIndexEntries(base, extra []string) []string {
seen := make(map[string]bool, len(base)+len(extra))
merged := make([]string, 0, len(base)+len(extra))
for _, e := range base {
if seen[e] {
continue
}
seen[e] = true
merged = append(merged, e)
}
for _, e := range extra {
if seen[e] {
continue
}
seen[e] = true
merged = append(merged, e)
}
sort.Strings(merged)
return merged
}
// checkFileIndexAnalyzeNewFiles classifies a list of newly-detected file paths
// and returns alert.Finding values for those that match known threat patterns.
// Called by both the normal incremental path and the audit (ForceFileIndex) path.
func checkFileIndexAnalyzeNewFiles(ctx context.Context, cfg *config.Config, newFiles []string) []alert.Finding {
var findings []alert.Finding
for _, path := range newFiles {
name := filepath.Base(path)
nameLower := strings.ToLower(name)
// Bypassed for explicit full-scan / audit requests.
suppressed := false
if scanRespectsIgnores(ctx, cfg) {
for _, ignore := range cfg.Suppressions.IgnorePaths {
if matchGlob(path, ignore) {
suppressed = true
break
}
}
}
if suppressed {
continue
}
severity := alert.Severity(-1)
check := ""
message := ""
contentSHA256 := ""
if strings.Contains(path, "/wp-content/uploads/") && phpPathExecutes(path, nameLower) {
// Content decides, never the path or name. A negative
// severity is a content-verified inert stub (e.g. the
// WordPress "silence is golden" index.php, or BackWPup's
// "<?php //<json>" working files) and is suppressed; any
// real code surfaces, malicious or merely present.
sev, ck, msg, hash, readOK := classifyUploadPHPWithFingerprint(path)
if !readOK {
recordCoverageGapPaths(ctx, "file_index", coveragePathAliases(path))
reportFileIndexFailure(ctx)
}
if sev >= 0 {
severity = sev
check = ck
message = msg
contentSHA256 = hash
} else {
continue
}
}
// PHP files in wp-content/languages and wp-content/upgrade: content-first.
// Path-only Critical buried real alerts under location noise (WPML
// translation queues, WP auto-update staging). See classifySensitiveDirPHP.
sev, ck, msg, hash, readOK := classifySensitiveDirPHPWithFingerprint(path, name)
if !readOK {
recordCoverageGapPaths(ctx, "file_index", coveragePathAliases(path))
reportFileIndexFailure(ctx)
}
if sev >= 0 {
severity = sev
check = ck
message = msg
contentSHA256 = hash
}
if strings.Contains(path, "/.config/") {
severity = alert.Critical
check = "new_executable_in_config"
message = fmt.Sprintf("New executable in .config: %s", path)
}
if isWebshellName(nameLower) {
severity = alert.Critical
check = "new_webshell_file"
message = fmt.Sprintf("New file with webshell name: %s", path)
}
if isExecutablePHPName(nameLower) && isSuspiciousPHPName(nameLower) {
if severity < 0 {
severity = alert.High
check = "new_suspicious_php"
message = fmt.Sprintf("New suspicious PHP file: %s", path)
}
}
if severity >= 0 {
details := ""
if info, err := osFS.Stat(path); err == nil {
details = fmt.Sprintf("Size: %d, Mtime: %s", info.Size(), info.ModTime().Format("2006-01-02 15:04:05"))
}
f := alert.Finding{
Severity: severity,
Check: check,
Message: message,
Details: details,
FilePath: path,
}
if IsContentReverifiable(f.Check) {
f.ContentSHA256 = contentSHA256
f.DetectLogic = ContentDetectionVersion()
}
findings = append(findings, f)
}
}
return findings
}
// classifySensitiveDirPHP returns (severity, check, message) for a PHP file
// in /wp-content/languages/ or /wp-content/upgrade/. Returns a negative
// severity when the path is not in a sensitive dir.
//
// Every PHP file in a sensitive dir is content-analysed -- there is no
// filename allowlist, so an attacker cannot hide a backdoor by naming it like
// a translation or index file. A real indicator keeps Critical severity and
// the content-based check name (obfuscated_php / suspicious_php_content, both
// already wired into autoresponse, remediate, correlation, and attackdb). A
// clean file is demoted to Warning with check new_php_in_sensitive_dir_clean,
// which is intentionally NOT in any of those maps -- a clean file is a
// visibility signal, not an attack. Mirrors the realtime path at fanotify.go.
func classifySensitiveDirPHP(path, name string) (alert.Severity, string, string) {
sev, check, message, _, _ := classifySensitiveDirPHPWithFingerprint(path, name)
return sev, check, message
}
func classifySensitiveDirPHPWithFingerprint(path, name string) (alert.Severity, string, string, string, bool) {
nameLower := strings.ToLower(name)
if !phpPathExecutes(path, nameLower) {
return -1, "", "", "", true
}
isLanguages := strings.Contains(path, "/wp-content/languages/")
isUpgrade := strings.Contains(path, "/wp-content/upgrade/")
if !isLanguages && !isUpgrade {
return -1, "", "", "", true
}
locLabel := "wp-content/languages"
if isUpgrade {
locLabel = "wp-content/upgrade"
}
result, contentSHA256 := analyzePHPContentWithFingerprint(path)
if result.severity >= 0 {
return result.severity, result.check, fmt.Sprintf("%s: %s", result.message, path), contentSHA256, result.readOK
}
// Fail closed: an unreadable body (attacker racing the scanner with rm or
// chmod 000) must not be demoted to a clean Warning. Mirrors classifyUploadPHP.
if !result.readOK {
return alert.High, "new_php_in_sensitive_dir",
fmt.Sprintf("New unreadable PHP file in %s: %s", locLabel, path), "", false
}
// A zero-byte body reaching here was verified stable across the read (a
// truncation under the scanner fails the post-read stat and is handled
// above), so it holds no code and is visibility, not an attack.
if result.empty {
return alert.Warning, "new_php_in_sensitive_dir_clean",
fmt.Sprintf("New empty PHP file in %s (no content): %s", locLabel, path), "", true
}
// Content-verified inert stub (e.g. the "silence is golden" index.php) is
// suppressed; any real code surfaces as a non-actionable visibility Warning.
if IsBenignPHPStub(path) {
return -1, "", "", "", true
}
// WordPress 6.5+ auto-generates *.l10n.php translation caches as pure data
// return arrays. Recognized by content structure (not filename), they carry
// no executable construct, so suppress rather than warn on every locale file.
if isWPTranslationCache(path) {
return -1, "", "", "", true
}
return alert.Warning, "new_php_in_sensitive_dir_clean",
fmt.Sprintf("New PHP file in %s (content clean): %s", locLabel, path), "", true
}
// groupEntriesByUploadDir groups index entries by each ancestor directory.
// Used to carry forward the whole cached subtree when an unchanged directory
// is skipped before walking into its children.
func groupEntriesByUploadDir(entries []string) map[string][]string {
grouped := make(map[string][]string)
for _, path := range entries {
for dir := filepath.Dir(path); dir != "." && dir != string(filepath.Separator); dir = filepath.Dir(dir) {
grouped[dir] = append(grouped[dir], path)
}
}
return grouped
}
// buildFileIndex scans targeted directory subtrees using ReadDir.
// Skips directories whose mtime hasn't changed - carries forward
// their entries from the previous index instead.
// If forceFullScan is true, all directories are re-scanned regardless of mtime.
func buildFileIndex(ctx context.Context, dirCache dirMtimeCache, prevByDir map[string][]string, forceFullScan bool) []string {
var entries []string
if ctx == nil {
ctx = context.Background()
}
if ctx.Err() != nil {
return nil
}
homeDirs, err := fileIndexHomeDirs(ctx)
if err != nil {
markCheckIncomplete(ctx, "file_index")
}
var subtreeChanges *subtreeChangeTracker
if !forceFullScan {
subtreeChanges = newSubtreeChangeTracker(dirCache)
}
for _, homeEntry := range homeDirs {
// A cancelled scan must stop walking and let the caller discard the
// partial index rather than promote it as the new baseline.
if ctx.Err() != nil {
return entries
}
if !homeEntry.IsDir() {
continue
}
homeDir := scanHomeDirPath(homeEntry)
// Scan wp-content/uploads for PHP files
uploadDirs := []string{
filepath.Join(homeDir, "public_html", "wp-content", "uploads"),
}
// Scan directories that shouldn't normally contain user PHP:
// languages (translation files only), upgrade (temp dir), mu-plugins
sensitiveWPDirs := []string{
filepath.Join(homeDir, "public_html", "wp-content", "languages"),
filepath.Join(homeDir, "public_html", "wp-content", "upgrade"),
filepath.Join(homeDir, "public_html", "wp-content", "mu-plugins"),
}
subDirs, err := osFS.ReadDir(homeDir)
if err != nil && !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
for _, sd := range subDirs {
if ctx.Err() != nil {
return entries
}
if sd.IsDir() && sd.Name() != "public_html" && sd.Name() != "mail" &&
!strings.HasPrefix(sd.Name(), ".") && sd.Name() != "etc" &&
sd.Name() != "logs" && sd.Name() != "ssl" && sd.Name() != "tmp" {
uploadsPath := filepath.Join(homeDir, sd.Name(), "wp-content", "uploads")
if info, err := osFS.Stat(uploadsPath); err == nil && info.IsDir() {
uploadDirs = append(uploadDirs, uploadsPath)
} else if err != nil && !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
// Also track sensitive dirs for addon domains
for _, subDir := range []string{"languages", "upgrade", "mu-plugins"} {
if ctx.Err() != nil {
return entries
}
sensitiveDir := filepath.Join(homeDir, sd.Name(), "wp-content", subDir)
if info, err := osFS.Stat(sensitiveDir); err == nil && info.IsDir() {
sensitiveWPDirs = append(sensitiveWPDirs, sensitiveDir)
} else if err != nil && !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
}
}
}
for _, uploadsDir := range uploadDirs {
if ctx.Err() != nil {
return entries
}
scanDirForPHPContextWithTracker(ctx, uploadsDir, 6, dirCache, prevByDir, forceFullScan, phpHandlerOverlay{}, subtreeChanges, &entries)
}
// Scan sensitive WP directories for any PHP files
for _, sensitiveDir := range sensitiveWPDirs {
if ctx.Err() != nil {
return entries
}
scanDirForPHPContextWithTracker(ctx, sensitiveDir, 4, dirCache, prevByDir, forceFullScan, phpHandlerOverlay{}, subtreeChanges, &entries)
}
// Scan .config for executables
configDir := filepath.Join(homeDir, ".config")
if ctx.Err() != nil {
return entries
}
scanDirForExecutablesContextWithTracker(ctx, configDir, 3, dirCache, prevByDir, forceFullScan, subtreeChanges, &entries)
}
if AccountFromContext(ctx) == "" {
// Scan tmp dirs
for _, tmpDir := range []string{"/tmp", "/dev/shm", "/var/tmp"} {
if ctx.Err() != nil {
return entries
}
scanDirForSuspiciousExtContextWithTracker(ctx, tmpDir, 2, dirCache, prevByDir, forceFullScan, subtreeChanges, &entries)
}
}
sort.Strings(entries)
return entries
}
func fileIndexHomeDirs(ctx context.Context) ([]os.DirEntry, error) {
if AccountFromContext(ctx) != "" {
return GetScanHomeDirs(ctx)
}
homes, err := readAccountHomes()
entries := make([]os.DirEntry, 0, len(homes))
for _, home := range homes {
entries = append(entries, rootedDirEntry{DirEntry: home.Entry, root: home.Root})
}
return entries, err
}
// scanDirForPHP recursively reads directories for PHP-executable files.
// If directory mtime is unchanged, carries forward previous entries.
func scanDirForPHP(dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, overlay phpHandlerOverlay, entries *[]string) {
scanDirForPHPContext(context.Background(), dir, maxDepth, cache, prev, forceFullScan, overlay, entries)
}
func scanDirForPHPContext(ctx context.Context, dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, overlay phpHandlerOverlay, entries *[]string) {
var subtreeChanges *subtreeChangeTracker
if !forceFullScan {
subtreeChanges = newSubtreeChangeTracker(cache)
}
scanDirForPHPContextWithTracker(ctx, dir, maxDepth, cache, prev, forceFullScan, overlay, subtreeChanges, entries)
}
func scanDirForPHPContextWithTracker(ctx context.Context, dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, overlay phpHandlerOverlay, subtreeChanges *subtreeChangeTracker, entries *[]string) {
if maxDepth <= 0 || ctx.Err() != nil {
return
}
if htaccess, err := osFS.ReadFile(filepath.Join(dir, ".htaccess")); err == nil {
overlay = overlay.mergeHtaccess(htaccess)
} else if !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
if ctx.Err() != nil {
return
}
changed := dirChanged(dir, cache, forceFullScan || overlay.active())
if ctx.Err() != nil {
return
}
subtreeChanged := false
if !changed {
subtreeChanged = subtreeChanges.hasChangedDir(ctx, dir)
}
if ctx.Err() != nil {
return
}
if !changed && !subtreeChanged {
// dir's mtime is unchanged and no cached subdirectory changed, so the
// whole subtree is stable: carry the previous entries forward without
// re-reading disk.
*entries = append(*entries, prev[dir]...)
return
}
// Either dir changed, or a file was dropped deep in the subtree (bumping
// only a subdirectory's mtime). Walk it so the new file is indexed now
// instead of hiding until the periodic forced full scan.
dirEntries, err := osFS.ReadDir(dir)
if err != nil {
if !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
return
}
for _, entry := range dirEntries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
if entry.IsDir() {
scanDirForPHPContextWithTracker(ctx, fullPath, maxDepth-1, cache, prev, forceFullScan, overlay, subtreeChanges, entries)
continue
}
nameLower := strings.ToLower(name)
// index.php is indexed too: a webshell named index.php must not hide
// behind the WordPress silence-stub convention. The inert stub itself
// is suppressed later by content analysis. All PHP-executable
// extensions are indexed, not just .php, so a .phtml/.php7 backdoor
// cannot dodge the index by extension. The two predicates overlap
// (the PHP family is also "suspicious"), so index each file once.
if overlay.executes(nameLower) || suspiciousExtensions[filepath.Ext(nameLower)] {
*entries = append(*entries, fullPath)
}
}
}
// scanDirForExecutables reads .config dirs for executable files.
func scanDirForExecutables(dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, entries *[]string) {
scanDirForExecutablesContext(context.Background(), dir, maxDepth, cache, prev, forceFullScan, entries)
}
func scanDirForExecutablesContext(ctx context.Context, dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, entries *[]string) {
var subtreeChanges *subtreeChangeTracker
if !forceFullScan {
subtreeChanges = newSubtreeChangeTracker(cache)
}
scanDirForExecutablesContextWithTracker(ctx, dir, maxDepth, cache, prev, forceFullScan, subtreeChanges, entries)
}
func scanDirForExecutablesContextWithTracker(ctx context.Context, dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, subtreeChanges *subtreeChangeTracker, entries *[]string) {
if maxDepth <= 0 || ctx.Err() != nil {
return
}
changed := dirChanged(dir, cache, forceFullScan)
if ctx.Err() != nil {
return
}
subtreeChanged := false
if !changed {
subtreeChanged = subtreeChanges.hasChangedDir(ctx, dir)
}
if ctx.Err() != nil {
return
}
if !changed && !subtreeChanged {
*entries = append(*entries, prev[dir]...)
return
}
dirEntries, err := osFS.ReadDir(dir)
if err != nil {
if !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
return
}
for _, entry := range dirEntries {
if ctx.Err() != nil {
return
}
fullPath := filepath.Join(dir, entry.Name())
if entry.IsDir() {
scanDirForExecutablesContextWithTracker(ctx, fullPath, maxDepth-1, cache, prev, forceFullScan, subtreeChanges, entries)
continue
}
info, err := entry.Info()
if err != nil {
if !os.IsNotExist(err) {
reportFileIndexFailure(ctx)
}
continue
}
if info.Mode()&0111 != 0 {
*entries = append(*entries, fullPath)
}
}
}
// scanDirForSuspiciousExt reads tmp dirs for files with suspicious extensions.
func scanDirForSuspiciousExt(dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, entries *[]string) {
scanDirForSuspiciousExtContext(context.Background(), dir, maxDepth, cache, prev, forceFullScan, entries)
}
func scanDirForSuspiciousExtContext(ctx context.Context, dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, entries *[]string) {
var subtreeChanges *subtreeChangeTracker
if !forceFullScan {
subtreeChanges = newSubtreeChangeTracker(cache)
}
scanDirForSuspiciousExtContextWithTracker(ctx, dir, maxDepth, cache, prev, forceFullScan, subtreeChanges, entries)
}
func scanDirForSuspiciousExtContextWithTracker(ctx context.Context, dir string, maxDepth int, cache dirMtimeCache, prev map[string][]string, forceFullScan bool, subtreeChanges *subtreeChangeTracker, entries *[]string) {
if maxDepth <= 0 || ctx.Err() != nil {
return
}
changed := dirChanged(dir, cache, forceFullScan)
if ctx.Err() != nil {
return
}
subtreeChanged := false
if !changed {
subtreeChanged = subtreeChanges.hasChangedDir(ctx, dir)
}
if ctx.Err() != nil {
return
}
if !changed && !subtreeChanged {
*entries = append(*entries, prev[dir]...)
return
}
dirEntries, err := osFS.ReadDir(dir)
if err != nil {
if !os.IsNotExist(err) {
markCheckIncomplete(ctx, "file_index")
}
return
}
for _, entry := range dirEntries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
if entry.IsDir() {
scanDirForSuspiciousExtContextWithTracker(ctx, fullPath, maxDepth-1, cache, prev, forceFullScan, subtreeChanges, entries)
continue
}
ext := filepath.Ext(strings.ToLower(name))
if suspiciousExtensions[ext] {
*entries = append(*entries, fullPath)
}
}
}
func isWebshellName(name string) bool {
webshells := map[string]bool{
"h4x0r.php": true, "c99.php": true, "r57.php": true,
"wso.php": true, "alfa.php": true, "b374k.php": true,
"shell.php": true, "cmd.php": true, "backdoor.php": true,
"webshell.php": true, "hack.php": true, "0x.php": true,
"up.php": true, "uploader.php": true, "filemanager.php": true,
}
return webshells[name]
}
func isSuspiciousPHPName(name string) bool {
suspicious := []string{
"shell", "cmd", "exec", "hack", "backdoor", "upload",
"exploit", "reverse", "connect", "proxy", "tunnel",
"0x", "x0", "eval", "assert", "passthru",
}
for _, s := range suspicious {
if strings.Contains(name, s) {
return true
}
}
nameNoExt := strings.TrimSuffix(name, ".php")
if len(nameNoExt) <= 5 && strings.ContainsAny(nameNoExt, "0123456789") {
return true
}
return false
}
func writeIndex(path string, entries []string) error {
tmpPath := path + ".tmp"
// #nosec G304 -- path is filepath.Join under operator-configured StatePath.
f, err := os.Create(tmpPath)
if err != nil {
return err
}
defer func() { _ = f.Close() }()
w := bufio.NewWriter(f)
for _, e := range entries {
if _, err := w.WriteString(e + "\n"); err != nil {
return err
}
}
if err := w.Flush(); err != nil {
return err
}
if err := f.Close(); err != nil {
return err
}
return os.Rename(tmpPath, path)
}
func loadIndex(path string) ([]string, error) {
f, err := osFS.Open(path)
if os.IsNotExist(err) {
return nil, nil
}
if err != nil {
return nil, err
}
defer func() { _ = f.Close() }()
var entries []string
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 0, 1024*1024), 1024*1024)
for scanner.Scan() {
line := scanner.Text()
if line != "" {
entries = append(entries, line)
}
}
return entries, scanner.Err()
}
func copyFile(src, dst string) error {
data, err := osFS.ReadFile(src)
if err != nil {
return err
}
return os.WriteFile(dst, data, 0600)
}
package checks
import (
"context"
"errors"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
)
var fileIndexQueues = newFileIndexQueue()
type fileIndexWorkKey struct{}
// Content and metadata failures retain existing findings and baseline policy.
// They still belong to the live scan, even when a walker can continue.
func reportFileIndexFailure(ctx context.Context) {
if work, ok := ctx.Value(fileIndexWorkKey{}).(*fileIndexWork); ok {
work.fail()
}
}
type fileIndexQueueMonitor struct {
mu sync.Mutex
waiting map[*fileIndexWork]struct{}
active *fileIndexWork
idleSince time.Time
waitingLoss, activeLoss *queuehealth.Tracker
}
type fileIndexWork struct {
queue *fileIndexQueueMonitor
at, deadline, executionDeadline time.Time
running, failed bool
}
func newFileIndexQueue() *fileIndexQueueMonitor {
return &fileIndexQueueMonitor{
waiting: make(map[*fileIndexWork]struct{}),
waitingLoss: queuehealth.New(0, time.Minute),
activeLoss: queuehealth.New(1, time.Minute),
}
}
func (q *fileIndexQueueMonitor) run(ctx context.Context, fn func(*fileIndexWork) []alert.Finding) []alert.Finding {
if ctx.Err() != nil {
return nil
}
now := time.Now()
budget := timeoutFor("file_index")
parent, hasParent := ctx.Deadline()
deadline := now.Add(budget)
if hasParent && parent.Before(deadline) {
deadline = parent
}
w := &fileIndexWork{queue: q, at: now, deadline: deadline}
q.mu.Lock()
if q.active == nil && len(q.waiting) == 0 {
q.idleSince = now
}
q.waiting[w] = struct{}{}
q.mu.Unlock()
completed := false
defer func() { w.finish(completed) }()
select {
case fileIndexLiveScanGate <- struct{}{}:
now = time.Now()
executionDeadline := now.Add(budget)
if hasParent && parent.Before(executionDeadline) {
executionDeadline = parent
}
q.mu.Lock()
delete(q.waiting, w)
q.active = w
q.idleSince = time.Time{}
w.running = true
w.at = now
w.deadline = now.Add(time.Minute)
w.executionDeadline = executionDeadline
q.mu.Unlock()
case <-ctx.Done():
w.withdraw(ctx.Err())
completed = true
return nil
}
findings := fn(w)
completed = true
return findings
}
func (w *fileIndexWork) execution() {
w.queue.mu.Lock()
w.at = time.Now()
w.deadline = w.executionDeadline
w.queue.mu.Unlock()
}
func (w *fileIndexWork) local() {
w.queue.mu.Lock()
w.at = time.Now()
w.deadline = w.at.Add(time.Minute)
w.queue.mu.Unlock()
}
func (w *fileIndexWork) failLocked() {
if w.failed {
return
}
w.failed = true
losses := w.queue.waitingLoss
if w.running {
losses = w.queue.activeLoss
}
losses.Lose(time.Now(), 1)
}
func (w *fileIndexWork) fail() {
w.queue.mu.Lock()
w.failLocked()
w.queue.mu.Unlock()
}
func (w *fileIndexWork) observe(err error) {
if err != nil {
w.fail()
}
}
func (w *fileIndexWork) withdraw(err error) {
if errors.Is(err, context.Canceled) {
return
}
w.queue.mu.Lock()
w.failLocked()
w.queue.mu.Unlock()
}
func (w *fileIndexWork) finish(completed bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
if !completed {
w.failLocked()
}
if w.running {
q.active = nil
q.idleSince = time.Now()
// The scan owns this token. Release its health owner and token under
// one lock so its successor cannot publish before this owner leaves.
<-fileIndexLiveScanGate
} else {
delete(q.waiting, w)
}
}
// FileIndexQueueStatuses reads memory without waiting for a scan or filesystem.
func FileIndexQueueStatuses(now time.Time) map[string]queuehealth.Status {
q := fileIndexQueues
q.mu.Lock()
defer q.mu.Unlock()
waiting := q.waitingLoss.Snapshot(now)
waiting.CapacityUnavailable = true
active := q.activeLoss.Snapshot(now)
for w := range q.waiting {
waiting.Depth++
waiting.LagSeconds = max(waiting.LagSeconds, now.Sub(w.at).Seconds())
if !now.Before(w.deadline) || (q.active == nil && now.Sub(q.idleSince) >= time.Minute) {
waiting.Status, waiting.Reason = "degraded", "backlog_lag"
}
}
if w := q.active; w != nil {
active.InFlight = 1
active.ProcessingSeconds = max(0, now.Sub(w.at).Seconds())
if !now.Before(w.deadline) {
active.Status, active.Reason = "degraded", "processing_lag"
}
}
return map[string]queuehealth.Status{"waiting": waiting, "active": active}
}
package checks
import (
"fmt"
"github.com/pidginhost/csm/internal/alert"
)
// classifyUploadPHP decides severity, check name, and message for a fresh PHP
// file under wp-content/uploads using its CONTENT, never its path or name.
// Uploads should hold media, not PHP, so any new PHP is at least a visibility
// signal; the body decides whether it is an attack.
//
// A negative severity means "suppress" (a content-verified inert stub). It
// mirrors classifySensitiveDirPHP so the two anomalous-PHP-location detectors
// behave identically. Path/name allowlists are intentionally absent: skipping
// a file because it sits under /cache/ or is named index.php is exactly how an
// attacker hides a webshell in a "safe" location.
//
// An unreadable body fails closed at High: an attacker who races the scanner
// with `rm` or chmod 000 must not earn a demote. A body truncated mid-read
// fails the post-read stat comparison and lands in that same branch, so a
// zero-byte result carries proof the file really was empty for the whole read.
// Empty means no code, which cannot execute, so it joins content-clean
// real-code files as a non-actionable Warning under a check name that is
// intentionally absent from the correlation and auto-response maps -- neither
// is an attack, and rating them High buries the findings that are.
func classifyUploadPHP(path string) (alert.Severity, string, string) {
sev, check, message, _, _ := classifyUploadPHPWithFingerprint(path)
return sev, check, message
}
func classifyUploadPHPWithFingerprint(path string) (alert.Severity, string, string, string, bool) {
r, contentSHA256 := analyzePHPContentWithFingerprint(path)
if r.severity >= 0 {
return r.severity, r.check, fmt.Sprintf("%s: %s", r.message, path), contentSHA256, r.readOK
}
if !r.readOK {
return alert.High, "new_php_in_uploads", fmt.Sprintf("New unreadable PHP file in uploads: %s", path), "", false
}
if r.empty {
return alert.Warning, "new_php_in_uploads_clean", fmt.Sprintf("New empty PHP file in uploads (no content): %s", path), "", true
}
if IsBenignPHPStub(path) {
return -1, "", "", "", true
}
return alert.Warning, "new_php_in_uploads_clean", fmt.Sprintf("New PHP file in uploads (content clean): %s", path), "", true
}
package checks
import (
"context"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// CheckFilesystem uses globs and targeted ReadDir to check for backdoors,
// hidden files, and SUID binaries. No `find` command needed.
func CheckFilesystem(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
// GSocket / backdoor binaries in .config dirs - glob (instant).
// Rank by mtime desc so recently-touched accounts process first
// when the check timeout cuts iteration short.
backdoorNames := map[string]bool{
"defunct": true, "defunct.dat": true, "gs-netcat": true,
"gs-sftp": true, "gs-mount": true, "gsocket": true,
}
configGlobs := [][]string{
{".config", "htop", "*"},
{".config", "*", "*"},
}
// The htop glob is a subset of the wider .config glob. Deduplicate
// before ranking so the per-account cap applies to this scanner once.
configCandidates := make([]string, 0)
seenConfigCandidate := make(map[string]struct{})
for _, pattern := range configGlobs {
if ctx.Err() != nil {
return findings
}
matches, err := homeGlob(ctx, pattern...)
markScanReadError(ctx, "filesystem", err)
for _, path := range matches {
if ctx.Err() != nil {
return findings
}
if backdoorNames[filepath.Base(path)] {
if _, seen := seenConfigCandidate[path]; seen {
continue
}
seenConfigCandidate[path] = struct{}{}
configCandidates = append(configCandidates, path)
}
}
}
rankedConfigCandidates := rankPathsByMtimeDesc(ctx, configCandidates, accountScanMaxFiles(ctx, cfg))
if len(rankedConfigCandidates) < len(configCandidates) {
markCheckIncomplete(ctx, "filesystem")
}
if ctx.Err() != nil {
return findings
}
for _, path := range rankedConfigCandidates {
if ctx.Err() != nil {
return findings
}
info, _ := osFS.Stat(path)
var details string
if info != nil {
details = fmt.Sprintf("Size: %d bytes, Mtime: %s", info.Size(), info.ModTime().Format("2006-01-02 15:04:05"))
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "backdoor_binary",
Message: fmt.Sprintf("Backdoor binary found: %s", path),
Details: details,
FilePath: path,
})
}
if AccountFromContext(ctx) == "" {
// Hidden files in /tmp, /dev/shm, /var/tmp - glob (instant)
safeHiddenPrefixes := []string{
".s.PGSQL", ".font-unix", ".ICE-unix", ".X11-unix",
".XIM-unix", ".crontab.", ".Test-unix",
}
// One candidate set across all three roots: on CloudLinux /var/tmp is
// the same filesystem as /tmp, so the same physical file is reachable
// through two of these patterns and was reported once per pattern.
var candidates []string
for _, pattern := range []string{"/tmp/.*", "/dev/shm/.*", "/var/tmp/.*"} {
if ctx.Err() != nil {
return findings
}
matches, err := osFS.Glob(pattern)
markScanReadError(ctx, "filesystem", err)
for _, match := range matches {
if ctx.Err() != nil {
return findings
}
base := filepath.Base(match)
safe := false
for _, prefix := range safeHiddenPrefixes {
if strings.HasPrefix(base, prefix) {
safe = true
break
}
}
if safe {
continue
}
candidates = append(candidates, match)
}
}
// These are global temp locations, not account paths; do not let
// account_scan_max_files hide older suspicious files here.
ranked := rankPathsByMtimeDesc(ctx, candidates, 0)
if ctx.Err() != nil {
return findings
}
var reported []os.FileInfo
for _, match := range ranked {
if ctx.Err() != nil {
return findings
}
info, err := osFS.Stat(match)
markScanReadError(ctx, "filesystem", err)
if err != nil || info.IsDir() {
continue
}
// A leading dot is not by itself a signal: these directories are
// full of root-owned infrastructure state. What matters is whether
// the file could execute.
canExecute, err := hiddenTempFileCanExecute(match, info)
markScanReadError(ctx, "filesystem", err)
if !canExecute {
continue
}
seen := false
for _, prev := range reported {
if os.SameFile(prev, info) {
seen = true
break
}
}
if seen {
continue
}
reported = append(reported, info)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "suspicious_file",
Message: fmt.Sprintf("Suspicious hidden file: %s", match),
Details: fmt.Sprintf("Size: %d, Mtime: %s", info.Size(), info.ModTime()),
FilePath: match,
})
}
// SUID binaries in tmp dirs - ReadDir + stat (small dirs, fast)
for _, dir := range []string{"/tmp", "/var/tmp", "/dev/shm"} {
if ctx.Err() != nil {
return findings
}
scanForSUID(ctx, dir, 3, &findings)
}
}
// SUID in /home - shallow scan only
if ctx.Err() != nil {
return findings
}
homeDirs := scanHomeDirsWithCoverage(ctx, "filesystem")
for _, entry := range homeDirs {
if ctx.Err() != nil {
return findings
}
if !entry.IsDir() {
continue
}
scanForSUID(ctx, scanHomeDirPath(entry), 3, &findings)
}
return findings
}
// hiddenTempFileCanExecute reports whether a hidden file in a world-writable
// temp directory could run: an executable bit or ELF magic, or a script marker
// that an interpreter would honour. Inert data written there by system
// components is not a finding.
func hiddenTempFileCanExecute(path string, info os.FileInfo) (bool, error) {
if !info.Mode().IsRegular() {
return false, nil
}
if info.Mode()&0o111 != 0 {
return true, nil
}
// Stat alone cannot protect a world-writable path: it can be replaced with
// a FIFO before open. Use the provider's nonblocking, fd-verified opener
// and bind the content decision to the same inode used for deduplication.
opener, ok := osFS.(interface {
openRegularFile(string, int) (*os.File, error)
})
if !ok {
return false, fmt.Errorf("filesystem provider cannot safely open %s", path)
}
f, err := opener.openRegularFile(path, 0)
if err != nil {
return false, err
}
defer func() { _ = f.Close() }()
opened, err := f.Stat()
if err != nil {
return false, err
}
if !sameFileSnapshot(info, opened) {
return false, errFileChanged
}
var head [8]byte
n, err := io.ReadFull(f, head[:])
if err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, io.ErrUnexpectedEOF) {
return false, err
}
after, err := f.Stat()
if err != nil {
return false, err
}
if !sameFileSnapshot(opened, after) {
return false, errFileChanged
}
prefix := strings.ToLower(string(head[:n]))
return strings.HasPrefix(string(head[:n]), "\x7fELF") ||
strings.HasPrefix(prefix, "#!") ||
strings.HasPrefix(prefix, "<?php") ||
strings.HasPrefix(prefix, "<?="), nil
}
// scanForSUID checks for SUID binaries using ReadDir.
func scanForSUID(ctx context.Context, dir string, maxDepth int, findings *[]alert.Finding) {
if ctx.Err() != nil {
return
}
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
markScanReadError(ctx, "filesystem", err)
return
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
fullPath := filepath.Join(dir, entry.Name())
if entry.IsDir() {
// Skip virtfs and known large dirs
if entry.Name() == "virtfs" || entry.Name() == "mail" || entry.Name() == "public_html" {
continue
}
scanForSUID(ctx, fullPath, maxDepth-1, findings)
continue
}
info, err := entry.Info()
if err != nil {
markScanReadError(ctx, "filesystem", err)
continue
}
if info.Mode()&os.ModeSetuid != 0 {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "suid_binary",
Message: fmt.Sprintf("SUID binary in unusual location: %s", fullPath),
Details: fmt.Sprintf("Mode: %s, Size: %d", info.Mode(), info.Size()),
FilePath: fullPath,
})
}
}
}
// CheckWebshells uses pure Go ReadDir to scan for known webshell files
// and directories. No `find` command needed.
func CheckWebshells(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
webshellNames := map[string]bool{
"h4x0r.php": true, "c99.php": true, "r57.php": true,
"wso.php": true, "alfa.php": true, "b374k.php": true,
"mini.php": true, "adminer.php": true,
}
webshellDirs := map[string]bool{
"LEVIATHAN": true, "haxorcgiapi": true,
}
// Scan each user's public_html and addon domains
homeDirs := scanHomeDirsWithCoverage(ctx, "webshells")
for _, homeEntry := range homeDirs {
if ctx.Err() != nil {
return findings
}
if !homeEntry.IsDir() {
continue
}
homeDir := scanHomeDirPath(homeEntry)
// Get all potential document roots
docRoots := []string{filepath.Join(homeDir, "public_html")}
subDirs, err := osFS.ReadDir(homeDir)
markScanReadError(ctx, "webshells", err)
for _, sd := range subDirs {
if sd.IsDir() && sd.Name() != "public_html" && sd.Name() != "mail" &&
!strings.HasPrefix(sd.Name(), ".") && sd.Name() != "etc" &&
sd.Name() != "logs" && sd.Name() != "ssl" && sd.Name() != "tmp" {
docRoots = append(docRoots, filepath.Join(homeDir, sd.Name()))
}
}
for _, docRoot := range docRoots {
scanForWebshells(ctx, docRoot, 8, webshellNames, webshellDirs, cfg, &findings)
if ctx.Err() != nil {
return findings
}
}
}
return findings
}
// scanForWebshells recursively reads directories looking for known webshell
// files and directories. Uses ReadDir (getdents) - no stat unless matched.
func scanForWebshells(ctx context.Context, dir string, maxDepth int, names map[string]bool, dirs map[string]bool, cfg *config.Config, findings *[]alert.Finding) {
if ctx.Err() != nil {
return
}
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
markScanReadError(ctx, "webshells", err)
return
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
// Check suppressed paths (bypassed for explicit full-scan / audit requests).
suppressed := false
if scanRespectsIgnores(ctx, cfg) {
for _, ignore := range cfg.Suppressions.IgnorePaths {
if matchGlob(fullPath, ignore) {
suppressed = true
break
}
}
}
if suppressed {
continue
}
if entry.IsDir() {
if dirs[name] {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "webshell",
Message: fmt.Sprintf("Webshell directory found: %s", fullPath),
FilePath: fullPath,
})
}
scanForWebshells(ctx, fullPath, maxDepth-1, names, dirs, cfg, findings)
continue
}
nameLower := strings.ToLower(name)
if names[nameLower] {
info, _ := osFS.Stat(fullPath)
var details string
if info != nil {
details = fmt.Sprintf("Size: %d, Mtime: %s", info.Size(), info.ModTime())
}
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "webshell",
Message: fmt.Sprintf("Known webshell found: %s", fullPath),
Details: details,
FilePath: fullPath,
})
}
// .haxor extension
if strings.HasSuffix(nameLower, ".haxor") || strings.HasSuffix(nameLower, ".cgix") {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "webshell",
Message: fmt.Sprintf("Suspicious CGI file: %s", fullPath),
FilePath: fullPath,
})
}
// File permission anomalies - only check PHP-executable files to keep it fast
if isExecutablePHPName(nameLower) {
info, err := entry.Info()
markScanReadError(ctx, "webshells", err)
if err == nil {
mode := info.Mode()
// World-writable PHP
if mode&0002 != 0 {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "world_writable_php",
Message: fmt.Sprintf("World-writable PHP file: %s", fullPath),
Details: fmt.Sprintf("Mode: %s", mode),
FilePath: fullPath,
})
}
// Note: executable PHP check removed - most PHP files on cPanel
// have +x due to suPHP/lsapi, making this too noisy.
}
}
}
}
// matchGlob reports whether path is covered by an operator suppression pattern.
//
// Matching is tried in order:
// 1. filepath.Match against the basename ("*.php", "*.log") and the full path.
// 2. For a leading-any-depth glob pattern, a substring match of the
// wildcard-stripped residue -- but ONLY when that residue still contains a
// path separator with literal content (e.g. "*/node_modules/*" ->
// "/node_modules/"). This preserves the "directory anywhere in the path, at
// any depth" intent without broadening anchored full-path globs like
// "/tmp/safe/*" into recursive subtree suppressions.
// 3. For a pattern with no wildcards, a literal substring match, so an operator
// can suppress a directory ("/uploads/") or a filename ("adminer.php").
//
// The separator requirement in step 2 is the fix for an over-suppression
// footgun: the previous code stripped every "*" and substring-matched the
// remainder, so "*.php" became the bare token ".php" and silenced every file
// whose path merely contained ".php" -- turning a narrow pattern into a
// whole-subtree allowlist an attacker could hide a webshell in.
// PathMatchesIgnore reports whether path is covered by any of the operator's
// suppressions.ignore_paths patterns, using the same glob semantics the
// content checks apply. Exported so the real-time watchers honour the same
// suppression list instead of maintaining a second interpretation of it.
func PathMatchesIgnore(path string, ignores []string) bool {
for _, ignore := range ignores {
if matchGlob(path, ignore) {
return true
}
}
return false
}
func matchGlob(path, pattern string) bool {
if pattern == "" {
return false
}
if strings.ContainsAny(pattern, "*?[") {
if matched, _ := filepath.Match(pattern, filepath.Base(path)); matched {
return true
}
if matched, _ := filepath.Match(pattern, path); matched {
return true
}
if strings.ContainsAny(pattern, "?[") || !hasLeadingAnyDepthGlob(pattern) {
return false
}
residue := strings.ReplaceAll(pattern, "*", "")
if strings.Contains(residue, "/") && strings.Trim(residue, "/") != "" {
return strings.Contains(path, residue)
}
return false
}
return strings.Contains(path, pattern)
}
func hasLeadingAnyDepthGlob(pattern string) bool {
firstSlash := strings.Index(pattern, "/")
if firstSlash <= 0 {
return false
}
return strings.Trim(pattern[:firstSlash], "*") == ""
}
package checks
import (
"context"
"fmt"
"net"
"net/netip"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// firewallHostIPv6Addrs is the host's global-unicast IPv6 address source;
// swapped in tests.
var firewallHostIPv6Addrs = hostGlobalIPv6Addrs
// hostGlobalIPv6Addrs returns the host's global-unicast IPv6 addresses.
// Loopback, link-local, ULA, and IPv4 do not count: only globally routable
// IPv6 makes the unmanaged-family bypass reachable from the internet.
func hostGlobalIPv6Addrs() ([]string, error) {
addrs, err := net.InterfaceAddrs()
if err != nil {
return nil, fmt.Errorf("enumerate interface addresses: %w", err)
}
var out []string
for _, a := range addrs {
if ip := globalUnicastIPv6FromCIDR(a.String()); ip != "" {
out = append(out, ip)
}
}
return out, nil
}
func globalUnicastIPv6FromCIDR(cidr string) string {
addrText, prefixText, hasPrefix := strings.Cut(cidr, "/")
ip, err := netip.ParseAddr(addrText)
if err != nil {
return ""
}
if hasPrefix {
bits, err := strconv.Atoi(prefixText)
if err != nil || bits < 0 || bits > ip.BitLen() {
return ""
}
}
if ip.Is4() || ip.Is4In6() {
return ""
}
ip = ip.WithZone("")
if !ip.IsGlobalUnicast() || ip.IsPrivate() {
return ""
}
return ip.String()
}
func CheckFirewall(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
if !cfg.Firewall.Enabled {
// Firewall not managed by CSM - skip nftables checks
return findings
}
// When firewall.ipv6 is false the engine inserts a blanket NFPROTO
// ipv6 accept ahead of the blocked sets and every port rule, so on a
// dual-stack host the entire IPv6 attack surface bypasses the
// DROP-policy firewall. That trade-off must never be silent.
if !cfg.Firewall.IPv6 {
addrs, err := firewallHostIPv6Addrs()
if err != nil {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "firewall_ipv6_unmanaged",
Message: "Unable to inspect host IPv6 addresses while firewall.ipv6 is disabled; CSM cannot determine whether the unmanaged IPv6 path exposes this host",
Timestamp: time.Now(),
})
} else if len(addrs) > 0 {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "firewall_ipv6_unmanaged",
Message: "IPv6 traffic bypasses the CSM firewall: firewall.ipv6 is disabled but this host has a global IPv6 address; all IPv6 traffic is accepted unfiltered - set firewall.ipv6: true",
Timestamp: time.Now(),
})
}
}
// The running engine pairs live rules with its applied baseline. Standalone
// checks have no engine and inspect the existing table through cmdExec.
var out []byte
var applied string
var err error
monitor, managed := getIPBlocker().(interface {
RulesetSnapshot() (current, applied string, err error)
})
if managed {
var current string
current, applied, err = monitor.RulesetSnapshot()
out = []byte(current)
} else {
out, err = cmdExec.RunAllowNonZero("nft", "list", "table", "inet", "csm")
}
if err != nil {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "firewall",
Message: "CSM firewall table not found in nftables - rules may not be active",
Timestamp: time.Now(),
})
return findings
}
output := string(out)
required := []string{"chain input", "chain output", "set blocked_ips", "set allowed_ips", "set infra_ips"}
for _, component := range required {
if !strings.Contains(output, component) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "firewall",
Message: fmt.Sprintf("Firewall missing expected component: %s", component),
Timestamp: time.Now(),
})
}
}
// Hash the rule structure, excluding dynamic set members which change on
// every block/unblock.
hash := nftRulesetStructureHash(out)
// Only a successful engine Apply establishes a ruleset baseline. The
// unkeyed digest in csm.yaml can change without any firewall
// transaction, including when SIGHUP defers restart-required settings.
prev, exists := store.GetRaw("_nftables_rules_hash")
if applied != "" {
prev, exists = nftRulesetStructureHash([]byte(applied)), true
} else if managed {
findings = append(findings, alert.Finding{
Severity: alert.Warning, Check: "firewall", Timestamp: time.Now(),
Message: "Firewall integrity baseline unavailable after applying rules; re-apply the firewall to restore monitoring",
})
}
if exists && prev != hash {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "firewall",
Message: "nftables ruleset modified outside of CSM",
Timestamp: time.Now(),
})
}
// Retain the trusted baseline while a mismatch persists so acknowledging
// one finding cannot make the next check treat the modified rules as clean.
if applied != "" {
store.SetRaw("_nftables_rules_hash", prev)
} else if !exists {
store.SetRaw("_nftables_rules_hash", hash)
}
// Check for dangerous ports in config
findings = append(findings, checkDangerousPorts(cfg)...)
return findings
}
func nftRulesetStructureHash(out []byte) string {
var stableLines []byte
inElements := false
elementBraceDepth := 0
for _, line := range strings.Split(string(out), "\n") {
trimmed := strings.TrimSpace(line)
if inElements {
elementBraceDepth += nftBraceDelta(trimmed)
if elementBraceDepth <= 0 {
inElements = false
elementBraceDepth = 0
}
continue
}
if depth, ok := nftElementsBlockDepth(trimmed); ok {
if depth > 0 {
inElements = true
elementBraceDepth = depth
}
continue
}
stableLines = append(stableLines, line...)
stableLines = append(stableLines, '\n')
}
return hashBytes(stableLines)
}
func nftElementsBlockDepth(line string) (int, bool) {
const key = "elements"
if !strings.HasPrefix(line, key) {
return 0, false
}
rest := line[len(key):]
if rest != "" && rest[0] != '=' && rest[0] != ' ' && rest[0] != '\t' {
return 0, false
}
rest = strings.TrimSpace(rest)
if !strings.HasPrefix(rest, "=") {
return 0, false
}
rest = strings.TrimSpace(rest[1:])
if !strings.HasPrefix(rest, "{") {
return 0, false
}
return nftBraceDelta(rest), true
}
func nftBraceDelta(line string) int {
depth := 0
inQuote := false
escaped := false
for _, r := range line {
if inQuote {
if escaped {
escaped = false
continue
}
switch r {
case '\\':
escaped = true
case '"':
inQuote = false
}
continue
}
switch r {
case '"':
inQuote = true
case '#':
return depth
case '{':
depth++
case '}':
depth--
}
}
return depth
}
func checkDangerousPorts(cfg *config.Config) []alert.Finding {
var findings []alert.Finding
dangerousPorts := make(map[int]bool)
for _, p := range cfg.BackdoorPorts {
dangerousPorts[p] = true
}
restricted := make(map[int]bool)
for _, p := range cfg.Firewall.RestrictedTCP {
restricted[p] = true
}
for _, port := range cfg.Firewall.TCPIn {
if dangerousPorts[port] {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "firewall_ports",
Message: fmt.Sprintf("Known backdoor port %d is open in firewall TCP_IN", port),
Timestamp: time.Now(),
})
}
if restricted[port] {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "firewall_ports",
Message: fmt.Sprintf("Restricted port %d found in public TCP_IN - should be infra-only", port),
Timestamp: time.Now(),
})
}
}
return findings
}
package checks
import (
"bufio"
"context"
"crypto/sha256"
"fmt"
"io"
"path/filepath"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// splitValiasLine splits a valiases line "key: destinations" into its key
// and raw, still-quoted destination list. Returns empty strings for
// comments, blank lines, or malformed lines.
func splitValiasLine(line string) (key, rawDest string) {
line = strings.Trim(line, " \t\r\n\v\f")
if line == "" || strings.HasPrefix(line, "#") {
return "", ""
}
idx := strings.IndexByte(line, ':')
if idx < 0 {
return "", ""
}
return strings.TrimSpace(line[:idx]), strings.Trim(line[idx+1:], " \t\r\n\v\f")
}
// ValiasEntry is one destination of one valiases alias.
type ValiasEntry struct {
LocalPart string
Domain string
Dest string
}
// ParseValiasEntries reads a valiases file for fileDomain and returns one
// entry per destination. cPanel keys aliases by full address
// ("bob@example.com"); a bare key ("bob", "*") belongs to fileDomain.
func ParseValiasEntries(r io.Reader, fileDomain string) ([]ValiasEntry, error) {
var entries []ValiasEntry
scanner := bufio.NewScanner(r)
for scanner.Scan() {
key, rawDest := splitValiasLine(scanner.Text())
if key == "" || rawDest == "" {
continue
}
localPart, domain := key, fileDomain
if at := strings.LastIndexByte(key, '@'); at > 0 && at < len(key)-1 {
localPart, domain = key[:at], key[at+1:]
}
for _, d := range splitValiasDests(rawDest) {
if d == "" {
continue
}
entries = append(entries, ValiasEntry{LocalPart: localPart, Domain: domain, Dest: d})
}
}
return entries, scanner.Err()
}
// unquoteValiasDest strips one layer of matching double or single quotes.
// cPanel writes pipe and command destinations quoted ("|/path args"), and
// the pipe detector keys on the leading "|".
func unquoteValiasDest(s string) string {
s = strings.Trim(s, " \t\r\n\v\f")
if len(s) >= 2 && (s[0] == '"' || s[0] == '\'') && s[len(s)-1] == s[0] {
quote := s[0]
s = strings.Trim(s[1:len(s)-1], " \t\r\n\v\f")
// The redirect router removes one backslash layer in double-quoted
// pipes/files before the pipe transport splits command arguments.
if quote == '"' && (strings.HasPrefix(s, "|") || strings.HasPrefix(s, "/")) {
var b strings.Builder
for i := 0; i < len(s); i++ {
if s[i] == '\\' && i+1 < len(s) {
i++
}
b.WriteByte(s[i])
}
return b.String()
}
}
return s
}
// splitValiasDests splits a comma-separated destination list, keeping
// commas inside quotes with their destination, and unquotes each item.
func splitValiasDests(dest string) []string {
var out []string
var cur strings.Builder
var quote byte
for i := 0; i < len(dest); i++ {
c := dest[i]
switch {
case quote != 0:
if quote == '"' && c == '\\' && i+1 < len(dest) {
cur.WriteByte(c)
i++
cur.WriteByte(dest[i])
continue
}
if c == quote {
quote = 0
}
cur.WriteByte(c)
case c == '"' || c == '\'':
quote = c
cur.WriteByte(c)
case c == ',':
out = append(out, unquoteValiasDest(cur.String()))
cur.Reset()
default:
cur.WriteByte(c)
}
}
if quote != 0 {
// A malformed unmatched quote must not turn every later comma into
// quoted data and hide a pipe or external forwarder. Fall back to a
// conservative raw split; false positives are preferable to dropping
// the rest of an attacker-controlled valias line.
parts := strings.Split(dest, ",")
out = out[:0]
for _, part := range parts {
part = strings.Trim(part, " \t\r\n\v\f")
part = strings.Trim(strings.Trim(part, "\"'"), " \t\r\n\v\f")
out = append(out, part)
}
return out
}
out = append(out, unquoteValiasDest(cur.String()))
return out
}
// IsPipeForwarder returns true if the destination is a pipe forwarder,
// excluding pipes whose executed command is a cPanel built-in.
func IsPipeForwarder(dest string) bool {
return strings.HasPrefix(dest, "|") && !isSafePipe(dest)
}
// isDevNullForwarder returns true if the destination is /dev/null.
func isDevNullForwarder(dest string) bool {
return dest == "/dev/null"
}
// IsExternalDest returns true if the destination is an email address
// with a domain not in the local domains set. A pipe is a command, even when
// its arguments contain an address.
func IsExternalDest(dest string, localDomains map[string]bool) bool {
if strings.HasPrefix(dest, "|") {
return false
}
atIdx := strings.LastIndexByte(dest, '@')
if atIdx < 0 || atIdx >= len(dest)-1 {
return false
}
domain := strings.ToLower(dest[atIdx+1:])
return !localDomains[domain]
}
// parseVfilterExternalDests extracts external email destinations from vfilter content.
// Looks for `to "dest@domain"` directives.
func parseVfilterExternalDests(content string, localDomains map[string]bool) []string {
var external []string
lines := strings.Split(content, "\n")
for _, line := range lines {
line = strings.TrimSpace(line)
// Look for: to "email@domain"
if !strings.HasPrefix(line, "to ") && !strings.HasPrefix(line, "to\t") {
continue
}
// Extract the quoted destination
quoteStart := strings.IndexByte(line, '"')
if quoteStart < 0 {
continue
}
rest := line[quoteStart+1:]
quoteEnd := strings.IndexByte(rest, '"')
if quoteEnd < 0 {
continue
}
dest := rest[:quoteEnd]
if IsExternalDest(dest, localDomains) {
external = append(external, dest)
}
}
return external
}
// parseLocalDomainsContent parses the content of /etc/localdomains or /etc/virtualdomains.
func parseLocalDomainsContent(content string) map[string]bool {
domains := make(map[string]bool)
lines := strings.Split(content, "\n")
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
// virtualdomains format: "domain: user" - take domain part
if idx := strings.IndexByte(line, ':'); idx > 0 {
line = strings.TrimSpace(line[:idx])
}
domains[strings.ToLower(line)] = true
}
return domains
}
// loadLocalDomains reads /etc/localdomains and /etc/virtualdomains.
func loadLocalDomains() map[string]bool {
domains := make(map[string]bool)
for _, path := range []string{"/etc/localdomains", "/etc/virtualdomains"} {
data, err := osFS.ReadFile(path)
if err != nil {
continue
}
for k, v := range parseLocalDomainsContent(string(data)) {
domains[k] = v
}
}
return domains
}
// IsKnownForwarder checks if a forwarder rule matches the known forwarders suppression list.
func IsKnownForwarder(localPart, domain, dest string, knownForwarders []string) bool {
entry := fmt.Sprintf("%s@%s: %s", localPart, domain, dest)
for _, known := range knownForwarders {
if strings.EqualFold(strings.TrimSpace(known), entry) {
return true
}
}
return false
}
// fileContentHash returns the SHA256 hex hash of a file's content.
func fileContentHash(path string) (string, error) {
data, err := osFS.ReadFile(path)
if err != nil {
return "", err
}
h := sha256.Sum256(data)
return fmt.Sprintf("%x", h[:]), nil
}
// CheckForwarders audits all valiases and vfilters files for dangerous forwarder
// patterns. Uses internal throttle: skips if last refresh was less than
// password_check_interval_min ago (reuses the same interval).
func CheckForwarders(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
db := store.Global()
if db == nil {
return nil
}
// Internal throttle (24h) - reuse PasswordCheckIntervalMin
if !ForceAll {
lastRefreshStr := db.GetMetaString("email:fwd_last_refresh")
if lastRefreshStr != "" {
if lastRefresh, err := time.Parse(time.RFC3339, lastRefreshStr); err == nil {
interval := time.Duration(cfg.EmailProtection.PasswordCheckIntervalMin) * time.Minute
if time.Since(lastRefresh) < interval {
// A skipped cycle examined nothing, so it must not let the
// runner retire what the last run found.
markCheckSkipped(ctx, "email_forwarder_audit")
return nil
}
}
}
}
if ctx.Err() != nil {
return nil
}
localDomains := loadLocalDomains()
var findings []alert.Finding
// Audit valiases. Rank by mtime desc so recently-changed mail
// domains process first when the check timeout cuts iteration short.
if ctx.Err() != nil {
return findings
}
maxFiles := accountScanMaxFiles(ctx, cfg)
baselineComplete := true
valiasFiles, err := osFS.Glob("/etc/valiases/*")
if err != nil {
baselineComplete = false
}
baselineComplete = baselineComplete && scanCoversAllFiles(valiasFiles, maxFiles)
rankedValiasFiles := rankPathsByMtimeDesc(ctx, valiasFiles, maxFiles)
if ctx.Err() != nil {
return findings
}
for _, path := range rankedValiasFiles {
if ctx.Err() != nil {
return findings
}
domain := filepath.Base(path)
entries, ok := auditValiasFileWithStatus(path, domain, localDomains, cfg)
if !ok {
baselineComplete = false
}
findings = append(findings, entries...)
}
// Audit vfilters
if ctx.Err() != nil {
return findings
}
vfilterFiles, err := osFS.Glob("/etc/vfilters/*")
if err != nil {
baselineComplete = false
}
baselineComplete = baselineComplete && scanCoversAllFiles(vfilterFiles, maxFiles)
rankedVfilterFiles := rankPathsByMtimeDesc(ctx, vfilterFiles, maxFiles)
if ctx.Err() != nil {
return findings
}
for _, path := range rankedVfilterFiles {
if ctx.Err() != nil {
return findings
}
domain := filepath.Base(path)
entries, ok := auditVfilterFileWithStatus(path, domain, localDomains, cfg)
if !ok {
baselineComplete = false
}
findings = append(findings, entries...)
}
if ctx.Err() != nil {
return findings
}
baselineExists := db.GetMetaString("email:fwd_last_refresh") != ""
if baselineComplete || baselineExists {
_ = db.SetMetaString("email:fwd_last_refresh", time.Now().Format(time.RFC3339))
}
return findings
}
// forwarderFileIsNew reports whether a forwarder/filter file should be
// treated as newly added. A stored hash that differs means the file changed.
// No stored hash means one of two things: before the first complete audit
// (no baseline marker) it is pre-existing install backlog and stays quiet;
// after the baseline it genuinely appeared post-audit -- the classic BEC
// drop the first-sight suppression used to silence forever, because the
// next scan saw an unchanged hash.
func forwarderFileIsNew(db *store.DB, baselineKey, hashKey, currentHash string) bool {
old, found := db.GetForwarderHash(hashKey)
if found {
return old != currentHash
}
return db.GetMetaString(baselineKey) != ""
}
func scanCoversAllFiles(paths []string, maxFiles int) bool {
return maxFiles <= 0 || len(paths) <= maxFiles
}
// auditValiasFile parses a valiases file and returns findings for dangerous entries.
func auditValiasFile(path, domain string, localDomains map[string]bool, cfg *config.Config) []alert.Finding {
findings, _ := auditValiasFileWithStatus(path, domain, localDomains, cfg)
return findings
}
func auditValiasFileWithStatus(path, domain string, localDomains map[string]bool, cfg *config.Config) ([]alert.Finding, bool) {
f, err := osFS.Open(path)
if err != nil {
return nil, false
}
defer f.Close()
db := store.Global()
var findings []alert.Finding
complete := true
isNew := false
hashKey := "valiases:" + domain
var currentHash string
if db != nil {
var hashErr error
currentHash, hashErr = fileContentHash(path)
if hashErr == nil {
isNew = forwarderFileIsNew(db, "email:fwd_last_refresh", hashKey, currentHash)
} else {
complete = false
}
}
entries, err := ParseValiasEntries(f, domain)
if err != nil {
complete = false
}
for _, e := range entries {
localPart, mailDomain, d := e.LocalPart, e.Domain, e.Dest
if IsKnownForwarder(localPart, mailDomain, d, cfg.EmailProtection.KnownForwarders) {
continue
}
newContext := ""
if isNew {
newContext = " (newly added)"
}
if IsPipeForwarder(d) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_pipe_forwarder",
Message: fmt.Sprintf("Pipe forwarder detected: %s@%s -> %s%s", localPart, mailDomain, d, newContext),
Details: fmt.Sprintf("Domain: %s\nLocal part: %s\nDestination: %s\nFile: %s\nPipe forwarders execute arbitrary commands on incoming mail.", mailDomain, localPart, d, path),
Domain: domain,
TenantID: MailOwner(domain),
})
continue
}
if isDevNullForwarder(d) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_suspicious_forwarder",
Message: fmt.Sprintf("Mail blackhole: %s@%s -> /dev/null%s", localPart, mailDomain, newContext),
Details: fmt.Sprintf("Domain: %s\nLocal part: %s\nDestination: /dev/null\nFile: %s\nAll mail to this address is silently discarded.", mailDomain, localPart, path),
Domain: domain,
TenantID: MailOwner(domain),
})
continue
}
if IsExternalDest(d, localDomains) && isNew {
severity := alert.High
msg := fmt.Sprintf("External forwarder: %s@%s -> %s%s", localPart, mailDomain, d, newContext)
if localPart == "*" {
msg = fmt.Sprintf("Wildcard catch-all to external: *@%s -> %s%s", mailDomain, d, newContext)
}
findings = append(findings, alert.Finding{
Severity: severity,
Check: "email_suspicious_forwarder",
Message: msg,
Details: fmt.Sprintf("Domain: %s\nLocal part: %s\nDestination: %s\nFile: %s", mailDomain, localPart, d, path),
Domain: domain,
TenantID: MailOwner(domain),
})
}
}
if complete && db != nil {
if err := db.SetForwarderHash(hashKey, currentHash); err != nil {
complete = false
}
}
return findings, complete
}
// auditVfilterFile parses a vfilters file and returns findings for external destinations.
func auditVfilterFile(path, domain string, localDomains map[string]bool, cfg *config.Config) []alert.Finding {
findings, _ := auditVfilterFileWithStatus(path, domain, localDomains, cfg)
return findings
}
func auditVfilterFileWithStatus(path, domain string, localDomains map[string]bool, cfg *config.Config) ([]alert.Finding, bool) {
data, err := osFS.ReadFile(path)
if err != nil {
return nil, false
}
db := store.Global()
content := string(data)
complete := true
// Same newness logic as valiases above.
isNew := false
if db != nil {
currentHash := fmt.Sprintf("%x", sha256.Sum256(data))
isNew = forwarderFileIsNew(db, "email:fwd_last_refresh", "vfilters:"+domain, currentHash)
if err := db.SetForwarderHash("vfilters:"+domain, currentHash); err != nil {
complete = false
}
}
externalDests := parseVfilterExternalDests(content, localDomains)
var findings []alert.Finding
for _, dest := range externalDests {
// Suppression check - use "*" as localPart for vfilter entries
if IsKnownForwarder("*", domain, dest, cfg.EmailProtection.KnownForwarders) {
continue
}
// Only alert when the vfilter file actually changed; existing forwarders
// are normal customer configuration, not an attack indicator.
if !isNew {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_suspicious_forwarder",
Message: fmt.Sprintf("External destination in vfilter: %s -> %s (newly added)", domain, dest),
Details: fmt.Sprintf("Domain: %s\nDestination: %s\nFile: %s\nA mail filter rule forwards messages to an external address.", domain, dest, path),
})
}
return findings, complete
}
package checks
import (
"bytes"
"encoding/json"
"fmt"
"hash/fnv"
"io"
"os"
"sort"
"strconv"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/state"
)
const (
maxCatchUpBytes = 8 << 20
maxTrackedIPs = 4096
fingerprintBytes = 512
)
// followState is the persisted position+identity of the syslog follower.
type followState struct {
Offset int64 `json:"offset"`
HeadLen int `json:"head_len"`
HeadFP string `json:"head_fp"`
AnchorLen int `json:"anchor_len"`
AnchorFP string `json:"anchor_fp"`
}
// ftpFailTracker is the persisted detector state: where we last read to, and a
// per-IP sliding window of pure-ftpd auth-failure counts bucketed by minute.
type ftpFailTracker struct {
Follow followState `json:"follow"`
Buckets map[string]map[int64]int `json:"buckets"`
}
func newFTPFailTracker() *ftpFailTracker {
return &ftpFailTracker{Buckets: map[string]map[int64]int{}}
}
func (t *ftpFailTracker) record(ip string, at time.Time) {
if t.Buckets == nil {
t.Buckets = map[string]map[int64]int{}
}
m := t.Buckets[ip]
if m == nil {
m = map[int64]int{}
t.Buckets[ip] = m
}
m[at.Unix()/60]++
}
// evict drops minute buckets older than windowMin minutes and any IP left empty.
func (t *ftpFailTracker) evict(now time.Time, windowMin int) {
cutoff := now.Unix()/60 - int64(windowMin)
for ip, m := range t.Buckets {
for minute := range m {
if minute < cutoff {
delete(m, minute)
}
}
if len(m) == 0 {
delete(t.Buckets, ip)
}
}
}
// capIPs bounds the tracked-IP set, evicting the weakest candidates first.
// Count wins over recency so a flood of one-off sources cannot push an IP that
// is close to the brute-force threshold out of the retained window.
func (t *ftpFailTracker) capIPs(max int) {
if len(t.Buckets) <= max {
return
}
type ipScore struct {
ip string
recent int64
count int
}
scores := make([]ipScore, 0, len(t.Buckets))
for ip, m := range t.Buckets {
var recent int64
var count int
for minute := range m {
if minute > recent {
recent = minute
}
count += m[minute]
}
scores = append(scores, ipScore{ip: ip, recent: recent, count: count})
}
sort.Slice(scores, func(i, j int) bool {
if scores[i].count != scores[j].count {
return scores[i].count < scores[j].count
}
if scores[i].recent != scores[j].recent {
return scores[i].recent < scores[j].recent
}
return scores[i].ip < scores[j].ip
})
for i := 0; i < len(scores)-max; i++ {
delete(t.Buckets, scores[i].ip)
}
}
// count returns the summed failure count for ip across the retained buckets.
func (t *ftpFailTracker) count(ip string) int {
sum := 0
for _, c := range t.Buckets[ip] {
sum += c
}
return sum
}
type ftpOffender struct {
IP string
Count int
}
// offenders returns IPs whose summed failure count over the retained buckets is
// at least threshold, sorted by IP for stable output.
func (t *ftpFailTracker) offenders(threshold int) []ftpOffender {
var out []ftpOffender
for ip, m := range t.Buckets {
sum := 0
for _, c := range m {
sum += c
}
if sum >= threshold {
out = append(out, ftpOffender{ip, sum})
}
}
sort.Slice(out, func(i, j int) bool { return out[i].IP < out[j].IP })
return out
}
// fpAt returns the fnv64a fingerprint (hex) of n bytes at offset off.
// n <= 0 returns "".
func fpAt(f *os.File, off, n int64) (string, error) {
if n <= 0 {
return "", nil
}
buf := make([]byte, n)
read, err := f.ReadAt(buf, off)
if err != nil {
return "", err
}
if read != len(buf) {
return "", io.ErrUnexpectedEOF
}
h := fnv.New64a()
_, _ = h.Write(buf)
return strconv.FormatUint(h.Sum64(), 16), nil
}
// nextNewline returns the byte offset just after the first '\n' at or after
// from, or size if none is found before size.
func nextNewline(f *os.File, from, size int64) (int64, error) {
buf := make([]byte, 1)
for pos := from; pos < size; pos++ {
read, err := f.ReadAt(buf, pos)
if err != nil {
return 0, err
}
if read != len(buf) {
return 0, io.ErrUnexpectedEOF
}
if buf[0] == '\n' {
return pos + 1, nil
}
}
return size, nil
}
// completeLines splits data into complete newline-terminated lines and returns
// the number of bytes consumed (through the last '\n'). Trailing partial bytes
// are not returned and not consumed.
func completeLines(data []byte) ([]string, int64) {
lastNL := bytes.LastIndexByte(data, '\n')
if lastNL < 0 {
return nil, 0
}
var lines []string
for _, ln := range bytes.Split(data[:lastNL+1], []byte{'\n'}) {
if len(ln) == 0 {
continue
}
lines = append(lines, string(ln))
}
return lines, int64(lastNL + 1)
}
// chooseStart returns the byte offset to begin reading from, plus skipped bytes
// for first-run catch-up. It restarts at 0 on truncate / rotate / anchor
// mismatch, and follows from st.Offset on a verified append.
func chooseStart(f *os.File, st followState, curSize int64) (int64, int64, error) {
zero := st.Offset == 0 && st.HeadLen == 0 && st.HeadFP == "" && st.AnchorLen == 0 && st.AnchorFP == ""
if zero {
start := curSize - maxCatchUpBytes
if start <= 0 {
return 0, 0, nil
}
aligned, err := nextNewline(f, start, curSize)
if err != nil {
return 0, 0, err
}
return aligned, aligned, nil // skipped == bytes before aligned
}
if curSize < st.Offset || curSize < int64(st.HeadLen) {
return 0, 0, nil // truncation / copytruncate
}
if st.HeadLen > 0 {
fp, err := fpAt(f, 0, int64(st.HeadLen))
if err != nil {
return 0, 0, err
}
if fp != st.HeadFP {
return 0, 0, nil // different-head rotate
}
}
if st.Offset > 0 && (st.AnchorLen == 0 || st.AnchorFP == "") {
return 0, 0, nil // incomplete stored anchor
}
if st.Offset > 0 {
fp, err := fpAt(f, st.Offset-int64(st.AnchorLen), int64(st.AnchorLen))
if err != nil {
return 0, 0, err
}
if fp != st.AnchorFP {
return 0, 0, nil // same-head replacement
}
}
return st.Offset, 0, nil
}
// fillIdentity sets next's head and anchor fingerprints from the file content.
func fillIdentity(f *os.File, st *followState, curSize int64) error {
headLen := int64(fingerprintBytes)
if curSize < headLen {
headLen = curSize
}
hfp, err := fpAt(f, 0, headLen)
if err != nil {
return err
}
st.HeadLen = int(headLen)
st.HeadFP = hfp
anchorLen := int64(fingerprintBytes)
if st.Offset < anchorLen {
anchorLen = st.Offset
}
afp, err := fpAt(f, st.Offset-anchorLen, anchorLen)
if err != nil {
return err
}
st.AnchorLen = int(anchorLen)
st.AnchorFP = afp
return nil
}
// readNewSyslogLines reads complete lines appended to path since st, returning
// the new lines, the next follow state, bytes skipped by the catch-up cap, and
// any I/O error. On error, next == st so the caller leaves stored state intact.
func readNewSyslogLines(path string, st followState) ([]string, followState, int64, error) {
f, err := osFS.Open(path)
if err != nil {
return nil, st, 0, err
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil, st, 0, err
}
curSize := info.Size()
start, skipped, err := chooseStart(f, st, curSize)
if err != nil {
return nil, st, 0, err
}
if curSize-start > maxCatchUpBytes {
capped := curSize - maxCatchUpBytes
aligned, aerr := nextNewline(f, capped, curSize)
if aerr != nil {
return nil, st, 0, aerr
}
skipped += aligned - start
start = aligned
}
var lines []string
var consumed int64
if curSize-start > 0 {
data := make([]byte, curSize-start)
read, rerr := f.ReadAt(data, start)
if rerr != nil {
return nil, st, 0, rerr
}
if read != len(data) {
return nil, st, 0, io.ErrUnexpectedEOF
}
lines, consumed = completeLines(data)
}
next := followState{Offset: start + consumed}
if err := fillIdentity(f, &next, curSize); err != nil {
return nil, st, 0, err
}
return lines, next, skipped, nil
}
// ftpTrackerKey is underscore-prefixed so state.Store.Update does not prune it
// as a stale non-finding key after 24h.
const ftpTrackerKey = "_ftp_fail_tracker"
func loadFTPFailTracker(store *state.Store) *ftpFailTracker {
raw, ok := store.GetRaw(ftpTrackerKey)
if !ok || raw == "" {
return newFTPFailTracker()
}
var decoded ftpFailTracker
if err := json.Unmarshal([]byte(raw), &decoded); err != nil {
return newFTPFailTracker()
}
if invalidFollowState(decoded.Follow) {
return newFTPFailTracker()
}
if decoded.Buckets == nil {
decoded.Buckets = map[string]map[int64]int{}
}
return &decoded
}
// invalidFollowState reports stored follow state that is impossible or lacks
// the identity needed to verify a claimed offset; such state is dropped so the
// reader falls back to a bounded first-run catch-up.
func invalidFollowState(st followState) bool {
if st.Offset < 0 || st.HeadLen < 0 || st.HeadLen > fingerprintBytes ||
st.AnchorLen < 0 || st.AnchorLen > fingerprintBytes ||
int64(st.AnchorLen) > st.Offset {
return true
}
if st.Offset <= 0 {
return false
}
return st.HeadLen <= 0 || st.HeadFP == "" || st.AnchorLen <= 0 || st.AnchorFP == ""
}
func (t *ftpFailTracker) save(store *state.Store) {
b, err := json.Marshal(t)
if err != nil {
return
}
if err := store.SetRawAndSave(ftpTrackerKey, string(b)); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving FTP fail tracker: %v\n", err)
}
}
const ftpSyslogPath = "/var/log/messages"
// effectiveFTPFailWindowMin returns the operator-configured sliding-window
// length in minutes, or the built-in default (30) when unset.
func effectiveFTPFailWindowMin(cfg *config.Config) int {
if cfg == nil || cfg.Thresholds.FTPFailWindowMin <= 0 {
return 30
}
return cfg.Thresholds.FTPFailWindowMin
}
var (
ftpTrackerMu sync.Mutex
ftpSkippedBytes *metrics.Counter
ftpSkippedBytesOnce sync.Once
)
func observeFTPSkippedBytes(n int64) {
ftpSkippedBytesOnce.Do(func() {
ftpSkippedBytes = metrics.NewCounter(
"csm_checks_ftp_syslog_skipped_bytes_total",
"Bytes of /var/log/messages skipped by the FTP detector catch-up cap (8 MiB). Steady growth means the detector is falling behind the syslog write rate or the log has large bursts between cycles.",
)
metrics.MustRegister("csm_checks_ftp_syslog_skipped_bytes_total", ftpSkippedBytes)
})
ftpSkippedBytes.Add(float64(n))
}
package checks
import (
"fmt"
"os"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// eligibleFullScanChecks is the set of check types that map to a pure file
// quarantine (fixQuarantine) in ApplyFix. These are the only checks eligible
// for full-scan --quarantine remediation.
//
// Explicitly excluded:
// - backdoor_binary, new_executable_in_config → fixKillAndQuarantine (process kill forbidden)
// - htaccess_* → file edit, not a move
// - email_phishing_content → Exim spool, not a regular file
// - suspicious_crontab → crontab truncate, not a pure file move
// - world_writable_php, group_writable_php → chmod, not a move
var eligibleFullScanChecks = map[string]bool{
"webshell": true,
"new_webshell_file": true,
"obfuscated_php": true,
"suspicious_php_content": true,
"new_php_in_languages": true,
"new_php_in_upgrade": true,
"phishing_page": true,
"phishing_directory": true,
}
// QuarantineFindingFile quarantines the file a malware/webshell finding points
// at, for the full-scan --quarantine path. It reuses fixQuarantine (move to the
// quarantine dir + .meta sidecar) and deliberately covers ONLY the pure
// file-quarantine check set — it never kills processes, cleans databases, or
// touches the firewall. Returns eligible=false for any finding that is not a
// quarantinable malware/webshell FILE finding (caller marks those
// "left_for_review").
//
// The job runs unattended, so it gets the same bar the scheduled auto-response
// applies: only a Critical finding (two converging indicators) may act, only on
// a regular file that is not a symlink, never on a whole directory, and a
// WordPress core, plugin or theme file is cleaned in place rather than moved,
// because moving it takes the site down while only the injected code had to go.
func QuarantineFindingFile(f alert.Finding) (RemediationResult, bool) {
if !eligibleFullScanChecks[f.Check] || f.FilePath == "" || f.Severity != alert.Critical {
return RemediationResult{}, false
}
info, err := osFS.Lstat(f.FilePath)
if err != nil || info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() {
return RemediationResult{}, false
}
if ShouldCleanInsteadOfQuarantine(f.FilePath) {
clean := CleanInfectedFile(f.FilePath)
switch {
case clean.Cleaned:
return RemediationResult{
Success: true,
Action: fmt.Sprintf("cleaned %s in place", f.FilePath),
Description: fmt.Sprintf("Removed: %s (backup: %s)", strings.Join(clean.Removals, "; "), clean.BackupPath),
RemediationStatus: "cleaned",
}, true
case clean.Refused || clean.Error == "":
// Nothing the cleaner recognises: a core file with no removable
// injection is an operator decision, not a move.
return RemediationResult{}, false
default:
// Moving a WordPress core, plugin or theme file after the safer
// clean failed defeats this branch's purpose and can take the site
// down. Leave the file in place and report the failed remediation.
return RemediationResult{Error: fmt.Sprintf("cleaning %s in place: %s", f.FilePath, clean.Error)}, true
}
}
return fixQuarantine(f.FilePath), true
}
package checks
import (
"context"
"fmt"
"os"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// CheckOpenBasedir verifies that each cPanel account has proper PHP
// isolation via CageFS and/or open_basedir.
// Flags accounts where CageFS is disabled AND open_basedir is not set.
func CheckOpenBasedir(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Check global CageFS mode
cageFSMode := getCageFSMode()
// Get list of disabled CageFS users (if mode is "Enable All", check exceptions)
disabledUsers := getCageFSDisabledUsers(cageFSMode)
users := cPanelUsersForOpenBasedirScan(ctx)
if cageFSMode == "unknown" && len(disabledUsers) == 0 {
for _, user := range users {
disabledUsers[user] = true
}
}
// For users without CageFS, check if open_basedir is set
for _, user := range users {
// Skip if CageFS is active for this user
if !disabledUsers[user] {
continue
}
// CageFS is disabled for this user - check open_basedir
if !hasOpenBasedir(user) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "open_basedir",
Message: fmt.Sprintf("Account %s has no PHP isolation: CageFS disabled and no open_basedir", user),
Details: "This account's PHP scripts can read any file on the server including other accounts' wp-config.php and /etc/shadow",
})
}
}
return findings
}
func cPanelUsersForOpenBasedirScan(ctx context.Context) []string {
if account := AccountFromContext(ctx); account != "" {
if _, err := osFS.Stat(filepath.Join("/var/cpanel/users", account)); err != nil {
return nil
}
return []string{account}
}
userDirs, _ := osFS.ReadDir("/var/cpanel/users")
users := make([]string, 0, len(userDirs))
for _, userEntry := range userDirs {
users = append(users, userEntry.Name())
}
return users
}
func getCageFSMode() string {
// CageFS mode file
data, err := osFS.ReadFile("/etc/cagefs/cagefs.mp")
if err != nil {
return "unknown"
}
content := strings.TrimSpace(string(data))
if content != "" {
return "enabled"
}
return "unknown"
}
func getCageFSDisabledUsers(mode string) map[string]bool {
disabled := make(map[string]bool)
// Check cagefsctl --list-disabled
out, _ := runCmd("cagefsctl", "--list-disabled")
if out != nil {
for _, line := range strings.Split(strings.TrimSpace(string(out)), "\n") {
user := strings.TrimSpace(line)
if user != "" {
disabled[user] = true
}
}
}
// If mode is unknown (no CageFS), all users are "disabled"
if mode == "unknown" && len(disabled) == 0 {
userDirs, _ := osFS.ReadDir("/var/cpanel/users")
for _, u := range userDirs {
disabled[u.Name()] = true
}
}
return disabled
}
func hasOpenBasedir(user string) bool {
// Check .user.ini in public_html
userIni := filepath.Join(accountHomeDir(user), "public_html", ".user.ini")
if data, err := osFS.ReadFile(userIni); err == nil {
if strings.Contains(strings.ToLower(string(data)), "open_basedir") {
return true
}
}
// Check .htaccess for php_value open_basedir
htaccess := filepath.Join(accountHomeDir(user), "public_html", ".htaccess")
if data, err := osFS.ReadFile(htaccess); err == nil {
if strings.Contains(strings.ToLower(string(data)), "open_basedir") {
return true
}
}
// Check per-user PHP config set via cPanel MultiPHP
phpConfDirs, _ := osFS.Glob("/opt/cpanel/ea-php*/root/etc/php.d/")
for _, confDir := range phpConfDirs {
userConf := filepath.Join(confDir, "local.ini")
if data, err := osFS.ReadFile(userConf); err == nil {
if strings.Contains(string(data), "open_basedir") {
return true
}
}
}
return false
}
// CheckSymlinkAttacks detects symbolic links inside user public_html
// directories that point outside the account's own directory.
// This is a classic shared hosting attack to read other users' files.
func CheckSymlinkAttacks(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
homeDirs, _ := GetScanHomeDirs(ctx)
for _, homeEntry := range homeDirs {
if !homeEntry.IsDir() {
continue
}
user := homeEntry.Name()
homeDir := scanHomeDirPath(homeEntry)
docRoot := filepath.Join(homeDir, "public_html")
scanForMaliciousSymlinks(docRoot, user, homeDir, 4, &findings)
}
return findings
}
func scanForMaliciousSymlinks(dir, user, homeDir string, maxDepth int, findings *[]alert.Finding) {
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, entry := range entries {
fullPath := filepath.Join(dir, entry.Name())
// Check if it's a symlink
if entry.Type()&os.ModeSymlink == 0 {
if entry.IsDir() {
scanForMaliciousSymlinks(fullPath, user, homeDir, maxDepth-1, findings)
}
continue
}
// It's a symlink - read the target
target, err := osFS.Readlink(fullPath)
if err != nil {
continue
}
// Resolve relative targets
if !filepath.IsAbs(target) {
target = filepath.Join(filepath.Dir(fullPath), target)
}
target = filepath.Clean(target)
// Check if target is outside the user's home
if isSymlinkSafe(target, user, homeDir) {
continue
}
// Check if target points to another user's home
if _, targetUser, ok := accountRootOf(target); ok {
if targetUser != user {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "symlink_attack",
Message: fmt.Sprintf("Symlink to another user's directory: %s -> %s", fullPath, target),
Details: fmt.Sprintf("User: %s, target user: %s\nThis could be used to read other accounts' files", user, targetUser),
FilePath: fullPath,
TenantID: user,
})
continue
}
}
// Check if target points to sensitive system files
sensitiveTargets := []string{"/etc/shadow", "/etc/passwd", "/root/"}
for _, sens := range sensitiveTargets {
if strings.HasPrefix(target, sens) {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "symlink_attack",
Message: fmt.Sprintf("Symlink to sensitive system file: %s -> %s", fullPath, target),
Details: fmt.Sprintf("User: %s", user),
FilePath: fullPath,
TenantID: user,
})
break
}
}
}
}
func isSymlinkSafe(target, user, homeDir string) bool {
// Inside own home directory
if strings.HasPrefix(target, homeDir+"/") || target == homeDir {
return true
}
// Standard cPanel-created symlinks
safeTargets := []string{
"/etc/apache2/logs/",
"/usr/local/apache/logs/",
"/var/cpanel/",
"/opt/cpanel/",
"/usr/",
"/var/lib/mysql/",
"/var/run/",
}
for _, safe := range safeTargets {
if strings.HasPrefix(target, safe) {
return true
}
}
return false
}
package checks
import (
"context"
"encoding/hex"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"slices"
"strconv"
"strings"
"syscall"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/sshdconf"
"github.com/pidginhost/csm/internal/store"
)
// auditCmdTimeout is the per-subprocess timeout for audit checks.
// Audit checks are fast config reads, not heavy scans.
const auditCmdTimeout = 10 * time.Second
// RunHardeningAudit runs all hardening checks and returns a report.
// Pure function — reads system state only, no store access.
func RunHardeningAudit(cfg *config.Config) *store.AuditReport {
serverType := detectServerType()
var results []store.AuditResult
results = append(results, auditSSH()...)
results = append(results, auditPHP(serverType)...)
results = append(results, auditWebServer()...)
results = append(results, auditMail()...)
if serverType != "bare" {
results = append(results, auditCPanel(serverType)...)
}
results = append(results, auditOS()...)
results = append(results, auditFirewall()...)
score := 0
for _, r := range results {
if r.Status == "pass" {
score++
}
}
return &store.AuditReport{
Timestamp: time.Now(),
ServerType: serverType,
Results: results,
Score: score,
Total: len(results),
}
}
func detectServerType() string {
info := platform.Detect()
if info.IsCPanel() {
if info.OS == platform.OSCloudLinux {
return "cloudlinux"
}
return "cpanel"
}
return "bare"
}
// auditRunCmd executes a command with the audit-specific timeout via the
// cmdExec injector so tests can mock systemctl/cagefsctl/etc. without
// invoking real binaries on the host.
func auditRunCmd(name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), auditCmdTimeout)
defer cancel()
out, err := cmdExec.RunContext(ctx, name, args...)
if ctx.Err() == context.DeadlineExceeded {
return nil, fmt.Errorf("command timed out: %s", name)
}
return out, err
}
// --- SSH checks ---
var sshdConfigPath = sshdconf.DefaultPath
type sshdSettings struct {
PasswordAuthentication string
PermitRootLogin string
X11Forwarding string
}
func parseSSHDConfig() *sshdconf.Config {
return sshdconf.Parse(osFS, sshdConfigPath)
}
func settingsFromSSHDConfig(parsed *sshdconf.Config) sshdSettings {
return sshdSettings{
PasswordAuthentication: parsed.Value("passwordauthentication"),
PermitRootLogin: parsed.Value("permitrootlogin"),
X11Forwarding: parsed.Value("x11forwarding"),
}
}
// portsLabel renders a port list for operator-facing audit messages.
func portsLabel(ports []int) string {
parts := make([]string, 0, len(ports))
for _, p := range ports {
parts = append(parts, strconv.Itoa(p))
}
if len(parts) == 1 {
return "port " + parts[0]
}
return "ports " + strings.Join(parts, ", ")
}
func auditSSH() []store.AuditResult {
parsed := parseSSHDConfig()
var results []store.AuditResult
// ssh_port. sshd binds every Port directive, so exposure on 22 counts
// even when the config also names an alternate port.
ports := parsed.ListenPorts()
if slices.Contains(ports, 22) {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_port", Title: "SSH Port",
Status: "warn", Message: "SSH is running on default port 22",
Fix: "Change to a non-standard port in /etc/ssh/sshd_config to reduce automated scan noise. Update firewall rules before changing.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_port", Title: "SSH Port",
Status: "pass", Message: fmt.Sprintf("SSH on non-standard %s", portsLabel(ports)),
})
}
// ssh_protocol
proto := parsed.Value("protocol")
if strings.Contains(proto, "1") {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_protocol", Title: "SSH Protocol",
Status: "fail", Message: "SSHv1 protocol is enabled",
Fix: "Set 'Protocol 2' in /etc/ssh/sshd_config. SSHv1 has known cryptographic weaknesses.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_protocol", Title: "SSH Protocol",
Status: "pass", Message: "SSHv1 disabled",
})
}
// ssh_password_auth
passAuth := parsed.Value("passwordauthentication")
if passAuth != "no" {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_password_auth", Title: "SSH PasswordAuthentication",
Status: "fail", Message: "Password authentication is enabled",
Fix: "Set 'PasswordAuthentication no' in /etc/ssh/sshd_config and use SSH key authentication only.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_password_auth", Title: "SSH PasswordAuthentication",
Status: "pass", Message: "Password authentication disabled",
})
}
// ssh_root_login
rootLogin := parsed.Value("permitrootlogin")
if rootLogin == "yes" {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_root_login", Title: "SSH PermitRootLogin",
Status: "fail", Message: "Direct root login is permitted",
Fix: "Set 'PermitRootLogin no' or 'PermitRootLogin prohibit-password' in /etc/ssh/sshd_config.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_root_login", Title: "SSH PermitRootLogin",
Status: "pass", Message: fmt.Sprintf("PermitRootLogin set to %s", rootLogin),
})
}
// ssh_max_auth_tries
maxTries := parsed.Value("maxauthtries")
n, _ := strconv.Atoi(maxTries)
if n == 0 {
n = 6 // default
}
if n > 4 {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_max_auth_tries", Title: "SSH MaxAuthTries",
Status: "warn", Message: fmt.Sprintf("MaxAuthTries is %d (recommended: 4 or less)", n),
Fix: "Set 'MaxAuthTries 4' in /etc/ssh/sshd_config to limit brute-force attempts per connection.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_max_auth_tries", Title: "SSH MaxAuthTries",
Status: "pass", Message: fmt.Sprintf("MaxAuthTries set to %d", n),
})
}
// ssh_x11_forwarding
x11 := parsed.Value("x11forwarding")
if x11 != "no" {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_x11_forwarding", Title: "SSH X11Forwarding",
Status: "warn", Message: "X11 forwarding is enabled",
Fix: "Set 'X11Forwarding no' in /etc/ssh/sshd_config unless X11 forwarding is actively needed.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_x11_forwarding", Title: "SSH X11Forwarding",
Status: "pass", Message: "X11 forwarding disabled",
})
}
// ssh_use_dns
useDNS := parsed.Value("usedns")
if useDNS != "no" {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_use_dns", Title: "SSH UseDNS",
Status: "warn", Message: "UseDNS is enabled",
Fix: "Set 'UseDNS no' in /etc/ssh/sshd_config. Otherwise lfd may not track SSH login failures by IP.",
})
} else {
results = append(results, store.AuditResult{
Category: "ssh", Name: "ssh_use_dns", Title: "SSH UseDNS",
Status: "pass", Message: "UseDNS disabled",
})
}
return results
}
// --- OS hardening checks ---
func auditOS() []store.AuditResult {
var results []store.AuditResult
// /tmp and /var/tmp permissions
for _, dir := range []struct {
path, id, title string
}{
{"/tmp", "os_tmp_permissions", "/tmp Permissions"},
{"/var/tmp", "os_var_tmp_permissions", "/var/tmp Permissions"},
} {
info, err := osFS.Stat(dir.path)
if err != nil {
results = append(results, store.AuditResult{
Category: "os", Name: dir.id, Title: dir.title,
Status: "warn", Message: fmt.Sprintf("Cannot stat %s: %v", dir.path, err),
})
continue
}
// Use the raw Unix mode bits from syscall to get the traditional
// permission representation (sticky=01000, setuid=04000, etc.).
// Go's os.ModeSticky uses high bits that don't map to Unix octal,
// so os.FileMode math produces wrong values for comparison.
// Only check the lower 12 bits (sticky + rwx) — ignore setuid/setgid
// which CloudLinux/CageFS may set on virtmp mounts.
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
results = append(results, store.AuditResult{
Category: "os", Name: dir.id, Title: dir.title,
Status: "warn", Message: fmt.Sprintf("Cannot read ownership of %s", dir.path),
})
continue
}
mode := stat.Mode & 01777 // sticky + rwxrwxrwx, ignore setuid/setgid
if mode != 01777 || stat.Uid != 0 || stat.Gid != 0 {
results = append(results, store.AuditResult{
Category: "os", Name: dir.id, Title: dir.title,
Status: "fail",
Message: fmt.Sprintf("%s has mode %04o uid=%d gid=%d (expected 1777 root:root)", dir.path, mode, stat.Uid, stat.Gid),
Fix: fmt.Sprintf("chmod 1777 %s && chown root:root %s", dir.path, dir.path),
})
} else {
results = append(results, store.AuditResult{
Category: "os", Name: dir.id, Title: dir.title,
Status: "pass", Message: fmt.Sprintf("%s is 1777 root:root", dir.path),
})
}
}
// /etc/shadow permissions
// Accept 0000, 0600 (RHEL/CentOS default), and 0640 (Debian default).
// All three restrict access to root only. 0600 is the standard on
// CentOS/CloudLinux — changing it can break passwd/chage.
if info, err := osFS.Stat("/etc/shadow"); err == nil {
perm := info.Mode().Perm()
if perm == 0 || perm == 0o600 || perm == 0o640 {
results = append(results, store.AuditResult{
Category: "os", Name: "os_shadow_permissions", Title: "/etc/shadow Permissions",
Status: "pass", Message: fmt.Sprintf("/etc/shadow has mode %04o", perm),
})
} else {
results = append(results, store.AuditResult{
Category: "os", Name: "os_shadow_permissions", Title: "/etc/shadow Permissions",
Status: "fail", Message: fmt.Sprintf("/etc/shadow has mode %04o (expected 0000, 0600, or 0640)", perm),
Fix: "chmod 0600 /etc/shadow",
})
}
} else {
results = append(results, store.AuditResult{
Category: "os", Name: "os_shadow_permissions", Title: "/etc/shadow Permissions",
Status: "warn", Message: "Cannot stat /etc/shadow",
})
}
// Swap
if data, err := osFS.ReadFile("/proc/swaps"); err == nil {
lines := strings.Split(strings.TrimSpace(string(data)), "\n")
if len(lines) > 1 {
results = append(results, store.AuditResult{
Category: "os", Name: "os_swap", Title: "Swap Configured",
Status: "pass", Message: fmt.Sprintf("%d swap device(s) active", len(lines)-1),
})
} else {
results = append(results, store.AuditResult{
Category: "os", Name: "os_swap", Title: "Swap Configured",
Status: "warn", Message: "No swap configured",
Fix: "Configure swap space to prevent OOM kills: fallocate -l 2G /swapfile && chmod 600 /swapfile && mkswap /swapfile && swapon /swapfile",
})
}
}
// Distro EOL
results = append(results, checkDistroEOL()...)
// Web-server user crontab: the web server's account (cPanel nobody,
// Debian www-data, RHEL apache/nginx) never needs a crontab; one with
// content is a persistence mechanism planted through the web tier.
// nobody is always included: it is the suEXEC/CGI fallback everywhere.
results = append(results, auditWebUserCrontab(cronSpoolDir(), webUsersWithNobody(webServerUsers())))
// Unnecessary services
results = append(results, checkUnnecessaryServices()...)
// CVE-2026-31431 "Copy Fail" — algif_aead is the AF_ALG submodule the
// exploit chains through. Blacklisting it neutralises the attack on
// unpatched kernels.
results = append(results, auditAlgifAEAD())
// Sysctl checks (table-driven)
sysctlChecks := []struct {
id, title, path, expected string
}{
{"os_sysctl_syncookies", "TCP SYN Cookies", "/proc/sys/net/ipv4/tcp_syncookies", "1"},
{"os_sysctl_aslr", "Address Space Layout Randomization", "/proc/sys/kernel/randomize_va_space", "2"},
{"os_sysctl_rp_filter", "Reverse Path Filtering", "/proc/sys/net/ipv4/conf/all/rp_filter", "1"},
{"os_sysctl_icmp_broadcast", "ICMP Broadcast Ignore", "/proc/sys/net/ipv4/icmp_echo_ignore_broadcasts", "1"},
{"os_sysctl_symlinks", "Protected Symlinks", "/proc/sys/fs/protected_symlinks", "1"},
{"os_sysctl_hardlinks", "Protected Hardlinks", "/proc/sys/fs/protected_hardlinks", "1"},
}
for _, sc := range sysctlChecks {
data, err := osFS.ReadFile(sc.path)
if err != nil {
results = append(results, store.AuditResult{
Category: "os", Name: sc.id, Title: sc.title,
Status: "warn", Message: fmt.Sprintf("Cannot read %s", sc.path),
})
continue
}
val := strings.TrimSpace(string(data))
// Convert /proc/sys path to sysctl dotted notation for fix command
sysctlKey := strings.TrimPrefix(sc.path, "/proc/sys/")
sysctlKey = strings.ReplaceAll(sysctlKey, "/", ".")
if val == sc.expected {
results = append(results, store.AuditResult{
Category: "os", Name: sc.id, Title: sc.title,
Status: "pass", Message: fmt.Sprintf("%s = %s", sysctlKey, val),
})
} else {
results = append(results, store.AuditResult{
Category: "os", Name: sc.id, Title: sc.title,
Status: "fail", Message: fmt.Sprintf("%s = %s (expected %s)", sysctlKey, val, sc.expected),
Fix: fmt.Sprintf("sysctl -w %s=%s && echo '%s = %s' >> /etc/sysctl.d/99-csm-hardening.conf", sysctlKey, sc.expected, sysctlKey, sc.expected),
})
}
}
return results
}
// algifAEADBlacklisted reports whether any of the supplied modprobe.d files
// contain a non-comment directive that prevents algif_aead from loading.
// Either of the following blocks is sufficient:
//
// blacklist algif_aead
// install algif_aead /bin/false
// blacklist af_alg (parent — algif_aead depends on it)
// install af_alg /bin/false (parent — same reason)
//
// The dependency relationship matters for hand-hardened images: a sysadmin
// who blocked only the parent `af_alg` has correctly mitigated Copy Fail
// without needing to also block the AEAD submodule. Reporting "no blacklist
// exists" on such hosts would be a false-fail alert.
//
// `install <module> /sbin/modprobe --ignore-install <module>` is the
// idiomatic re-load form and explicitly does NOT block the module — we
// detect that by skipping any install replacement whose first token's
// basename is "modprobe".
func algifAEADBlacklisted(confs map[string]string) bool {
for _, body := range confs {
for _, line := range strings.Split(body, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
if fields[1] != "algif_aead" && fields[1] != "af_alg" {
continue
}
switch fields[0] {
case "blacklist":
return true
case "install":
if len(fields) < 3 {
// Malformed (no replacement command). Don't claim a
// pass on a half-written directive.
continue
}
// Match on the basename of the first token, not a substring
// anywhere in the line. That correctly classifies
// /sbin/modprobe --ignore-install <module> → re-load
// /bin/false → block
// /usr/local/bin/my-modprobe-wrapper → block (wrapper, not modprobe itself)
// Substring matching would have lumped the wrapper case in
// with the re-load case, producing a false-fail alert.
if filepath.Base(fields[2]) == "modprobe" {
continue
}
return true
}
}
}
return false
}
// algifAEAD identifiers shared by the pure helper and the impure wrapper —
// hoisted to package scope so a future edit cannot drift the two copies out
// of sync (which would silently mismatch Name fields between the pass/fail
// path and the warn path).
const (
algifAEADAuditID = "os_algif_aead_blocked"
algifAEADAuditTitle = "AF_ALG (algif_aead) Blocked — CVE-2026-31431"
)
// evaluateAlgifAEAD is the pure, testable core of the algif_aead hardening
// check. `loaded` reports whether algif_aead currently shows up in
// /proc/modules; `confs` is a map of modprobe.d file path → contents.
func evaluateAlgifAEAD(loaded bool, confs map[string]string) store.AuditResult {
blocked := algifAEADBlacklisted(confs)
switch {
case !loaded && blocked:
msg := "algif_aead is blacklisted and not loaded"
if content, ok := confs[afAlgMarkerPath]; ok && validateMarkerContent([]byte(content)) {
msg = "algif_aead is blocked by CSM-managed enforcement (csm harden --copy-fail)"
}
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "pass", Message: msg,
}
case loaded:
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "fail",
Message: "algif_aead is currently loaded — Copy Fail (CVE-2026-31431) exploitable",
Fix: "echo 'install algif_aead /bin/false' > /etc/modprobe.d/csm-disable-algif.conf && modprobe -r algif_aead af_alg",
}
default:
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "fail",
Message: "algif_aead is not loaded but no modprobe.d blacklist exists — module can be loaded on demand",
Fix: "echo 'install algif_aead /bin/false' > /etc/modprobe.d/csm-disable-algif.conf",
}
}
}
// auditAlgifAEAD is the impure wrapper: it reads the running kernel's
// build configuration, KernelCare/livepatch state, /proc/modules, and
// /etc/modprobe.d/*.conf via osFS/cmdExec, then produces an AuditResult.
//
// Decision order:
// 1. KernelCare has applied a Copy Fail livepatch -> pass.
// 2. Kernel has CONFIG_CRYPTO_USER_API_AEAD=y -> fail with a
// truthful message; the modprobe blacklist is ineffective on this
// kernel because the AEAD code is statically linked.
// 3. Otherwise, fall through to the modprobe-state evaluator (the
// existing logic for hosts where AF_ALG is a loadable module).
//
// If any modprobe.d file is unreadable, return a "warn" AuditResult
// naming the offending file rather than silently misreporting.
func auditAlgifAEAD() store.AuditResult {
kernelState := observeAFAlgKernelState()
if kernelState.LivepatchActive {
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "pass",
Message: kernelState.String(),
}
}
if kernelState.BuiltIn {
// On built-in kernels the modprobe blacklist is ineffective. Two
// interim options exist: KernelCare livepatch (handled above) or
// per-service seccomp drop-ins. Recognize the seccomp coverage
// here so an operator who ran `csm harden --copy-fail-seccomp`
// gets a truthful pass.
seccomp := SummarizeAFAlgSeccompCoverage()
if len(seccomp.Covered) > 0 && len(seccomp.Uncovered) == 0 {
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "pass",
Message: fmt.Sprintf(
"kernel built-in but Copy Fail blocked by seccomp drop-ins on %d units (%s)",
len(seccomp.Covered), strings.Join(seccomp.Covered, ", "),
),
}
}
fixMsg := "Apply KernelCare/kpatch when the CVE-2026-31431 patch ships (kcarectl --update); " +
"or run `csm harden --copy-fail-seccomp` to apply per-service seccomp drop-ins now. " +
"The modprobe blacklist file is harmless but does not protect this kernel."
messageDetail := "modprobe blacklist is ineffective on this kernel and Copy Fail (CVE-2026-31431) is exploitable"
if len(seccomp.Covered) > 0 {
messageDetail = fmt.Sprintf(
"seccomp drop-ins present on %d units but %d candidate units still uncovered (%s)",
len(seccomp.Covered), len(seccomp.Uncovered), strings.Join(seccomp.Uncovered, ", "),
)
}
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "fail",
Message: "AF_ALG aead is built into the kernel (CONFIG_CRYPTO_USER_API_AEAD=y); " +
messageDetail,
Fix: fixMsg,
}
}
loaded := false
for _, mod := range loadModuleList() {
if mod == "algif_aead" {
loaded = true
break
}
}
confs := make(map[string]string)
matches, err := osFS.Glob("/etc/modprobe.d/*.conf")
if err == nil {
for _, p := range matches {
data, err := osFS.ReadFile(p)
if err != nil {
return store.AuditResult{
Category: "os", Name: algifAEADAuditID, Title: algifAEADAuditTitle,
Status: "warn",
Message: fmt.Sprintf("Cannot read %s: %v — blacklist state undetermined", p, err),
}
}
confs[p] = string(data)
}
}
return evaluateAlgifAEAD(loaded, confs)
}
// distroEOLPolicy encodes the oldest supported major version per known OS.
// Anything below the minimum is considered EOL by this check.
var distroEOLPolicy = map[platform.OSFamily]int{
platform.OSAlma: 8,
platform.OSRocky: 8,
platform.OSRHEL: 8,
platform.OSCloudLinux: 7,
platform.OSUbuntu: 20, // 20.04 is the oldest non-EOL LTS
platform.OSDebian: 11, // Debian 11 "bullseye"
}
func checkDistroEOL() []store.AuditResult {
return evaluateDistroEOL(platform.Detect(), readOSReleasePretty())
}
// evaluateDistroEOL is the pure, testable core of checkDistroEOL. It returns
// an AuditResult based purely on the supplied platform info and
// PRETTY_NAME string (either may be empty).
func evaluateDistroEOL(info platform.Info, prettyName string) []store.AuditResult {
if prettyName == "" && info.OSVersion != "" {
prettyName = fmt.Sprintf("%s %s", info.OS, info.OSVersion)
}
if info.OS == platform.OSUnknown || info.OSVersion == "" {
return []store.AuditResult{{
Category: "os", Name: "os_distro_eol", Title: "Distribution End of Life",
Status: "warn", Message: "Unable to determine distribution version",
}}
}
if info.OS == platform.OSCentOS {
return []store.AuditResult{{
Category: "os", Name: "os_distro_eol", Title: "Distribution End of Life",
Status: "fail",
Message: fmt.Sprintf("%s — CentOS is end-of-life", prettyName),
Fix: "Migrate to a supported replacement such as AlmaLinux, Rocky Linux, or RHEL. CentOS no longer receives security patches.",
}}
}
// Extract the major version. Ubuntu/Debian use "24.04" / "12", RHEL
// family uses "8.6" / "10", etc. Taking the integer prefix handles both.
majorStr, _, _ := strings.Cut(info.OSVersion, ".")
major, err := strconv.Atoi(majorStr)
if err != nil {
return []store.AuditResult{{
Category: "os", Name: "os_distro_eol", Title: "Distribution End of Life",
Status: "warn", Message: fmt.Sprintf("%s — unable to parse version %q", prettyName, info.OSVersion),
}}
}
minVersion, known := distroEOLPolicy[info.OS]
if !known {
return []store.AuditResult{{
Category: "os", Name: "os_distro_eol", Title: "Distribution End of Life",
Status: "warn", Message: fmt.Sprintf("%s — no EOL policy configured for this distro", prettyName),
}}
}
if major < minVersion {
fix := "Upgrade to a supported release. EOL distributions receive no security patches."
if info.IsRHELFamily() {
fix = fmt.Sprintf("Upgrade to %s %d+ or newer. EOL distributions receive no security patches.", info.OS, minVersion)
}
if info.IsDebianFamily() {
fix = fmt.Sprintf("Upgrade to %s %d+ or newer LTS. EOL distributions receive no security patches.", info.OS, minVersion)
}
return []store.AuditResult{{
Category: "os", Name: "os_distro_eol", Title: "Distribution End of Life",
Status: "fail",
Message: fmt.Sprintf("%s — major version %d is EOL", prettyName, major),
Fix: fix,
}}
}
return []store.AuditResult{{
Category: "os", Name: "os_distro_eol", Title: "Distribution End of Life",
Status: "pass", Message: prettyName,
}}
}
// readOSReleasePretty returns the PRETTY_NAME from /etc/os-release or "".
func readOSReleasePretty() string {
data, err := osFS.ReadFile("/etc/os-release")
if err != nil {
return ""
}
for _, line := range strings.Split(string(data), "\n") {
if strings.HasPrefix(line, "PRETTY_NAME=") {
return strings.Trim(strings.TrimPrefix(line, "PRETTY_NAME="), `"'`)
}
}
return ""
}
func checkUnnecessaryServices() []store.AuditResult {
badServices := []string{
"avahi-daemon", "bluetooth", "cups", "cupsd", "gdm",
"ModemManager", "packagekit", "rpcbind", "wpa_supplicant", "firewalld",
}
out, err := auditRunCmd("systemctl", "list-unit-files", "--state=enabled", "--no-pager", "--no-legend")
if err != nil {
return []store.AuditResult{{
Category: "os", Name: "os_services", Title: "Unnecessary Services",
Status: "warn", Message: "Cannot query systemd unit files",
}}
}
var found []string
lines := strings.Split(string(out), "\n")
for _, line := range lines {
fields := strings.Fields(line)
if len(fields) == 0 {
continue
}
unit := strings.TrimSuffix(fields[0], ".service")
for _, bad := range badServices {
if unit == bad {
found = append(found, bad)
}
}
}
if len(found) == 0 {
return []store.AuditResult{{
Category: "os", Name: "os_services", Title: "Unnecessary Services",
Status: "pass", Message: "No unnecessary services enabled",
}}
}
return []store.AuditResult{{
Category: "os", Name: "os_services", Title: "Unnecessary Services",
Status: "warn",
Message: fmt.Sprintf("Unnecessary services enabled: %s", strings.Join(found, ", ")),
Fix: fmt.Sprintf("systemctl disable --now %s", strings.Join(found, " ")),
}}
}
// --- Firewall checks ---
func auditFirewall() []store.AuditResult {
var results []store.AuditResult
// Gather nft and iptables state
nftOut, nftErr := auditRunCmd("nft", "list", "ruleset")
nftRules := string(nftOut)
hasNft := nftErr == nil && strings.TrimSpace(nftRules) != ""
iptOut, iptErr := auditRunCmd("iptables", "-L", "INPUT", "-n")
iptRules := string(iptOut)
hasIpt := iptErr == nil && strings.TrimSpace(iptRules) != ""
// fw_active
if hasNft || hasIpt {
results = append(results, store.AuditResult{
Category: "firewall", Name: "fw_active", Title: "Firewall Active",
Status: "pass", Message: "Firewall has active rules",
})
} else {
results = append(results, store.AuditResult{
Category: "firewall", Name: "fw_active", Title: "Firewall Active",
Status: "fail", Message: "No active firewall rules detected",
Fix: "Install and configure nftables or iptables with a default-deny policy.",
})
}
results = append(results, checkFirewallDefaultPolicy(hasNft, nftRules, hasIpt, iptRules))
// fw_mysql_exposed
results = append(results, checkMySQLExposed(hasNft, nftRules, hasIpt, iptRules)...)
// fw_telnet
if isPortListening(23) {
results = append(results, store.AuditResult{
Category: "firewall", Name: "fw_telnet", Title: "Telnet Service",
Status: "fail", Message: "Something is listening on port 23 (telnet)",
Fix: "Disable telnet: systemctl disable --now telnet.socket xinetd; use SSH instead.",
})
} else {
results = append(results, store.AuditResult{
Category: "firewall", Name: "fw_telnet", Title: "Telnet Service",
Status: "pass", Message: "Nothing listening on port 23",
})
}
// fw_ipv6
results = append(results, checkIPv6Firewall()...)
return results
}
// getListeningAddr reads /proc/net/tcp for a port in LISTEN state (0A)
// and returns the hex-encoded local IP, or "" if not found.
func getListeningAddr(port int) string {
hexPort := fmt.Sprintf("%04X", port)
for _, path := range []string{"/proc/net/tcp", "/proc/net/tcp6"} {
data, err := osFS.ReadFile(path)
if err != nil {
continue
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) < 4 {
continue
}
// fields[1] = local_address (hex_ip:hex_port), fields[3] = state
if fields[3] != "0A" { // 0A = LISTEN
continue
}
parts := strings.SplitN(fields[1], ":", 2)
if len(parts) != 2 {
continue
}
if parts[1] == hexPort {
return parts[0]
}
}
}
return ""
}
// hexToIPv4 converts a /proc/net/tcp hex IP (little-endian 32-bit) to dotted notation.
func hexToIPv4(h string) string {
if len(h) != 8 {
return h
}
b, err := hex.DecodeString(h)
if err != nil || len(b) != 4 {
return h
}
// /proc/net/tcp stores IPs in little-endian byte order
return fmt.Sprintf("%d.%d.%d.%d", b[3], b[2], b[1], b[0])
}
// isPrivateOrLoopback returns true for loopback, RFC1918, and RFC4193 addresses.
func isPrivateOrLoopback(ipStr string) bool {
ip := net.ParseIP(ipStr)
if ip == nil {
return false
}
if ip.IsLoopback() {
return true
}
// Check private ranges
privateRanges := []string{
"10.0.0.0/8",
"172.16.0.0/12",
"192.168.0.0/16",
"fc00::/7",
}
for _, cidr := range privateRanges {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
continue
}
if network.Contains(ip) {
return true
}
}
return false
}
// isPortListening checks /proc/net/tcp and /proc/net/tcp6 for a port in LISTEN state.
func isPortListening(port int) bool {
hexPort := fmt.Sprintf("%04X", port)
for _, path := range []string{"/proc/net/tcp", "/proc/net/tcp6"} {
data, err := osFS.ReadFile(path)
if err != nil {
continue
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) < 4 {
continue
}
if fields[3] != "0A" {
continue
}
parts := strings.SplitN(fields[1], ":", 2)
if len(parts) == 2 && parts[1] == hexPort {
return true
}
}
}
return false
}
func checkMySQLExposed(hasNft bool, nftRules string, hasIpt bool, iptRules string) []store.AuditResult {
hexAddr := getListeningAddr(3306)
if hexAddr == "" {
return []store.AuditResult{{
Category: "firewall", Name: "fw_mysql_exposed", Title: "MySQL Exposure",
Status: "pass", Message: "MySQL is not listening on any port",
}}
}
// Convert hex addr to IP and check if private/loopback
var ip string
if len(hexAddr) == 8 {
ip = hexToIPv4(hexAddr)
} else {
// IPv6: 32 hex chars, little-endian 4-byte groups
if b, err := hex.DecodeString(hexAddr); err == nil {
ipBytes := make(net.IP, len(b))
// Reverse each 4-byte group for /proc/net/tcp6 little-endian encoding
for i := 0; i+4 <= len(b); i += 4 {
ipBytes[i] = b[i+3]
ipBytes[i+1] = b[i+2]
ipBytes[i+2] = b[i+1]
ipBytes[i+3] = b[i]
}
ip = ipBytes.String()
}
}
// All zeros = wildcard bind
allZero := true
for _, c := range hexAddr {
if c != '0' {
allZero = false
break
}
}
if !allZero && ip != "" && isPrivateOrLoopback(ip) {
return []store.AuditResult{{
Category: "firewall", Name: "fw_mysql_exposed", Title: "MySQL Exposure",
Status: "pass", Message: fmt.Sprintf("MySQL bound to private/loopback address %s", ip),
}}
}
// Wildcard or public bind — check if firewall blocks 3306
// nft has rules, none mention 3306, and the input hook denies by
// default: the port is blocked.
fwBlocks3306 := hasNft && !strings.Contains(nftRules, "3306") && nftInputDefaultDeny(nftRules)
if !fwBlocks3306 && hasIpt && !strings.Contains(iptRules, "3306") {
for _, line := range strings.Split(iptRules, "\n") {
if strings.HasPrefix(line, "Chain INPUT") {
upper := strings.ToUpper(line)
if strings.Contains(upper, "POLICY DROP") || strings.Contains(upper, "POLICY REJECT") {
fwBlocks3306 = true
}
break
}
}
}
bindDesc := "wildcard (0.0.0.0)"
if !allZero && ip != "" {
bindDesc = ip
}
if fwBlocks3306 {
return []store.AuditResult{{
Category: "firewall", Name: "fw_mysql_exposed", Title: "MySQL Exposure",
Status: "warn",
Message: fmt.Sprintf("MySQL bound to %s but firewall blocks port 3306", bindDesc),
Fix: "Bind MySQL to 127.0.0.1 in /etc/my.cnf: bind-address = 127.0.0.1",
}}
}
return []store.AuditResult{{
Category: "firewall", Name: "fw_mysql_exposed", Title: "MySQL Exposure",
Status: "fail",
Message: fmt.Sprintf("MySQL bound to %s and port 3306 appears accessible", bindDesc),
Fix: "Bind MySQL to 127.0.0.1 in /etc/my.cnf and/or block port 3306 in firewall.",
}}
}
func checkIPv6Firewall() []store.AuditResult {
// Check if any non-loopback, non-link-local IPv6 addresses exist
data, err := osFS.ReadFile("/proc/net/if_inet6")
if err != nil {
return []store.AuditResult{{
Category: "firewall", Name: "fw_ipv6", Title: "IPv6 Firewall",
Status: "pass", Message: "IPv6 not active (cannot read /proc/net/if_inet6)",
}}
}
hasIPv6 := false
for _, line := range strings.Split(strings.TrimSpace(string(data)), "\n") {
fields := strings.Fields(line)
if len(fields) < 6 {
continue
}
addr := fields[0]
iface := fields[5]
// Skip loopback
if iface == "lo" {
continue
}
// Skip link-local (fe80::/10)
if strings.HasPrefix(strings.ToLower(addr), "fe80") {
continue
}
hasIPv6 = true
break
}
if !hasIPv6 {
return []store.AuditResult{{
Category: "firewall", Name: "fw_ipv6", Title: "IPv6 Firewall",
Status: "pass", Message: "No non-link-local IPv6 addresses found",
}}
}
// Check nftables for inet/ip6 family chain with input hook and default deny
nftChains, err := auditRunCmd("nft", "list", "chains")
if err == nil {
chainsStr := strings.ToLower(string(nftChains))
// Look for chains in inet or ip6 family with filter hook input
// nft list chains output looks like:
// table inet filter {
// chain input {
// type filter hook input priority filter; policy drop;
// }
// }
var currentFamily string
for _, line := range strings.Split(chainsStr, "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "table ") {
parts := strings.Fields(trimmed)
if len(parts) >= 3 {
currentFamily = parts[1]
}
}
if (currentFamily == "inet" || currentFamily == "ip6") &&
strings.Contains(trimmed, "hook input") {
if strings.Contains(trimmed, "policy drop") || strings.Contains(trimmed, "policy reject") {
return []store.AuditResult{{
Category: "firewall", Name: "fw_ipv6", Title: "IPv6 Firewall",
Status: "pass", Message: fmt.Sprintf("IPv6 active; nftables %s family has default-deny input chain", currentFamily),
}}
}
}
}
}
// Fallback: check ip6tables
ip6Out, err := auditRunCmd("ip6tables", "-L", "INPUT", "-n")
if err == nil {
for _, line := range strings.Split(string(ip6Out), "\n") {
if strings.HasPrefix(line, "Chain INPUT") {
upper := strings.ToUpper(line)
if strings.Contains(upper, "POLICY DROP") || strings.Contains(upper, "POLICY REJECT") {
return []store.AuditResult{{
Category: "firewall", Name: "fw_ipv6", Title: "IPv6 Firewall",
Status: "pass", Message: "IPv6 active; ip6tables INPUT chain has default-deny policy",
}}
}
break
}
}
}
return []store.AuditResult{{
Category: "firewall", Name: "fw_ipv6", Title: "IPv6 Firewall",
Status: "fail", Message: "IPv6 is active but no default-deny input policy found",
Fix: "Configure ip6tables or nftables inet family with a default DROP policy for INPUT.",
}}
}
// --- cPanel/WHM and CloudLinux checks ---
func auditCPanel(serverType string) []store.AuditResult {
var results []store.AuditResult
cpConf := parseCpanelConfig("/var/cpanel/cpanel.config")
// Table-driven boolean checks on cpanel.config.
// fix is the human-readable remediation shown in the UI.
type cpCheck struct {
id, title, key, wantVal string
invert bool // true = fail when value matches wantVal
fix string
}
checks := []cpCheck{
{"cp_ssl_only", "Always Redirect to SSL", "alwaysredirecttossl", "1", false,
"In WHM > Tweak Settings > Redirection, set 'Always redirect to SSL/TLS' to On."},
{"cp_boxtrapper", "BoxTrapper Disabled", "skipboxtrapper", "1", false,
"In WHM > Tweak Settings > Mail, set 'Enable BoxTrapper spam trap' to Off."},
{"cp_password_reset", "Password Reset Disabled", "resetpass", "1", true,
"In WHM > Tweak Settings > System, set 'Reset Password for cPanel accounts' to Off."},
{"cp_password_reset_sub", "Subaccount Password Reset Disabled", "resetpass_sub", "1", true,
"In WHM > Tweak Settings > System, set 'Reset Password for Subaccounts' to Off."},
{"cp_email_passwords", "Email Passwords Disabled", "emailpasswords", "1", true,
"In WHM > Tweak Settings > Security, set 'Send passwords when creating a new account' to Off."},
{"cp_cookie_validation", "Cookie IP Validation", "cookieipvalidation", "strict", false,
"In WHM > Tweak Settings > Security, set 'Cookie IP validation' to strict."},
{"cp_remote_domains", "Remote Domains Disabled", "allowremotedomains", "1", true,
"In WHM > Tweak Settings > Domains, set 'Allow Remote Domains' to Off."},
{"cp_core_dumps", "Core Dumps Disabled", "coredump", "1", true,
"In WHM > Tweak Settings > Security, set 'Generate core dumps' to Off."},
{"cp_nobodyspam", "Nobody Spam Prevention", "nobodyspam", "1", false,
"In WHM > Tweak Settings > Mail, set 'Prevent nobody from sending mail' to On."},
}
for _, c := range checks {
val := cpConf[c.key]
var pass bool
if c.invert {
pass = val != c.wantVal
} else {
pass = val == c.wantVal
}
if pass {
results = append(results, store.AuditResult{
Category: "cpanel", Name: c.id, Title: c.title,
Status: "pass", Message: fmt.Sprintf("%s = %s", c.key, val),
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: c.id, Title: c.title,
Status: "fail", Message: fmt.Sprintf("%s = %s", c.key, val),
Fix: c.fix,
})
}
}
// cp_max_emails_hour
maxEmail := cpConf["maxemailsperhour"]
if maxEmail != "" && maxEmail != "0" {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_max_emails_hour", Title: "Max Emails Per Hour",
Status: "pass", Message: fmt.Sprintf("maxemailsperhour = %s", maxEmail),
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_max_emails_hour", Title: "Max Emails Per Hour",
Status: "fail", Message: "maxemailsperhour is not set or is 0",
Fix: "In WHM > Tweak Settings, set 'Max emails per hour per domain' to a reasonable limit (e.g., 200).",
})
}
// cp_compilers: check /usr/bin/cc permissions
if info, err := osFS.Stat("/usr/bin/cc"); err == nil {
perm := info.Mode().Perm()
if perm <= 0o750 {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_compilers", Title: "Compiler Access Restricted",
Status: "pass", Message: fmt.Sprintf("/usr/bin/cc has mode %04o", perm),
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_compilers", Title: "Compiler Access Restricted",
Status: "fail", Message: fmt.Sprintf("/usr/bin/cc has mode %04o (should be <= 0750)", perm),
Fix: "WHM > Security Center > Compiler Access, or: chmod 750 /usr/bin/cc",
})
}
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_compilers", Title: "Compiler Access Restricted",
Status: "pass", Message: "No compiler found at /usr/bin/cc",
})
}
// cp_ftp_anonymous: parse /etc/pure-ftpd.conf
if data, err := osFS.ReadFile("/etc/pure-ftpd.conf"); err == nil {
noAnon := false
for _, line := range strings.Split(string(data), "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "#") {
continue
}
if strings.HasPrefix(trimmed, "NoAnonymous") {
parts := strings.Fields(trimmed)
if len(parts) >= 2 && strings.EqualFold(parts[1], "yes") {
noAnon = true
}
}
}
if noAnon {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_ftp_anonymous", Title: "Anonymous FTP Disabled",
Status: "pass", Message: "NoAnonymous is enabled in pure-ftpd",
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_ftp_anonymous", Title: "Anonymous FTP Disabled",
Status: "fail", Message: "Anonymous FTP may be enabled",
Fix: "Set 'NoAnonymous yes' in /etc/pure-ftpd.conf and restart pure-ftpd.",
})
}
}
// cp_updates: parse /etc/cpupdate.conf
if data, err := osFS.ReadFile("/etc/cpupdate.conf"); err == nil {
updatesDaily := false
for _, line := range strings.Split(string(data), "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "#") {
continue
}
if strings.HasPrefix(strings.ToUpper(trimmed), "UPDATES=") {
val := strings.TrimPrefix(trimmed, trimmed[:strings.Index(trimmed, "=")+1])
if strings.EqualFold(strings.TrimSpace(val), "daily") {
updatesDaily = true
}
}
}
if updatesDaily {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_updates", Title: "cPanel Auto-Updates",
Status: "pass", Message: "UPDATES=daily in cpupdate.conf",
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cp_updates", Title: "cPanel Auto-Updates",
Status: "warn", Message: "cPanel auto-updates not set to daily",
Fix: "Set UPDATES=daily in /etc/cpupdate.conf or WHM > Update Preferences.",
})
}
}
// CloudLinux-specific checks
if serverType == "cloudlinux" {
results = append(results, auditCloudLinux()...)
}
return results
}
func parseCpanelConfig(path string) map[string]string {
conf := make(map[string]string)
data, err := osFS.ReadFile(path)
if err != nil {
return conf
}
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if idx := strings.Index(line, "="); idx > 0 {
key := strings.TrimSpace(line[:idx])
val := strings.TrimSpace(line[idx+1:])
conf[key] = val
}
}
return conf
}
func auditCloudLinux() []store.AuditResult {
var results []store.AuditResult
// cl_cagefs
out, err := auditRunCmd("cagefsctl", "--cagefs-status")
switch {
case err != nil:
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_cagefs", Title: "CageFS Enabled",
Status: "warn", Message: "Cannot check CageFS status",
})
case strings.Contains(string(out), "Enabled"):
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_cagefs", Title: "CageFS Enabled",
Status: "pass", Message: "CageFS is enabled",
})
default:
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_cagefs", Title: "CageFS Enabled",
Status: "fail", Message: "CageFS is not enabled",
Fix: "Enable CageFS: cagefsctl --enable-all",
})
}
// cl_symlink_protection
if data, err := osFS.ReadFile("/proc/sys/fs/enforce_symlinksifowner"); err == nil {
val := strings.TrimSpace(string(data))
n, _ := strconv.Atoi(val)
if n >= 1 {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_symlink_protection", Title: "CloudLinux Symlink Protection",
Status: "pass", Message: fmt.Sprintf("enforce_symlinksifowner = %s", val),
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_symlink_protection", Title: "CloudLinux Symlink Protection",
Status: "fail", Message: fmt.Sprintf("enforce_symlinksifowner = %s (expected >= 1)", val),
Fix: "sysctl -w fs.enforce_symlinksifowner=1",
})
}
}
// cl_proc_virtualization
if data, err := osFS.ReadFile("/proc/sys/fs/proc_can_see_other_uid"); err == nil {
val := strings.TrimSpace(string(data))
if val == "0" {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_proc_virtualization", Title: "CloudLinux /proc Virtualization",
Status: "pass", Message: "proc_can_see_other_uid = 0",
})
} else {
results = append(results, store.AuditResult{
Category: "cpanel", Name: "cl_proc_virtualization", Title: "CloudLinux /proc Virtualization",
Status: "fail", Message: fmt.Sprintf("proc_can_see_other_uid = %s (expected 0)", val),
Fix: "sysctl -w fs.proc_can_see_other_uid=0",
})
}
}
return results
}
// --- PHP checks ---
func auditPHP(serverType string) []store.AuditResult {
var results []store.AuditResult
type phpInstall struct {
version string // e.g. "8.1"
shortID string // e.g. "81"
iniPath string
fpmDir string // for FPM pool override merging
}
var installs []phpInstall
// cPanel EA4 PHP installs
eaInis, _ := osFS.Glob("/opt/cpanel/ea-php*/root/etc/php.ini")
for _, ini := range eaInis {
// Extract version from path: /opt/cpanel/ea-php81/root/etc/php.ini -> "81"
dir := filepath.Dir(filepath.Dir(filepath.Dir(ini))) // /opt/cpanel/ea-php81/root -> /opt/cpanel/ea-php81
base := filepath.Base(dir) // ea-php81
shortID := strings.TrimPrefix(base, "ea-php")
if len(shortID) >= 2 {
ver := shortID[:len(shortID)-1] + "." + shortID[len(shortID)-1:]
fpmDir := filepath.Join(dir, "root", "etc", "php-fpm.d")
installs = append(installs, phpInstall{
version: ver,
shortID: shortID,
iniPath: ini,
fpmDir: fpmDir,
})
}
}
// CloudLinux alt-php installs (skip Imunify360's internal PHP builds)
if serverType == "cloudlinux" {
altInis, _ := osFS.Glob("/opt/alt/php*/etc/php.ini")
for _, ini := range altInis {
dir := filepath.Dir(filepath.Dir(ini)) // /opt/alt/php81
base := filepath.Base(dir) // php81
if strings.Contains(base, "-") {
continue // skip php74-imunify, php81-hardened, etc.
}
shortID := strings.TrimPrefix(base, "php")
if len(shortID) >= 2 {
ver := shortID[:len(shortID)-1] + "." + shortID[len(shortID)-1:]
installs = append(installs, phpInstall{
version: ver,
shortID: shortID,
iniPath: ini,
})
}
}
}
// Bare server fallback
if len(installs) == 0 {
out, err := auditRunCmd("php", "-i")
if err == nil {
for _, line := range strings.Split(string(out), "\n") {
if strings.HasPrefix(line, "Loaded Configuration File") {
parts := strings.SplitN(line, "=>", 2)
if len(parts) == 2 {
iniPath := strings.TrimSpace(parts[1])
if iniPath != "(none)" && iniPath != "" {
installs = append(installs, phpInstall{
version: "unknown",
shortID: "system",
iniPath: iniPath,
})
}
}
}
}
}
// Try to get version for bare
if len(installs) > 0 && installs[0].version == "unknown" {
vout, verr := auditRunCmd("php", "-v")
if verr == nil {
first := strings.SplitN(string(vout), "\n", 2)[0]
// "PHP 8.2.15 (cli) ..."
fields := strings.Fields(first)
if len(fields) >= 2 {
verParts := strings.SplitN(fields[1], ".", 3)
if len(verParts) >= 2 {
installs[0].version = verParts[0] + "." + verParts[1]
installs[0].shortID = verParts[0] + verParts[1]
}
}
}
}
}
for _, inst := range installs {
data, err := osFS.ReadFile(inst.iniPath)
if err != nil {
continue
}
ini := parsePHPIni(string(data))
// Merge FPM pool overrides if available
if inst.fpmDir != "" {
poolConfs, _ := osFS.Glob(filepath.Join(inst.fpmDir, "*.conf"))
for _, pc := range poolConfs {
pdata, perr := osFS.ReadFile(pc)
if perr != nil {
continue
}
for _, line := range strings.Split(string(pdata), "\n") {
line = strings.TrimSpace(line)
if strings.HasPrefix(line, ";") {
continue
}
// php_admin_value[key] = val or php_value[key] = val
for _, prefix := range []string{"php_admin_value[", "php_value["} {
if strings.HasPrefix(line, prefix) {
rest := strings.TrimPrefix(line, prefix)
if idx := strings.Index(rest, "]"); idx > 0 {
key := rest[:idx]
valPart := rest[idx+1:]
if eqIdx := strings.Index(valPart, "="); eqIdx >= 0 {
val := strings.TrimSpace(valPart[eqIdx+1:])
ini[key] = val
}
}
}
}
}
}
}
suffix := inst.shortID
// php_version check
major, minor := parsePHPVersion(inst.version)
if major > 0 {
if major < 8 || (major == 8 && minor < 1) {
results = append(results, store.AuditResult{
Category: "php", Name: "php_version_" + suffix, Title: fmt.Sprintf("PHP %s Version", inst.version),
Status: "fail", Message: fmt.Sprintf("PHP %s is end-of-life", inst.version),
Fix: fmt.Sprintf("Upgrade PHP %s to 8.1 or later. EOL versions receive no security patches.", inst.version),
})
} else {
results = append(results, store.AuditResult{
Category: "php", Name: "php_version_" + suffix, Title: fmt.Sprintf("PHP %s Version", inst.version),
Status: "pass", Message: fmt.Sprintf("PHP %s is supported", inst.version),
})
}
}
// php_disable_functions
df := strings.TrimSpace(ini["disable_functions"])
if df == "" || strings.EqualFold(df, "none") {
results = append(results, store.AuditResult{
Category: "php", Name: "php_disable_functions_" + suffix, Title: fmt.Sprintf("PHP %s disable_functions", inst.version),
Status: "fail", Message: "disable_functions is empty",
Fix: fmt.Sprintf("Set disable_functions in %s to include dangerous functions like exec, system, passthru, shell_exec, popen, proc_open.", inst.iniPath),
})
} else {
results = append(results, store.AuditResult{
Category: "php", Name: "php_disable_functions_" + suffix, Title: fmt.Sprintf("PHP %s disable_functions", inst.version),
Status: "pass", Message: "disable_functions is configured",
})
}
// php_expose
expose := strings.TrimSpace(strings.ToLower(ini["expose_php"]))
if expose == "off" || expose == "0" {
results = append(results, store.AuditResult{
Category: "php", Name: "php_expose_" + suffix, Title: fmt.Sprintf("PHP %s expose_php", inst.version),
Status: "pass", Message: "expose_php is off",
})
} else {
results = append(results, store.AuditResult{
Category: "php", Name: "php_expose_" + suffix, Title: fmt.Sprintf("PHP %s expose_php", inst.version),
Status: "warn", Message: "expose_php is on — PHP version disclosed in headers",
Fix: fmt.Sprintf("Set expose_php = Off in %s", inst.iniPath),
})
}
// php_allow_url_fopen
auf := strings.TrimSpace(strings.ToLower(ini["allow_url_fopen"]))
if auf == "off" || auf == "0" {
results = append(results, store.AuditResult{
Category: "php", Name: "php_allow_url_fopen_" + suffix, Title: fmt.Sprintf("PHP %s allow_url_fopen", inst.version),
Status: "pass", Message: "allow_url_fopen is off",
})
} else {
results = append(results, store.AuditResult{
Category: "php", Name: "php_allow_url_fopen_" + suffix, Title: fmt.Sprintf("PHP %s allow_url_fopen", inst.version),
Status: "warn", Message: "allow_url_fopen is on — remote file inclusion risk",
Fix: fmt.Sprintf("Set allow_url_fopen = Off in %s", inst.iniPath),
})
}
// php_enable_dl
edl := strings.TrimSpace(strings.ToLower(ini["enable_dl"]))
if edl == "on" || edl == "1" {
results = append(results, store.AuditResult{
Category: "php", Name: "php_enable_dl_" + suffix, Title: fmt.Sprintf("PHP %s enable_dl", inst.version),
Status: "fail", Message: "enable_dl is on — allows loading arbitrary shared objects",
Fix: fmt.Sprintf("Set enable_dl = Off in %s", inst.iniPath),
})
} else {
results = append(results, store.AuditResult{
Category: "php", Name: "php_enable_dl_" + suffix, Title: fmt.Sprintf("PHP %s enable_dl", inst.version),
Status: "pass", Message: "enable_dl is off",
})
}
}
return results
}
func parsePHPIni(content string) map[string]string {
ini := make(map[string]string)
for _, line := range strings.Split(content, "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, ";") || strings.HasPrefix(line, "[") {
continue
}
if idx := strings.Index(line, "="); idx > 0 {
key := strings.TrimSpace(line[:idx])
val := strings.TrimSpace(line[idx+1:])
ini[key] = val
}
}
return ini
}
func parsePHPVersion(ver string) (int, int) {
parts := strings.SplitN(ver, ".", 3)
if len(parts) < 2 {
return 0, 0
}
major, _ := strconv.Atoi(parts[0])
minor, _ := strconv.Atoi(parts[1])
return major, minor
}
// --- Web server checks ---
func auditWebServer() []store.AuditResult {
var results []store.AuditResult
if configPath := findWebServerConfigPath(platform.Detect()); configPath != "" {
results = append(results, auditApacheDirectives(configPath)...)
}
return append(results, auditWebTLS()...)
}
// findWebServerConfigPath returns the active Apache-compatible main file.
// Native LiteSpeed XML uses different syntax; cPanel LiteSpeed consumes the
// EA4 Apache tree and is the only LiteSpeed case this audit can inspect.
func findWebServerConfigPath(info platform.Info) string {
if info.WebServer == platform.WSLiteSpeed && !info.IsCPanel() {
return ""
}
var candidates []string
if info.IsCPanel() {
candidates = append(candidates, "/etc/apache2/conf/httpd.conf")
}
if info.ApacheConfigDir != "" {
switch {
case info.IsDebianFamily() && !info.IsCPanel():
candidates = append(candidates, filepath.Join(info.ApacheConfigDir, "apache2.conf"))
case info.IsRHELFamily() && !info.IsCPanel():
candidates = append(candidates, filepath.Join(info.ApacheConfigDir, "conf", "httpd.conf"))
default:
candidates = append(candidates, filepath.Join(info.ApacheConfigDir, "httpd.conf"))
}
}
switch {
case info.IsDebianFamily():
candidates = append(candidates, "/etc/apache2/apache2.conf")
case info.IsRHELFamily():
candidates = append(candidates, "/etc/httpd/conf/httpd.conf")
case len(candidates) == 0:
candidates = append(candidates, "/etc/httpd/conf/httpd.conf", "/etc/apache2/apache2.conf")
}
seen := make(map[string]bool)
for _, p := range candidates {
if seen[p] {
continue
}
seen[p] = true
if _, err := osFS.Stat(p); err == nil {
return p
} else if !errors.Is(err, os.ErrNotExist) {
return p
}
}
return ""
}
// auditApacheDirectivesMaxScopes caps how many indexing scopes a single
// finding names, so a host with hundreds of vhosts still gets a readable
// message.
const auditApacheDirectivesMaxScopes = 5
// auditApacheDirectives checks the effective Apache configuration after
// Include and IncludeOptional targets have been spliced into place.
func auditApacheDirectives(configPath string) []store.AuditResult {
var results []store.AuditResult
lines, assemblyComplete := assembleApacheConfigWithStatus(configPath)
type directiveCheck struct {
id, title, directive string
goodValues []string
allowScoped bool
}
dirChecks := []directiveCheck{
{"web_server_tokens", "ServerTokens", "ServerTokens", []string{"prod", "productonly"}, false},
{"web_server_signature", "ServerSignature", "ServerSignature", []string{"off"}, true},
{"web_trace_enable", "TraceEnable", "TraceEnable", []string{"off"}, false},
{"web_file_etag", "FileETag", "FileETag", []string{"none"}, true},
}
for _, dc := range dirChecks {
values, configValid := apacheDirectiveValues(lines, dc.directive)
var serverValue *apacheDirectiveValue
var unsafeValue *apacheDirectiveValue
var unsafeConditional *apacheDirectiveValue
for i := range values {
value := &values[i]
if value.Scope == "server config" && !value.Conditional {
serverValue = value
}
if value.Scope != "server config" && !dc.allowScoped {
configValid = false
continue
}
if apacheDirectiveValueIsGood(value.Value, dc.goodValues) {
continue
}
if value.Conditional {
unsafeConditional = value
} else {
unsafeValue = value
break
}
}
if unsafeValue != nil {
results = append(results, store.AuditResult{
Category: "webserver", Name: dc.id, Title: dc.title,
Status: "fail", Message: fmt.Sprintf("%s = %s in %s (%s)", dc.directive, unsafeValue.Value, unsafeValue.Scope, unsafeValue.File),
Fix: fmt.Sprintf("Set '%s %s' in %s", dc.directive, dc.goodValues[0], unsafeValue.File),
})
continue
}
if unsafeConditional != nil {
results = append(results, store.AuditResult{
Category: "webserver", Name: dc.id, Title: dc.title,
Status: "warn", Message: fmt.Sprintf("%s may be %s in conditional %s scope (%s)", dc.directive, unsafeConditional.Value, unsafeConditional.Scope, unsafeConditional.File),
Fix: fmt.Sprintf("Set '%s %s' in %s", dc.directive, dc.goodValues[0], unsafeConditional.File),
})
continue
}
if serverValue == nil {
results = append(results, store.AuditResult{
Category: "webserver", Name: dc.id, Title: dc.title,
Status: "warn", Message: fmt.Sprintf("%s not set in %s or its included snippets", dc.directive, configPath),
Fix: fmt.Sprintf("Add '%s %s' to %s", dc.directive, dc.goodValues[0], configPath),
})
continue
}
if assemblyComplete && configValid {
results = append(results, store.AuditResult{
Category: "webserver", Name: dc.id, Title: dc.title,
Status: "pass", Message: fmt.Sprintf("%s = %s (%s)", dc.directive, serverValue.Value, serverValue.File),
})
continue
}
results = append(results, store.AuditResult{
Category: "webserver", Name: dc.id, Title: dc.title,
Status: "warn", Message: fmt.Sprintf("%s appears secure, but the Apache configuration could not be fully evaluated", dc.directive),
})
}
scopes, optionsValid := apacheIndexesScopesWithStatus(lines)
if len(scopes) == 0 {
if !assemblyComplete || !optionsValid {
results = append(results, store.AuditResult{
Category: "webserver", Name: "web_directory_listing", Title: "Directory Listing",
Status: "warn", Message: "Apache Options could not be fully evaluated",
})
return results
}
results = append(results, store.AuditResult{
Category: "webserver", Name: "web_directory_listing", Title: "Directory Listing",
Status: "pass", Message: "No configured scope enables directory listing",
})
return results
}
shown := scopes
suffix := ""
if len(shown) > auditApacheDirectivesMaxScopes {
suffix = fmt.Sprintf(" (+%d more)", len(shown)-auditApacheDirectivesMaxScopes)
shown = shown[:auditApacheDirectivesMaxScopes]
}
results = append(results, store.AuditResult{
Category: "webserver", Name: "web_directory_listing", Title: "Directory Listing",
Status: "warn",
Message: fmt.Sprintf("Directory listing enabled in: %s%s", strings.Join(shown, ", "), suffix),
Fix: "Replace 'Indexes' with '-Indexes' in the Options directive of each scope listed.",
})
return results
}
func apacheDirectiveValueIsGood(value string, goodValues []string) bool {
for _, good := range goodValues {
if strings.EqualFold(value, good) {
return true
}
}
return false
}
// auditWebTLS probes the local HTTPS listener for legacy TLS versions.
func auditWebTLS() []store.AuditResult {
var results []store.AuditResult
for _, tc := range []struct {
id, title, flag, version string
}{
{"web_tls_version", "Legacy TLS Disabled", "-tls1", "TLSv1.0"},
{"web_tls11_version", "TLS 1.1 Disabled", "-tls1_1", "TLSv1.1"},
} {
out, err := auditRunCmd("openssl", "s_client", "-connect", "localhost:443", tc.flag)
output := string(out)
// If the handshake succeeds, output contains "SSL-Session:" without ":error:" on the same handshake
succeeded := err == nil && strings.Contains(output, "SSL-Session:") && !strings.Contains(output, ":error:")
if succeeded {
results = append(results, store.AuditResult{
Category: "webserver", Name: tc.id, Title: tc.title,
Status: "fail", Message: fmt.Sprintf("Server accepts %s connections", tc.version),
Fix: fmt.Sprintf("Disable %s in your web server's SSL configuration. Minimum should be TLSv1.2.", tc.version),
})
} else {
results = append(results, store.AuditResult{
Category: "webserver", Name: tc.id, Title: tc.title,
Status: "pass", Message: fmt.Sprintf("%s is rejected", tc.version),
})
}
}
return results
}
// --- Mail checks ---
// detectMTA is a seam so tests can pin the host's mail stack.
var detectMTA = platform.DetectMTA
var isCPanelHost = func() bool { return platform.Detect().IsCPanel() }
func auditMail() []store.AuditResult {
var results []store.AuditResult
// mail_root_forwarder
if info, err := osFS.Stat("/root/.forward"); err != nil {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_root_forwarder", Title: "Root Mail Forwarder",
Status: "warn", Message: "/root/.forward does not exist — root mail may go unread",
Fix: "Create /root/.forward with an email address to receive root's mail.",
})
} else if info.Size() == 0 {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_root_forwarder", Title: "Root Mail Forwarder",
Status: "warn", Message: "/root/.forward is empty — root mail may go unread",
Fix: "Add an email address to /root/.forward to receive root's mail.",
})
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_root_forwarder", Title: "Root Mail Forwarder",
Status: "pass", Message: "Root mail forwarding is configured",
})
}
// Exim-specific checks only run where exim is the delivery agent.
// Reporting them on a postfix host produced both phantom warnings
// and a fabricated pass for a cPanel-only override file.
if detectMTA() == platform.MTAExim {
// Get exim config for multiple checks
eximOut, eximErr := auditRunCmd("exim", "-bP")
eximConfig := ""
if eximErr == nil {
eximConfig = string(eximOut)
}
// mail_exim_logging
if eximConfig != "" {
lower := strings.ToLower(eximConfig)
if strings.Contains(lower, "+arguments") || strings.Contains(lower, "+all") {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_exim_logging", Title: "Exim Argument Logging",
Status: "pass", Message: "Exim logs include +arguments",
})
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_exim_logging", Title: "Exim Argument Logging",
Status: "warn", Message: "Exim log_selector does not include +arguments",
Fix: "Add '+arguments' to log_selector in exim configuration for better forensics.",
})
}
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_exim_logging", Title: "Exim Argument Logging",
Status: "warn", Message: "Cannot query exim configuration",
})
}
// mail_exim_tls: check for SSLv2 in tls_require_ciphers
// +no_sslv2 in openssl_options means SSLv2 is DISABLED (good).
// Only flag if SSLv2 is referenced WITHOUT +no_ prefix.
if eximConfig != "" {
lower := strings.ToLower(eximConfig)
hasSslv2 := strings.Contains(lower, "sslv2")
isDisabled := strings.Contains(lower, "+no_sslv2") || strings.Contains(lower, "no_sslv2")
if hasSslv2 && !isDisabled {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_exim_tls", Title: "Exim TLS Ciphers",
Status: "fail", Message: "Exim allows SSLv2 connections",
Fix: "Add '+no_sslv2' to openssl_options in exim configuration.",
})
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_exim_tls", Title: "Exim TLS Ciphers",
Status: "pass", Message: "SSLv2 is disabled in exim TLS configuration",
})
}
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_exim_tls", Title: "Exim TLS Ciphers",
Status: "warn", Message: "Cannot query exim TLS configuration",
})
}
// require_secure_auth is a cPanel-managed Exim setting. Other Exim
// packages do not use this file or option.
if isCPanelHost() {
if data, err := osFS.ReadFile("/etc/exim.conf.localopts"); err == nil {
disabled, valid := cpanelSecureAuthDisabled(data)
switch {
case !valid:
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_secure_auth", Title: "Exim Secure Authentication",
Status: "warn", Message: "Cannot evaluate require_secure_auth in /etc/exim.conf.localopts",
})
case disabled:
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_secure_auth", Title: "Exim Secure Authentication",
Status: "fail", Message: "require_secure_auth is disabled in /etc/exim.conf.localopts",
Fix: "Remove or set require_secure_auth=1 in /etc/exim.conf.localopts.",
})
default:
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_secure_auth", Title: "Exim Secure Authentication",
Status: "pass", Message: "Secure authentication is not disabled",
})
}
} else if errors.Is(err, os.ErrNotExist) {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_secure_auth", Title: "Exim Secure Authentication",
Status: "pass", Message: "No local exim overrides file found (default is secure)",
})
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_secure_auth", Title: "Exim Secure Authentication",
Status: "warn", Message: "Cannot read /etc/exim.conf.localopts",
})
}
}
}
if detectMTA() == platform.MTAPostfix {
results = append(results, auditPostfix()...)
}
// mail_dovecot_tls: check ssl_min_protocol
// Use 'doveconf -a' for the effective config — cPanel manages Dovecot
// settings outside the standard config files, so file parsing misses
// them. Routed through cmdExec so tests can mock the doveconf output.
dovecotTLS := false
if out, err := cmdExec.Run("doveconf", "-a"); err == nil {
for _, line := range strings.Split(string(out), "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "ssl_min_protocol") {
val := strings.TrimSpace(strings.TrimPrefix(trimmed, "ssl_min_protocol"))
val = strings.TrimLeft(val, "= ")
if strings.Contains(val, "TLSv1.2") || strings.Contains(val, "TLSv1.3") {
dovecotTLS = true
}
}
}
} else {
// Fallback: try config files
for _, path := range []string{"/etc/dovecot/conf.d/10-ssl.conf", "/etc/dovecot/dovecot.conf"} {
data, readErr := osFS.ReadFile(path)
if readErr != nil {
continue
}
for _, line := range strings.Split(string(data), "\n") {
trimmed := strings.TrimSpace(line)
if strings.HasPrefix(trimmed, "#") {
continue
}
if strings.HasPrefix(trimmed, "ssl_min_protocol") {
val := strings.TrimSpace(strings.TrimPrefix(trimmed, "ssl_min_protocol"))
val = strings.TrimLeft(val, "= ")
if strings.Contains(val, "TLSv1.2") || strings.Contains(val, "TLSv1.3") {
dovecotTLS = true
}
}
}
}
}
if dovecotTLS {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_dovecot_tls", Title: "Dovecot TLS Minimum",
Status: "pass", Message: "Dovecot ssl_min_protocol is TLSv1.2 or higher",
})
} else {
if _, err := osFS.Stat("/etc/dovecot/dovecot.conf"); err != nil {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_dovecot_tls", Title: "Dovecot TLS Minimum",
Status: "warn", Message: "Dovecot configuration not found",
})
} else {
results = append(results, store.AuditResult{
Category: "mail", Name: "mail_dovecot_tls", Title: "Dovecot TLS Minimum",
Status: "fail", Message: "Dovecot ssl_min_protocol not set to TLSv1.2 or higher",
Fix: "Set 'ssl_min_protocol = TLSv1.2' in /etc/dovecot/conf.d/10-ssl.conf.",
})
}
}
return results
}
func cpanelSecureAuthDisabled(data []byte) (disabled, valid bool) {
valid = true
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
assignment, _, _ := strings.Cut(line, "#")
key, value, found := strings.Cut(assignment, "=")
if !found || !strings.EqualFold(strings.TrimSpace(key), "require_secure_auth") {
continue
}
switch strings.TrimSpace(value) {
case "0":
disabled, valid = true, true
case "1":
disabled, valid = false, true
default:
disabled, valid = false, false
}
}
return disabled, valid
}
func webUsersWithNobody(users []string) []string {
if slices.Contains(users, "nobody") {
return users
}
return append(append([]string(nil), users...), "nobody")
}
// auditWebUserCrontab reports a crontab with content for any web-server
// user. The result keeps the historical os_nobody_cron name so stored
// audits and dashboards stay comparable across platforms.
func auditWebUserCrontab(spoolDir string, users []string) store.AuditResult {
for _, user := range users {
info, err := osFS.Stat(filepath.Join(spoolDir, user))
if err != nil || info.Size() == 0 {
continue
}
return store.AuditResult{
Category: "os", Name: "os_nobody_cron", Title: "Web User Crontab",
Status: "fail", Message: fmt.Sprintf("web server user %s has a crontab with content", user),
Fix: fmt.Sprintf("Review and remove: crontab -u %s -r", user),
}
}
return store.AuditResult{
Category: "os", Name: "os_nobody_cron", Title: "Web User Crontab",
Status: "pass", Message: fmt.Sprintf("No crontab with content for web server user(s) %s", strings.Join(users, ", ")),
}
}
// checkFirewallDefaultPolicy audits the default policy of the chain that
// filters inbound traffic. Only an nft chain on the input hook counts: a
// Docker host carries "policy drop" on its FORWARD chain while INPUT stays
// at accept, and that used to pass this audit.
func checkFirewallDefaultPolicy(hasNft bool, nftRules string, hasIpt bool, iptRules string) store.AuditResult {
defaultDeny := hasNft && nftInputDefaultDeny(nftRules)
if !defaultDeny && hasIpt {
for _, line := range strings.Split(iptRules, "\n") {
if strings.HasPrefix(line, "Chain INPUT") {
upper := strings.ToUpper(line)
if strings.Contains(upper, "POLICY DROP") || strings.Contains(upper, "POLICY REJECT") {
defaultDeny = true
}
break
}
}
}
if defaultDeny {
return store.AuditResult{
Category: "firewall", Name: "fw_default_policy", Title: "Default INPUT Policy",
Status: "pass", Message: "INPUT chain has default-deny policy",
}
}
return store.AuditResult{
Category: "firewall", Name: "fw_default_policy", Title: "Default INPUT Policy",
Status: "fail", Message: "INPUT chain does not have a DROP/REJECT policy",
Fix: "Set the default INPUT policy to DROP: iptables -P INPUT DROP (or nft equivalent).",
}
}
// nftInputDefaultDeny reports whether an nft ruleset has a drop or reject
// policy on a chain hooked at input. nft prints the hook and the policy on
// one line; a policy on a following line inside the same chain also counts.
func nftInputDefaultDeny(nftRules string) bool {
inChain := false
chainDepth := 0
inputHook := false
denyPolicy := false
for _, raw := range strings.Split(nftRules, "\n") {
line := strings.ToLower(strings.TrimSpace(nftCodeOnly(raw)))
opens := strings.Count(line, "{")
closes := strings.Count(line, "}")
if strings.HasPrefix(line, "chain ") {
inChain = true
chainDepth = 0
inputHook = false
denyPolicy = false
}
if !inChain {
continue
}
inputHook = inputHook || nftHasWordPair(line, "hook", "input")
denyPolicy = denyPolicy || nftHasWordPair(line, "policy", "drop") || nftHasWordPair(line, "policy", "reject")
if inputHook && denyPolicy {
return true
}
chainDepth += opens - closes
if chainDepth <= 0 && closes > 0 {
inChain = false
}
}
return false
}
func nftCodeOnly(line string) string {
var b strings.Builder
var quote byte
escaped := false
for i := 0; i < len(line); i++ {
c := line[i]
if quote != 0 {
if escaped {
escaped = false
continue
}
if c == '\\' {
escaped = true
continue
}
if c == quote {
quote = 0
}
continue
}
if c == '#' {
break
}
if c == '\'' || c == '"' {
quote = c
b.WriteByte(' ')
continue
}
b.WriteByte(c)
}
return b.String()
}
func nftHasWordPair(line, first, second string) bool {
words := strings.FieldsFunc(line, func(r rune) bool {
return (r < 'a' || r > 'z') && (r < '0' || r > '9') && r != '_'
})
for i := 1; i < len(words); i++ {
if words[i-1] == first && words[i] == second {
return true
}
}
return false
}
package checks
import (
"context"
"fmt"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// CheckHealth verifies that CSM's dependencies are working.
// Reports on missing external commands, broken auditd, etc.
func CheckHealth(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
info := platform.Detect()
// Required commands depend on the platform. On plain Linux hosts we
// don't need Exim/cPanel-specific tooling.
requiredCmds := platformRequiredCommands(info)
for _, cmd := range requiredCmds {
if _, err := cmdExec.LookPath(cmd); err != nil {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "csm_health",
Message: fmt.Sprintf("Required command not found: %s", cmd),
Details: "Some checks will be skipped",
})
}
}
// Optional commands. Only complain about cPanel tools on cPanel hosts.
optionalCmds := map[string]string{
"wp": "WordPress core integrity check will be skipped",
}
if info.IsCPanel() {
optionalCmds["whmapi1"] = "WHM API token check will be skipped"
}
for cmd, impact := range optionalCmds {
if _, err := cmdExec.LookPath(cmd); err != nil {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "csm_health",
Message: fmt.Sprintf("Optional command not found: %s", cmd),
Details: impact,
})
}
}
// Check auditd is running and has CSM rules
out, _ := runCmd("auditctl", "-l")
if out != nil {
rules := string(out)
if !strings.Contains(rules, "csm_shadow_change") {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "csm_health",
Message: "auditd CSM rules not loaded",
Details: "Run 'csm install' to deploy auditd rules, then 'service auditd restart'",
})
}
}
if cfg != nil && cfg.BPFEnforcement.Enabled && cfg.BPFEnforcement.DirectSMTPEgress {
switch active := bpf.ActiveKind("connection_tracker"); active {
case bpf.BackendLegacy, bpf.BackendNone:
message := "BPF enforcement enabled but connection tracker is running on legacy backend"
if active == bpf.BackendNone {
message = "BPF enforcement enabled but connection tracker has no active backend"
}
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "csm_health",
Message: message,
Details: "bpf_enforcement.direct_smtp_egress requires the connection tracker BPF backend. Check kernel version, LSM availability, or CAP_BPF.",
})
}
}
// Check state directory is writable
stateDir := "/var/lib/csm/state"
if cfg != nil && cfg.StatePath != "" {
stateDir = cfg.StatePath
}
testFile := filepath.Join(stateDir, ".health_check")
if err := osFS.WriteFile(testFile, []byte("ok"), 0600); err != nil {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "csm_health",
Message: fmt.Sprintf("State directory not writable: %s", stateDir),
Details: err.Error(),
})
} else {
_ = osFS.Remove(testFile)
}
return findings
}
// platformRequiredCommands returns the external commands CSM needs on the
// detected platform. On plain Linux hosts Exim is not required.
func platformRequiredCommands(info platform.Info) []string {
cmds := []string{"find", "auditctl"}
if info.IsCPanel() {
cmds = append(cmds, "exim")
}
return cmds
}
package checks
import (
"context"
"crypto/sha256"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"strings"
"time"
)
// cmdTimeout bounds one external command. A variable so a test can exercise
// the real deadline without waiting two minutes for it.
var cmdTimeout = 2 * time.Minute
var systemCommandSearchDirs = []string{
"/usr/local/sbin",
"/usr/sbin",
"/sbin",
"/usr/local/bin",
"/usr/bin",
"/bin",
}
func hashFileContent(path string) (string, error) {
data, err := osFS.ReadFile(path)
if err != nil {
return "", err
}
h := sha256.Sum256(data)
return fmt.Sprintf("%x", h[:]), nil
}
func hashBytes(data []byte) string {
h := sha256.Sum256(data)
return fmt.Sprintf("%x", h[:])
}
// runCmd delegates to the package-level cmdExec provider.
// Check functions call runCmd; tests swap cmdExec via SetCmdRunner.
func runCmd(name string, args ...string) ([]byte, error) {
return cmdExec.Run(name, args...)
}
func runCmdAllowNonZero(name string, args ...string) ([]byte, error) {
return cmdExec.RunAllowNonZero(name, args...)
}
func runCmdCombinedContext(parent context.Context, name string, args ...string) ([]byte, error) {
return cmdExec.RunContext(parent, name, args...)
}
func lookupSystemCommand(name string) (string, error) {
if strings.ContainsRune(name, os.PathSeparator) {
return exec.LookPath(name)
}
path, err := exec.LookPath(name)
if err == nil {
return path, nil
}
for _, dir := range systemCommandSearchDirs {
candidate := filepath.Join(dir, name)
info, statErr := os.Stat(candidate)
if statErr == nil && !info.IsDir() && info.Mode()&0111 != 0 {
return candidate, nil
}
}
return "", err
}
func resolveSystemCommand(name string) string {
path, err := lookupSystemCommand(name)
if err != nil {
return name
}
return path
}
// ---------------------------------------------------------------------------
// Real implementations — used by realCmd in provider.go
//
// Every caller of these helpers passes a constant system command name
// (nft, rpm, wp-cli, doveadm, systemctl, etc.) with arguments built from
// either config, filesystem state, or CSM-generated data. Nothing here
// reaches out to HTTP request bodies or webui form inputs. The gosec
// G204 suppressions below are in that trust model.
// ---------------------------------------------------------------------------
func runCmdReal(name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), cmdTimeout)
defer cancel()
// #nosec G204 -- see package-level trust note above.
out, err := exec.CommandContext(ctx, resolveSystemCommand(name), args...).Output()
if ctx.Err() == context.DeadlineExceeded {
fmt.Fprintf(os.Stderr, "Command timed out: %s %v\n", name, args)
return nil, nil
}
return out, err
}
func runCmdAllowNonZeroReal(name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), cmdTimeout)
defer cancel()
// #nosec G204 -- see package-level trust note above.
out, err := exec.CommandContext(ctx, resolveSystemCommand(name), args...).Output()
if ctx.Err() == context.DeadlineExceeded {
fmt.Fprintf(os.Stderr, "Command timed out: %s %v\n", name, args)
return nil, nil
}
var exitErr *exec.ExitError
if err != nil && errors.As(err, &exitErr) {
return out, nil
}
return out, err
}
// commandRefused reports whether a command ran to completion and answered with
// a failure, such as wp-cli on a tree that is not a WordPress installation or
// one whose wp-config.php fatals. The check has its answer, so the queue lost
// no work. A command that never started, was killed by a signal or ran out of
// time is still lost work.
func commandRefused(err error) bool {
var exit *exec.ExitError
return errors.As(err, &exit) && exit.ExitCode() > 0
}
func runCmdCombinedContextReal(parent context.Context, name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(parent, cmdTimeout)
defer cancel()
// #nosec G204 -- see package-level trust note above.
out, err := exec.CommandContext(ctx, resolveSystemCommand(name), args...).CombinedOutput()
if ctx.Err() == context.DeadlineExceeded {
fmt.Fprintf(os.Stderr, "Command timed out: %s %v\n", name, args)
return nil, context.DeadlineExceeded
}
if parent.Err() != nil {
return nil, parent.Err()
}
return out, err
}
// runCmdStdoutContextReal runs a command with a per-call timeout and returns
// stdout only. Stderr is discarded so chatter from the child process (PHP
// warnings, MySQL deprecation notices, wp-cli plugin backtraces, ...) cannot
// poison parsers that expect JSON/URL bytes on stdout. On timeout the caller
// receives context.DeadlineExceeded rather than a silent (nil, nil), so it
// can distinguish a hung command from a legitimately empty result.
func runCmdStdoutContextReal(parent context.Context, name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(parent, cmdTimeout)
defer cancel()
// #nosec G204 -- see package-level trust note above.
out, err := exec.CommandContext(ctx, resolveSystemCommand(name), args...).Output()
if ctx.Err() == context.DeadlineExceeded {
return nil, context.DeadlineExceeded
}
if parent.Err() != nil {
return nil, parent.Err()
}
return out, err
}
func runCmdWithEnvReal(name string, args []string, extraEnv ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), cmdTimeout)
defer cancel()
// #nosec G204 -- see package-level trust note above.
cmd := exec.CommandContext(ctx, resolveSystemCommand(name), args...)
cmd.Env = append(os.Environ(), extraEnv...)
out, err := cmd.Output()
if ctx.Err() == context.DeadlineExceeded {
fmt.Fprintf(os.Stderr, "Command timed out: %s\n", name)
return nil, nil
}
return out, err
}
package checks
import (
"io"
)
// htaccessMaxFileBytes bounds every scheduled .htaccess read. A real
// .htaccess is a few kilobytes; the ceiling matches the per-line limit the
// directive scanner already enforces. Reading the whole file with no bound
// let a tenant park a multi-gigabyte .htaccess and turn the deep scan into
// an out-of-memory crash loop.
const htaccessMaxFileBytes = htaccessMaxLineBytes
// readHtaccessBounded returns the file's content when it is at most
// htaccessMaxFileBytes. An oversized file yields ok=false with no error: the
// caller knows the file exists but cannot be judged, which differs from a
// missing file. Any other failure is returned as err.
func readHtaccessBounded(path string) (data []byte, ok bool, err error) {
f, err := osFS.Open(path)
if err != nil {
return nil, false, err
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil, false, err
}
if !info.Mode().IsRegular() || info.Size() > htaccessMaxFileBytes {
return nil, false, nil
}
// Read one byte past the limit: a file that grows between Stat and
// Read must not slip through as complete.
data, err = io.ReadAll(io.LimitReader(f, htaccessMaxFileBytes+1))
if err != nil {
return nil, false, err
}
if int64(len(data)) > htaccessMaxFileBytes {
return nil, false, nil
}
return data, true, nil
}
// htaccessOversized reports whether err/ok from readHtaccessBounded mean
// "present but too large".
func htaccessOversized(ok bool, err error) bool {
return !ok && err == nil
}
package checks
import (
"errors"
"fmt"
"io"
"net/url"
"os"
"path/filepath"
"regexp"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/alert"
)
// .htaccess hardened detection / cleaning.
//
// Detection emits a specific finding name for each pattern so operators
// can suppress, route, or auto-respond per attack pattern instead of
// relying on the generic htaccess_injection / htaccess_handler_abuse
// categories. Cleaning is gated by AutoResponse.CleanHtaccess and is
// always backed up under /opt/csm/quarantine/pre_clean/.
//
// Each detector returns matches as byte ranges into the file
// content; the cleaner merges all ranges (deduplicating overlaps),
// removes them, and writes the result atomically. If post-clean
// content is identical to pre-clean (no detector matched anything
// new), no write happens and no backup is created.
// htaccessBackupDirRoot is the parent directory under which
// CleanHtaccessFile writes unique recovery backups. Exposed
// as a package var so tests can redirect it to a t.TempDir().
var htaccessBackupDirRoot = "/opt/csm/quarantine/pre_clean"
// htaccessByteRange is a half-open [start, end) byte slice into the
// file content. Cleaning removes the bytes; the line ending after
// `end` is included when `end` falls just before a `\n` so we do
// not leave a blank line behind.
type htaccessByteRange struct {
Start int
End int
}
// htaccessMatch is one finding-worthy hit returned by a detector.
type htaccessMatch struct {
Range htaccessByteRange
Excerpt string // the offending line(s), trimmed for the finding details
// Severity overrides the detector's default for this match. Nil keeps the
// default. Used where one detector covers directives of differing effect.
Severity *alert.Severity
// Retain keeps the match out of the cleaner's removal set. A match that is
// reported for visibility but has no effect on the server must not cause an
// edit to a customer's file.
Retain bool
}
// htaccessDetector pairs a finding category with a function that
// scans the content and reports every range the detector wants
// removed. Adding a pattern is one entry in this slice.
type htaccessDetector struct {
Name string
Severity alert.Severity
Detect func(content []byte, path string) []htaccessMatch
}
// htaccessSpamTLDs lists TLDs commonly abused for spam-redirect
// .htaccess injections. Operators on legitimate hosts in these TLDs
// can suppress the finding by file path.
var htaccessSpamTLDs = []string{
".xyz", ".tk", ".ml", ".ga", ".cf", ".gq", ".click",
".country", ".loan", ".work", ".top",
}
// htaccessNonScriptDirHints names directory components where PHP
// execution is rarely legitimate. .htaccess files inside one of
// these get the htaccess_php_in_uploads finding when they map
// non-PHP extensions to a PHP handler.
//
// /tmp/ is intentionally NOT in this list even though attackers
// drop payloads there: a real-world .htaccess inside /tmp/ would
// only be reached if the webserver served /tmp/, which is rare
// outside misconfigurations -- and including it caused
// false-positive matches against Linux t.TempDir() paths under
// /tmp/ in the test suite. The auto_prepend detector covers the
// /tmp/ payload-target angle separately.
var htaccessNonScriptDirHints = []string{
"/uploads/", "/images/", "/cache/",
"/wp-content/uploads/", "/wp-content/cache/",
"/files/", "/media/",
}
// htaccessSuspiciousAutoPrependPaths are scratch locations that a prelude
// script is never legitimately served from.
var htaccessSuspiciousAutoPrependPaths = []string{
"/tmp/", "/dev/shm/", "/var/tmp/",
}
// Apache accepts quoted directive arguments, including paths with spaces.
// Keep the quotes in the capture; autoPrependTargetSuspicious removes them
// after the parser has found the complete target.
const htaccessPreludeTargetPattern = `("[^"\r\n]*"|'[^'\r\n]*'|\S+)`
// A raw newline ends the directive. Apache continuation is explicit, so only
// horizontal whitespace or a backslash-newline may separate its arguments.
const htaccessPreludeSeparatorPattern = `(?:[\t ]|\\\r?\n)+`
// reAutoPrependTarget captures the argument of either prelude directive in
// any of the forms .htaccess and php.ini fragments use.
var reAutoPrependTarget = regexp.MustCompile(`(?i)auto_(?:prepend|append)_file(?:[\t ]*=[\t ]*|[\t ]+)` + htaccessPreludeTargetPattern)
// autoPrependTargetIsKnownPrelude reports whether target names a prelude
// script shipped by a security plugin. Only the basename is consulted: every
// other byte of the directive is text the account owner types, so a
// substring test anywhere else on the line is an exemption the attacker
// controls.
func autoPrependTargetIsKnownPrelude(target string) bool {
base := strings.ToLower(filepath.Base(strings.ReplaceAll(target, `\`, "/")))
switch base {
case "wordfence-waf.php", "advanced-headers.php", "malcare-waf.php":
return true
}
return strings.HasPrefix(base, "sucuri") && strings.HasSuffix(base, ".php")
}
// autoPrependTargetSuspicious reports whether an auto_prepend_file or
// auto_append_file target can point at code the account owner controls.
// htaccessPath is the file carrying the directive. Unless the target is a
// known plugin prelude, it is suspicious when it sits in a scratch location,
// is not a PHP file at all, is relative (it resolves inside the docroot),
// lives under any home directory, or shares the .htaccess file's own account
// tree. A root-owned path elsewhere (/etc, /opt, /usr) needs root to write
// and is left alone; "none" merely disables an inherited prelude.
func autoPrependTargetSuspicious(target, htaccessPath string) bool {
target = strings.Trim(strings.TrimSpace(target), `"'`)
lower := strings.ToLower(target)
if lower == "" || lower == "none" || autoPrependTargetIsKnownPrelude(lower) {
return false
}
// PHP resolves lexical dot segments before opening the file. Classify the
// same normalized path so an account-controlled target cannot hide behind
// an apparently root-owned prefix such as /etc/../home/user/prelude.php.
lower = strings.ToLower(filepath.Clean(target))
for _, p := range htaccessSuspiciousAutoPrependPaths {
if strings.HasPrefix(lower, p) {
return true
}
}
if !strings.HasSuffix(lower, ".php") || !strings.HasPrefix(lower, "/") || underAccountRoot(lower) {
return true
}
if tree := htaccessAccountTree(htaccessPath); tree != "" && strings.HasPrefix(lower, strings.ToLower(tree)+"/") {
return true
}
return false
}
// htaccessAccountTree returns the first two path components of the .htaccess
// location (/home/<user>, /var/www/<vhosts>), which is the tree the account
// can write to under every supported panel layout.
func htaccessAccountTree(htaccessPath string) string {
parts := strings.Split(filepath.Clean(htaccessPath), "/")
if len(parts) < 4 || parts[0] != "" {
return ""
}
return "/" + parts[1] + "/" + parts[2]
}
// htaccessTrackingHeaders is a small allowlist of header *names*
// known to be used in injection campaigns. Scoped intentionally;
// false positives on legitimate analytics/CDN headers are worse
// than missed detections here.
var htaccessTrackingHeaders = []string{
"X-Track-", "X-Affiliate-", "X-Promo-", "X-Click-ID",
}
var (
// SetHandler / ForceType take effect with a single argument: they map
// EVERY file in the directory to the PHP interpreter, so an uploaded
// image runs as PHP. That directory-wide form is the worst case and has
// no extension list, so the second argument is optional for them.
// AddHandler is inert without an extension list, so it still requires the
// trailing token to be a genuine remap.
rePHPHandlerMap = regexp.MustCompile(`(?im)^\s*(?:(?:SetHandler|ForceType)\s+\S*php\S*(?:\s+\S[^\n]*)?|AddHandler\s+\S*php\S*\s+\S[^\n]*)\s*$`)
// Match both forms because mod_php and some LSAPI builds honor either
// directive in .htaccess.
reAutoPrepend = regexp.MustCompile(`(?im)^[\t ]*php(?:_admin)?_value[\t ]+auto_(?:prepend|append)_file` + htaccessPreludeSeparatorPattern + htaccessPreludeTargetPattern)
reUACloakCond = regexp.MustCompile(`(?im)^\s*RewriteCond\s+%\{HTTP_USER_AGENT\}\s+([^\n]+)`)
reSpamRedirect = regexp.MustCompile(`(?im)^\s*RewriteRule\s+\S+\s+(https?://[^\s\[]+)`)
reFilesMatchOpen = regexp.MustCompile(`(?im)^\s*<FilesMatch\s+["']?[^"'>]*\\\.(php|phtml|ph[2-7])[^"'>]*["']?\s*>`)
reFilesMatchClose = regexp.MustCompile(`(?im)^\s*</FilesMatch>`)
reHeaderSetAdd = regexp.MustCompile(`(?im)^\s*Header\s+(set|add)\s+([A-Za-z0-9_-]+)`)
reErrorDocument = regexp.MustCompile(`(?im)^\s*ErrorDocument\s+\d+\s+(https?://[^\s]+)`)
// crawlerUARegex matches the UA strings frequently used in cloak
// conditions: search-engine bots and social-share scrapers. Used
// as a positive filter on the htaccess_user_agent_cloak finding;
// matching one of these names is what makes a UA-keyed redirect
// suspicious.
crawlerUARegex = regexp.MustCompile(`(?i)(googlebot|bingbot|baiduspider|yandex|facebookexternalhit|slurp|duckduckbot)`)
// searchCrawlerUARegex names the indexers a cloak targets. Operator
// blocklists are made of scrapers and SEO tools; a list whose members
// are mostly search engines is a cloak whatever its length, because no
// site blocks Googlebot, Bingbot, Yandex and Baidu together.
searchCrawlerUARegex = regexp.MustCompile(`(?i)googlebot|bingbot|baiduspider|yandex|slurp|duckduckbot|applebot|sogou|seznambot|petalbot`)
)
// htaccessDetectors is the registry. Detectors run in slice order for
// deterministic finding emission. Overlapping removal ranges are merged later.
var htaccessDetectors = []htaccessDetector{
{
Name: "htaccess_php_in_uploads",
Severity: alert.Critical,
Detect: detectPHPInUploads,
},
{
Name: "htaccess_auto_prepend",
Severity: alert.Critical,
Detect: detectAutoPrepend,
},
{
Name: "htaccess_user_agent_cloak",
Severity: alert.High,
Detect: detectUserAgentCloak,
},
{
Name: "htaccess_spam_redirect",
Severity: alert.High,
Detect: detectSpamRedirect,
},
{
Name: "htaccess_filesmatch_shield",
Severity: alert.Critical,
Detect: detectFilesMatchShield,
},
{
Name: "htaccess_header_injection",
Severity: alert.High,
Detect: detectHeaderInjection,
},
{
Name: "htaccess_errordocument_hijack",
Severity: alert.High,
Detect: detectErrorDocumentHijack,
},
{
Name: "htaccess_cgi_handler_abuse",
Severity: alert.Critical,
Detect: detectCGIHandlerAbuse,
},
{
Name: "htaccess_security_disabled",
Severity: alert.High,
Detect: detectSecurityDisabled,
},
}
func htaccessDetectorNames() []string {
names := make([]string, 0, len(htaccessDetectors))
for _, detector := range htaccessDetectors {
names = append(names, detector.Name)
}
return names
}
// handlerIsCGI reports whether an Apache handler/MIME token routes matching
// files to a CGI interpreter (as opposed to serving them statically). This is
// the CGI counterpart of handlerIsPHP.
func handlerIsCGI(token string) bool {
token = strings.ToLower(strings.Trim(strings.TrimSpace(token), `"'`))
switch token {
case "cgi-script",
"fcgid-script",
"fastcgi-script",
"application/x-httpd-cgi",
"application/x-httpd-fcgi":
return true
}
return false
}
// conventionalCGIExt reports whether ext (leading dot, lowercase) is one a
// legitimate cgi-bin routes to a CGI handler. A CGI handler mapped onto any
// other extension is the webshell-arming remap.
func conventionalCGIExt(ext string) bool {
switch ext {
case ".cgi", ".fcg", ".fcgi", ".pl", ".py", ".rb":
return true
}
return false
}
type cgiHandlerContext struct {
onlyConventional bool
}
func openCGIHandlerContext(line string) (cgiHandlerContext, bool) {
lower := strings.ToLower(strings.TrimSpace(line))
switch {
case strings.HasPrefix(lower, "<filesmatch"):
return cgiHandlerContext{
onlyConventional: filesMatchTargetsOnlyConventionalCGI(apacheContainerArgument(line)),
}, true
case strings.HasPrefix(lower, "<files"):
name := strings.ToLower(strings.TrimSpace(apacheContainerArgument(line)))
if name == "" || strings.Contains(name, "/") {
return cgiHandlerContext{}, true
}
if strings.HasPrefix(name, "~") {
return cgiHandlerContext{
onlyConventional: filesMatchTargetsOnlyConventionalCGI(name),
}, true
}
ext := normalizeExt(filepath.Ext(name))
if ext == "" {
return cgiHandlerContext{}, true
}
return cgiHandlerContext{
onlyConventional: conventionalCGIExt(ext),
}, true
default:
return cgiHandlerContext{}, false
}
}
func filesMatchTargetsOnlyConventionalCGI(pattern string) bool {
pattern = strings.TrimSpace(pattern)
if strings.HasPrefix(pattern, "~") {
pattern = strings.TrimSpace(strings.TrimPrefix(pattern, "~"))
}
if len(pattern) >= 2 {
quote := pattern[0]
if (quote == '"' || quote == '\'') && pattern[len(pattern)-1] == quote {
pattern = pattern[1 : len(pattern)-1]
}
}
pattern = strings.ToLower(strings.TrimSpace(pattern))
switch {
case strings.HasSuffix(pattern, `\z`):
pattern = pattern[:len(pattern)-2]
case strings.HasSuffix(pattern, "$"):
pattern = pattern[:len(pattern)-1]
default:
// An unanchored "\.cgi" also matches names such as shell.cgi.jpg.
return false
}
dot := lastRegexLiteralDot(pattern)
if dot < 0 || strings.Contains(pattern[:dot], "|") {
return false
}
extPattern := pattern[dot+2:]
switch {
case strings.HasPrefix(extPattern, "(?:") && strings.HasSuffix(extPattern, ")"):
extPattern = extPattern[3 : len(extPattern)-1]
case strings.HasPrefix(extPattern, "(") && strings.HasSuffix(extPattern, ")"):
extPattern = extPattern[1 : len(extPattern)-1]
}
parts := strings.Split(extPattern, "|")
exts := make([]string, 0, len(parts))
for _, part := range parts {
ext := normalizeExt(part)
if ext == "" {
return false
}
exts = append(exts, ext)
}
return allConventionalCGIExts(exts)
}
// lastRegexLiteralDot returns the escape backslash before the final literal
// dot. An odd backslash run escapes the dot; an even run leaves it as a regex
// wildcard and cannot prove an extension-only match.
func lastRegexLiteralDot(pattern string) int {
for dot := len(pattern) - 1; dot >= 0; dot-- {
if pattern[dot] != '.' {
continue
}
backslashes := 0
for i := dot - 1; i >= 0 && pattern[i] == '\\'; i-- {
backslashes++
}
if backslashes%2 == 1 {
return dot - 1
}
}
return -1
}
func allConventionalCGIExts(exts []string) bool {
if len(exts) == 0 {
return false
}
for _, ext := range exts {
if !conventionalCGIExt(ext) {
return false
}
}
return true
}
func cgiHandlerTargetsOnlyConventional(contexts []cgiHandlerContext) bool {
for _, context := range contexts {
// Nested file containers intersect. Any conventional-only container
// therefore prevents the handler from reaching a custom extension even
// when another container is broader.
if context.onlyConventional {
return true
}
}
return false
}
func cgiExtensionListIsSuspicious(tokens []string) bool {
for _, token := range tokens {
ext := normalizeExt(token)
// AddHandler and AddType treat every remaining token as an
// extension. If the shared plain-extension parser cannot normalize
// one, it is still a custom target and must fail closed.
if ext == "" || !conventionalCGIExt(ext) {
return true
}
}
return false
}
type cgiOptionsScope struct {
name string
neutralized bool
childMayEnable bool
}
func apacheContainerTag(line string) (name string, closing bool, ok bool) {
line = strings.TrimSpace(line)
if len(line) < 3 || line[0] != '<' {
return "", false, false
}
end := strings.IndexByte(line, '>')
if end < 0 {
return "", false, false
}
tag := strings.TrimSpace(line[1:end])
if tag == "" || strings.HasPrefix(tag, "!") {
return "", false, false
}
if strings.HasPrefix(tag, "/") {
closing = true
tag = strings.TrimSpace(tag[1:])
}
if fields := strings.Fields(tag); len(fields) > 0 {
return strings.ToLower(fields[0]), closing, true
}
return "", false, false
}
func unquoteApacheOption(token string) (string, bool) {
token = strings.TrimSpace(token)
if token == "" {
return "", false
}
if token[0] == '"' || token[0] == '\'' {
if len(token) < 2 || token[len(token)-1] != token[0] {
return "", false
}
token = token[1 : len(token)-1]
}
if token == "" || strings.ContainsAny(token, `"'`) {
return "", false
}
return strings.ToLower(token), true
}
// applyCGIOptions follows Apache's Options parser for the ExecCGI bit. Bare
// option lists replace the inherited set, while lists made only of +/- tokens
// merge with it. All and None may start a list before relative adjustments.
func applyCGIOptions(fields []string, neutralized bool) (bool, bool) {
if len(fields) < 2 {
return false, false
}
first := true
merge := false
allOrNone := false
for _, field := range fields[1:] {
token, ok := unquoteApacheOption(field)
if !ok {
return false, false
}
var action byte
switch {
case token[0] == '+' || token[0] == '-':
action = token[0]
token = token[1:]
if token == "" || (!merge && !first && !allOrNone) {
return false, false
}
merge = true
case first:
neutralized = true
case merge:
return false, false
}
switch token {
case "none":
if !first || merge {
return false, false
}
neutralized = true
allOrNone = true
case "all":
if !first || merge {
return false, false
}
neutralized = false
allOrNone = true
case "execcgi", "runscripts":
neutralized = action == '-'
case "indexes", "includes", "includesnoexec", "followsymlinks",
"symlinksifownermatch", "multiviews":
default:
return false, false
}
first = false
}
return neutralized, true
}
// cgiExecutionNeutralized reports whether CGI execution is disabled for every
// request in this directory. Conditional or scoped Options blocks cannot prove
// a global disable and may re-enable ExecCGI after directory options merge.
func cgiExecutionNeutralized(content []byte) bool {
neutralized := false
nestedMayEnable := false
var scopes []cgiOptionsScope
for _, logical := range htaccessLogicalByteLines(content) {
line := strings.TrimSpace(logical.text)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if name, closing, ok := apacheContainerTag(line); ok {
if !closing {
scopes = append(scopes, cgiOptionsScope{
name: name,
neutralized: true,
})
continue
}
if len(scopes) == 0 || scopes[len(scopes)-1].name != name {
return false
}
scope := scopes[len(scopes)-1]
scopes = scopes[:len(scopes)-1]
mayEnable := !scope.neutralized || scope.childMayEnable
if len(scopes) == 0 {
nestedMayEnable = nestedMayEnable || mayEnable
} else {
scopes[len(scopes)-1].childMayEnable =
scopes[len(scopes)-1].childMayEnable || mayEnable
}
continue
}
fields := apacheDirectiveFields(line)
if len(fields) == 0 || !strings.EqualFold(fields[0], "Options") {
continue
}
var valid bool
if len(scopes) == 0 {
neutralized, valid = applyCGIOptions(fields, neutralized)
} else {
scope := &scopes[len(scopes)-1]
scope.neutralized, valid = applyCGIOptions(fields, scope.neutralized)
}
if !valid {
return false
}
}
return len(scopes) == 0 && neutralized && !nestedMayEnable
}
// detectCGIHandlerAbuse flags an .htaccess that maps a non-conventional
// extension (or the whole directory) to a CGI interpreter. An attacker uses
// this to make an uploaded Perl/binary file execute: e.g.
// `AddHandler cgi-script .alfa` / `AddType application/x-httpd-cgi .alfa`.
func detectCGIHandlerAbuse(content []byte, _ string) []htaccessMatch {
// A cgi-script handler mapping is inert when the directory disables CGI
// execution, so a hardening block that maps .php to cgi-script alongside
// Options -ExecCGI is not webshell arming and must not be flagged or
// auto-cleaned out of a legitimate security plugin's .htaccess.
if cgiExecutionNeutralized(content) {
return nil
}
var out []htaccessMatch
var contexts []cgiHandlerContext
for _, logical := range htaccessLogicalByteLines(content) {
line := strings.TrimSpace(logical.text)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if context, ok := openCGIHandlerContext(line); ok {
contexts = append(contexts, context)
continue
}
if closesPHPHandlerContext(line) {
if len(contexts) > 0 {
contexts = contexts[:len(contexts)-1]
}
continue
}
fields := apacheDirectiveFields(line)
if len(fields) < 2 || !handlerIsCGI(fields[1]) {
continue
}
var flagged bool
switch strings.ToLower(fields[0]) {
case "addhandler", "addtype":
flagged = cgiExtensionListIsSuspicious(fields[2:])
case "sethandler", "forcetype":
flagged = !cgiHandlerTargetsOnlyConventional(contexts)
}
if flagged {
out = append(out, htaccessMatch{
Range: logical.span,
Excerpt: trimExcerpt(content, logical.span.Start, logical.span.End),
})
}
}
return out
}
func normalizeModSecurityDirective(name string) string {
return strings.Map(func(r rune) rune {
switch {
case r >= 'A' && r <= 'Z':
return r + ('a' - 'A')
case r >= 'a' && r <= 'z', r >= '0' && r <= '9':
return r
default:
return -1
}
}, name)
}
func modSecurityDirectiveDisablesWAF(name, value string) bool {
name = normalizeModSecurityDirective(name)
value = strings.ToLower(strings.Trim(strings.TrimSpace(value), `"'`))
switch value {
case "off":
switch name {
case "secruleengine",
"secfilterengine",
"secfilterscanpost",
"secruleinheritance",
"secrequestbodyaccess",
"secengine",
"secscanpost":
return true
}
case "detectiononly":
switch name {
case "secruleengine", "secfilterengine", "secengine":
return true
}
}
return false
}
// detectSecurityDisabled flags a per-account .htaccess that disables
// ModSecurity enforcement, request inspection, or inherited rules. A customer
// .htaccess never legitimately weakens those controls; attackers do it to mask
// an intrusion. Punctuation is removed from directive names so known
// obfuscated spellings are classified without matching unrelated settings such
// as SecAuditEngine or SecStatusEngine.
func detectSecurityDisabled(content []byte, _ string) []htaccessMatch {
var out []htaccessMatch
warning := alert.Warning
for _, logical := range htaccessLogicalByteLines(content) {
fields := apacheDirectiveFields(strings.TrimSpace(logical.text))
if len(fields) < 2 || !modSecurityDirectiveDisablesWAF(fields[0], fields[1]) {
continue
}
match := htaccessMatch{
Range: logical.span,
Excerpt: trimExcerpt(content, logical.span.Start, logical.span.End),
}
if legacyModSecurityDirective(fields[0]) {
match.Severity = &warning
match.Retain = true
}
out = append(out, match)
}
return out
}
// legacyModSecurityDirective reports whether a directive belongs to
// mod_security 1.x, which no supported server reads. Apache runs
// mod_security2 or mod_security3, LiteSpeed reads its own equivalent, and
// Nginx has neither, so these lines change nothing wherever CSM runs. They are
// still worth surfacing -- an intruder who pastes one is announcing an attempt
// -- but they did not disable anything, and stripping them from a legacy shop
// or Magento .htaccess edits a customer file to no effect.
func legacyModSecurityDirective(name string) bool {
switch normalizeModSecurityDirective(name) {
case "secfilterengine", "secfilterscanpost", "secscanpost":
return true
}
return false
}
// AuditHtaccessFile runs every registered detector against the file
// at path. Returns the alert findings (one per detector hit) and
// the merged byte ranges that the cleaner would remove. The two
// outputs travel together so cleaning never disagrees with what
// the operator was alerted about.
func AuditHtaccessFile(path string) ([]alert.Finding, []htaccessByteRange) {
findings, ranges, _ := auditHtaccessFile(path)
return findings, ranges
}
func auditHtaccessFile(path string) ([]alert.Finding, []htaccessByteRange, bool) {
if filepath.Base(path) != ".htaccess" {
return nil, nil, true
}
content, ok, err := readHtaccessBounded(path)
if htaccessOversized(ok, err) {
return []alert.Finding{{
Severity: alert.High,
Check: "htaccess_injection",
Message: fmt.Sprintf(".htaccess too large to audit: %s", path),
Details: fmt.Sprintf("File exceeds %d bytes; a real .htaccess is a few kilobytes. Inspect it by hand.", htaccessMaxFileBytes),
FilePath: path,
Timestamp: time.Now(),
}}, nil, false
}
if err != nil || !ok {
return nil, nil, os.IsNotExist(err)
}
findings, ranges := AuditHtaccessContent(path, content)
return findings, ranges, true
}
func AuditHtaccessContent(path string, content []byte) ([]alert.Finding, []htaccessByteRange) {
var findings []alert.Finding
var ranges []htaccessByteRange
for _, d := range htaccessDetectors {
matches := d.Detect(content, path)
for _, m := range matches {
severity := d.Severity
if m.Severity != nil {
severity = *m.Severity
}
findings = append(findings, alert.Finding{
Severity: severity,
Check: d.Name,
Message: fmt.Sprintf("%s in %s", d.Name, path),
Details: fmt.Sprintf("File: %s\nMatch: %s", path, m.Excerpt),
FilePath: path,
Timestamp: time.Now(),
})
if !m.Retain {
ranges = append(ranges, m.Range)
}
}
}
return findings, mergeRanges(ranges)
}
// CleanHtaccessFile audits the file, computes the removal range
// set, backs up the original, and writes the trimmed content.
// Returns success=false with no Action when no detector matched
// (i.e., nothing to clean).
//
// Caller is responsible for gating on cfg.AutoResponse.CleanHtaccess
// before invoking; this function will clean unconditionally if
// detectors find anything.
func CleanHtaccessFile(path string) RemediationResult {
return cleanHtaccessFileIdentified(path, nil)
}
func cleanHtaccessFileIdentified(path string, expected os.FileInfo) (result RemediationResult) {
audit := newCleanAction(path)
defer func() { audit.finish(result.Error) }()
if filepath.Base(path) != ".htaccess" {
return RemediationResult{Refused: true, Error: "automated .htaccess remediation only applies to .htaccess files"}
}
resolved, _, err := resolveExistingFixPath(path, effectiveFixRoots(fixHtaccessAllowedRoots))
if err != nil {
return RemediationResult{Refused: errors.Is(fileResponseSourceError(err), errFileResponseRefused), Error: err.Error()}
}
// The account owner controls this directory and we run as root, so the
// file is pinned by inode and replaced through a random O_EXCL name under
// the pinned parent fd. A guessable staging name would let the owner
// plant a symlink there and have the cleaned bytes written anywhere.
target, err := openCleanTarget(resolved)
if err != nil {
return RemediationResult{Refused: errors.Is(fileResponseSourceError(err), errFileResponseRefused), Error: fmt.Sprintf("cannot open: %v", err)}
}
defer target.Close()
if expected != nil && (!sameFileIdentity(expected, target.Info) || !sameContentShape(expected, target.Info)) {
result.Refused = true
result.Error = "file changed before automatic cleaning"
return result
}
audit.rec.Result = actionlog.Failed
original, err := io.ReadAll(target.File)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot read: %v", err)}
}
audit.capture(target, original)
audit.rec.Result = actionlog.Refused
_, ranges := AuditHtaccessContent(resolved, original)
if len(ranges) == 0 {
return RemediationResult{Refused: true, Error: "no malicious directives found to remove"}
}
cleaned := applyRangeRemoval(original, ranges)
if len(cleaned) == len(original) {
return RemediationResult{Refused: true, Error: "no bytes removed (range computation produced empty diff)"}
}
backupDir := htaccessBackupDirRoot
backupPath := newQuarantinePath(backupDir, resolved)
meta := quarantineMetadata(resolved, target.Info, fmt.Sprintf("htaccess clean: %d ranges removed (%d -> %d bytes)", len(ranges), len(original), len(cleaned)))
audit.rec.Result = actionlog.Failed
audit.rec.Reason = meta.Reason
if err := storeQuarantineBackup(backupPath, original, meta, 0640); err != nil {
return RemediationResult{Error: fmt.Sprintf("writing durable backup: %v", err)}
}
if err := audit.replace(target, cleaned, backupPath); err != nil {
return RemediationResult{Refused: errors.Is(err, errFileResponseRefused), Error: fmt.Sprintf("atomic replace: %v", err)}
}
bytesRemoved := len(original) - len(cleaned)
return RemediationResult{
Success: true,
Action: fmt.Sprintf("removed %d malicious byte(s) from %s", bytesRemoved, resolved),
Description: fmt.Sprintf("Cleaned .htaccess: %d ranges, %d bytes removed (backup: %s)", len(ranges), bytesRemoved, backupPath),
}
}
// mergeRanges normalises the input slice: sort by start, then merge
// overlapping or adjacent ranges so cleaning produces deterministic
// output regardless of detector order.
func mergeRanges(in []htaccessByteRange) []htaccessByteRange {
if len(in) == 0 {
return nil
}
cp := make([]htaccessByteRange, len(in))
copy(cp, in)
sort.Slice(cp, func(i, j int) bool { return cp[i].Start < cp[j].Start })
out := []htaccessByteRange{cp[0]}
for _, r := range cp[1:] {
last := &out[len(out)-1]
if r.Start <= last.End {
if r.End > last.End {
last.End = r.End
}
continue
}
out = append(out, r)
}
return out
}
// applyRangeRemoval slices `content` minus every range in `ranges`
// (which mergeRanges has already sorted/merged). Each removal also
// includes the trailing newline if `end` lands on one, so we do not
// leave a blank line behind.
func applyRangeRemoval(content []byte, ranges []htaccessByteRange) []byte {
out := make([]byte, 0, len(content))
cursor := 0
for _, r := range ranges {
if r.Start > cursor {
out = append(out, content[cursor:r.Start]...)
}
end := r.End
if end < len(content) && content[end] == '\n' {
end++
}
cursor = end
}
if cursor < len(content) {
out = append(out, content[cursor:]...)
}
return out
}
type htaccessPhysicalByteLine struct {
text string
start int
end int
}
type htaccessLogicalByteLine struct {
text string
span htaccessByteRange
}
func splitHtaccessPhysicalByteLines(content []byte) []htaccessPhysicalByteLine {
var lines []htaccessPhysicalByteLine
for start := 0; start <= len(content); {
if start == len(content) {
lines = append(lines, htaccessPhysicalByteLine{
text: "",
start: start,
end: start,
})
break
}
end := start
for end < len(content) && content[end] != '\n' {
end++
}
lines = append(lines, htaccessPhysicalByteLine{
text: string(content[start:end]),
start: start,
end: end,
})
if end == len(content) {
break
}
start = end + 1
}
return lines
}
func htaccessLogicalByteLines(content []byte) []htaccessLogicalByteLine {
physical := splitHtaccessPhysicalByteLines(content)
var out []htaccessLogicalByteLine
for i := 0; i < len(physical); {
start := physical[i].start
var end int
var sb strings.Builder
for {
body, continues := htaccessContinuationBody(physical[i].text, i < len(physical)-1)
end = physical[i].end
if continues {
sb.WriteString(body)
i++
continue
}
sb.WriteString(body)
break
}
out = append(out, htaccessLogicalByteLine{
text: sb.String(),
span: htaccessByteRange{Start: start, End: end},
})
i++
}
return out
}
func matchesFromLogicalLineRegex(content []byte, re *regexp.Regexp) []htaccessMatch {
var out []htaccessMatch
for _, logical := range htaccessLogicalByteLines(content) {
if !re.MatchString(logical.text) {
continue
}
out = append(out, htaccessMatch{
Range: logical.span,
Excerpt: trimExcerpt(content, logical.span.Start, logical.span.End),
})
}
return out
}
// detectPHPInUploads flags AddHandler/SetHandler/ForceType lines
// that map to PHP when the .htaccess lives inside a directory where
// PHP execution is rarely legitimate.
func detectPHPInUploads(content []byte, path string) []htaccessMatch {
if !pathInNonScriptDir(path) {
return nil
}
return matchesFromLogicalLineRegex(content, rePHPHandlerMap)
}
func pathInNonScriptDir(path string) bool {
lower := strings.ToLower(path)
for _, dir := range htaccessNonScriptDirHints {
if strings.Contains(lower, dir) {
return true
}
}
return false
}
// detectAutoPrepend flags PHP auto_prepend_file and auto_append_file
// directives whose target the account owner can write to (see
// autoPrependTargetSuspicious).
func detectAutoPrepend(content []byte, path string) []htaccessMatch {
idxs := reAutoPrepend.FindAllSubmatchIndex(content, -1)
var out []htaccessMatch
for _, idx := range idxs {
if len(idx) < 4 {
continue
}
target := string(content[idx[2]:idx[3]])
if !autoPrependTargetSuspicious(target, path) {
continue
}
out = append(out, htaccessMatch{
Range: lineRange(content, idx[0], idx[1]),
Excerpt: trimExcerpt(content, idx[0], idx[1]),
})
}
return out
}
// reUARewriteRuleLine captures the substitution and flag list of a single
// RewriteRule directive. Pairing is handled by uaCloakPairedRule so a later
// unrelated rule is not mistaken for the rule fed by a RewriteCond chain.
var reUARewriteRuleLine = regexp.MustCompile(`(?i)^\s*RewriteRule\s+\S+\s+(\S+)(?:\s+\[([^\]]+)\])?\s*$`)
type uaRewriteRulePair struct {
start int
end int
substitution string
flags string
parsed bool
}
// uaCloakDefensiveFlags lists RewriteRule flags that, when set on the
// rule paired with a UA cond, indicate defensive blocking rather than
// content cloaking. Apache combines these flags with a comma so each
// flag is matched as a substring of the bracketed flag list.
var uaCloakDefensiveFlags = []string{"f", "g"}
// uaCloakBlocklistThreshold is the number of OR-list entries in the
// UA cond's regex that converts the cond from "potential cloak" to
// "operator-installed bot blocklist". A cond with 4+ alternation
// entries is overwhelmingly a defensive block (the canonical
// SoftAculous / Apache Bad Bots list ships ~20+ entries).
const uaCloakBlocklistThreshold = 4
// uaCloakAlternationCount counts the top-level "|" alternation
// branches in a UA cond regex pattern, ignoring "|" characters inside
// nested parentheses. Used to identify long bot blocklists.
func uaCloakAlternationCount(pattern string) int {
return len(uaCloakAlternationBranches(pattern))
}
// uaCloakAlternationBranches splits a UA cond regex pattern on its
// top-level "|" alternations, ignoring "|" inside nested parentheses.
func uaCloakAlternationBranches(pattern string) []string {
depth := 0
inClass := false
prevEscape := false
start := 0
var branches []string
for i := 0; i < len(pattern); i++ {
c := pattern[i]
if prevEscape {
prevEscape = false
continue
}
switch c {
case '\\':
prevEscape = true
case '[':
if !inClass {
inClass = true
}
case ']':
if inClass {
inClass = false
}
case '(':
if !inClass {
depth++
}
case ')':
if !inClass && depth > 0 {
depth--
}
case '|':
if !inClass && depth <= 1 {
branches = append(branches, pattern[start:i])
start = i + 1
}
}
}
return append(branches, pattern[start:])
}
// uaCloakSearchCrawlerMajority reports whether more than half of the UA
// alternatives across the given cond patterns name search-engine crawlers.
// Such a chain is a cloak target list, not a scraper blocklist, so the
// long-list gates must not silence it.
func uaCloakSearchCrawlerMajority(patterns []string) bool {
total, search := 0, 0
for _, p := range patterns {
for _, branch := range uaCloakAlternationBranches(p) {
total++
if searchCrawlerUARegex.MatchString(branch) {
search++
}
}
}
return total > 0 && search*2 > total
}
// uaCloakPairedRuleIsDefensive scans forward from condEnd for the
// RewriteRule directive that Apache will pair with the cond. Returns
// true when that rule is a no-op ("-" substitution), a forbid
// ([F]/[G] flag), or absent. Both shapes are defensive, not cloaking.
//
// Apache's chaining model: a chain of RewriteCond lines applies to
// the FIRST RewriteRule that follows them in the file. We walk
// line-by-line and stop at the first RewriteRule. A blank line, a
// non-RewriteCond directive, or end-of-file means there is no paired
// rule - the cond is dead text and not actively cloaking anything.
func uaCloakPairedRuleIsDefensive(content []byte, condEnd int) bool {
rule, ok := uaCloakPairedRule(content, condEnd)
if !ok {
return true
}
if !rule.parsed {
return false
}
if rule.substitution == "-" {
return true
}
for _, f := range uaCloakDefensiveFlags {
// Match flag as a comma-bounded token: "F" matches "[F]",
// "[F,L]", "[L,F]"; does NOT match "[NC]" or "[QSA]".
for _, token := range strings.Split(rule.flags, ",") {
if strings.TrimSpace(token) == f {
return true
}
}
}
return false
}
// uaCloakPairedRule returns the first RewriteRule in the contiguous
// RewriteCond chain that starts after condEnd. Blank lines and comments are
// ignored by Apache, so they do not break the chain; any other directive does.
// Otherwise a stale UA condition can be paired with an unrelated later
// RewriteRule and cleaning would remove too much.
func uaCloakPairedRule(content []byte, condEnd int) (uaRewriteRulePair, bool) {
pos := condEnd
for {
lineStart, lineEnd, rawEnd, ok := htaccessLineAfter(content, pos)
if !ok {
return uaRewriteRulePair{}, false
}
trimmed := strings.TrimSpace(string(content[lineStart:lineEnd]))
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
pos = rawEnd
continue
}
switch htaccessDirectiveName(trimmed) {
case "rewritecond":
pos = rawEnd
continue
case "rewriterule":
rule := uaRewriteRulePair{start: lineStart, end: lineEnd}
if idx := reUARewriteRuleLine.FindStringSubmatchIndex(string(content[lineStart:lineEnd])); idx != nil {
line := string(content[lineStart:lineEnd])
rule.substitution = line[idx[2]:idx[3]]
if idx[4] != -1 {
rule.flags = strings.ToLower(line[idx[4]:idx[5]])
}
rule.parsed = true
}
return rule, true
default:
return uaRewriteRulePair{}, false
}
}
}
func htaccessLineAfter(content []byte, pos int) (start, lineEnd, rawEnd int, ok bool) {
for pos < len(content) && content[pos] != '\n' {
pos++
}
if pos >= len(content) {
return 0, 0, 0, false
}
pos++
if pos >= len(content) {
return 0, 0, 0, false
}
start = pos
rawEnd = pos
for rawEnd < len(content) && content[rawEnd] != '\n' {
rawEnd++
}
lineEnd = rawEnd
if lineEnd > start && content[lineEnd-1] == '\r' {
lineEnd--
}
return start, lineEnd, rawEnd, true
}
func htaccessLineBefore(content []byte, pos int) (start, lineEnd int, ok bool) {
if pos <= 0 {
return 0, 0, false
}
lineEnd = pos
for lineEnd > 0 && (content[lineEnd-1] == '\n' || content[lineEnd-1] == '\r') {
lineEnd--
}
if lineEnd <= 0 {
return 0, 0, false
}
start = lineEnd
for start > 0 && content[start-1] != '\n' {
start--
}
return start, lineEnd, true
}
func htaccessDirectiveName(line string) string {
fields := strings.Fields(line)
if len(fields) == 0 {
return ""
}
return strings.ToLower(fields[0])
}
// detectUserAgentCloak flags RewriteCond %{HTTP_USER_AGENT}
// directives that match a known crawler UA AND are part of an active
// content-cloaking rule. Four suppression gates filter legitimate
// shapes before emitting the High alert:
//
// 1. Negated cond ("RewriteCond %{HTTP_USER_AGENT} !..."): the rule
// applies only when the UA is NOT this crawler. Cache plugins
// (WP Fastest Cache, WP Super Cache) ship long negated lists to
// exclude social-share scrapers from the cached-content rewrite.
//
// 2. Long alternation (>= uaCloakBlocklistThreshold OR-branches):
// operator-installed defensive blocklists ship many bot UAs in
// a single OR chain paired with a [F] forbid or sinkhole rewrite.
// Cloakers use one or two crawler names.
//
// 3. Long multi-line chain: the canonical Apache "Bad Bots" snippet
// puts each scraper UA on its own RewriteCond line, terminated
// by one RewriteRule. The chain length is the blocklist signal;
// per-cond alternation is always 1.
//
// 4. Paired RewriteRule is defensive ("-" substitution, [F]/[G]
// flag, or absent): the cond is part of a forbid / env-var-set
// block, not a content swap.
//
// All gates fail-closed: any uncertainty (no paired rule found,
// parse failure) keeps the original alert firing.
func detectUserAgentCloak(content []byte, _ string) []htaccessMatch {
idxs := reUACloakCond.FindAllSubmatchIndex(content, -1)
chainSize := uaCloakChainSizes(content, idxs)
// A chain whose UA alternatives are mostly search-engine crawlers is a
// cloak target list, not a scraper blocklist: gates 2 and 3 do not
// apply to it and only the paired rule (gate 4) can clear it.
searchList := make([]bool, len(idxs))
for _, group := range uaCloakChainGroups(content, idxs) {
patterns := make([]string, 0, len(group))
for _, i := range group {
if len(idxs[i]) >= 4 {
patterns = append(patterns, string(content[idxs[i][2]:idxs[i][3]]))
}
}
majority := uaCloakSearchCrawlerMajority(patterns)
for _, i := range group {
searchList[i] = majority
}
}
var out []htaccessMatch
for i, idx := range idxs {
if len(idx) < 4 {
continue
}
uaPattern := string(content[idx[2]:idx[3]])
if !crawlerUARegex.MatchString(uaPattern) {
continue
}
// Gate 1: negated cond. Strip leading whitespace and look
// for "!" before the rest of the pattern.
if strings.HasPrefix(strings.TrimSpace(uaPattern), "!") {
continue
}
// Gate 2: long alternation list = bot blocklist.
if !searchList[i] && uaCloakAlternationCount(uaPattern) >= uaCloakBlocklistThreshold {
continue
}
// Gate 3: long multi-line chain = bot blocklist.
if !searchList[i] && chainSize[i] >= uaCloakBlocklistThreshold {
continue
}
// Gate 4: paired RewriteRule is defensive.
if uaCloakPairedRuleIsDefensive(content, idx[1]) {
continue
}
// Removal must take the whole cloak unit: the full contiguous
// RewriteCond chain plus the RewriteRule it feeds. Removing only the
// cond would leave the rule unconditional, turning a crawler-only
// cloak into a redirect of every visitor to the attacker URL.
block := uaCloakBlockRange(content, idx[0], idx[1])
out = append(out, htaccessMatch{
Range: block,
Excerpt: trimExcerpt(content, block.Start, block.End),
})
}
return out
}
// uaCloakBlockRange returns the byte range covering the entire cloak unit for
// the UA cond at [condStart, condEnd): the full contiguous RewriteCond chain it
// belongs to (Apache applies a run of conds to the first following
// RewriteRule) plus that paired RewriteRule. Falls back to the cond's own line
// when no paired rule follows - a dead cond redirects nobody, so there is
// nothing extra to strip.
func uaCloakBlockRange(content []byte, condStart, condEnd int) htaccessByteRange {
chainStart := uaCondChainStart(content, condStart)
rule, ok := uaCloakPairedRule(content, condEnd)
if !ok {
return lineRange(content, chainStart, condEnd)
}
return lineRange(content, chainStart, rule.end)
}
// uaCondChainStart walks backward from the start of a RewriteCond line over
// immediately preceding RewriteCond lines and returns the byte offset where the
// chain begins. Comments and blank lines between conds do not break Apache's
// chain, but they also are not swallowed when they merely precede the first
// cond. A non-RewriteCond directive terminates the walk.
func uaCondChainStart(content []byte, condStart int) int {
start := condStart
scan := condStart
for {
prevStart, prevEnd, ok := htaccessLineBefore(content, scan)
if !ok {
return start
}
trimmed := strings.TrimSpace(string(content[prevStart:prevEnd]))
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
scan = prevStart
continue
}
if htaccessDirectiveName(trimmed) != "rewritecond" {
return start
}
start = prevStart
scan = prevStart
}
}
// uaCloakChainSizes returns, for each UA cond match, the number of
// UA conds in its chain. Two UA conds belong to the same chain when
// only whitespace, comments, or blank lines sit between them; any
// other directive breaks the chain. That mirrors how Apache groups
// conds onto the next RewriteRule.
func uaCloakChainSizes(content []byte, idxs [][]int) []int {
sizes := make([]int, len(idxs))
for _, group := range uaCloakChainGroups(content, idxs) {
for _, i := range group {
sizes[i] = len(group)
}
}
return sizes
}
// uaCloakChainGroups partitions the UA cond matches into chains (see
// uaCloakChainSizes) and returns each chain as the match indexes it holds.
func uaCloakChainGroups(content []byte, idxs [][]int) [][]int {
var groups [][]int
var current []int
for i := range idxs {
if i > 0 && !uaCondsAreAdjacent(content, idxs[i-1][1], idxs[i][0]) {
groups = append(groups, current)
current = nil
}
current = append(current, i)
}
if len(current) > 0 {
groups = append(groups, current)
}
return groups
}
// uaCondsAreAdjacent reports whether the gap between two UA cond matches
// contains only line breaks, whitespace, blank lines, or comments.
func uaCondsAreAdjacent(content []byte, prevEnd, nextStart int) bool {
if prevEnd > nextStart || nextStart > len(content) {
return false
}
seenNewline := false
for pos := prevEnd; pos < nextStart; {
lineEnd := pos
for lineEnd < nextStart && content[lineEnd] != '\n' {
lineEnd++
}
trimmed := strings.TrimSpace(string(content[pos:lineEnd]))
if trimmed != "" && !strings.HasPrefix(trimmed, "#") {
return false
}
if lineEnd < nextStart {
seenNewline = true
lineEnd++
}
pos = lineEnd
}
return seenNewline
}
// detectSpamRedirect flags RewriteRule directives whose target host
// is on a known spam TLD. Operator-supplied legitimate hosts in
// these TLDs need a per-path suppression.
func detectSpamRedirect(content []byte, _ string) []htaccessMatch {
idxs := reSpamRedirect.FindAllSubmatchIndex(content, -1)
var out []htaccessMatch
for _, idx := range idxs {
if len(idx) < 4 {
continue
}
target := string(content[idx[2]:idx[3]])
host := extractHost(target)
if !hostOnSpamTLD(host) {
continue
}
out = append(out, htaccessMatch{
Range: lineRange(content, idx[0], idx[1]),
Excerpt: trimExcerpt(content, idx[0], idx[1]),
})
}
return out
}
func extractHost(rawURL string) string {
u, err := url.Parse(rawURL)
if err != nil {
return ""
}
return strings.ToLower(u.Hostname())
}
func hostOnSpamTLD(host string) bool {
for _, tld := range htaccessSpamTLDs {
if strings.HasSuffix(host, tld) {
return true
}
}
return false
}
// reFilesMatchPattern captures the regex pattern inside the FilesMatch
// quotes. The capture group is everything between the optional quote
// characters, which is the Apache-side regex applied to filenames.
var reFilesMatchPattern = regexp.MustCompile(`(?im)^\s*<FilesMatch\s+["']?([^"'>]+?)["']?\s*>`)
// reFilesMatchExtensionTail strips the canonical PHP extension suffix
// from a FilesMatch pattern so the remaining literal can be examined.
// Anchors and end-of-string markers are left to the caller.
var reFilesMatchExtensionTail = regexp.MustCompile(`(?i)\\\.(?:php|phtml|ph[2-7])\$?\)?$`)
// filesMatchPatternIsTargeted reports whether the FilesMatch regex
// names at least one specific PHP filename rather than granting access
// to every .php file in the directory. Stock plugins ship targeted
// patterns ("wpc\.php$", "ps_facetedsearch-.+\.php$",
// "(webp-on-demand\.php|webp-realizer\.php)$"); the malicious shape
// is a bare wildcard ("\.php$", ".*\.php$", "[^/]+\.php$").
//
// The check is character-class based: any literal alphanumeric, dash,
// or underscore in the pattern (after stripping the trailing
// "\.php$" / "\.phtml$" extension) means the pattern names something
// specific. A pattern composed only of regex meta-characters (".",
// "*", "^", "$", "[", "]", "(", ")", "|", "+", "?", "\\") is treated
// as a wildcard and continues to the wildcard-context check.
//
// The test is applied per alternative: `^(a|.*)\.php$` carries a literal
// yet still matches every .php through its second branch, and treating it
// as targeted let an attacker write a shield the detector skipped. Every
// top-level alternative (after unwrapping one outer group) must name
// something; escapes such as `\w` and character classes do not count.
func filesMatchPatternIsTargeted(pattern string) bool {
stripped := reFilesMatchExtensionTail.ReplaceAllString(pattern, "")
stripped = strings.TrimSuffix(strings.TrimPrefix(stripped, "^"), "$")
if inner, ok := unwrapOuterGroup(stripped); ok {
stripped = inner
}
alternatives := topLevelAlternatives(stripped)
if len(alternatives) == 0 {
return false
}
for _, alt := range alternatives {
if !regexAlternativeNamesSomething(alt) {
return false
}
}
return true
}
// htaccessParentPHPFileCount counts ".php" files (and other handler
// extensions FilesMatch covers) sitting alongside the .htaccess at
// path. Used to differentiate a plugin directory full of legitimate
// PHP dispatchers from a freshly-attacker-written upload directory
// containing one or zero PHP files.
//
// Errors (parent missing, permission denied, race) return 0 - the
// caller treats that as "not enough sibling PHP" and keeps firing.
func htaccessParentPHPFileCount(htaccessPath string) int {
parent := filepath.Dir(htaccessPath)
entries, err := os.ReadDir(parent)
if err != nil {
return 0
}
n := 0
for _, e := range entries {
if e.IsDir() {
continue
}
ext := strings.ToLower(filepath.Ext(e.Name()))
switch ext {
case ".php", ".phtml", ".ph2", ".ph3", ".ph4", ".ph5", ".ph6", ".ph7":
n++
}
}
return n
}
// filesMatchShieldSiblingThreshold is the number of sibling PHP files
// that converts a bare-wildcard FilesMatch shield from "fire" to
// "treat as legitimate plugin allowlist". A directory with 3+ stock
// PHP dispatchers existing alongside the shield is overwhelmingly
// likely to be a legitimate webapp module (KCFinder ships ~5,
// PrestaShop modules ship ~10+, vendor dispatchers ship a handful).
// The attacker drop pattern is .htaccess + one or zero PHP files.
const filesMatchShieldSiblingThreshold = 3
// detectFilesMatchShield finds <FilesMatch ...\.php(tml)?> blocks
// that grant Allow from all or Require all granted -- the canonical
// "let everyone execute everything we just dropped" pattern. The
// returned range covers the full block, opening tag through closing
// tag inclusive.
//
// Two suppression gates run before emitting the finding:
//
// 1. Targeted pattern: if the FilesMatch regex names a specific
// filename ("wpc\.php$", named allowlist, or prefix pattern), it
// is a legitimate plugin allowlist and is skipped.
//
// 2. Bare wildcard with sibling PHP context: if the FilesMatch is a
// bare wildcard but the .htaccess parent directory contains
// multiple sibling PHP dispatchers, the shield is protecting an
// existing legitimate plugin layout, not a freshly-dropped
// dropper. Threshold is filesMatchShieldSiblingThreshold.
//
// Both gates fail-open: any uncertainty (parse failure, IO error)
// keeps the original Critical alert firing. The sibling-PHP gate
// requires the htaccess path so the detector now reads it from the
// caller (passed as the second argument to all htaccessDetector.Detect
// implementations).
func detectFilesMatchShield(content []byte, path string) []htaccessMatch {
openIdxs := reFilesMatchOpen.FindAllIndex(content, -1)
closeIdxs := reFilesMatchClose.FindAllIndex(content, -1)
if len(openIdxs) == 0 || len(closeIdxs) == 0 {
return nil
}
patternIdxs := reFilesMatchPattern.FindAllSubmatchIndex(content, -1)
var out []htaccessMatch
for _, open := range openIdxs {
// pair this opening tag with the next closing tag after it
var paired []int
for _, c := range closeIdxs {
if c[0] >= open[1] {
paired = c
break
}
}
if paired == nil {
continue
}
body := content[open[1]:paired[0]]
bodyLower := strings.ToLower(string(body))
if !strings.Contains(bodyLower, "allow from all") && !strings.Contains(bodyLower, "require all granted") {
continue
}
// Look up the FilesMatch pattern that opened at this position
// so we can apply the targeted-vs-wildcard discriminator.
var openPattern string
for _, pIdx := range patternIdxs {
if len(pIdx) < 4 {
continue
}
if pIdx[0] == open[0] {
openPattern = string(content[pIdx[2]:pIdx[3]])
break
}
}
if openPattern != "" && filesMatchPatternIsTargeted(openPattern) {
continue
}
// Bare wildcard: check sibling PHP count. A directory with
// multiple stock PHP dispatchers is a legitimate plugin layout.
// Not inside an upload-style tree, though: there the siblings
// are whatever the uploader chose to drop, and three dummy .php
// files are the cheapest way to silence this finding.
if path != "" && !htaccessInUploadTree(path) && htaccessParentPHPFileCount(path) >= filesMatchShieldSiblingThreshold {
continue
}
out = append(out, htaccessMatch{
Range: blockRange(content, open[0], paired[1]),
Excerpt: trimExcerpt(content, open[0], paired[1]),
})
}
return out
}
// detectHeaderInjection flags Header set / Header add directives
// whose name is on the small tracking-header allowlist. Generic
// CSP / HSTS / X-Frame-Options headers do not match.
func detectHeaderInjection(content []byte, _ string) []htaccessMatch {
idxs := reHeaderSetAdd.FindAllSubmatchIndex(content, -1)
var out []htaccessMatch
for _, idx := range idxs {
if len(idx) < 6 {
continue
}
name := string(content[idx[4]:idx[5]])
if !headerNameSuspicious(name) {
continue
}
out = append(out, htaccessMatch{
Range: lineRange(content, idx[0], idx[1]),
Excerpt: trimExcerpt(content, idx[0], idx[1]),
})
}
return out
}
func headerNameSuspicious(name string) bool {
lower := strings.ToLower(name)
for _, h := range htaccessTrackingHeaders {
if strings.HasPrefix(lower, strings.ToLower(h)) {
return true
}
}
return false
}
// errorDocumentHostShareThreshold bounds how short a host-vs-path
// substring match can be while still treating the redirect as
// "same-brand". Three characters is the floor: anything shorter would
// match incidental segments ("us" inside "user", "co" inside "co.uk")
// and let an attacker tunnel through with a name like "co.evil.com".
const errorDocumentHostShareThreshold = 4
// errorDocumentHostIsSameBrand reports whether the URL host's
// "registrable label" (the leftmost segment of the public suffix +
// 1) shares an alphanumeric stem of >= errorDocumentHostShareThreshold
// chars with any path component of the .htaccess file. This catches
// the dominant legitimate shape: a custom 404 redirect to the site's
// own homepage on the same brand domain.
//
// Examples:
//
// /home/flores/public_html/.htaccess + https://floresgrup.ro
// -> account "flores" is a substring of label "floresgrup" -> same-brand
//
// /home/shop/example-shop.com/.htaccess + https://www.example-shop.com/404
// -> domain dir "example-shop.com" contains label "example-shop" -> same-brand
//
// /home/victim/public_html/.htaccess + https://attacker.com/landing
// -> "attacker" shares no >=4-char stem with any path component -> different-brand
func errorDocumentHostIsSameBrand(htaccessPath, urlHost string) bool {
label := registrableLabel(urlHost)
if len(label) < errorDocumentHostShareThreshold {
return false
}
labelLower := strings.ToLower(label)
for _, component := range strings.Split(htaccessPath, string(filepath.Separator)) {
if component == "" || component == "public_html" || component == "home" {
continue
}
comp := strings.ToLower(component)
if longestCommonAlnumRun(comp, labelLower) >= errorDocumentHostShareThreshold {
return true
}
}
return false
}
// registrableLabel extracts the leftmost segment of the public
// suffix + 1: for "www.example-shop.com" returns "example-shop", for
// "floresgrup.ro" returns "floresgrup". This is heuristic - we treat
// the last dot-separated segment as the TLD - but it is robust to
// the common cases (single-segment TLD, two-segment country code TLD
// like "co.uk" handled by stripping known double-segment suffixes).
//
// Returns "" for inputs that look like an IPv4 dotted quad: numeric
// targets are flagged separately and never qualify as same-brand.
func registrableLabel(host string) string {
host = strings.ToLower(strings.TrimSpace(host))
if host == "" {
return ""
}
// IPv4 dotted quad: each segment must be 1-3 digits.
if isIPv4(host) {
return ""
}
// Strip a leading "www." for canonical comparison.
host = strings.TrimPrefix(host, "www.")
parts := strings.Split(host, ".")
if len(parts) < 2 {
return host
}
// Two-segment public suffix heuristic: "co.uk", "co.za",
// "com.au" etc. If the second-to-last segment is one of these
// short country-code prefixes, the registrable label is the
// THIRD-from-last segment.
if len(parts) >= 3 {
penultimate := parts[len(parts)-2]
twoSegmentCC := map[string]bool{
"co": true, "com": true, "net": true, "org": true,
"ac": true, "gov": true, "edu": true,
}
if twoSegmentCC[penultimate] {
return parts[len(parts)-3]
}
}
return parts[len(parts)-2]
}
// isIPv4 reports whether s parses as a dotted-quad IPv4 address.
// IPv6 / hex / mixed forms are caught separately by IP-target
// signalling above the same-brand check.
func isIPv4(s string) bool {
parts := strings.Split(s, ".")
if len(parts) != 4 {
return false
}
for _, p := range parts {
if p == "" || len(p) > 3 {
return false
}
for _, c := range p {
if c < '0' || c > '9' {
return false
}
}
}
return true
}
// longestCommonAlnumRun returns the length of the longest common
// substring of a and b that consists entirely of alphanumeric
// characters. Symbols (".", "-", "_") segment the run so an
// "example-shop" path component can match a "example-shop.com" URL
// host on the "example" stem without crediting the dash.
func longestCommonAlnumRun(a, b string) int {
best := 0
// Walk a, accumulate alnum tokens, check each token against b.
tokens := splitAlnumTokens(a)
bLower := b
for _, tok := range tokens {
if len(tok) < best {
continue
}
// Find tok inside b: contiguous alnum substring match.
if strings.Contains(bLower, tok) {
if len(tok) > best {
best = len(tok)
}
}
}
return best
}
// splitAlnumTokens splits s on any non-alphanumeric character, returning
// the maximal alphanumeric runs.
func splitAlnumTokens(s string) []string {
var out []string
start := -1
for i := 0; i < len(s); i++ {
c := s[i]
isAlnum := (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9')
if isAlnum && start == -1 {
start = i
} else if !isAlnum && start != -1 {
out = append(out, s[start:i])
start = -1
}
}
if start != -1 {
out = append(out, s[start:])
}
return out
}
// detectErrorDocumentHijack flags ErrorDocument directives whose
// target is an external http(s) URL pointing at a host that does NOT
// share a brand stem with the .htaccess file's path. Same-brand
// redirects (custom 404 -> site homepage) are extremely common and
// must not be flagged. Spam-TLD targets and IP-address targets are
// always flagged regardless of any brand match.
func detectErrorDocumentHijack(content []byte, path string) []htaccessMatch {
idxs := reErrorDocument.FindAllSubmatchIndex(content, -1)
var out []htaccessMatch
for _, idx := range idxs {
if len(idx) < 4 {
continue
}
target := string(content[idx[2]:idx[3]])
host := extractHost(target)
// Spam TLDs and IP-address targets always fire. Both are
// signals of compromise even when the path-share heuristic
// would otherwise consider the host same-brand.
if hostOnSpamTLD(host) || isIPv4(host) {
out = append(out, htaccessMatch{
Range: lineRange(content, idx[0], idx[1]),
Excerpt: trimExcerpt(content, idx[0], idx[1]),
})
continue
}
// Same-brand check only when we have a path to compare
// against. Detector is sometimes called with an empty path
// (unit tests of the regex alone); fail-closed.
if path != "" && errorDocumentHostIsSameBrand(path, host) {
continue
}
out = append(out, htaccessMatch{
Range: lineRange(content, idx[0], idx[1]),
Excerpt: trimExcerpt(content, idx[0], idx[1]),
})
}
return out
}
// lineRange computes a range covering the full line(s) that contain
// [start, end]. Line boundaries are at '\n'; the returned end is
// the position of the trailing '\n' (exclusive) or len(content) if
// the match is at EOF without a trailing newline.
func lineRange(content []byte, start, end int) htaccessByteRange {
if start < 0 {
start = 0
}
if end > len(content) {
end = len(content)
}
for start > 0 && content[start-1] != '\n' {
start--
}
for end < len(content) && content[end] != '\n' {
end++
}
return htaccessByteRange{Start: start, End: end}
}
// blockRange covers the whole block including the opening and
// closing directive lines. Same line-snapping rules as lineRange.
func blockRange(content []byte, start, end int) htaccessByteRange {
r := lineRange(content, start, end)
return r
}
// trimExcerpt returns up to ~200 bytes of the matched content,
// trimmed and with newlines collapsed, for finding details.
func trimExcerpt(content []byte, start, end int) string {
if start < 0 {
start = 0
}
if end > len(content) {
end = len(content)
}
s := strings.TrimSpace(string(content[start:end]))
s = strings.ReplaceAll(s, "\n", " | ")
if len(s) > 200 {
s = s[:200] + "..."
}
return s
}
package checks
import (
"path/filepath"
"strings"
)
// unwrapOuterGroup returns the inside of pattern when the whole pattern is
// one parenthesised group, so its alternatives can be judged one by one.
func unwrapOuterGroup(pattern string) (string, bool) {
if len(pattern) < 2 || pattern[0] != '(' || pattern[len(pattern)-1] != ')' {
return pattern, false
}
depth := 0
for i := 0; i < len(pattern); i++ {
switch pattern[i] {
case '\\':
i++
case '(':
depth++
case ')':
depth--
if depth == 0 && i != len(pattern)-1 {
// The first group closes before the end: not one outer group.
return pattern, false
}
}
}
if depth != 0 {
return pattern, false
}
return pattern[1 : len(pattern)-1], true
}
// regexAlternativeNamesSomething reports whether one regex alternative
// carries a literal name character. Escaped sequences (`\w`, `\.`) and
// character classes are skipped: they describe shapes, not names.
func regexAlternativeNamesSomething(alt string) bool {
inClass := false
for i := 0; i < len(alt); i++ {
c := alt[i]
switch {
case c == '\\':
i++
case inClass:
if c == ']' {
inClass = false
}
case c == '[':
inClass = true
case (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') || (c >= '0' && c <= '9') || c == '_' || c == '-':
return true
}
}
return false
}
// uploadTreeDirNames are directory names under which web-writable content
// lands. Findings there are judged without the sibling-PHP gate.
var uploadTreeDirNames = map[string]bool{
"uploads": true, "upload": true, "cache": true, "tmp": true, "temp": true,
"files": true, "media": true, "attachments": true, "images": true,
}
// htaccessInUploadTree reports whether the .htaccess at path sits under an
// upload-style directory. The top-level directory (/home, /var, /tmp) is
// never a docroot-relative upload dir and is skipped, so a system temp root
// does not make every path below it an upload tree.
func htaccessInUploadTree(path string) bool {
for _, part := range webTreeComponents(filepath.Dir(path)) {
if uploadTreeDirNames[strings.ToLower(part)] {
return true
}
}
return false
}
// docrootMarkerNames are directory names that begin a web tree on hosts
// whose layout is not one of the configured account roots.
var docrootMarkerNames = map[string]bool{
"public_html": true, "www": true, "htdocs": true, "httpdocs": true,
"web": true, "html": true, "public": true,
}
// webTreeComponents returns the directory components that lie inside the
// account's web tree: everything below <root>/<account> when dir is under
// an account root, otherwise everything below the first document-root
// marker. The prefix above the web tree (a temp root, /var/www, an account
// literally named tmp) never counts: it is not attacker-writable content.
func webTreeComponents(dir string) []string {
clean := filepath.ToSlash(filepath.Clean(dir))
if root, account, ok := accountRootOf(clean); ok {
rel := strings.TrimPrefix(clean, filepath.ToSlash(filepath.Join(root, account))+"/")
return strings.Split(rel, "/")
}
parts := strings.Split(strings.Trim(clean, "/"), "/")
for i, part := range parts {
if docrootMarkerNames[strings.ToLower(part)] {
return parts[i+1:]
}
}
return nil
}
// HTTP abuse detection.
//
// This file holds the access-log line parser, the per-scan aggregator
// struct (domlogStats), the UA classifier, the bot-classifier interface,
// and the static allowlist classifier that consults embedded bot IP ranges.
// The rDNS verifying classifier arrives in Task 5.
package checks
import (
"fmt"
"net"
"net/url"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/threatintel"
)
// accessLogRecord is the parsed shape of one access-log line. Combined
// Log Format plus the cPanel final-vhost extension:
//
// IP - - [time] "METHOD URI PROTO" status bytes "referer" "ua" "vhost"
//
// The parser tolerates either the 9-field plain CLF or the 10-field
// cPanel variant. Bad lines return ok=false; callers must skip them.
type accessLogRecord struct {
RemoteIP string
Time time.Time
Method string
URI string
Status int
UserAgent string
XFF string // optional; only trusted when RemoteIP is a trusted proxy
Domain string // vhost the line came from (per-domain domlog); empty for the central log
Account string // cPanel account owning Domain; empty when unknown/non-cPanel
}
// uaKind is the User-Agent classification produced by classifyUA and
// consumed by domlogStats.scan and the http_ua_spoof emit logic.
type uaKind int
const (
uaKindBrowser uaKind = iota
uaKindClaimedBot
uaKindClaimedBotNegative
// uaKindClaimedBotPending is a claimed-bot UA whose rDNS verification has
// not resolved yet (cache miss, async job in flight). It is neither trusted
// (would have returned early via IsVerifiedBot) nor a confirmed spoof. Such
// an IP's flood/scanner abuse is routed to the reversible
// http_claimed_bot_unverified check instead of a hard block, so a real
// crawler is not blocked during the verification window.
uaKindClaimedBotPending
uaKindKnownScanner
uaKindWPSpoofPingback
uaKindScriptingLang
uaKindHeadless
uaKindEmpty
)
// botClassifier decides whether the source IP is a known verified bot
// the detector should NOT count or flag. Returns true to skip. The
// real implementation lives in internal/threatintel; tests use the
// nopBotClassifier below.
type botClassifier interface {
IsVerifiedBot(ip string, ua string) bool
}
type confirmedNegativeBotClassifier interface {
ConfirmedNegative(ip, ua string) bool
}
type pendingBotVerificationClassifier interface {
VerificationPending(ip, ua string) bool
}
type nopBotClassifier struct{}
func (nopBotClassifier) IsVerifiedBot(string, string) bool { return false }
// httpSample is one representative request kept per IP for forensic
// display in the finding Details field. First-seen wins; subsequent
// requests increment counters only.
type httpSample struct {
Method string
URI string
UA string
}
// domlogStats is the per-scan aggregator. One instance per CheckWPBruteForce
// invocation. Every counter is map[ip] -> int so emit can produce findings
// per source IP without a second pass.
type domlogStats struct {
wpLogin map[string]int
xmlrpc map[string]int
userEnum map[string]int
httpReqs map[string]int
uaCat map[string]map[uaKind]int
samples map[string]httpSample
// domains tracks the set of distinct vhosts each IP touched, so the
// per-IP aggregate findings can report cross-site spread (one IP
// scanning many vhosts on a shared host). Empty-domain records (the
// central access log) do not contribute.
domains map[string]map[string]struct{}
// abuseDomains tracks the in-window vhosts where each source made the
// request shape for a specific HTTP-abuse finding. The distributed
// rollup uses this instead of all in-window touches so an IP that was
// abusive elsewhere cannot make a normal hit count against a vhost.
abuseDomains map[string]map[string]map[string]struct{}
// scannerErr counts in-window non-asset requests per IP whose status
// matched the configured probe-error set. scannerReqs is the matching
// denominator, so the two counters must be gated by the same window
// check and the same display-asset exclusion.
scannerErr map[string]int
// scannerReqs is the per-IP scanner-profile denominator: in-window
// requests excluding static display assets (see isDisplayAssetProbe),
// so broken images, styles, scripts, and fonts neither dilute the
// error rate nor pad the volume gate. Separate from httpReqs so the
// flood detector's denominator is unchanged.
scannerReqs map[string]int
// scannerPaths is the distinct error-status path set per IP, query
// strings stripped, capped at httpScannerMaxTrackedPaths to bound
// memory under a flood of unique probe URLs.
scannerPaths map[string]map[string]struct{}
// scannerSamples keeps the first error-status request per IP for the
// finding Details field; the generic samples map can hold a 200 hit.
scannerSamples map[string]httpSample
// scannerDomainReqs and friends track the scanner profile per vhost so
// the distributed rollup only attributes a source to domains where the
// scanner shape actually crossed the detector gates. scannerDomainReqs is
// flat-keyed by ip+"\x00"+domain so a busy host's tens of thousands of
// one-off client IPs do not each allocate a nested map that the gates only
// read for the few IPs with errors.
scannerDomainReqs map[string]int
scannerDomainErr map[string]map[string]int
scanTime time.Time
// asnCrawl accumulates per-(scope, ASN) crawl fingerprints. Populated by
// observeASNCrawl from scan(); read by emitASNCrawl. Lazily allocated.
asnCrawl map[string]*asnCrawlScope
// Scanner thresholds are derived from cfg once per scan -- cfg is stable
// across a single domlogStats lifetime -- instead of on every parsed record.
scannerThreshComputed bool
scannerMinReq int
scannerPct int
scannerMinPaths int
scannerEnabled bool
}
func newDomlogStats() *domlogStats {
return newDomlogStatsAt(time.Now())
}
func newDomlogStatsAt(t time.Time) *domlogStats {
return &domlogStats{
wpLogin: make(map[string]int),
xmlrpc: make(map[string]int),
userEnum: make(map[string]int),
httpReqs: make(map[string]int),
uaCat: make(map[string]map[uaKind]int),
samples: make(map[string]httpSample),
domains: make(map[string]map[string]struct{}),
abuseDomains: make(map[string]map[string]map[string]struct{}),
scannerErr: make(map[string]int),
scannerReqs: make(map[string]int),
scannerPaths: make(map[string]map[string]struct{}),
scannerSamples: make(map[string]httpSample),
scannerDomainReqs: make(map[string]int),
scannerDomainErr: make(map[string]map[string]int),
scanTime: t,
}
}
// scan updates counters for one parsed record. cfg is allowed to be nil
// at Task 1 (parity tests pass nil); later tasks read thresholds from
// it. bot is consulted before any count so a verified Googlebot does
// not contribute to either legacy or new metrics.
func (s *domlogStats) scan(rec accessLogRecord, cfg *config.Config, bot botClassifier) {
ip := normalizeHTTPClientIP(clientIPForRecord(rec, cfg))
if ip == "" {
return
}
if cfg != nil && isInfraIP(ip, cfg.InfraIPs) {
return
}
if bot != nil && bot.IsVerifiedBot(ip, rec.UserAgent) {
removeChallengeIP(ip)
return
}
wpLoginHit := rec.Method == "POST" && strings.Contains(rec.URI, "wp-login.php")
xmlrpcHit := rec.Method == "POST" && strings.Contains(rec.URI, "xmlrpc.php")
userEnumHit := false
if rec.Method == "POST" {
if wpLoginHit {
s.wpLogin[ip]++
}
if xmlrpcHit {
s.xmlrpc[ip]++
}
}
if strings.Contains(rec.URI, "?author=") {
s.userEnum[ip]++
userEnumHit = true
} else if strings.Contains(rec.URI, "/wp-json/wp/v2/users") &&
!strings.Contains(rec.URI, "/users/me") {
s.userEnum[ip]++
userEnumHit = true
}
// Rate and UA counters only fire for requests that fall inside the
// flood window. Malformed timestamp lines still feed the legacy POST
// counters above but not rate or UA findings.
if withinHTTPFloodWindow(rec.Time, cfg, s.scanTime) {
if rec.Domain != "" {
set := s.domains[ip]
if set == nil {
set = make(map[string]struct{})
s.domains[ip] = set
}
set[rec.Domain] = struct{}{}
if wpLoginHit {
s.recordAbuseDomain("wp_login_bruteforce", ip, rec.Domain)
}
if xmlrpcHit {
s.recordAbuseDomain("xmlrpc_abuse", ip, rec.Domain)
}
if userEnumHit {
s.recordAbuseDomain("wp_user_enumeration", ip, rec.Domain)
}
s.recordAbuseDomain("http_request_flood", ip, rec.Domain)
}
if _, ok := s.samples[ip]; !ok {
s.samples[ip] = httpSample{Method: rec.Method, URI: rec.URI, UA: rec.UserAgent}
}
s.httpReqs[ip]++
_, _, _, scannerEnabled := s.scannerThresholds(cfg)
// Static display assets (images, styles, scripts, fonts, media)
// are excluded from the scanner profile entirely: a 404 on a
// browser sub-resource is a broken-asset signal, not URL
// enumeration. Dropping them from both the numerator (scannerErr)
// and the denominator (scannerReqs) stops a site with a missing
// CDN from making ordinary visitors look like scanners, without
// blinding the profile to real probes in the same window.
if scannerEnabled && !isDisplayAssetProbe(rec.URI) {
s.scannerReqs[ip]++
if rec.Domain != "" {
s.recordScannerDomainRequest(ip, rec.Domain)
}
if isScannerErrorStatus(rec.Status, cfg.Thresholds.HTTPScannerStatusCodes) {
path := probePath(rec.URI)
s.scannerErr[ip]++
if _, ok := s.scannerSamples[ip]; !ok {
s.scannerSamples[ip] = httpSample{Method: rec.Method, URI: rec.URI, UA: rec.UserAgent}
}
paths := s.scannerPaths[ip]
if paths == nil {
paths = make(map[string]struct{})
s.scannerPaths[ip] = paths
}
if len(paths) < httpScannerMaxTrackedPaths {
paths[path] = struct{}{}
}
if rec.Domain != "" {
s.recordScannerDomainError(ip, rec.Domain)
}
}
}
kind := classifyUA(rec.UserAgent, rec.Method)
if kind == uaKindClaimedBot {
// Static allowlist hits returned early through IsVerifiedBot
// above. For IPs outside the static range, check whether the
// async rDNS verifier has confirmed a negative result. Only
// promote to uaKindClaimedBotNegative when the cache has a
// definitive negative. The reversible pending-bot route is only
// safe when an active verifier can resolve the cache miss later;
// static-only/no-verifier paths keep the old flood/scanner
// behavior instead of softening spoofed bot UAs forever.
if cv, ok := bot.(confirmedNegativeBotClassifier); ok && cv.ConfirmedNegative(ip, rec.UserAgent) {
kind = uaKindClaimedBotNegative
} else if pv, ok := bot.(pendingBotVerificationClassifier); ok && pv.VerificationPending(ip, rec.UserAgent) {
// Verification still pending: not a confirmed spoof, so do not
// flag ua_spoof. Tracked separately so a flood/scan crawl from
// this IP routes to the reversible challenge, not a hard block.
kind = uaKindClaimedBotPending
} else {
kind = uaKindBrowser
}
}
if _, ok := s.uaCat[ip]; !ok {
s.uaCat[ip] = make(map[uaKind]int)
}
s.uaCat[ip][kind]++
if rec.Domain != "" && kind != uaKindBrowser && kind != uaKindClaimedBotPending {
s.recordAbuseDomain("http_ua_spoof", ip, rec.Domain)
}
}
// http_asn_crawl uses its OWN lookback window (default 60 min), which is
// wider than the flood window (default 5 min), so it must gate on its own
// window OUTSIDE the flood block; nesting it would cap the detector at the
// flood window and defeat spec section 4.1. ip is already infra/verified-bot
// filtered and proxy-resolved above, so it is valid here.
if asnCrawlWithinWindow(rec.Time, cfg, s.scanTime) {
s.observeASNCrawl(ip, rec, cfg)
}
}
func (s *domlogStats) recordAbuseDomain(check, ip, domain string) {
byIP := s.abuseDomains[check]
if byIP == nil {
byIP = make(map[string]map[string]struct{})
s.abuseDomains[check] = byIP
}
set := byIP[ip]
if set == nil {
set = make(map[string]struct{})
byIP[ip] = set
}
set[domain] = struct{}{}
}
// scannerDomainKey joins an IP and a domain into the flat scannerDomainReqs
// key. The NUL separator cannot appear in either part, so distinct (ip,domain)
// pairs never collide.
func scannerDomainKey(ip, domain string) string {
return ip + "\x00" + domain
}
// scannerThresholds returns the scanner-profile gates derived from cfg, caching
// them on first use. cfg does not change across a single scan, so this avoids
// re-deriving the thresholds on every parsed access-log record.
func (s *domlogStats) scannerThresholds(cfg *config.Config) (minReq, pct, minPaths int, ok bool) {
if !s.scannerThreshComputed {
s.scannerMinReq, s.scannerPct, s.scannerMinPaths, s.scannerEnabled = scannerProfileThresholds(cfg)
s.scannerThreshComputed = true
}
return s.scannerMinReq, s.scannerPct, s.scannerMinPaths, s.scannerEnabled
}
func (s *domlogStats) recordScannerDomainRequest(ip, domain string) {
s.scannerDomainReqs[scannerDomainKey(ip, domain)]++
}
func (s *domlogStats) recordScannerDomainError(ip, domain string) {
byDomain := s.scannerDomainErr[ip]
if byDomain == nil {
byDomain = make(map[string]int)
s.scannerDomainErr[ip] = byDomain
}
byDomain[domain]++
}
// emitLegacy returns the three pre-existing finding kinds. Kept
// separate from the new emit() (Tasks 3/4) so the parity test can
// assert "no new findings yet".
func (s *domlogStats) emitLegacy(cfg *config.Config) []alert.Finding {
var out []alert.Finding
// xmlrpc_threshold is operator-tunable; <= 0 disables the check entirely.
// Fall back to the package default only when called without a config (tests).
xmlrpcThr := xmlrpcThreshold
if cfg != nil {
xmlrpcThr = cfg.Thresholds.XMLRPCThreshold
}
for ip, count := range s.wpLogin {
if count >= wpLoginThreshold {
out = append(out, alert.Finding{
Severity: alert.Critical,
Check: "wp_login_bruteforce",
SourceIP: ip,
Message: formatLegacyMessage("WordPress login brute force", ip, count, "attempts"),
Details: "Aggregated across per-vhost access logs",
})
}
}
for ip, count := range s.xmlrpc {
if xmlrpcThr > 0 && count >= xmlrpcThr {
out = append(out, alert.Finding{
Severity: alert.Critical,
Check: "xmlrpc_abuse",
SourceIP: ip,
Message: formatLegacyMessage("XML-RPC abuse", ip, count, "requests"),
Details: "Aggregated across per-vhost access logs",
})
}
}
for ip, count := range s.userEnum {
if count >= 5 {
out = append(out, alert.Finding{
Severity: alert.High,
Check: "wp_user_enumeration",
SourceIP: ip,
Message: formatLegacyMessage("WordPress user enumeration", ip, count, "requests"),
Details: "Requests to /wp-json/wp/v2/users or ?author=",
})
}
}
return out
}
// emit produces all finding kinds from a single populated domlogStats.
// Legacy three kinds come first (so existing callers still get them when
// running through emit), then http_request_flood, then http_ua_spoof.
func (s *domlogStats) emit(cfg *config.Config) []alert.Finding {
out := s.emitLegacy(cfg)
if cfg != nil && cfg.Thresholds.HTTPFloodThreshold > 0 {
for ip, count := range s.httpReqs {
if count < cfg.Thresholds.HTTPFloodThreshold {
continue
}
if s.isPendingClaimedBot(ip) {
// Routed to http_claimed_bot_unverified (challenge) below so a
// real crawler mid-verification is not hard-blocked.
continue
}
sample := s.samples[ip]
out = append(out, alert.Finding{
Severity: alert.High,
Check: "http_request_flood",
SourceIP: ip,
Domain: s.singleDomain(ip),
Message: "HTTP request flood from " + ip + ": " + itoa(count) + " requests" + s.vhostSuffix(ip),
Details: "Sample: " + sample.Method + " " + sample.URI + " UA=" + truncate(sample.UA, 120),
})
}
}
if cfg != nil {
threshold := cfg.Thresholds.HTTPUASpoofThreshold
if threshold <= 0 {
threshold = 30
}
for ip, byKind := range s.uaCat {
// Require sustained spoofing before a hard block. A single
// rDNS-failed bot-UA request is too often a legitimate residential
// or mobile client that happens to send a bot-like User-Agent;
// gating on the same threshold as the other spoof kinds keeps real
// spoof-crawlers (which make many requests) blocked.
if hits, ok := byKind[uaKindClaimedBotNegative]; ok && hits >= threshold {
out = append(out, s.makeUASpoofFinding(ip,
"claimed search-engine bot failed rDNS verification",
s.samples[ip], hits))
continue
}
if hits, ok := byKind[uaKindKnownScanner]; ok && hits > 0 {
out = append(out, s.makeUASpoofFinding(ip, "known scanner UA",
s.samples[ip], hits))
continue
}
if hits := byKind[uaKindWPSpoofPingback]; hits >= threshold {
out = append(out, s.makeUASpoofFinding(ip,
"WordPress/<ver> UA on GET (pingback spoof)",
s.samples[ip], hits))
continue
}
if cfg.Thresholds.HTTPUAScriptingEnabled {
if hits := byKind[uaKindScriptingLang]; hits >= threshold {
out = append(out, s.makeUASpoofFinding(ip,
"scripting-language UA (curl/python/etc.)",
s.samples[ip], hits))
continue
}
}
if cfg.Thresholds.HTTPUAHeadlessEnabled {
if hits := byKind[uaKindHeadless]; hits >= threshold {
out = append(out, s.makeUASpoofFinding(ip,
"headless browser UA", s.samples[ip], hits))
continue
}
}
if cfg.Thresholds.HTTPUAEmptyEnabled {
if hits := byKind[uaKindEmpty]; hits >= threshold {
out = append(out, s.makeUASpoofFinding(ip,
"empty/dash User-Agent", s.samples[ip], hits))
continue
}
}
}
}
out = append(out, s.emitScannerProfile(cfg)...)
out = append(out, s.emitClaimedBotUnverified(cfg)...)
// Distributed attack: many distinct abusive IPs hitting one vhost.
// Built from the per-IP findings already emitted above, so only IPs
// that crossed an abuse threshold count -- a popular site's normal
// visitor spread never trips it.
out = append(out, s.emitDistributedFlood(cfg, out)...)
out = append(out, s.emitASNCrawl(cfg)...)
return out
}
// httpScannerMaxTrackedPaths bounds the per-IP distinct probe-path set.
// A scanner cycling unique URLs past the cap keeps incrementing the
// error counter, and any sane min_distinct_paths threshold sits far
// below the cap, so detection quality does not depend on growth beyond it.
const httpScannerMaxTrackedPaths = config.HTTPScannerMaxDistinctPaths
// emitScannerProfile produces http_scanner_profile findings: source IPs
// whose in-window traffic is almost entirely probe-error responses spread
// across many distinct paths -- the shape of URL enumeration for
// downloadable files, exposed backups, and dormant shells. Volume,
// error-rate, and path-breadth gates must all pass so that dead
// bookmarks, broken assets, and site migrations stay out of scope.
func (s *domlogStats) emitScannerProfile(cfg *config.Config) []alert.Finding {
minReq, pct, minPaths, ok := s.scannerThresholds(cfg)
if !ok {
return nil
}
var out []alert.Finding
for ip, errs := range s.scannerErr {
total := s.scannerReqs[ip]
paths := len(s.scannerPaths[ip])
if !scannerProfilePasses(total, errs, paths, minReq, pct, minPaths) {
continue
}
if s.isPendingClaimedBot(ip) {
// Routed to http_claimed_bot_unverified (challenge) instead of a
// hard scanner-profile block while verification is pending.
continue
}
for _, domain := range s.scannerProfileDomains(ip, minReq, pct) {
s.recordAbuseDomain("http_scanner_profile", ip, domain)
}
sample := s.scannerSamples[ip]
out = append(out, alert.Finding{
Severity: alert.High,
Check: "http_scanner_profile",
SourceIP: ip,
Domain: s.singleDomain(ip),
Message: "URL scanner profile from " + ip + ": " + itoa(errs) + " of " + itoa(total) +
" requests hit probe-error statuses across " + itoa(paths) + " distinct paths" + s.vhostSuffix(ip),
Details: "Sample: " + sample.Method + " " + sample.URI + " UA=" + truncate(sample.UA, 120),
})
}
return out
}
// isPendingClaimedBot reports whether ip's in-window traffic is dominated by a
// claimed-bot UA whose rDNS verification has not resolved. The claimed-bot
// requests must be a strict majority so an attacker cannot downgrade a hard
// block to the softer challenge route by mixing in bot-UA requests.
func (s *domlogStats) isPendingClaimedBot(ip string) bool {
pending := s.uaCat[ip][uaKindClaimedBotPending]
if pending == 0 {
return false
}
return pending*2 > s.httpReqs[ip]
}
// emitClaimedBotUnverified emits one http_claimed_bot_unverified finding per
// pending-claimed-bot IP whose flood or scanner-profile volume crossed a hard
// threshold. The check is challengeable (routes to the PoW gate when challenge
// is enabled) and blockable (hard-blocked when it is not), so a real crawler
// mid-verification clears itself while a spoofer that cannot solve the
// challenge stays blocked.
func (s *domlogStats) emitClaimedBotUnverified(cfg *config.Config) []alert.Finding {
if cfg == nil {
return nil
}
minReq, pct, minPaths, scannerOK := s.scannerThresholds(cfg)
floodThreshold := cfg.Thresholds.HTTPFloodThreshold
var out []alert.Finding
for ip := range s.uaCat {
if !s.isPendingClaimedBot(ip) {
continue
}
total := s.httpReqs[ip]
flood := floodThreshold > 0 && total >= floodThreshold
scanner := scannerOK && scannerProfilePasses(s.scannerReqs[ip], s.scannerErr[ip], len(s.scannerPaths[ip]), minReq, pct, minPaths)
if !flood && !scanner {
continue
}
reason := "request flood"
sample := s.samples[ip]
if scanner && !flood {
reason = "scanner-profile crawl"
sample = s.scannerSamples[ip]
}
s.recordClaimedBotUnverifiedDomains(ip, flood, scanner, minReq, pct)
out = append(out, alert.Finding{
Severity: alert.High,
Check: "http_claimed_bot_unverified",
SourceIP: ip,
Domain: s.singleDomain(ip),
Message: "Unverified claimed bot from " + ip + ": " + reason + " (" + itoa(total) +
" requests, rDNS not yet confirmed)" + s.vhostSuffix(ip),
Details: "Sample: " + sample.Method + " " + sample.URI + " UA=" + truncate(sample.UA, 120),
})
}
return out
}
func (s *domlogStats) recordClaimedBotUnverifiedDomains(ip string, flood, scanner bool, minReq, pct int) {
if flood {
for domain := range s.abuseDomains["http_request_flood"][ip] {
s.recordAbuseDomain("http_claimed_bot_unverified", ip, domain)
}
}
if scanner {
for _, domain := range s.scannerProfileDomains(ip, minReq, pct) {
s.recordAbuseDomain("http_claimed_bot_unverified", ip, domain)
}
}
}
func scannerProfileThresholds(cfg *config.Config) (minReq, pct, minPaths int, ok bool) {
if cfg == nil {
return 0, 0, 0, false
}
minReq = cfg.Thresholds.HTTPScannerMinRequests
if minReq <= 0 {
return 0, 0, 0, false
}
pct = cfg.Thresholds.HTTPScannerErrorPct
if pct <= 0 {
pct = config.DefaultHTTPScannerErrorPct
}
if pct > 100 {
pct = 100
}
minPaths = cfg.Thresholds.HTTPScannerMinDistinctPaths
if minPaths <= 0 {
minPaths = config.DefaultHTTPScannerMinDistinctPaths
}
if minPaths > httpScannerMaxTrackedPaths {
minPaths = httpScannerMaxTrackedPaths
}
return minReq, pct, minPaths, true
}
func scannerProfilePasses(total, errs, paths, minReq, pct, minPaths int) bool {
if total < minReq {
return false
}
if errs*100 < total*pct {
return false
}
return paths >= minPaths
}
// scannerDomainDefaultMinErrors is the small absolute floor of probe-error
// hits a single vhost must receive before a confirmed per-IP scanner is
// attributed to it for the distributed rollup. The per-IP gates (minReq,
// minPaths) already proved the source is a scanner; per vhost we only require
// a few errors plus the same error-rate gate, so a scanner spread thin across
// many vhosts still feeds the rollup while the default scanner thresholds do
// not attribute a vhost that caught one incidental 404.
const scannerDomainDefaultMinErrors = 3
func scannerDomainErrorFloor(minReq int) int {
if minReq <= 0 {
return scannerDomainDefaultMinErrors
}
if minReq < scannerDomainDefaultMinErrors {
return minReq
}
return scannerDomainDefaultMinErrors
}
// scannerProfileDomains returns the vhosts a confirmed scanner IP should be
// attributed to in the distributed rollup. The full per-IP minimum-request and
// min-path gates are deliberately NOT reused per vhost: a scanner that sprays a
// few probes across many vhosts trips the per-IP profile in aggregate but lands
// only a handful of hits on each vhost, so requiring the per-IP minimums per
// vhost would drop it from the rollup. Instead each vhost needs a small
// absolute error floor and the same error-rate gate.
func (s *domlogStats) scannerProfileDomains(ip string, minReq, pct int) []string {
errsByDomain := s.scannerDomainErr[ip]
if len(errsByDomain) == 0 {
return nil
}
var domains []string
errorFloor := scannerDomainErrorFloor(minReq)
for domain, errs := range errsByDomain {
total := s.scannerDomainReqs[scannerDomainKey(ip, domain)]
if scannerDomainAttributes(total, errs, pct, errorFloor) {
domains = append(domains, domain)
}
}
sort.Strings(domains)
return domains
}
func scannerDomainAttributes(total, errs, pct, minErrors int) bool {
if errs < minErrors {
return false
}
return errs*100 >= total*pct
}
// isScannerErrorStatus reports whether status belongs to the configured
// probe-error set. Linear scan: the set is a handful of entries, so this
// beats building a lookup map on a per-record hot path.
func isScannerErrorStatus(status int, codes []int) bool {
if status == 0 {
return false
}
if len(codes) == 0 {
return status == 404 || status == 403
}
for _, c := range codes {
if status == c {
return true
}
}
return false
}
// probePath strips the query string and fragment so cache-buster style
// queries on one missing endpoint count as a single probe path.
func probePath(uri string) string {
if i := strings.IndexAny(uri, "?#"); i >= 0 {
uri = uri[:i]
}
if uri == "" {
return "/"
}
return uri
}
// displayAssetExts are static, browser-rendered sub-resource extensions.
// A 404 on one of these is a broken-asset signal -- the static-file
// handler looked for a file and did not find it, executing no code and
// disclosing nothing -- not URL enumeration. Excluding them keeps a site
// whose CDN is missing its images, styles, or scripts from making every
// ordinary visitor look like a scanner. Archives, code, configs, dumps,
// and extensionless paths -- the targets a real scanner enumerates -- are
// deliberately absent and keep counting toward the profile.
var displayAssetExts = map[string]struct{}{
".gif": {}, ".jpg": {}, ".jpeg": {}, ".png": {}, ".webp": {}, ".bmp": {},
".svg": {}, ".ico": {}, ".avif": {}, ".tif": {}, ".tiff": {},
".css": {}, ".js": {}, ".mjs": {}, ".cjs": {},
".woff": {}, ".woff2": {}, ".ttf": {}, ".eot": {}, ".otf": {},
".mp4": {}, ".webm": {}, ".ogg": {}, ".mp3": {}, ".wav": {},
".m4a": {}, ".mov": {}, ".avi": {}, ".flac": {},
}
// scannerProbeExts are extensions commonly used when probing for code,
// configs, source maps, dumps, archives, and backups. If one appears before a
// final display-asset suffix (for example shell.php.gif), the request still
// has scanner shape and must not be hidden by the asset exclusion.
var scannerProbeExts = map[string]struct{}{
".php": {}, ".phtml": {}, ".phps": {}, ".php3": {}, ".php4": {}, ".php5": {}, ".php7": {},
".asp": {}, ".aspx": {}, ".ashx": {}, ".asmx": {}, ".jsp": {}, ".jspx": {}, ".cfm": {},
".cgi": {}, ".pl": {}, ".py": {}, ".rb": {}, ".sh": {}, ".bash": {}, ".zsh": {},
".env": {}, ".conf": {}, ".config": {}, ".ini": {}, ".yaml": {}, ".yml": {}, ".json": {}, ".xml": {},
".sql": {}, ".dump": {}, ".db": {}, ".sqlite": {}, ".sqlite3": {}, ".log": {},
".zip": {}, ".rar": {}, ".7z": {}, ".tar": {}, ".tgz": {}, ".gz": {}, ".bz2": {}, ".xz": {},
".bak": {}, ".backup": {}, ".old": {}, ".orig": {}, ".save": {}, ".swp": {}, ".tmp": {},
".map": {},
}
// isDisplayAssetProbe reports whether uri's final path segment ends in one
// of displayAssetExts. The query string and fragment are stripped first so
// a cache-buster suffix cannot hide the extension.
func isDisplayAssetProbe(uri string) bool {
path := probePath(uri)
if i := strings.LastIndexByte(path, '/'); i >= 0 {
path = path[i+1:]
}
if unescaped, err := url.PathUnescape(path); err == nil {
path = unescaped
}
dot := strings.LastIndexByte(path, '.')
if dot <= 0 {
return false
}
if _, ok := displayAssetExts[strings.ToLower(path[dot:])]; !ok {
return false
}
return !hasEmbeddedScannerProbeExt(path[:dot])
}
func hasEmbeddedScannerProbeExt(pathPrefix string) bool {
parts := strings.Split(strings.ToLower(pathPrefix), ".")
for _, part := range parts[1:] {
if part == "" {
continue
}
if _, ok := scannerProbeExts["."+part]; ok {
return true
}
}
return false
}
// httpAbuseChecks are the per-IP finding kinds that mark a source IP as
// abusive for the distributed-attack rollup.
var httpAbuseChecks = map[string]struct{}{
"wp_login_bruteforce": {},
"xmlrpc_abuse": {},
"wp_user_enumeration": {},
"http_request_flood": {},
"http_claimed_bot_unverified": {},
"http_ua_spoof": {},
"http_scanner_profile": {},
}
// emitDistributedFlood rolls the per-IP HTTP-abuse findings up per vhost:
// when at least HTTPDistributedMinIPs distinct abusive IPs hit one vhost
// in this window, it emits a single http_distributed_flood finding for
// that vhost. The per-IP findings still stand; this adds the
// many-IPs-one-target view (botnet / distributed brute-force) that the
// per-IP and per-source-IP-spray paths cannot see. Disabled when the
// threshold is <= 0.
func (s *domlogStats) emitDistributedFlood(cfg *config.Config, prior []alert.Finding) []alert.Finding {
if cfg == nil {
return nil
}
minIPs := cfg.Thresholds.HTTPDistributedMinIPs
if minIPs <= 0 {
return nil
}
domainIPs := map[string]map[string]struct{}{}
for _, f := range prior {
if _, ok := httpAbuseChecks[f.Check]; !ok || f.SourceIP == "" {
continue
}
for dom := range s.abuseDomains[f.Check][f.SourceIP] {
if domainIPs[dom] == nil {
domainIPs[dom] = map[string]struct{}{}
}
domainIPs[dom][f.SourceIP] = struct{}{}
}
}
domains := make([]string, 0, len(domainIPs))
for dom := range domainIPs {
domains = append(domains, dom)
}
sort.Strings(domains)
var out []alert.Finding
for _, dom := range domains {
n := len(domainIPs[dom])
if n < minIPs {
continue
}
out = append(out, alert.Finding{
Severity: alert.High,
Check: "http_distributed_flood",
Domain: dom,
Message: fmt.Sprintf("Distributed HTTP attack on %s: %d distinct abusive source IPs", dom, n),
Details: fmt.Sprintf("%d source IPs each tripped an HTTP-abuse threshold against %s in this window. "+
"Likely a botnet or distributed brute-force; consider edge rate-limiting or a challenge.", n, dom),
Timestamp: time.Now(),
})
}
return out
}
func (s *domlogStats) makeUASpoofFinding(ip, reason string, sample httpSample, hits int) alert.Finding {
return alert.Finding{
Severity: alert.Critical,
Check: "http_ua_spoof",
SourceIP: ip,
Domain: s.singleDomain(ip),
Message: "User-Agent spoof from " + ip + ": " + reason + s.vhostSuffix(ip),
Details: "Hits: " + itoa(hits) + ", Sample: " + sample.Method + " " +
sample.URI + " UA=" + truncate(sample.UA, 200),
}
}
// singleDomain returns the one vhost ip touched, or "" when zero or more
// than one (the per-IP aggregate spans several vhosts, so no single domain
// attributes the finding). vhostSuffix renders the cross-site spread for
// the operator-facing message.
func (s *domlogStats) singleDomain(ip string) string {
set := s.domains[ip]
if len(set) != 1 {
return ""
}
for d := range set {
return d
}
return ""
}
func (s *domlogStats) vhostSuffix(ip string) string {
if n := len(s.domains[ip]); n > 1 {
return " across " + itoa(n) + " vhosts"
}
return ""
}
// withinHTTPFloodWindow reports whether a log timestamp falls inside the
// configured flood rate window relative to the scan start time. Timestamps
// in the future (up to one clock-skew minute) are accepted. Zero timestamps
// from malformed lines return false so they do not contribute to rate counts.
func withinHTTPFloodWindow(ts time.Time, cfg *config.Config, now time.Time) bool {
if ts.IsZero() || cfg == nil {
return false
}
windowMin := cfg.Thresholds.HTTPFloodWindowMin
if windowMin <= 0 {
windowMin = 5
}
cutoff := now.Add(-time.Duration(windowMin) * time.Minute)
return !ts.Before(cutoff) && !ts.After(now.Add(time.Minute))
}
func formatLegacyMessage(kind, ip string, n int, unit string) string {
return kind + " from " + ip + ": " + itoa(n) + " " + unit
}
func itoa(n int) string {
if n == 0 {
return "0"
}
var buf [20]byte
pos := len(buf)
neg := n < 0
if neg {
n = -n
}
for n > 0 {
pos--
buf[pos] = byte('0' + n%10)
n /= 10
}
if neg {
pos--
buf[pos] = '-'
}
return string(buf[pos:])
}
// parseAccessLogRecord parses one Combined Log Format line into an
// accessLogRecord. It does NOT use strings.Fields because quoted
// request/referer/user-agent fields can contain spaces.
//
// Format:
//
// <ip> <ident> <user> [<time>] "<method> <uri> <proto>" <status> <bytes> "<ref>" "<ua>" ["<vhost>"]
//
// Returns ok=false for any line that cannot be parsed. Never panics.
func parseAccessLogRecord(line string) (accessLogRecord, bool) {
const maxUALen = 512
const maxURILen = 4096
var rec accessLogRecord
// IP is everything up to the first space.
sp := strings.IndexByte(line, ' ')
if sp <= 0 {
return rec, false
}
rec.RemoteIP = line[:sp]
rest := line[sp+1:]
// Skip ident, user (two single-token fields). Loose: we just need to
// land at the [time] bracket.
br := strings.IndexByte(rest, '[')
if br < 0 {
return rec, false
}
rest = rest[br+1:]
closeBr := strings.IndexByte(rest, ']')
if closeBr < 0 {
return rec, false
}
timeStr := rest[:closeBr]
rest = rest[closeBr+1:]
// time format: 02/Jan/2006:15:04:05 -0700
t, err := time.Parse("02/Jan/2006:15:04:05 -0700", timeStr)
if err == nil {
rec.Time = t
}
// Request quoted field.
q1 := strings.IndexByte(rest, '"')
if q1 < 0 {
return rec, false
}
rest = rest[q1+1:]
q2 := strings.IndexByte(rest, '"')
if q2 < 0 {
return rec, false
}
request := rest[:q2]
rest = rest[q2+1:]
parts := strings.SplitN(request, " ", 3)
if len(parts) >= 1 {
rec.Method = parts[0]
}
if len(parts) >= 2 {
uri := parts[1]
if len(uri) > maxURILen {
uri = uri[:maxURILen]
}
rec.URI = uri
}
// status (skip leading spaces)
rest = strings.TrimLeft(rest, " ")
end := strings.IndexByte(rest, ' ')
if end > 0 {
rec.Status = atoiSafe(rest[:end])
rest = rest[end+1:]
}
// bytes field -- skip leading spaces then advance past the token
rest = strings.TrimLeft(rest, " ")
end = strings.IndexByte(rest, ' ')
if end > 0 {
rest = rest[end+1:]
} else {
// no more fields after bytes
return rec, true
}
// referer quoted field (skipped).
q1 = strings.IndexByte(rest, '"')
if q1 < 0 {
return rec, false
}
rest = rest[q1+1:]
q2 = strings.IndexByte(rest, '"')
if q2 < 0 {
return rec, false
}
rest = rest[q2+1:]
// UA quoted field.
q1 = strings.IndexByte(rest, '"')
if q1 < 0 {
return rec, true // no UA present is fine
}
rest = rest[q1+1:]
q2 = strings.IndexByte(rest, '"')
if q2 < 0 {
return rec, false
}
ua := rest[:q2]
if len(ua) > maxUALen {
ua = ua[:maxUALen]
}
rec.UserAgent = ua
rest = rest[q2+1:]
// Optional quoted extensions. cPanel may append a quoted vhost after
// UA. Custom proxy formats may append an X-Forwarded-For value. Only
// retain a quoted extension that parses as an IP list; clientIPForRecord
// still ignores it unless RemoteIP is a configured trusted proxy.
for {
q1 = strings.IndexByte(rest, '"')
if q1 < 0 {
break
}
rest = rest[q1+1:]
q2 = strings.IndexByte(rest, '"')
if q2 < 0 {
return rec, false
}
extra := rest[:q2]
if looksLikeXFF(extra) {
rec.XFF = extra
}
rest = rest[q2+1:]
}
return rec, true
}
func atoiSafe(s string) int {
n := 0
for i := 0; i < len(s); i++ {
c := s[i]
if c < '0' || c > '9' {
break
}
n = n*10 + int(c-'0')
}
return n
}
func clientIPForRecord(rec accessLogRecord, cfg *config.Config) string {
if cfg == nil || len(cfg.WebServer.TrustedProxies) == 0 || rec.XFF == "" {
return rec.RemoteIP
}
if !isTrustedProxy(rec.RemoteIP, cfg.WebServer.TrustedProxies) {
return rec.RemoteIP
}
// A trusted direct proxy appends the peer it observed to the end of
// X-Forwarded-For. Use that entry only; earlier entries can come from
// the client.
parts := strings.Split(rec.XFF, ",")
for i := len(parts) - 1; i >= 0; i-- {
ip := strings.TrimSpace(parts[i])
if net.ParseIP(ip) == nil {
continue
}
return ip
}
return rec.RemoteIP
}
func normalizeHTTPClientIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if host, _, err := net.SplitHostPort(raw); err == nil {
raw = host
}
raw = strings.Trim(raw, "[]")
ip := net.ParseIP(raw)
if ip == nil || ip.IsLoopback() || ip.IsUnspecified() {
return ""
}
return ip.String()
}
// isTrustedProxy returns true when addr matches any entry in proxies (exact
// IP or CIDR). Entries that fail to parse are skipped.
func isTrustedProxy(addr string, proxies []string) bool {
addr = strings.TrimSpace(addr)
parsed := net.ParseIP(addr)
if parsed == nil {
return false
}
for _, entry := range proxies {
entry = strings.TrimSpace(entry)
if entry == "" {
continue
}
if _, cidr, err := net.ParseCIDR(entry); err == nil {
if cidr.Contains(parsed) {
return true
}
continue
}
if ip := net.ParseIP(entry); ip != nil && ip.Equal(parsed) {
return true
}
}
return false
}
func looksLikeXFF(raw string) bool {
for _, part := range strings.Split(raw, ",") {
if net.ParseIP(strings.TrimSpace(part)) != nil {
return true
}
}
return false
}
// classifyUA maps a User-Agent string to a uaKind. Matching is performed
// on a lower-cased copy of the UA because scanner and impersonation tools
// routinely vary case. Precedence order follows spec section 6: scanner
// signatures win over claimed-bot, claimed-bot wins over headless, etc.
func classifyUA(ua, method string) uaKind {
const maxUALen = 512
if len(ua) > maxUALen {
ua = ua[:maxUALen]
}
if ua == "" || ua == "-" {
return uaKindEmpty
}
low := strings.ToLower(ua)
for _, s := range knownScannerSubstrings {
if strings.Contains(low, s) {
return uaKindKnownScanner
}
}
// WordPress pingback UA on a GET request is illegal: legitimate
// pingback clients always POST. A GET with this UA is a content
// scraper or probe spoofing the pingback agent.
if method == "GET" && strings.HasPrefix(low, "wordpress/") {
return uaKindWPSpoofPingback
}
for _, s := range claimedBotSubstrings {
if strings.Contains(low, s) {
return uaKindClaimedBot
}
}
// Operator-configured bots (reputation.verified_bots) classify as
// claimed bots too, so an impostor reusing the UA is caught as a spoof.
if threatintel.OperatorBotFromUA(low) != "" {
return uaKindClaimedBot
}
for _, s := range headlessSubstrings {
if strings.Contains(low, s) {
return uaKindHeadless
}
}
for _, s := range scriptingSubstrings {
if strings.Contains(low, s) {
return uaKindScriptingLang
}
}
return uaKindBrowser
}
var (
knownScannerSubstrings = []string{
"nikto", "sqlmap", "acunetix", "nmap ", "masscan", "wpscan",
"nuclei", "dirbuster", "gobuster", "feroxbuster",
}
claimedBotSubstrings = []string{
"googlebot", "bingbot", "applebot", "duckduckbot", "yandexbot",
"baiduspider", "facebookexternalhit", "twitterbot",
// Appendix A bots plus AI crawlers verified by published ranges.
"amazonbot", "gptbot", "chatgpt-user", "oai-searchbot",
"claudebot", "claude-user", "claude-searchbot",
"perplexitybot", "meta-externalagent", "meta-webindexer", "bravebot",
"seranking",
}
headlessSubstrings = []string{
"headlesschrome", "phantomjs", "puppeteer", "playwright",
}
scriptingSubstrings = []string{
"python-requests/", "curl/", "go-http-client/", "java/", "wget/",
"libwww-perl/", "node-fetch/",
}
)
// staticAllowlistClassifier consults only DNS-free bot verification: embedded
// static IP ranges and operator-configured IP ranges. rDNS verification for
// IPs outside those ranges is handled by verifyingClassifier.
type staticAllowlistClassifier struct{}
func (staticAllowlistClassifier) IsVerifiedBot(ipStr, ua string) bool {
bot := threatintel.ClaimedBotFromUA(ua)
if bot == "" {
return false
}
ip := net.ParseIP(ipStr)
if ip == nil {
return false
}
if threatintel.DefaultRanges().IPInBot(ip, bot) {
return true
}
return threatintel.OperatorBotIPVerified(bot, ip)
}
// verifyingClassifier consults the static allowlist first, then the
// rDNS verify cache. Cache misses enqueue an async job and return
// false (treat as unverified for this scan cycle).
type verifyingClassifier struct {
async *threatintel.AsyncBotVerifier
cacheGet func(net.IP, string) (bool, bool)
}
func newVerifyingClassifier(async *threatintel.AsyncBotVerifier,
cacheGet func(net.IP, string) (bool, bool)) verifyingClassifier {
return verifyingClassifier{async: async, cacheGet: cacheGet}
}
func (c verifyingClassifier) IsVerifiedBot(ipStr, ua string) bool {
bot := threatintel.ClaimedBotFromUA(ua)
if bot == "" {
return false
}
ip := net.ParseIP(ipStr)
if ip == nil {
return false
}
// Static range is the fast positive path.
if threatintel.DefaultRanges().IPInBot(ip, bot) {
return true
}
// Operator IP-range bots (AI agents) verify synchronously, no rDNS.
if threatintel.OperatorBotIPVerified(bot, ip) {
return true
}
// Cache lookup: valid positive -> verified, valid negative -> false
// (ConfirmedNegative handles it), no entry -> enqueue and fail open.
if c.cacheGet != nil {
if verified, valid := c.cacheGet(ip, bot); valid {
return verified
}
}
// Enqueue async verification only when a later scan can read the result;
// otherwise there is no bounded pending window to route through challenge.
if c.async != nil && c.cacheGet != nil {
c.async.Enqueue(ip, bot)
}
return false
}
// ConfirmedNegative reports whether the rDNS cache has a definitive
// negative result for this IP+UA pair. Called from scan() to decide
// whether to promote uaKindClaimedBot to uaKindClaimedBotNegative.
func (c verifyingClassifier) ConfirmedNegative(ipStr, ua string) bool {
bot := threatintel.ClaimedBotFromUA(ua)
if bot == "" {
return false
}
ip := net.ParseIP(ipStr)
if ip == nil || c.cacheGet == nil {
return false
}
verified, valid := c.cacheGet(ip, bot)
return valid && !verified
}
func (c verifyingClassifier) VerificationPending(ipStr, ua string) bool {
bot := threatintel.ClaimedBotFromUA(ua)
if bot == "" {
return false
}
ip := net.ParseIP(ipStr)
if ip == nil || c.async == nil || c.cacheGet == nil {
return false
}
if threatintel.DefaultRanges().IPInBot(ip, bot) || threatintel.OperatorBotIPVerified(bot, ip) {
return false
}
if _, valid := c.cacheGet(ip, bot); valid {
return false
}
return c.async.Pending(ip, bot)
}
var (
globalBotVerifier *threatintel.AsyncBotVerifier
globalBotGet func(net.IP, string) (bool, bool)
botMu sync.RWMutex
)
// SetBotVerifier installs the daemon-lifetime async verifier and cache
// reader. Called from daemon.go after the store and goroutine are ready.
func SetBotVerifier(v *threatintel.AsyncBotVerifier, get func(net.IP, string) (bool, bool)) {
botMu.Lock()
defer botMu.Unlock()
globalBotVerifier = v
globalBotGet = get
}
// currentBotClassifier returns the appropriate botClassifier based on
// config. When bot_verify_enabled is false, falls back to the
// static-only classifier so DNS calls are never made.
func currentBotClassifier(cfg *config.Config) botClassifier {
if cfg == nil || !cfg.BotVerifyEnabled() {
return staticAllowlistClassifier{}
}
botMu.RLock()
defer botMu.RUnlock()
return newVerifyingClassifier(globalBotVerifier, globalBotGet)
}
// Package checks: http_asn_crawl detector — single-ASN distributed crawl of
// uncacheable URLs saturating one account's PHP pool. See
// docs/superpowers/specs/2026-06-24-http-asn-crawl-detector-design.md.
package checks
import (
"fmt"
"net"
"net/url"
"path"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
)
// httpASNCrawlStaticExts are path extensions whose responses are cacheable
// static assets; requests for them never reach PHP, so they are not
// "expensive" for this detector.
var httpASNCrawlStaticExts = map[string]struct{}{
"jpg": {}, "jpeg": {}, "png": {}, "gif": {}, "webp": {}, "svg": {}, "ico": {},
"bmp": {}, "css": {}, "js": {}, "mjs": {}, "map": {}, "woff": {}, "woff2": {},
"ttf": {}, "eot": {}, "otf": {}, "mp4": {}, "webm": {}, "ogg": {}, "mp3": {},
"pdf": {}, "zip": {}, "gz": {}, "avif": {},
}
// httpASNCrawlExpensive reports whether a request is a dynamic, uncacheable
// hit that reaches PHP: a GET or HEAD with a query string whose path extension
// is not a static asset.
func httpASNCrawlExpensive(rec accessLogRecord) bool {
if rec.Method != "GET" && rec.Method != "HEAD" {
return false
}
q := strings.IndexByte(rec.URI, '?')
if q < 0 || q == len(rec.URI)-1 {
return false
}
p := rec.URI[:q]
ext := strings.ToLower(strings.TrimPrefix(path.Ext(p), "."))
if ext == "" {
return true
}
_, isStatic := httpASNCrawlStaticExts[ext]
return !isStatic
}
// httpASNCrawlAmplifyKeys are query parameter names that signal an
// expensive layered-nav/search request; their presence raises severity.
var httpASNCrawlAmplifyKeys = map[string]struct{}{
"orderby": {}, "add-to-cart": {}, "s": {}, "paged": {}, "product-page": {},
}
// httpASNCrawlAmplified reports whether the URI's query carries a known
// expensive layered-nav/search parameter. Key names are matched
// case-insensitively; values alone never match.
func httpASNCrawlAmplified(uri string) bool {
q := strings.IndexByte(uri, '?')
if q < 0 {
return false
}
for raw := range strings.SplitSeq(uri[q+1:], "&") {
key, _, _ := strings.Cut(raw, "=")
key, err := url.QueryUnescape(key)
if err != nil {
continue
}
lk := strings.ToLower(key)
if strings.HasPrefix(lk, "filter_") || strings.HasPrefix(lk, "query_type_") {
return true
}
if _, ok := httpASNCrawlAmplifyKeys[lk]; ok {
return true
}
}
return false
}
// asnCrawlASN accumulates one ASN's footprint within one scope.
type asnCrawlASN struct {
org string
expensive int
total int
amplified int
domains map[string]struct{}
samples []string // up to 5 distinct expensive URIs (<=200 bytes)
ips map[string]struct{} // distinct source IPs, capped
ipsCapped bool
cidr24 map[string]int // observed /24 (v4) or /64 (v6) -> count
}
func (a *asnCrawlASN) distinctIPs() int {
// saturated at cap; callers may render as ">=cap" at emit
return len(a.ips)
}
// asnCrawlScope is one account-or-domain scope's per-ASN map plus the
// scope-wide expensive denominator (excludes reverse-proxy traffic).
type asnCrawlScope struct {
byASN map[uint]*asnCrawlASN
scopeExpensive int
}
func (s *domlogStats) observeASNCrawl(ip string, rec accessLogRecord, cfg *config.Config) {
if cfg == nil || cfg.Thresholds.HTTPASNCrawlMinIPs <= 0 {
return // detector disabled
}
if rec.Domain == "" {
// Central access logs (and malformed paths) carry no domain/account,
// so they cannot be PHP-pool correlated and must not feed this
// detector (spec section 4). Per-domain domlogs always set Domain;
// without this guard they key a degenerate catch-all "domain:" scope
// and double-count the share denominator.
return
}
if !httpASNCrawlExpensive(rec) {
return
}
lookup := CurrentASNLookup()
if lookup == nil {
return
}
asn, org := lookup(ip)
if asn == 0 {
// Unknown ASN: still counts in the scope denominator so known-ASN
// share is not inflated, but cannot form its own fingerprint.
s.asnCrawlScopeFor(rec).scopeExpensive++
return
}
if uintInSlice(asn, cfg.Thresholds.HTTPASNCrawlReverseProxyASNs) {
return // reverse-proxy: dropped from numerator AND denominator
}
// v1 limitation (spec section 4.3): real-client attribution behind a
// reverse proxy/CDN depends on clientIPForRecord restoring the client from
// XFF, which only happens when web_server.trusted_proxies is configured.
// Operators fronting hosts with a CDN MUST set trusted_proxies so the real
// client IP (and its ASN) is attributed here; without it a CDN-fronted
// distributed crawl is seen as proxy-ASN traffic and not flagged (a false
// negative). The catastrophic false positive is still prevented because the
// proxy ASN is never emitted or tempbanned. A general CDN-aware fix spans
// all http-abuse detectors and is out of scope for this detector.
scope := s.asnCrawlScopeFor(rec)
scope.scopeExpensive++
a := scope.byASN[asn]
if a == nil {
a = &asnCrawlASN{
org: org,
domains: map[string]struct{}{},
ips: map[string]struct{}{},
cidr24: map[string]int{},
}
scope.byASN[asn] = a
}
a.total++
a.expensive++
if httpASNCrawlAmplified(rec.URI) {
a.amplified++
}
if rec.Domain != "" {
a.domains[rec.Domain] = struct{}{}
}
if len(a.samples) < 5 {
a.samples = append(a.samples, truncate(rec.URI, 200))
}
if _, ok := a.ips[ip]; !ok {
maxIPs := cfg.Thresholds.HTTPASNCrawlMaxTrackedIPs
if maxIPs <= 0 {
maxIPs = config.DefaultHTTPASNCrawlMaxTrackedIPs
}
if len(a.ips) < maxIPs {
a.ips[ip] = struct{}{}
if g := asnCrawlGroupCIDR(ip); g != "" {
a.cidr24[g]++
}
} else {
a.ipsCapped = true
}
}
}
func (s *domlogStats) asnCrawlScopeFor(rec accessLogRecord) *asnCrawlScope {
if s.asnCrawl == nil {
s.asnCrawl = map[string]*asnCrawlScope{}
}
key := rec.Account
if key == "" {
key = "domain:" + rec.Domain
}
sc := s.asnCrawl[key]
if sc == nil {
sc = &asnCrawlScope{byASN: map[uint]*asnCrawlASN{}}
s.asnCrawl[key] = sc
}
return sc
}
func uintInSlice(v uint, list []uint) bool {
for _, x := range list {
if x == v {
return true
}
}
return false
}
// asnCrawlGroupCIDR returns the /24 (IPv4) or /64 (IPv6) the ip belongs to.
func asnCrawlGroupCIDR(ip string) string {
p := net.ParseIP(ip)
if p == nil {
return ""
}
if v4 := p.To4(); v4 != nil {
return net.IP(v4.Mask(net.CIDRMask(24, 32))).String() + "/24"
}
return p.Mask(net.CIDRMask(64, 128)).String() + "/64"
}
// asnCrawlWithinWindow gates a record to the detector's own lookback window
// (thresholds.http_asn_crawl_window_min). A zero/absent timestamp is excluded,
// and future timestamps beyond a small clock-skew allowance are ignored.
func asnCrawlWithinWindow(ts time.Time, cfg *config.Config, now time.Time) bool {
if ts.IsZero() || cfg == nil {
return false
}
win := cfg.Thresholds.HTTPASNCrawlWindowMin
if win <= 0 {
win = config.DefaultHTTPASNCrawlWindowMin
}
cutoff := now.Add(-time.Duration(win) * time.Minute)
return !ts.Before(cutoff) && !ts.After(now.Add(time.Minute))
}
// phpWorkersByUserFn is a seam for testing; production code uses phpWorkersByUser.
var phpWorkersByUserFn = phpWorkersByUser
// asnCrawlSaturated reports whether account's live lsphp worker count meets or
// exceeds the saturation threshold. Returns false for empty account (domain-scoped
// findings never escalate) or when no positive threshold is configured.
func asnCrawlSaturated(account string, workers map[string][]string, cfg *config.Config) bool {
if account == "" {
return false
}
threshold := cfg.Thresholds.HTTPASNCrawlSaturation
if threshold <= 0 {
threshold = cfg.Performance.PHPProcessWarnPerUser
}
if threshold <= 0 {
return false
}
return len(workers[account]) >= threshold
}
// emitASNCrawl produces Warning/High findings for each (scope, ASN) pair that
// passes all stage-1 gates. When an accounted scope is PHP-pool saturated the
// severity is escalated to Critical (stage 2).
func (s *domlogStats) emitASNCrawl(cfg *config.Config) []alert.Finding {
if cfg == nil || cfg.Thresholds.HTTPASNCrawlMinIPs <= 0 || s.asnCrawl == nil {
return nil
}
th := cfg.Thresholds
minIPs := th.HTTPASNCrawlMinIPs
minExpensive := th.HTTPASNCrawlMinExpensive
if minExpensive <= 0 {
minExpensive = config.DefaultHTTPASNCrawlMinExpensive
}
minSharePct := th.HTTPASNCrawlMinSharePct
if minSharePct <= 0 {
minSharePct = config.DefaultHTTPASNCrawlMinSharePct
}
highAmpPct := th.HTTPASNCrawlHighAmpPct
if highAmpPct <= 0 {
highAmpPct = config.DefaultHTTPASNCrawlHighAmpPct
}
highVolMult := th.HTTPASNCrawlHighVolumeMult
if highVolMult <= 0 {
highVolMult = config.DefaultHTTPASNCrawlHighVolMult
}
var out []alert.Finding
var workers map[string][]string
workersLoaded := false
for scopeKey, scope := range s.asnCrawl {
for asn, a := range scope.byASN {
if uintInSlice(asn, th.HTTPASNCrawlAllowlistASNs) ||
uintInSlice(asn, th.HTTPASNCrawlReverseProxyASNs) {
continue
}
if a.distinctIPs() < minIPs {
continue
}
if a.expensive < minExpensive {
continue
}
if scope.scopeExpensive == 0 ||
a.expensive*100 < minSharePct*scope.scopeExpensive {
continue
}
sev := alert.Warning
ampHigh := a.expensive > 0 && a.amplified*100 >= highAmpPct*a.expensive
volHigh := a.expensive >= highVolMult*minExpensive
if ampHigh || volHigh {
sev = alert.High
}
account, domain := asnCrawlScopeParts(scopeKey, a)
if account != "" {
if !workersLoaded {
workers = phpWorkersByUserFn()
workersLoaded = true
}
if asnCrawlSaturated(account, workers, cfg) {
sev = alert.Critical
}
}
cidrs := dropFirewallAllowedCIDRs(collapseASNCrawlCIDRs(a, cfg), a)
out = append(out, alert.Finding{
Severity: sev,
Check: "http_asn_crawl",
TenantID: account,
Domain: domain,
Message: fmt.Sprintf("Distributed crawl from AS%d (%s) against %s", asn, a.org, asnCrawlScopeLabel(scopeKey)),
Details: asnCrawlDetails(asn, a, scope, cfg, cidrs),
CIDRs: cidrs,
Timestamp: time.Now(),
})
}
}
sort.Slice(out, func(i, j int) bool { return out[i].Message < out[j].Message })
return out
}
// asnCrawlScopeParts returns (account, domain) for a finding. account is the
// scope key unless it is a domain-scoped fallback ("domain:<d>"); domain is set
// only when exactly one domain was observed.
func asnCrawlScopeParts(scopeKey string, a *asnCrawlASN) (account, domain string) {
if !strings.HasPrefix(scopeKey, "domain:") {
account = scopeKey
}
if len(a.domains) == 1 {
for d := range a.domains {
domain = d
}
}
return account, domain
}
func asnCrawlScopeLabel(scopeKey string) string {
return strings.TrimPrefix(scopeKey, "domain:")
}
func asnCrawlDetails(asn uint, a *asnCrawlASN, scope *asnCrawlScope, cfg *config.Config, cidrs []string) string {
maxTracked := cfg.Thresholds.HTTPASNCrawlMaxTrackedIPs
if maxTracked <= 0 {
maxTracked = config.DefaultHTTPASNCrawlMaxTrackedIPs
}
ipCount := fmt.Sprintf("%d", a.distinctIPs())
if a.ipsCapped {
ipCount = fmt.Sprintf(">=%d", maxTracked)
}
share := 0
if scope.scopeExpensive > 0 {
share = a.expensive * 100 / scope.scopeExpensive
}
var b strings.Builder
fmt.Fprintf(&b, "ASN: AS%d (%s)\n", asn, a.org)
fmt.Fprintf(&b, "Distinct source IPs: %s\n", ipCount)
fmt.Fprintf(&b, "Expensive/total reqs: %d/%d (%d%% of scope expensive)\n", a.expensive, a.total, share)
fmt.Fprintf(&b, "Amplified (layered-nav/search) reqs: %d\n", a.amplified)
if len(a.samples) > 0 {
fmt.Fprintf(&b, "Sample URIs: %s\n", strings.Join(a.samples, " | "))
}
if len(cidrs) > 0 {
fmt.Fprintf(&b, "Suggested subnets: %s\n", strings.Join(cidrs, ", "))
}
return b.String()
}
// collapseASNCrawlCIDRs reduces an ASN's observed IPs to a sorted, capped set
// of /24 (IPv4) / /64 (IPv6) subnets. When a single /16 covers >=
// http_asn_crawl_16_pref_pct of the IPv4 IPs across >= 4 distinct /24s, that
// /16 replaces its member /24s. Remaining /24s and any /64s are kept as-is.
// The result is sorted and capped to http_asn_crawl_max_prefix entries.
func collapseASNCrawlCIDRs(a *asnCrawlASN, cfg *config.Config) []string {
if len(a.cidr24) == 0 {
return nil
}
th := cfg.Thresholds
pct16 := th.HTTPASNCrawl16PrefPct
if pct16 <= 0 {
pct16 = config.DefaultHTTPASNCrawl16PrefPct
}
// Count total IPv4 IPs and group /24s by their /16 parent.
totalV4 := 0
bySixteen := map[string][]string{} // "x.y.0.0/16" -> member /24 CIDRs
for cidr, n := range a.cidr24 {
if strings.HasSuffix(cidr, "/24") {
totalV4 += n
sl := strings.SplitN(cidr, ".", 4)
s16 := sl[0] + "." + sl[1] + ".0.0/16"
bySixteen[s16] = append(bySixteen[s16], cidr)
}
}
chosen := map[string]struct{}{}
covered := map[string]struct{}{} // /24s replaced by their /16
// Promote a /16 when >= 4 member /24s cover >= pct16% of all IPv4 IPs.
for s16, members := range bySixteen {
if len(members) < 4 {
continue
}
count := 0
for _, m := range members {
count += a.cidr24[m]
}
if totalV4 > 0 && count*100 >= pct16*totalV4 {
chosen[s16] = struct{}{}
for _, m := range members {
covered[m] = struct{}{}
}
}
}
// Add uncovered /24s and all /64s.
for cidr := range a.cidr24 {
if _, done := covered[cidr]; done {
continue
}
chosen[cidr] = struct{}{}
}
out := make([]string, 0, len(chosen))
for c := range chosen {
out = append(out, c)
}
sort.Strings(out)
maxPrefix := th.HTTPASNCrawlMaxPrefix
if maxPrefix <= 0 {
maxPrefix = config.DefaultHTTPASNCrawlMaxPrefix
}
if len(out) > maxPrefix {
out = out[:maxPrefix]
}
return out
}
// dropFirewallAllowedCIDRs removes any candidate CIDR that contains an observed
// source IP already firewall-allowed (whitelisted), per spec section 6.2. This
// only ever REMOVES subnets, so it is strictly more conservative and can never
// cause over-blocking. When no blocker is wired, or it does not implement
// allowChecker, all CIDRs are kept. Each allowed observed IP is tested against
// the (<=8) collapsed CIDRs only.
func dropFirewallAllowedCIDRs(cidrs []string, a *asnCrawlASN) []string {
if len(cidrs) == 0 {
return cidrs
}
ac, ok := getIPBlocker().(allowChecker)
if !ok {
return cidrs
}
drop := map[string]struct{}{}
for ipStr := range a.ips {
if !ac.IsAllowed(ipStr) {
continue
}
parsed := net.ParseIP(ipStr)
if parsed == nil {
continue
}
for _, cidr := range cidrs {
if _, dropped := drop[cidr]; dropped {
continue
}
_, ipnet, err := net.ParseCIDR(cidr)
if err != nil {
continue
}
if ipnet.Contains(parsed) {
drop[cidr] = struct{}{}
}
}
}
if len(drop) == 0 {
return cidrs
}
kept := make([]string, 0, len(cidrs))
for _, cidr := range cidrs {
if _, dropped := drop[cidr]; !dropped {
kept = append(kept, cidr)
}
}
return kept
}
package checks
import (
"errors"
"fmt"
)
// MySQL identifier limit per the manual is 64 bytes. Anything beyond
// is silently truncated by the server, but for the cleaner we want
// the failure to be loud at the validation step.
const mysqlIdentMaxLen = 64
// errEmptyIdent and errInvalidIdent are surfaced to the operator so
// the CLI prints a clear "the name you typed is not a valid MySQL
// identifier" instead of a SQL syntax error from the server.
var (
errEmptyIdent = errors.New("identifier is empty")
errInvalidIdent = errors.New("identifier contains characters outside [A-Za-z0-9_$]")
errLongIdent = fmt.Errorf("identifier exceeds %d bytes (MySQL limit)", mysqlIdentMaxLen)
)
// QuoteIdent returns a backtick-quoted MySQL identifier, or an error
// if the input is empty, longer than 64 bytes, or contains characters
// outside the safe class. Used at every site where an attacker-
// controlled object name (trigger / event / routine / schema) would
// otherwise reach a SQL string concatenation.
//
// The safe class is intentionally narrow: standard MySQL allows more
// (digits-only names, dotted names, $-prefixed) but the cleaner only
// needs to handle CMS-shaped identifiers and operator-typed schema
// names. Rejecting anything weirder is cheaper than reasoning about
// edge cases in dynamic SQL.
func QuoteIdent(name string) (string, error) {
if name == "" {
return "", errEmptyIdent
}
if len(name) > mysqlIdentMaxLen {
return "", errLongIdent
}
for _, r := range name {
switch {
case r >= 'A' && r <= 'Z':
case r >= 'a' && r <= 'z':
case r >= '0' && r <= '9':
case r == '_' || r == '$':
default:
return "", fmt.Errorf("%w: %q", errInvalidIdent, name)
}
}
return "`" + name + "`", nil
}
package checks
import (
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
)
func init() {
// CLI scans and standalone web UI dispatches also need this policy,
// before any daemon is constructed. It depends only on check metadata
// and the config passed to each call, not on a daemon's lifetime.
alert.SetIPResponsePolicy(IPResponseAnswersFinding)
}
// IPResponseAnswersFinding is the alert.IPResponsePolicy for
// suppress_blocked_alerts. A block or challenge on the source address answers
// only attacker-side findings: attacker activity or attempted access that is
// not evidence of compromise. Once the source is stopped the operator has
// nothing left to do. Compromise evidence, successful logins and audit events
// stay visible even when their source address is already blocked.
//
// A challenge answers only the findings challenge routing would send to the
// gate. The address being on the challenge list is itself proof the gate is
// in use, so the policy does not consult challenge.enabled.
func IPResponseAnswersFinding(cfg *config.Config, f alert.Finding, blocked bool) bool {
if correlationReasonOf(f.Check) != reasonAttackerSide {
return false
}
switch f.Check {
case "email_phishing_content", "email_malware":
// Mail content may be evidence of a compromised local sender, and
// any IP in its message may come from an untrusted mail header.
return false
case "mail_account_spray", "smtp_account_spray", "mail_subnet_spray", "smtp_subnet_spray",
"http_distributed_flood", "http_asn_crawl":
// These summarize many sources. SourceIP can be the latest sender
// or a subnet; blocking one address does not answer the finding.
return false
}
if blocked {
return true
}
return challengeRoutesFinding(cfg, f)
}
package checks
import (
"context"
"time"
"github.com/pidginhost/csm/internal/jstaint"
)
// jsTaintAnalyze is indirected so tests can substitute the analyzer at the
// adapter boundary: forcing a panic or a specific status proves each owner's
// containment without crafting pathological JavaScript.
var jsTaintAnalyze = jstaint.Analyze
// jsTaintReverifyTimeout is the per-file safety net for a re-verification
// analysis. The analyzer checks its context between AST nodes, and its node,
// depth, and fact caps bound normal work far below this; the deadline only
// stops a defect from stalling an operator-triggered re-check.
const jsTaintReverifyTimeout = 15 * time.Second
// runJSTaintAnalysis is the adapter recovery boundary around the analyzer
// call. The analyzer converts its own internal panics to StatusPanic, but a
// panic in adapter-side code around the call must have the same contained
// outcome instead of escaping through the owning check.
func runJSTaintAnalysis(ctx context.Context, data []byte) (report jstaint.Report) {
defer func() {
if r := recover(); r != nil {
report = jstaint.Report{Status: jstaint.StatusPanic, Reason: "panic"}
}
}()
return jsTaintAnalyze(ctx, data)
}
package checks
import (
"context"
"fmt"
"sort"
"strings"
"time"
"unicode"
"unicode/utf8"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/jstaint"
)
// jsTaintDeepCursorCheck is the host-scope scan-cursor key for the scheduled
// JS taint consumer of the shared deep-content walk.
const jsTaintDeepCursorCheck = logicalOwnerJSTaintDeep
// jsTaintDeepPerFileTimeout is the per-file safety net for one deep-scan
// analysis. The engine's node, depth, and fact caps bound normal work far
// below this; the deadline only stops an analyzer defect from eating the
// walk's soft-deadline margin on a single file.
const jsTaintDeepPerFileTimeout = 20 * time.Second
// Display bounds: the sanitized path in the finding message, the rendered
// evidence in its details, and diagnostic example paths (spec: 512, 2048,
// and 256 bytes, markers included).
const (
jsTaintMessageMaxBytes = 512
jsTaintDetailsMaxBytes = 2048
jsTaintExampleMaxBytes = 256
)
// sanitizeJSTaintDisplay renders untrusted path or evidence text for
// operator-facing fields: valid UTF-8, control and invalid bytes become '?',
// truncation happens at a rune boundary and the marker counts into the cap.
func sanitizeJSTaintDisplay(s string, maxBytes int) string {
s = strings.ToValidUTF8(s, "?")
s = strings.Map(func(r rune) rune {
if unicode.IsControl(r) {
return '?'
}
return r
}, s)
if len(s) <= maxBytes {
return s
}
cut := maxBytes - 3
for cut > 0 && !utf8.RuneStart(s[cut]) {
cut--
}
return s[:cut] + "..."
}
// jsTaintGapCollector aggregates per-path JS coverage gaps for one deep run:
// exact paths feed the carry-forward, counts and one example per status feed
// the js_taint_scan_incomplete diagnostic. Non-analyzed statuses are never
// counted as clean files.
type jsTaintGapCollector struct {
paths map[string]struct{}
pathAliases map[string]struct{}
aliasesByPath map[string][]string
byStatus map[string]int
example map[string]string
recordCoverage func([]string)
resolveAliases func(string) ([]string, bool)
unknown int
unknownExample string
}
func newJSTaintGapCollector() *jsTaintGapCollector {
return &jsTaintGapCollector{
paths: map[string]struct{}{},
pathAliases: map[string]struct{}{},
aliasesByPath: map[string][]string{},
byStatus: map[string]int{},
example: map[string]string{},
}
}
func (g *jsTaintGapCollector) record(path, status string) {
aliases, retained := g.aliasesByPath[path]
if !retained {
stable := true
if g.resolveAliases != nil {
aliases, stable = g.resolveAliases(path)
} else {
aliases = []string{coverageLexicalPath(path)}
}
if !stable {
g.recordUnknownRange(fmt.Sprintf("%s changed while its path identity was captured", path))
g.byStatus[status]++
if _, ok := g.example[status]; !ok {
g.example[status] = sanitizeJSTaintDisplay(path, jsTaintExampleMaxBytes)
}
return
}
g.paths[path] = struct{}{}
g.aliasesByPath[path] = aliases
for _, alias := range aliases {
g.pathAliases[alias] = struct{}{}
}
}
if g.recordCoverage != nil {
g.recordCoverage(aliases)
}
g.byStatus[status]++
if _, ok := g.example[status]; !ok {
g.example[status] = sanitizeJSTaintDisplay(path, jsTaintExampleMaxBytes)
}
}
func (g *jsTaintGapCollector) recordUnknownRange(detail string) {
g.unknown++
if g.unknownExample == "" {
g.unknownExample = sanitizeJSTaintDisplay(detail, jsTaintExampleMaxBytes)
}
}
func (g *jsTaintGapCollector) pathsIncomplete() bool { return g.unknown > 0 }
func (g *jsTaintGapCollector) empty() bool { return len(g.byStatus) == 0 && g.unknown == 0 }
func (g *jsTaintGapCollector) hasPath(path string) bool {
if _, ok := g.paths[path]; ok {
return true
}
for _, alias := range coveragePathAliases(path) {
if _, ok := g.pathAliases[alias]; ok {
return true
}
}
return false
}
func (g *jsTaintGapCollector) finding() alert.Finding {
total := 0
statuses := make([]string, 0, len(g.byStatus))
for status, n := range g.byStatus {
total += n
statuses = append(statuses, status)
}
sort.Strings(statuses)
parts := make([]string, 0, len(statuses)+1)
for _, status := range statuses {
parts = append(parts, fmt.Sprintf("%s=%d (example: %s)", status, g.byStatus[status], g.example[status]))
}
if g.unknown > 0 {
parts = append(parts, fmt.Sprintf("unreadable-range=%d (example: %s)", g.unknown, g.unknownExample))
}
message := fmt.Sprintf("JavaScript taint deep scan could not analyze %d file(s)", total)
if total == 0 {
message = fmt.Sprintf("JavaScript taint deep scan could not cover %d location(s)", g.unknown)
}
return alert.Finding{
Severity: alert.Warning,
Check: "js_taint_scan_incomplete",
Message: message,
Details: strings.Join(parts, "; "),
// One host-wide condition: counts and examples vary per cycle and
// must not re-alert while coverage stays degraded.
DedupKey: "coverage_gap",
}
}
// carryForwardJSTaintFindings keeps at most one prior state finding for each
// path the current full cycle could not analyze. Partial rolling windows can
// leave more than one historical variant at a path because their owner is not
// purgeable; the most recently observed variant is the one that represents
// current state when a later full cycle hits a path-specific gap.
func carryForwardJSTaintFindings(prior []alert.Finding, gaps *jsTaintGapCollector) []alert.Finding {
byPath := make(map[string]alert.Finding)
for _, finding := range prior {
if finding.Check != "js_keylogger_dataflow" || !gaps.hasPath(finding.FilePath) {
continue
}
current, exists := byPath[finding.FilePath]
if !exists || finding.Timestamp.After(current.Timestamp) ||
(finding.Timestamp.Equal(current.Timestamp) && finding.Key() < current.Key()) {
byPath[finding.FilePath] = finding
}
}
paths := make([]string, 0, len(byPath))
for path := range byPath {
paths = append(paths, path)
}
sort.Strings(paths)
carried := make([]alert.Finding, 0, len(paths))
for _, path := range paths {
finding := byPath[path]
finding.ScanCarryForward = true
carried = append(carried, finding)
}
return carried
}
// analyzeJSTaintSnapshot runs the JS consumer on one complete in-memory
// snapshot and converts the result into at most one finding. A non-completed
// status is recorded as a known-path coverage gap, never as a clean file.
func analyzeJSTaintSnapshot(ctx context.Context, path, contentSHA256 string, data []byte, gaps *jsTaintGapCollector) []alert.Finding {
fileCtx, cancel := context.WithTimeout(ctx, jsTaintDeepPerFileTimeout)
report := runJSTaintAnalysis(fileCtx, data)
cancel()
switch report.Status {
case jstaint.StatusAnalyzed:
if len(report.Results) == 0 {
return nil
}
return []alert.Finding{jsTaintDeepFinding(path, contentSHA256, report)}
case jstaint.StatusNotCandidate:
return nil
default:
gaps.record(path, report.Status.String())
return nil
}
}
// jsTaintDeepFinding renders the single finding for one analyzed file: the
// evidence flows the engine returned plus the count of endpoint flows beyond
// them, with every display field sanitized and bounded. The fingerprint is
// the hash of the exact analyzed bytes; FilePath keeps the exact live path
// for remediation while only its display copy is sanitized.
func jsTaintDeepFinding(path, contentSHA256 string, report jstaint.Report) alert.Finding {
flows := make([]string, 0, len(report.Results))
for _, res := range report.Results {
segs := make([]string, 0, len(res.Via)+2)
segs = append(segs, res.Source)
segs = append(segs, res.Via...)
segs = append(segs, res.Sink)
flows = append(flows, strings.Join(segs, " -> "))
}
details := "Keystroke data reaches a network sink. Evidence: " + strings.Join(flows, "; ")
if extra := report.TotalResults - len(report.Results); extra > 0 {
details += fmt.Sprintf("; %d additional flow(s) beyond returned evidence", extra)
}
if report.EvidenceTruncated {
details += " [evidence truncated]"
}
return alert.Finding{
Severity: alert.Critical,
Check: "js_keylogger_dataflow",
Message: "JavaScript keystroke exfiltration data flow: " + sanitizeJSTaintDisplay(path, jsTaintMessageMaxBytes),
Details: sanitizeJSTaintDisplay(details, jsTaintDetailsMaxBytes),
FilePath: path,
ContentSHA256: contentSHA256,
DetectLogic: ContentDetectionVersion(),
}
}
package checks
import (
"math"
"os"
"path/filepath"
"strconv"
"strings"
"time"
)
// procClockTicks is USER_HZ, 100 on the platforms CSM runs on.
const procClockTicks = 100.0
// processStartedBefore reports whether pid's process began before t.
//
// A PID is recycled freely on a busy host, and a finding can be acted on long
// after it was raised -- immediately by auto-response, or whenever an operator
// clicks Fix. Killing by number alone therefore risks destroying a process that
// merely inherited the PID. A process that started after the finding cannot be
// the one the finding describes.
//
// Unverifiable input fails closed: no uptime, no stat, or an unparsable
// field all report false, and the caller does not kill.
//
// (internal/daemon carries its own copy of this parse for the af_alg reaction:
// it reads procfs directly rather than through this package's filesystem seam,
// and sharing one implementation would mean mixing two abstractions.)
func processStartedBefore(pid string, t time.Time) bool {
pidInt, ok := parseProcessPID(pid)
if !ok || t.IsZero() {
return false
}
ticks, ok := procStartTicks(strconv.Itoa(pidInt))
if !ok {
return false
}
uptime, ok := procUptime()
if !ok {
return false
}
// Read wall time after uptime so elapsed is conservatively rounded up. A
// process close enough to the boundary to be ambiguous is not killed.
elapsed := time.Since(t).Seconds()
if elapsed < 0 {
return false
}
eventUptime := uptime - elapsed
return eventUptime >= 0 && float64(ticks)/procClockTicks <= eventUptime
}
func parseProcessPID(pid string) (int, bool) {
n, err := strconv.Atoi(pid)
return n, err == nil && n > 1 && n <= math.MaxInt32
}
func procUptime() (float64, bool) {
data, err := osFS.ReadFile("/proc/uptime")
if err != nil {
return 0, false
}
fields := strings.Fields(string(data))
if len(fields) == 0 {
return 0, false
}
uptime, err := strconv.ParseFloat(fields[0], 64)
return uptime, err == nil && !math.IsNaN(uptime) && !math.IsInf(uptime, 0) && uptime >= 0
}
func procStartTicks(pid string) (uint64, bool) {
data, err := osFS.ReadFile(filepath.Join("/proc", pid, "stat"))
if err != nil {
return 0, false
}
// The comm field is parenthesised and may contain spaces, so fields are
// counted from after its closing parenthesis.
closing := strings.LastIndex(string(data), ")")
if closing < 0 {
return 0, false
}
fields := strings.Fields(string(data)[closing+1:])
// starttime is field 22 overall, index 19 once pid and comm are behind us.
if len(fields) < 20 {
return 0, false
}
ticks, err := strconv.ParseUint(fields[19], 10, 64)
if err != nil {
return 0, false
}
return ticks, true
}
// processUsesFile reports whether pid still references path, as its executable
// or through an open descriptor. The kill in fixKillAndQuarantine exists to
// release the file being quarantined; a process that no longer references that
// file is not the process the finding meant, so it is not killed.
func processUsesFile(pid, path string) bool {
pidInt, ok := parseProcessPID(pid)
if !ok || path == "" {
return false
}
target, err := osFS.Lstat(path)
if err != nil || target.Mode()&os.ModeSymlink != 0 {
return false
}
return processUsesFileIdentity(pidInt, target)
}
func processUsesFileIdentity(pid int, target os.FileInfo) bool {
if pid <= 1 || target == nil || target.Mode()&os.ModeSymlink != 0 {
return false
}
procDir := filepath.Join("/proc", strconv.Itoa(pid))
if info, statErr := osFS.Stat(filepath.Join(procDir, "exe")); statErr == nil && sameObject(info, target) {
return true
}
fdDir := filepath.Join(procDir, "fd")
entries, err := osFS.ReadDir(fdDir)
if err != nil {
return false
}
for _, entry := range entries {
info, err := osFS.Stat(filepath.Join(fdDir, entry.Name()))
if err == nil && sameObject(info, target) {
return true
}
}
return false
}
// sameObject reports whether two stats describe the same file, requiring the
// content shape to agree as well as device and inode. A deleted inode is handed
// straight back to the next file created in the same directory, so dev+ino
// alone would call a descriptor for the removed file the same object as its
// replacement -- the same inode-reuse hole the quarantine move already guards
// against.
func sameObject(a, b os.FileInfo) bool {
return sameFileIdentity(a, b) && sameContentShape(a, b)
}
package checks
import (
"context"
"fmt"
"math"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/attackdb"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
const localThreatScoreFindingLimit = 50
// CheckLocalThreatScore generates findings for IPs that have accumulated
// a high local threat score but have not yet been blocked.
// Runs every 10 minutes as part of TierCritical.
func CheckLocalThreatScore(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
adb := attackdb.Global()
if adb == nil {
return nil
}
alreadyBlocked := loadAllBlockedIPs(cfg.StatePath)
var findings []alert.Finding
for _, rec := range adb.TopAttackers(math.MaxInt) {
if ctx.Err() != nil {
return findings
}
if alreadyBlocked[rec.IP] {
continue
}
// Local addresses can represent proxied attacks or compromised local
// processes. Firewall self-protection is not a reason to hide evidence.
score := attackdb.ComputeScore(rec)
if score >= 70 {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "local_threat_score",
Message: fmt.Sprintf("High local threat score: %s (score %d/100, %d attacks)", rec.IP, score, rec.EventCount),
Details: fmt.Sprintf("Attack types: %v\nAccounts targeted: %d\nFirst seen: %s\nLast seen: %s", rec.AttackCounts, len(rec.Accounts), rec.FirstSeen.Format("2006-01-02 15:04"), rec.LastSeen.Format("2006-01-02 15:04")),
Timestamp: time.Now(),
SourceIP: rec.IP,
})
if len(findings) >= localThreatScoreFindingLimit {
break
}
}
}
return findings
}
package checks
import (
"crypto/sha256"
"fmt"
)
// loginRecordKey retains session identity even when the syslog prefix is long
// enough to push the PID or client port past the display-details truncation.
func loginRecordKey(line string) string {
return fmt.Sprintf("%x", sha256.Sum256([]byte(line)))
}
package checks
import (
"context"
"errors"
"fmt"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/systemdrun"
)
// eximQueueTimeout bounds the queue probe. `exim -bpc` walks the spool, so it
// is slower on a deep queue but never long-running.
const eximQueueTimeout = 30 * time.Second
var errEximNotInstalled = errors.New("exim is not installed")
func CheckMailQueue(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
count, err := eximQueueCount(ctx)
if err != nil {
if errors.Is(err, errEximNotInstalled) {
return nil
}
// A queue depth CSM cannot read is a monitoring blind spot, not a
// clean bill of health. Reporting nothing here is what let a fully
// broken probe run unnoticed on a production host for months.
return []alert.Finding{{
Severity: alert.Warning,
Check: "mail_queue_unavailable",
Message: "Exim mail queue depth could not be read",
Details: fmt.Sprintf("Spam-outbreak detection via queue depth is inactive until this succeeds. %v",
err),
}}
}
if count >= cfg.Thresholds.MailQueueCrit {
return []alert.Finding{{
Severity: alert.Critical,
Check: "mail_queue",
Message: fmt.Sprintf("Exim mail queue critical: %d messages", count),
Details: "Possible spam outbreak from compromised account",
}}
}
if count >= cfg.Thresholds.MailQueueWarn {
return []alert.Finding{{
Severity: alert.Warning,
Check: "mail_queue",
Message: fmt.Sprintf("Exim mail queue elevated: %d messages", count),
}}
}
return nil
}
// eximQueueCount returns the number of messages in the Exim queue.
//
// exim opens its own main log for append on every invocation and aborts when it
// cannot. Under the daemon's ProtectSystem=strict sandbox /var/log is read-only,
// so a direct call fails with "Cannot open main log file" and yields no count at
// all. Running the probe as a transient unit hands it to PID 1, outside the
// sandbox. Hosts without systemd-run have no sandbox to escape and call exim
// directly.
func eximQueueCount(parent context.Context) (int, error) {
eximPath, err := cmdExec.LookPath("exim")
if err != nil {
return 0, fmt.Errorf("%w: %v", errEximNotInstalled, err)
}
ctx, cancel := context.WithTimeout(parent, eximQueueTimeout)
defer cancel()
out, err := systemdrun.Run(ctx, cmdExec.LookPath, cmdExec.RunContextStdout, systemdrun.Options{
Pipe: true,
RuntimeMax: eximQueueTimeout,
}, eximPath, "-bpc")
if err != nil {
return 0, fmt.Errorf("exim -bpc: %w", err)
}
return parseEximQueueCount(out)
}
func parseEximQueueCount(out []byte) (int, error) {
trimmed := strings.TrimSpace(string(out))
count, err := strconv.Atoi(trimmed)
if err != nil {
return 0, fmt.Errorf("exim -bpc returned %q, not a queue count", truncateForDetail(trimmed))
}
return count, nil
}
// truncateForDetail keeps an unexpected command output short enough to sit in a
// finding's Details without pasting a whole error page into the alert.
func truncateForDetail(s string) string {
const max = 120
if len(s) <= max {
return s
}
return s[:max] + "..."
}
package checks
import (
"context"
"crypto/sha256"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// Mail-filter exfiltration detector.
//
// cPanel stores per-mailbox Exim filters at
// /home/<user>/etc/<domain>/<localpart>/filter and domain-wide defaults at
// /etc/vfilters/<domain>. A compromised webmail account is commonly weaponised
// by writing a filter that copies every inbound message to an external dropbox
// while keeping a local copy, so the victim never notices the interception
// (business email compromise). This check parses those Exim filters and scores
// the deliver/save actions for that stealth pattern.
//
// Unlike CheckForwarders (valiases redirects), interception-shaped rules are
// reported even when the filter predates CSM. Newness gating only applies to
// plain external forwards that are frequently legitimate customer
// configuration; severity for copy-forwards is assigned after corroboration.
// filterAction is a single Exim filter action (deliver/save/pipe/finish/...).
type filterAction struct {
verb string
arg string
unseen bool
knownSuppressible bool
matchesAll bool
}
// filterRule is one branch of an Exim filter: the condition that guards it and
// the actions it performs. The unconditional top level is represented as a rule
// with an empty condition.
type filterRule struct {
condition string
matchesAll bool
actions []filterAction
}
// filterMailbox identifies the mailbox a filter file belongs to. localPart is
// "*" for a domain-wide /etc/vfilters file.
type filterMailbox struct {
localPart string
domain string
}
func (m filterMailbox) String() string {
return m.localPart + "@" + m.domain
}
// filterFinding is the scorer's intermediate result before it is turned into an
// alert.Finding (which needs file path and newness context).
type filterFinding struct {
severity alert.Severity
check string
kind string // "exfil" | "forwarder" | "pipe" | "blackhole"
dest string // external destination, when applicable, for correlation
reason string
onlyIfNew bool
// retainsLocalCopy marks an external delivery the mailbox still receives a
// copy of. Alone that is indistinguishable from a user-created forward.
retainsLocalCopy bool
}
// safePipeCommands are cPanel built-in pipe targets that are not attacker code.
// cPanel writes Mailman list aliases as pipes to the 3rdparty mail binaries.
var safePipeCommands = []string{
"/usr/local/cpanel/bin/autorespond",
"/usr/local/cpanel/bin/boxtrapper",
"/usr/local/cpanel/bin/mailman",
"/usr/local/cpanel/3rdparty/mailman/mail/mailman",
"/usr/local/cpanel/3rdparty/mailman/mail/wrapper",
}
// ---------------------------------------------------------------------------
// Parser
// ---------------------------------------------------------------------------
type eximToken struct {
text string
str bool
}
// tokenizeExim splits Exim filter source into tokens. Quoted strings (with \"
// and \\ escapes) become a single string token with the quotes removed;
// parentheses are standalone tokens; everything else is a bareword. Comments
// (# to end of line, outside strings) are dropped.
func tokenizeExim(s string) []eximToken {
var toks []eximToken
runes := []rune(s)
i := 0
for i < len(runes) {
c := runes[i]
switch c {
case ' ', '\t', '\n', '\r':
i++
case '#':
for i < len(runes) && runes[i] != '\n' {
i++
}
case '(', ')':
toks = append(toks, eximToken{text: string(c)})
i++
case '"':
i++
var b strings.Builder
for i < len(runes) && runes[i] != '"' {
if runes[i] == '\\' && i+1 < len(runes) {
i++
}
b.WriteRune(runes[i])
i++
}
if i < len(runes) {
i++ // closing quote
}
toks = append(toks, eximToken{text: b.String(), str: true})
default:
start := i
for i < len(runes) {
r := runes[i]
if r == ' ' || r == '\t' || r == '\n' || r == '\r' || r == '(' || r == ')' || r == '"' || r == '#' {
break
}
i++
}
toks = append(toks, eximToken{text: string(runes[start:i])})
}
}
return toks
}
// renderCondition reconstructs a condition string from its tokens so the
// match-all heuristic can run against text that matches the source form
// (string tokens are re-quoted).
func renderCondition(toks []eximToken) string {
var parts []string
for _, t := range toks {
if t.str {
parts = append(parts, `"`+t.text+`"`)
} else {
parts = append(parts, t.text)
}
}
return strings.Join(parts, " ")
}
var actionVerbs = map[string]bool{
"deliver": true,
"save": true,
"pipe": true,
"finish": true,
"mail": true,
"vacation": true,
}
var actionArgs = map[string]bool{
"deliver": true,
"save": true,
"pipe": true,
"mail": true,
"vacation": true,
}
var controlWords = map[string]bool{
"if": true,
"elif": true,
"else": true,
"endif": true,
"then": true,
"unseen": true,
}
type filterRuleNode struct {
rule filterRule
parent *filterRuleNode
}
// parseEximFilter parses Exim filter source into a flat list of rules, one per
// if/elif/else branch plus one for any unconditional top-level actions. Nested
// branches include ancestor actions so split deliver/save patterns still score
// as one executed branch.
func parseEximFilter(content string) []filterRule {
toks := tokenizeExim(content)
top := &filterRuleNode{rule: filterRule{matchesAll: true}}
stack := []*filterRuleNode{top}
rules := []*filterRuleNode{top}
pendingUnseen := false
i := 0
for i < len(toks) {
t := toks[i]
if t.str {
i++
continue
}
kw := strings.ToLower(t.text)
switch kw {
case "if", "elif":
if kw == "elif" && len(stack) > 1 {
stack = stack[:len(stack)-1]
}
parent := stack[len(stack)-1]
i++
condStart := i
for i < len(toks) && !tokenIs(toks[i], "then") {
i++
}
cond := renderCondition(toks[condStart:i])
if i < len(toks) {
i++ // consume "then"
}
r := &filterRuleNode{
rule: filterRule{condition: cond, matchesAll: conditionMatchesAll(cond)},
parent: parent,
}
rules = append(rules, r)
stack = append(stack, r)
pendingUnseen = false
case "else":
if len(stack) > 1 {
stack = stack[:len(stack)-1]
}
parent := stack[len(stack)-1]
i++
r := &filterRuleNode{
rule: filterRule{condition: "else"},
parent: parent,
}
rules = append(rules, r)
stack = append(stack, r)
pendingUnseen = false
case "endif":
if len(stack) > 1 {
stack = stack[:len(stack)-1]
}
i++
pendingUnseen = false
case "unseen":
pendingUnseen = true
i++
default:
if !actionVerbs[kw] {
i++
continue
}
verb := kw
i++
arg := ""
if actionArgs[verb] && i < len(toks) {
if toks[i].str || isBareActionArg(toks[i]) {
arg = toks[i].text
i++
}
}
cur := stack[len(stack)-1]
cur.rule.actions = append(cur.rule.actions, filterAction{verb: verb, arg: arg, unseen: pendingUnseen})
pendingUnseen = false
}
}
out := make([]filterRule, 0, len(rules))
for _, r := range rules {
if len(r.rule.actions) > 0 {
out = append(out, flattenRuleNode(r))
}
}
return out
}
func tokenIs(t eximToken, word string) bool {
return !t.str && strings.EqualFold(t.text, word)
}
func isBareActionArg(t eximToken) bool {
if t.str {
return true
}
lower := strings.ToLower(t.text)
return !controlWords[lower] && !actionVerbs[lower]
}
func flattenRuleNode(node *filterRuleNode) filterRule {
var chain []*filterRuleNode
for n := node; n != nil; n = n.parent {
chain = append(chain, n)
}
out := filterRule{matchesAll: true}
var conditions []string
for i := len(chain) - 1; i >= 0; i-- {
r := chain[i].rule
if r.condition != "" {
conditions = append(conditions, r.condition)
}
if !r.matchesAll {
out.matchesAll = false
}
out.actions = append(out.actions, r.actions...)
}
out.condition = strings.Join(conditions, " && ")
return out
}
// conditionMatchesAll reports whether a filter condition fires on effectively
// all mail: an unconditional rule, or one that only tests that an address or
// header comparison is true for every normal email address.
func conditionMatchesAll(cond string) bool {
c := strings.TrimSpace(cond)
if c == "" {
return true
}
return tokenExpressionMatchesAll(tokenizeExim(c))
}
func tokenExpressionMatchesAll(toks []eximToken) bool {
toks = trimOuterParens(toks)
if len(toks) == 0 {
return false
}
for _, term := range splitTopLevel(toks, "or") {
if tokenConjunctionMatchesAll(term) {
return true
}
}
return false
}
func tokenConjunctionMatchesAll(toks []eximToken) bool {
toks = trimOuterParens(toks)
if len(toks) == 0 {
return false
}
parts := splitTopLevel(toks, "and")
for _, part := range parts {
if !tokenTermMatchesAll(part) {
return false
}
}
return len(parts) > 0
}
func tokenTermMatchesAll(toks []eximToken) bool {
toks = trimOuterParens(toks)
if len(toks) == 0 {
return false
}
if parts := splitTopLevel(toks, "or"); len(parts) > 1 {
return tokenExpressionMatchesAll(toks)
}
if parts := splitTopLevel(toks, "and"); len(parts) > 1 {
return tokenConjunctionMatchesAll(toks)
}
hasMatchAllComparison := false
for i := 0; i < len(toks); i++ {
if toks[i].str {
continue
}
if strings.EqualFold(toks[i].text, "not") {
return false
}
}
for i := 0; i+1 < len(toks); i++ {
if toks[i].str {
continue
}
op := strings.ToLower(toks[i].text)
if !isAddressComparisonOperator(op) {
continue
}
if !comparisonMatchesAllAddress(op, toks[i+1].text) || !comparisonHasAddressOperand(toks, i) {
return false
}
hasMatchAllComparison = true
}
return hasMatchAllComparison
}
func splitTopLevel(toks []eximToken, word string) [][]eximToken {
var parts [][]eximToken
start := 0
depth := 0
for i, t := range toks {
if t.str {
continue
}
switch t.text {
case "(":
depth++
case ")":
if depth > 0 {
depth--
}
default:
if depth == 0 && strings.EqualFold(t.text, word) {
if start < i {
parts = append(parts, toks[start:i])
}
start = i + 1
}
}
}
if start < len(toks) {
parts = append(parts, toks[start:])
}
return parts
}
func trimOuterParens(toks []eximToken) []eximToken {
for len(toks) >= 2 && tokenIs(toks[0], "(") && tokenIs(toks[len(toks)-1], ")") && outerParensEncloseAll(toks) {
toks = toks[1 : len(toks)-1]
}
return toks
}
func outerParensEncloseAll(toks []eximToken) bool {
depth := 0
for i, t := range toks {
if t.str {
continue
}
switch t.text {
case "(":
depth++
case ")":
depth--
if depth == 0 && i != len(toks)-1 {
return false
}
if depth < 0 {
return false
}
}
}
return depth == 0
}
func isAddressComparisonOperator(op string) bool {
switch op {
case "contains", "matches", "is":
return true
}
return false
}
func comparisonMatchesAllAddress(op, value string) bool {
v := strings.ToLower(strings.TrimSpace(value))
switch op {
case "contains":
return v == "@"
case "matches":
switch v {
case "@", ".*@.*", ".+@.+", "^.*@.*$", "^.+@.+$":
return true
}
case "is":
return v == "*@*"
}
return false
}
func comparisonHasAddressOperand(toks []eximToken, opIndex int) bool {
for i := opIndex - 1; i >= 0; i-- {
if toks[i].str {
continue
}
word := strings.ToLower(toks[i].text)
switch word {
case "and", "or", "then", "else":
return false
case "(", ")":
return false
case "not":
continue
}
if tokenLooksAddressOperand(word) {
return true
}
}
return false
}
func tokenLooksAddressOperand(token string) bool {
addressOperands := []string{
"$thisaddress",
"foranyaddress",
"$sender_address",
"$return_path",
"$header_from",
"$h_from",
"$header_to",
"$h_to",
"$header_cc",
"$h_cc",
"$header_bcc",
"$h_bcc",
"$header_reply-to",
"$h_reply-to",
"$header_sender",
"$h_sender",
"$header_return-path",
"$h_return-path",
}
for _, operand := range addressOperands {
if strings.Contains(token, operand) {
return true
}
}
return false
}
// ---------------------------------------------------------------------------
// Scorer
// ---------------------------------------------------------------------------
// destIsExternal reports whether an Exim deliver destination leaves the local
// mail system. Exim variables ($domain etc.) and same-domain/local-domain
// addresses are not external.
func destIsExternal(dest string, mb filterMailbox, localDomains map[string]bool) bool {
_, dom, ok := splitDeliverDest(dest)
if !ok {
return false
}
if deliverDomainIsLocal(dom, mb, localDomains) {
return false
}
return true
}
// destIsLocalSelf reports whether a deliver destination routes back into the
// local mail system (a self re-delivery that keeps a copy for the victim).
func destIsLocalSelf(dest string, mb filterMailbox, localDomains map[string]bool) bool {
_, dom, ok := splitDeliverDest(dest)
if !ok {
return false
}
return deliverDomainIsLocal(dom, mb, localDomains)
}
func splitDeliverDest(dest string) (string, string, bool) {
clean := strings.Trim(strings.TrimSpace(dest), `"`)
at := strings.LastIndexByte(clean, '@')
if at < 0 || at == len(clean)-1 {
return "", "", false
}
local := strings.Trim(strings.TrimSpace(clean[:at]), `"`)
dom := strings.ToLower(strings.Trim(strings.TrimSpace(clean[at+1:]), `"`))
if local == "" || dom == "" {
return "", "", false
}
return local, dom, true
}
func deliverDomainIsLocal(dom string, mb filterMailbox, localDomains map[string]bool) bool {
d := strings.ToLower(strings.TrimSpace(dom))
if d == "$domain" || d == "${domain}" {
return true
}
if d == strings.ToLower(mb.domain) {
return true
}
return localDomains[d]
}
func isSafePipe(cmd string) bool {
first := firstPipeCommandWord(cmd)
for _, s := range safePipeCommands {
if first == s {
return true
}
}
return false
}
func firstPipeCommandWord(cmd string) string {
// Exim's transport_set_up_command uses byte whitespace and only treats
// a quote at the start of an argument specially. Shell-style quote
// concatenation or Unicode trimming can turn a different path into a
// trusted executable here.
const whitespace = " \t\r\n\v\f"
s := strings.TrimLeft(strings.TrimPrefix(strings.TrimLeft(cmd, whitespace), "|"), whitespace)
if s == "" || strings.IndexByte(s, 0) >= 0 {
return ""
}
quote := s[0]
if quote != '\'' && quote != '"' {
if end := strings.IndexAny(s, whitespace); end >= 0 {
return s[:end]
}
return s
}
var b strings.Builder
for i := 1; i < len(s); i++ {
c := s[i]
if c == quote {
return b.String()
}
if quote == '"' && c == '\\' {
if i+1 == len(s) {
return ""
}
var consumed int
c, consumed = pipeCommandEscape(s[i+1:])
i += consumed
if c == 0 {
return ""
}
}
b.WriteByte(c)
}
// Incomplete quoting is not evidence of a trusted command.
return ""
}
// pipeCommandEscape follows Exim's string_interpret_escape: up to three
// octal digits, up to two hex digits, C control escapes, or a literal byte.
// s begins immediately after the backslash and is nonempty.
func pipeCommandEscape(s string) (byte, int) {
switch s[0] {
case 'b':
return '\b', 1
case 'f':
return '\f', 1
case 'n':
return '\n', 1
case 'r':
return '\r', 1
case 't':
return '\t', 1
case 'v':
return '\v', 1
}
base := byte(8)
start := 0
if s[0] == 'x' {
base, start = 16, 1
} else if s[0] < '0' || s[0] > '7' {
return s[0], 1
}
var value byte
i := start
for ; i < len(s) && i < 3; i++ {
c := s[i] | 0x20
var digit byte
switch {
case c >= '0' && c <= '9':
digit = c - '0'
case c >= 'a' && c <= 'f':
digit = c - 'a' + 10
default:
return value, i
}
if digit >= base {
break
}
value = value*base + digit
}
return value, i
}
// scoreFilterRules evaluates one mailbox's parsed filter rules and returns the
// dangerous patterns found. Suppression entries in known (format
// "local@domain: dest") drop matching expected plain destinations and Sieve
// :copy redirects. Other stealth patterns remain non-suppressible.
func scoreFilterRules(rules []filterRule, mb filterMailbox, localDomains map[string]bool, known []string) []filterFinding {
var out []filterFinding
seen := map[string]bool{}
add := func(f filterFinding) {
key := f.kind + "|" + f.dest
if seen[key] {
// Multiple rules in one file can target the same destination. If
// any occurrence prevents local delivery, keep that stronger
// evidence instead of letting an earlier copy-forward hide it.
if f.kind == "exfil" && !f.retainsLocalCopy {
for i := range out {
if out[i].kind == f.kind && out[i].dest == f.dest {
out[i] = f
break
}
}
}
return
}
if f.kind == "forwarder" && seen["exfil|"+f.dest] {
return
}
if f.kind == "exfil" && seen["forwarder|"+f.dest] {
delete(seen, "forwarder|"+f.dest)
for i := range out {
if out[i].kind == "forwarder" && out[i].dest == f.dest {
out = append(out[:i], out[i+1:]...)
break
}
}
}
seen[key] = true
out = append(out, f)
}
type externalDelivery struct {
dest string
unseen bool
knownSuppressible bool
matchesAll bool
}
for _, r := range rules {
var external []externalDelivery
hasLocalCopy := false
hasDevNull := false
hasMatchAllDevNull := false
for _, a := range r.actions {
switch a.verb {
case "deliver":
switch {
case destIsExternal(a.arg, mb, localDomains):
external = append(external, externalDelivery{
dest: a.arg,
unseen: a.unseen,
knownSuppressible: a.knownSuppressible,
matchesAll: a.matchesAll,
})
case destIsLocalSelf(a.arg, mb, localDomains):
hasLocalCopy = true
}
case "save":
if strings.TrimSpace(a.arg) == "/dev/null" {
hasDevNull = true
hasMatchAllDevNull = hasMatchAllDevNull || a.matchesAll
} else {
hasLocalCopy = true
}
case "pipe":
if !isSafePipe(a.arg) {
add(filterFinding{
severity: alert.Critical,
check: "email_filter_pipe",
kind: "pipe",
dest: a.arg,
reason: fmt.Sprintf("filter pipes mail to a command: %s", a.arg),
})
}
}
}
if len(external) > 0 {
for _, delivery := range external {
matchesAll := r.matchesAll || delivery.matchesAll
stealth := hasLocalCopy || hasDevNull || matchesAll || delivery.unseen
// An explicit local delivery still happens alongside /dev/null.
// Only an implicit keep created by unseen/:copy is canceled by a
// destructive action in the same rule.
retainsLocalCopy := hasLocalCopy || (delivery.unseen && !hasDevNull)
knownAllowed := !stealth || (delivery.knownSuppressible && !hasDevNull)
if knownAllowed && IsKnownForwarder(mb.localPart, mb.domain, delivery.dest, known) {
continue
}
if stealth {
add(filterFinding{
severity: alert.Critical,
check: "email_filter_exfil",
kind: "exfil",
dest: delivery.dest,
reason: stealthReason(retainsLocalCopy, hasDevNull, matchesAll),
retainsLocalCopy: retainsLocalCopy,
})
} else {
add(filterFinding{
severity: alert.High,
check: "email_filter_forwarder",
kind: "forwarder",
dest: delivery.dest,
reason: "filter forwards mail to an external address",
onlyIfNew: true,
})
}
}
continue
}
if hasDevNull && (r.matchesAll || hasMatchAllDevNull) {
add(filterFinding{
severity: alert.High,
check: "email_filter_blackhole",
kind: "blackhole",
reason: "filter discards all mail to /dev/null",
})
}
}
return out
}
func stealthReason(retainsLocalCopy, devNull, matchAll bool) string {
switch {
case retainsLocalCopy:
if !matchAll {
return "filter sends matching mail to an external address while keeping a local copy"
}
return "filter copies every message to an external address while keeping a local copy"
case devNull:
if matchAll {
return "filter forwards all mail externally and discards the local copy to hide it"
}
return "filter forwards mail externally and discards the local copy to hide it"
case matchAll:
return "filter forwards all mail to an external address"
}
return "filter forwards mail to an external address"
}
// ---------------------------------------------------------------------------
// Check
// ---------------------------------------------------------------------------
// mailboxFromFilterPath derives the mailbox from a filter file path. Per-mailbox
// filters live at /home/<user>/etc/<domain>/<localpart>/filter; domain-wide
// filters at /etc/vfilters/<domain>.
func mailboxFromFilterPath(path string) filterMailbox {
if dir := filepath.Dir(path); dir == "/etc/vfilters" {
return filterMailbox{localPart: "*", domain: filepath.Base(path)}
}
parts := strings.Split(path, "/")
for i := 0; i+2 < len(parts); i++ {
if parts[i] == "etc" && parts[i+1] != "" && parts[i+2] != "" {
return filterMailbox{localPart: parts[i+2], domain: parts[i+1]}
}
}
return filterMailbox{localPart: "*", domain: filepath.Base(filepath.Dir(path))}
}
// CheckMailFilters scans per-mailbox and domain-wide Exim filters for BEC-style
// exfiltration rules. Throttled to PasswordCheckIntervalMin, like the forwarder
// audit it complements.
func CheckMailFilters(ctx context.Context, cfg *config.Config, st *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
db := store.Global()
if db == nil {
if !shouldReportMailFilterStoreUnavailable(st, cfg) {
return nil
}
// Without the store the whole check is inoperative (hashes and the
// throttle live there). Say so instead of looking like a clean host.
return []alert.Finding{{
Severity: alert.Warning,
Check: "email_mail_filters",
Message: "Mail filter audit skipped: state store unavailable",
Timestamp: time.Now(),
}}
}
// A host upgraded from an Exim-only release already has the shared
// mail-filter throttle marker but no Sieve baseline. Run once immediately
// so the new source is inventoried instead of waiting a full audit interval.
// Explicit account scans also bypass the host-wide throttle; they write only
// their per-file hashes and never establish either shared baseline below.
if !ForceAll && AccountFromContext(ctx) == "" && db.GetMetaString("email:mailsieve_last_refresh") != "" {
if last := db.GetMetaString("email:mailfilter_last_refresh"); last != "" {
if ts, err := time.Parse(time.RFC3339, last); err == nil {
interval := time.Duration(cfg.EmailProtection.PasswordCheckIntervalMin) * time.Minute
if time.Since(ts) < interval {
return nil
}
}
}
}
if ctx.Err() != nil {
return nil
}
localDomains := loadLocalDomains()
var files []string
baselineComplete := true
if perMailbox, err := homeGlob(ctx, "etc", "*", "*", "filter"); err == nil {
files = append(files, perMailbox...)
} else {
baselineComplete = false
}
if AccountFromContext(ctx) == "" {
if vfilters, err := osFS.Glob("/etc/vfilters/*"); err == nil {
files = append(files, vfilters...)
} else {
baselineComplete = false
}
}
maxFiles := accountScanMaxFiles(ctx, cfg)
baselineComplete = baselineComplete && scanCoversAllFiles(files, maxFiles)
ranked := rankPathsByMtimeDesc(ctx, files, maxFiles)
if ctx.Err() != nil {
return nil
}
var collected []mailFilterPending
for _, path := range ranked {
if ctx.Err() != nil {
applyMailFilterCorroboration(collected)
return findingsFromPending(collected)
}
data, err := osFS.ReadFile(path)
if err != nil {
baselineComplete = false
continue
}
currentHash := sha256Hex(data)
isNew := forwarderFileIsNew(db, "email:mailfilter_last_refresh", "mailfilter:"+path, currentHash)
if err := db.SetForwarderHash("mailfilter:"+path, currentHash); err != nil {
baselineComplete = false
}
mb := mailboxFromFilterPath(path)
rules := parseEximFilter(string(data))
collected = append(collected, mailFilterPendings(path, mb, rules, localDomains, cfg.EmailProtection.KnownForwarders, isNew, mailFilterMechanismExim)...)
}
sievePending, sieveComplete := scanSieveMailFilters(ctx, db, cfg, localDomains, maxFiles)
collected = append(collected, sievePending...)
baselineComplete = baselineComplete && sieveComplete
applyMailFilterCorroboration(collected)
if ctx.Err() != nil {
return nil
}
// Only a full scan establishes the baseline / refreshes the throttle:
// an account-scoped scan hashes one account's files, and marking it
// complete would make the next full scan treat every other account's
// existing filters as newly created.
baselineExists := db.GetMetaString("email:mailfilter_last_refresh") != ""
if AccountFromContext(ctx) == "" && (sieveComplete || db.GetMetaString("email:mailsieve_last_refresh") != "") {
_ = db.SetMetaString("email:mailsieve_last_refresh", time.Now().Format(time.RFC3339))
}
if AccountFromContext(ctx) == "" && (baselineComplete || baselineExists) {
_ = db.SetMetaString("email:mailfilter_last_refresh", time.Now().Format(time.RFC3339))
}
return findingsFromPending(collected)
}
// mailFilterPendings scores one parsed filter/sieve file and turns each
// dangerous pattern into a pending finding. Shared by the Exim filter and Sieve
// scan loops so both mechanisms report identically.
func mailFilterPendings(path string, mb filterMailbox, rules []filterRule, localDomains map[string]bool, known []string, isNew bool, mechanism mailFilterMechanism) []mailFilterPending {
var out []mailFilterPending
for _, ff := range scoreFilterRules(rules, mb, localDomains, known) {
if ff.onlyIfNew && !isNew {
continue
}
out = append(out, mailFilterPending{
finding: alert.Finding{
Severity: ff.severity,
Check: ff.check,
Message: fmt.Sprintf("%s: %s", mb.String(), ff.reason),
Details: mailFilterDetails(mb, path, ff),
FilePath: path,
Domain: mb.domain,
Mailbox: mailboxField(mb),
},
dest: ff.dest,
mailbox: mb.String(),
mechanism: mechanism,
retainsLocalCopy: ff.retainsLocalCopy,
})
}
return out
}
// scanSieveMailFilters scans per-mailbox Sieve scripts (the rules webmail
// actually executes) for the same BEC exfil patterns the Exim filter audit
// covers. Returns the pending findings plus whether the scan saw every file it
// meant to (a partial scan must not establish the baseline).
func scanSieveMailFilters(ctx context.Context, db *store.DB, cfg *config.Config, localDomains map[string]bool, maxFiles int) ([]mailFilterPending, bool) {
type activePointer struct {
path string
hash string
}
complete := true
var files []string
seen := make(map[string]bool)
activePaths := make(map[string]bool)
activePointers := make(map[string][]activePointer)
addFile := func(path string) {
if !seen[path] {
seen[path] = true
files = append(files, path)
}
}
if scripts, err := homeGlob(ctx, "mail", "*", "*", "sieve", "*.sieve"); err == nil {
for _, path := range scripts {
addFile(path)
}
} else {
complete = false
}
if active, err := homeGlob(ctx, "mail", "*", "*", ".dovecot.sieve"); err == nil {
for _, path := range active {
// The active file normally links to a source already included by the
// sieve/*.sieve glob. Remove that alias before quota accounting so it
// cannot crowd the real source out of every capped scan. A link to an
// unusual target must itself be scanned; on any metadata error, retain
// the path and let ReadFile make the final decision.
if info, statErr := osFS.Lstat(path); statErr == nil && info.Mode()&os.ModeSymlink != 0 {
if target, linkErr := osFS.Readlink(path); linkErr == nil {
if !filepath.IsAbs(target) {
target = filepath.Join(filepath.Dir(path), target)
}
target = filepath.Clean(target)
if seen[target] && mailboxFromSievePath(path) == mailboxFromSievePath(target) {
activePaths[target] = true
activePointers[target] = append(activePointers[target], activePointer{
path: path,
hash: sha256Hex([]byte(target)),
})
continue
}
}
}
addFile(path)
activePaths[path] = true
}
} else {
complete = false
}
complete = complete && scanCoversAllFiles(files, maxFiles)
rankedByMtime := rankPathsByMtimeDesc(ctx, files, 0)
ranked := make([]string, 0, len(rankedByMtime))
for _, path := range rankedByMtime {
if activePaths[path] {
ranked = append(ranked, path)
}
}
for _, path := range rankedByMtime {
if !activePaths[path] {
ranked = append(ranked, path)
}
}
if maxFiles > 0 && len(ranked) > maxFiles {
recordAccountScanTruncatedPaths(ctx, ranked[maxFiles:], maxFiles)
ranked = ranked[:maxFiles]
}
var out []mailFilterPending
reportedContent := make(map[string]bool)
for _, path := range ranked {
if ctx.Err() != nil {
return out, false
}
data, err := osFS.ReadFile(path)
if err != nil {
complete = false
continue
}
currentHash := sha256Hex(data)
isNew := forwarderFileIsNew(db, "email:mailsieve_last_refresh", "mailsieve:"+path, currentHash)
for _, pointer := range activePointers[path] {
if forwarderFileIsNew(db, "email:mailsieve_last_refresh", "mailsieve-active:"+pointer.path, pointer.hash) {
isNew = true
}
if err := db.SetForwarderHash("mailsieve-active:"+pointer.path, pointer.hash); err != nil {
complete = false
}
}
if err := db.SetForwarderHash("mailsieve:"+path, currentHash); err != nil {
complete = false
}
mb := mailboxFromSievePath(path)
contentKey := mb.String() + "\x00" + currentHash
rules := parseSieveFilter(string(data))
pending := mailFilterPendings(path, mb, rules, localDomains, cfg.EmailProtection.KnownForwarders, isNew, mailFilterMechanismSieve)
if len(pending) > 0 && !reportedContent[contentKey] {
reportedContent[contentKey] = true
out = append(out, pending...)
}
}
return out, complete
}
func shouldReportMailFilterStoreUnavailable(st *state.Store, cfg *config.Config) bool {
if st == nil {
return true
}
return st.ShouldRunThrottled("email_mail_filters_store_unavailable", mailFilterAuditIntervalMin(cfg))
}
func mailFilterAuditIntervalMin(cfg *config.Config) int {
if cfg == nil || cfg.EmailProtection.PasswordCheckIntervalMin <= 0 {
return 1440
}
return cfg.EmailProtection.PasswordCheckIntervalMin
}
func mailFilterDetails(mb filterMailbox, path string, ff filterFinding) string {
var b strings.Builder
fmt.Fprintf(&b, "Mailbox: %s\nDomain: %s\nFile: %s\n", mb.String(), mb.domain, path)
if ff.dest != "" {
fmt.Fprintf(&b, "Destination: %s\n", ff.dest)
}
b.WriteString(ff.reason)
return b.String()
}
func mailboxField(mb filterMailbox) string {
if mb.localPart == "*" {
return ""
}
return mb.String()
}
func sha256Hex(data []byte) string {
h := sha256.Sum256(data)
return fmt.Sprintf("%x", h[:])
}
type mailFilterMechanism string
const (
mailFilterMechanismExim mailFilterMechanism = "exim"
mailFilterMechanismSieve mailFilterMechanism = "sieve"
)
// mailFilterPending is an in-flight finding plus the fields needed for the
// corroboration passes before findings are emitted.
type mailFilterPending struct {
finding alert.Finding
dest string
mailbox string
mechanism mailFilterMechanism
// retainsLocalCopy marks an external delivery the mailbox still receives a
// copy of, which is indistinguishable from a user-created forward.
retainsLocalCopy bool
uncorroboratedCopyForward bool
}
// downgradeUncorroboratedCopyExfil lowers an exfil finding to Warning when its
// only stealth evidence is that the forward keeps a local copy. That is exactly
// what a webmail "forward and keep a copy" rule produces, so by itself it cannot
// tell an interception apart from the owner forwarding their own mail. A second
// forwarding mechanism on the same mailbox keeps it Critical, as does mail the
// mailbox never receives; cross-account reuse re-escalates separately.
func downgradeUncorroboratedCopyExfil(collected []mailFilterPending) {
// Count only interception-shaped rules: a plain selective forwarder is
// ordinary mail routing and cannot corroborate anything. Destinations are
// deliberately not part of the key, because planting redundant persistence
// is the signature even when each mechanism aims at a different drop.
mechanisms := map[string]map[mailFilterMechanism]bool{}
for _, p := range collected {
if p.finding.Check != "email_filter_exfil" {
continue
}
if mechanisms[p.mailbox] == nil {
mechanisms[p.mailbox] = map[mailFilterMechanism]bool{}
}
mechanisms[p.mailbox][p.mechanism] = true
}
for i := range collected {
p := &collected[i]
if !p.retainsLocalCopy {
continue
}
if len(mechanisms[p.mailbox]) > 1 {
continue
}
p.finding.Severity = alert.Warning
p.uncorroboratedCopyForward = true
}
}
func applyMailFilterCorroboration(collected []mailFilterPending) {
downgradeUncorroboratedCopyExfil(collected)
// Campaign evidence must run last so it can re-escalate copy-forwards.
annotateCrossAccount(collected)
for i := range collected {
p := &collected[i]
if !p.uncorroboratedCopyForward || p.finding.Severity != alert.Warning {
continue
}
p.finding.Details += "\nThe mailbox still receives this mail and nothing else corroborates the forward, so it is reported for review rather than as a confirmed interception. Confirm with the account owner that the forward is theirs."
}
}
// annotateCrossAccount marks exfil findings whose external destination appears
// across two or more distinct mailboxes -- a strong campaign signal.
func annotateCrossAccount(collected []mailFilterPending) {
byDest := map[string]map[string]bool{}
for _, p := range collected {
if p.finding.Check != "email_filter_exfil" || p.dest == "" {
continue
}
destKey := mailFilterDestinationKey(p.dest)
if byDest[destKey] == nil {
byDest[destKey] = map[string]bool{}
}
byDest[destKey][p.mailbox] = true
}
for i := range collected {
if collected[i].finding.Check != "email_filter_exfil" {
continue
}
dest := collected[i].dest
boxes := byDest[mailFilterDestinationKey(dest)]
if len(boxes) < 2 {
continue
}
others := make([]string, 0, len(boxes))
for b := range boxes {
if b != collected[i].mailbox {
others = append(others, b)
}
}
sort.Strings(others)
collected[i].finding.Severity = alert.Critical
collected[i].finding.Details += fmt.Sprintf(
"\nCross-account: the same destination %s is used by %d mailboxes (also %s). This indicates a coordinated campaign.",
dest, len(boxes), strings.Join(others, ", "))
}
}
func mailFilterDestinationKey(dest string) string {
localPart, domain, ok := splitDeliverDest(dest)
if !ok {
return dest
}
return localPart + "@" + domain
}
func findingsFromPending(collected []mailFilterPending) []alert.Finding {
out := make([]alert.Finding, 0, len(collected))
for _, p := range collected {
out = append(out, p.finding)
}
return out
}
package checks
import (
"context"
"fmt"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/eximlog"
"github.com/pidginhost/csm/internal/state"
)
const perAccountMailThreshold = 100 // emails per recent log window
// mailLogTailLinesDefault is the built-in fallback for how many trailing
// lines of /var/log/exim_mainlog CheckMailPerAccount tails per cycle.
// Operator override: cfg.Thresholds.MailLogTailLines.
const mailLogTailLinesDefault = 500
// CheckMailPerAccount counts recent Exim arrivals per envelope-sender domain.
// Ownership requires the same verified submitter across the entire count.
func CheckMailPerAccount(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
tailLines := mailLogTailLinesDefault
if cfg != nil && cfg.Thresholds.MailLogTailLines > 0 {
tailLines = cfg.Thresholds.MailLogTailLines
}
lines := tailFile("/var/log/exim_mainlog", tailLines)
// Keep the existing sender-domain volume calculation.
type volume struct {
count int
owner string
}
counts := make(map[string]volume)
for _, line := range lines {
// Look for message arrivals (<=).
idx := strings.Index(line, " <= ")
if idx < 0 {
continue
}
// Extract sender address
rest := line[idx+4:]
fields := strings.Fields(rest)
if len(fields) < 1 {
continue
}
sender := fields[0]
// Extract the domain part
atIdx := strings.LastIndex(sender, "@")
if atIdx < 0 {
continue
}
domain := sender[atIdx+1:]
// Skip system/bounce messages
if domain == "" || sender == "<>" || strings.HasPrefix(sender, "cPanel") {
continue
}
identity := eximlog.Submitter(line)
owner := ""
if strings.Contains(identity, "@") {
owner = MailOwner(identity)
} else {
owner = HostingAccountForUser(identity)
}
v := counts[domain]
if v.count == 0 {
v.owner = owner
} else if v.owner != owner {
v.owner = ""
}
v.count++
counts[domain] = v
}
// Keep the sender-domain volume signal, including unverified messages.
// Only a unanimous verified submitter establishes ownership of its count.
for domain, v := range counts {
if v.count >= perAccountMailThreshold {
message := fmt.Sprintf("High email volume from %s: %d messages in recent log", domain, v.count)
if v.owner != "" {
message += fmt.Sprintf(" (account %s)", v.owner)
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "mail_per_account",
Message: message,
Details: "Possible spam outbreak or compromised email account",
Domain: domain,
TenantID: v.owner,
})
}
}
return findings
}
package checks
import (
"net/mail"
"sort"
"strconv"
"strings"
)
// Dovecot/Sieve exfiltration detector.
//
// cPanel webmail (Roundcube) stores per-mailbox Sieve scripts at
// /home/<user>/mail/<domain>/<localpart>/sieve/*.sieve, with the active script
// pointed to by .dovecot.sieve in the mailbox root. A compromised webmail
// account is commonly weaponised with a Sieve rule that redirects every message
// to an external dropbox while keeping a local copy, so the victim never
// notices the interception (business email compromise). Unlike the Exim filter
// path, Sieve is what actually executes for webmail-managed rules, so an
// exfil-only Exim audit misses it entirely.
//
// Sieve is a different grammar from Exim's filter language, so it has its own
// tokenizer and parser here, but it lowers into the same filterRule model the
// Exim scorer already understands (scoreFilterRules): a `redirect` becomes a
// deliver, `:copy` maps to Exim's `unseen` (deliver externally while the
// implicit keep preserves a local copy), `keep`/`fileinto` become a local
// save, and `discard` becomes a /dev/null save. Stealth scoring and reasons
// therefore stay aligned across both mechanisms; an exact KnownForwarders
// entry may additionally acknowledge Roundcube's ordinary :copy forwards.
type sieveTokenKind int
const (
sieveWord sieveTokenKind = iota
sieveString
sieveTag // :copy, :contains, ...
sievePunct
)
type sieveToken struct {
text string
kind sieveTokenKind
}
// tokenizeSieve splits Sieve source into tokens. Quoted strings (with \" and \\
// escapes) become a single string token; tags (:copy) become tag tokens;
// braces, brackets, parens, commas and semicolons are punctuation; line (#) and
// block (/* */) comments are dropped.
func tokenizeSieve(s string) []sieveToken {
var toks []sieveToken
runes := []rune(s)
i := 0
for i < len(runes) {
c := runes[i]
switch {
case c == ' ' || c == '\t' || c == '\n' || c == '\r':
i++
case c == '#':
for i < len(runes) && runes[i] != '\n' {
i++
}
case c == '/' && i+1 < len(runes) && runes[i+1] == '*':
i += 2
for i+1 < len(runes) && (runes[i] != '*' || runes[i+1] != '/') {
i++
}
i += 2
if i > len(runes) {
i = len(runes)
}
case isSieveTextLiteralStart(runes, i):
var text string
text, i = consumeSieveTextLiteral(runes, i)
toks = append(toks, sieveToken{text: decodeSieveEncodedCharacters(text), kind: sieveString})
case c == '"':
i++
var b strings.Builder
for i < len(runes) && runes[i] != '"' {
if runes[i] == '\\' && i+1 < len(runes) {
i++
}
b.WriteRune(runes[i])
i++
}
if i < len(runes) {
i++ // closing quote
}
toks = append(toks, sieveToken{text: decodeSieveEncodedCharacters(b.String()), kind: sieveString})
case c == '{' || c == '}' || c == '[' || c == ']' || c == '(' || c == ')' || c == ',' || c == ';':
toks = append(toks, sieveToken{text: string(c), kind: sievePunct})
i++
case c == ':':
start := i
i++
for i < len(runes) && isSieveIdentRune(runes[i]) {
i++
}
toks = append(toks, sieveToken{text: string(runes[start:i]), kind: sieveTag})
default:
start := i
for i < len(runes) && isSieveIdentRune(runes[i]) {
i++
}
if i == start {
// Unknown punctuation (e.g. a stray operator); skip one rune so
// the tokenizer always makes progress on hostile input.
i++
continue
}
toks = append(toks, sieveToken{text: string(runes[start:i]), kind: sieveWord})
}
}
return toks
}
func isSieveTextLiteralStart(runes []rune, start int) bool {
const prefix = "text:"
if len(runes)-start < len(prefix) {
return false
}
for i, want := range prefix {
got := runes[start+i]
if got >= 'A' && got <= 'Z' {
got += 'a' - 'A'
}
if got != want {
return false
}
}
next := start + len(prefix)
return next == len(runes) || runes[next] == ' ' || runes[next] == '\t' ||
runes[next] == '\r' || runes[next] == '\n' || runes[next] == '#'
}
// consumeSieveTextLiteral consumes Sieve's text: multi-line string form. Its
// contents must stay opaque to the parser: prose in a vacation response can
// legitimately contain text such as `redirect "address";` without being an
// executable action. Unterminated literals consume the rest of the input.
func consumeSieveTextLiteral(runes []rune, start int) (string, int) {
i := start + len("text:")
for i < len(runes) && runes[i] != '\n' {
i++
}
if i == len(runes) {
return "", i
}
i++
var b strings.Builder
for i < len(runes) {
lineStart := i
for i < len(runes) && runes[i] != '\n' {
i++
}
lineEnd := i
if lineEnd > lineStart && runes[lineEnd-1] == '\r' {
lineEnd--
}
line := runes[lineStart:lineEnd]
if len(line) == 1 && line[0] == '.' {
if i < len(runes) {
i++
}
return b.String(), i
}
if len(line) >= 2 && line[0] == '.' && line[1] == '.' {
line = line[1:]
}
b.WriteString(string(line))
if i < len(runes) {
b.WriteByte('\n')
i++
}
}
return b.String(), i
}
// decodeSieveEncodedCharacters handles the standard encoded-character
// extension. Without this normalization, `${hex:40}` can hide the @ in a
// redirect destination from the external-address scorer.
func decodeSieveEncodedCharacters(s string) string {
var out strings.Builder
for i := 0; i < len(s); {
if s[i] != '$' || i+2 >= len(s) || s[i+1] != '{' {
out.WriteByte(s[i])
i++
continue
}
endRel := strings.IndexByte(s[i+2:], '}')
if endRel < 0 {
out.WriteString(s[i:])
break
}
end := i + 2 + endRel
body := s[i+2 : end]
colon := strings.IndexByte(body, ':')
if colon < 0 {
out.WriteByte(s[i])
i++
continue
}
decoded, ok := decodeSieveCharacterSequence(strings.ToLower(strings.TrimSpace(body[:colon])), body[colon+1:])
if !ok {
out.WriteString(s[i : end+1])
} else {
out.WriteString(decoded)
}
i = end + 1
}
return out.String()
}
func decodeSieveCharacterSequence(kind, payload string) (string, bool) {
fields := strings.Fields(payload)
if len(fields) == 0 {
return "", false
}
var out strings.Builder
for _, field := range fields {
if !isSieveHex(field) {
return "", false
}
switch kind {
case "hex":
if len(field) > 2 {
return "", false
}
value, err := strconv.ParseUint(field, 16, 8)
if err != nil {
return "", false
}
out.WriteByte(byte(value))
case "unicode":
if len(field) > 6 {
return "", false
}
value, err := strconv.ParseUint(field, 16, 32)
if err != nil || value > 0x10ffff || (value >= 0xd800 && value <= 0xdfff) {
return "", false
}
out.WriteRune(rune(value))
default:
return "", false
}
}
return out.String(), true
}
func isSieveHex(s string) bool {
if s == "" {
return false
}
for i := range len(s) {
c := s[i]
switch {
case c >= '0' && c <= '9':
case c >= 'a' && c <= 'f':
case c >= 'A' && c <= 'F':
default:
return false
}
}
return true
}
func isSieveIdentRune(r rune) bool {
return r == '_' || r == '.' || r == '-' ||
(r >= 'a' && r <= 'z') || (r >= 'A' && r <= 'Z') || (r >= '0' && r <= '9')
}
func isSievePunct(t sieveToken, text string) bool {
return t.kind == sievePunct && t.text == text
}
var sieveActions = map[string]bool{
"redirect": true,
"keep": true,
"fileinto": true,
"discard": true,
"stop": true,
"return": true,
}
type sieveOrderedAction struct {
position int
action filterAction
}
type sieveBranchChain struct {
parent *sieveNode
allPreviousNever bool
previousAlways bool
continuePossible bool
stopPossible bool
}
type sieveNode struct {
rule filterRule
parent *sieveNode
branch *sieveBranchChain
actions []sieveOrderedAction
reachable bool
flow sieveTruth
}
func (c *sieveBranchChain) recordBranch(flow sieveTruth) {
c.continuePossible = c.continuePossible || flow != sieveNever
c.stopPossible = c.stopPossible || flow != sieveAlways
}
func finishSieveBranchChain(c *sieveBranchChain) {
if c == nil || c.parent == nil {
return
}
// With no final else, messages that match none of the branch tests flow
// through the chain unchanged.
if !c.previousAlways {
c.continuePossible = true
}
continuation := sieveConditional
switch {
case c.continuePossible && !c.stopPossible:
continuation = sieveAlways
case !c.continuePossible && c.stopPossible:
continuation = sieveNever
}
c.parent.flow = composeSieveFlow(c.parent.flow, continuation)
}
func composeSieveFlow(first, second sieveTruth) sieveTruth {
if first == sieveNever || second == sieveNever {
return sieveNever
}
if first == sieveAlways {
return second
}
return sieveConditional
}
// parseSieveFilter parses Sieve source into the flat filterRule list the Exim
// scorer consumes: one rule per if/elsif/else branch plus the unconditional top
// level, with ancestor actions folded into each branch so a keep in an outer
// block still pairs with a redirect in an inner branch.
func parseSieveFilter(content string) []filterRule {
toks := tokenizeSieve(content)
top := &sieveNode{rule: filterRule{matchesAll: true}, reachable: true, flow: sieveAlways}
stack := []*sieveNode{top}
nodes := []*sieveNode{top}
var pendingChain *sieveBranchChain
i := 0
for i < len(toks) {
t := toks[i]
continuesBranch := pendingChain != nil && pendingChain.parent == stack[len(stack)-1] &&
t.kind == sieveWord && (strings.EqualFold(t.text, "elsif") || strings.EqualFold(t.text, "else"))
if pendingChain != nil && !continuesBranch {
finishSieveBranchChain(pendingChain)
pendingChain = nil
}
if isSievePunct(t, "}") {
if len(stack) > 1 {
closed := stack[len(stack)-1]
if closed.reachable {
closed.branch.recordBranch(closed.flow)
}
pendingChain = closed.branch
stack = stack[:len(stack)-1]
} else {
pendingChain = nil
}
i++
continue
}
if t.kind == sieveWord && (strings.EqualFold(t.text, "if") || strings.EqualFold(t.text, "elsif") || strings.EqualFold(t.text, "else")) {
keyword := strings.ToLower(t.text)
isElse := keyword == "else"
parent := stack[len(stack)-1]
chain := pendingChain
if keyword == "if" || chain == nil || chain.parent != parent {
chain = &sieveBranchChain{parent: parent, allPreviousNever: keyword == "if"}
}
pendingChain = nil
i++
condStart := i
for i < len(toks) && !isSievePunct(toks[i], "{") {
i++
}
branchReachable := !chain.previousAlways
branchMatchesAll := chain.allPreviousNever
if isElse {
chain.previousAlways = true
chain.allPreviousNever = false
} else {
truth := sieveTestTruthValue(toks[condStart:i])
branchReachable = branchReachable && truth != sieveNever
branchMatchesAll = branchMatchesAll && truth == sieveAlways
chain.previousAlways = chain.previousAlways || truth == sieveAlways
chain.allPreviousNever = chain.allPreviousNever && truth == sieveNever
}
if i < len(toks) {
i++ // consume "{"
}
n := &sieveNode{
rule: filterRule{matchesAll: parent.rule.matchesAll && parent.flow == sieveAlways && branchMatchesAll},
parent: parent,
branch: chain,
reachable: parent.reachable && parent.flow != sieveNever && branchReachable,
flow: sieveAlways,
}
nodes = append(nodes, n)
stack = append(stack, n)
continue
}
pendingChain = nil
if t.kind == sieveWord && sieveActions[strings.ToLower(t.text)] {
verb := strings.ToLower(t.text)
position := i
i++
var args []sieveToken
for i < len(toks) && !isSievePunct(toks[i], ";") {
if isSievePunct(toks[i], "{") || isSievePunct(toks[i], "}") {
break
}
args = append(args, toks[i])
i++
}
if i >= len(toks) || !isSievePunct(toks[i], ";") {
continue
}
i++
cur := stack[len(stack)-1]
if cur.flow == sieveNever {
continue
}
if act, ok := sieveActionToFilter(verb, args); ok {
act.matchesAll = cur.rule.matchesAll && cur.flow == sieveAlways
cur.actions = append(cur.actions, sieveOrderedAction{position: position, action: act})
if act.verb == "finish" {
cur.flow = sieveNever
}
}
continue
}
i++
}
finishSieveBranchChain(pendingChain)
out := make([]filterRule, 0, len(nodes))
for _, n := range nodes {
if n.reachable && len(n.actions) > 0 {
out = append(out, flattenSieveNode(n))
}
}
return out
}
// sieveActionToFilter lowers one Sieve action into the Exim filterAction the
// scorer understands. Returns ok=false for actions that carry no exfil signal.
func sieveActionToFilter(verb string, args []sieveToken) (filterAction, bool) {
switch verb {
case "redirect":
dest := lastSieveString(args)
if dest == "" {
return filterAction{}, false
}
if address, err := mail.ParseAddress(dest); err == nil {
dest = address.Address
}
// :copy keeps the implicit local copy alongside the forward -- the same
// stealth signal as Exim's `unseen`.
copyRedirect := sieveHasTag(args, ":copy")
return filterAction{
verb: "deliver",
arg: dest,
unseen: copyRedirect,
knownSuppressible: copyRedirect,
}, true
case "fileinto":
dest := lastSieveString(args)
if dest == "" {
return filterAction{}, false
}
return filterAction{verb: "save", arg: dest}, true
case "keep":
return filterAction{verb: "save", arg: "$home/mail/INBOX"}, true
case "discard":
return filterAction{verb: "save", arg: "/dev/null"}, true
case "stop", "return":
return filterAction{verb: "finish"}, true
}
return filterAction{}, false
}
func lastSieveString(args []sieveToken) string {
for i := len(args) - 1; i >= 0; i-- {
if args[i].kind == sieveString {
return args[i].text
}
}
return ""
}
func sieveHasTag(args []sieveToken, tag string) bool {
for _, a := range args {
if a.kind == sieveTag && strings.EqualFold(a.text, tag) {
return true
}
}
return false
}
func flattenSieveNode(node *sieveNode) filterRule {
var chain []*sieveNode
for n := node; n != nil; n = n.parent {
chain = append(chain, n)
}
var actions []sieveOrderedAction
for _, n := range chain {
actions = append(actions, n.actions...)
}
sort.SliceStable(actions, func(i, j int) bool {
return actions[i].position < actions[j].position
})
// Sieve actions carry their own match-all reachability because a prior
// conditional stop can make later actions selective within the same scope.
// Exim rules retain their rule-wide matchesAll representation.
out := filterRule{}
for _, ordered := range actions {
out.actions = append(out.actions, ordered.action)
if ordered.action.verb == "finish" {
break
}
}
return out
}
type sieveTruth int
const (
sieveNever sieveTruth = iota
sieveConditional
sieveAlways
)
// sieveTestMatchesAll reports whether a Sieve test fires on effectively all
// mail. It evaluates anyof/allof structure instead of combining unrelated
// operands from different subtests, and stays conservative on malformed or
// excessively nested input.
func sieveTestMatchesAll(test []sieveToken) bool {
return sieveTestTruthValue(test) == sieveAlways
}
func sieveTestTruthValue(test []sieveToken) sieveTruth {
return sieveExpressionTruth(test, 0)
}
func sieveExpressionTruth(test []sieveToken, depth int) sieveTruth {
const maxSieveTestDepth = 256
if depth >= maxSieveTestDepth {
return sieveConditional
}
test, ok := trimSieveOuterParens(test)
if !ok || len(test) == 0 || test[0].kind != sieveWord {
return sieveConditional
}
keyword := strings.ToLower(test[0].text)
switch keyword {
case "true":
if len(test) == 1 {
return sieveAlways
}
case "false":
if len(test) == 1 {
return sieveNever
}
case "not":
inner := sieveExpressionTruth(test[1:], depth+1)
switch inner {
case sieveAlways:
return sieveNever
case sieveNever:
return sieveAlways
default:
return sieveConditional
}
case "anyof", "allof":
parts, ok := sieveCallArguments(test[1:])
if !ok || len(parts) == 0 {
return sieveConditional
}
if keyword == "anyof" {
allNever := true
for _, part := range parts {
truth := sieveExpressionTruth(part, depth+1)
if truth == sieveAlways {
return sieveAlways
}
allNever = allNever && truth == sieveNever
}
if allNever {
return sieveNever
}
return sieveConditional
}
allAlways := true
for _, part := range parts {
truth := sieveExpressionTruth(part, depth+1)
if truth == sieveNever {
return sieveNever
}
allAlways = allAlways && truth == sieveAlways
}
if allAlways {
return sieveAlways
}
return sieveConditional
case "address", "header":
if sieveAddressTestMatchesAll(test) {
return sieveAlways
}
}
return sieveConditional
}
func trimSieveOuterParens(test []sieveToken) ([]sieveToken, bool) {
matchingClose := make([]int, len(test))
for i := range matchingClose {
matchingClose[i] = -1
}
var stack []int
for i, tok := range test {
switch {
case isSievePunct(tok, "("):
stack = append(stack, i)
case isSievePunct(tok, ")"):
if len(stack) == 0 {
return nil, false
}
open := stack[len(stack)-1]
stack = stack[:len(stack)-1]
matchingClose[open] = i
}
}
if len(stack) != 0 {
return nil, false
}
left, right := 0, len(test)-1
for left < right && isSievePunct(test[left], "(") && matchingClose[left] == right {
left++
right--
}
return test[left : right+1], true
}
func sieveCallArguments(test []sieveToken) ([][]sieveToken, bool) {
if len(test) < 2 || !isSievePunct(test[0], "(") || !isSievePunct(test[len(test)-1], ")") {
return nil, false
}
return splitSieveTopLevel(test[1:len(test)-1], ",")
}
func splitSieveTopLevel(test []sieveToken, separator string) ([][]sieveToken, bool) {
var parts [][]sieveToken
start := 0
parenDepth := 0
listDepth := 0
for i, tok := range test {
switch {
case isSievePunct(tok, "("):
parenDepth++
case isSievePunct(tok, ")"):
parenDepth--
case isSievePunct(tok, "["):
listDepth++
case isSievePunct(tok, "]"):
listDepth--
case parenDepth == 0 && listDepth == 0 && isSievePunct(tok, separator):
if start == i {
return nil, false
}
parts = append(parts, test[start:i])
start = i + 1
}
if parenDepth < 0 || listDepth < 0 {
return nil, false
}
}
if parenDepth != 0 || listDepth != 0 || start >= len(test) {
return nil, false
}
parts = append(parts, test[start:])
return parts, true
}
func sieveAddressTestMatchesAll(test []sieveToken) bool {
testKind := strings.ToLower(test[0].text)
matchType := ""
addressPart := ":all"
addressPartSeen := false
comparatorSeen := false
var operands [][]string
for i := 1; i < len(test); {
tok := test[i]
switch {
case tok.kind == sieveTag:
tag := strings.ToLower(tok.text)
switch tag {
case ":contains", ":matches", ":is":
if matchType != "" {
return false
}
matchType = tag
case ":all", ":localpart", ":domain", ":detail":
if testKind != "address" {
return false
}
if addressPartSeen {
return false
}
addressPartSeen = true
addressPart = tag
case ":comparator":
if comparatorSeen {
return false
}
comparatorSeen = true
i++
if i >= len(test) || test[i].kind != sieveString {
return false
}
default:
// Extensions such as :index/:last and :mime/:anychild alter
// which header occurrence or MIME part is examined. Even From
// is not guaranteed to exist in that narrowed scope, so it is
// unsafe to classify the test as match-all. Unknown modifiers
// stay conservative for the same reason.
return false
}
i++
case tok.kind == sieveString:
operands = append(operands, []string{tok.text})
i++
case isSievePunct(tok, "["):
var values []string
i++
for i < len(test) && !isSievePunct(test[i], "]") {
if test[i].kind == sieveString {
values = append(values, test[i].text)
} else if !isSievePunct(test[i], ",") {
return false
}
i++
}
if i >= len(test) || len(values) == 0 {
return false
}
i++
operands = append(operands, values)
default:
return false
}
}
if len(operands) != 2 {
return false
}
// The first string-list is the header name and the final string-list is
// the match key. Only From is reliably present and address-bearing on every
// normal message; Subject/To and arbitrary headers are not match-all.
hasFrom := false
for _, header := range operands[0] {
if strings.EqualFold(strings.TrimSpace(header), "from") {
hasFrom = true
break
}
}
if !hasFrom {
return false
}
for _, key := range operands[1] {
switch matchType {
case ":contains":
if key == "" || (addressPart == ":all" && key == "@") {
return true
}
case ":matches":
if key == "*" || (addressPart == ":all" && key == "*@*") {
return true
}
}
}
return false
}
// mailboxFromSievePath derives the mailbox from a sieve script path of the form
// /home/<user>/mail/<domain>/<localpart>/... .
func mailboxFromSievePath(path string) filterMailbox {
parts := strings.Split(path, "/")
for i := 0; i+4 < len(parts); i++ {
if parts[i] == "home" && parts[i+1] != "" && parts[i+2] == "mail" && parts[i+3] != "" && parts[i+4] != "" {
return filterMailbox{localPart: parts[i+4], domain: parts[i+3]}
}
}
return filterMailbox{localPart: "*", domain: ""}
}
package checks
import (
"context"
"fmt"
"net"
"strconv"
"strings"
"sync"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
func CheckOutboundConnections(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Parse /proc/net/tcp for established connections
// Format: sl local_address rem_address st ...
// local_address = IP:port we are listening/connecting from
// rem_address = IP:port of the remote end
data, err := osFS.ReadFile("/proc/net/tcp")
if err != nil {
return nil
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Fields(line)
if len(fields) < 4 {
continue
}
if fields[0] == "sl" {
continue
}
// State 01 = ESTABLISHED
if fields[3] != "01" {
continue
}
firstFinding := len(findings)
// Ownership belongs to this kernel row, not to a later PID lookup.
// Keep detections visible when an incomplete row has no usable UID.
var uid uint64
var uidKnown bool
if len(fields) >= 8 {
var err error
uid, err = strconv.ParseUint(fields[7], 10, 32)
uidKnown = err == nil
}
localAddr := fields[1]
remoteAddr := fields[2]
_, localPort := parseHexAddr(localAddr)
remoteIP, remotePort := parseHexAddr(remoteAddr)
if remoteIP == "" || remoteIP == "127.0.0.1" || remoteIP == "0.0.0.0" {
continue
}
// Check remote IP against C2 blocklist
for _, blocked := range cfg.C2Blocklist {
if remoteIP == blocked {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "c2_connection",
Message: fmt.Sprintf("Connection to known C2 IP: %s:%d", remoteIP, remotePort),
Details: fmt.Sprintf("Local port: %d", localPort),
SourceIP: remoteIP,
})
}
}
// Check if OUR LOCAL port is a backdoor port (we're listening on it)
// This catches backdoor listeners, not clients connecting from high ports
for _, bp := range cfg.BackdoorPorts {
if localPort == bp {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "backdoor_port",
Message: fmt.Sprintf("Listening on known backdoor port %d, connected from %s:%d", localPort, remoteIP, remotePort),
SourceIP: remoteIP,
})
}
}
// Also check if we're connecting OUT to a backdoor port on a remote host
// (e.g. reverse shell calling back to attacker's listener)
// Skip if our local port is a known service (the remote port is just
// the client's ephemeral port, not a backdoor listener)
knownServicePorts := map[int]bool{
21: true, 25: true, 26: true, 53: true, 80: true, 110: true,
143: true, 443: true, 465: true, 587: true, 993: true, 995: true,
2082: true, 2083: true, 2086: true, 2087: true, 2095: true, 2096: true,
3306: true, 4190: true,
}
if !knownServicePorts[localPort] {
for _, bp := range cfg.BackdoorPorts {
if remotePort == bp {
if isInfraIP(remoteIP, cfg.InfraIPs) {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "backdoor_port_outbound",
Message: fmt.Sprintf("Outbound connection to backdoor port: %s:%d", remoteIP, remotePort),
Details: fmt.Sprintf("Local port: %d", localPort),
SourceIP: remoteIP,
})
}
}
}
if uidKnown {
for i := firstFinding; i < len(findings); i++ {
AttributeSocketOwner(&findings[i], uint32(uid))
}
}
}
return findings
}
// IsInfraIP reports whether ip is infrastructure the daemon must never act
// against: an operator infra entry (CIDR or address) or a Cloudflare edge.
// One implementation serves scans and the realtime path alike.
func IsInfraIP(ip string, infraNets []string) bool { return isInfraIP(ip, infraNets) }
// IsCloudflareIP reports whether ip is inside the published Cloudflare
// ranges the daemon last refreshed.
func IsCloudflareIP(ip net.IP) bool { return isCloudflareIP(ip) }
func isInfraIP(ip string, infraNets []string) bool {
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
for _, entry := range infraNets {
entry = strings.TrimSpace(entry)
if entry == "" {
continue
}
// Try CIDR first (e.g. "10.0.0.0/8")
_, network, err := net.ParseCIDR(entry)
if err == nil {
if network.Contains(parsed) {
return true
}
continue
}
// Fall back to plain IP match (e.g. "1.2.3.4")
if entryIP := net.ParseIP(entry); entryIP != nil && entryIP.Equal(parsed) {
return true
}
}
// Also check Cloudflare IPs - these must never be blocked/challenged
// because blocking a CF edge IP blocks thousands of legitimate users.
// The detection/alert still fires; only the nftables action is skipped.
if isCloudflareIP(parsed) {
return true
}
return false
}
var (
cfNets []*net.IPNet
cfNetsMu sync.RWMutex
)
// SetCloudflareNets updates the cached Cloudflare IP ranges.
// Called by the daemon after fetching CF IPs.
func SetCloudflareNets(cidrs []string) {
var nets []*net.IPNet
for _, cidr := range cidrs {
_, network, err := net.ParseCIDR(cidr)
if err == nil {
nets = append(nets, network)
}
}
cfNetsMu.Lock()
cfNets = nets
cfNetsMu.Unlock()
}
func isCloudflareIP(ip net.IP) bool {
cfNetsMu.RLock()
defer cfNetsMu.RUnlock()
for _, network := range cfNets {
if network.Contains(ip) {
return true
}
}
return false
}
func parseHexAddr(hexAddr string) (string, int) {
parts := strings.Split(hexAddr, ":")
if len(parts) != 2 {
return "", 0
}
hexIP := parts[0]
hexPort := parts[1]
if len(hexIP) != 8 {
return "", 0
}
// Parse little-endian hex IP
var octets [4]byte
for i := 0; i < 4; i++ {
val := hexToByte(hexIP[6-2*i : 8-2*i])
octets[i] = val
}
ip := net.IPv4(octets[0], octets[1], octets[2], octets[3]).String()
port, ok := parseProcNetHexPort(hexPort)
if !ok {
return ip, 0
}
return ip, port
}
func parseProcNetHexPort(hexPort string) (int, bool) {
if len(hexPort) != 4 {
return 0, false
}
port := 0
for i := 0; i < len(hexPort); i++ {
c := hexPort[i]
if !isHexDigit(c) {
return 0, false
}
port = port*16 + hexVal(c)
}
return port, true
}
func hexToByte(s string) byte {
if len(s) != 2 {
return 0
}
// #nosec G115 -- hexVal returns 0..15; (h<<4)|h fits in a byte (0..255).
return byte(hexVal(s[0])<<4 | hexVal(s[1]))
}
func hexVal(c byte) int {
switch {
case c >= '0' && c <= '9':
return int(c - '0')
case c >= 'a' && c <= 'f':
return int(c-'a') + 10
case c >= 'A' && c <= 'F':
return int(c-'A') + 10
}
return 0
}
package checks
import (
"bytes"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"syscall"
)
// fixPerfAllowedRoots scopes performance remediations to per-account web
// content. Same shape as the other fix*AllowedRoots so tests can swap in
// a t.TempDir(); nil means the platform's account roots.
var fixPerfAllowedRoots []string
// FixErrorLogBloat truncates an account-owned error_log file in place.
// Truncating preserves the inode and file ownership so any PHP process
// holding the descriptor keeps appending to the same file without an
// open/reopen race; this is also the safest action because nothing in
// the host needs the historical lines to keep serving traffic.
func FixErrorLogBloat(path string) RemediationResult {
return FixErrorLogBloatInRoots(path, effectiveFixRoots(fixPerfAllowedRoots))
}
// FixErrorLogBloatInRoots is FixErrorLogBloat with caller-supplied roots.
// The Web UI uses this to include configured account_roots while tests can
// keep writes under t.TempDir().
func FixErrorLogBloatInRoots(path string, allowedRoots []string) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
resolved, info, err := resolveExistingFixPath(path, allowedRoots)
if err != nil {
return RemediationResult{Error: err.Error()}
}
if info.IsDir() {
return RemediationResult{Error: "refusing to truncate a directory"}
}
if filepath.Base(resolved) != "error_log" {
return RemediationResult{Error: fmt.Sprintf("refusing to truncate non error_log file: %s", resolved)}
}
oldSize := info.Size()
if err := truncateFilePreservingIdentity(resolved, info); err != nil {
return RemediationResult{Error: fmt.Sprintf("truncate failed: %v", err)}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("truncate %s", resolved),
Description: fmt.Sprintf("Emptied error_log (was %s)", humanBytes(oldSize)),
}
}
// FixDisplayErrorsOn rewrites an INI / .htaccess / .user.ini file so the
// display_errors directive is set to Off. The original line is preserved
// commented out for operator review; an override line is appended at the
// end of the file so the last-write-wins semantics of every supported
// config format land on Off regardless of earlier statements.
//
// Only .user.ini, php.ini, and .htaccess are accepted. wp-config.php and
// other PHP source files require code-level edits this routine does not
// attempt. The caller (web UI) should not advertise the fix for those.
func FixDisplayErrorsOn(path string) RemediationResult {
return FixDisplayErrorsOnInRoots(path, effectiveFixRoots(fixPerfAllowedRoots))
}
// FixDisplayErrorsOnInRoots is FixDisplayErrorsOn with caller-supplied
// roots. The Web UI uses this to honor account_roots outside /home.
func FixDisplayErrorsOnInRoots(path string, allowedRoots []string) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
resolved, info, err := resolveExistingFixPath(path, allowedRoots)
if err != nil {
return RemediationResult{Error: err.Error()}
}
if info.IsDir() {
return RemediationResult{Error: "refusing to edit a directory"}
}
base := filepath.Base(resolved)
var (
isHtaccess bool
supported bool
)
switch {
case base == ".user.ini" || base == "php.ini" || strings.HasSuffix(base, ".ini"):
supported = true
case base == ".htaccess":
supported = true
isHtaccess = true
}
if !supported {
return RemediationResult{Error: fmt.Sprintf("automated display_errors fix only supports .user.ini, php.ini, and .htaccess (got %s)", base)}
}
// #nosec G304 -- path was validated by resolveExistingFixPath against
// the supplied remediation roots; symlinks already rejected.
data, err := readFilePreservingIdentity(resolved, info)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("read failed: %v", err)}
}
rewritten, changedLines := commentDisplayErrorsLines(data)
if changedLines == 0 {
return RemediationResult{Error: "no display_errors directive found in file"}
}
if isHtaccess {
rewritten = appendHtaccessOverride(rewritten)
} else {
rewritten = appendIniOverride(rewritten)
}
// Preserve ownership + mode. Write atomically via a sibling temp file +
// rename so a partial write does not leave the operator with a broken
// config.
if err := writeFilePreservingOwner(resolved, rewritten, info); err != nil {
return RemediationResult{Error: err.Error()}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("disable display_errors in %s", resolved),
Description: fmt.Sprintf(
"Commented %d display_errors line(s) and appended an Off override at end of file",
changedLines,
),
}
}
// commentDisplayErrorsLines walks the file line-by-line, comments out any
// non-comment line whose directive is display_errors (matching the
// detector's logic). Returns the rewritten bytes and the count of lines
// changed.
func commentDisplayErrorsLines(data []byte) ([]byte, int) {
var out bytes.Buffer
changed := 0
for _, raw := range bytes.SplitAfter(data, []byte("\n")) {
if len(raw) == 0 {
continue
}
line := raw
newline := []byte(nil)
if bytes.HasSuffix(raw, []byte("\n")) {
line = raw[:len(raw)-1]
newline = []byte("\n")
}
lineText := string(line)
trimmed := strings.TrimSpace(lineText)
if trimmed == "" || strings.HasPrefix(trimmed, "#") || strings.HasPrefix(trimmed, ";") {
out.Write(raw)
continue
}
if !strings.Contains(strings.ToLower(trimmed), "display_errors") {
out.Write(raw)
continue
}
out.WriteString("# csm: disabled by remediation -- ")
out.WriteString(lineText)
out.Write(newline)
changed++
}
return out.Bytes(), changed
}
func appendIniOverride(data []byte) []byte {
override := "display_errors = Off"
return appendOverrideLine(data, override, "; csm: appended by remediation")
}
func appendHtaccessOverride(data []byte) []byte {
override := "php_flag display_errors Off"
return appendOverrideLine(data, override, "# csm: appended by remediation")
}
func appendOverrideLine(data []byte, directive, marker string) []byte {
var out bytes.Buffer
out.Write(data)
if len(data) > 0 && data[len(data)-1] != '\n' {
out.WriteByte('\n')
}
out.WriteString(marker)
out.WriteByte('\n')
out.WriteString(directive)
out.WriteByte('\n')
return out.Bytes()
}
// These helpers re-check the target inode after the initial path validation
// so a file swap between validation and mutation fails closed.
func truncateFilePreservingIdentity(path string, expected os.FileInfo) error {
// #nosec G304 -- path was validated by resolveExistingFixPath against
// the per-account remediation roots; the open uses O_NOFOLLOW and the
// subsequent sameFileIdentity check fails closed on inode swap.
f, err := os.OpenFile(path, os.O_WRONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return err
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return err
}
if !sameFileIdentity(info, expected) {
return fmt.Errorf("file changed during remediation")
}
return f.Truncate(0)
}
func readFilePreservingIdentity(path string, expected os.FileInfo) ([]byte, error) {
// #nosec G304 -- path was validated by resolveExistingFixPath against
// the per-account remediation roots; the open uses O_NOFOLLOW and the
// subsequent sameFileIdentity check fails closed on inode swap.
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return nil, err
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil, err
}
if !sameFileIdentity(info, expected) {
return nil, fmt.Errorf("file changed during remediation")
}
return io.ReadAll(f)
}
func writeFilePreservingOwner(path string, data []byte, original os.FileInfo) error {
dir := filepath.Dir(path)
tmp, createErr := os.CreateTemp(dir, ".csm-perf-fix-*")
if createErr != nil {
return fmt.Errorf("create temp: %v", createErr)
}
tmpPath := tmp.Name()
cleanup := func() { _ = os.Remove(tmpPath) }
if _, werr := tmp.Write(data); werr != nil {
_ = tmp.Close()
cleanup()
return fmt.Errorf("write temp: %v", werr)
}
if cerr := tmp.Close(); cerr != nil {
cleanup()
return fmt.Errorf("close temp: %v", cerr)
}
if merr := os.Chmod(tmpPath, original.Mode().Perm()); merr != nil {
cleanup()
return fmt.Errorf("chmod temp: %v", merr)
}
info, statErr := os.Lstat(path)
if statErr != nil {
cleanup()
return fmt.Errorf("stat original: %v", statErr)
}
if info.Mode()&os.ModeSymlink != 0 {
cleanup()
return fmt.Errorf("original became a symlink during remediation")
}
if !sameFileIdentity(info, original) {
cleanup()
return fmt.Errorf("file changed during remediation")
}
if stat, ok := info.Sys().(*syscall.Stat_t); ok {
if chownErr := os.Chown(tmpPath, int(stat.Uid), int(stat.Gid)); chownErr != nil {
cleanup()
return fmt.Errorf("chown temp: %v", chownErr)
}
}
if renameErr := os.Rename(tmpPath, path); renameErr != nil {
cleanup()
return fmt.Errorf("rename: %v", renameErr)
}
return nil
}
func sameFileIdentity(a, b os.FileInfo) bool {
if a == nil || b == nil {
return false
}
if os.SameFile(a, b) {
return true
}
as, aok := a.Sys().(*syscall.Stat_t)
bs, bok := b.Sys().(*syscall.Stat_t)
return aok && bok && as.Dev == bs.Dev && as.Ino == bs.Ino
}
package checks
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/mysqlclient"
"github.com/pidginhost/csm/internal/redisinfo"
"github.com/pidginhost/csm/internal/state"
)
// perfEnabled returns false only if Performance.Enabled is explicitly set to false.
// nil (unset) is treated as enabled.
func perfEnabled(cfg *config.Config) bool {
if cfg.Performance.Enabled == nil {
return true
}
return *cfg.Performance.Enabled
}
// cpuCoresOnce guards the cached CPU core count.
var (
cpuCoresOnce sync.Once
cpuCoresCache int
)
// getCPUCores reads /proc/cpuinfo and counts "processor\t" lines.
// The result is cached after the first call. Returns 1 on error.
func getCPUCores() int {
cpuCoresOnce.Do(func() {
f, err := osFS.Open("/proc/cpuinfo")
if err != nil {
cpuCoresCache = 1
return
}
defer func() { _ = f.Close() }()
count := 0
scanner := bufio.NewScanner(f)
for scanner.Scan() {
if strings.HasPrefix(scanner.Text(), "processor\t") {
count++
}
}
if count == 0 {
count = 1
}
cpuCoresCache = count
})
return cpuCoresCache
}
// parseLoadAvg reads /proc/loadavg and returns the first three load average
// values (1m, 5m, 15m).
func parseLoadAvg() ([3]float64, error) {
var result [3]float64
data, err := osFS.ReadFile("/proc/loadavg")
if err != nil {
return result, fmt.Errorf("reading /proc/loadavg: %w", err)
}
fields := strings.Fields(string(data))
if len(fields) < 3 {
return result, fmt.Errorf("unexpected /proc/loadavg format: %q", string(data))
}
for i := 0; i < 3; i++ {
v, err := strconv.ParseFloat(fields[i], 64)
if err != nil {
return result, fmt.Errorf("parsing load avg field %d: %w", i, err)
}
result[i] = v
}
return result, nil
}
// parseMemInfo reads /proc/meminfo and returns total memory, available memory,
// swap total, and swap free - all in kilobytes.
func parseMemInfo() (total, available, swapTotal, swapFree uint64) {
f, err := osFS.Open("/proc/meminfo")
if err != nil {
return
}
defer func() { _ = f.Close() }()
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := scanner.Text()
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
key := fields[0]
val, err := strconv.ParseUint(fields[1], 10, 64)
if err != nil {
continue
}
switch key {
case "MemTotal:":
total = val
case "MemAvailable:":
available = val
case "SwapTotal:":
swapTotal = val
case "SwapFree:":
swapFree = val
}
}
return
}
const (
maxInt64Value = int64(1<<63 - 1)
maxUint64Value = ^uint64(0)
)
// uint64ToInt64Clamped narrows a uint64 to int64 for byte display.
func uint64ToInt64Clamped(v uint64) int64 {
if v > uint64(maxInt64Value) {
return maxInt64Value
}
return int64(v)
}
func redisLargeDatasetThresholdBytes(gb int) uint64 {
if gb <= 0 {
return 0
}
const gbBytes uint64 = 1024 * 1024 * 1024
thresholdGB := uint64(gb)
if thresholdGB > maxUint64Value/gbBytes {
return maxUint64Value
}
return thresholdGB * gbBytes
}
// humanBytes formats a byte count as a human-readable string.
// Thresholds: >=1G → "1.0G", >=1M → "1M", >=1K → "1K", else "0B".
func humanBytes(b int64) string {
const (
KB = 1024
MB = 1024 * KB
GB = 1024 * MB
)
switch {
case b >= GB:
return fmt.Sprintf("%.1fG", float64(b)/float64(GB))
case b >= MB:
return fmt.Sprintf("%dM", b/MB)
case b >= KB:
return fmt.Sprintf("%dK", b/KB)
default:
return "0B"
}
}
func kibToDisplayBytes(kib uint64) int64 {
const maxKiBForInt64Bytes = uint64(maxInt64Value / 1024)
if kib > maxKiBForInt64Bytes {
return maxInt64Value
}
return int64(kib) * 1024
}
// CheckLoadAverage compares load averages against per-core thresholds
// from config. The 1-minute load drives the Critical / High findings;
// when 1-minute is below the High threshold we additionally check the
// 5- and 15-minute averages for sustained pressure (>= 0.7 * High
// threshold on both) and emit a Warning. The sustained variant catches
// "constant 22%-of-cores busy for 15 minutes" which is invisible to a
// 1-minute spike check but is what operators actually want to see.
func CheckLoadAverage(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
loads, err := parseLoadAvg()
if err != nil {
return nil
}
cores := getCPUCores()
load1 := loads[0]
critThreshold := float64(cores) * cfg.Performance.LoadCriticalMultiplier
highThreshold := float64(cores) * cfg.Performance.LoadHighMultiplier
switch {
case load1 > critThreshold:
return []alert.Finding{{
Severity: alert.Critical,
Check: "perf_load",
Message: "High load average exceeds critical threshold",
Details: fmt.Sprintf("Load: %.1f/%.1f/%.1f, Cores: %d, Threshold: %.1f",
loads[0], loads[1], loads[2], cores, critThreshold),
Timestamp: time.Now(),
}}
case load1 > highThreshold:
return []alert.Finding{{
Severity: alert.High,
Check: "perf_load",
Message: "High load average exceeds high threshold",
Details: fmt.Sprintf("Load: %.1f/%.1f/%.1f, Cores: %d, Threshold: %.1f",
loads[0], loads[1], loads[2], cores, highThreshold),
Timestamp: time.Now(),
}}
}
// Sustained pressure: 1-minute is calm but 5- and 15-minute
// averages are both above 70% of the High threshold. This is the
// "load 9 on 40 cores for 15 minutes" shape -- below the spike
// threshold but a real operator concern.
sustainedThreshold := highThreshold * 0.7
if loads[1] > sustainedThreshold && loads[2] > sustainedThreshold {
return []alert.Finding{{
Severity: alert.Warning,
Check: "perf_load",
Message: "Sustained load (5m + 15m) above 70% of high threshold",
Details: fmt.Sprintf("Load: %.1f/%.1f/%.1f, Cores: %d, Sustained threshold: %.1f",
loads[0], loads[1], loads[2], cores, sustainedThreshold),
Timestamp: time.Now(),
}}
}
return nil
}
// phpWorkersByUser walks /proc and returns, per username, the cmdline samples
// of that user's live PHP web-worker processes. Instantaneous snapshot; callers that
// need a count use len(result[user]).
func phpWorkersByUser() map[string][]string {
cmdlinePaths, _ := osFS.Glob("/proc/[0-9]*/cmdline")
userProcs := make(map[string][]string)
for _, cmdPath := range cmdlinePaths {
pid := filepath.Base(filepath.Dir(cmdPath))
data, err := osFS.ReadFile(cmdPath)
if err != nil {
continue
}
cmdStr := strings.ReplaceAll(string(data), "\x00", " ")
cmdStr = strings.TrimSpace(cmdStr)
safeCmdStr := redactProcCommandLine(data)
if !isPHPWorkerCommand(cmdStr) {
continue
}
// Read UID from status
statusData, _ := osFS.ReadFile(filepath.Join("/proc", pid, "status"))
var uid string
for _, line := range strings.Split(string(statusData), "\n") {
if strings.HasPrefix(line, "Uid:\t") {
fields := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(fields) > 0 {
uid = fields[0]
}
break
}
}
if uid == "" {
uid = "unknown"
}
username := uidStringToUser(uid)
userProcs[username] = append(userProcs[username], safeCmdStr)
}
return userProcs
}
func isPHPWorkerCommand(command string) bool {
fields := strings.Fields(command)
if len(fields) == 0 {
return false
}
name := strings.TrimSuffix(strings.ToLower(filepath.Base(fields[0])), ":")
return isVersionedPHPWorkerBinary(name, "lsphp") ||
isVersionedPHPWorkerBinary(name, "php-cgi") ||
isVersionedPHPWorkerBinary(name, "php-fpm")
}
func isVersionedPHPWorkerBinary(name, base string) bool {
if name == base {
return true
}
suffix := strings.TrimPrefix(name, base)
if suffix == name || suffix == "" || suffix[0] < '0' || suffix[0] > '9' {
return false
}
for _, char := range suffix {
if (char < '0' || char > '9') && char != '.' {
return false
}
}
return suffix[len(suffix)-1] != '.'
}
// CheckPHPProcessLoad scans /proc for PHP web workers, groups them by user,
// and fires Critical if total exceeds cores*multiplier, High per user if
// individual count exceeds threshold.
func CheckPHPProcessLoad(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
cores := getCPUCores()
userProcs := phpWorkersByUser()
total := 0
for _, procs := range userProcs {
total += len(procs)
}
if total == 0 {
return nil
}
var findings []alert.Finding
// Critical: total PHP worker count exceeds cores * multiplier
critTotalThreshold := cores * cfg.Performance.PHPProcessCriticalTotalMult
if total > critTotalThreshold {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "perf_php_processes",
Message: "Total PHP worker process count exceeds critical threshold",
Details: fmt.Sprintf("Count: %d, Threshold: %d (cores: %d × %d)", total, critTotalThreshold, cores, cfg.Performance.PHPProcessCriticalTotalMult),
Timestamp: time.Now(),
})
}
// High: per-user count exceeds threshold
for username, procs := range userProcs {
if len(procs) > cfg.Performance.PHPProcessWarnPerUser {
// Collect up to 3 sample cmdlines
samples := procs
if len(samples) > 3 {
samples = samples[:3]
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_php_processes",
Message: fmt.Sprintf("Excessive PHP worker processes for user %s", username),
Details: fmt.Sprintf("Count: %d, Threshold: %d, Sample cmdlines: %s", len(procs), cfg.Performance.PHPProcessWarnPerUser, alert.RedactCommandLine(strings.Join(samples, " | "))),
Timestamp: time.Now(),
})
}
}
return findings
}
// CheckSwapAndOOM checks for OOM killer events in dmesg and elevated swap
// usage from /proc/meminfo. Host OOM is Critical, cgroup OOM Warning, and
// swap usage above 50% High.
func CheckSwapAndOOM(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
var findings []alert.Finding
// Check dmesg for OOM events
// Prefer ISO timestamps so we can filter to the last hour.
// Fall back to -T (human-readable) on older kernels that don't support --time-format.
dmesgOut, isoErr := runCmd("dmesg", "--time-format", "iso", "--level=err")
useISO := isoErr == nil && dmesgOut != nil
if !useISO {
dmesgOut, _ = runCmd("dmesg", "--level=err", "-T")
}
if dmesgOut != nil {
cutoff := time.Now().Add(-1 * time.Hour)
seen := make(map[string]bool)
for _, line := range strings.Split(string(dmesgOut), "\n") {
lower := strings.ToLower(line)
if !strings.Contains(lower, "out of memory") && !strings.Contains(lower, "oom_reaper") {
continue
}
// Both the ISO and the -T fallback are filtered to the last
// hour. A line whose timestamp cannot be parsed is skipped, not
// reported: an OOM event we cannot date is exactly the stale
// finding that previously fired a Critical on every scan.
when, ok := parseDmesgOOMTime(line, useISO)
if !ok || when.Before(cutoff) {
continue
}
key := oomDedupKey(line)
if seen[key] {
continue
}
seen[key] = true
severity, accountScoped := classifyOOMLine(line)
message := "OOM killer invoked in the last hour"
if accountScoped {
message = "Account memory limit reached in the last hour"
}
findings = append(findings, alert.Finding{
Severity: severity,
Check: "perf_memory",
Message: message,
Details: strings.TrimSpace(line),
// Every kill logs a fresh pid and byte counts; keying dedup on
// the victim process name keeps an ongoing OOM loop to one
// finding per state-expiry window instead of one per scan.
DedupKey: key,
Timestamp: time.Now(),
})
}
}
// Check swap usage
_, _, swapTotal, swapFree := parseMemInfo()
if swapTotal > 0 {
swapUsed := swapTotal - swapFree
usagePct := float64(swapUsed) / float64(swapTotal) * 100
if usagePct > 50 {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_memory",
Message: "High swap usage",
Details: fmt.Sprintf("Swap used: %s / %s (%.0f%%)", humanBytes(kibToDisplayBytes(swapUsed)), humanBytes(kibToDisplayBytes(swapTotal)), usagePct),
// The percentage moves every scan; without a pinned identity
// each drift re-alerts the same sustained condition.
DedupKey: "swap_high",
Timestamp: time.Now(),
})
}
}
return findings
}
// oomVictimProcess returns the killed process name from a dmesg OOM line
// ("... Killed process 2845662 (lsphp) ..."), or "host" when the line names
// no victim, so the dedup identity always has a stable value. Parenthesized
// OOM context before the process marker must not be mistaken for the victim.
func oomVictimProcess(line string) string {
for _, marker := range []string{"Killed process ", "reaped process "} {
markerIdx := strings.Index(line, marker)
if markerIdx < 0 {
continue
}
rest := line[markerIdx+len(marker):]
pidEnd := strings.IndexAny(rest, " \t")
if pidEnd <= 0 {
continue
}
if _, err := strconv.ParseUint(rest[:pidEnd], 10, 64); err != nil {
continue
}
rest = strings.TrimLeft(rest[pidEnd:], " \t")
if len(rest) < 3 || rest[0] != '(' {
continue
}
closeIdx := strings.IndexByte(rest[1:], ')')
if closeIdx <= 0 {
continue
}
if process := strings.TrimSpace(rest[1 : closeIdx+1]); process != "" {
return process
}
}
return "host"
}
// classifyOOMLine separates a host-wide OOM from a cgroup one. On a shared
// host a cgroup kill is an account reaching the memory limit its plan sets:
// routine, and not evidence about host health. Only real memory exhaustion is
// Critical, or the two become indistinguishable in the alert stream.
func classifyOOMLine(line string) (alert.Severity, bool) {
if strings.Contains(strings.ToLower(line), "memory cgroup out of memory") {
return alert.Warning, true
}
return alert.Critical, false
}
// oomDedupKey keeps the account-scoped and host-wide cases on separate dedup
// identities, so one account repeatedly hitting its limit cannot suppress the
// host-wide alert that follows it.
func oomDedupKey(line string) string {
if _, accountScoped := classifyOOMLine(line); accountScoped {
return "oom:cgroup:" + oomVictimProcess(line)
}
return "oom:host:" + oomVictimProcess(line)
}
// parseDmesgOOMTime extracts the event time from a dmesg line. ISO lines
// (--time-format iso) carry an absolute timestamp with a timezone offset as
// the first field. The -T fallback carries a bracketed ctime in local time
// ("[Mon Jan _2 15:04:05 2006] ..."). Returns ok=false when no timestamp can
// be parsed, so the caller drops the line rather than reporting an undatable
// (and therefore possibly stale) OOM event.
func parseDmesgOOMTime(line string, useISO bool) (time.Time, bool) {
if useISO {
// 2006-01-02T15:04:05,000000+0300 -- comma decimal, first field.
ts := strings.Replace(strings.SplitN(line, " ", 2)[0], ",", ".", 1)
for _, layout := range []string{"2006-01-02T15:04:05.000000-0700", "2006-01-02T15:04:05.000000-07:00"} {
if parsed, err := time.Parse(layout, ts); err == nil {
return parsed, true
}
}
return time.Time{}, false
}
open := strings.IndexByte(line, '[')
closeIdx := strings.IndexByte(line, ']')
if open != 0 || closeIdx <= open {
return time.Time{}, false
}
// dmesg -T prints local wall-clock time with no zone, so parse in Local.
parsed, err := time.ParseInLocation("Mon Jan _2 15:04:05 2006", strings.TrimSpace(line[open+1:closeIdx]), time.Local)
if err != nil {
return time.Time{}, false
}
return parsed, true
}
// CheckPHPHandler detects PHP CGI handler usage on LiteSpeed servers.
// On LiteSpeed, CGI is significantly slower than LSAPI; this check fires
// a Critical finding for each PHP version using the CGI handler.
func CheckPHPHandler(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
// Only relevant on LiteSpeed
if _, err := osFS.Stat("/usr/local/lsws/bin/litespeed"); err != nil {
return nil
}
var cgiVersions []string
// Try whmapi1 first
out, err := runCmd("whmapi1", "php_get_handlers", "--output=json")
if err == nil && len(out) > 0 {
// Parse JSON: look for handler entries with type "cgi"
var result struct {
Data struct {
Handlers []struct {
Version string `json:"version"`
Handler string `json:"handler"`
Type string `json:"type"`
} `json:"handlers"`
} `json:"data"`
}
if jsonErr := json.Unmarshal(out, &result); jsonErr == nil {
for _, h := range result.Data.Handlers {
t := strings.ToLower(h.Handler + " " + h.Type)
if strings.Contains(t, "cgi") && !strings.Contains(t, "lsapi") && !strings.Contains(t, "fpm") {
cgiVersions = append(cgiVersions, h.Version)
}
}
}
} else {
// Fallback: read /etc/cpanel/ea4/ea4.conf
data, readErr := osFS.ReadFile("/etc/cpanel/ea4/ea4.conf")
if readErr == nil {
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
// Lines like: ea-php74.handler = cgi
if !strings.Contains(line, ".handler") {
continue
}
parts := strings.SplitN(line, "=", 2)
if len(parts) != 2 {
continue
}
val := strings.TrimSpace(parts[1])
if val == "cgi" {
versionPart := strings.TrimSpace(parts[0])
cgiVersions = append(cgiVersions, versionPart)
}
}
}
}
if len(cgiVersions) == 0 {
return nil
}
return []alert.Finding{{
Severity: alert.Critical,
Check: "perf_php_handler",
Message: "PHP handler set to CGI instead of LSAPI on LiteSpeed",
Details: fmt.Sprintf("Affected PHP versions: %s", strings.Join(cgiVersions, ", ")),
Timestamp: time.Now(),
}}
}
// CheckMySQLConfig inspects MySQL global variables and runtime status for
// performance-impacting misconfigurations. Each issue emits its own finding
// with a stable message so deduplication works correctly.
func CheckMySQLConfig(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
var findings []alert.Finding
// --- Global variables ---
varRows, err := mysqlclient.RootQuery(ctx,
"SHOW GLOBAL VARIABLES WHERE Variable_name IN ('join_buffer_size','wait_timeout','interactive_timeout','max_user_connections','slow_query_log')")
varOut := []byte(strings.Join(varRows, "\n"))
if err == nil && len(varOut) > 0 {
joinBufThresholdBytes := int64(cfg.Performance.MySQLJoinBufferMaxMB) * 1024 * 1024
waitTimeoutMax := cfg.Performance.MySQLWaitTimeoutMax
for _, line := range strings.Split(string(varOut), "\n") {
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
name := fields[0]
val := fields[1]
switch name {
case "join_buffer_size":
v, convErr := strconv.ParseInt(val, 10, 64)
if convErr == nil && v > joinBufThresholdBytes {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "perf_mysql_config",
Message: "MySQL join_buffer_size exceeds safe maximum",
Details: fmt.Sprintf("Current: %s, Max: %s", humanBytes(v), humanBytes(joinBufThresholdBytes)),
Timestamp: time.Now(),
})
}
case "wait_timeout":
v, convErr := strconv.Atoi(val)
if convErr == nil && v > waitTimeoutMax {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_mysql_config",
Message: "MySQL wait_timeout is too high",
Details: fmt.Sprintf("Current: %ds, Max: %ds", v, waitTimeoutMax),
Timestamp: time.Now(),
})
}
case "interactive_timeout":
v, convErr := strconv.Atoi(val)
if convErr == nil && v > waitTimeoutMax {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_mysql_config",
Message: "MySQL interactive_timeout is too high",
Details: fmt.Sprintf("Current: %ds, Max: %ds", v, waitTimeoutMax),
Timestamp: time.Now(),
})
}
case "max_user_connections":
if val == "0" {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_mysql_config",
Message: "MySQL max_user_connections is unlimited",
Details: fmt.Sprintf("Current: 0 (unlimited), Recommended: %d", cfg.Performance.MySQLMaxConnectionsPerUser),
Timestamp: time.Now(),
})
}
case "slow_query_log":
if strings.ToUpper(val) == "OFF" {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_mysql_config",
Message: "MySQL slow query log is disabled",
Details: "Set slow_query_log=ON to help diagnose performance issues",
Timestamp: time.Now(),
})
}
}
}
}
// --- InnoDB buffer pool hit ratio + temporary disk tables ---
statusRows, err := mysqlclient.RootQuery(ctx,
"SHOW GLOBAL STATUS WHERE Variable_name IN ('Innodb_buffer_pool_read_requests','Innodb_buffer_pool_reads','Created_tmp_disk_tables','Created_tmp_tables')")
statusOut := []byte(strings.Join(statusRows, "\n"))
if err == nil && len(statusOut) > 0 {
var readRequests, reads, tmpDiskTables, tmpTables int64
for _, line := range strings.Split(string(statusOut), "\n") {
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
v, convErr := strconv.ParseInt(fields[1], 10, 64)
if convErr != nil {
continue
}
switch fields[0] {
case "Innodb_buffer_pool_read_requests":
readRequests = v
case "Innodb_buffer_pool_reads":
reads = v
case "Created_tmp_disk_tables":
tmpDiskTables = v
case "Created_tmp_tables":
tmpTables = v
}
}
if tmpTables > 0 && tmpDiskTables > 0 {
diskRatio := float64(tmpDiskTables) / float64(tmpTables) * 100
if diskRatio > 25.0 {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_mysql_config",
Message: "MySQL creating excessive temporary tables on disk",
Details: fmt.Sprintf("Disk ratio: %.1f%% (%d disk tables / %d total tables)", diskRatio, tmpDiskTables, tmpTables),
Timestamp: time.Now(),
})
}
}
if readRequests > 0 {
hitRatio := float64(readRequests-reads) / float64(readRequests) * 100
if hitRatio < 95.0 {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_mysql_config",
Message: "InnoDB buffer pool hit ratio is low",
Details: fmt.Sprintf("Hit ratio: %.1f%% (threshold: 95%%), disk reads: %d", hitRatio, reads),
Timestamp: time.Now(),
})
}
}
}
// --- Per-user connection counts ---
plRows, err := mysqlclient.RootQuery(ctx, "SHOW PROCESSLIST")
if err == nil && len(plRows) > 0 {
userCounts := make(map[string]int)
for _, line := range plRows {
fields := strings.Fields(line)
// SHOW PROCESSLIST columns: Id, User, Host, db, Command, Time, State, Info
if len(fields) < 2 {
continue
}
user := fields[1]
if user == "" || user == "User" {
continue
}
userCounts[user]++
}
maxConn := cfg.Performance.MySQLMaxConnectionsPerUser
for dbUser, count := range userCounts {
if count > maxConn {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_mysql_config",
Message: fmt.Sprintf("MySQL user %s holding excessive connections", dbUser),
Details: fmt.Sprintf("Connections: %d, Threshold: %d", count, maxConn),
Timestamp: time.Now(),
})
}
}
}
return findings
}
// CheckRedisConfig inspects a local Redis instance for performance-impacting
// misconfigurations: unset maxmemory, noeviction policy, non-expiring keys,
// and an overly aggressive bgsave schedule for the dataset size.
func CheckRedisConfig(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
// Skip on hosts without a local redis: peek at the redisinfo
// client by trying a fast INFO server. Any error short-circuits
// the whole check (matches the historical behaviour where a
// missing redis-cli binary made every redis check a no-op).
if _, _, err := redisinfo.MemoryUsage(ctx); err != nil {
return nil
}
var findings []alert.Finding
// --- maxmemory ---
if maxMem, err := redisinfo.ConfigGet(ctx, "maxmemory"); err == nil && maxMem == "0" {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "perf_redis_config",
Message: "Redis maxmemory is not set",
Details: "maxmemory=0 means Redis will use all available system memory without bound",
Timestamp: time.Now(),
})
}
// --- maxmemory-policy ---
policy, policyErr := redisinfo.ConfigGet(ctx, "maxmemory-policy")
policyLower := ""
if policyErr == nil {
policyLower = strings.ToLower(strings.TrimSpace(policy))
}
if policyLower == "noeviction" {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_redis_config",
Message: "Redis maxmemory-policy is noeviction",
Details: "noeviction causes Redis to return errors when memory is full instead of evicting keys",
Timestamp: time.Now(),
})
}
// --- Non-expiring keys ratio via keyspace ---
// A high non-expiring ratio only breaks eviction under volatile-* policies
// (which evict keys carrying a TTL) or noeviction. Under allkeys-* Redis
// evicts any key, so non-expiring keys are reclaimable and the ratio is
// benign.
if stats, err := redisinfo.KeyspaceStats(ctx); err == nil && stats.TotalKeys > 0 && !strings.HasPrefix(policyLower, "allkeys-") {
nonExpiring := stats.TotalKeys - stats.TotalExpires
ratio := float64(nonExpiring) / float64(stats.TotalKeys) * 100
if ratio > 95.0 {
details := fmt.Sprintf("Non-expiring: %d / %d total keys (%.1f%%); %s",
nonExpiring, stats.TotalKeys, ratio, redisNonExpiringPolicyDetail(policyLower))
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_redis_config",
Message: "Redis has excessive non-expiring keys",
Details: details,
Timestamp: time.Now(),
})
}
}
// --- bgsave interval vs dataset size ---
saveSpec, _ := redisinfo.ConfigGet(ctx, "save")
usedBytes, _, _ := redisinfo.MemoryUsage(ctx)
largeDatasetBytes := redisLargeDatasetThresholdBytes(cfg.Performance.RedisLargeDatasetGB)
bgsaveMinInterval := cfg.Performance.RedisBgsaveMinInterval
if usedBytes > largeDatasetBytes && saveSpec != "" {
// `CONFIG GET save` returns the spec as a single space-separated
// string of alternating "<seconds> <changes>" pairs, e.g.
// "900 1 300 10 60 10000". Walk the seconds tokens (every
// other field) and flag any below the configured floor.
fields := strings.Fields(saveSpec)
aggressiveSave := false
for i := 0; i < len(fields); i += 2 {
seconds, convErr := strconv.Atoi(fields[i])
if convErr == nil && seconds < bgsaveMinInterval {
aggressiveSave = true
break
}
}
if aggressiveSave {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_redis_config",
Message: "Redis bgsave interval too aggressive for dataset size",
Details: fmt.Sprintf(
"Used memory: %s, Threshold: %s, Minimum safe bgsave interval: %ds",
humanBytes(uint64ToInt64Clamped(usedBytes)),
humanBytes(uint64ToInt64Clamped(largeDatasetBytes)),
bgsaveMinInterval,
),
Timestamp: time.Now(),
})
}
}
// --- used_memory vs maxmemory headroom ---
// The maxmemory==0 branch above flags the unset case. When maxmemory
// IS set, used/max ratio is the operator-meaningful signal: at 80%
// the eviction policy is about to start churning hot keys; at 90%
// noeviction-policy instances start returning OOM errors.
_, maxBytes, _ := redisinfo.MemoryUsage(ctx)
if maxBytes > 0 && usedBytes > 0 {
pct := float64(usedBytes) / float64(maxBytes) * 100
switch {
case pct >= 90:
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "perf_redis_config",
Message: "Redis used memory >= 90% of maxmemory",
Details: fmt.Sprintf(
"Used: %s / Max: %s (%.1f%%)",
humanBytes(uint64ToInt64Clamped(usedBytes)),
humanBytes(uint64ToInt64Clamped(maxBytes)),
pct,
),
Timestamp: time.Now(),
})
case pct >= 80:
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_redis_config",
Message: "Redis used memory >= 80% of maxmemory",
Details: fmt.Sprintf(
"Used: %s / Max: %s (%.1f%%)",
humanBytes(uint64ToInt64Clamped(usedBytes)),
humanBytes(uint64ToInt64Clamped(maxBytes)),
pct,
),
Timestamp: time.Now(),
})
}
}
return findings
}
func redisNonExpiringPolicyDetail(policyLower string) string {
switch {
case policyLower == "":
return "maxmemory-policy is unavailable, so non-expiring keys may be unsafe under memory pressure"
case policyLower == "noeviction":
return "maxmemory-policy noeviction does not evict keys under memory pressure"
case strings.HasPrefix(policyLower, "volatile-"):
return fmt.Sprintf("maxmemory-policy %q only evicts keys with a TTL under memory pressure", policyLower)
default:
return fmt.Sprintf("maxmemory-policy %q may leave non-expiring keys unreclaimable under memory pressure", policyLower)
}
}
// ---------------------------------------------------------------------------
// Performance check helpers (WP-specific)
// ---------------------------------------------------------------------------
// safeIdentifier returns true if s matches ^[a-zA-Z0-9_]+$ (non-empty).
// Used to reject values with shell metacharacters before use in commands/SQL.
var safeIdentRe = regexp.MustCompile(`^[a-zA-Z0-9_]+$`)
func safeIdentifier(s string) bool {
return s != "" && safeIdentRe.MatchString(s)
}
// extractPHPDefine extracts the value argument from a PHP define() line:
//
// define('KEY', 'value'); or define("KEY", "value");
//
// It is distinct from extractDefine (dbscan.go) which requires a key parameter.
// Returns the empty string if no value can be extracted.
func extractPHPDefine(line string) string {
// Trim whitespace and trailing semicolons/comments.
line = strings.TrimSpace(line)
// Find the opening parenthesis.
parenIdx := strings.Index(line, "(")
if parenIdx < 0 {
return ""
}
inner := line[parenIdx+1:]
// Strip closing paren and anything after.
if closeIdx := strings.LastIndex(inner, ")"); closeIdx >= 0 {
inner = inner[:closeIdx]
}
// inner is now like: 'KEY', 'value' or "KEY", "value"
// Split on the first comma, ignoring the key part.
commaIdx := strings.Index(inner, ",")
if commaIdx < 0 {
return ""
}
valuePart := strings.TrimSpace(inner[commaIdx+1:])
if valuePart == "" {
return ""
}
// Strip surrounding quotes (single or double) when present.
q := valuePart[0]
if q == '\'' || q == '"' {
if len(valuePart) < 2 {
return ""
}
end := strings.LastIndexByte(valuePart, q)
if end <= 0 {
return ""
}
return valuePart[1:end]
}
// Unquoted literal (boolean/number constant). Strip a trailing ); or
// whitespace and return the bare token. Examples wp-config.php uses:
// define('DISABLE_WP_CRON', true);
// define('WP_DEBUG', false);
// define('WP_MEMORY_LIMIT', 256);
for i, c := range valuePart {
if c == ' ' || c == '\t' || c == ';' || c == ')' || c == ',' {
return strings.TrimSpace(valuePart[:i])
}
}
return strings.TrimSpace(valuePart)
}
// ---------------------------------------------------------------------------
// Subdirs to skip in recursive helpers.
// ---------------------------------------------------------------------------
var skipDirs = map[string]bool{
"wp-admin": true,
"wp-content": true,
"wp-includes": true,
"cache": true,
"node_modules": true,
"vendor": true,
}
// ---------------------------------------------------------------------------
// CheckErrorLogBloat
// ---------------------------------------------------------------------------
// scanErrorLogs recursively walks dir up to maxDepth looking for error_log
// files larger than threshold bytes. Results are appended to *findings (capped
// at 20).
func scanErrorLogs(dir string, thresholdBytes int64, depth int, findings *[]alert.Finding) {
if depth < 0 || len(*findings) >= 20 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, e := range entries {
if len(*findings) >= 20 {
return
}
name := e.Name()
fullPath := filepath.Join(dir, name)
if e.IsDir() {
if skipDirs[name] {
continue
}
scanErrorLogs(fullPath, thresholdBytes, depth-1, findings)
continue
}
if name != "error_log" {
continue
}
info, statErr := e.Info()
if statErr != nil {
continue
}
if info.Size() > thresholdBytes {
*findings = append(*findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_error_logs",
Message: fmt.Sprintf("Bloated error_log: %s", fullPath),
Details: fmt.Sprintf("Size: %s", humanBytes(info.Size())),
Timestamp: time.Now(),
})
}
}
}
// CheckErrorLogBloat walks configured web roots (default /home/*/public_html
// on cPanel) looking for error_log files that exceed the configured size
// threshold. The runner enforces a 60-minute throttle via checkThrottleMin.
func CheckErrorLogBloat(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
thresholdBytes := int64(cfg.Performance.ErrorLogWarnSizeMB) * 1024 * 1024
homeDirs := ResolveWebRoots(cfg)
var findings []alert.Finding
for _, dir := range homeDirs {
scanErrorLogs(dir, thresholdBytes, 3, &findings)
if len(findings) >= 20 {
break
}
}
return findings
}
// ---------------------------------------------------------------------------
// CheckWPConfig
// ---------------------------------------------------------------------------
// parseMemoryLimit converts a PHP memory_limit string (e.g. "256M", "1G")
// to megabytes. Returns 0 if the value cannot be parsed.
func parseMemoryLimit(s string) int {
s = strings.TrimSpace(strings.ToUpper(s))
if s == "" || s == "-1" {
return 0
}
suffix := s[len(s)-1]
numStr := s
mult := 1
switch suffix {
case 'K':
numStr = s[:len(s)-1]
v, err := strconv.Atoi(numStr)
if err != nil {
return 0
}
return v / 1024
case 'M':
numStr = s[:len(s)-1]
mult = 1
case 'G':
numStr = s[:len(s)-1]
mult = 1024
}
v, err := strconv.Atoi(numStr)
if err != nil {
return 0
}
return v * mult
}
// scanWPConfigs recursively searches dir (max depth) for wp-config.php files
// and checks WP_MEMORY_LIMIT and co-located config files for issues.
func scanWPConfigs(dir, account string, cfg *config.Config, depth int, findings *[]alert.Finding) {
if depth < 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, e := range entries {
name := e.Name()
fullPath := filepath.Join(dir, name)
if e.IsDir() {
if skipDirs[name] {
continue
}
scanWPConfigs(fullPath, account, cfg, depth-1, findings)
continue
}
if name != "wp-config.php" {
continue
}
// --- WP_MEMORY_LIMIT ---
wpData, readErr := osFS.ReadFile(fullPath)
if readErr == nil {
for _, line := range strings.Split(string(wpData), "\n") {
if strings.Contains(line, "WP_MEMORY_LIMIT") {
val := extractPHPDefine(strings.TrimSpace(line))
if mb := parseMemoryLimit(val); mb > cfg.Performance.WPMemoryLimitMaxMB {
*findings = append(*findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_wp_config",
Message: fmt.Sprintf("Excessive WP_MEMORY_LIMIT for %s", account),
Details: fmt.Sprintf("File: %s, Value: %s", fullPath, val),
Timestamp: time.Now(),
})
}
break
}
}
}
// --- Co-located PHP config files ---
wpDir := filepath.Dir(fullPath)
for _, cfgFile := range []string{".htaccess", "php.ini", ".user.ini"} {
cfgPath := filepath.Join(wpDir, cfgFile)
data, readErr2 := osFS.ReadFile(cfgPath)
if readErr2 != nil {
continue
}
// cPanel MultiPHP INI Editor writes .user.ini with a fixed
// header and owns the file's content. Values inside a
// cPanel-managed .user.ini (max_execution_time=0 for a
// backup importer, display_errors=On for a staging account)
// reflect operator choices made through the cPanel UI and
// are not attacker actions. Suppress findings for this
// file in that case — operators do not need alerts for
// their own configuration. The suppression is scoped
// strictly to .user.ini: the same signature in php.ini or
// .htaccess is not authoritative (cPanel does not write
// those files) and the scanner treats it normally.
if cfgFile == ".user.ini" && isCpanelManagedUserIni(data) {
continue
}
for _, line := range strings.Split(string(data), "\n") {
trimmed := strings.TrimSpace(line)
// Skip comment lines
if strings.HasPrefix(trimmed, "#") || strings.HasPrefix(trimmed, ";") {
continue
}
lc := strings.ToLower(trimmed)
switch {
case strings.Contains(lc, "max_execution_time"):
// max_execution_time = 0 (or php_value max_execution_time 0)
parts := strings.FieldsFunc(trimmed, func(r rune) bool { return r == '=' || r == ' ' || r == '\t' })
if len(parts) >= 2 && parts[len(parts)-1] == "0" {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "perf_wp_config",
Message: fmt.Sprintf("Unlimited max_execution_time for %s", account),
Details: fmt.Sprintf("File: %s, Value: 0", cfgPath),
Timestamp: time.Now(),
})
}
case strings.Contains(lc, "display_errors"):
parts := strings.FieldsFunc(trimmed, func(r rune) bool { return r == '=' || r == ' ' || r == '\t' })
if len(parts) >= 2 && strings.ToLower(parts[len(parts)-1]) == "on" {
*findings = append(*findings, alert.Finding{
Severity: alert.Warning,
Check: "perf_wp_config",
Message: fmt.Sprintf("display_errors enabled in production for %s", account),
Details: fmt.Sprintf("File: %s, Value: On", cfgPath),
Timestamp: time.Now(),
})
}
}
}
}
}
}
// CheckWPConfig scans /home/*/public_html (max depth 2) for wp-config.php
// files and reports excessive WP_MEMORY_LIMIT values, unlimited
// max_execution_time, and display_errors enabled in production.
// The runner enforces a 60-minute throttle via checkThrottleMin.
func CheckWPConfig(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
homeDirs := ResolveWebRoots(cfg)
var findings []alert.Finding
for _, dir := range homeDirs {
scanWPConfigs(dir, accountFromPath(dir), cfg, 2, &findings)
}
return findings
}
// accountFromPath extracts a best-effort account name from a web root path.
// On cPanel (/home/USER/public_html) it returns USER. On other layouts it
// returns the parent directory name, or the final path component if there
// is no parent. Used for reporting only — never for authorization.
func accountFromPath(dir string) string {
parts := strings.Split(dir, string(filepath.Separator))
// cPanel shape: /home/<account>/public_html
for i, p := range parts {
if p == "home" && i+1 < len(parts) {
return parts[i+1]
}
}
// Generic shape: /var/www/<site>, /srv/http/<site>, etc.
if len(parts) >= 2 && parts[len(parts)-1] != "" {
return parts[len(parts)-2]
}
return filepath.Base(dir)
}
// ---------------------------------------------------------------------------
// CheckWPTransientBloat
// ---------------------------------------------------------------------------
// findWPTransients recursively searches dir for wp-config.php files and
// queries the WordPress database for bloated transients.
func findWPTransients(dir string, cfg *config.Config, warnBytes, critBytes int64, depth int, findings *[]alert.Finding) {
if depth < 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, e := range entries {
name := e.Name()
fullPath := filepath.Join(dir, name)
if e.IsDir() {
if skipDirs[name] {
continue
}
findWPTransients(fullPath, cfg, warnBytes, critBytes, depth-1, findings)
continue
}
if name != "wp-config.php" {
continue
}
info := parseWPConfig(fullPath)
if info.dbName == "" || info.dbUser == "" {
continue
}
// Apply default table prefix when not set.
if info.tablePrefix == "" {
info.tablePrefix = "wp_"
}
// Security: validate identifiers before use in SQL.
if !safeIdentifier(info.dbName) || !safeIdentifier(info.dbUser) || !safeIdentifier(info.tablePrefix) {
continue
}
query := fmt.Sprintf(
"SELECT option_name, LENGTH(option_value) as size FROM %soptions WHERE option_name LIKE '_transient_%%' AND LENGTH(option_value) > %d ORDER BY size DESC LIMIT 5",
info.tablePrefix,
warnBytes,
)
rows, runErr := mysqlclient.PerAccountQuery(context.Background(), mysqlclient.Creds{
User: info.dbUser,
Password: info.dbPass,
Host: info.dbHost,
DBName: info.dbName,
}, query)
if runErr != nil || len(rows) == 0 {
continue
}
for _, line := range rows {
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
optionName := fields[0]
sizeBytes, convErr := strconv.ParseInt(fields[1], 10, 64)
if convErr != nil {
continue
}
var sev alert.Severity
switch {
case sizeBytes > critBytes:
sev = alert.High
case sizeBytes > warnBytes:
sev = alert.Warning
default:
continue
}
*findings = append(*findings, alert.Finding{
Severity: sev,
Check: "perf_wp_transients",
Message: fmt.Sprintf("Bloated transient %s in %s", optionName, info.dbName),
Details: fmt.Sprintf("Size: %s", humanBytes(sizeBytes)),
Timestamp: time.Now(),
})
}
}
}
// CheckWPTransientBloat scans configured web roots (default /home/*/public_html
// on cPanel) for WordPress installs and queries each database for oversized
// transients. DB credentials are read from wp-config.php; the password is
// passed via MYSQL_PWD environment variable (never on the command line).
// The runner enforces a 60-minute throttle via checkThrottleMin.
func CheckWPTransientBloat(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
warnBytes := int64(cfg.Performance.WPTransientWarnMB) * 1024 * 1024
critBytes := int64(cfg.Performance.WPTransientCriticalMB) * 1024 * 1024
homeDirs := ResolveWebRoots(cfg)
var findings []alert.Finding
for _, dir := range homeDirs {
findWPTransients(dir, cfg, warnBytes, critBytes, 2, &findings)
}
return findings
}
// ---------------------------------------------------------------------------
// CheckWPCron
// ---------------------------------------------------------------------------
const (
wpCronFindingLimit = 30
wpCronCursorKey = "_wpcron_scan_cursor"
)
var (
wpCronTablePrefixAssignRe = regexp.MustCompile(`(?im)^[\t ]*\$table_prefix[\t ]*=`)
wpCronRequireRe = regexp.MustCompile(`(?im)^[\t ]*require(?:_once)?\b`)
)
type wpCronScanRoot struct {
path string
account string
}
type wpCronCandidate struct {
path string
account string
}
type wpCronScanCursor struct {
Root string `json:"root"`
Path string `json:"path,omitempty"`
}
// scanWPCronCandidates recursively finds real WordPress installs. When a
// nested vhost is also a scan root, the broader root leaves that subtree to
// the more-specific root. This preserves the full depth allowance for both
// roots without reading or reporting the same install twice.
func scanWPCronCandidates(
ctx context.Context,
scanRoot, dir, account string,
depth int,
allRoots map[string]struct{},
candidates map[string]wpCronCandidate,
) bool {
if depth < 0 {
return true
}
if ctx.Err() != nil {
return false
}
cleanDir := filepath.Clean(dir)
if cleanDir != filepath.Clean(scanRoot) {
if _, nestedRoot := allRoots[cleanDir]; nestedRoot {
return true
}
}
entries, err := osFS.ReadDir(cleanDir)
if err != nil {
return errors.Is(err, fs.ErrNotExist)
}
complete := true
for _, entry := range entries {
if ctx.Err() != nil {
return false
}
name := entry.Name()
fullPath := filepath.Join(cleanDir, name)
if entry.IsDir() {
if skipDirs[name] {
continue
}
if !scanWPCronCandidates(ctx, scanRoot, fullPath, account, depth-1, allRoots, candidates) {
complete = false
}
continue
}
if name != "wp-config.php" || entry.Type()&os.ModeSymlink != 0 {
continue
}
info, infoErr := entry.Info()
if infoErr != nil {
complete = false
continue
}
if !info.Mode().IsRegular() {
continue
}
data, readErr := osFS.ReadFile(fullPath)
if readErr != nil {
if !errors.Is(readErr, fs.ErrNotExist) {
complete = false
}
continue
}
if wpCronHasActiveDisableDefine(data) {
continue
}
isWordPress, validateErr := wpCronInstallIsValid(fullPath, data)
if validateErr != nil {
complete = false
continue
}
if !isWordPress {
continue
}
candidates[fullPath] = wpCronCandidate{path: fullPath, account: account}
}
return complete
}
// wpCronInstallIsValid requires both the WordPress bootstrap shape and the
// core files the remediation will invoke. A stray or backup wp-config.php is
// not enough evidence to edit customer data or install a crontab entry.
func wpCronInstallIsValid(configPath string, data []byte) (bool, error) {
code := stripPHPCommentsFromCode(phpCodeOnly(string(data)))
codeWithoutStrings := stripPHPStringsFromCode(code)
if !wpCronTablePrefixAssignRe.MatchString(codeWithoutStrings) ||
!wpCronHasSettingsRequire(code, codeWithoutStrings) {
return false, nil
}
docroot := filepath.Dir(configPath)
for _, name := range []string{"wp-settings.php", "wp-cron.php", "wp-load.php"} {
info, err := osFS.Lstat(filepath.Join(docroot, name))
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return false, nil
}
return false, err
}
if !info.Mode().IsRegular() {
return false, nil
}
}
for _, name := range []string{"wp-admin", "wp-includes"} {
info, err := osFS.Lstat(filepath.Join(docroot, name))
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return false, nil
}
return false, err
}
if !info.IsDir() {
return false, nil
}
}
return true, nil
}
// wpCronHasSettingsRequire ties the wp-settings.php literal to an active
// require statement. Looking for both tokens independently lets quoted sample
// text masquerade as a WordPress bootstrap and can trigger destructive fixes.
func wpCronHasSettingsRequire(code, codeWithoutStrings string) bool {
for _, loc := range wpCronRequireRe.FindAllStringIndex(codeWithoutStrings, -1) {
statementEnd := strings.IndexByte(codeWithoutStrings[loc[1]:], ';')
if statementEnd < 0 {
continue
}
statementEnd += loc[1]
if strings.Contains(strings.ToLower(code[loc[0]:statementEnd]), "wp-settings.php") {
return true
}
}
return false
}
func sortedWPCronCandidates(candidates map[string]wpCronCandidate) []wpCronCandidate {
out := make([]wpCronCandidate, 0, len(candidates))
for _, candidate := range candidates {
out = append(out, candidate)
}
sort.Slice(out, func(i, j int) bool { return out[i].path < out[j].path })
return out
}
func newWPCronFinding(candidate wpCronCandidate) alert.Finding {
return alert.Finding{
Severity: alert.Warning,
Check: "perf_wp_cron",
Message: fmt.Sprintf("WP-Cron not disabled for %s", candidate.account),
Details: fmt.Sprintf(
"File: %s - add define('DISABLE_WP_CRON', true); and use a real cron job instead",
candidate.path,
),
Timestamp: time.Now(),
}
}
// CheckWPCron scans configured web roots plus validated cPanel document roots
// for WordPress installs that have not disabled the built-in WP-Cron
// mechanism. Running WP-Cron via HTTP is a common cause of high load on busy
// sites.
// The runner enforces a 60-minute throttle via checkThrottleMin.
func CheckWPCron(ctx context.Context, cfg *config.Config, scanState *state.Store) []alert.Finding {
if !perfEnabled(cfg) {
return nil
}
roots, rootsComplete := wpCronScanRootSet(cfg)
if !rootsComplete {
markCheckIncomplete(ctx, "perf_wp_cron")
}
if len(roots) == 0 {
clearWPCronCursor(scanState)
return nil
}
allRoots := make(map[string]struct{}, len(roots))
for _, root := range roots {
allRoots[root.path] = struct{}{}
}
cursor := loadWPCronCursor(scanState)
start, resumeCurrent := wpCronStartRoot(roots, cursor)
var findings []alert.Finding
var nextCursor wpCronScanCursor
capped := false
for visited := 0; visited < len(roots); visited++ {
if ctx.Err() != nil {
return findings
}
root := roots[(start+visited)%len(roots)]
candidates := make(map[string]wpCronCandidate)
if !scanWPCronCandidates(ctx, root.path, root.path, root.account, 2, allRoots, candidates) {
markCheckIncomplete(ctx, "perf_wp_cron")
}
sorted := sortedWPCronCandidates(candidates)
lastPath := ""
if visited == 0 && resumeCurrent && root.path == cursor.Root {
lastPath = cursor.Path
}
first := sort.Search(len(sorted), func(i int) bool { return sorted[i].path > lastPath })
eligible := sorted[first:]
remaining := wpCronFindingLimit - len(findings)
selected := len(eligible)
if selected > remaining {
selected = remaining
}
for _, candidate := range eligible[:selected] {
findings = append(findings, newWPCronFinding(candidate))
}
if len(findings) == wpCronFindingLimit {
nextCursor.Root = root.path
if selected < len(eligible) {
nextCursor.Path = eligible[selected-1].path
}
capped = true
break
}
}
if ctx.Err() != nil {
return findings
}
if capped {
storeWPCronCursor(scanState, nextCursor)
} else {
clearWPCronCursor(scanState)
}
sort.Slice(findings, func(i, j int) bool { return findings[i].Details < findings[j].Details })
return findings
}
// wpCronScanRoots lists the document roots to search for WordPress installs.
//
// ResolveWebRoots alone resolves to /home/*/public_html on cPanel, which misses
// every addon domain: those are served from /home/<user>/<domain>/ and never
// appear beneath public_html. On a live host that hid 177 of 277 installs from
// this check, so the fix could never reach them. cPanel's own domain map is the
// authoritative list, and is already how the exposed-files and php-config
// checks enumerate document roots.
//
// Roots are de-duplicated and path-sorted. Nested roots stay in the list so
// each gets its own full depth allowance; the scanner assigns the nested
// subtree to the more-specific root to avoid duplicate work and findings.
func wpCronScanRoots(cfg *config.Config) []string {
rootSet, _ := wpCronScanRootSet(cfg)
roots := make([]string, 0, len(rootSet))
for _, root := range rootSet {
roots = append(roots, root.path)
}
return roots
}
// ResolveWPCronRoots returns the same validated roots used by CheckWPCron.
// Remediation callers use it so a finding from an addon domain remains inside
// the exact root authorized by cPanel's map.
func ResolveWPCronRoots(cfg *config.Config) []string {
return wpCronScanRoots(cfg)
}
func wpCronScanRootSet(cfg *config.Config) ([]wpCronScanRoot, bool) {
configured := ResolveWebRoots(cfg)
byPath := make(map[string]wpCronScanRoot, len(configured))
for _, root := range configured {
clean := filepath.Clean(root)
byPath[clean] = wpCronScanRoot{path: clean, account: accountFromPath(clean)}
}
complete := true
data, err := osFS.ReadFile(userdataDomainsPath)
if err != nil {
if vhostMapFailureIsIncomplete(err) {
complete = false
}
return sortedWPCronRoots(byPath), complete
}
vhosts, parsedComplete := parseUserdataDomainRootsChecked(string(data))
complete = complete && parsedComplete
attempted := make(map[string]struct{}, len(vhosts))
for _, vhost := range vhosts {
clean := filepath.Clean(vhost.docroot)
key := vhost.user + "\x00" + clean
if _, duplicate := attempted[key]; duplicate {
continue
}
attempted[key] = struct{}{}
homeBase, accountHome, safe := wpCronMapAccountHome(vhost.user, clean, configured)
if !safe {
complete = false
continue
}
exists, pathComplete := wpCronMapRootIsSafe(homeBase, accountHome, clean)
if !pathComplete {
complete = false
}
if !exists {
continue
}
byPath[clean] = wpCronScanRoot{path: clean, account: vhost.user}
}
return sortedWPCronRoots(byPath), complete
}
func sortedWPCronRoots(byPath map[string]wpCronScanRoot) []wpCronScanRoot {
roots := make([]wpCronScanRoot, 0, len(byPath))
for _, root := range byPath {
roots = append(roots, root)
}
sort.Slice(roots, func(i, j int) bool { return roots[i].path < roots[j].path })
return roots
}
func wpCronMapAccountHome(user, docroot string, configured []string) (homeBase, accountHome string, ok bool) {
clean := filepath.Clean(docroot)
if !filepath.IsAbs(clean) {
return "", "", false
}
parts := strings.Split(strings.TrimPrefix(clean, string(filepath.Separator)), string(filepath.Separator))
if len(parts) >= 3 && wpCronHomeVolume(parts[0]) && parts[1] == user {
homeBase = filepath.Join(string(filepath.Separator), parts[0])
accountHome = filepath.Join(homeBase, user)
return homeBase, accountHome, clean != accountHome
}
// Explicit account_roots may point into a chroot-style test or custom
// mount. Only extend a map root to a sibling when the configured root
// already anchors that same /home/USER boundary.
for _, configuredRoot := range configured {
rootParts := strings.Split(strings.TrimPrefix(filepath.Clean(configuredRoot), string(filepath.Separator)), string(filepath.Separator))
for i := 0; i+1 < len(rootParts); i++ {
if rootParts[i] != "home" || rootParts[i+1] != user {
continue
}
homeBaseParts := append([]string{string(filepath.Separator)}, rootParts[:i+1]...)
homeBase = filepath.Join(homeBaseParts...)
accountHome = filepath.Join(homeBase, user)
if filepath.Clean(configuredRoot) != accountHome &&
isPathWithinOrEqual(configuredRoot, accountHome) &&
clean != accountHome && isPathWithinOrEqual(clean, accountHome) {
return homeBase, accountHome, true
}
}
}
return "", "", false
}
func wpCronHomeVolume(name string) bool {
if name == "home" {
return true
}
if !strings.HasPrefix(name, "home") || len(name) == len("home") {
return false
}
for _, r := range name[len("home"):] {
if r < '0' || r > '9' {
return false
}
}
return true
}
// wpCronMapRootIsSafe rejects symlinks at every component from the home mount
// through the docroot. Stat on only the final path would follow an intermediate
// symlink and let a malformed map redirect the scan outside the account.
func wpCronMapRootIsSafe(homeBase, accountHome, root string) (exists, complete bool) {
if !isPathWithinOrEqual(accountHome, homeBase) ||
!isPathWithinOrEqual(root, accountHome) || root == accountHome {
return false, false
}
rel, err := filepath.Rel(homeBase, root)
if err != nil || rel == "." || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return false, false
}
paths := []string{homeBase}
current := homeBase
for _, part := range strings.Split(rel, string(filepath.Separator)) {
current = filepath.Join(current, part)
paths = append(paths, current)
}
for _, path := range paths {
info, statErr := osFS.Lstat(path)
if statErr != nil {
if errors.Is(statErr, fs.ErrNotExist) {
return false, true
}
return false, false
}
if info.Mode()&os.ModeSymlink != 0 || !info.IsDir() {
return false, false
}
}
return true, true
}
func loadWPCronCursor(scanState *state.Store) wpCronScanCursor {
if scanState == nil {
return wpCronScanCursor{}
}
raw, ok := scanState.GetRaw(wpCronCursorKey)
if !ok {
return wpCronScanCursor{}
}
var cursor wpCronScanCursor
if json.Unmarshal([]byte(raw), &cursor) != nil {
return wpCronScanCursor{}
}
return cursor
}
func wpCronStartRoot(roots []wpCronScanRoot, cursor wpCronScanCursor) (start int, resume bool) {
if cursor.Root == "" {
return 0, false
}
if cursor.Path != "" {
i := sort.Search(len(roots), func(i int) bool { return roots[i].path >= cursor.Root })
if i < len(roots) && roots[i].path == cursor.Root {
return i, true
}
}
i := sort.Search(len(roots), func(i int) bool { return roots[i].path > cursor.Root })
if i == len(roots) {
i = 0
}
return i, false
}
func storeWPCronCursor(scanState *state.Store, cursor wpCronScanCursor) {
if scanState == nil {
return
}
raw, err := json.Marshal(cursor)
if err != nil {
return
}
if err := scanState.SetRawAndSave(wpCronCursorKey, string(raw)); err != nil {
fmt.Fprintf(os.Stderr, "wpcron: cursor write: %v\n", err)
}
}
func clearWPCronCursor(scanState *state.Store) {
if scanState == nil {
return
}
if err := scanState.DeleteRawAndSave(wpCronCursorKey); err != nil {
fmt.Fprintf(os.Stderr, "wpcron: cursor clear: %v\n", err)
}
}
package checks
import (
"archive/zip"
"context"
"encoding/binary"
"fmt"
"io"
"path/filepath"
"regexp"
"strings"
"unicode"
"unicode/utf16"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"golang.org/x/net/html"
)
// phishingReadSize is how much of an accepted HTML page is analysed. It
// matches the 100 KB acceptance ceiling: reading only the first 16 KB let a
// kit that opens with a large inline stylesheet keep its form past the read
// window and pass as clean.
const phishingReadSize = 100000
// phishingScanMaxDepth bounds how deep CheckPhishing recurses below each doc
// root. Real kits land in date-nested WordPress upload folders
// (wp-content/uploads/YYYY/MM/<kit>/), six directory levels below the root, so
// the budget must clear that. Heavy/transient dirs (node_modules, vendor, WP
// core, caches) are pruned by isKnownSafeDir before recursion, keeping the
// deeper walk affordable.
const phishingScanMaxDepth = 8
// ---------------------------------------------------------------------------
// Brand impersonation patterns
// ---------------------------------------------------------------------------
// phishingGenericBrandScoreFloor is the score a page matching only the
// "Generic Login" pseudo-brand must reach before it is flagged. The generic
// title patterns ("Sign In", "Secure Access") are ubiquitous on legitimate
// login pages and, unlike a real brand, are not themselves evidence of
// impersonation. A single equally-ubiquitous JS token (window.location,
// fetch) would otherwise clear the normal brand floor and mislabel a
// customer's own login page. Requiring this higher floor forces multiple
// independent signals (external exfil, trust badge, urgency, server-side
// credential capture) that a plain login page does not carry.
const phishingGenericBrandScoreFloor = 7
var phishingBrands = []struct {
name string
titlePatterns []string
bodyPatterns []string
generic bool
}{
{
name: "Microsoft/SharePoint",
titlePatterns: []string{"sharepoint", "onedrive", "microsoft 365", "outlook web", "office 365", "ms online"},
bodyPatterns: []string{"sharepoint", "onedrive", "secured by microsoft", "microsoft corporation"},
},
{
name: "Google",
titlePatterns: []string{"google drive", "google docs", "google sign", "gmail", "google workspace"},
bodyPatterns: []string{"google drive", "google docs", "accounts.google", "secured by google"},
},
{
name: "Dropbox",
titlePatterns: []string{"dropbox", "shared file", "shared folder"},
bodyPatterns: []string{"dropbox", "dropbox.com", "secured by dropbox"},
},
{
name: "DocuSign",
titlePatterns: []string{"docusign", "document signing", "e-signature"},
bodyPatterns: []string{"docusign", "please review and sign", "e-signature"},
},
{
name: "Adobe",
titlePatterns: []string{"adobe sign", "adobe document", "adobe acrobat"},
bodyPatterns: []string{"adobe sign", "adobe document cloud", "secured by adobe"},
},
{
name: "WeTransfer",
titlePatterns: []string{"wetransfer", "file transfer"},
bodyPatterns: []string{"wetransfer", "download your files"},
},
{
name: "Apple/iCloud",
titlePatterns: []string{"icloud", "apple id", "find my"},
bodyPatterns: []string{"icloud.com", "apple id", "secured by apple"},
},
{
name: "PayPal",
titlePatterns: []string{"paypal", "pay pal"},
bodyPatterns: []string{"paypal.com", "secured by paypal"},
},
{
name: "Webmail/Roundcube",
titlePatterns: []string{"roundcube", "horde", "webmail login", "webmail ::", "squirrelmail", "zimbra"},
bodyPatterns: []string{"roundcube webmail", "horde login", "zimbra web client", "webmail login"},
// Note: bare "squirrelmail"/"roundcube" removed from body - sites legitimately
// link to their server's webmail (e.g. href="squirrelmail/index.php").
},
{
name: "cPanel/WHM",
titlePatterns: []string{"cpanel", "whm login", "webhost manager"},
bodyPatterns: []string{"cpanel login", "whm login", "webhost manager"},
},
{
name: "Banking/Financial",
titlePatterns: []string{"online banking", "bank login", "secure banking", "account login"},
bodyPatterns: []string{"online banking", "bank account", "transaction verification"},
},
{
name: "Generic Login",
titlePatterns: []string{"secure access", "verify your", "confirm your identity", "account verification", "email verification", "sign in"},
bodyPatterns: []string{"verify your identity", "confirm your account", "unusual activity"},
generic: true,
},
}
// ---------------------------------------------------------------------------
// Content-based indicators
// ---------------------------------------------------------------------------
// Credential harvesting patterns in page body.
var harvestIndicators = []string{
"window.location.href",
"window.location.replace",
"window.location =",
"document.location.href",
"form.submit()",
".workers.dev",
"confirm access",
"verify your email",
"confirm your email",
"verify identity",
"continue to document",
"access confirmed, redirecting",
"secured by microsoft",
"secured by google",
"secured by apple",
"256-bit encrypted",
"256‑bit encrypted",
// fetch/XHR exfiltration - silent credential POST without redirect
"fetch(",
"xmlhttprequest",
"$.ajax(",
"$.post(",
"navigator.sendbeacon(",
}
// Redirect/exfiltration URL patterns.
var exfilPatterns = []string{
".workers.dev",
"//t.co/",
"/redir?",
"/redirect?",
"effi.redir",
"link?url=",
"goto_url=",
"//bit.ly/",
"//tinyurl.com/",
"//rb.gy/",
"//is.gd/",
"/servlet/effi.redir",
}
// Fake trust badge patterns - security claims in pages not on the brand's domain.
var trustBadgePatterns = []string{
"secured by microsoft",
"secured by google",
"secured by apple",
"secured by dropbox",
"secured by adobe",
"verified by microsoft",
"protected by microsoft",
"256-bit encrypted",
"256‑bit encrypted",
"ssl secured",
"bank-level encryption",
"enterprise security",
}
// Urgency language used to pressure victims.
var urgencyPatterns = []string{
"expires in",
"temporary hold",
"limited time",
"unusual activity detected",
"suspicious activity",
"your account will be",
"verify within",
"action required",
"immediate action",
"account suspended",
"access will be revoked",
}
// Embedded asset indicators - phishing kits embed logos to avoid external loading.
var embeddedAssetPatterns = []string{
"data:image/png;base64,",
"data:image/svg+xml;base64,",
"data:image/jpeg;base64,",
}
// ---------------------------------------------------------------------------
// Main check entry point
// ---------------------------------------------------------------------------
// CheckPhishing scans HTML files in user document roots for phishing pages.
// Uses three detection layers:
// 1. Content analysis - brand impersonation + credential harvesting patterns
// 2. Structural analysis - self-contained HTML with embedded assets
// 3. Directory anomaly - lone HTML files in otherwise empty directories
func CheckPhishing(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
homeDirs := scanHomeDirsWithCoverage(ctx, "phishing")
for _, homeEntry := range homeDirs {
if ctx.Err() != nil {
return findings
}
if !homeEntry.IsDir() {
continue
}
user := homeEntry.Name()
if strings.HasPrefix(user, ".") || user == "virtfs" {
continue
}
homeDir := scanHomeDirPath(homeEntry)
docRoots := []string{filepath.Join(homeDir, "public_html")}
subDirs, err := osFS.ReadDir(homeDir)
markScanReadError(ctx, "phishing", err)
for _, sd := range subDirs {
if sd.IsDir() && sd.Name() != "public_html" && sd.Name() != "mail" &&
!strings.HasPrefix(sd.Name(), ".") && sd.Name() != "etc" &&
sd.Name() != "logs" && sd.Name() != "ssl" && sd.Name() != "tmp" {
docRoots = append(docRoots, filepath.Join(homeDir, sd.Name()))
}
}
for _, docRoot := range docRoots {
scanForPhishing(ctx, docRoot, phishingScanMaxDepth, user, cfg, &findings)
if ctx.Err() != nil {
return findings
}
}
}
return findings
}
// ---------------------------------------------------------------------------
// Directory scanner
// ---------------------------------------------------------------------------
func scanForPhishing(ctx context.Context, dir string, maxDepth int, user string, cfg *config.Config, findings *[]alert.Finding) {
if ctx.Err() != nil {
return
}
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
markScanReadError(ctx, "phishing", err)
return
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
// Bypassed for explicit full-scan / audit requests.
suppressed := false
if scanRespectsIgnores(ctx, cfg) {
for _, ignore := range cfg.Suppressions.IgnorePaths {
if matchGlob(fullPath, ignore) {
suppressed = true
break
}
}
}
if suppressed {
continue
}
if entry.IsDir() {
if isKnownSafeDir(name) {
continue
}
// --- Directory anomaly detection ---
dirResult := analyzeDirectoryStructure(ctx, fullPath, user)
if dirResult != nil {
*findings = append(*findings, *dirResult)
}
scanForPhishing(ctx, fullPath, maxDepth-1, user, cfg, findings)
continue
}
nameLower := strings.ToLower(name)
info, err := entry.Info()
if err != nil {
markScanReadError(ctx, "phishing", err)
continue
}
size := info.Size()
// --- HTML/HTM phishing pages ---
if strings.HasSuffix(nameLower, ".html") || strings.HasSuffix(nameLower, ".htm") {
// Standard phishing page check (3KB-100KB)
if size >= 3000 && size <= 100000 {
result := analyzeHTMLForPhishing(ctx, fullPath)
if result != nil {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "phishing_page",
Message: fmt.Sprintf("Phishing page detected (%s impersonation): %s", result.brand, fullPath),
Details: fmt.Sprintf("Account: %s\nBrand: %s\nScore: %d/10\nIndicators:\n- %s\nSize: %d bytes",
user, result.brand, result.score, strings.Join(result.indicators, "\n- "), size),
FilePath: fullPath,
})
}
}
// --- iframe phishing (tiny HTML files that embed external phishing) ---
if size > 0 && size < 3000 {
if result := checkIframePhishing(ctx, fullPath); result != "" {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "phishing_iframe",
Message: fmt.Sprintf("Iframe phishing page detected: %s", fullPath),
Details: fmt.Sprintf("Account: %s\n%s", user, result),
FilePath: fullPath,
})
}
}
continue
}
// --- PHP phishing pages and open redirectors ---
if isExecutablePHPName(nameLower) {
// Skip known CMS files
if isKnownCMSFile(nameLower) {
continue
}
// PHP phishing (3KB-100KB) - same brand/content analysis as HTML
if size >= 3000 && size <= 100000 {
result := analyzePHPForPhishing(ctx, fullPath)
if result != nil {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "phishing_php",
Message: fmt.Sprintf("PHP phishing page detected (%s): %s", result.brand, fullPath),
Details: fmt.Sprintf("Account: %s\nBrand: %s\nScore: %d/10\nIndicators:\n- %s\nSize: %d bytes",
user, result.brand, result.score, strings.Join(result.indicators, "\n- "), size),
FilePath: fullPath,
})
}
}
// PHP open redirector (tiny PHP files under 1KB)
if size > 0 && size < 1024 {
if result := checkPHPRedirector(ctx, fullPath); result != "" {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "phishing_redirector",
Message: fmt.Sprintf("PHP open redirector detected: %s", fullPath),
Details: fmt.Sprintf("Account: %s\n%s", user, result),
FilePath: fullPath,
})
}
}
continue
}
// --- Credential log files ---
if !strings.HasSuffix(nameLower, ".zip") && isCredentialLogName(nameLower) &&
size > 0 && size < 10*1024*1024 {
if result := checkCredentialLog(ctx, fullPath); result != "" {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "phishing_credential_log",
Message: fmt.Sprintf("Harvested credential log file detected: %s", fullPath),
Details: fmt.Sprintf("Account: %s\n%s", user, result),
FilePath: fullPath,
})
}
continue
}
// --- Phishing kit ZIP archives ---
if strings.HasSuffix(nameLower, ".zip") && size > 1000 && size < 50*1024*1024 {
if isPhishingKitZipName(nameLower) && zipLooksLikeKit(ctx, fullPath) {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "phishing_kit_archive",
Message: fmt.Sprintf("Suspected phishing kit archive: %s", fullPath),
Details: fmt.Sprintf("Account: %s\nFilename: %s\nSize: %d bytes",
user, name, size),
FilePath: fullPath,
})
}
}
}
}
// ---------------------------------------------------------------------------
// Layer 1: Content analysis
// ---------------------------------------------------------------------------
type phishingResult struct {
brand string
score int
indicators []string
}
func analyzeHTMLForPhishing(ctx context.Context, path string) *phishingResult {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return nil
}
defer func() { _ = f.Close() }()
buf := make([]byte, phishingReadSize)
n, err := io.ReadFull(f, buf)
if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF {
markCheckIncomplete(ctx, "phishing")
}
if n == 0 {
return nil
}
content := string(buf[:n])
contentLower := strings.ToLower(content)
// Must contain a form or input
hasForm := strings.Contains(contentLower, "<form") ||
strings.Contains(contentLower, "<input")
if !hasForm {
return nil
}
// Must contain email/password input
hasCredentialInput := hasHTMLCredentialInput(contentLower)
if !hasCredentialInput {
return nil
}
var indicators []string
score := 0
// --- Brand impersonation ---
brandMatch := ""
matchedGeneric := false
titleContent := extractTitle(contentLower)
for _, brand := range phishingBrands {
titleHit := false
bodyHit := false
for _, tp := range brand.titlePatterns {
if strings.Contains(titleContent, tp) {
titleHit = true
indicators = append(indicators, fmt.Sprintf("title impersonates '%s'", tp))
score += 3
break
}
}
for _, bp := range brand.bodyPatterns {
if strings.Contains(contentLower, bp) {
bodyHit = true
if !titleHit {
indicators = append(indicators, fmt.Sprintf("body impersonates '%s'", bp))
score += 2
}
break
}
}
if titleHit || bodyHit {
brandMatch = brand.name
matchedGeneric = brand.generic
break
}
}
// --- Credential harvesting / redirect indicators ---
for _, pattern := range harvestIndicators {
if strings.Contains(contentLower, pattern) {
score++
indicators = append(indicators, fmt.Sprintf("harvest: '%s'", pattern))
}
}
// --- Exfiltration URL patterns ---
for _, pattern := range exfilPatterns {
if strings.Contains(contentLower, pattern) {
score += 2
indicators = append(indicators, fmt.Sprintf("exfiltration: '%s'", pattern))
}
}
// --- Fake trust badges ---
for _, pattern := range trustBadgePatterns {
if strings.Contains(contentLower, pattern) {
score++
indicators = append(indicators, fmt.Sprintf("fake trust badge: '%s'", pattern))
}
}
// --- Urgency language ---
urgencyCount := 0
for _, pattern := range urgencyPatterns {
if strings.Contains(contentLower, pattern) {
urgencyCount++
}
}
if urgencyCount > 0 {
score++
indicators = append(indicators, fmt.Sprintf("urgency language (%d patterns)", urgencyCount))
}
// --- Embedded Base64 assets (logos embedded to avoid external loading) ---
embeddedCount := 0
for _, pattern := range embeddedAssetPatterns {
embeddedCount += countOccurrences(contentLower, pattern)
}
if embeddedCount > 0 {
score++
indicators = append(indicators, fmt.Sprintf("embedded base64 assets (%d)", embeddedCount))
}
// --- Form action pointing to external domain ---
if hasExternalFormAction(content) {
score += 2
indicators = append(indicators, "form action points to external domain")
}
// --- Self-contained page (all CSS inline, no external stylesheets except CDN) ---
if isSelfContainedHTML(contentLower) {
score++
indicators = append(indicators, "self-contained HTML (all styles inline)")
}
// --- Person name as filename ---
baseName := strings.TrimSuffix(strings.TrimSuffix(filepath.Base(path), ".html"), ".htm")
if looksLikePersonName(baseName) {
score += 2
indicators = append(indicators, fmt.Sprintf("filename looks like person name: '%s'", baseName))
}
// --- Decision ---
// Real brand match: need score >= 4 (brand gives 2-3 + at least 1 other
// signal). Generic pseudo-brand: need a higher floor since its title
// patterns are ubiquitous on legitimate login pages. No brand: need
// score >= 6 (multiple strong signals).
if brandMatch != "" {
floor := 4
if matchedGeneric {
floor = phishingGenericBrandScoreFloor
}
if score >= floor {
return &phishingResult{brand: brandMatch, score: score, indicators: indicators}
}
}
if brandMatch == "" && score >= 6 {
return &phishingResult{brand: "Unknown", score: score, indicators: indicators}
}
return nil
}
// ---------------------------------------------------------------------------
// Layer 2: Structural analysis helpers
// ---------------------------------------------------------------------------
// extractTitle pulls the <title> content from HTML.
func extractTitle(contentLower string) string {
for offset := 0; offset < len(contentLower); {
idx := strings.Index(contentLower[offset:], "<title")
if idx < 0 {
return ""
}
start := offset + idx
afterName := start + len("<title")
if afterName < len(contentLower) && !isTagBoundary(contentLower[afterName]) {
offset = afterName
continue
}
openEnd := findTagEnd(contentLower, afterName)
if openEnd < 0 {
return ""
}
closeOffset := strings.Index(contentLower[openEnd+1:], "</title>")
if closeOffset < 0 {
return ""
}
return strings.TrimSpace(contentLower[openEnd+1 : openEnd+1+closeOffset])
}
return ""
}
// hasExternalFormAction checks if a <form> action points to a different domain.
func hasExternalFormAction(content string) bool {
for _, url := range htmlAttrValues(strings.ToLower(content), "form", "action", false) {
// External if it starts with http:// or https:// (not relative)
if strings.HasPrefix(url, "http://") || strings.HasPrefix(url, "https://") {
return true
}
}
return false
}
func hasHTMLCredentialInput(contentLower string) bool {
for _, value := range htmlAttrValues(contentLower, "input", "type", true) {
switch strings.TrimSpace(value) {
case "email", "password":
return true
}
}
for _, value := range htmlAttrValues(contentLower, "input", "name", true) {
switch strings.TrimSpace(value) {
case "email", "pass", "password", "login":
return true
}
}
for _, value := range htmlAttrValues(contentLower, "input", "placeholder", true) {
value = strings.TrimSpace(value)
if strings.HasPrefix(value, "email") ||
strings.HasPrefix(value, "you@") ||
strings.HasPrefix(value, "your email") {
return true
}
}
return strings.Contains(contentLower, "work or school email") ||
strings.Contains(contentLower, "corporate email")
}
func htmlAttrValues(contentLower, tagName, attrName string, allowUnquoted bool) []string {
var values []string
needle := "<" + tagName
for offset := 0; offset < len(contentLower); {
idx := strings.Index(contentLower[offset:], needle)
if idx < 0 {
break
}
start := offset + idx
afterName := start + len(needle)
if afterName < len(contentLower) && !isTagBoundary(contentLower[afterName]) {
offset = afterName
continue
}
end := findTagEnd(contentLower, afterName)
if end < 0 {
break
}
if value, ok := tagAttrValue(contentLower[afterName:end], attrName, allowUnquoted); ok {
values = append(values, value)
}
offset = end + 1
}
return values
}
func isTagBoundary(c byte) bool {
return c == '>' || c == '/' || unicode.IsSpace(rune(c))
}
func findTagEnd(content string, start int) int {
var quote byte
for i := start; i < len(content); i++ {
c := content[i]
if quote != 0 {
if c == quote {
quote = 0
}
continue
}
if c == '"' || c == '\'' {
quote = c
continue
}
if c == '>' {
return i
}
}
return -1
}
func tagAttrValue(attrs, attrName string, allowUnquoted bool) (string, bool) {
for i := 0; i < len(attrs); {
for i < len(attrs) && (unicode.IsSpace(rune(attrs[i])) || attrs[i] == '/') {
i++
}
nameStart := i
for i < len(attrs) && attrs[i] != '=' && attrs[i] != '>' &&
attrs[i] != '/' && !unicode.IsSpace(rune(attrs[i])) {
i++
}
if nameStart == i {
i++
continue
}
name := attrs[nameStart:i]
for i < len(attrs) && unicode.IsSpace(rune(attrs[i])) {
i++
}
if i >= len(attrs) || attrs[i] != '=' {
continue
}
i++
for i < len(attrs) && unicode.IsSpace(rune(attrs[i])) {
i++
}
if i >= len(attrs) {
return "", false
}
value := ""
if attrs[i] == '"' || attrs[i] == '\'' {
quote := attrs[i]
i++
valueStart := i
for i < len(attrs) && attrs[i] != quote {
i++
}
value = attrs[valueStart:i]
if i < len(attrs) {
i++
}
} else {
valueStart := i
for i < len(attrs) && attrs[i] != '>' && !unicode.IsSpace(rune(attrs[i])) {
i++
}
if !allowUnquoted {
continue
}
value = attrs[valueStart:i]
}
if name == attrName {
return strings.TrimSpace(value), true
}
}
return "", false
}
// isSelfContainedHTML checks if a page has all its CSS inline (embedded <style> tags)
// with no or minimal external stylesheet references - typical of phishing kits.
func isSelfContainedHTML(contentLower string) bool {
hasInlineStyle := strings.Contains(contentLower, "<style")
externalCSS := countOccurrences(contentLower, "rel=\"stylesheet\"") +
countOccurrences(contentLower, "rel='stylesheet'")
// Allow 1 external CSS (e.g., Font Awesome CDN) - phishing kits often use one
return hasInlineStyle && externalCSS <= 1
}
// looksLikePersonName checks if a filename looks like a person name (CamelCase
// with 2+ capitalized words, e.g., "PalmerHamilton", "MarilynEsguerra").
func looksLikePersonName(name string) bool {
if len(name) < 6 {
return false
}
// Count uppercase transitions (start of words in CamelCase)
upperCount := 0
for i, c := range name {
if unicode.IsUpper(c) {
if i == 0 || unicode.IsLower(rune(name[i-1])) {
upperCount++
}
}
}
// 2+ capitalized words, all letters, no common web words
if upperCount < 2 {
return false
}
allLetters := true
for _, c := range name {
if !unicode.IsLetter(c) {
allLetters = false
break
}
}
if !allLetters {
return false
}
// Exclude common web filenames
nameLower := strings.ToLower(name)
webNames := []string{"index", "default", "portal", "login", "home", "main",
"readme", "changelog", "license", "manifest", "service"}
for _, w := range webNames {
if nameLower == w {
return false
}
}
return true
}
// ---------------------------------------------------------------------------
// Layer 3: Directory anomaly detection
// ---------------------------------------------------------------------------
// analyzeDirectoryStructure checks if a directory looks like a phishing drop:
// - Contains only 1-3 HTML files and nothing else significant
// - Directory name looks like a business/organization name
// - No CMS markers (wp-config, index.php, etc.)
func analyzeDirectoryStructure(ctx context.Context, dir string, user string) *alert.Finding {
entries, err := osFS.ReadDir(dir)
if err != nil {
markScanReadError(ctx, "phishing", err)
return nil
}
var htmlFiles []string
otherFiles := 0
totalFiles := 0
for _, entry := range entries {
if entry.IsDir() {
return nil // Has subdirectories - likely not a simple phishing drop
}
name := entry.Name()
if strings.HasPrefix(name, ".") {
continue // Skip dotfiles (.htaccess etc.)
}
totalFiles++
nameLower := strings.ToLower(name)
if strings.HasSuffix(nameLower, ".html") || strings.HasSuffix(nameLower, ".htm") {
htmlFiles = append(htmlFiles, name)
} else {
otherFiles++
}
}
// Must have exactly 1-3 HTML files and at most 1 other file
if len(htmlFiles) == 0 || len(htmlFiles) > 3 || otherFiles > 1 {
return nil
}
// Directory name should look like a business/organization (CamelCase or multi-word)
dirName := filepath.Base(dir)
if !looksLikeBusinessName(dirName) {
return nil
}
// Verify at least one HTML file has credential inputs (quick check)
hasPhishingContent := false
for _, htmlFile := range htmlFiles {
fullPath := filepath.Join(dir, htmlFile)
if quickPhishingCheck(ctx, fullPath) {
hasPhishingContent = true
break
}
}
if !hasPhishingContent {
return nil
}
// Build indicators
indicators := []string{
fmt.Sprintf("directory '%s' contains only %d HTML file(s)", dirName, len(htmlFiles)),
fmt.Sprintf("directory name resembles business/organization: '%s'", dirName),
}
for _, h := range htmlFiles {
baseName := strings.TrimSuffix(strings.TrimSuffix(h, ".html"), ".htm")
if looksLikePersonName(baseName) {
indicators = append(indicators, fmt.Sprintf("HTML filename looks like person name: '%s'", baseName))
}
}
return &alert.Finding{
Severity: alert.High,
Check: "phishing_directory",
Message: fmt.Sprintf("Suspected phishing directory (lone HTML in business-named folder): %s", dir),
Details: fmt.Sprintf("Account: %s\nDirectory: %s\nHTML files: %s\nIndicators:\n- %s",
user, dirName, strings.Join(htmlFiles, ", "), strings.Join(indicators, "\n- ")),
FilePath: dir,
}
}
// looksLikeBusinessName checks if a directory name looks like a business or
// organization name rather than a standard web directory.
func looksLikeBusinessName(name string) bool {
if len(name) < 5 {
return false
}
nameLower := strings.ToLower(name)
// Skip names that start with tech/dev terms - these are tutorial
// or test directories, not business names (e.g. "php-email-form",
// "PHP-Login", "JavaScript Login")
techPrefixes := []string{
"php", "javascript", "js-", "css", "html", "python",
"java", "node", "react", "vue", "angular", "jquery",
"bootstrap", "wordpress", "wp-", "laravel",
}
for _, prefix := range techPrefixes {
if strings.HasPrefix(nameLower, prefix) {
return false
}
}
// Skip standard web directories
standardDirs := []string{
"images", "img", "css", "js", "fonts", "assets", "static",
"media", "uploads", "downloads", "files", "docs", "data",
"api", "admin", "config", "templates", "scripts", "lib",
"src", "dist", "build", "public", "private", "backup",
"old", "new", "test", "dev", "staging", "demo",
}
for _, sd := range standardDirs {
if nameLower == sd {
return false
}
}
// CamelCase detection (e.g., WashingtonGolf, XRFScientificAmericasInc)
upperTransitions := 0
for i, c := range name {
if unicode.IsUpper(c) && i > 0 && unicode.IsLower(rune(name[i-1])) {
upperTransitions++
}
}
if upperTransitions >= 1 && unicode.IsUpper(rune(name[0])) {
return true
}
// Multi-word with separators (e.g., federated-lighting, northwest_crawlspace)
if strings.ContainsAny(name, "-_") {
parts := strings.FieldsFunc(name, func(r rune) bool { return r == '-' || r == '_' })
if len(parts) >= 2 {
return true
}
}
// Long lowercase name that doesn't match standard dirs (e.g., "healthcornerpediattrics")
allLower := true
for _, c := range name {
if !unicode.IsLetter(c) {
allLower = false
break
}
}
if allLower && len(name) >= 12 {
return true
}
return false
}
// quickPhishingCheck does a fast read of an HTML file and confirms phishing
// shape: a credential-collection form AND at least one phishing-kit signal
// in the page body itself. Used by the directory-anomaly heuristic to
// avoid flagging benign HTML that happens to ship an <input> tag.
//
// A bare "<form> + email/password keyword" gate matches developer demo
// pages, JavaScript login tutorials, contact forms, and password-reset
// stubs. Phishing kits add at least one of:
// - a form action that posts to an external host (exfiltration target),
// - a fully self-contained inline-styled HTML body (kits ship one file),
// - real brand impersonation in the title or visible body (Office/PayPal/etc).
//
// Requiring credential intake plus one of those signals keeps real
// phishing kits in scope while letting tutorials and trivial forms drop
// out without consulting any path-name allowlist.
func quickPhishingCheck(ctx context.Context, path string) bool {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return false
}
defer func() { _ = f.Close() }()
buf := make([]byte, 4096) // first 4KB is enough for the head, form attrs, brand strings
n, err := io.ReadFull(f, buf)
if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF {
markCheckIncomplete(ctx, "phishing")
}
if n == 0 {
return false
}
content := string(buf[:n])
contentLower := strings.ToLower(content)
if !hasHTMLCredentialInput(contentLower) {
return false
}
if hasExternalFormAction(content) {
return true
}
if isSelfContainedHTML(contentLower) {
return true
}
titleContent := extractTitle(contentLower)
for _, brand := range phishingBrands {
if brand.generic {
continue
}
for _, tp := range brand.titlePatterns {
if titleContent != "" && strings.Contains(titleContent, tp) {
return true
}
}
for _, bp := range brand.bodyPatterns {
if strings.Contains(contentLower, bp) {
return true
}
}
}
return false
}
// ---------------------------------------------------------------------------
// Layer 4: PHP phishing pages
// ---------------------------------------------------------------------------
// analyzePHPForPhishing reads a PHP file and checks for embedded HTML with
// brand impersonation. PHP phishing kits often have PHP code at the top
// (credential handling, emailing) and HTML output below.
func analyzePHPForPhishing(ctx context.Context, path string) *phishingResult {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return nil
}
defer func() { _ = f.Close() }()
buf := make([]byte, phishingReadSize)
n, err := io.ReadFull(f, buf)
if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF {
markCheckIncomplete(ctx, "phishing")
}
if n == 0 {
return nil
}
content := string(buf[:n])
contentLower := strings.ToLower(content)
// PHP phishing indicators: credential handling code.
// Only truly specific patterns belong here - generic functions like
// mail() and fwrite() are handled separately with context checks below.
phpPhishingPatterns := []string{
"$_post['email']",
"$_post['password']",
"$_post['pass']",
"$_post[\"email\"]",
"$_post[\"password\"]",
"$_post[\"pass\"]",
"$_request['email']",
"$_request['password']",
}
var indicators []string
score := 0
// Check for PHP credential handling
phpCredHandling := false
for _, pattern := range phpPhishingPatterns {
if strings.Contains(contentLower, pattern) {
phpCredHandling = true
indicators = append(indicators, fmt.Sprintf("PHP credential handling: '%s'", pattern))
score += 2
break
}
}
// Must have either PHP credential handling OR HTML form output
hasForm := strings.Contains(contentLower, "<form") || strings.Contains(contentLower, "<input")
hasCredentialInput := hasHTMLCredentialInput(contentLower)
if !phpCredHandling && (!hasForm || !hasCredentialInput) {
return nil
}
// A title brand is strong evidence on its own. A visible body or logo brand
// is only accepted when the PHP reads a submitted password. Provider names
// in backend backup, payment, and AJAX integration code are not page
// impersonation and must not establish a brand.
brandMatch := ""
matchedGeneric := false
pageMarkup := stripPHPBlocks(contentLower)
titleContent := extractTitle(pageMarkup)
if titleContent != "" {
for _, brand := range phishingBrands {
for _, tp := range brand.titlePatterns {
if strings.Contains(titleContent, tp) {
brandMatch = brand.name
matchedGeneric = brand.generic
indicators = append(indicators, fmt.Sprintf("title impersonates '%s'", tp))
score += 3
break
}
}
if brandMatch != "" {
break
}
}
}
if brandMatch == "" && hasPHPSubmittedPassword(contentLower) {
bodyContent := visiblePageContent(pageMarkup, true)
for _, brand := range phishingBrands {
if brand.generic {
continue
}
for _, bp := range brand.bodyPatterns {
if strings.Contains(bodyContent, bp) {
brandMatch = brand.name
indicators = append(indicators, fmt.Sprintf("body impersonates '%s'", bp))
score += 2
break
}
}
if brandMatch != "" {
break
}
}
}
// Check harvest/exfil patterns
for _, pattern := range harvestIndicators {
if strings.Contains(contentLower, pattern) {
score++
indicators = append(indicators, fmt.Sprintf("harvest: '%s'", pattern))
}
}
for _, pattern := range exfilPatterns {
if strings.Contains(contentLower, pattern) {
score += 2
indicators = append(indicators, fmt.Sprintf("exfiltration: '%s'", pattern))
}
}
// PHP-specific exfil: emailing or writing harvested credentials.
// These only fire when the file also reads from $_POST/$_REQUEST,
// because a phishing kit must capture form data before exfiltrating it.
// Without this gate, any PHP file using fwrite()+config keywords triggers.
hasPostData := strings.Contains(contentLower, "$_post") || strings.Contains(contentLower, "$_request")
if hasPostData && strings.Contains(contentLower, "mail(") &&
(strings.Contains(contentLower, "password") || strings.Contains(contentLower, "email")) {
score += 2
indicators = append(indicators, "PHP mail() with credential data")
}
if hasPostData &&
(strings.Contains(contentLower, "fwrite(") || strings.Contains(contentLower, "file_put_contents(")) {
if strings.Contains(contentLower, "password") || strings.Contains(contentLower, "email") ||
strings.Contains(contentLower, "result") || strings.Contains(contentLower, "log") {
score += 2
indicators = append(indicators, "PHP writes credential data to file")
}
}
// Require brand impersonation to flag as phishing.
// PHP files with $_POST['email'] + mail() are normal (contact forms, CMS user
// admin, gallery software). Without brand impersonation, these are almost
// always legitimate applications.
if brandMatch == "" {
return nil
}
// The generic pseudo-brand ("Sign In" titles) matches a customer's own
// login.php, which legitimately reads $_POST credentials. Require the
// higher floor so a real brand or genuine exfil behaviour is needed.
floor := 4
if matchedGeneric {
floor = phishingGenericBrandScoreFloor
}
if score >= floor {
return &phishingResult{brand: brandMatch, score: score, indicators: indicators}
}
return nil
}
func hasPHPSubmittedPassword(contentLower string) bool {
passwordKeys := map[string]bool{
"password": true,
"pass": true,
"passwd": true,
"pwd": true,
"passcode": true,
}
code := stripPHPCommentsFromCode(phpCodeOnly(contentLower))
for i := 0; i < len(code); {
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label)
continue
}
if isPHPQuote(code[i]) {
i = skipPHPString(code, i) + 1
continue
}
nameLen := 0
switch {
case strings.HasPrefix(code[i:], "$_post"):
nameLen = len("$_post")
case strings.HasPrefix(code[i:], "$_request"):
nameLen = len("$_request")
}
if nameLen == 0 || (i+nameLen < len(code) && isPHPIdentifierPart(code[i+nameLen])) {
i++
continue
}
j := skipPHPWhitespace(code, i+nameLen)
if j >= len(code) || code[j] != '[' {
i += nameLen
continue
}
j = skipPHPWhitespace(code, j+1)
if j >= len(code) || !isPHPQuote(code[j]) {
i += nameLen
continue
}
keyEnd := skipPHPString(code, j)
if keyEnd <= j || keyEnd >= len(code) || code[keyEnd] != code[j] {
return false
}
key := code[j+1 : keyEnd]
j = skipPHPWhitespace(code, keyEnd+1)
if j < len(code) && code[j] == ']' && passwordKeys[key] {
return true
}
i += nameLen
}
return false
}
func stripPHPBlocks(content string) string {
codeOnly := phpCodeOnly(content)
visible := []byte(content)
for i := range visible {
switch codeOnly[i] {
case ' ', '\t', '\n', '\r':
default:
visible[i] = ' '
}
}
return string(visible)
}
// isKnownCMSFile returns true for PHP files that are standard CMS files
// and should not be scanned for phishing (too many false positives).
func isKnownCMSFile(nameLower string) bool {
cmsFiles := map[string]bool{
"index.php": true, "wp-config.php": true, "wp-login.php": true,
"wp-cron.php": true, "wp-settings.php": true, "wp-load.php": true,
"wp-blog-header.php": true, "wp-links-opml.php": true,
"xmlrpc.php": true, "wp-signup.php": true, "wp-activate.php": true,
"wp-trackback.php": true, "wp-comments-post.php": true,
"wp-mail.php": true, "configuration.php": true,
"config.php": true, "settings.php": true,
}
return cmsFiles[nameLower]
}
// ---------------------------------------------------------------------------
// Layer 5: PHP open redirectors
// ---------------------------------------------------------------------------
// checkPHPRedirector reads a small PHP file and checks if it's an open
// redirector - a file that redirects the visitor to a URL from a parameter.
func checkPHPRedirector(ctx context.Context, path string) string {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return ""
}
defer func() { _ = f.Close() }()
buf := make([]byte, 1024)
n, err := io.ReadFull(f, buf)
if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF {
markCheckIncomplete(ctx, "phishing")
}
if n == 0 {
return ""
}
content := strings.ToLower(string(buf[:n]))
hasHeader := strings.Contains(content, "header(") &&
(strings.Contains(content, "location:") || strings.Contains(content, "location :"))
if !hasHeader {
return ""
}
// Pattern 1: user-controlled redirect target - the URL in header() must
// come from user input. Just having $_GET anywhere + header() is too broad;
// normal form handlers use $_POST for data then header() for redirect.
// Only flag when the redirect URL itself is parameterized.
userControlledRedirect := false
redirectPatterns := []string{
"$_get['url']", "$_get[\"url\"]",
"$_get['redirect']", "$_get[\"redirect\"]",
"$_get['r']", "$_get[\"r\"]",
"$_get['return']", "$_get[\"return\"]",
"$_get['next']", "$_get[\"next\"]",
"$_get['goto']", "$_get[\"goto\"]",
"$_get['link']", "$_get[\"link\"]",
"$_request['url']", "$_request[\"url\"]",
"$_request['redirect']", "$_request[\"redirect\"]",
"header(\"location: \".$_get", "header(\"location: \".$_request",
"header('location: '.$_get", "header('location: '.$_request",
"header(\"location:\".$_get", "header('location:'.$_get",
}
for _, p := range redirectPatterns {
if strings.Contains(content, p) {
userControlledRedirect = true
break
}
}
if userControlledRedirect {
return "PHP open redirector: header(Location) with user-supplied URL"
}
// Pattern 2: Hardcoded redirect to suspicious domain
for _, pattern := range exfilPatterns {
if strings.Contains(content, pattern) {
return fmt.Sprintf("PHP redirect to suspicious destination matching '%s'", pattern)
}
}
return ""
}
// ---------------------------------------------------------------------------
// Layer 6: Credential log files
// ---------------------------------------------------------------------------
// isCredentialLogName checks if a filename matches patterns used by phishing
// kits to store harvested credentials.
func isCredentialLogName(nameLower string) bool {
// Exact names commonly used by phishing kits
exactNames := map[string]bool{
"results.txt": true, "result.txt": true, "log.txt": true,
"logs.txt": true, "emails.txt": true, "data.txt": true,
"passwords.txt": true, "creds.txt": true, "credentials.txt": true,
"victims.txt": true, "output.txt": true, "harvested.txt": true,
"results.log": true, "emails.log": true, "data.log": true,
"results.csv": true, "emails.csv": true, "data.csv": true,
"results.html": true,
}
if exactNames[nameLower] {
return true
}
// Pattern: contains "result", "victim", "harvested", "credential" in name
suspiciousWords := []string{"result", "victim", "harvest", "credential", "creds", "stolen"}
for _, word := range suspiciousWords {
if strings.Contains(nameLower, word) {
return true
}
}
return false
}
// emailPattern matches a plausibly-real email address. The old heuristic
// counted any line containing '@', which random binary bytes and stray code
// tokens trip constantly. Requiring a local part, host and TLD keeps the count
// tied to actual addresses.
var (
emailPattern = regexp.MustCompile(`[A-Za-z0-9._%+\-]+@[A-Za-z0-9.\-]+\.[A-Za-z]{2,}`)
emailAddressPattern = regexp.MustCompile(`^[A-Za-z0-9._%+\-]+@[A-Za-z0-9.\-]+\.[A-Za-z]{2,}$`)
)
// credentialLogReadLimit bounds how much of a candidate file is read for
// credential analysis.
const credentialLogReadLimit = 10 * 1024 * 1024
// looksBinary reports whether normalized text bytes still look binary. Images,
// video, fonts, and other files can carry result/harvest in their names and
// incidental email delimiters in their bytes. A NUL or a high share of control
// bytes keeps those files out of the credential parser.
func looksBinary(data []byte) bool {
control := 0
for _, b := range data {
if b == 0x00 {
return true
}
// Bytes >= 0x80 are kept as text so UTF-8 (accented names in an email
// list) is not misread as binary.
if b < 0x09 || (b > 0x0d && b < 0x20) || b == 0x7f {
control++
}
}
return control*10 > len(data)
}
func normalizeCredentialLogText(data []byte) []byte {
if len(data) >= 3 && data[0] == 0xef && data[1] == 0xbb && data[2] == 0xbf {
return data[3:]
}
isLittleEndian := len(data) >= 2 && data[0] == 0xff && data[1] == 0xfe
isBigEndian := len(data) >= 2 && data[0] == 0xfe && data[1] == 0xff
body := data
var order binary.ByteOrder
switch {
case isLittleEndian:
body = data[2:]
order = binary.LittleEndian
case isBigEndian:
body = data[2:]
order = binary.BigEndian
default:
pairs := len(data) / 2
if pairs < 4 {
return data
}
evenNUL, oddNUL := 0, 0
for i := 0; i+1 < len(data); i += 2 {
if data[i] == 0 {
evenNUL++
}
if data[i+1] == 0 {
oddNUL++
}
}
switch {
case oddNUL >= 4 && oddNUL*2 >= pairs && evenNUL*20 <= pairs:
order = binary.LittleEndian
case evenNUL >= 4 && evenNUL*2 >= pairs && oddNUL*20 <= pairs:
order = binary.BigEndian
default:
return data
}
}
units := make([]uint16, len(body)/2)
for i := range units {
units[i] = order.Uint16(body[i*2:])
}
return []byte(string(utf16.Decode(units)))
}
// checkCredentialLog reads a text file and checks if it contains harvested
// credentials (email:password pairs, one per line) or a harvested address list.
func checkCredentialLog(ctx context.Context, path string) string {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return ""
}
defer func() { _ = f.Close() }()
// Reject binaries from the head before reading the whole file into memory:
// images, video, and fonts that carry a "result"/"harvest" keyword in their
// name would otherwise be slurped in full only to be discarded. The head is
// decoded first so a UTF-16 dump, whose raw bytes look binary, survives the
// peek.
head := make([]byte, 8192)
hn, err := io.ReadFull(f, head)
if err != nil && err != io.ErrUnexpectedEOF && err != io.EOF {
markCheckIncomplete(ctx, "phishing")
return ""
}
head = head[:hn]
if hn == 0 || looksBinary(normalizeCredentialLogText(head)) {
return ""
}
rest, err := io.ReadAll(io.LimitReader(f, credentialLogReadLimit+1-int64(hn)))
if err != nil {
markScanReadError(ctx, "phishing", err)
return ""
}
return analyzeCredentialLog(append(head, rest...), path)
}
// analyzeCredentialLog classifies bounded, already-read file content as a
// harvested credential dump (email:password pairs) or address list, returning
// the finding detail or an empty string.
func analyzeCredentialLog(data []byte, path string) string {
if len(data) == 0 || len(data) > credentialLogReadLimit {
return ""
}
data = normalizeCredentialLogText(data)
if looksBinary(data) {
return ""
}
lines := strings.Split(string(data), "\n")
credentialLines := 0
emailLines := 0
nonEmpty := 0
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" {
continue
}
nonEmpty++
if !emailPattern.MatchString(line) {
continue
}
emailLines++
// Pattern: email:password or email|password or email,password, where the
// field before the delimiter holds the address and the field after is a
// non-empty secret token with no whitespace.
for _, delim := range []string{":", "|", "\t", ","} {
parts := strings.SplitN(line, delim, 3)
if len(parts) >= 2 {
secret := strings.TrimSpace(parts[1])
address := strings.TrimSpace(parts[0])
if emailAddressPattern.MatchString(address) && secret != "" && !strings.ContainsAny(secret, " \t") {
credentialLines++
break
}
}
}
}
// 3+ lines that look like email:password pairs = credential log
if credentialLines >= 3 {
return fmt.Sprintf("File contains %d credential-like lines (email:password format) out of %d email lines",
credentialLines, emailLines)
}
// A harvested address list is dominated by addresses (about one per line).
// Source files (JavaScript modules, Drupal handlers) legitimately embed a
// handful of contributor/support addresses among mostly-code lines; those
// are not harvested lists, so require the addresses to be the majority of
// non-empty lines. .csv exports of a contact list are excluded outright.
if emailLines >= 10 && emailLines*2 >= nonEmpty && !strings.HasSuffix(strings.ToLower(path), ".csv") {
return fmt.Sprintf("File contains %d email addresses - possible harvested email list", emailLines)
}
return ""
}
// ---------------------------------------------------------------------------
// Layer 7: Iframe phishing
// ---------------------------------------------------------------------------
// checkIframePhishing checks small HTML files for iframe-based phishing -
// a minimal HTML page that just loads an external phishing page in a full-screen iframe.
func checkIframePhishing(ctx context.Context, path string) string {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return ""
}
defer func() { _ = f.Close() }()
buf := make([]byte, 3000)
n, err := io.ReadFull(f, buf)
if err != nil && err != io.EOF && err != io.ErrUnexpectedEOF {
markCheckIncomplete(ctx, "phishing")
}
if n == 0 {
return ""
}
contentLower := strings.ToLower(string(buf[:n]))
// Must contain an iframe
if !strings.Contains(contentLower, "<iframe") {
return ""
}
// Check if iframe src points to an external URL
idx := strings.Index(contentLower, "<iframe")
if idx < 0 {
return ""
}
rest := contentLower[idx:]
endTag := strings.Index(rest, ">")
if endTag < 0 {
return ""
}
iframeTag := rest[:endTag+1]
// Extract src
srcIdx := strings.Index(iframeTag, "src=")
if srcIdx < 0 {
return ""
}
srcRest := iframeTag[srcIdx+4:]
if len(srcRest) == 0 {
return ""
}
quote := srcRest[0]
if quote != '"' && quote != '\'' {
return ""
}
srcEnd := strings.IndexByte(srcRest[1:], quote)
if srcEnd < 0 {
return ""
}
src := srcRest[1 : srcEnd+1]
// Must be external (http:// or https://)
if !strings.HasPrefix(src, "http://") && !strings.HasPrefix(src, "https://") {
return ""
}
// Check if iframe is fullscreen (width/height 100% or style covers viewport)
isFullscreen := strings.Contains(iframeTag, "100%") ||
strings.Contains(contentLower, "width:100%") ||
strings.Contains(contentLower, "width: 100%") ||
strings.Contains(contentLower, "position:fixed") ||
strings.Contains(contentLower, "position: fixed")
// A phishing wrapper is essentially just the iframe - its only purpose is to
// fill the screen with the external page. A documented embed or demo (e.g.
// software shipping an "iframe-example.html" with an explanatory paragraph)
// carries prose around the iframe, so real visible text means this is not a
// bare redirect wrapper.
if isFullscreen && visibleTextLen(contentLower) <= 40 {
return fmt.Sprintf("Full-screen iframe loading external URL: %s", src)
}
// Even non-fullscreen, check if URL matches known phishing/exfil patterns
for _, pattern := range exfilPatterns {
if strings.Contains(src, pattern) {
return fmt.Sprintf("Iframe loading suspicious external URL matching '%s': %s", pattern, src)
}
}
return ""
}
// ---------------------------------------------------------------------------
// Layer 8: Phishing kit ZIP archives
// ---------------------------------------------------------------------------
// isPhishingKitZipName checks if a ZIP filename matches common phishing kit names.
// Requires 2+ keyword matches to reduce false positives (e.g. "CssCheckboxKit"
// matched "kit" alone, but legitimate UI kits, CSS kits, etc. are common).
func isPhishingKitZipName(nameLower string) bool {
// High-confidence single-match keywords (brand impersonation in filename)
singleMatch := []string{
"office365", "office 365", "sharepoint", "onedrive",
"microsoft", "outlook", "gmail",
"dropbox", "docusign", "wetransfer",
"paypal", "icloud", "netflix",
"facebook", "instagram", "linkedin",
"roundcube", "cpanel",
"phish", "scam",
}
for _, kw := range singleMatch {
if strings.Contains(nameLower, kw) {
return true
}
}
// Lower-confidence keywords - require 2+ matches to flag.
// Words like "login", "verify", "secure", "google", "apple", "bank"
// appear in legitimate archives too.
multiMatch := []string{
"login", "verify", "secure", "bank",
"google", "apple", "adobe", "webmail",
}
matches := 0
for _, kw := range multiMatch {
if strings.Contains(nameLower, kw) {
matches++
}
}
return matches >= 2
}
// kitCaptureScripts are filenames phishing kits use for the server-side
// credential-capture step. These names are kit idiom, not generic app files.
var kitCaptureScripts = map[string]bool{
"next.php": true, "post.php": true, "send.php": true, "grab.php": true,
"result.php": true, "results.php": true,
}
var kitAntibotScripts = map[string]bool{
"antibots": true, "blocker": true,
"antibot.php": true, "blocker.php": true, "bots.php": true, "killbot.php": true,
}
// kitLoginPageWords identify login or verification pages inside an archive.
// The archive prefilter already supplies the brand signal, so repeating brand
// names here would misclassify ordinary provider plugins as login pages.
var kitLoginPageWords = []string{
"signin", "sign-in", "verify", "login", "log-in", "secure",
}
func isKitCredentialSinkName(base string) bool {
switch strings.ToLower(filepath.Ext(base)) {
case ".txt", ".log", ".csv":
return isCredentialLogName(base)
default:
return false
}
}
// zipLooksLikeKit inspects a ZIP's central directory (entry names only, no
// decompression) and reports whether the contents match a phishing kit rather
// than a legitimate plugin/theme distribution archive. Filename brand keywords
// alone flagged legitimate plugin zips (Instagram Feed, Facebook articles,
// cPanel eCRM); the archive body is what actually distinguishes a kit: a
// credential sink file, a capture script, a brand-login page, an anti-bot
// blocker. Two signal categories from at least two distinct entries are
// required so one generic result filename cannot decide the archive alone.
func zipLooksLikeKit(ctx context.Context, path string) bool {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "phishing", err)
return false
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
markScanReadError(ctx, "phishing", err)
return false
}
// Out of bounds or not a regular file is a scope decision, not a coverage
// failure: the caller already gates on the same range, and marking the
// check incomplete here would stop it ever retiring a stale finding on a
// host that simply keeps a large archive under a docroot.
if !info.Mode().IsRegular() || info.Size() <= 1000 || info.Size() >= 50*1024*1024 {
return false
}
return phishingKitZipHasEvidence(ctx, f, info.Size())
}
// phishingKitZipHasEvidence confirms a bounded ZIP using independent entry-name
// signals without decompressing attacker-controlled archive contents.
func phishingKitZipHasEvidence(ctx context.Context, reader io.ReaderAt, size int64) bool {
if size <= 1000 || size >= 50*1024*1024 {
return false
}
zr, err := zip.NewReader(reader, size)
if err != nil {
markCheckIncomplete(ctx, "phishing")
return false
}
capture, credSink, antibot := false, false, false
brandPages := 0
evidenceEntries := make(map[string]bool)
for _, e := range zr.File {
if e.FileInfo().IsDir() {
continue
}
base := e.Name
if i := strings.LastIndexAny(base, "/\\"); i >= 0 {
base = base[i+1:]
}
base = strings.ToLower(base)
hasEvidence := false
if kitCaptureScripts[base] {
capture = true
hasEvidence = true
}
if isKitCredentialSinkName(base) {
credSink = true
hasEvidence = true
}
if kitAntibotScripts[base] || strings.Contains(base, "antibot") {
antibot = true
hasEvidence = true
}
if strings.HasSuffix(base, ".html") || strings.HasSuffix(base, ".htm") || isExecutablePHPName(base) {
for _, w := range kitLoginPageWords {
if strings.Contains(base, w) {
brandPages++
hasEvidence = true
break
}
}
}
if hasEvidence {
evidenceEntries[e.Name] = true
}
}
signals := 0
if capture {
signals++
}
if credSink {
signals++
}
if antibot {
signals++
}
if brandPages >= 1 {
signals++
}
return signals >= 2 && len(evidenceEntries) >= 2
}
func visiblePageContent(s string, includeImageAttrs bool) string {
doc, err := html.Parse(strings.NewReader(s))
if err != nil {
return ""
}
hiddenClasses, hiddenIDs := stylesheetHiddenSelectors(doc)
var content strings.Builder
var walk func(*html.Node, bool)
walk = func(node *html.Node, hidden bool) {
if node.Type == html.ElementNode {
switch node.Data {
case "head", "script", "style", "template", "noscript", "iframe":
hidden = true
}
if htmlElementHidden(node, hiddenClasses, hiddenIDs) {
hidden = true
}
if includeImageAttrs && !hidden && node.Data == "img" {
for _, attr := range node.Attr {
switch attr.Key {
case "alt", "title", "aria-label", "src":
content.WriteByte(' ')
content.WriteString(attr.Val)
}
}
}
}
if node.Type == html.TextNode && !hidden {
content.WriteByte(' ')
content.WriteString(node.Data)
}
for child := node.FirstChild; child != nil; child = child.NextSibling {
walk(child, hidden)
}
}
walk(doc, false)
return content.String()
}
func htmlElementHidden(node *html.Node, hiddenClasses, hiddenIDs map[string]bool) bool {
for _, attr := range node.Attr {
switch attr.Key {
case "hidden":
return true
case "class":
for _, class := range strings.Fields(attr.Val) {
if hiddenClasses[class] {
return true
}
}
case "id":
if hiddenIDs[attr.Val] {
return true
}
case "aria-hidden":
if strings.EqualFold(strings.TrimSpace(attr.Val), "true") {
return true
}
case "style":
if cssDeclarationsHide(attr.Val) {
return true
}
}
}
return false
}
func stylesheetHiddenSelectors(doc *html.Node) (map[string]bool, map[string]bool) {
hiddenClasses := make(map[string]bool)
hiddenIDs := make(map[string]bool)
var walk func(*html.Node)
walk = func(node *html.Node) {
if node.Type == html.ElementNode && node.Data == "style" {
var css strings.Builder
for child := node.FirstChild; child != nil; child = child.NextSibling {
if child.Type == html.TextNode {
css.WriteString(child.Data)
}
}
collectHiddenStylesheetSelectors(css.String(), hiddenClasses, hiddenIDs)
}
for child := node.FirstChild; child != nil; child = child.NextSibling {
walk(child)
}
}
walk(doc)
return hiddenClasses, hiddenIDs
}
func collectHiddenStylesheetSelectors(css string, hiddenClasses, hiddenIDs map[string]bool) {
css = stripCSSComments(css)
for {
open := strings.IndexByte(css, '{')
if open < 0 {
return
}
closeOffset := strings.IndexByte(css[open+1:], '}')
if closeOffset < 0 {
return
}
closeIndex := open + 1 + closeOffset
if cssDeclarationsHide(css[open+1 : closeIndex]) {
for _, selector := range strings.Split(css[:open], ",") {
selector = strings.TrimSpace(selector)
if len(selector) < 2 || strings.ContainsAny(selector, " >+~:[]*") {
continue
}
marker := strings.LastIndexAny(selector, ".#")
if marker < 0 || !isSimpleCSSIdentifier(selector[marker+1:]) {
continue
}
prefix := selector[:marker]
if prefix != "" && !isSimpleCSSIdentifier(prefix) {
continue
}
switch selector[marker] {
case '.':
hiddenClasses[selector[marker+1:]] = true
case '#':
hiddenIDs[selector[marker+1:]] = true
}
}
}
css = css[closeIndex+1:]
}
}
func stripCSSComments(css string) string {
var stripped strings.Builder
for len(css) > 0 {
start := strings.Index(css, "/*")
if start < 0 {
stripped.WriteString(css)
break
}
stripped.WriteString(css[:start])
end := strings.Index(css[start+2:], "*/")
if end < 0 {
break
}
css = css[start+2+end+2:]
}
return stripped.String()
}
func cssDeclarationsHide(declarations string) bool {
for _, declaration := range strings.Split(strings.ToLower(declarations), ";") {
parts := strings.SplitN(declaration, ":", 2)
if len(parts) != 2 {
continue
}
property := strings.TrimSpace(parts[0])
value := strings.TrimSpace(strings.TrimSuffix(strings.TrimSpace(parts[1]), "!important"))
if (property == "display" && value == "none") ||
(property == "visibility" && (value == "hidden" || value == "collapse")) ||
(property == "opacity" && value == "0") {
return true
}
}
return false
}
func isSimpleCSSIdentifier(s string) bool {
if s == "" {
return false
}
for _, r := range s {
if !unicode.IsLetter(r) && !unicode.IsDigit(r) && r != '-' && r != '_' {
return false
}
}
return true
}
// visibleTextLen counts rendered non-whitespace characters. Script, style,
// template, noscript, iframe fallback, comment, and head contents do not
// document an embed.
func visibleTextLen(s string) int {
n := 0
for _, r := range visiblePageContent(s, false) {
if !unicode.IsSpace(r) {
n++
}
}
return n
}
// ---------------------------------------------------------------------------
// Safe directory list
// ---------------------------------------------------------------------------
// isKnownSafeDir names directories that CheckPhishing does not recurse into.
// The list is deliberately narrow: dependency trees whose bundled HTML
// documentation contains legitimate login-form examples, and VCS metadata.
// It is not an allowlist of "trusted" paths. WordPress core directories,
// caches, tmp and logs used to be pruned too, and kits were found under
// wp-includes and wp-admin precisely because scanners skip them; stock core
// ships no login-form HTML there, so a kit in those trees is as anomalous as
// one under uploads. wp-content and .well-known are prime drop paths and are
// always scanned. Never widen this list to skip a path where a file could be
// dropped and served; fix detection instead.
func isKnownSafeDir(name string) bool {
safeDirs := map[string]bool{
"node_modules": true, "vendor": true, ".git": true,
}
return safeDirs[name]
}
package checks
import (
"bytes"
"context"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/state"
)
// nestedEvalDecodeRe matches the PHP token sequence
// `<eval|assert> [ws] ( [ws] [@] [\] <ident> [ws] (`, with DOTALL so the
// source can have line breaks inside the whitespace gaps. Attackers wedge
// comments or common call modifiers between the sink and decoder; comments and
// strings are stripped before matching so only executable token structure is
// evaluated.
var nestedEvalDecodeRe = regexp.MustCompile(`(?is)\b(eval|assert)\s*\(\s*(?:(?:"\s*"|'\s*')?\s*\.\s*)*@?\s*\\?\s*(\w+)\s*\(`)
// reEvalRequestInput matches eval/assert applied to request input with no
// decoder at all (`eval($_POST['c'])`), the simplest possible shell. Strings
// are stripped before matching, so the superglobal name is what remains.
var reEvalRequestInput = regexp.MustCompile(`(?is)\b(eval|assert)\s*\(\s*(?:(?:"\s*"|'\s*')?\s*\.\s*)*@?\s*\$_(?:post|get|request|cookie)\b`)
// reEvalVarCallee matches eval wrapping a variable function call,
// e.g. `eval($f(...))`. The literal-callee form above cannot capture a
// `$var` callee, yet eval'ing the result of a dynamic function call is a
// near-certain dropper signal in user web directories.
var reEvalVarCallee = regexp.MustCompile(`(?is)\beval\s*\(\s*@?\s*\$\w+\s*\(`)
// evalExecWrapInner lists code-construction primitives that, when wrapped
// directly by eval, indicate dynamic code execution rather than the
// decoder/decompressor chain nestedEvalDecodeRe already covers.
var evalExecWrapInner = map[string]struct{}{
"create_function": {},
"call_user_func": {},
"call_user_func_array": {},
}
var callbackFirstArgFuncs = map[string]struct{}{
"array_map": {},
"call_user_func": {},
"call_user_func_array": {},
"register_shutdown_function": {},
"register_tick_function": {},
}
// callbackExecNames are names that execute code when used as a callback. They
// are RCE regardless of where the call's arguments come from, so they flag
// unconditionally.
var callbackExecNames = map[string]struct{}{
"assert": {},
"call_user_func": {},
"create_function": {},
"eval": {},
"exec": {},
"passthru": {},
"popen": {},
"proc_open": {},
"shell_exec": {},
"system": {},
}
// callbackDecoderNames are decode/decompress primitives. As a callback they
// only transform data (array_map('base64_decode', $data) returns decoded
// bytes, it executes nothing), so legitimate plugins use them constantly. They
// only signal a dropper when the same call is fed request input; the
// decode-then-eval shape is covered separately by the eval-chain detectors.
var callbackDecoderNames = map[string]struct{}{
"base64_decode": {},
"gzinflate": {},
"gzuncompress": {},
"str_rot13": {},
}
// reVarVarCall matches a variable-variable or dynamic-expression function
// invocation, e.g. `$$h(...)` or `${$x}(...)`. On its own this shows up in
// some dispatcher code, so it is only treated as an indicator when the same
// line also carries a request superglobal (the RCE shape).
var reVarVarCall = regexp.MustCompile(`(?:\$\$\w+|\$\{[^}]{1,64}\})\s*\(`)
// Goto-obfuscation discrimination. A label an obfuscator generates carries no
// meaning: it is short, or it is a stem plus a counter. A hand-written state
// machine names its labels after what they mean, which is why WordPress core's
// HTML5 insertion-mode labels must not count.
var reGotoLabel = regexp.MustCompile(`(?i)\bgoto\s+([A-Za-z_][A-Za-z0-9_]{0,63})\s*;`)
// reGotoExecSink is the evidence half of the goto heuristic, kept in step with
// the php_goto_obfuscation signature in configs/. call_user_func is absent on
// purpose: plugin loaders dispatch their own callables through it.
var reGotoExecSink = regexp.MustCompile(`(?i)\b(?:eval|assert|create_function|system|exec|passthru|shell_exec|proc_open|popen|pcntl_exec|base64_decode|gzinflate|gzuncompress|gzdecode|str_rot13|hex2bin|convert_uudecode)(?:\s|/\*[^*]*\*+(?:[^/*][^*]*\*+)*/|//[^\r\n]*[\r\n]|#[^\r\n]*[\r\n])*\(` +
`|(?-i:\$_(?:GET|POST|REQUEST|COOKIE|FILES)\b)` +
`|\$[A-Za-z_][A-Za-z0-9_]*(?:\s*\[[^\]\r\n]+\])*(?:\s|/\*[^*]*\*+(?:[^/*][^*]*\*+)*/|//[^\r\n]*[\r\n]|#[^\r\n]*[\r\n])*(?:\)(?:\s|/\*[^*]*\*+(?:[^/*][^*]*\*+)*/|//[^\r\n]*[\r\n]|#[^\r\n]*[\r\n])*)?\(` +
`|\b(?:include|require)(?:_once)?\b(?:\s|/\*[^*]*\*+(?:[^/*][^*]*\*+)*/|//[^\r\n]*[\r\n]|#[^\r\n]*[\r\n])*(?:\(\s*)?\$`)
// gotoLabelIsGenerated reports whether a goto label looks machine-generated.
// Digits are the strongest tell (lbl0, x9k, a1); anything shorter than four
// characters cannot carry meaning either.
func gotoLabelIsGenerated(label string) bool {
if len(label) < 4 {
return true
}
return strings.ContainsAny(label, "0123456789")
}
// includeDangerWrappers are stream wrappers / remote schemes that, as an
// include/require target, mean remote-file inclusion or php://input code
// execution. Matched on the comment-stripped (strings preserved) source.
var includeDangerWrappers = []string{"data://", "php://", "phar://", "http://", "https://", "ftp://"}
var includeKeywords = []string{"include_once", "require_once", "include", "require"}
var (
pregReplaceCallName = map[string]struct{}{"preg_replace": {}}
codeEvalPrimitiveCallNames = map[string]struct{}{"assert": {}, "create_function": {}}
callUserFuncCallNames = map[string]struct{}{"call_user_func": {}, "call_user_func_array": {}}
)
// phpContentReadSize bounds the head window read from a PHP file for content
// analysis. Most real PHP is far smaller; the window is large enough that an
// attacker cannot cheaply hide a payload by prepending benign padding.
const phpContentReadSize = 1 << 20 // 1 MiB head window
// phpContentTailSize is also scanned for files larger than the head window so
// a payload appended after >1 MiB of padding is still seen. The two windows
// are scanned together; bounded total memory is head+tail.
const phpContentTailSize = 64 << 10 // 64 KiB tail window
// readPHPContentWindows returns the bytes to analyse for a PHP file: the
// first phpContentReadSize bytes, plus the last phpContentTailSize bytes when
// the file is larger than the head window. tailOffset is -1 when tail is empty.
// readOK is false only on a real read error.
func readPHPContentWindows(f io.ReaderAt, size int64) (head, tail []byte, tailOffset int64, readOK bool) {
tailOffset = -1
headLen := int64(phpContentReadSize)
if size >= 0 && size < headLen {
headLen = size
}
headBuf := make([]byte, headLen)
n, err := f.ReadAt(headBuf, 0)
if err != nil && !errors.Is(err, io.EOF) {
return nil, nil, -1, false
}
head = headBuf[:n]
if size > int64(phpContentReadSize) {
tailLen := int64(phpContentTailSize)
off := size - tailLen
if off < int64(phpContentReadSize) {
// Overlap into the head window is fine; it only re-scans bytes.
off = int64(phpContentReadSize)
tailLen = size - off
}
if tailLen > 0 {
tailBuf := make([]byte, tailLen)
tn, terr := f.ReadAt(tailBuf, off)
if terr != nil && !errors.Is(terr, io.EOF) {
return nil, nil, -1, false
}
tail = tailBuf[:tn]
tailOffset = off
}
}
return head, tail, tailOffset, true
}
func combinePHPContentWindows(head, tail []byte) []byte {
if len(tail) == 0 {
return head
}
capacity := len(head) + len(tail)
raw := make([]byte, 0, capacity)
raw = append(raw, head...)
raw = append(raw, tail...)
return raw
}
func phpCodeOnlyWindows(head, tail []byte) string {
return phpCodeOnly(string(combinePHPContentWindows(head, tail)))
}
func hasBacktickSuperglobal(code string) bool {
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
if code[i] != '`' {
continue
}
end := i + 1
for end < len(code) {
if code[end] == '\\' && end+1 < len(code) {
end += 2
continue
}
if code[end] == '`' {
break
}
end++
}
if end >= len(code) {
break
}
if containsRequestSuperglobal(code[i+1 : end]) {
return true
}
i = end
}
return false
}
func hasCallbackExecName(code string) bool {
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
nameStart := i
if code[i] == '\\' {
if i+1 >= len(code) || !isPHPIdentifierStart(code[i+1]) || !canStartGlobalPHPFunction(code, i) {
continue
}
nameStart = i + 1
} else if !isPHPIdentifierStart(code[i]) || !canStartPHPFunctionName(code, i) {
continue
}
nameEnd := nameStart + 1
for nameEnd < len(code) && isPHPIdentifierPart(code[nameEnd]) {
nameEnd++
}
name := strings.ToLower(code[nameStart:nameEnd])
if _, ok := callbackFirstArgFuncs[name]; !ok {
i = nameEnd - 1
continue
}
openParen := skipPHPWhitespace(code, nameEnd)
if openParen >= len(code) || code[openParen] != '(' {
i = nameEnd - 1
continue
}
firstArg := skipPHPWhitespace(code, openParen+1)
if firstArg >= len(code) || !isPHPQuote(code[firstArg]) {
i = nameEnd - 1
continue
}
callbackName, _, ok := readPHPFunctionString(code, firstArg)
if !ok {
i = nameEnd - 1
continue
}
if _, dangerous := callbackExecNames[callbackName]; dangerous {
return true
}
if _, decoder := callbackDecoderNames[callbackName]; decoder {
closeParen := matchingParen(code, openParen)
if containsRequestSuperglobalExpression(code[openParen:closeParen]) {
return true
}
}
i = nameEnd - 1
}
return false
}
// matchingParen returns the index of the close paren matching the open paren
// at openParen, or len(code) if unbalanced. Quoted strings are skipped so
// parens inside string literals do not throw off the depth count.
func matchingParen(code string, openParen int) int {
depth := 0
for i := openParen; i < len(code); i++ {
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label) - 1
continue
}
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
switch code[i] {
case '(':
depth++
case ')':
depth--
if depth == 0 {
return i
}
}
}
return len(code)
}
// hasDangerousInclude reports an include/require whose target is request
// input or a remote/stream wrapper -- the LFI/RFI and php://input code-exec
// shapes. Scans the include expression, not the whole line, so unrelated
// request reads or URLs after a static include do not trip the detector.
func hasDangerousInclude(code string) bool {
searchFrom := 0
for {
exprStart, exprEnd, ok := nextIncludeExpression(code, searchFrom)
if !ok {
break
}
expr := code[exprStart:exprEnd]
if containsIncludeTargetExpression(expr) {
return true
}
exprLower := strings.ToLower(expr)
for _, w := range includeDangerWrappers {
if strings.Contains(exprLower, w) {
return true
}
}
searchFrom = exprEnd
if searchFrom >= len(code) {
break
}
}
return false
}
// hasCodeEvalPrimitiveWithRequest reports assert()/create_function() fed
// request input. create_function() evals its body argument, so any request
// superglobal inside the call is a code-eval sink. assert() (pre-8.0) only
// evaluates its first argument as PHP when that argument is a string, so the
// argument expression is classified first: expressions that provably yield a
// non-string value (boolean/comparison operators, negation, builtins that
// never return strings) are not code-eval sinks. Anything not provably
// non-string stays flagged, fail closed.
func hasCodeEvalPrimitiveWithRequest(code string) bool {
// Whether an unqualified builtin name can be trusted is only consulted on
// the assert()-with-request-argument path, which almost no file hits.
// phpMayRedirectUnqualifiedBuiltin tokenizes the whole scan window, so
// defer that pass until the first time it is actually needed, then cache.
trustComputed := false
trustUnqualified := false
resolveTrust := func() bool {
if !trustComputed {
trustUnqualified = !phpMayRedirectUnqualifiedBuiltin(code)
trustComputed = true
}
return trustUnqualified
}
searchFrom := 0
for {
callStart, openParen, closeParen, ok := nextStandalonePHPCall(code, searchFrom, codeEvalPrimitiveCallNames)
if !ok {
break
}
name := strings.ToLower(strings.TrimPrefix(strings.TrimSpace(code[callStart:openParen]), `\`))
exprEnd := closeParen
if exprEnd < len(code) {
exprEnd++
}
if containsRequestSuperglobal(code[openParen:exprEnd]) {
if name != "assert" {
return true
}
if closeParen >= len(code) {
// The call is unbalanced in the scan window, so the
// argument cannot be classified. Fail closed.
return true
}
args := phpCallArguments(code, openParen+1, closeParen)
if len(args) > 0 && containsRequestSuperglobal(args[0]) &&
!phpExpressionYieldsNonString(args[0], resolveTrust()) {
return true
}
}
searchFrom = nextSearchOffset(closeParen, len(code))
if searchFrom >= len(code) {
break
}
}
return false
}
// Precedence ranks for depth-0 binary operators inside an assert() argument,
// lowest precedence first. The lowest-precedence operator present produces
// the expression's final value, so it alone decides whether the result can
// be a string.
const (
phpRankWordLogical = 1 // and / or / xor -- boolean result
phpRankAssign = 2 // = and compound assignments -- value passes through
phpRankTernary = 3 // ?: / ?? / ?-> -- string-capable value
phpRankLogicalOr = 4 // || -- boolean result
phpRankLogicalAnd = 5 // && -- boolean result
phpRankBitwiseOr = 6 // | -- string-capable when both operands are strings
phpRankBitwiseXor = 7 // ^ -- string-capable when both operands are strings
phpRankBitwiseAnd = 8 // & -- string-capable when both operands are strings
phpRankComparison = 9 // == != === !== <> <=> < > <= >= instanceof -- non-string
phpRankShift = 10 // << >> -- integer result; tied with . for PHP 7/8 precedence drift
phpRankConcat = 10 // . -- string result
phpRankAdditive = 12 // + - -- numeric/array result, never string
phpRankMultiplicative = 13 // * / % -- numeric result
phpRankPower = 14 // ** -- numeric result
phpRankNone = 100 // no depth-0 binary operator: single atom
)
type phpOperatorValue int
const (
phpOperatorNone phpOperatorValue = iota
phpOperatorNonString
phpOperatorStringCapable
)
// phpExpressionYieldsNonString reports whether a PHP expression provably
// evaluates to a non-string value, so assert() cannot treat it as code.
// trustUnqualified is false when the file declares a namespace or imports a
// function alias, because then an unqualified builtin name can resolve to
// attacker-defined code. Anything unrecognized counts as string-capable.
//
// Wrapping parentheses are peeled in one parser pass so a payload with tens of
// thousands of nested parens cannot exhaust the stack or trigger quadratic
// rescans.
func phpExpressionYieldsNonString(expr string, trustUnqualified bool) bool {
expr = peelPHPAssertionExpression(expr)
if expr == "" {
return false
}
switch phpTopLevelOperatorValue(expr) {
case phpOperatorNonString:
return true
case phpOperatorStringCapable:
return false
}
if expr[0] == '!' {
// Negation of a whole atom is always boolean. Operators inside the
// negated expression would have been classified above, so this point
// is only reached for a single negated atom.
return true
}
return phpIsNonStringBuiltinCall(expr, trustUnqualified)
}
func peelPHPAssertionExpression(expr string) string {
start, end := trimPHPExpressionBounds(expr, 0, len(expr))
if start >= end || expr[start] != '(' {
return expr[start:end]
}
prefixOpens := phpLeadingParenPositions(expr, start, end)
if len(prefixOpens) == 0 {
return expr[start:end]
}
matches, ok := phpMatchPrefixParens(expr, start, end, prefixOpens)
if !ok {
return expr[start:end]
}
innerStart, innerEnd := start, end
for idx, open := range prefixOpens {
if open != skipPHPWhitespaceUntil(expr, innerStart, innerEnd) {
break
}
close := matches[idx]
if close != skipPHPWhitespaceBack(expr, innerStart, innerEnd)-1 {
break
}
innerStart = open + 1
innerEnd = close
}
start, end = trimPHPExpressionBounds(expr, innerStart, innerEnd)
return expr[start:end]
}
func phpLeadingParenPositions(expr string, start, end int) []int {
var positions []int
for {
start = skipPHPWhitespaceUntil(expr, start, end)
if start >= end || expr[start] != '(' {
return positions
}
positions = append(positions, start)
start++
}
}
func phpMatchPrefixParens(expr string, start, end int, prefixOpens []int) ([]int, bool) {
matches := make([]int, len(prefixOpens))
for i := range matches {
matches[i] = -1
}
depth := 0
for i := start; i < end; i++ {
if label, bodyStart, ok := phpHeredocOpen(expr, i); ok {
i = phpHeredocEnd(expr, bodyStart, label) - 1
continue
}
if isPHPQuote(expr[i]) {
i = skipPHPString(expr, i)
continue
}
switch expr[i] {
case '(':
depth++
case ')':
if depth == 0 {
return nil, false
}
depth--
if depth < len(matches) && matches[depth] == -1 {
matches[depth] = i
}
}
}
if depth != 0 {
return nil, false
}
for _, match := range matches {
if match == -1 {
return nil, false
}
}
return matches, true
}
func trimPHPExpressionBounds(expr string, start, end int) (int, int) {
if start < 0 {
start = 0
}
if end > len(expr) {
end = len(expr)
}
start = skipPHPWhitespaceUntil(expr, start, end)
end = skipPHPWhitespaceBack(expr, start, end)
return start, end
}
func skipPHPWhitespaceUntil(expr string, start, end int) int {
for start < end && isPHPSpace(expr[start]) {
start++
}
return start
}
func skipPHPWhitespaceBack(expr string, start, end int) int {
for end > start && isPHPSpace(expr[end-1]) {
end--
}
return end
}
// phpTopLevelOperatorValue scans expr outside string literals at paren depth 0
// and returns whether the lowest-precedence operator found has a provably
// non-string result. String-capable operators win ties to keep the classifier
// fail-closed across PHP version precedence differences.
func phpTopLevelOperatorValue(expr string) phpOperatorValue {
rank := phpRankNone
value := phpOperatorNone
observe := func(opRank int, opValue phpOperatorValue) {
if opRank < rank || (opRank == rank && opValue == phpOperatorStringCapable) {
rank = opRank
value = opValue
}
}
depth := 0
prev := byte(0) // last significant byte
i := 0
for i < len(expr) {
c := expr[i]
if label, bodyStart, ok := phpHeredocOpen(expr, i); ok {
observe(phpRankNone, phpOperatorStringCapable)
i = phpHeredocEnd(expr, bodyStart, label)
prev = '"'
continue
}
if isPHPQuote(c) {
i = skipPHPString(expr, i) + 1
prev = '"'
continue
}
if isPHPSpace(c) {
i++
continue
}
switch c {
case '(', '[', '{':
depth++
prev = c
i++
continue
case ')', ']', '}':
if depth > 0 {
depth--
}
prev = c
i++
continue
}
if depth > 0 {
prev = c
i++
continue
}
if isPHPIdentifierStart(c) {
start := i
for i < len(expr) && isPHPIdentifierPart(expr[i]) {
i++
}
// A word preceded by $, ->, ::, or \ is a variable, member, or
// namespaced name, never the operator keyword.
if !isPHPIdentifierPart(prev) && prev != '$' && prev != '>' && prev != ':' && prev != '\\' {
switch strings.ToLower(expr[start:i]) {
case "and", "or", "xor":
observe(phpRankWordLogical, phpOperatorNonString)
case "instanceof":
observe(phpRankComparison, phpOperatorNonString)
}
}
prev = expr[i-1]
continue
}
next := byte(0)
if i+1 < len(expr) {
next = expr[i+1]
}
consumed := 1
switch c {
case '=':
switch next {
case '=':
observe(phpRankComparison, phpOperatorNonString)
consumed = 2
if i+2 < len(expr) && expr[i+2] == '=' {
consumed = 3
}
case '>':
// A stray => at depth 0 is not valid expression syntax;
// treat it like an assignment and fail closed.
observe(phpRankAssign, phpOperatorStringCapable)
consumed = 2
default:
observe(phpRankAssign, phpOperatorStringCapable)
}
case '!':
if next == '=' {
observe(phpRankComparison, phpOperatorNonString)
consumed = 2
if i+2 < len(expr) && expr[i+2] == '=' {
consumed = 3
}
}
case '?':
observe(phpRankTernary, phpOperatorStringCapable)
if next == '?' {
consumed = 2
} else if next == '-' && i+2 < len(expr) && expr[i+2] == '>' {
consumed = 3
}
case '<':
switch {
case next == '<':
observe(phpRankShift, phpOperatorNonString)
consumed = 2
case next == '=' && i+2 < len(expr) && expr[i+2] == '>':
observe(phpRankComparison, phpOperatorNonString)
consumed = 3
default:
observe(phpRankComparison, phpOperatorNonString)
if next == '=' || next == '>' {
consumed = 2
}
}
case '>':
switch next {
case '>':
observe(phpRankShift, phpOperatorNonString)
consumed = 2
case '=':
observe(phpRankComparison, phpOperatorNonString)
consumed = 2
default:
observe(phpRankComparison, phpOperatorNonString)
}
case '-':
switch next {
case '>':
consumed = 2 // property or method access, not an operator
case '-':
consumed = 2 // string increment/decrement is string-capable; fail closed as an atom
default:
observe(phpRankAdditive, phpOperatorNonString)
}
case '+':
if next == '+' {
consumed = 2 // string increment is string-capable; fail closed as an atom
} else {
observe(phpRankAdditive, phpOperatorNonString)
}
case '|':
if next == '|' {
observe(phpRankLogicalOr, phpOperatorNonString)
consumed = 2
} else {
observe(phpRankBitwiseOr, phpOperatorStringCapable)
}
case '&':
if next == '&' {
observe(phpRankLogicalAnd, phpOperatorNonString)
consumed = 2
} else {
observe(phpRankBitwiseAnd, phpOperatorStringCapable)
}
case '^':
observe(phpRankBitwiseXor, phpOperatorStringCapable)
case '.':
if !phpDotBelongsToNumberLiteral(expr, i) {
observe(phpRankConcat, phpOperatorStringCapable)
}
case '*':
if next == '*' {
observe(phpRankPower, phpOperatorNonString)
consumed = 2
} else {
observe(phpRankMultiplicative, phpOperatorNonString)
}
case '/', '%':
observe(phpRankMultiplicative, phpOperatorNonString)
}
prev = expr[i+consumed-1]
i += consumed
}
return value
}
func phpDotBelongsToNumberLiteral(expr string, dot int) bool {
if dot > 0 && isPHPDecimalDigit(expr[dot-1]) {
return phpDotContinuesDecimalLiteral(expr, dot)
}
if dot+1 >= len(expr) || !isPHPDecimalDigit(expr[dot+1]) {
return false
}
for i := dot - 1; i >= 0; i-- {
if isPHPSpace(expr[i]) {
continue
}
switch expr[i] {
case '(', '[', '{', ',', '?', ':', '=', '!', '<', '>', '+', '-', '*', '/', '%', '&', '|', '^', '~':
return true
default:
return false
}
}
return true
}
func phpDotContinuesDecimalLiteral(expr string, dot int) bool {
for i := dot - 1; i >= 0; i-- {
if isPHPDecimalDigit(expr[i]) || expr[i] == '_' {
continue
}
switch expr[i] {
case '.', 'e', 'E', 'x', 'X', 'b', 'B', 'o', 'O':
return false
case '+', '-':
return i == 0 || (expr[i-1] != 'e' && expr[i-1] != 'E')
default:
return !isPHPIdentifierPart(expr[i]) && expr[i] != '$' && expr[i] != '\\'
}
}
return true
}
func isPHPDecimalDigit(c byte) bool {
return c >= '0' && c <= '9'
}
// phpNonStringReturnBuiltins lists PHP builtins (plus the isset/empty
// language constructs) whose return value can never be a string, so feeding
// their result to assert() cannot evaluate attacker input as code. Functions
// that can return a string (filter_var, base64_decode, ...) must never be
// added here.
var phpNonStringReturnBuiltins = map[string]struct{}{
"isset": {}, "empty": {}, "defined": {},
"file_exists": {}, "is_file": {}, "is_dir": {}, "is_link": {},
"is_readable": {}, "is_writable": {}, "is_writeable": {}, "is_executable": {},
"is_uploaded_file": {},
"is_numeric": {}, "is_string": {}, "is_array": {}, "is_int": {},
"is_integer": {}, "is_long": {}, "is_bool": {}, "is_float": {},
"is_double": {}, "is_null": {}, "is_object": {}, "is_callable": {},
"is_iterable": {}, "is_countable": {}, "is_resource": {}, "is_a": {},
"is_subclass_of": {},
"in_array": {}, "array_key_exists": {}, "key_exists": {},
"property_exists": {}, "method_exists": {}, "function_exists": {},
"class_exists": {}, "interface_exists": {}, "trait_exists": {}, "enum_exists": {},
"str_contains": {}, "str_starts_with": {}, "str_ends_with": {},
"preg_match": {}, "preg_match_all": {},
"count": {}, "sizeof": {}, "strlen": {}, "mb_strlen": {},
"strcmp": {}, "strcasecmp": {}, "strncmp": {}, "strncasecmp": {},
"substr_compare": {}, "version_compare": {}, "hash_equals": {},
"password_verify": {}, "checkdate": {},
"ctype_alnum": {}, "ctype_alpha": {}, "ctype_cntrl": {}, "ctype_digit": {},
"ctype_graph": {}, "ctype_lower": {}, "ctype_print": {}, "ctype_punct": {},
"ctype_space": {}, "ctype_upper": {}, "ctype_xdigit": {},
}
// phpIsNonStringBuiltinCall reports whether expr is exactly one call to a
// builtin that never returns a string. isset/empty are language constructs
// and cannot be shadowed; other names are trusted only when fully qualified
// (\file_exists) or when the file cannot redirect unqualified names.
func phpIsNonStringBuiltinCall(expr string, trustUnqualified bool) bool {
i := 0
qualified := false
if i < len(expr) && expr[i] == '\\' {
qualified = true
i++
}
if i >= len(expr) || !isPHPIdentifierStart(expr[i]) {
return false
}
nameStart := i
for i < len(expr) && isPHPIdentifierPart(expr[i]) {
i++
}
name := strings.ToLower(expr[nameStart:i])
if _, ok := phpNonStringReturnBuiltins[name]; !ok {
return false
}
open := skipPHPWhitespace(expr, i)
if open >= len(expr) || expr[open] != '(' {
return false
}
if matchingParen(expr, open) != len(expr)-1 {
return false
}
if name == "isset" || name == "empty" {
return true
}
return qualified || trustUnqualified
}
// phpMayRedirectUnqualifiedBuiltin reports whether the source can make an
// unqualified trusted builtin name resolve to attacker-defined code: namespace
// fallback/shadowing, use-function imports, or same-file global polyfills for
// names that are builtins only on newer PHP versions. Closure captures
// (function () use ($x)), class methods, and trait use inside classes do not
// count.
func phpMayRedirectUnqualifiedBuiltin(code string) bool {
prev := byte(0) // last significant byte
classDepth := 0
var braceStack []bool
pendingClassScope := false
i := 0
for i < len(code) {
c := code[i]
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label)
prev = '"'
continue
}
if isPHPQuote(c) {
i = skipPHPString(code, i) + 1
prev = '"'
continue
}
if c == ' ' || c == '\t' || c == '\n' || c == '\r' {
i++
continue
}
switch c {
case '{':
braceStack = append(braceStack, pendingClassScope)
if pendingClassScope {
classDepth++
}
pendingClassScope = false
prev = c
i++
continue
case '}':
if len(braceStack) > 0 {
last := len(braceStack) - 1
if braceStack[last] && classDepth > 0 {
classDepth--
}
braceStack = braceStack[:last]
}
pendingClassScope = false
prev = c
i++
continue
case ';':
pendingClassScope = false
}
if !isPHPIdentifierStart(c) {
prev = c
i++
continue
}
start := i
for i < len(code) && isPHPIdentifierPart(code[i]) {
i++
}
word := strings.ToLower(code[start:i])
guarded := isPHPIdentifierPart(prev) || prev == '$' || prev == '>' || prev == ':' || prev == '\\'
prev = code[i-1]
if guarded {
continue
}
switch word {
case "class", "interface", "trait", "enum":
pendingClassScope = true
case "namespace":
return true
case "function":
if classDepth == 0 && phpFunctionDeclarationShadowsNonStringBuiltin(code, i) {
return true
}
case "use":
if classDepth > 0 {
continue
}
j := skipPHPWhitespace(code, i)
if j < len(code) && code[j] == '(' {
continue // closure capture
}
if phpUseStatementImportsFunction(code, j) {
return true
}
}
}
return false
}
func phpFunctionDeclarationShadowsNonStringBuiltin(code string, from int) bool {
i := skipPHPWhitespace(code, from)
if i < len(code) && code[i] == '&' {
i = skipPHPWhitespace(code, i+1)
}
if i >= len(code) || !isPHPIdentifierStart(code[i]) {
return false
}
start := i
for i < len(code) && isPHPIdentifierPart(code[i]) {
i++
}
_, shadows := phpNonStringReturnBuiltins[strings.ToLower(code[start:i])]
return shadows
}
// phpUseStatementImportsFunction scans one use statement (until the first
// semicolon) for the function keyword, covering both use function a\b and
// use a\{function b} forms.
func phpUseStatementImportsFunction(code string, from int) bool {
i := from
for i < len(code) && code[i] != ';' {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i) + 1
continue
}
if !isPHPIdentifierStart(code[i]) {
i++
continue
}
start := i
for i < len(code) && isPHPIdentifierPart(code[i]) {
i++
}
if strings.ToLower(code[start:i]) == "function" {
return true
}
}
return false
}
// hasPregReplaceEvalWithRequest reports a preg_replace() call whose pattern
// carries the /e modifier and whose evaluated replacement or subject reads
// request input. The /e modifier (removed in PHP 7.0) evaluates the replacement
// as PHP after backreferences from the subject are interpolated, so pattern-only
// request input is not enough to call this a code-execution sink. Bare /e with
// static arguments is the legacy WordPress serialize-fix / autolink idiom --
// real /e usage, but not a dropper -- so it is gated on request-input
// correlation, the same way the shell, include, and assert detectors here gate
// their sinks.
// Walks the source skipping string literals so a documentation example in a
// string does not trip, then inspects the literal first argument of each call.
func hasPregReplaceEvalWithRequest(code string) bool {
searchFrom := 0
for {
callStart, openParen, closeParen, ok := nextStandalonePHPCall(code, searchFrom, pregReplaceCallName)
if !ok {
break
}
args := phpCallArguments(code, openParen+1, closeParen)
if len(args) < 3 || !pregPatternArgumentHasEvalModifier(args[0]) {
searchFrom = nextSearchOffset(closeParen, len(code))
continue
}
tainted := requestTaintedVariablesBefore(code, callStart)
if pregReplacementReadsRequest(args[1], tainted) || phpExpressionReadsRequest(args[2], tainted) {
return true
}
searchFrom = nextSearchOffset(closeParen, len(code))
}
return false
}
func pregPatternArgumentHasEvalModifier(expr string) bool {
expr = strings.TrimSpace(expr)
if expr == "" || !isPHPQuote(expr[0]) {
return false
}
end := skipPHPString(expr, 0)
// skipPHPString returns the closing-quote index, or the last index for an
// unterminated literal. Guard the slice: a quote in the final position
// leaves no string body to inspect.
if end <= 0 {
return false
}
return pregPatternHasEvalModifier(expr[1:end])
}
func pregReplacementReadsRequest(expr string, tainted map[string]struct{}) bool {
if phpExpressionReadsRequest(expr, tainted) {
return true
}
body, _, ok := phpStringLiteralExpression(expr)
if !ok {
return false
}
return phpExpressionReadsRequest(body, tainted)
}
func phpExpressionReadsRequest(expr string, tainted map[string]struct{}) bool {
return containsRequestSuperglobalExpression(expr) || phpExpressionReferencesTaintedVariable(expr, tainted)
}
func phpExpressionReferencesTaintedVariable(expr string, tainted map[string]struct{}) bool {
if len(tainted) == 0 {
return false
}
for i := 0; i < len(expr); i++ {
if isPHPQuote(expr[i]) {
i = skipPHPString(expr, i)
continue
}
if expr[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(expr, i)
if !ok {
continue
}
if _, found := tainted[variable]; found {
return true
}
i = next - 1
}
return false
}
func requestTaintedVariablesBefore(code string, limit int) map[string]struct{} {
if limit > len(code) {
limit = len(code)
}
if limit <= 0 {
return nil
}
scan := code[:limit]
taintStack := []map[string]struct{}{{}}
functionBraceStack := []bool{}
pendingFunctionScope := false
for i := 0; i < len(scan); i++ {
if isPHPQuote(scan[i]) {
i = skipPHPString(scan, i)
continue
}
if end, ok := phpKeywordAt(scan, i, "function"); ok {
pendingFunctionScope = true
i = end - 1
continue
}
switch scan[i] {
case '{':
functionBraceStack = append(functionBraceStack, pendingFunctionScope)
if pendingFunctionScope {
taintStack = append(taintStack, map[string]struct{}{})
}
pendingFunctionScope = false
continue
case '}':
if len(functionBraceStack) > 0 {
last := len(functionBraceStack) - 1
if functionBraceStack[last] && len(taintStack) > 1 {
taintStack = taintStack[:len(taintStack)-1]
}
functionBraceStack = functionBraceStack[:last]
}
pendingFunctionScope = false
continue
case ';':
pendingFunctionScope = false
}
if scan[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(scan, i)
if !ok || isRequestSuperglobalVariable(variable) {
continue
}
j := skipPHPWhitespace(scan, next)
opLen, directAssign, appendAssign, ok := phpAssignmentOperator(scan, j)
if !ok {
i = next - 1
continue
}
exprStart := skipPHPWhitespace(scan, j+opLen)
exprEnd := phpExpressionEnd(scan, exprStart)
expr := scan[exprStart:exprEnd]
tainted := taintStack[len(taintStack)-1]
exprTainted := phpExpressionReadsRequest(expr, tainted)
switch {
case directAssign:
if exprTainted {
tainted[variable] = struct{}{}
} else {
delete(tainted, variable)
}
case appendAssign:
if exprTainted {
tainted[variable] = struct{}{}
}
default:
if exprTainted {
tainted[variable] = struct{}{}
} else {
delete(tainted, variable)
}
}
i = exprEnd - 1
}
return taintStack[len(taintStack)-1]
}
func phpKeywordAt(code string, start int, keyword string) (int, bool) {
end := start + len(keyword)
if end > len(code) || !strings.EqualFold(code[start:end], keyword) {
return 0, false
}
if start > 0 {
prev := code[start-1]
if isPHPIdentifierPart(prev) || prev == '$' || prev == '>' || prev == ':' || prev == '\\' {
return 0, false
}
}
if end < len(code) && isPHPIdentifierPart(code[end]) {
return 0, false
}
return end, true
}
func isRequestSuperglobalVariable(variable string) bool {
switch strings.ToLower(variable) {
case "_request", "_post", "_get", "_cookie", "_server":
return true
default:
return false
}
}
func phpStringLiteralExpression(expr string) (string, byte, bool) {
expr = strings.TrimSpace(expr)
if expr == "" || !isPHPQuote(expr[0]) {
return "", 0, false
}
end := skipPHPString(expr, 0)
if end <= 0 || end >= len(expr) || expr[end] != expr[0] {
return "", 0, false
}
if skipPHPWhitespace(expr, end+1) != len(expr) {
return "", 0, false
}
return expr[1:end], expr[0], true
}
func phpCallArguments(code string, start, end int) []string {
if start < 0 {
start = 0
}
if end > len(code) {
end = len(code)
}
if start > end {
return nil
}
var args []string
argStart := skipPHPWhitespace(code, start)
depth := 0
for i := start; i < end; i++ {
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label) - 1
continue
}
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
switch code[i] {
case '(', '[', '{':
depth++
case ')', ']', '}':
if depth > 0 {
depth--
}
case ',':
if depth == 0 {
args = append(args, strings.TrimSpace(code[argStart:i]))
argStart = skipPHPWhitespace(code, i+1)
}
}
}
if tail := strings.TrimSpace(code[argStart:end]); tail != "" || len(args) > 0 {
args = append(args, tail)
}
return args
}
// pregPatternHasEvalModifier returns true when the PCRE pattern string carries
// an "e" modifier after its closing delimiter.
func pregPatternHasEvalModifier(pat string) bool {
if len(pat) < 2 {
return false
}
open := pat[0]
// PHP forbids alphanumeric, backslash, and whitespace delimiters.
if isPHPIdentifierPart(open) || open == '\\' || open == ' ' {
return false
}
closeDelim := open
switch open {
case '(':
closeDelim = ')'
case '[':
closeDelim = ']'
case '{':
closeDelim = '}'
case '<':
closeDelim = '>'
}
idx := strings.LastIndexByte(pat, closeDelim)
if idx <= 0 {
return false
}
for _, c := range pat[idx+1:] {
if (c < 'a' || c > 'z') && (c < 'A' || c > 'Z') {
return false
}
if c == 'e' {
return true
}
}
return false
}
func nextIncludeExpression(line string, searchFrom int) (int, int, bool) {
for i := searchFrom; i < len(line); i++ {
if isPHPQuote(line[i]) {
i = skipPHPString(line, i)
continue
}
keywordEnd, ok := includeKeywordEnd(line, i)
if !ok {
continue
}
exprStart := skipPHPWhitespace(line, keywordEnd)
return exprStart, phpExpressionEnd(line, exprStart), true
}
return 0, 0, false
}
func includeKeywordEnd(line string, start int) (int, bool) {
if !canStartIncludeKeyword(line, start) {
return 0, false
}
for _, keyword := range includeKeywords {
end := start + len(keyword)
if end > len(line) || !strings.EqualFold(line[start:end], keyword) {
continue
}
if end < len(line) && isPHPIdentifierPart(line[end]) {
continue
}
if precededByFunctionKeyword(line, start) {
continue
}
return end, true
}
return 0, false
}
func canStartIncludeKeyword(line string, start int) bool {
if start == 0 {
return true
}
prev := line[start-1]
return !isPHPIdentifierPart(prev) && prev != '$' && prev != '>' && prev != ':' && prev != '\\'
}
// precededByFunctionKeyword reports whether the token before start is the
// `function` keyword, i.e. start begins a declared function or method name
// rather than a language construct. PHP 7 allows a reserved word as a method
// name, and a bootstrap method called include() is a common plugin idiom;
// treating that declaration as an include statement makes the method body its
// target expression, so any request input inside the method reads as an
// include of request input.
func precededByFunctionKeyword(line string, start int) bool {
const kw = "function"
// Only inspect the separator after a keyword match, never at each byte
// offset. Separators have no length limit in PHP; each whitespace run is
// visited at most once per matched keyword, keeping the full scan linear.
end := skipPHPWhitespaceBack(line, 0, start)
if end > 0 && line[end-1] == '&' {
end = skipPHPWhitespaceBack(line, 0, end-1)
}
begin := end - len(kw)
return begin >= 0 && strings.EqualFold(line[begin:end], kw) && canStartIncludeKeyword(line, begin)
}
func phpExpressionEnd(code string, start int) int {
depth := 0
for i := start; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
switch code[i] {
case '(', '[', '{':
depth++
case ')', ']', '}':
if depth == 0 {
return i
}
depth--
case ';':
if depth == 0 {
return i
}
case ',':
if depth == 0 {
return i
}
}
}
return len(code)
}
// phpCodeOnly blanks the inline-HTML regions of a PHP source so only the code
// inside <?php ... ?> (and <?= ... ?>) spans is analysed for execution sinks.
// Inline HTML is literal output and cannot execute PHP, so scanning it as code
// only yields false positives: apostrophes in prose ("don't", "you're") desync
// the string scanner, and href URLs, JS backtick template literals, and English
// words like "include"/"require" in markup then read as PHP execution sinks.
//
// HTML bytes become spaces (newlines kept so line-oriented detectors keep their
// line structure); a closing "?>" becomes "; " so it still bounds the preceding
// statement and consecutive <?php?> blocks do not run their expressions
// together. The "?>" scan skips PHP strings, heredoc/nowdoc bodies, and block
// comments so a "?>" inside them does not end PHP mode and blank real code. A
// file with no PHP open tag yields all blanks -- it executes nothing.
func phpCodeOnly(src string) string {
var b strings.Builder
b.Grow(len(src))
n := len(src)
i := 0
for i < n {
// HTML mode: blank up to the next "<?" open tag.
htmlStart := i
for i < n && (src[i] != '<' || i+1 >= n || src[i+1] != '?') {
i++
}
blankInlineHTML(&b, src[htmlStart:i])
if i >= n {
break
}
// Blank the opening tag: "<?php" (needs trailing whitespace/EOF), "<?=",
// or a bare "<?" short tag.
i += 2
b.WriteString(" ")
if i+3 <= n && strings.EqualFold(src[i:i+3], "php") && (i+3 == n || isPHPSpace(src[i+3])) {
b.WriteString(" ")
i += 3
} else if i < n && src[i] == '=' {
b.WriteByte(' ')
i++
}
// PHP mode: copy verbatim until a top-level "?>".
i = copyPHPModeRegion(&b, src, i)
}
return b.String()
}
// copyPHPModeRegion copies src[start:] into b verbatim until a top-level "?>"
// (which it replaces with "; ") or EOF, and returns the resume index. Strings,
// heredoc/nowdoc bodies, and block comments are copied whole so a "?>" inside
// them is not mistaken for a closing tag. A "?>" inside a // or # line comment
// does end PHP mode, matching PHP's own tokeniser.
func copyPHPModeRegion(b *strings.Builder, src string, start int) int {
n := len(src)
i := start
for i < n {
if label, bodyStart, ok := phpHeredocOpen(src, i); ok {
end := phpHeredocEnd(src, bodyStart, label)
b.WriteString(src[i:end])
i = end
continue
}
if isPHPQuote(src[i]) {
i = copyPHPString(b, src, i) + 1
continue
}
if src[i] == '/' && i+1 < n && src[i+1] == '*' {
b.WriteString("/*")
i += 2
for i < n {
if src[i] == '*' && i+1 < n && src[i+1] == '/' {
b.WriteString("*/")
i += 2
break
}
b.WriteByte(src[i])
i++
}
continue
}
if isPHPLineCommentStart(src, i) {
end := skipPHPLineComment(src, i)
b.WriteString(src[i:end])
i = end
if i+1 < n && src[i] == '?' && src[i+1] == '>' {
b.WriteString("; ")
return i + 2
}
continue
}
if src[i] == '?' && i+1 < n && src[i+1] == '>' {
b.WriteString("; ")
return i + 2
}
b.WriteByte(src[i])
i++
}
return i
}
// blankInlineHTML writes s as spaces, preserving newlines so line-based
// detectors keep their line boundaries across blanked template text.
func blankInlineHTML(b *strings.Builder, s string) {
for i := 0; i < len(s); i++ {
if s[i] == '\n' || s[i] == '\r' {
b.WriteByte(s[i])
} else {
b.WriteByte(' ')
}
}
}
func nextStandalonePHPCall(code string, searchFrom int, names map[string]struct{}) (int, int, int, bool) {
for i := searchFrom; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
nameStart := i
if code[i] == '\\' {
if i+1 >= len(code) || !isPHPIdentifierStart(code[i+1]) || !canStartGlobalPHPFunction(code, i) {
continue
}
nameStart = i + 1
} else if !isPHPIdentifierStart(code[i]) || !canStartPHPFunctionName(code, i) {
continue
}
nameEnd := nameStart + 1
for nameEnd < len(code) && isPHPIdentifierPart(code[nameEnd]) {
nameEnd++
}
if _, ok := names[strings.ToLower(code[nameStart:nameEnd])]; !ok {
i = nameEnd - 1
continue
}
openParen := skipPHPWhitespace(code, nameEnd)
if openParen >= len(code) || code[openParen] != '(' {
i = nameEnd - 1
continue
}
return i, openParen, matchingParen(code, openParen), true
}
return 0, 0, 0, false
}
func nextSearchOffset(pos, codeLen int) int {
if pos >= codeLen {
return codeLen
}
return pos + 1
}
var requestSuperglobalNames = []string{"$_request", "$_post", "$_get", "$_cookie", "$_server"}
// includeTargetSuperglobalNames omits $_SERVER path keys: including a path
// built from the server document root or script filename is the standard
// WordPress bootstrap idiom, not an LFI/RFI primitive. Header-derived
// $_SERVER keys are handled separately below.
var includeTargetSuperglobalNames = []string{"$_request", "$_post", "$_get", "$_cookie"}
func containsRequestSuperglobal(code string) bool {
return containsAnySuperglobal(code, requestSuperglobalNames)
}
func containsAnySuperglobal(code string, names []string) bool {
code = strings.ToLower(code)
for _, requestVar := range names {
searchFrom := 0
for {
pos := strings.Index(code[searchFrom:], requestVar)
if pos < 0 {
break
}
end := searchFrom + pos + len(requestVar)
if end >= len(code) || !isPHPIdentifierPart(code[end]) {
return true
}
searchFrom = end
}
}
return false
}
// PHP single-quoted strings do not interpolate variables, but double-quoted
// strings do. The decoder callback gate needs that distinction so a literal
// '$_POST' data value does not look like request input.
func containsRequestSuperglobalExpression(code string) bool {
return containsSuperglobalExpression(code, requestSuperglobalNames)
}
// containsIncludeTargetExpression reports whether an include/require target
// expression reads an attacker-controlled superglobal. Non-header $_SERVER
// path keys are left alone (see includeTargetSuperglobalNames).
func containsIncludeTargetExpression(code string) bool {
return containsSuperglobalExpression(code, includeTargetSuperglobalNames) ||
containsServerHeaderIncludeExpression(code)
}
func containsSuperglobalExpression(code string, names []string) bool {
start := 0
for i := 0; i < len(code); i++ {
if !isPHPQuote(code[i]) {
continue
}
if containsAnySuperglobal(code[start:i], names) {
return true
}
end := skipPHPString(code, i)
if code[i] == '"' && containsAnySuperglobal(code[i:end+1], names) {
return true
}
i = end
start = end + 1
}
return containsAnySuperglobal(code[start:], names)
}
func containsServerHeaderIncludeExpression(code string) bool {
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
end := skipPHPString(code, i)
if code[i] == '"' && end > i && containsServerHeaderReference(code[i+1:end]) {
return true
}
i = end
continue
}
dangerous, next, ok := serverIncludeReferenceAt(code, i)
if !ok {
continue
}
if dangerous {
return true
}
i = next - 1
}
return false
}
func containsServerHeaderReference(code string) bool {
for i := 0; i < len(code); i++ {
dangerous, next, ok := serverIncludeReferenceAt(code, i)
if !ok {
continue
}
if dangerous {
return true
}
i = next - 1
}
return false
}
func serverIncludeReferenceAt(code string, start int) (bool, int, bool) {
const serverName = "$_server"
end := start + len(serverName)
if end > len(code) || !strings.EqualFold(code[start:end], serverName) {
return false, start, false
}
if end < len(code) && isPHPIdentifierPart(code[end]) {
return false, start, false
}
bracket := skipPHPWhitespace(code, end)
if bracket >= len(code) || code[bracket] != '[' {
return true, end, true
}
keyStart := skipPHPWhitespace(code, bracket+1)
if keyStart >= len(code) {
return true, len(code), true
}
if isPHPQuote(code[keyStart]) {
keyEnd := skipPHPString(code, keyStart)
key := phpStringLiteralValue(code, keyStart, keyEnd)
next := skipPHPWhitespace(code, keyEnd+1)
if next < len(code) && code[next] == ']' {
return isAttackerControlledServerKey(key), next + 1, true
}
return true, next, true
}
if isPHPIdentifierStart(code[keyStart]) {
keyEnd := keyStart + 1
for keyEnd < len(code) && isPHPIdentifierPart(code[keyEnd]) {
keyEnd++
}
key := code[keyStart:keyEnd]
next := skipPHPWhitespace(code, keyEnd)
if next < len(code) && code[next] == ']' {
return isAttackerControlledServerKey(key), next + 1, true
}
return true, next, true
}
return true, keyStart + 1, true
}
func phpStringLiteralValue(code string, start, end int) string {
var b strings.Builder
quote := code[start]
for i := start + 1; i < end; i++ {
if code[i] != '\\' || i+1 >= end {
b.WriteByte(code[i])
continue
}
i++
esc := code[i]
if quote == '\'' {
if esc == '\'' || esc == '\\' {
b.WriteByte(esc)
} else {
b.WriteByte('\\')
b.WriteByte(esc)
}
continue
}
switch esc {
case 'x':
if i+1 >= end || !isHexDigit(code[i+1]) {
b.WriteByte('\\')
b.WriteByte(esc)
continue
}
value := hexVal(code[i+1])
i++
if i+1 < end && isHexDigit(code[i+1]) {
value = value*16 + hexVal(code[i+1])
i++
}
// #nosec G115 -- PHP hex string escapes are at most one byte.
b.WriteByte(byte(value))
case 'u':
if i+2 >= end || code[i+1] != '{' || !isHexDigit(code[i+2]) {
b.WriteByte('\\')
b.WriteByte(esc)
continue
}
value := 0
j := i + 2
for ; j < end && isHexDigit(code[j]); j++ {
value = value*16 + hexVal(code[j])
}
if j >= end || code[j] != '}' {
b.WriteByte('\\')
b.WriteByte(esc)
continue
}
if value > 0x10ffff {
b.WriteByte('\\')
b.WriteByte(esc)
continue
}
// #nosec G115 -- value is capped at the largest valid Unicode code point.
b.WriteRune(rune(value))
i = j
case '0', '1', '2', '3', '4', '5', '6', '7':
value := int(esc - '0')
digits := 1
for i+1 < end && digits < 3 && code[i+1] >= '0' && code[i+1] <= '7' {
value = value*8 + int(code[i+1]-'0')
i++
digits++
}
// #nosec G115 -- PHP octal string escapes are byte escapes.
b.WriteByte(byte(value))
case 'n':
b.WriteByte('\n')
case 'r':
b.WriteByte('\r')
case 't':
b.WriteByte('\t')
case 'v':
b.WriteByte('\v')
case 'e':
b.WriteByte(0x1b)
case 'f':
b.WriteByte('\f')
case '\\', '$', '"':
b.WriteByte(esc)
default:
b.WriteByte('\\')
b.WriteByte(esc)
}
}
return b.String()
}
func isAttackerControlledServerKey(key string) bool {
key = strings.ToUpper(strings.TrimSpace(key))
if strings.HasPrefix(key, "HTTP_") {
return true
}
switch key {
case "CONTENT_LENGTH", "CONTENT_TYPE", "PHP_AUTH_DIGEST", "PHP_AUTH_PW", "PHP_AUTH_USER":
return true
default:
return false
}
}
func canStartGlobalPHPFunction(code string, slash int) bool {
if slash == 0 {
return true
}
prev := code[slash-1]
return !isPHPIdentifierPart(prev) && prev != '$' && prev != '>' && prev != ':' && prev != '\\'
}
func canStartPHPFunctionName(code string, start int) bool {
if start == 0 {
return true
}
prev := code[start-1]
return !isPHPIdentifierPart(prev) && prev != '$' && prev != '>' && prev != ':' && prev != '\\'
}
// CheckPHPContent scans new/suspicious PHP files for obfuscation patterns,
// remote payload fetching, and eval chains. This is designed to catch droppers
// like the LEVIATHAN attack's file.php and files.php that use goto spaghetti,
// hex-encoded strings, and call_user_func with string-built function names.
//
// This check scans PHP files in directories that shouldn't normally contain
// user-authored PHP: wp-content/languages, wp-content/upgrade, wp-content/mu-plugins,
// and also checks any PHP files flagged by the file index as new.
// phpFileStamp is the cheap content-version key for a scanned PHP file. A file
// whose stamp matches the previous cycle is treated as unchanged.
type phpFileStamp struct {
Mtime int64 `json:"m"`
Size int64 `json:"s"`
// Unlike mtime, change time cannot be restored by a file's owner.
Dev uint64 `json:"d,omitempty"`
Inode uint64 `json:"i,omitempty"`
Ctime int64 `json:"c,omitempty"`
}
func phpFileStampOf(info os.FileInfo) phpFileStamp {
stamp := phpFileStamp{Mtime: info.ModTime().Unix(), Size: info.Size()}
if id, ok := selfWriteIdentityFromFileInfo(info); ok {
stamp.Dev, stamp.Inode = id.Device, id.Inode
stamp.Ctime = id.ChangeSec*1_000_000_000 + id.ChangeNsec
}
return stamp
}
// phpContentNow is indirected so tests can treat fixtures they just wrote as
// settled without sleeping out the change-time window.
var phpContentNow = time.Now
func (stamp phpFileStamp) cacheableAt(start time.Time) bool {
// Missing identity must fail closed, including legacy mtime+size entries.
// A recent read cannot authorize future skips: Linux can give a later
// write the same coarse ctime. Wait a full second before the clean read,
// also covering filesystems with whole-second timestamps. Do not sleep in
// the scan; a later visit will read the file again and can then cache it.
return stamp.Inode != 0 && stamp.Ctime != 0 &&
time.Unix(0, stamp.Ctime).Before(start.Add(-time.Second))
}
// phpContentCache maps a file path to the stamp it carried when last confirmed
// clean. Only clean files are stored, so a present, matching entry means
// "unchanged and previously produced no finding."
type phpContentCache map[string]phpFileStamp
func loadPHPContentCache(stateDir string) phpContentCache {
cache := phpContentCache{}
if stateDir == "" {
return cache
}
data, err := osFS.ReadFile(filepath.Join(stateDir, "phpcontentcache.json"))
if err == nil {
_ = json.Unmarshal(data, &cache)
}
return cache
}
func savePHPContentCache(stateDir string, cache phpContentCache) {
if stateDir == "" {
return
}
data, _ := json.Marshal(cache)
tmpPath := filepath.Join(stateDir, "phpcontentcache.json.tmp")
_ = os.WriteFile(tmpPath, data, 0600)
_ = os.Rename(tmpPath, filepath.Join(stateDir, "phpcontentcache.json"))
}
// A timed-out check can outlive the runner's drain grace. Only the latest
// host scan may publish its snapshot; an older scan must not restore clean
// stamps that its successor invalidated. Account and audit scans do not write
// this cache and therefore do not take ownership away from a host scan.
var phpContentCacheWriter struct {
sync.Mutex
generation uint64
}
func beginPHPContentCacheRun() uint64 {
phpContentCacheWriter.Lock()
defer phpContentCacheWriter.Unlock()
phpContentCacheWriter.generation++
return phpContentCacheWriter.generation
}
func savePHPContentCacheRun(stateDir string, scan *phpContentScan, generation uint64) {
phpContentCacheWriter.Lock()
defer phpContentCacheWriter.Unlock()
if generation == phpContentCacheWriter.generation {
savePHPContentCache(stateDir, scan.merged())
}
}
// phpContentHostScanCount drives a periodic forced full rescan that bypasses
// the content cache, mirroring the file-index cadence. This remains a backstop
// for filesystems whose metadata does not reliably distinguish content changes.
var phpContentHostScanCount int32
// phpContentAccountScanCount keeps account-scoped scans from consuming the
// host-wide full-rescan cadence. Account scans are subsets and never save the
// shared cache, so mixing the counters can make the host-wide backstop miss its
// intended cycle.
var phpContentAccountScanCount int32
func phpContentForceFull(ctx context.Context) bool {
if AccountFromContext(ctx) != "" {
return atomic.AddInt32(&phpContentAccountScanCount, 1)%6 == 0
}
return atomic.AddInt32(&phpContentHostScanCount, 1)%6 == 0
}
// phpContentScan carries the per-cycle cache state through the recursive walk.
// prev holds prior clean stamps not yet invalidated by this run; next holds
// files confirmed clean during this run.
type phpContentScan struct {
cfg *config.Config
prev phpContentCache
next phpContentCache
forceFull bool
// visited holds directories whose full entry list this run read. Their
// prior stamps are authoritative-by-absence: a cached file the walk did
// not see again is gone. Everything else is simply unreached.
visited map[string]bool
}
func newPHPContentScan(cfg *config.Config, prev phpContentCache, forceFull bool) *phpContentScan {
if prev == nil {
prev = phpContentCache{}
}
return &phpContentScan{
cfg: cfg,
prev: prev,
next: phpContentCache{},
forceFull: forceFull,
visited: map[string]bool{},
}
}
// merged is the cache to persist after a run that may have been cut short. It
// keeps every stamp confirmed this cycle and carries forward prior stamps from
// directories the run never finished reading. A run that walked a directory
// prunes what vanished from it; a run that never reached one leaves its files
// cached so the next cycle does not re-read the whole host from scratch.
func (s *phpContentScan) merged() phpContentCache {
out := make(phpContentCache, len(s.next)+len(s.prev))
for path, stamp := range s.prev {
// The periodic forced pass must also expire stamps outside its fixed
// directories, so rolling coverage re-reads them on its next visit.
if s.forceFull || s.visited[filepath.Dir(path)] {
continue
}
out[path] = stamp
}
for path, stamp := range s.next {
out[path] = stamp
}
return out
}
// pruneMissing drops cached stamps under roots for paths that are no longer on
// disk. Enumeration can omit unreadable directories and remapped extensions,
// so absence from its list alone is not proof that a cached path is gone.
func (s *phpContentScan) pruneMissing(roots []string, present []string) {
live := make(map[string]struct{}, len(present))
for _, path := range present {
live[path] = struct{}{}
}
for path := range s.prev {
if _, ok := live[path]; ok {
continue
}
for _, root := range roots {
if strings.HasPrefix(path, root+string(filepath.Separator)) {
if _, err := osFS.Stat(path); os.IsNotExist(err) {
delete(s.prev, path)
}
break
}
}
}
}
// accountDocRoots lists an account's document roots: public_html plus the
// addon-domain directories beside it, skipping the mail, control and temp
// directories that never serve web content.
func accountDocRoots(homeDir string) []string {
docRoots := []string{filepath.Join(homeDir, "public_html")}
subDirs, _ := osFS.ReadDir(homeDir)
for _, sd := range subDirs {
if sd.IsDir() && sd.Name() != "public_html" && sd.Name() != "mail" &&
!strings.HasPrefix(sd.Name(), ".") && sd.Name() != "etc" &&
sd.Name() != "logs" && sd.Name() != "ssl" && sd.Name() != "tmp" {
docRoots = append(docRoots, filepath.Join(homeDir, sd.Name()))
}
}
return docRoots
}
func CheckPHPContent(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
homeDirs, err := GetScanHomeDirs(ctx)
if err != nil {
return nil
}
forcedFull := phpContentForceFull(ctx) || scanForceContent(ctx)
persistCache := AccountFromContext(ctx) == "" && !scanForceContent(ctx)
var cacheGeneration uint64
if persistCache {
cacheGeneration = beginPHPContentCacheRun()
}
scan := newPHPContentScan(cfg, loadPHPContentCache(cfg.StatePath), forcedFull)
// A run cut short still falls through to the cache write below, so the
// stamps it confirmed are not thrown away with the rest of the cycle.
accounts:
for _, homeEntry := range homeDirs {
if ctx.Err() != nil {
break
}
if !homeEntry.IsDir() {
continue
}
homeDir := scanHomeDirPath(homeEntry)
docRoots := accountDocRoots(homeDir)
for _, docRoot := range docRoots {
// Scan directories that shouldn't contain user PHP
suspiciousDirs := []string{
filepath.Join(docRoot, "wp-content", "languages"),
filepath.Join(docRoot, "wp-content", "upgrade"),
filepath.Join(docRoot, "wp-content", "mu-plugins"),
filepath.Join(docRoot, "wp-content", "plugins"),
filepath.Join(docRoot, "wp-content", "themes"),
}
for _, dir := range suspiciousDirs {
scan.scanDir(ctx, dir, 4, phpHandlerOverlay{}, &findings)
if ctx.Err() != nil {
break accounts
}
}
}
}
if rollingContentEnabled(ctx, cfg, forcedFull) {
rollingContentPass(ctx, cfg, scan, homeDirs, &findings)
}
// Persist on every host-wide run, including one cut short by the check
// budget: scan.merged() carries forward the stamps for directories this run
// never read, so a host too large to finish in one cycle still makes
// progress instead of re-reading everything next cycle. An account-scoped
// run (account_scan) only walks one account, so its stamps must not
// overwrite the host-wide cache. A forced-content scan (ForceContent=true)
// re-reads every file regardless of the cache, so scan.next reflects only
// the files visited this run and must not overwrite the host-wide live cache.
if persistCache {
savePHPContentCacheRun(cfg.StatePath, scan, cacheGeneration)
}
return findings
}
// scanDirForObfuscatedPHP scans dir without the content cache: every PHP file
// is read and analysed. Used where caching does not apply (no prior cycle to
// compare against).
func scanDirForObfuscatedPHP(ctx context.Context, dir string, maxDepth int, cfg *config.Config, findings *[]alert.Finding) {
newPHPContentScan(cfg, nil, true).scanDir(ctx, dir, maxDepth, phpHandlerOverlay{}, findings)
}
// scanDir recursively scans dir for PHP files with malicious content patterns.
// A file that was clean last cycle and has an unchanged stable stamp skips the
// read+parse, unless this is a forced full rescan. Files that produce a finding
// are never cached, so they re-surface on every cycle for the alert pipeline.
func (s *phpContentScan) scanDir(ctx context.Context, dir string, maxDepth int, overlay phpHandlerOverlay, findings *[]alert.Finding) {
if ctx.Err() != nil {
return
}
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
// A directory that is gone has no files left to cache. One we cannot
// read may still hold them, so its stamps stay.
if os.IsNotExist(err) {
s.visited[dir] = true
}
return
}
// Layer this directory's .htaccess PHP handler remappings onto the set
// inherited from parent directories. An attacker who maps a non-PHP
// extension to the PHP interpreter (the LEVIATHAN .htaccess trick) or who
// SetHandlers the whole directory must not be able to hide executable PHP
// behind a name the stock handler would not run. Read once per directory.
if htaccess, herr := osFS.ReadFile(filepath.Join(dir, ".htaccess")); herr == nil {
overlay = overlay.mergeHtaccess(htaccess)
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
// Check suppressed paths (bypassed for explicit full-scan / audit requests).
suppressed := false
if scanRespectsIgnores(ctx, s.cfg) {
for _, ignore := range s.cfg.Suppressions.IgnorePaths {
if matchGlob(fullPath, ignore) {
suppressed = true
break
}
}
}
if suppressed {
continue
}
if entry.IsDir() {
s.scanDir(ctx, fullPath, maxDepth-1, overlay, findings)
continue
}
s.scanFile(ctx, fullPath, overlay, findings)
}
s.visited[dir] = true
}
// scanFile content-analyses one PHP file using the same cache, size-guard, and
// finding logic as scanDir's per-entry loop. It re-stats fullPath (it no longer
// has the DirEntry) and decides whether it is PHP source or executes under the
// overlay. Files that produce a finding are never cached, so they re-surface
// each cycle for the alert pipeline. scanDir calls this for every non-directory
// entry; the rolling driver calls it for each path in its bounded slice.
func (s *phpContentScan) scanFile(ctx context.Context, fullPath string, overlay phpHandlerOverlay, findings *[]alert.Finding) {
// Once attempted, only a fresh clean result may retain a stamp. Rolling
// and interrupted directory walks cannot rely on directory completion to
// invalidate a finding, read failure, or an earlier read in the same run.
previous, cached := s.prev[fullPath]
delete(s.next, fullPath)
defer func() {
if _, clean := s.next[fullPath]; !clean {
delete(s.prev, fullPath)
}
}()
nameLower := strings.ToLower(filepath.Base(fullPath))
if !contenttype.IsPHPSourceName(nameLower) && !overlay.executes(nameLower) {
return
}
started := phpContentNow()
info, statErr := osFS.Stat(fullPath)
var stamp phpFileStamp
canCache := statErr == nil
if canCache {
stamp = phpFileStampOf(info)
canCache = stamp.cacheableAt(started)
}
// Full-scan file-size guard: when the caller sets a per-file byte cap
// (MaxFileBytes > 0), skip content analysis for oversized files and
// record a job-scoped warning instead. The warning is persisted in
// scan_job_findings by the scan-job manager and never dispatched to the
// live alert pipeline. Normal scheduled scans set MaxFileBytes=0 and
// reach this block unchanged.
if limit := scanMaxFileBytes(ctx); limit > 0 && info != nil && info.Size() > limit {
*findings = append(*findings, alert.Finding{
Severity: alert.Warning,
Check: "full_scan_file_too_large",
Message: fmt.Sprintf("Full scan skipped oversized file: %s", fullPath),
Details: fmt.Sprintf("Size: %d, Limit: %d", info.Size(), limit),
FilePath: fullPath,
})
return
}
f, err := osFS.Open(fullPath)
if err != nil {
return
}
defer func() { _ = f.Close() }()
opened, openStatErr := f.Stat()
canCache = canCache && openStatErr == nil && opened.Mode().IsRegular() && phpFileStampOf(opened) == stamp
// Opening alone proves readability, not identity: a path can be swapped
// between Stat and Open. Only reuse the clean result for the same inode.
if canCache && !s.forceFull && cached && previous == stamp {
s.next[fullPath] = stamp
return
}
var size int64 = -1
if openStatErr == nil {
size = opened.Size()
}
// Every PHP source file is content-analysed. No filename/path allowlist:
// clean files produce no finding, so there is no benefit to skipping
// them, and any skip is a place an attacker can hide a backdoor.
result := analyzePHPContentReaderAt(fullPath, f, size)
if result.severity >= 0 {
details := result.details
if info != nil {
details += fmt.Sprintf("\nSize: %d, Mtime: %s", info.Size(), info.ModTime().Format("2006-01-02 15:04:05"))
}
*findings = append(*findings, alert.Finding{
Severity: result.severity,
Check: result.check,
Message: fmt.Sprintf("%s: %s", result.message, fullPath),
Details: details,
FilePath: fullPath,
ContentSHA256: FileContentSHA256(fullPath),
DetectLogic: ContentDetectionVersion(),
})
return
}
// A clean read only validates the stamp if both the path and the opened
// inode still match. Never attach one file's result to another's stamp or
// retain a clean result for an inode changed during the read.
if canCache && result.readOK {
afterPath, pathErr := osFS.Stat(fullPath)
afterFile, fileErr := f.Stat()
if pathErr == nil && fileErr == nil &&
phpFileStampOf(afterPath) == stamp && phpFileStampOf(afterFile) == stamp {
s.next[fullPath] = stamp
}
}
}
type phpAnalysisResult struct {
severity alert.Severity
check string
message string
details string
indicators []string
readOK bool
empty bool
}
// analyzePHPContent reads a PHP file's head window (and tail window for large
// files) and checks for obfuscation and malicious patterns.
func analyzePHPContent(path string) phpAnalysisResult {
f, err := osFS.Open(path)
if err != nil {
return phpAnalysisResult{severity: -1}
}
defer func() { _ = f.Close() }()
var size int64 = -1
if info, statErr := f.Stat(); statErr == nil {
size = info.Size()
}
return analyzePHPContentReaderAt(path, f, size)
}
func analyzePHPContentWithFingerprint(path string) (phpAnalysisResult, string) {
f, err := osFS.Open(path)
if err != nil {
return phpAnalysisResult{severity: -1}, ""
}
defer func() { _ = f.Close() }()
var size int64 = -1
info, statErr := f.Stat()
if statErr == nil {
size = info.Size()
}
if statErr == nil && info.Mode().IsRegular() && size <= contentFingerprintMaxBytes {
data, err := io.ReadAll(io.LimitReader(f, contentFingerprintMaxBytes+1))
if err != nil {
return phpAnalysisResult{severity: -1, readOK: false}, ""
}
after, err := f.Stat()
if err != nil || !sameCleanContentShape(after, info) || int64(len(data)) != size {
return phpAnalysisResult{severity: -1, readOK: false}, ""
}
result := analyzePHPContentReaderAt(path, bytes.NewReader(data), int64(len(data)))
if !result.readOK {
return result, ""
}
sum := sha256.Sum256(data)
return result, fmt.Sprintf("%x", sum)
}
return analyzePHPContentReaderAt(path, f, size), ""
}
func analyzePHPContentReaderAt(path string, f io.ReaderAt, size int64) phpAnalysisResult {
head, tail, tailOffset, readOK := readPHPContentWindows(f, size)
if !readOK {
return phpAnalysisResult{severity: -1, readOK: false}
}
if len(head) == 0 && len(tail) == 0 {
return phpAnalysisResult{severity: -1, readOK: true, empty: true}
}
if len(tail) > 0 && tailOffset > int64(len(head)) {
// The skipped bytes may close a quote or comment that starts in the
// head. Score windows independently so head parser state cannot hide
// executable tail code we did read.
return mergePHPAnalysisResults(
analyzePHPCode(path, phpCodeOnly(string(head)), readOK),
analyzePHPCode(path, phpCodeOnly("<?php\n"+string(tail)), readOK),
)
}
return analyzePHPCode(path, phpCodeOnlyWindows(head, tail), readOK)
}
func mergePHPAnalysisResults(results ...phpAnalysisResult) phpAnalysisResult {
readOK := true
var indicators []string
for _, result := range results {
readOK = readOK && result.readOK
if result.severity >= 0 {
indicators = append(indicators, result.indicators...)
}
}
return phpAnalysisFromIndicators(indicators, readOK)
}
func analyzePHPCode(path, content string, readOK bool) phpAnalysisResult {
contentLower := strings.ToLower(content)
var indicators []string
// --- Critical: Remote payload fetching ---
// Paste sites are always suspicious in PHP files.
// GitHub raw URLs are common in legitimate plugin update checkers,
// so only count them as indicators when they appear on the same line
// as a dangerous PHP function call.
pasteHosts := []string{
"pastebin.com/raw",
"paste.ee/r/",
"ghostbin.co/paste/",
"hastebin.com/raw/",
}
for _, host := range pasteHosts {
if strings.Contains(contentLower, host) {
indicators = append(indicators, fmt.Sprintf("remote payload URL: %s", host))
}
}
githubHosts := []string{"gist.githubusercontent.com", "raw.githubusercontent.com"}
dangerousCalls := []string{"file_put_contents(", "fwrite(", "shell_", "passthru(", "popen("}
for _, host := range githubHosts {
if !strings.Contains(contentLower, host) {
continue
}
// Same-line = strong signal (critical)
sameLine := false
for _, line := range strings.Split(contentLower, "\n") {
if !strings.Contains(line, host) {
continue
}
for _, fn := range dangerousCalls {
if strings.Contains(line, fn) {
indicators = append(indicators, fmt.Sprintf("remote payload URL with dangerous call: %s", host))
sameLine = true
break
}
}
if sameLine {
break
}
}
// Co-presence (different lines, same 32 KB window) was previously
// emitted as a weaker indicator. It generated standing FPs on
// legit plugins that fetch upstream resources from github mirrors
// (wp-statistics GeoLite2 updates, unyson font fetcher, polylang
// language packs). Same-line is the strong signal kept above; the
// co-presence path is removed entirely.
}
// --- Critical: eval() chains with decoding ---
// Only flag when eval directly wraps a decoder (structural nesting),
// not when they merely co-exist in the same file (which causes false
// positives on legitimate plugins that use eval for templates and
// base64_decode for unrelated data processing).
decoders := []string{
"base64_decode", "gzinflate", "gzuncompress", "str_rot13",
"rawurldecode", "gzdecode", "bzdecompress",
// Droppers rotate through every reversible transform PHP ships;
// the four below were the unlisted ones seen wrapping eval in the
// wild (hex payloads, reversed strings, URL-encoded blobs, uuencode).
"hex2bin", "strrev", "urldecode", "convert_uudecode",
}
hasDecoder := false
hasNestedEvalDecode := false
for _, d := range decoders {
if strings.Contains(contentLower, d) {
hasDecoder = true
}
}
// Check for structural nesting: eval(base64_decode(...)), eval(gzinflate(...)), etc.
// PHP tolerates inline comments and arbitrary whitespace (including line
// breaks) between the keyword and its open paren, so a naive line-by-line
// `eval(` substring scan misses `eval /*x*/ ( base64_decode(...))` and
// `eval // bypass\n( base64_decode(...))`. Strip PHP comments and
// strings first, then match the structural pattern across whitespace,
// and require the inner callee to be one of the known decoders /
// decompressors.
commentStripped := stripPHPCommentsFromCode(content)
codeLower := strings.ToLower(stripPHPStringsFromCode(commentStripped))
for _, m := range nestedEvalDecodeRe.FindAllStringSubmatch(codeLower, -1) {
if len(m) < 3 {
continue
}
inner := m[2]
for _, d := range decoders {
if inner == d {
hasNestedEvalDecode = true
break
}
}
if hasNestedEvalDecode {
break
}
}
if hasNestedEvalDecode {
indicators = append(indicators, "eval() directly wrapping encoding/compression function")
}
if reEvalRequestInput.MatchString(codeLower) {
indicators = append(indicators, "eval() of request input")
}
// eval wrapping dynamic code construction the decoder loop above
// ignores: a variable callee (eval($f(...))) or a code-building
// primitive (eval(create_function(...)), eval(call_user_func(...))).
// These never appear in legitimate user-directory PHP; a single hit is
// surfaced as a High signal (the >=2 gate still governs quarantine).
hasEvalExecWrap := reEvalVarCallee.MatchString(codeLower)
if !hasEvalExecWrap {
for _, m := range nestedEvalDecodeRe.FindAllStringSubmatch(codeLower, -1) {
if len(m) < 3 {
continue
}
if m[1] != "eval" {
continue
}
if _, ok := evalExecWrapInner[m[2]]; ok {
hasEvalExecWrap = true
break
}
}
}
if hasEvalExecWrap {
indicators = append(indicators, "eval() wrapping a dynamic code-execution primitive")
}
// Backtick shell execution with request input -- `...$_GET...`.
// Match only executable backtick spans, not quoted examples.
if hasBacktickSuperglobal(commentStripped) {
indicators = append(indicators, "backtick shell execution with request input")
}
// Callback-position exec: an exec/decoder function name passed as a
// string callback (array_map("system", ...), register_shutdown_function(
// "passthru", ...)). Runs on the comment-stripped source so the literal
// callback name is preserved.
if hasCallbackExecName(commentStripped) {
indicators = append(indicators, "exec/decoder function name passed as a callback")
}
// Variable-variable / dynamic-expression function call co-located with
// request input on the same line -- $$h($_GET[...]) -- a dynamic-dispatch
// RCE shape. The same-line request-var gate keeps benign dispatcher code
// (which uses $$var without attacker input) from tripping.
for _, line := range strings.Split(codeLower, "\n") {
if reVarVarCall.MatchString(line) && lineContainsRequestVar(line) {
indicators = append(indicators, "variable-variable function call with request input")
break
}
}
// preg_replace() with the /e modifier evaluates its replacement as PHP.
// Removed in PHP 7.0; flag it when request input reaches the evaluated
// replacement or its subject backreferences.
if hasPregReplaceEvalWithRequest(commentStripped) {
indicators = append(indicators, "preg_replace with /e modifier (code execution)")
}
// include/require of request input or a remote/stream wrapper -- LFI,
// RFI, and php://input code execution.
if hasDangerousInclude(commentStripped) {
indicators = append(indicators, "include/require of request input or remote/data wrapper")
}
// assert()/create_function() driven by request input -- both evaluate a
// string argument as PHP.
if hasCodeEvalPrimitiveWithRequest(commentStripped) {
indicators = append(indicators, "code-eval primitive (assert/create_function) with request input")
}
// --- Critical: call_user_func with string-built function names ---
// LEVIATHAN droppers build the target function name on the call itself:
// call_user_func("\x63"."\x75"."\x72"."\x6c", $payload) == call_user_func("curl", ...)
// File-wide hex/concat counts are unsafe here: WPML bundles PHPZip
// (inc/wpml_zip.php) which declares 20+ ZIP-format signature constants
// as hex literals ("\x50\x4b\x03\x04" etc.) and makes a single benign
// call_user_func(self::$temp) call to invoke a temp-file factory.
// Match the obfuscation on the callable target argument itself.
if strings.Contains(contentLower, "call_user_func") {
foundCallUserFuncObfuscation := false
for _, line := range strings.Split(content, "\n") {
if !strings.Contains(strings.ToLower(line), "call_user_func") {
continue
}
for _, targetArg := range phpCallUserFuncTargetArgs(line) {
// PHP 7+ accepts both "\xNN" hex and "\u{NN}" unicode-codepoint
// escapes inside double-quoted strings. Treat them as
// equivalent obfuscation forms so an attacker cannot bypass
// the detector by swapping syntax.
lineHex := countOccurrences(targetArg, `"\x`) + countOccurrences(targetArg, `"\u{`)
lineConcat := countOccurrences(targetArg, `" . "`) + countOccurrences(targetArg, `"."`)
// Typical shortest obfuscated name is 3-4 bytes ("exec", "curl",
// "eval"); require >=3 escapes AND >=2 concatenations on the
// call target argument.
if lineHex >= 3 && lineConcat >= 2 {
indicators = append(indicators, "call_user_func with obfuscated function names")
foundCallUserFuncObfuscation = true
break
}
}
if foundCallUserFuncObfuscation {
break
}
}
}
// --- High: Goto obfuscation (LEVIATHAN signature) ---
// Counting goto statements alone measures the wrong thing. WordPress
// core's HTML API drives the HTML5 insertion-mode state machine with goto
// and names every label after its spec section, so a plain count reports
// authentic core on every site of every account. A descriptive label is
// evidence against obfuscation: an obfuscator emits generated labels
// precisely because they carry no meaning.
//
// Label shape is not enough on its own either. Commercial obfuscators
// sold to plugin vendors emit the same generated labels, and their
// output carries no payload: a paid-for plugin looked exactly like a
// dropper. Malware still has to decode, execute, or read request input
// somewhere, so both branches want a sink. This mirrors the evidence
// the php_goto_obfuscation signature requires.
var generatedGotos, alphaGotos int
for _, m := range reGotoLabel.FindAllStringSubmatch(content, -1) {
if gotoLabelIsGenerated(m[1]) {
generatedGotos++
} else {
alphaGotos++
}
}
switch {
case generatedGotos > 8 && reGotoExecSink.MatchString(content):
indicators = append(indicators, fmt.Sprintf("goto obfuscation (%d generated labels)", generatedGotos))
case alphaGotos > 10 && reGotoExecSink.MatchString(content):
indicators = append(indicators, fmt.Sprintf("goto obfuscation (%d labels reaching an execution sink)", alphaGotos))
}
// --- High: Hex-encoded string construction ---
// Only flag hex strings when accompanied by concatenation - real obfuscation
// builds function names like "\x63" . "\x75" . "\x72" . "\x6c" (= "curl").
// Standalone hex arrays (Wordfence IPv6 subnet masks, binary data) are benign.
hexStringCount := countOccurrences(content, `"\x`)
dotConcatCount := countOccurrences(content, `" . "`)
if hexStringCount > 20 && dotConcatCount > 10 {
indicators = append(indicators, fmt.Sprintf("heavy hex-encoded strings with concatenation (%d hex, %d concat - obfuscation pattern)", hexStringCount, dotConcatCount))
}
// The standalone "concat>30 alone" branch was removed: WordPress
// themes and page builders concatenate literal CSS/HTML tokens
// dozens of times in dynamic style/markup builders (sydney theme,
// elementor, beaver builder), producing FPs on every install. Real
// function-name obfuscation always pairs concat with hex escapes
// and is still caught by the combined branch above.
// --- Critical: variable-function indirection that resolves to a
// decoder. Attackers slip past the literal "eval(base64_decode("
// detector by binding the dangerous name to a variable on one line
// and calling it on another:
// $d = "base64_decode";
// $r = "eval";
// $r($d("AAAA"));
// The heuristic looks for an assignment $var = "decoder_or_exec"
// followed by a $var( invocation in the same file. Hits must
// reference at least one decoder OR one shell-exec primitive,
// since plain variable function calls show up in legitimate
// metaprogramming.
if detectVarFuncDangerousAssignment(content) {
indicators = append(indicators, "variable function name resolves to decoder or exec primitive")
}
// --- High: Variable function calls with obfuscated names ---
// call_user_func + decoder alone is too broad - Elementor, WooCommerce, and
// dozens of plugins use call_user_func_array with base64_decode legitimately.
// Only flag when combined with hex-built callable target obfuscation.
if strings.Contains(contentLower, "call_user_func") && hasDecoder {
if hasCallUserFuncHexNameBuild(content) {
indicators = append(indicators, "variable function call with decoder and obfuscation")
}
}
// --- High: Shell execution functions combined with request input ---
// Uses containsStandaloneFunc to avoid substring false positives
// (e.g. "WP_Filesystem(" matching "exec(", "preg_match(" matching "exec(")
shellFuncs := []string{"system(", "passthru(", "exec(", "shell_exec(", "popen(", "proc_open(", "pcntl_exec("}
// Two-tier detection:
// Same line = CRITICAL signal (auto-quarantine eligible)
// Co-presence = HIGH signal (alert only, not quarantined alone)
// This prevents bypass by splitting across lines while avoiding
// false-positive quarantine of legitimate plugins.
hasShellFunc := false
sameLineShellRequest := false
for _, sf := range shellFuncs {
if containsStandaloneFunc(contentLower, sf) {
hasShellFunc = true
break
}
}
hasRequestVar := containsRequestSuperglobal(contentLower)
// Same-line is a strong signal on its own: "$ret = system($_POST['cmd']);"
// almost never occurs in legitimate code. Co-presence is weaker: elFinder
// and other media-processing libraries legitimately call exec() for
// ImageMagick and also consume $_POST for AJAX routing, placing both
// tokens in the same 32 KB window. The co-presence finding is therefore
// only emitted as CORROBORATION after all the stronger indicators have
// been collected -- see the deferred append further below.
coPresenceCandidate := false
if hasShellFunc && hasRequestVar {
for _, line := range strings.Split(contentLower, "\n") {
lineHasShell := false
for _, sf := range shellFuncs {
if containsStandaloneFunc(line, sf) {
lineHasShell = true
break
}
}
if !lineHasShell {
continue
}
if containsRequestSuperglobal(line) {
sameLineShellRequest = true
break
}
}
if sameLineShellRequest {
indicators = append(indicators, "shell function with request input on same line")
} else if !IsVerifiedCMSFile(path) {
coPresenceCandidate = true
}
}
// --- High: base64 encoding/decoding with execution on same line ---
if strings.Contains(contentLower, "base64_decode") && strings.Contains(contentLower, "base64_encode") {
for _, line := range strings.Split(contentLower, "\n") {
hasBoth := strings.Contains(line, "base64_decode") && strings.Contains(line, "base64_encode")
hasExec := false
for _, sf := range shellFuncs {
if containsStandaloneFunc(line, sf) {
hasExec = true
break
}
}
if hasBoth && hasExec {
indicators = append(indicators, "base64 encode+decode with execution on same line (command relay)")
break
}
}
}
// Deferred corroboration: a lone co-presence is not enough. If a
// stronger indicator was produced above, the co-presence is appended
// both as extra context for the operator and to nudge the severity
// into the >=2 Critical band for obfuscated droppers.
if coPresenceCandidate && len(indicators) > 0 {
indicators = append(indicators, "shell function co-present with request input")
}
// --- Determine severity based on indicators ---
return phpAnalysisFromIndicators(indicators, readOK)
}
func phpAnalysisFromIndicators(indicators []string, readOK bool) phpAnalysisResult {
indicators = dedupePHPIndicators(indicators)
if len(indicators) == 0 {
return phpAnalysisResult{severity: -1, readOK: readOK}
}
// Auto-quarantine in autoresponse.AutoQuarantineFiles acts only on
// Critical findings. A single heuristic indicator has false-positive
// classes severe enough to rm live production files (WPML's PHPZip
// tripped the former "call_user_func with obfuscated" bypass on hex
// constants that build ZIP magic bytes; legitimate plugins embed
// pastebin URLs in support docstrings and release notes). Require
// two converging indicators before the severity crosses the
// destructive-action threshold; single hits surface as High and stay
// in the operator queue.
if len(indicators) >= 2 {
return phpAnalysisResult{
severity: alert.Critical,
check: "obfuscated_php",
message: "Obfuscated/malicious PHP detected",
details: fmt.Sprintf("Indicators found:\n- %s", strings.Join(indicators, "\n- ")),
readOK: readOK,
indicators: indicators,
}
}
return phpAnalysisResult{
severity: alert.High,
check: "suspicious_php_content",
message: "Suspicious PHP content detected",
details: fmt.Sprintf("Indicators found:\n- %s", strings.Join(indicators, "\n- ")),
readOK: readOK,
indicators: indicators,
}
}
func dedupePHPIndicators(indicators []string) []string {
seen := make(map[string]struct{}, len(indicators))
out := indicators[:0]
for _, indicator := range indicators {
if _, ok := seen[indicator]; ok {
continue
}
seen[indicator] = struct{}{}
out = append(out, indicator)
}
return out
}
func countOccurrences(s, substr string) int {
count := 0
offset := 0
for {
idx := strings.Index(s[offset:], substr)
if idx < 0 {
break
}
count++
offset += idx + len(substr)
}
return count
}
type hexNameAssignment struct {
obfuscated bool
hexEscapes int
concatOps int
pos int
}
func hasCallUserFuncHexNameBuild(content string) bool {
code := stripPHPCommentsFromCode(content)
assignments := findHexNameAssignments(code)
searchFrom := 0
for {
callStart, openParen, closeParen, ok := nextStandalonePHPCall(code, searchFrom, callUserFuncCallNames)
if !ok {
return false
}
arg, argOK := firstPHPCallArgument(code, openParen+1)
if argOK && phpExprHasHexNameBuild(arg) {
return true
}
if argOK {
if variable, varOK := singlePHPVariableExpr(arg); varOK && hexNameBuildAt(assignments[variable], callStart) {
return true
}
}
searchFrom = nextSearchOffset(closeParen, len(code))
}
}
func findHexNameAssignments(code string) map[string][]hexNameAssignment {
assignments := map[string][]hexNameAssignment{}
for i := 0; i < len(code); i++ {
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
if code[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(code, i)
if !ok {
continue
}
j := skipPHPWhitespace(code, next)
opLen, directAssign, appendAssign, ok := phpAssignmentOperator(code, j)
if !ok {
i = next - 1
continue
}
exprStart := skipPHPWhitespace(code, j+opLen)
exprEnd := phpExpressionEnd(code, exprStart)
hexEscapes := 0
concatOps := 0
if directAssign || appendAssign {
hexEscapes = countPHPStringHexEscapes(code[exprStart:exprEnd])
concatOps = countPHPConcatOperators(code[exprStart:exprEnd])
}
if appendAssign {
prev := lastHexNameAssignment(assignments[variable])
hexEscapes += prev.hexEscapes
concatOps += prev.concatOps + 1
}
assignments[variable] = append(assignments[variable], hexNameAssignment{
obfuscated: phpExprCountsHaveHexNameBuild(hexEscapes, concatOps),
hexEscapes: hexEscapes,
concatOps: concatOps,
pos: exprEnd,
})
i = next - 1
}
return assignments
}
func phpAssignmentOperator(code string, pos int) (int, bool, bool, bool) {
if pos >= len(code) {
return 0, false, false, false
}
if code[pos] == '=' && (pos+1 >= len(code) || (code[pos+1] != '=' && code[pos+1] != '>')) {
return 1, true, false, true
}
if pos+1 < len(code) && code[pos] == '.' && code[pos+1] == '=' {
return 2, false, true, true
}
if pos+1 < len(code) && strings.ContainsRune("+-*/%&|^", rune(code[pos])) && code[pos+1] == '=' {
return 2, false, false, true
}
if pos+2 < len(code) {
op := code[pos : pos+3]
if op == "??=" || op == "<<=" || op == ">>=" {
return 3, false, false, true
}
}
return 0, false, false, false
}
func lastHexNameAssignment(assignments []hexNameAssignment) hexNameAssignment {
if len(assignments) == 0 {
return hexNameAssignment{}
}
return assignments[len(assignments)-1]
}
func hexNameBuildAt(assignments []hexNameAssignment, callPos int) bool {
var last hexNameAssignment
found := false
for _, assignment := range assignments {
if assignment.pos > callPos {
break
}
last = assignment
found = true
}
return found && last.obfuscated
}
func singlePHPVariableExpr(expr string) (string, bool) {
expr = strings.TrimSpace(expr)
if expr == "" || expr[0] != '$' {
return "", false
}
variable, next, ok := readPHPVariableName(expr, 0)
if !ok || skipPHPWhitespace(expr, next) != len(expr) {
return "", false
}
return variable, true
}
func phpExprHasHexNameBuild(expr string) bool {
return phpExprCountsHaveHexNameBuild(countPHPStringHexEscapes(expr), countPHPConcatOperators(expr))
}
func phpExprCountsHaveHexNameBuild(hexEscapes, concatOps int) bool {
return hexEscapes >= 3 && concatOps >= 2
}
func countPHPStringHexEscapes(expr string) int {
count := 0
for i := 0; i < len(expr); i++ {
if expr[i] != '"' {
if isPHPQuote(expr[i]) {
i = skipPHPString(expr, i)
}
continue
}
end := skipPHPString(expr, i)
if end > i {
literal := strings.ToLower(expr[i+1 : end])
count += countOccurrences(literal, `\x`)
count += countOccurrences(literal, `\u{`)
}
i = end
}
return count
}
func countPHPConcatOperators(expr string) int {
count := 0
for i := 0; i < len(expr); i++ {
if isPHPQuote(expr[i]) {
i = skipPHPString(expr, i)
continue
}
if expr[i] == '.' {
count++
}
}
return count
}
func phpCallUserFuncTargetArgs(line string) []string {
lower := strings.ToLower(line)
var args []string
for _, name := range []string{"call_user_func_array", "call_user_func"} {
searchFrom := 0
for {
pos := strings.Index(lower[searchFrom:], name)
if pos < 0 {
break
}
pos += searchFrom
next := pos + len(name)
searchFrom = next
if !isPHPFuncNameBoundary(line, pos, next) {
continue
}
i := skipPHPSpaceString(line, next)
if i >= len(line) || line[i] != '(' {
continue
}
if arg, ok := firstPHPCallArgument(line, i+1); ok {
args = append(args, arg)
}
}
}
return args
}
func isPHPFuncNameBoundary(s string, start, end int) bool {
if start > 0 {
prev := s[start-1]
if isIdentCont(prev) || prev == '>' || prev == ':' {
return false
}
}
return end >= len(s) || !isIdentCont(s[end])
}
func skipPHPSpaceString(s string, i int) int {
for i < len(s) && isPHPSpace(s[i]) {
i++
}
return i
}
func firstPHPCallArgument(s string, start int) (string, bool) {
start = skipPHPSpaceString(s, start)
i := start
depth := 0
var quote byte
escaped := false
for i < len(s) {
c := s[i]
if quote != 0 {
if escaped {
escaped = false
i++
continue
}
if c == '\\' {
escaped = true
i++
continue
}
if c == quote {
quote = 0
}
i++
continue
}
switch c {
case '"', '\'':
quote = c
case '(', '[', '{':
depth++
case ')':
if depth == 0 {
return s[start:i], true
}
depth--
case ',':
if depth == 0 {
return s[start:i], true
}
}
i++
}
return "", false
}
// containsStandaloneFunc reports whether content contains an occurrence of
// funcCall (e.g. "exec(") that is a real call to the named PHP function
// rather than something that shares the same suffix.
//
// Four shapes must be rejected:
//
// - embedded identifiers: "doubleval(" must not match "eval("; the
// preceding character is a letter/digit/underscore;
// - method invocations: "$this->DB->exec(" must not match "exec(" even
// though the preceding ">" is non-alphanumeric;
// - static invocations: "Foo::exec(" must not match for the same reason;
// - function declarations: "function exec(" names a local function of
// the same name and must not be counted as a call site.
//
// The earlier implementation only guarded against the first case and was
// the source of false positives on elFinder volume drivers that call
// "$this->DB->exec(...)" (SQLite) alongside $_SERVER references on the
// same line.
func containsStandaloneFunc(content, funcCall string) bool {
idx := 0
for {
pos := strings.Index(content[idx:], funcCall)
if pos < 0 {
return false
}
absPos := idx + pos
nextIdx := absPos + len(funcCall)
advance := func() bool {
if nextIdx >= len(content) {
return false
}
idx = nextIdx
return true
}
if absPos == 0 {
return true
}
prev := content[absPos-1]
isAlnum := (prev >= 'a' && prev <= 'z') || (prev >= 'A' && prev <= 'Z') ||
(prev >= '0' && prev <= '9') || prev == '_'
if isAlnum {
if !advance() {
return false
}
continue
}
if absPos >= 2 {
op := content[absPos-2 : absPos]
if op == "->" || op == "::" {
if !advance() {
return false
}
continue
}
} else if absPos == 1 && (prev == '>' || prev == ':') {
// Degenerate position: only one preceding byte, and it is
// the tail char of a possible method ("->") or static
// ("::") operator. We cannot confirm the second char
// because there is no second char. The conservative choice
// is to skip, so a truncated "->exec(" or "::exec(" at the
// very start of a buffer does not get flagged as a real
// shell-function call.
if !advance() {
return false
}
continue
}
const decl = "function "
if absPos >= len(decl) && content[absPos-len(decl):absPos] == decl {
if !advance() {
return false
}
continue
}
return true
}
}
func containsAny(strs []string, substrs ...string) bool {
for _, s := range strs {
for _, sub := range substrs {
if strings.Contains(s, sub) {
return true
}
}
}
return false
}
// benignPHPStubMaxScan caps how many bytes of a candidate stub the
// recogniser will read. Stub files in the wild (BackWPup folder.php
// caches at ~160 KB, WP "silence is golden" index.php at ~30 B, plugin
// 404-stub headers under 1 KB) fit comfortably under this bound;
// anything larger must surface for normal alerting rather than be
// accepted on faith.
const benignPHPStubMaxScan = 4 * 1024 * 1024
// MaxInertPHPScanBytes is the largest file the inert-content recognizers read
// in full. Translation caches and comment-only stubs need every byte to prove
// that no code follows; a proven PHP terminator can still be accepted from an
// incomplete prefix because its tail is unreachable.
const MaxInertPHPScanBytes = benignPHPStubMaxScan
// IsBenignPHPStub reports whether the reachable code region of a PHP
// file consists only of whitespace and comments, or terminates with a
// literal-argument die / exit, or __halt_compiler before any other statement.
// Files matching either shape cannot execute attacker-controlled code
// via a web request: PHP either runs to EOF emitting nothing, or hits
// the terminator and stops with the remaining bytes unreachable.
//
// The recogniser is content-shape only -- it does not look at the path,
// filename, parent directory, or whether a plugin is installed. An
// attacker cannot bypass it by naming a payload to mimic a known-plugin
// file because the gate fails the moment any executable statement
// appears before a terminator. Conversely a legitimate plugin that
// writes a stub-shaped working file (BackWPup writes
// "<?php //<json>" for job state and "<?php\n//path1\n//path2..." for
// folder caches) is recognised regardless of where it puts the file.
//
// Other detectors -- signature scans, YARA, suspicious filename, the
// webshell name list -- still run on the file in their own pipelines.
// Only the path-only "anomalous PHP location" warning is suppressed
// for files that this recogniser accepts.
func IsBenignPHPStub(path string) bool {
f, err := osFS.Open(path)
if err != nil {
return false
}
defer func() { _ = f.Close() }()
buf := make([]byte, benignPHPStubMaxScan)
n, _ := f.Read(buf)
if n == 0 {
return false
}
info, err := f.Stat()
complete := err == nil && info.Size() <= int64(n)
return IsBenignPHPStubBytesComplete(buf[:n], complete)
}
// IsBenignPHPStubBytes is the buffer-only variant. The realtime fanotify
// path uses it on the bytes it already read from the file descriptor;
// IsBenignPHPStub provides the path-based entry point for the polled
// fileindex scan. Both rely on the same parser so realtime and scheduled
// scans agree on which files are stubs.
//
// The parser tokenises the leading region of the buffer:
//
// - Optional UTF-8 BOM and whitespace, then the literal "<?php" opener.
// The short-echo opener "<?=" is rejected because it emits output.
// A "<?phpfoo" run-together opener is rejected because PHP requires
// whitespace (or EOF) after the tag.
// - Repeatedly accept whitespace, line comments ("//..." or "#..." up to
// newline or "?>"), and balanced block comments ("/* ... */"). A "/*"
// without a matching "*/" inside the scanned window is rejected -- we
// cannot prove the rest of the file is comment.
// - Accept die, exit, and __halt_compiler as terminators. die and exit may
// carry a single literal argument (a non-interpolating string or a
// decimal integer); __halt_compiler takes none. Once seen, the rest of
// the buffer is treated as unreachable.
// - Reject any closing "?>" tag (would allow HTML escape and a later
// "<?php" re-entry that this gate does not analyse).
// - Reject any other identifier (return, if, system, eval, function,
// class, ...) and any stray punctuation ("$", "(", "=", ";", ...).
// Those are statements we cannot prove benign.
// - If the loop reaches EOF in a complete buffer having only seen
// whitespace and comments, accept: PHP outputs nothing and executes
// nothing.
func IsBenignPHPStubBytes(buf []byte) bool {
return IsBenignPHPStubBytesComplete(buf, true)
}
// IsBenignPHPStubBytesComplete is like IsBenignPHPStubBytes, but complete
// tells the parser whether buf contains the entire file. Comment-only stubs
// require a complete buffer; terminators do not, because bytes after them are
// unreachable to PHP.
func IsBenignPHPStubBytesComplete(buf []byte, complete bool) bool {
if len(buf) >= 3 && buf[0] == 0xEF && buf[1] == 0xBB && buf[2] == 0xBF {
buf = buf[3:]
}
i := 0
for i < len(buf) && isPHPSpace(buf[i]) {
i++
}
const opener = "<?php"
if !bytes.HasPrefix(buf[i:], []byte(opener)) {
return false
}
i += len(opener)
if i < len(buf) && !isPHPOpenTagSpace(buf[i]) {
return false
}
for i < len(buf) {
c := buf[i]
if isPHPSpace(c) {
i++
continue
}
if isPHPLineCommentStart(buf, i) {
i = skipPHPLineComment(buf, i)
continue
}
if c == '/' && i+1 < len(buf) && buf[i+1] == '*' {
i += 2
end := bytes.Index(buf[i:], []byte("*/"))
if end < 0 {
return false
}
i += end + 2
continue
}
if c == '?' && i+1 < len(buf) && buf[i+1] == '>' {
return false
}
if isIdentStart(c) {
start := i
for i < len(buf) && isIdentCont(buf[i]) {
i++
}
word := strings.ToLower(string(buf[start:i]))
return isPHPTerminatorStatement(buf, i, word, complete)
}
return false
}
return complete
}
// PHPTerminatesImmediately recognizes a leading exit, die, or __halt_compiler
// in unencoded PHP source, with at most one plain literal argument. Callers
// must establish that PHP source conversion is disabled before using this
// to suppress findings: conversion can remove even a raw terminator keyword.
// Only the opening tag and whitespace may precede it. Unlike the broader
// stub parser, it never accepts comments before the terminator. A completed
// terminator can be recognized from a partial head; EOF alone is not proof.
func PHPTerminatesImmediately(buf []byte) bool {
_, ok := PHPTerminatesImmediatelyAt(buf)
return ok
}
// PHPTerminatesImmediatelyAt is PHPTerminatesImmediately plus the offset just
// past the terminator statement, so a caller can judge the unreachable tail.
func PHPTerminatesImmediatelyAt(buf []byte) (int, bool) {
buf = bytes.TrimPrefix(buf, []byte{0xEF, 0xBB, 0xBF})
i := skipPHPSpace(buf, 0)
const opener = "<?php"
if !bytes.HasPrefix(buf[i:], []byte(opener)) {
return 0, false
}
i += len(opener)
// PHP needs whitespace after the opening tag. Without it the tag is
// literal text, the file never enters code mode here, and a later
// `<?php` block is what actually runs.
if i >= len(buf) || !isPHPOpenTagSpace(buf[i]) {
return 0, false
}
i = skipPHPSpace(buf, i)
start := i
for i < len(buf) && isIdentCont(buf[i]) {
i++
}
if i == start || !isIdentStart(buf[start]) {
return 0, false
}
word := strings.ToLower(string(buf[start:i]))
if !isPHPTerminatorStatement(buf, i, word, false) {
return 0, false
}
return phpTerminatorStatementEnd(buf, i, word), true
}
// phpTerminatorStatementEnd returns the offset where the file's unreachable
// bytes begin. Past the validated terminator it also consumes any further
// terminator statements and the closing tag, because a state file commonly
// opens `exit('...'); __halt_compiler(); ?>` before its data.
func phpTerminatorStatementEnd(buf []byte, i int, word string) int {
for {
i = skipPHPSpace(buf, i)
if i < len(buf) && buf[i] == '(' {
if next, ok := consumeEmptyPHPParens(buf, i); ok {
i = next
} else if next, ok := consumeLiteralPHPParens(buf, i); ok {
i = next
}
}
i = skipPHPSpace(buf, i)
if i < len(buf) && buf[i] == ';' {
i++
}
i = skipPHPSpace(buf, i)
if i+1 < len(buf) && buf[i] == '?' && buf[i+1] == '>' {
i += 2
// PHP swallows one newline directly after the closing tag.
if i < len(buf) && buf[i] == '\n' {
i++
} else if i+1 < len(buf) && buf[i] == '\r' && buf[i+1] == '\n' {
i += 2
}
return i
}
start := i
for i < len(buf) && isIdentCont(buf[i]) {
i++
}
if i == start || !isIdentStart(buf[start]) {
return start
}
next := strings.ToLower(string(buf[start:i]))
if next != "die" && next != "exit" && next != "__halt_compiler" {
return start
}
}
}
func isPHPSpace(c byte) bool {
return c == ' ' || c == '\t' || c == '\n' || c == '\r' || c == '\v' || c == '\f'
}
// PHP's long opening tag excludes the form feed and vertical tab accepted
// by generic whitespace scanners. Accepting either would hide later PHP blocks.
func isPHPOpenTagSpace(c byte) bool {
return c == ' ' || c == '\t' || c == '\n' || c == '\r'
}
func isIdentStart(c byte) bool {
return c == '_' || (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z')
}
func isIdentCont(c byte) bool {
return isIdentStart(c) || (c >= '0' && c <= '9')
}
func skipPHPSpace(buf []byte, i int) int {
for i < len(buf) && isPHPSpace(buf[i]) {
i++
}
return i
}
// PHP 8 attributes begin with #[ and can precede executable declarations and
// statements on the same line. Every inert-content scanner must preserve them.
func isPHPLineCommentStart[T string | []byte](buf T, i int) bool {
if buf[i] == '#' {
return i+1 == len(buf) || buf[i+1] != '['
}
return buf[i] == '/' && i+1 < len(buf) && buf[i+1] == '/'
}
// skipPHPLineComment returns the index that ends the "//" or "#" comment at
// i. PHP ends one at a bare CR as well as at LF, so a CR must not be read as
// comment text: the statement after it runs.
func skipPHPLineComment[T string | []byte](buf T, i int) int {
for i < len(buf) {
if buf[i] == '\n' || buf[i] == '\r' {
return i
}
if buf[i] == '?' && i+1 < len(buf) && buf[i+1] == '>' {
return i
}
i++
}
return i
}
// isPHPTerminatorStatement reports whether the identifier at word, which
// starts the first statement of the buffer, ends execution. __halt_compiler
// takes no argument; die and exit may print one literal before stopping.
func isPHPTerminatorStatement(buf []byte, i int, word string, complete bool) bool {
if word != "die" && word != "exit" && word != "__halt_compiler" {
return false
}
i = skipPHPSpace(buf, i)
if word == "__halt_compiler" {
next, ok := consumeEmptyPHPParens(buf, i)
if !ok {
return false
}
return phpTerminatorStatementEnds(buf, next, complete)
}
if i >= len(buf) {
return complete
}
if buf[i] == ';' {
return true
}
if buf[i] == '?' && i+1 < len(buf) && buf[i+1] == '>' {
return true
}
if buf[i] != '(' {
return false
}
next, ok := consumeEmptyPHPParens(buf, i)
if !ok {
if next, ok = consumeLiteralPHPParens(buf, i); !ok {
return false
}
}
return phpTerminatorStatementEnds(buf, next, complete)
}
// consumeLiteralPHPParens accepts `( <literal> )` where the literal is a
// single-quoted string, a double-quoted string that interpolates nothing, or
// a decimal integer. exit and die evaluate their argument before stopping, so
// a literal is the only shape that proves no other code runs.
func consumeLiteralPHPParens(buf []byte, i int) (int, bool) {
if i >= len(buf) || buf[i] != '(' {
return i, false
}
i = skipPHPSpace(buf, i+1)
if i >= len(buf) {
return i, false
}
if isPHPQuote(buf[i]) {
end, ok := endOfPHPLiteralString(buf, i)
if !ok {
return i, false
}
i = end
} else {
start := i
for i < len(buf) && buf[i] >= '0' && buf[i] <= '9' {
i++
}
if i == start {
return i, false
}
}
i = skipPHPSpace(buf, i)
if i >= len(buf) || buf[i] != ')' {
return i, false
}
return skipPHPSpace(buf, i+1), true
}
// endOfPHPLiteralString returns the index just past the string literal that
// starts at i. It fails on a string the buffer does not terminate and on a
// double-quoted string carrying a variable, because PHP evaluates `$x` and
// `{$x}` inside double quotes. Keep literals plain ASCII without escape or
// encoding-shift bytes: the shared stub parser also runs without an encoding
// policy, and source conversion can expose expressions inside such strings.
func endOfPHPLiteralString(buf []byte, i int) (int, bool) {
quote := buf[i]
for j := i + 1; j < len(buf); j++ {
if buf[j] < ' ' || buf[j] > '~' || strings.ContainsRune("\\+=&~", rune(buf[j])) {
return j, false
}
switch buf[j] {
case '$':
if quote == '"' {
return j, false
}
case quote:
return j + 1, true
}
}
return len(buf), false
}
func consumeEmptyPHPParens(buf []byte, i int) (int, bool) {
if i >= len(buf) || buf[i] != '(' {
return i, false
}
i = skipPHPSpace(buf, i+1)
if i >= len(buf) || buf[i] != ')' {
return i, false
}
return skipPHPSpace(buf, i+1), true
}
func phpTerminatorStatementEnds(buf []byte, i int, complete bool) bool {
i = skipPHPSpace(buf, i)
if i >= len(buf) {
return complete
}
if buf[i] == ';' {
return true
}
return buf[i] == '?' && i+1 < len(buf) && buf[i+1] == '>'
}
package checks
import (
"os"
"sync"
"github.com/pidginhost/csm/internal/store"
)
// Keep unpersisted progress while the daemon is alive. Otherwise one failed
// cursor write repeatedly spends the host's entire budget on the same account.
// Successful writes evict the fallback; account removal and database changes
// discard it too, so historical accounts cannot accumulate in memory.
var rollingContentCursors = phpContentCursors{}
type phpContentCursors struct {
mu sync.Mutex
db *store.DB
pending map[string]store.ScanCursorRecord
}
func (c *phpContentCursors) resetLocked(db *store.DB) {
if c.db != db {
c.db = db
c.pending = make(map[string]store.ScanCursorRecord)
}
}
func (c *phpContentCursors) retain(db *store.DB, entries []os.DirEntry) {
c.mu.Lock()
defer c.mu.Unlock()
c.resetLocked(db)
live := make(map[string]bool, len(entries))
for _, entry := range entries {
if entry.IsDir() {
live[entry.Name()] = true
}
}
for account := range c.pending {
if !live[account] {
delete(c.pending, account)
}
}
}
func (c *phpContentCursors) load(db *store.DB, account string) (store.ScanCursorRecord, error) {
c.mu.Lock()
defer c.mu.Unlock()
c.resetLocked(db)
cur, _, err := db.GetScanCursor(account, rollingScanCheck)
if pending, ok := c.pending[account]; ok {
return pending, err
}
return cur, err
}
func (c *phpContentCursors) save(db *store.DB, cur store.ScanCursorRecord) error {
c.mu.Lock()
defer c.mu.Unlock()
c.resetLocked(db)
if err := db.PutScanCursor(cur); err != nil {
c.pending[cur.Account] = cur
return err
}
delete(c.pending, cur.Account)
return nil
}
package checks
import (
"context"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/store"
)
// rollingScanCheck is the cursor key under which rolling content coverage
// records its per-account progress. It reuses the existing content findings, so
// no new check name is registered.
const rollingScanCheck = "php_content"
// rollingWalkMaxDepth bounds how deep enumeratePHPFiles descends under each
// docroot, so a deeply nested or symlink-looped tree cannot make enumeration
// unbounded. The fixed suspicious-dir scan uses a shallow depth; rolling needs
// to reach app code nested a few levels down (e.g. wp-content/plugins/x/inc/y)
// but does not need to chase arbitrarily deep trees.
const rollingWalkMaxDepth = 12
// rollingContentEnabled gates rolling coverage to normal periodic host-wide
// runs: the knob is on, this is not an account-scoped run, and it is not a
// forced/audit full content scan (those already read every file). forcedFull is
// the per-cycle decision already computed by CheckPHPContent
// (phpContentForceFull || scanForceContent); it is threaded in rather than
// recomputed because phpContentForceFull advances a cadence counter on each
// call, so calling it twice per cycle would skew the forced-rescan cadence.
func rollingContentEnabled(ctx context.Context, cfg *config.Config, forcedFull bool) bool {
return cfg.Thresholds.RollingCoverage &&
AccountFromContext(ctx) == "" &&
!forcedFull
}
// rollingContentPass spends one cycle's file budget on rolling coverage across
// the host. The budget is the operator's per-scan file cap: sizing a window
// that large for every account asked for hundreds of windows inside one check
// budget, so the check never finished and its findings were never reported.
// Accounts are taken least-recently-covered first, so a host too large for one
// cycle keeps moving instead of resweeping the same alphabetical prefix.
func rollingContentPass(ctx context.Context, cfg *config.Config, scan *phpContentScan, homeDirs []os.DirEntry, findings *[]alert.Finding) {
db := store.Global()
if db == nil {
// Cannot persist a cursor, so rolling would scan from the start every
// cycle without making progress. Skip rather than spin in place.
return
}
rollingContentCursors.retain(db, homeDirs)
accounts := rollingAccountOrder(db, homeDirs)
budget := accountScanMaxFiles(ctx, cfg)
wrapped := 0
for _, account := range accounts {
if ctx.Err() != nil || budget <= 0 {
break
}
whole, used := rollingContentCoverage(ctx, cfg, scan, account.name, accountDocRoots(account.home), budget, findings)
budget -= used
if whole {
wrapped++
}
}
if wrapped != len(accounts) {
// The accounts this cycle did not finish are not re-emitting their
// earlier findings, and completing the check would purge them
// (mirrors yara_deep).
markCheckIncomplete(ctx, "php_content")
}
}
// rollingAccount pairs an account name with its home directory.
type rollingAccount struct {
name string
home string
}
// rollingAccountOrder sorts accounts by the time each last completed a full
// traversal, oldest first, so an account that has never been covered goes
// first and no account can be starved by the ones before it in the alphabet.
func rollingAccountOrder(db *store.DB, homeDirs []os.DirEntry) []rollingAccount {
type ordered struct {
rollingAccount
last time.Time
}
all := make([]ordered, 0, len(homeDirs))
for _, entry := range homeDirs {
if !entry.IsDir() {
continue
}
cur, _ := rollingContentCursors.load(db, entry.Name())
all = append(all, ordered{
rollingAccount: rollingAccount{name: entry.Name(), home: scanHomeDirPath(entry)},
last: cur.LastFullCycleTS,
})
}
sort.Slice(all, func(i, j int) bool {
if !all[i].last.Equal(all[j].last) {
return all[i].last.Before(all[j].last)
}
return all[i].name < all[j].name
})
out := make([]rollingAccount, 0, len(all))
for _, a := range all {
out = append(out, a.rollingAccount)
}
return out
}
// rollingContentCoverage sweeps a bounded path-sorted slice of the account's
// full docroot PHP-source set, advancing the per-account cursor so every stock
// PHP source or source-view file is eventually content-scanned over cycles.
// The caller guarantees the gate (rolling on, host-scope periodic, not a
// forced/audit run) and hands it what is left of the cycle's file budget.
// Findings append to the live findings slice (rolling is part of the periodic
// scan, not a report-only full-scan job). A canceled run leaves the prior
// cursor untouched. It returns how many files it read so the caller can charge
// them against the budget.
//
// Limitation: rolling enumerates only stock-PHP-executable filenames across the
// whole docroot. A file whose non-stock extension is remapped to PHP by an
// .htaccess handler (the LEVIATHAN trick) is NOT enumerated here; the fixed
// suspicious-dir scan (which layers per-directory overlays as it descends) and
// realtime fanotify still cover those.
// It reports whether this cycle covered the account's whole file list. A
// window that did not wrap leaves files from earlier windows unvisited, and
// their findings are not re-emitted this cycle, so the caller must mark the
// check incomplete or the runner purges them from the latest set.
func rollingContentCoverage(ctx context.Context, cfg *config.Config, scan *phpContentScan, account string, docRoots []string, limit int, findings *[]alert.Finding) (bool, int) {
db := store.Global()
files := enumeratePHPFiles(ctx, cfg, docRoots)
if ctx.Err() != nil {
return false, 0
}
// Enumeration already covers all roots independently of the content
// window. Prune even on partial or empty windows, including lists that
// keep growing before the cursor can ever wrap.
scan.pruneMissing(docRoots, files)
if len(files) == 0 {
return true, 0
}
cur, curErr := rollingContentCursors.load(db, account)
if curErr != nil {
// Keep storage failures visible even when in-memory progress lets
// this daemon continue covering the account.
fmt.Fprintf(os.Stderr, "php_content rolling: cursor read for %s: %v\n", account, curErr)
}
selected, newLast, wrapped := rollingCandidatesAfter(files, cur.LastPath, limit)
if len(selected) == 0 {
return true, 0
}
// Crossing the end of the list completes a traversal across several
// windows, but this run still did not re-emit findings from the earlier
// windows. Only a window containing the whole list is safe to report as a
// completed check to the runner.
windowComplete := len(selected) == len(files)
fullTraversal := wrapped || windowComplete
// Reconstruct the .htaccess handler overlay once per directory: every file
// in the slice that shares a directory shares the same overlay, and reading
// the ancestor .htaccess chain per file would multiply the read cost.
overlayCache := make(map[string]phpHandlerOverlay)
read := 0
for _, file := range selected {
if ctx.Err() != nil {
break
}
// Opening FIFOs or device nodes can block the scan; rolling only needs
// regular PHP files (including symlinks that resolve to regular files).
if !rollingRegularCandidate(file) {
// A failed stat or changed file type invalidates any earlier
// clean result just like an unsuccessful content read.
delete(scan.prev, file)
delete(scan.next, file)
continue
}
read++
dir := filepath.Dir(file)
overlay, ok := overlayCache[dir]
if !ok {
overlay = reconstructOverlay(rollingDocRootFor(file, docRoots), dir)
overlayCache[dir] = overlay
}
scan.scanFile(ctx, file, overlay, findings)
}
// Advance the cursor only on a complete, uncanceled run. A run cut short by
// ctx cancellation leaves the prior cursor so the next cycle resumes where
// this one stopped instead of skipping the unscanned tail.
if ctx.Err() != nil {
return false, read
}
cur.Account = account
cur.Check = rollingScanCheck
cur.LastPath = newLast
if fullTraversal {
now := time.Now().UTC()
cur.LastFullCycleTS = now
if wrapped {
cur.WrappedAt = now
}
}
if err := rollingContentCursors.save(db, cur); err != nil {
fmt.Fprintf(os.Stderr, "php_content rolling: cursor write for %s: %v\n", account, err)
}
return windowComplete, read
}
func rollingRegularCandidate(file string) bool {
info, err := osFS.Stat(file)
return err == nil && info.Mode().IsRegular()
}
// rollingDocRootFor returns the docRoot that contains file. file always sits
// under exactly one of docRoots (enumeratePHPFiles built it by descending from
// them); the longest matching prefix wins so nested account roots resolve to
// the most specific one.
func rollingDocRootFor(file string, docRoots []string) string {
best := ""
for _, root := range docRoots {
if root == file || strings.HasPrefix(file, root+string(filepath.Separator)) {
if len(root) > len(best) {
best = root
}
}
}
return best
}
// enumeratePHPFiles recursively collects, under each docRoot, candidate paths
// whose name contains stock PHP source. This includes source-view .phps files
// without classifying them as executable. The walk is bounded to
// rollingWalkMaxDepth, honours ctx cancellation, and
// respects suppressions.ignore_paths exactly like scanDir when the scan is not
// an explicit full-scan/audit. The result is ascending-sorted and de-duplicated
// so rollingCandidatesAfter can cursor through it stably.
func enumeratePHPFiles(ctx context.Context, cfg *config.Config, docRoots []string) []string {
seen := make(map[string]struct{})
respectIgnores := scanRespectsIgnores(ctx, cfg)
for _, root := range docRoots {
walkPHPFiles(ctx, cfg, root, rollingWalkMaxDepth, respectIgnores, seen)
if ctx.Err() != nil {
break
}
}
if len(seen) == 0 {
return nil
}
files := make([]string, 0, len(seen))
for f := range seen {
files = append(files, f)
}
sort.Strings(files)
return files
}
func walkPHPFiles(ctx context.Context, cfg *config.Config, dir string, maxDepth int, respectIgnores bool, seen map[string]struct{}) {
if ctx.Err() != nil || maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
fullPath := filepath.Join(dir, entry.Name())
// Same suppression gate as scanDir: this is not a path allowlist for
// "safe" files but an operator-configured ignore that the periodic scan
// already honours. It is bypassed for explicit full-scan/audit runs,
// which never reach rolling anyway (the gate excludes forced runs).
if respectIgnores && pathIsIgnored(cfg, fullPath) {
continue
}
if entry.IsDir() {
walkPHPFiles(ctx, cfg, fullPath, maxDepth-1, respectIgnores, seen)
continue
}
if contenttype.IsPHPSourceName(strings.ToLower(entry.Name())) {
seen[fullPath] = struct{}{}
}
}
}
func pathIsIgnored(cfg *config.Config, fullPath string) bool {
for _, ignore := range cfg.Suppressions.IgnorePaths {
if matchGlob(fullPath, ignore) {
return true
}
}
return false
}
// reconstructOverlay builds the handler overlay for fileDir by merging the
// .htaccess files from rootDir down through each ancestor to fileDir, starting
// from an empty overlay at rootDir. This matches how scanDir accumulates
// overlays as it descends. fileDir must be rootDir or a descendant; if rootDir
// is empty (file resolved to no docRoot, which should not happen) the overlay
// is built from fileDir alone.
func reconstructOverlay(rootDir, fileDir string) phpHandlerOverlay {
overlay := phpHandlerOverlay{}
if rootDir == "" {
if data, ok, err := readHtaccessBounded(filepath.Join(fileDir, ".htaccess")); err == nil && ok {
overlay = overlay.mergeHtaccess(data)
} else if htaccessOversized(ok, err) {
overlay.unrestricted = true
}
return overlay
}
// Build the ordered list of directories from rootDir down to fileDir
// inclusive by stripping the shared prefix and walking the relative
// components back on.
dirs := []string{rootDir}
rel, err := filepath.Rel(rootDir, fileDir)
if err == nil && rel != "." && rel != "" && !strings.HasPrefix(rel, "..") {
cur := rootDir
for _, part := range strings.Split(rel, string(filepath.Separator)) {
cur = filepath.Join(cur, part)
dirs = append(dirs, cur)
}
}
for _, d := range dirs {
if data, ok, err := readHtaccessBounded(filepath.Join(d, ".htaccess")); err == nil && ok {
overlay = overlay.mergeHtaccess(data)
} else if htaccessOversized(ok, err) {
// Unreadable handler configuration: scan every name rather
// than assume the default extension set.
overlay.unrestricted = true
}
}
return overlay
}
package checks
import (
"bytes"
"io"
"strings"
)
// wpTranslationMaxDepth bounds array nesting so a crafted file cannot drive
// unbounded recursion in the recognizer. Real translation caches nest one
// level deep (the "messages" map), so this leaves ample headroom.
const wpTranslationMaxDepth = 16
// IsWPTranslationCacheBytesComplete reports whether buf is exactly a WordPress
// PHP translation cache: the "<?php" opener, the keyword "return", a single PHP
// array literal whose elements are only string/integer scalars (optionally
// concatenated string literals, as GlotPress joins plural forms with a "\0"
// separator) or nested arrays of the same, then a ";" and nothing else.
// WordPress 6.5+ auto-generates these as pure data return maps (*.l10n.php);
// each one previously opened a sensitive-dir Warning incident.
//
// This is a content-structure recognizer, not a path or filename allowlist. A
// variable, a function call, string interpolation, a concatenation operand that
// is not a literal, a closing "?>" tag, or any statement after the array makes
// it return false, so an attacker cannot smuggle code into a file shaped like a
// translation cache. complete must be true: a truncated buffer cannot prove the
// unseen tail carries no code, so it is never suppressed.
func IsWPTranslationCacheBytesComplete(buf []byte, complete bool) bool {
if !complete || len(buf) == 0 {
return false
}
if len(buf) >= 3 && buf[0] == 0xEF && buf[1] == 0xBB && buf[2] == 0xBF {
buf = buf[3:]
}
s := &phpLiteralScanner{buf: buf}
s.skipSpace()
if !s.consumeOpener() {
return false
}
s.skipTrivia()
if id, ok := s.readIdent(); !ok || id != "return" {
return false
}
s.skipTrivia()
if !s.parseTopArray() {
return false
}
s.skipTrivia()
if s.i >= len(s.buf) || s.buf[s.i] != ';' {
return false
}
s.i++
s.skipTrivia()
return s.i == len(s.buf)
}
// phpLiteralScanner walks a byte buffer that is expected to be a constant PHP
// data literal. It never evaluates anything; it only proves the bytes contain
// no executable construct.
type phpLiteralScanner struct {
buf []byte
i int
}
// skipSpace advances past whitespace only. Used before the open tag, where any
// non-whitespace would be raw output (HTML) rather than a PHP comment.
func (s *phpLiteralScanner) skipSpace() {
for s.i < len(s.buf) && isPHPSpace(s.buf[s.i]) {
s.i++
}
}
// skipTrivia advances past whitespace and PHP comments. A "/*" without a
// matching "*/" consumes to EOF, leaving the scanner at the end so callers
// expecting a token fail closed.
func (s *phpLiteralScanner) skipTrivia() {
for s.i < len(s.buf) {
c := s.buf[s.i]
if isPHPSpace(c) {
s.i++
continue
}
if isPHPLineCommentStart(s.buf, s.i) {
s.i = skipPHPLineComment(s.buf, s.i)
continue
}
if c == '/' && s.i+1 < len(s.buf) && s.buf[s.i+1] == '*' {
end := bytes.Index(s.buf[s.i+2:], []byte("*/"))
if end < 0 {
s.i = len(s.buf)
return
}
s.i += 2 + end + 2
continue
}
return
}
}
func (s *phpLiteralScanner) consumeOpener() bool {
const opener = "<?php"
if !bytes.HasPrefix(s.buf[s.i:], []byte(opener)) {
return false
}
s.i += len(opener)
// PHP requires a space, tab, CR or LF (or EOF) after the tag; "<?phpreturn"
// is invalid, "<?=" is the short-echo opener, and after a vertical tab or
// form feed PHP prints the whole file as text.
if s.i < len(s.buf) && !isPHPOpenTagSpace(s.buf[s.i]) {
return false
}
return true
}
func (s *phpLiteralScanner) readIdent() (string, bool) {
if s.i >= len(s.buf) || !isIdentStart(s.buf[s.i]) {
return "", false
}
start := s.i
for s.i < len(s.buf) && isIdentCont(s.buf[s.i]) {
s.i++
}
return strings.ToLower(string(s.buf[start:s.i])), true
}
// parseTopArray requires the value at the cursor to be an array literal. A bare
// scalar return (return 'x';) is not a translation cache.
func (s *phpLiteralScanner) parseTopArray() bool {
if s.i >= len(s.buf) {
return false
}
if s.buf[s.i] == '[' {
return s.parseArray(0, '[', ']')
}
if isIdentStart(s.buf[s.i]) {
save := s.i
if id, _ := s.readIdent(); id == "array" {
s.skipTrivia()
return s.parseArray(0, '(', ')')
}
s.i = save
}
return false
}
// parseValue accepts one array element value: a nested array, or a chain of
// constant scalars joined by the "." concatenation operator.
func (s *phpLiteralScanner) parseValue(depth int) bool {
s.skipTrivia()
if s.i >= len(s.buf) {
return false
}
if s.buf[s.i] == '[' {
return s.parseArray(depth, '[', ']')
}
if isIdentStart(s.buf[s.i]) {
save := s.i
if id, _ := s.readIdent(); id == "array" {
s.skipTrivia()
return s.parseArray(depth, '(', ')')
}
s.i = save
}
return s.parseConcatChain()
}
// parseArray consumes an array literal delimited by open/close. Elements are
// "value" or "value => value"; an empty array and a trailing comma are allowed.
func (s *phpLiteralScanner) parseArray(depth int, open, close byte) bool {
if depth >= wpTranslationMaxDepth {
return false
}
if s.i >= len(s.buf) || s.buf[s.i] != open {
return false
}
s.i++
for {
s.skipTrivia()
if s.i >= len(s.buf) {
return false
}
if s.buf[s.i] == close {
s.i++
return true
}
if !s.parseValue(depth + 1) {
return false
}
s.skipTrivia()
if s.i+1 < len(s.buf) && s.buf[s.i] == '=' && s.buf[s.i+1] == '>' {
s.i += 2
if !s.parseValue(depth + 1) {
return false
}
s.skipTrivia()
}
if s.i >= len(s.buf) {
return false
}
switch s.buf[s.i] {
case ',':
s.i++
case close:
s.i++
return true
default:
return false
}
}
}
func (s *phpLiteralScanner) parseConcatChain() bool {
if !s.parseScalar() {
return false
}
for {
s.skipTrivia()
if s.i < len(s.buf) && s.buf[s.i] == '.' {
s.i++
if !s.parseScalar() {
return false
}
continue
}
return true
}
}
func (s *phpLiteralScanner) parseScalar() bool {
s.skipTrivia()
if s.i >= len(s.buf) {
return false
}
c := s.buf[s.i]
switch {
case c == '\'':
return s.parseSingleQuoted()
case c == '"':
return s.parseDoubleQuoted()
case c == '-' || (c >= '0' && c <= '9'):
return s.parseInt()
case isIdentStart(c):
id, ok := s.readIdent()
return ok && (id == "true" || id == "false" || id == "null")
default:
return false
}
}
// parseSingleQuoted consumes a PHP single-quoted string. Inside one, only "\\"
// and "\'" are escapes and "$" never interpolates, so the contents are inert.
func (s *phpLiteralScanner) parseSingleQuoted() bool {
s.i++ // opening quote
for s.i < len(s.buf) {
switch s.buf[s.i] {
case '\\':
s.i += 2
case '\'':
s.i++
return true
default:
s.i++
}
}
return false // unterminated
}
// parseDoubleQuoted consumes a PHP double-quoted string. A backslash escapes
// the next byte (so "\0", "\n", "\$" are literal). An unescaped "$" introduces
// variable interpolation, which is an execution vector, so it is rejected.
func (s *phpLiteralScanner) parseDoubleQuoted() bool {
s.i++ // opening quote
for s.i < len(s.buf) {
switch s.buf[s.i] {
case '\\':
s.i += 2
case '"':
s.i++
return true
case '$':
return false
default:
s.i++
}
}
return false // unterminated
}
func (s *phpLiteralScanner) parseInt() bool {
if s.buf[s.i] == '-' {
s.i++
}
start := s.i
for s.i < len(s.buf) && s.buf[s.i] >= '0' && s.buf[s.i] <= '9' {
s.i++
}
return s.i > start
}
// isWPTranslationCache is the path-based entry point for the polled scan. It
// reads up to the same window as IsBenignPHPStub and requires the whole file to
// fit, since the recognizer must see the entire body to prove it is inert.
func isWPTranslationCache(path string) bool {
f, err := osFS.Open(path)
if err != nil {
return false
}
defer func() { _ = f.Close() }()
buf, err := io.ReadAll(io.LimitReader(f, benignPHPStubMaxScan+1))
if err != nil || len(buf) == 0 || len(buf) > benignPHPStubMaxScan {
return false
}
info, err := f.Stat()
if err != nil || info.Size() != int64(len(buf)) {
return false
}
return IsWPTranslationCacheBytesComplete(buf, true)
}
package checks
// IsWPVersionDataBytesComplete recognizes only literal assignments to the
// variables WordPress reads from its version file. Matching an installed copy
// alone cannot establish that a short-lived version probe carried no payload.
func IsWPVersionDataBytesComplete(buf []byte, complete bool) bool {
if !complete || len(buf) == 0 {
return false
}
s := &phpLiteralScanner{buf: buf}
s.skipSpace()
if !s.consumeOpener() {
return false
}
version := false
for {
s.skipTrivia()
if s.i == len(s.buf) {
return version
}
if s.buf[s.i] != '$' {
return false
}
s.i++
name, ok := s.readIdent()
if !ok {
return false
}
switch name {
case "wp_version", "wp_db_version", "tinymce_version", "required_php_version",
"required_php_extensions", "required_mysql_version", "wp_local_package":
default:
return false
}
s.skipTrivia()
if s.i == len(s.buf) || s.buf[s.i] != '=' {
return false
}
s.i++
s.skipTrivia()
if name == "required_php_extensions" {
if !s.parseTopArray() {
return false
}
} else if !s.parseScalar() {
return false
}
s.skipTrivia()
if s.i == len(s.buf) || s.buf[s.i] != ';' {
return false
}
s.i++
version = version || name == "wp_version"
}
}
package checks
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"syscall"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
const (
phpIniWalkMaxDepth = -1
phpIniMaxRootsPerUser = 1024
phpIniReadDirBatch = 256
phpIniStateFormat = 1
phpIniFindingCheck = "php_config_change"
phpIniIncompleteCheck = "php_config_scan_incomplete"
phpIniIncompleteDetails = "The file exceeds the PHP configuration scan limit and was not parsed."
phpIniSpecialDetails = "The PHP configuration path is not a regular file and was not parsed."
phpIniStateLockShards = 64
// PHPConfigMaxBytes is the shared scheduled and realtime read ceiling.
PHPConfigMaxBytes = 1 << 20
)
// Walk limits. Per-root limits bound one document root; the per-account
// limits bound the whole account. They are vars so tests can shrink them.
//
// The per-root allowance is what the account limits used to be, because a
// single shared budget meant a large first root (a vendor or node_modules tree
// is thousands of directories on its own) consumed everything and every later
// addon domain was reported unscanned without being walked.
var (
phpIniWalkMaxDirs = 10000
phpIniWalkMaxEntries = 250000
phpIniWalkAccountMaxDirs = 150000
phpIniWalkAccountMaxEntries = 3000000
)
var (
errPHPIniNonRegular = errors.New("PHP configuration is not a regular file")
errPHPIniTooLarge = errors.New("PHP configuration exceeds scan limit")
phpIniStateLocks [phpIniStateLockShards]sync.Mutex
)
type phpIniFileState struct {
Version int `json:"version"`
Hash string `json:"hash"`
Assessed bool `json:"assessed"`
FullAnalysis bool `json:"full_analysis,omitempty"`
}
// phpIniWalkBudget carries both the allowance for the root currently being
// walked and the running total for the account. collectPHPIniFilesWithBudget
// resets the per-root counters on entry, so roots do not starve each other.
type phpIniWalkBudget struct {
dirs int
entries int
accountDirs int
accountEntries int
// maxDirs is the operator-set per-root ceiling. It rides on the budget
// rather than on the package var so concurrent per-account scans cannot
// overwrite each other's limit. Zero falls back to the package default.
maxDirs int
maxEntries int
// limitHit distinguishes a walk stopped by one of the configured ceilings
// from one stopped by an unreadable directory. Both leave the scan
// incomplete, but only the first means coverage can be bought back by
// raising a setting, so the operator has to be told which happened.
limitKind walkLimit
}
// walkLimit identifies the bound that stopped a walk. The two ceilings are
// raised by two different settings, so naming the wrong one sends the
// operator to a knob that cannot recover the lost coverage.
type walkLimit int
const (
walkLimitNone walkLimit = iota
walkLimitDirs
walkLimitEntries
)
// phpIniIncompleteReason explains why a walk below root did not finish. A
// ceiling that stopped it is reported with the distance covered and the setting
// that raises it; anything else is left as an unreadable-entry report.
func phpIniIncompleteReason(root string, b *phpIniWalkBudget) string {
if b == nil {
return fmt.Sprintf("Could not finish scanning PHP configuration files below %s.", root)
}
switch b.limitKind {
case walkLimitDirs:
return fmt.Sprintf(
"Scanning below %s stopped after %d directories, the per-root limit, so the rest was not examined for PHP configuration files. Raise thresholds.php_config_walk_max_dirs to cover the whole account.",
root, b.dirs,
)
case walkLimitEntries:
return fmt.Sprintf(
"Scanning below %s stopped after %d entries, the per-root limit, so the rest was not examined for PHP configuration files. Raise thresholds.php_config_walk_max_entries to cover the whole account.",
root, b.entries,
)
}
return fmt.Sprintf("Could not finish scanning PHP configuration files below %s.", root)
}
// dirLimit is the per-root directory ceiling in force for this walk.
// phpIniConfiguredMaxDirs reads the operator's per-root ceiling, tolerating a
// nil config so callers that scan without one keep the built-in default.
func phpIniConfiguredMaxDirs(cfg *config.Config) int {
if cfg == nil {
return 0
}
return cfg.Thresholds.PHPConfigWalkMaxDirs
}
// phpIniConfiguredMaxEntries reads the operator's per-root entry ceiling.
func phpIniConfiguredMaxEntries(cfg *config.Config) int {
if cfg == nil {
return 0
}
return cfg.Thresholds.PHPConfigWalkMaxEntries
}
func (b *phpIniWalkBudget) dirLimit() int {
if b.maxDirs > 0 {
return b.maxDirs
}
return phpIniWalkMaxDirs
}
// entryLimit is the per-root entry ceiling in force for this walk.
func (b *phpIniWalkBudget) entryLimit() int {
if b.maxEntries > 0 {
return b.maxEntries
}
return phpIniWalkMaxEntries
}
func (b *phpIniWalkBudget) startRoot() {
b.dirs = 0
b.entries = 0
}
func (b *phpIniWalkBudget) accountExhausted() bool {
return b.accountDirs >= phpIniWalkAccountMaxDirs || b.accountEntries >= phpIniWalkAccountMaxEntries
}
func (b *phpIniWalkBudget) addDir() {
b.dirs++
b.accountDirs++
}
func (b *phpIniWalkBudget) addEntry() {
b.entries++
b.accountEntries++
}
// CheckPHPConfigChanges monitors .user.ini and php.ini files anywhere under an
// account's document roots for settings that weaken PHP security (disable_functions
// cleared or neutralized, allow_url_include enabled, open_basedir removed). It runs
// as a deep check; the fanotify watcher also catches these writes in real-time.
func CheckPHPConfigChanges(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
accountScope := AccountFromContext(ctx)
var cpanelVhosts []vhost
vhostRootsComplete := true
vhostMapExpected := false
vhostData, vhostErr := osFS.ReadFile(userdataDomainsPath)
if vhostErr == nil {
vhostMapExpected = true
var complete bool
cpanelVhosts, complete = parseUserdataDomainRootsChecked(string(vhostData))
if !complete {
vhostRootsComplete = false
markCheckIncomplete(ctx, "php_config_changes")
if accountScope == "" {
findings = append(findings, phpIniScanIncompleteFinding(
"PHP configuration document-root map is incomplete",
"Some cPanel domain records were malformed, so their document roots could not be scanned.",
))
}
}
} else {
vhostRootsComplete = false
vhostMapExpected = vhostMapFailureIsIncomplete(vhostErr)
if vhostMapExpected {
markCheckIncomplete(ctx, "php_config_changes")
if accountScope == "" {
findings = append(findings, phpIniScanIncompleteFinding(
"PHP configuration document-root map is unavailable",
fmt.Sprintf("Could not read %s: %v", userdataDomainsPath, vhostErr),
))
}
}
}
homeDirs, homeErr := GetScanHomeDirs(ctx)
var users []string
userSet := make(map[string]struct{})
addUser := func(user string) {
if _, exists := userSet[user]; exists {
return
}
userSet[user] = struct{}{}
users = append(users, user)
}
if homeErr == nil {
for _, homeEntry := range homeDirs {
if homeEntry.IsDir() {
addUser(homeEntry.Name())
}
}
}
for _, vh := range cpanelVhosts {
if accountScope == "" || vh.user == accountScope {
addUser(vh.user)
}
}
if homeErr != nil &&
(!errors.Is(homeErr, os.ErrNotExist) || len(users) == 0 && vhostMapExpected) {
markCheckIncomplete(ctx, "php_config_changes")
if accountScope == "" {
findings = append(findings, phpIniScanIncompleteFinding(
"PHP configuration account discovery is incomplete",
fmt.Sprintf("Could not enumerate account home directories: %v", homeErr),
))
}
}
for _, user := range users {
if ctx.Err() != nil {
markCheckIncomplete(ctx, "php_config_changes")
return findings
}
// Collect every .user.ini and php.ini under the account's document
// roots. An attacker plants a php.ini (or a nested .user.ini) deep in
// the tree -- e.g. wp-includes/assets/php.ini -- to weaken PHP, so the
// scan cannot stop at the docroot root or watch .user.ini alone.
var incompleteReason string
recordIncomplete := func(reason string) {
markCheckIncomplete(ctx, "php_config_changes")
if incompleteReason == "" {
incompleteReason = reason
}
}
roots := []string{filepath.Join(accountHomeDir(user), "public_html")}
rootSet := map[string]struct{}{roots[0]: {}}
addRoot := func(root string) bool {
root = filepath.Clean(root)
if _, duplicate := rootSet[root]; duplicate {
return true
}
if len(roots) >= phpIniMaxRootsPerUser {
recordIncomplete(fmt.Sprintf(
"Account %s has more than %d candidate document roots.",
user,
phpIniMaxRootsPerUser,
))
return false
}
rootSet[root] = struct{}{}
roots = append(roots, root)
return true
}
hasAuthoritativeRoot := false
for _, vh := range cpanelVhosts {
if vh.user != user {
continue
}
hasAuthoritativeRoot = true
if !addRoot(vh.docroot) {
break
}
}
// Non-cPanel panels commonly keep domains below an account-owned
// top-level directory (for example domains/<name>/public_html).
// Preserve that layout as a fallback when no complete cPanel map is
// available instead of narrowing the scan to public_html.
if !hasAuthoritativeRoot || !vhostRootsComplete {
fallbackRoots, complete, err := phpIniFallbackRoots(
ctx,
user,
phpIniMaxRootsPerUser-len(roots),
phpIniWalkMaxEntries,
)
if err != nil {
recordIncomplete(fmt.Sprintf(
"Could not enumerate fallback document roots for account %s: %v",
user,
err,
))
} else {
if !complete {
recordIncomplete(fmt.Sprintf(
"Could not finish enumerating fallback document roots for account %s.",
user,
))
}
for _, root := range fallbackRoots {
if !addRoot(root) {
break
}
}
}
}
var iniPaths []string
walkBudget := &phpIniWalkBudget{
maxDirs: phpIniConfiguredMaxDirs(cfg),
maxEntries: phpIniConfiguredMaxEntries(cfg),
}
for _, root := range roots {
paths, complete := collectPHPIniFilesWithBudget(ctx, root, phpIniWalkMaxDepth, walkBudget)
if !complete {
recordIncomplete(phpIniIncompleteReason(root, walkBudget))
}
iniPaths = append(iniPaths, paths...)
}
seen := make(map[string]struct{}, len(iniPaths))
for _, iniPath := range iniPaths {
if ctx.Err() != nil {
markCheckIncomplete(ctx, "php_config_changes")
return findings
}
if _, duplicate := seen[iniPath]; duplicate {
continue
}
seen[iniPath] = struct{}{}
// Bound reads and reject special files before opening them. A user
// can create a FIFO named php.ini; reading it as an ordinary file
// would strand a scan worker after its context timed out.
dangerous, err := readAndAssessPHPIniFile(store, iniPath)
if err != nil {
switch {
case errors.Is(err, os.ErrNotExist):
// The candidate vanished after directory enumeration. A
// later stable scan can confirm removal; this run cannot
// safely clear an earlier finding.
markCheckIncomplete(ctx, "php_config_changes")
case errors.Is(err, errPHPIniNonRegular):
markCheckIncomplete(ctx, "php_config_changes")
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: phpIniFindingCheck,
Message: fmt.Sprintf("Special file used as PHP configuration: %s (user: %s)", iniPath, user),
Details: phpIniSpecialDetails,
FilePath: iniPath,
})
case errors.Is(err, errPHPIniTooLarge):
markCheckIncomplete(ctx, "php_config_changes")
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: phpIniFindingCheck,
Message: fmt.Sprintf("PHP configuration too large to inspect: %s (user: %s)", iniPath, user),
Details: phpIniIncompleteDetails,
FilePath: iniPath,
})
default:
recordIncomplete(fmt.Sprintf(
"Could not inspect PHP configuration %s: %v",
iniPath,
err,
))
}
continue
}
if len(dangerous) > 0 {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: phpIniFindingCheck,
Message: fmt.Sprintf("Dangerous PHP configuration: %s (user: %s)", iniPath, user),
Details: fmt.Sprintf("Dangerous settings:\n- %s", strings.Join(dangerous, "\n- ")),
FilePath: iniPath,
})
}
}
if incompleteReason != "" && ctx.Err() == nil {
findings = append(findings, phpIniScanIncompleteFinding(
fmt.Sprintf("PHP configuration scan incomplete for user: %s", user),
incompleteReason,
))
}
}
return findings
}
func phpIniScanIncompleteFinding(message, details string) alert.Finding {
return alert.Finding{
Severity: alert.Warning,
Check: phpIniIncompleteCheck,
Message: message,
Details: details,
}
}
func phpIniFallbackRoots(
ctx context.Context,
user string,
maxRoots int,
maxEntries int,
) ([]string, bool, error) {
if maxRoots <= 0 || maxEntries <= 0 {
return nil, false, nil
}
home := accountHomeDir(user)
var roots []string
entriesSeen := 0
complete, err := forEachPHPIniDirEntry(ctx, home, func(entry os.DirEntry) bool {
entriesSeen++
if entriesSeen > maxEntries {
return false
}
name := entry.Name()
if !entry.IsDir() ||
name == "public_html" ||
name == "mail" ||
name == "etc" ||
name == "logs" ||
name == "ssl" ||
name == "tmp" ||
strings.HasPrefix(name, ".") {
return true
}
if len(roots) >= maxRoots {
return false
}
roots = append(roots, filepath.Join(home, name))
return true
})
if err != nil {
return roots, false, err
}
return roots, complete, nil
}
func readPHPIniFile(path string) ([]byte, error) {
info, err := osFS.Stat(path)
if err != nil {
return nil, err
}
if !info.Mode().IsRegular() {
return nil, errPHPIniNonRegular
}
if info.Size() > PHPConfigMaxBytes {
return nil, errPHPIniTooLarge
}
var f *os.File
var openErr error
_, productionFS := osFS.(realOS)
if productionFS {
// A non-blocking open prevents a regular-file-to-FIFO swap from
// stranding the worker between the Stat above and this open.
// #nosec G304 -- read-only candidate discovered below an account web root.
f, openErr = os.OpenFile(path, os.O_RDONLY|syscall.O_NONBLOCK, 0)
} else {
f, openErr = osFS.Open(path)
}
if openErr == nil {
defer func() { _ = f.Close() }()
openedInfo, statErr := f.Stat()
if statErr != nil {
return nil, statErr
}
if !openedInfo.Mode().IsRegular() {
return nil, errPHPIniNonRegular
}
data, readErr := io.ReadAll(io.LimitReader(f, PHPConfigMaxBytes+1))
if readErr != nil {
return nil, readErr
}
if len(data) > PHPConfigMaxBytes {
return nil, errPHPIniTooLarge
}
return data, nil
}
if productionFS {
return nil, openErr
}
// Map-backed providers cannot return an *os.File, so their bounded test
// data uses the ReadFile hook.
data, readErr := osFS.ReadFile(path)
if readErr != nil {
return nil, readErr
}
if len(data) > PHPConfigMaxBytes {
return nil, errPHPIniTooLarge
}
return data, nil
}
func readAndAssessPHPIniFile(store *state.Store, path string) ([]string, error) {
// The read and state update are one critical section. Concurrent host and
// account scans must not let an older file snapshot overwrite the state
// recorded for a newer snapshot. Locks are sharded by path so unrelated
// account scans can still progress in parallel.
lock := phpIniStateLock(path)
lock.Lock()
defer lock.Unlock()
data, err := readPHPIniFile(path)
if err != nil {
return nil, err
}
return assessPHPIniFile(store, path, hashBytes(data), string(data)), nil
}
func phpIniStateLock(path string) *sync.Mutex {
var hash uint32 = 2166136261
for i := 0; i < len(path); i++ {
hash ^= uint32(path[i])
hash *= 16777619
}
return &phpIniStateLocks[hash%phpIniStateLockShards]
}
// assessPHPIniFile stores the content hash and whether that content needs the
// full change analysis. Re-running the selected analysis for unchanged files
// keeps active findings visible and automatically applies future parser fixes.
// The caller holds the path's state lock so the state update stays ordered
// with its file snapshot.
func assessPHPIniFile(store *state.Store, path, hash, content string) []string {
key := "_phpini:" + path
raw, exists := store.GetRaw(key)
previous := decodePHPIniFileState(raw)
if exists && previous.Assessed && previous.Hash == hash {
if previous.FullAnalysis {
return analyzePHPINI(content)
}
return PHPConfigSecurityBypasses(content)
}
fullAnalysis := exists && previous.Hash != "" && previous.Hash != hash
var dangerous []string
if fullAnalysis {
dangerous = analyzePHPINI(content)
} else {
// New files and legacy hash-only state are judged on the strong,
// low-false-positive bypass signals.
dangerous = PHPConfigSecurityBypasses(content)
}
next := phpIniFileState{
Version: phpIniStateFormat,
Hash: hash,
Assessed: true,
FullAnalysis: fullAnalysis,
}
encoded, err := json.Marshal(next)
if err != nil {
panic(err)
}
store.SetRaw(key, string(encoded))
return append([]string(nil), dangerous...)
}
func decodePHPIniFileState(raw string) phpIniFileState {
var decoded phpIniFileState
if err := json.Unmarshal([]byte(raw), &decoded); err == nil &&
decoded.Version == phpIniStateFormat && decoded.Hash != "" {
return decoded
}
return phpIniFileState{Hash: raw}
}
func analyzePHPINI(content string) []string {
var dangerous []string
seen := make(map[string]bool)
appendFinding := func(message string) {
if !seen[message] {
seen[message] = true
dangerous = append(dangerous, message)
}
}
for _, directives := range parsePHPIniDirectiveSets(content) {
if val, ok := directives["disable_functions"]; ok {
disabled := disabledPHPFunctions(val)
if len(disabled) == 0 {
appendFinding("disable_functions cleared or neutralized (all PHP functions enabled)")
}
for _, fn := range dangerousPHPFunctions {
if !disabled[fn] {
appendFinding(fmt.Sprintf("%s not in disable_functions", fn))
}
}
}
if val, ok := directives["allow_url_include"]; ok && phpIniBoolEnabled(val) {
appendFinding("allow_url_include enabled (remote code inclusion)")
}
if val, ok := directives["open_basedir"]; ok && openBasedirUnrestricted(val) {
appendFinding("open_basedir cleared or set to / (no restriction)")
}
}
return dangerous
}
// collectPHPIniFilesContext walks root up to maxDepth and returns every
// .user.ini and php.ini path found. A negative maxDepth walks all depths. No
// directory is skipped by name -- attackers hide these under node_modules and
// vendor trees. Total directory and entry limits keep a user-controlled tree
// from occupying a scan worker indefinitely.
func collectPHPIniFilesContext(ctx context.Context, root string, maxDepth int) ([]string, bool) {
return collectPHPIniFilesWithBudget(ctx, root, maxDepth, &phpIniWalkBudget{})
}
func collectPHPIniFilesWithBudget(
ctx context.Context,
root string,
maxDepth int,
budget *phpIniWalkBudget,
) ([]string, bool) {
type pendingDir struct {
path string
depth int
observed bool
}
if budget.accountExhausted() {
budget.limitKind = walkLimitDirs
return nil, false
}
budget.startRoot()
var out []string
queue := []pendingDir{{path: root}}
budget.addDir()
complete := true
for len(queue) > 0 {
if ctx.Err() != nil {
return out, false
}
dir := queue[0]
queue = queue[1:]
visitComplete, err := forEachPHPIniDirEntry(ctx, dir.path, func(e os.DirEntry) bool {
budget.addEntry()
if budget.entries > budget.entryLimit() || budget.accountEntries > phpIniWalkAccountMaxEntries {
budget.limitKind = walkLimitEntries
return false
}
name := e.Name()
full := filepath.Join(dir.path, name)
if e.IsDir() {
if maxDepth < 0 || dir.depth < maxDepth {
if budget.dirs >= budget.dirLimit() || budget.accountDirs >= phpIniWalkAccountMaxDirs {
complete = false
budget.limitKind = walkLimitDirs
return true
}
queue = append(queue, pendingDir{path: full, depth: dir.depth + 1, observed: true})
budget.addDir()
}
return true
}
if name == ".user.ini" || name == "php.ini" {
out = append(out, full)
}
return true
})
if err != nil {
if !errors.Is(err, os.ErrNotExist) || dir.observed {
complete = false
}
continue
}
if !visitComplete {
return out, false
}
}
return out, complete
}
func forEachPHPIniDirEntry(
ctx context.Context,
path string,
visit func(os.DirEntry) bool,
) (bool, error) {
if ctx.Err() != nil {
return false, ctx.Err()
}
if _, productionFS := osFS.(realOS); !productionFS {
entries, err := osFS.ReadDir(path)
if err != nil {
return false, err
}
for _, entry := range entries {
if ctx.Err() != nil {
return false, ctx.Err()
}
if !visit(entry) {
return false, nil
}
}
return true, nil
}
// Stream directory entries in bounded batches. os.ReadDir loads the whole
// directory, which lets one account force a large allocation before the
// entry budget can stop the scan.
// #nosec G304 -- directory discovered below an account web root.
dir, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NONBLOCK, 0)
if err != nil {
return false, err
}
defer func() { _ = dir.Close() }()
info, err := dir.Stat()
if err != nil {
return false, err
}
if !info.IsDir() {
return false, syscall.ENOTDIR
}
for {
entries, readErr := dir.ReadDir(phpIniReadDirBatch)
for _, entry := range entries {
if ctx.Err() != nil {
return false, ctx.Err()
}
if !visit(entry) {
return false, nil
}
}
switch {
case errors.Is(readErr, io.EOF):
return true, nil
case readErr != nil:
return false, readErr
}
}
}
// PHPConfigSecurityBypasses returns the strong, low-false-positive signals that a
// PHP ini file weakens security: disable_functions cleared or neutralized,
// allow_url_include enabled, or open_basedir removed. Used for newly-seen
// files, where the noisier per-function diffing of analyzePHPINI would flag
// benign partial disable lists.
func PHPConfigSecurityBypasses(content string) []string {
var out []string
var disableFunctions, allowURLInclude, openBasedir bool
for _, directives := range parsePHPIniDirectiveSets(content) {
if val, ok := directives["disable_functions"]; ok && DisableFunctionsNeutralized(val) {
disableFunctions = true
}
if val, ok := directives["allow_url_include"]; ok && phpIniBoolEnabled(val) {
allowURLInclude = true
}
if val, ok := directives["open_basedir"]; ok && openBasedirUnrestricted(val) {
openBasedir = true
}
}
if disableFunctions {
out = append(out, "disable_functions cleared or neutralized (dangerous PHP functions enabled)")
}
if allowURLInclude {
out = append(out, "allow_url_include enabled (remote code inclusion)")
}
if openBasedir {
out = append(out, "open_basedir cleared (filesystem sandbox removed)")
}
return out
}
// PHPConfigRealtimeRootPatterns returns the startup-time path patterns used to
// admit php.ini writes into the fanotify analyzer. Explicit account_roots are
// authoritative. On cPanel, account homes are intentionally broader than the
// primary public_html root because addon domains can live anywhere below an
// account and on alternate home mounts.
func PHPConfigRealtimeRootPatterns(cfg *config.Config) []string {
if cfg != nil && len(cfg.AccountRoots) > 0 {
return WebRootPatterns(cfg)
}
if len(WebRootPatterns(cfg)) == 0 {
return nil
}
patterns := accountHomePatterns()
seen := make(map[string]struct{}, len(patterns))
for _, p := range patterns {
seen[p] = struct{}{}
}
data, err := osFS.ReadFile(userdataDomainsPath)
if err != nil {
return patterns
}
vhosts, _ := parseUserdataDomainRootsChecked(string(data))
for _, vhost := range vhosts {
pattern := phpConfigRealtimeRootPattern(vhost.user, vhost.docroot)
if pattern == "" {
continue
}
if _, exists := seen[pattern]; exists {
continue
}
seen[pattern] = struct{}{}
patterns = append(patterns, pattern)
}
return patterns
}
// RealtimeDocumentRootPatterns returns the served roots used by realtime
// content detectors. cPanel's domain map is authoritative for addon domains
// and accounts on alternate home mounts; the conventional public_html glob
// remains the fallback when that map is unavailable.
func RealtimeDocumentRootPatterns(cfg *config.Config) []string {
patterns := WebRootPatterns(cfg)
if (cfg != nil && len(cfg.AccountRoots) > 0) || !platform.Detect().IsCPanel() {
return patterns
}
data, err := osFS.ReadFile(userdataDomainsPath)
if err != nil {
return patterns
}
vhosts, _ := parseUserdataDomainRootsChecked(string(data))
seen := make(map[string]struct{}, len(patterns)+len(vhosts))
for _, pattern := range patterns {
seen[filepath.Clean(pattern)] = struct{}{}
}
for _, vhost := range vhosts {
root := filepath.Clean(vhost.docroot)
if _, exists := seen[root]; exists {
continue
}
seen[root] = struct{}{}
patterns = append(patterns, root)
}
return patterns
}
func phpConfigRealtimeRootPattern(user, docroot string) string {
clean := filepath.Clean(docroot)
if !filepath.IsAbs(clean) {
return ""
}
parts := strings.Split(strings.TrimPrefix(clean, string(filepath.Separator)), string(filepath.Separator))
if len(parts) >= 2 && strings.HasPrefix(parts[0], "home") && parts[1] == user {
return filepath.Join(string(filepath.Separator), parts[0], "*")
}
return clean
}
var dangerousPHPFunctions = []string{
"exec",
"system",
"passthru",
"shell_exec",
"popen",
"proc_open",
"pcntl_exec",
}
// DisableFunctionsNeutralized reports whether a disable_functions value fails
// to actually disable any dangerous function -- empty, "none", or set to junk
// (the `disable_functions=ByPassed By 0xNix` camouflage). A genuine hardening
// list names at least one exact dangerous function; substrings such as
// "systemd" do not disable system().
func DisableFunctionsNeutralized(val string) bool {
return len(disabledPHPFunctions(val)) == 0
}
func disabledPHPFunctions(val string) map[string]bool {
disabled := make(map[string]bool)
val = strings.ToLower(normalizePHPIniValue(val))
items := strings.FieldsFunc(val, func(r rune) bool {
return r == ',' || r == ' '
})
for _, name := range items {
for _, dangerous := range dangerousPHPFunctions {
if name == dangerous {
disabled[dangerous] = true
break
}
}
}
return disabled
}
func parsePHPIniDirectiveSets(content string) []map[string]string {
sets := []map[string]string{{}}
sections := make(map[string]int)
current := 0
for _, raw := range strings.Split(content, "\n") {
line := strings.TrimSpace(raw)
line = strings.TrimPrefix(line, "\ufeff")
if line == "" || strings.HasPrefix(line, ";") || strings.HasPrefix(line, "#") {
continue
}
if strings.HasPrefix(line, "[") {
end := strings.IndexByte(line, ']')
if end <= 1 {
continue
}
section := strings.TrimSpace(line[1:end])
upperSection := strings.ToUpper(section)
var sectionKey string
switch {
case strings.HasPrefix(upperSection, "PATH="):
sectionKey = "PATH=" + strings.TrimSpace(section[len("PATH="):])
case strings.HasPrefix(upperSection, "HOST="):
sectionKey = "HOST=" + strings.ToLower(strings.TrimSpace(section[len("HOST="):]))
default:
// Ordinary php.ini section headings are ignored by PHP and
// return parsing to the global directive set. Only PATH and
// HOST headings create conditional settings.
current = 0
continue
}
if index, exists := sections[sectionKey]; exists {
current = index
} else {
current = len(sets)
sections[sectionKey] = current
sets = append(sets, make(map[string]string))
}
continue
}
parts := strings.SplitN(line, "=", 2)
if len(parts) != 2 {
continue
}
key := strings.ToLower(strings.TrimSpace(parts[0]))
switch key {
case "disable_functions", "allow_url_include", "open_basedir":
sets[current][key] = parts[1]
}
}
return sets
}
func normalizePHPIniValue(value string) string {
value, _ = normalizePHPIniValueWithExpression(value)
return value
}
func normalizePHPIniValueWithExpression(value string) (string, bool) {
value = strings.TrimSpace(stripPHPIniInlineComment(value))
joined, expression := joinPHPIniQuotedFragments(value)
return strings.TrimSpace(joined), expression
}
// joinPHPIniQuotedFragments mirrors the INI scanner's concatenation of quoted
// and unquoted fragments. For example, ex"ec",sys"tem" is the effective value
// exec,system. An unclosed quote makes PHP reject the value; return empty so
// security restrictions fail closed instead of treating the malformed token
// as an effective hardening value.
func joinPHPIniQuotedFragments(value string) (string, bool) {
var out strings.Builder
out.Grow(len(value))
var quote byte
expression := false
for i := 0; i < len(value); i++ {
ch := value[i]
if quote == 0 {
if ch == '"' || ch == '\'' {
quote = ch
continue
}
if strings.ContainsRune("|&^~!()", rune(ch)) {
expression = true
}
out.WriteByte(ch)
continue
}
if ch == quote {
quote = 0
continue
}
if ch == '\\' && i+1 < len(value) {
next := value[i+1]
if next == quote || next == '\\' {
out.WriteByte(next)
i++
continue
}
}
out.WriteByte(ch)
}
if quote != 0 {
return "", false
}
return out.String(), expression
}
func stripPHPIniInlineComment(value string) string {
var quote byte
escaped := false
for i := 0; i < len(value); i++ {
ch := value[i]
if quote != 0 {
if quote == '"' && ch == '\\' && !escaped {
escaped = true
continue
}
if ch == quote && !escaped {
quote = 0
}
escaped = false
continue
}
if ch == '"' || ch == '\'' {
quote = ch
continue
}
if ch == ';' {
return value[:i]
}
}
return value
}
func phpIniBoolEnabled(value string) bool {
value, expression := normalizePHPIniValueWithExpression(value)
value = strings.ToLower(value)
switch value {
case "1", "on", "true", "yes":
return true
}
if expression {
evaluated, ok := parsePHPIniIntExpression(value)
// A syntactically unusual expression is not a trustworthy disabled
// value. Treat it as enabled rather than let unsupported constants or
// excessive nesting bypass the security check.
return !ok || evaluated != 0
}
// PHP converts the leading signed integer portion to bool: "1foo" and
// "1e-2" are enabled, while "0.5" is not.
i := 0
if i < len(value) && (value[i] == '+' || value[i] == '-') {
i++
}
start := i
nonZero := false
for i < len(value) && value[i] >= '0' && value[i] <= '9' {
nonZero = nonZero || value[i] != '0'
i++
}
return i > start && nonZero
}
const phpIniExpressionMaxDepth = 64
type phpIniIntExprParser struct {
input string
pos int
limitExceeded bool
}
func parsePHPIniIntExpression(value string) (int64, bool) {
p := phpIniIntExprParser{input: value}
result, ok := p.parseBinary(0)
if p.limitExceeded {
return 1, true
}
p.skipSpace()
return result, ok && p.pos == len(p.input)
}
func (p *phpIniIntExprParser) parseBinary(depth int) (int64, bool) {
left, ok := p.parseUnary(depth)
if !ok {
return 0, false
}
for {
p.skipSpace()
if p.pos >= len(p.input) {
return left, true
}
op := p.input[p.pos]
if op != '|' && op != '&' && op != '^' {
return left, true
}
p.pos++
right, valid := p.parseUnary(depth)
if !valid {
return 0, false
}
switch op {
case '|':
left |= right
case '&':
left &= right
case '^':
left ^= right
}
}
}
func (p *phpIniIntExprParser) parseUnary(depth int) (int64, bool) {
if depth >= phpIniExpressionMaxDepth {
p.limitExceeded = true
return 0, false
}
p.skipSpace()
switch {
case p.take('!'):
value, ok := p.parseUnary(depth + 1)
if !ok {
return 0, false
}
if value == 0 {
return 1, true
}
return 0, true
case p.take('~'):
value, ok := p.parseUnary(depth + 1)
return ^value, ok
case p.take('+'):
return p.parseUnary(depth + 1)
case p.take('-'):
value, ok := p.parseUnary(depth + 1)
return -value, ok
case p.take('('):
value, ok := p.parseBinary(depth + 1)
p.skipSpace()
if !ok || !p.take(')') {
return 0, false
}
return value, true
default:
return p.parseDecimal()
}
}
func (p *phpIniIntExprParser) parseDecimal() (int64, bool) {
p.skipSpace()
start := p.pos
var value uint64
const maxInt64 = uint64(^uint64(0) >> 1)
for p.pos < len(p.input) {
ch := p.input[p.pos]
if ch < '0' || ch > '9' {
break
}
digit := uint64(ch - '0')
if value > (maxInt64-digit)/10 {
value = maxInt64
} else {
value = value*10 + digit
}
p.pos++
}
return int64(value), p.pos > start
}
func (p *phpIniIntExprParser) skipSpace() {
for p.pos < len(p.input) {
switch p.input[p.pos] {
case ' ', '\t', '\r', '\n', '\v', '\f':
p.pos++
default:
return
}
}
}
func (p *phpIniIntExprParser) take(ch byte) bool {
if p.pos >= len(p.input) || p.input[p.pos] != ch {
return false
}
p.pos++
return true
}
func openBasedirUnrestricted(value string) bool {
value = normalizePHPIniValue(value)
if value == "" {
return true
}
for _, dir := range strings.Split(value, string(os.PathListSeparator)) {
if filepath.Clean(strings.TrimSpace(dir)) == string(filepath.Separator) {
return true
}
}
return false
}
package checks
import (
"path/filepath"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/contenttype"
)
// PHPExecutionOverlay is an immutable snapshot of inherited .htaccess PHP
// handler mappings for one directory. Realtime monitors can cache it and test
// arbitrary filenames without duplicating the periodic scanner's parser.
type PHPExecutionOverlay struct {
overlay phpHandlerOverlay
}
// ResolvePHPExecutionOverlay reconstructs the PHP handler mappings inherited
// by fileDir from docroot through the directory's own .htaccess file.
func ResolvePHPExecutionOverlay(docroot, fileDir string) PHPExecutionOverlay {
return PHPExecutionOverlay{overlay: reconstructOverlay(docroot, fileDir)}
}
// Executes reports whether nameLower is handled as PHP by this overlay.
func (o PHPExecutionOverlay) Executes(nameLower string) bool {
return o.overlay.executes(nameLower)
}
// isExecutablePHPName is the stock-handler extension gate; per-directory
// .htaccess handler remappings are layered on top via phpHandlerOverlay.
func isExecutablePHPName(nameLower string) bool {
return contenttype.IsExecutablePHPName(nameLower)
}
// phpHandlerOverlay carries the extra PHP execution mappings discovered from
// .htaccess files while walking a directory tree. Apache merges a parent
// directory's directives into its children, so the overlay accumulates down
// the recursion: a mapping declared in a parent applies to every descendant.
type phpHandlerOverlay struct {
// exts holds extra ".ext" entries (lowercase, leading dot) that a local
// AddHandler/AddType maps to a PHP handler, e.g. "AddHandler
// application/x-httpd-php .inc".
exts map[string]struct{}
// names holds exact lowercase basenames matched by a <Files> container.
names map[string]struct{}
// scanAll is set when a SetHandler/ForceType routes the PHP interpreter
// for the whole directory with no extension filter. Every file in the
// subtree then executes as PHP and must be content-analysed.
scanAll bool
// unrestricted marks a <FilesMatch> context whose pattern selects files
// by name rather than by extension ("logo", "^(config|data)$", or any
// alternative without a "\." marker). A PHP handler inside it executes
// files no extension list describes, so it is treated like a
// directory-wide handler.
unrestricted bool
}
func (o phpHandlerOverlay) active() bool {
return o.scanAll || len(o.exts) > 0 || len(o.names) > 0
}
// executes reports whether a file named nameLower (lowercased) runs as PHP
// under this overlay, either by a stock extension, a directory-wide handler,
// or an .htaccess-mapped extension.
func (o phpHandlerOverlay) executes(nameLower string) bool {
if o.scanAll {
return true
}
if isExecutablePHPName(nameLower) {
return true
}
if _, ok := o.names[nameLower]; ok {
return true
}
for ext := range o.exts {
if strings.HasSuffix(nameLower, ext) {
return true
}
}
return false
}
// mergeHtaccess returns a new overlay combining the receiver (inherited from
// the parent directory) with any PHP handler directives found in the .htaccess
// at dirHtaccessContent. The receiver is never mutated, so sibling directories
// do not see each other's mappings.
func (o phpHandlerOverlay) mergeHtaccess(content []byte) phpHandlerOverlay {
if len(content) == 0 {
return o
}
parsed := parsePHPHandlerDirectives(content)
if !parsed.active() {
return o
}
merged := phpHandlerOverlay{scanAll: o.scanAll || parsed.scanAll}
if len(o.exts) > 0 || len(parsed.exts) > 0 {
merged.exts = make(map[string]struct{}, len(o.exts)+len(parsed.exts))
for e := range o.exts {
merged.exts[e] = struct{}{}
}
for e := range parsed.exts {
merged.exts[e] = struct{}{}
}
}
if len(o.names) > 0 || len(parsed.names) > 0 {
merged.names = make(map[string]struct{}, len(o.names)+len(parsed.names))
for name := range o.names {
merged.names[name] = struct{}{}
}
for name := range parsed.names {
merged.names[name] = struct{}{}
}
}
return merged
}
// parsePHPHandlerDirectives extracts extension-to-PHP mappings from .htaccess
// content. It recognises:
//
// AddHandler <php-handler> .ext [.ext...]
// AddType <php-mime> .ext [.ext...]
// SetHandler <php-handler> (no extension -> whole directory)
// ForceType <php-mime> (no extension -> whole directory)
//
// A directive counts as PHP when its handler/MIME token names PHP directly or
// routes through proxy-fcgi. x-httpd-php-source is excluded because it renders
// highlighted source instead of executing. Matching is case-insensitive.
func parsePHPHandlerDirectives(content []byte) phpHandlerOverlay {
var overlay phpHandlerOverlay
var contexts []phpHandlerOverlay
for _, logical := range joinHtaccessContinuations(strings.Split(string(content), "\n")) {
line := strings.TrimSpace(logical.text)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if ctx, ok := openPHPHandlerContext(line); ok {
contexts = append(contexts, ctx)
continue
}
if closesPHPHandlerContext(line) {
if len(contexts) > 0 {
contexts = contexts[:len(contexts)-1]
}
continue
}
fields := apacheDirectiveFields(line)
if len(fields) < 2 {
continue
}
directive := strings.ToLower(fields[0])
handler := fields[1]
switch directive {
case "addhandler", "addtype":
if !handlerIsPHP(handler) {
continue
}
// Remaining fields are extensions.
addExtensions(&overlay, normalizedExts(fields[2:]))
if len(fields) == 2 {
mergeContext(&overlay, contexts)
}
case "sethandler", "forcetype":
if !handlerIsPHP(handler) {
continue
}
exts := normalizedExts(fields[2:])
if len(exts) > 0 {
addExtensions(&overlay, exts)
continue
}
if mergeContext(&overlay, contexts) {
continue
}
if len(contexts) > 0 {
continue
}
// No extension or file filter: the handler applies to every
// file in this directory.
overlay.scanAll = true
}
}
return overlay
}
func handlerIsPHP(token string) bool {
token = strings.ToLower(strings.Trim(strings.TrimSpace(token), `"'`))
// The source viewer renders highlighted source instead of executing.
if strings.Contains(token, "php-source") {
return false
}
if strings.Contains(token, "php") {
return true
}
// cPanel/PHP-FPM .htaccess wiring can use a custom socket alias whose
// path does not include the literal "php". A proxy-fcgi handler still
// routes matching files to an executable backend, so remapped extensions
// must be treated as PHP-executed for scanning and .htaccess alerts.
return strings.HasPrefix(token, "proxy:") && strings.Contains(token, "fcgi://")
}
// normalizeExt turns an .htaccess extension token into a lowercase
// leading-dot extension, rejecting anything that is not a plain extension
// token (e.g. a stray flag or MIME fragment).
func normalizeExt(token string) string {
t := strings.ToLower(strings.Trim(strings.TrimSpace(token), `"'`))
t = strings.TrimPrefix(t, ".")
if t == "" {
return ""
}
for _, r := range t {
alnum := (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9')
if !alnum && r != '_' && r != '-' {
return ""
}
}
return "." + t
}
// htaccessLogicalLine is one Apache directive after joining physical
// continuation lines. text is the joined directive (continuation backslashes
// removed); lines are the original physical lines it spans, so a per-line
// rewrite can drop or keep them together. start is the 0-based index of the
// first physical line.
type htaccessLogicalLine struct {
text string
lines []string
start int
}
// joinHtaccessContinuations groups physical .htaccess lines into logical
// directives, honoring Apache's trailing-backslash line continuation: a line
// ending in "\" is joined with the next. Without this, a directive split as
// "AddHandler ...php \" + ".jpg" reads as two harmless physical lines and
// every per-line scanner misses the remap.
func joinHtaccessContinuations(physical []string) []htaccessLogicalLine {
var out []htaccessLogicalLine
for i := 0; i < len(physical); {
start := i
var sb strings.Builder
var span []string
for {
cur := physical[i]
span = append(span, cur)
body, continues := htaccessContinuationBody(cur, i < len(physical)-1)
if continues {
sb.WriteString(body)
i++
continue
}
sb.WriteString(body)
break
}
out = append(out, htaccessLogicalLine{text: sb.String(), lines: span, start: start})
i++
}
return out
}
func htaccessContinuationBody(line string, hasNext bool) (string, bool) {
// On CRLF input (strings.Split keeps the trailing "\r") the continuation
// backslash sits before the carriage return, so strip it before the suffix
// test and from the joined text.
body := strings.TrimRight(line, "\r")
// A trailing backslash continues only when a next line exists; a backslash on
// the final line is left literal, matching Apache.
if !hasNext || !strings.HasSuffix(body, `\`) {
return body, false
}
if htaccessContainerLineKeepsTrailingBackslash(body) {
return body, false
}
return body[:len(body)-1], true
}
func htaccessContainerLineKeepsTrailingBackslash(body string) bool {
candidate := strings.TrimSpace(strings.TrimSuffix(body, `\`))
if !strings.Contains(candidate, ">") {
return false
}
if _, ok := openPHPHandlerContext(candidate); ok {
return true
}
return closesPHPHandlerContext(candidate)
}
func apacheDirectiveFields(line string) []string {
fields := strings.Fields(line)
for i, field := range fields {
if strings.HasPrefix(field, "#") {
return fields[:i]
}
}
return fields
}
func normalizedExts(tokens []string) []string {
var exts []string
for _, token := range tokens {
if ext := normalizeExt(token); ext != "" {
exts = append(exts, ext)
}
}
return exts
}
func addExtensions(overlay *phpHandlerOverlay, exts []string) {
if len(exts) == 0 {
return
}
for _, ext := range exts {
if isExecutablePHPName("x" + ext) {
continue
}
if overlay.exts == nil {
overlay.exts = make(map[string]struct{}, len(exts))
}
overlay.exts[ext] = struct{}{}
}
}
func mergeContext(overlay *phpHandlerOverlay, contexts []phpHandlerOverlay) bool {
merged := false
for _, ctx := range contexts {
if ctx.unrestricted {
// The container selects files by name: the handler can reach
// any file in the directory, so every file must be analysed.
overlay.scanAll = true
merged = true
}
if len(ctx.exts) > 0 {
if overlay.exts == nil {
overlay.exts = make(map[string]struct{}, len(ctx.exts))
}
for ext := range ctx.exts {
overlay.exts[ext] = struct{}{}
merged = true
}
}
if len(ctx.names) > 0 {
if overlay.names == nil {
overlay.names = make(map[string]struct{}, len(ctx.names))
}
for name := range ctx.names {
overlay.names[name] = struct{}{}
merged = true
}
}
}
return merged
}
func openPHPHandlerContext(line string) (phpHandlerOverlay, bool) {
lower := strings.ToLower(line)
switch {
case strings.HasPrefix(lower, "<filesmatch"):
pattern := apacheContainerArgument(line)
return overlayForFilesMatch(pattern), true
case strings.HasPrefix(lower, "<files"):
name := apacheContainerArgument(line)
name = strings.ToLower(strings.TrimSpace(name))
if name == "" || strings.Contains(name, "/") {
return phpHandlerOverlay{}, true
}
if strings.ContainsAny(name, "*?[") {
overlay := phpHandlerOverlay{}
addExtensions(&overlay, []string{filepath.Ext(name)})
return overlay, true
}
overlay := phpHandlerOverlay{names: map[string]struct{}{name: {}}}
return overlay, true
default:
return phpHandlerOverlay{}, false
}
}
func closesPHPHandlerContext(line string) bool {
lower := strings.ToLower(line)
return strings.HasPrefix(lower, "</filesmatch") || strings.HasPrefix(lower, "</files")
}
func apacheContainerArgument(line string) string {
start := strings.IndexAny(line, " \t")
end := strings.LastIndexByte(line, '>')
if start < 0 || end <= start {
return ""
}
arg := strings.TrimSpace(line[start:end])
if len(arg) >= 2 {
quote := arg[0]
if (quote == '"' || quote == '\'') && arg[len(arg)-1] == quote {
arg = arg[1 : len(arg)-1]
}
}
return arg
}
func overlayForFilesMatch(pattern string) phpHandlerOverlay {
overlay := phpHandlerOverlay{}
addExtensions(&overlay, extensionsFromFilesMatchPattern(pattern))
overlay.unrestricted = filesMatchSelectsByName(pattern)
return overlay
}
// filesMatchSelectsByName reports whether a FilesMatch pattern can match a
// filename without requiring a literal dot. An empty or unsupported pattern
// is treated as unrestricted so the content scanner fails toward coverage.
func filesMatchSelectsByName(pattern string) bool {
pattern = strings.TrimSpace(pattern)
if pattern == "" {
return true
}
p := filesMatchRegexParser{pattern: pattern}
requiresDot, valid := p.alternation(0)
return !valid || p.pos != len(pattern) || !requiresDot
}
type filesMatchRegexParser struct {
pattern string
pos int
}
// alternation returns true only when every branch requires a literal dot.
// That conservative proof is enough to keep extension-only handlers narrow;
// regex constructs outside the small parser fall back to scanning all names.
func (p *filesMatchRegexParser) alternation(stop byte) (bool, bool) {
allRequireDot := true
for {
requiresDot, valid := p.sequence(stop)
if !valid {
return false, false
}
allRequireDot = allRequireDot && requiresDot
if p.pos < len(p.pattern) && p.pattern[p.pos] == '|' {
p.pos++
continue
}
if stop != 0 {
if p.pos >= len(p.pattern) || p.pattern[p.pos] != stop {
return false, false
}
p.pos++
}
return allRequireDot, true
}
}
func (p *filesMatchRegexParser) sequence(stop byte) (bool, bool) {
requiresDot := false
for p.pos < len(p.pattern) {
if p.pattern[p.pos] == '|' || (stop != 0 && p.pattern[p.pos] == stop) {
break
}
atomRequiresDot, valid := p.atom()
if !valid {
return false, false
}
if p.pos < len(p.pattern) {
switch p.pattern[p.pos] {
case '*', '?':
atomRequiresDot = false
p.pos++
case '+':
p.pos++
case '{':
minimum, next, ok := regexRepeatMinimum(p.pattern, p.pos)
if !ok {
return false, false
}
p.pos = next
if minimum == 0 {
atomRequiresDot = false
}
}
}
requiresDot = requiresDot || atomRequiresDot
}
return requiresDot, true
}
func (p *filesMatchRegexParser) atom() (bool, bool) {
if p.pos >= len(p.pattern) {
return false, false
}
c := p.pattern[p.pos]
p.pos++
switch c {
case '\\':
if p.pos >= len(p.pattern) {
return false, false
}
escaped := p.pattern[p.pos]
p.pos++
return escaped == '.', true
case '[':
return p.characterClass()
case '(':
parseRequirement := true
if p.pos < len(p.pattern) && p.pattern[p.pos] == '?' {
p.pos++
switch {
case p.pos < len(p.pattern) && p.pattern[p.pos] == ':':
p.pos++
case p.regexFlagGroup():
default:
// Lookarounds, named groups and other PCRE extensions are
// parsed for balance but not used as a coverage proof.
parseRequirement = false
}
}
requiresDot, valid := p.alternation(')')
return parseRequirement && requiresDot, valid
case ')':
return false, false
default:
return false, true
}
}
func (p *filesMatchRegexParser) regexFlagGroup() bool {
start := p.pos
for p.pos < len(p.pattern) {
c := p.pattern[p.pos]
if c == ':' {
if p.pos == start {
return false
}
p.pos++
return true
}
if (c < 'a' || c > 'z') && (c < 'A' || c > 'Z') && c != '-' {
p.pos = start
return false
}
p.pos++
}
p.pos = start
return false
}
func (p *filesMatchRegexParser) characterClass() (bool, bool) {
for p.pos < len(p.pattern) {
c := p.pattern[p.pos]
p.pos++
if c == ']' {
// Even a class that currently contains only a dot is not used as
// an extension proof: the extension extractor deliberately handles
// only escaped-dot forms. Treating it as unrestricted keeps the two
// decisions aligned and cannot lose content coverage.
return false, true
}
if c == '\\' {
if p.pos >= len(p.pattern) {
return false, false
}
p.pos++
}
}
return false, false
}
func regexRepeatMinimum(pattern string, start int) (int, int, bool) {
end := strings.IndexByte(pattern[start+1:], '}')
if end < 0 {
return 0, start, false
}
end += start + 1
fields := strings.Split(pattern[start+1:end], ",")
if len(fields) > 2 || fields[0] == "" {
return 0, start, false
}
minimum, err := strconv.Atoi(fields[0])
if err != nil || minimum < 0 {
return 0, start, false
}
if len(fields) == 2 && fields[1] != "" {
maximum, err := strconv.Atoi(fields[1])
if err != nil || maximum < minimum {
return 0, start, false
}
}
return minimum, end + 1, true
}
// topLevelAlternatives splits a regex on "|" at nesting depth zero only,
// ignoring "|" inside groups, character classes and after a backslash.
func topLevelAlternatives(pattern string) []string {
var out []string
depth := 0
inClass := false
start := 0
for i := 0; i < len(pattern); i++ {
switch c := pattern[i]; {
case c == '\\':
i++
case inClass:
if c == ']' {
inClass = false
}
case c == '[':
inClass = true
case c == '(':
depth++
case c == ')':
if depth > 0 {
depth--
}
case c == '|' && depth == 0:
out = append(out, pattern[start:i])
start = i + 1
}
}
return append(out, pattern[start:])
}
func extensionsFromFilesMatchPattern(pattern string) []string {
pattern = strings.ToLower(pattern)
seen := make(map[string]struct{})
var exts []string
add := func(ext string) {
ext = normalizeExt(ext)
if ext == "" {
return
}
if _, ok := seen[ext]; ok {
return
}
seen[ext] = struct{}{}
exts = append(exts, ext)
}
for i := 0; i+1 < len(pattern); i++ {
if pattern[i] != '\\' || pattern[i+1] != '.' {
continue
}
j := i + 2
if j < len(pattern) && pattern[j] == '(' {
end := strings.IndexByte(pattern[j+1:], ')')
if end >= 0 {
group := pattern[j+1 : j+1+end]
group = strings.TrimPrefix(group, "?:")
for _, part := range strings.Split(group, "|") {
add(part)
}
}
continue
}
start := j
for j < len(pattern) {
c := pattern[j]
alnum := (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9')
if !alnum && c != '_' && c != '-' {
break
}
j++
}
add(pattern[start:j])
i = j
}
return exts
}
func phpPathExecutes(path, nameLower string) bool {
if isExecutablePHPName(nameLower) {
return true
}
overlay := phpHandlerOverlay{}
for _, dir := range htaccessAncestorDirs(path) {
htaccess, ok, err := readHtaccessBounded(filepath.Join(dir, ".htaccess"))
if htaccessOversized(ok, err) {
// A .htaccess too large to read may carry any handler: fail
// toward coverage and treat the name as executable.
return true
}
if err == nil && ok {
overlay = overlay.mergeHtaccess(htaccess)
}
}
return overlay.executes(nameLower)
}
func htaccessAncestorDirs(path string) []string {
dir := filepath.Clean(filepath.Dir(path))
if dir == "." {
return nil
}
var dirs []string
for {
dirs = append(dirs, dir)
if stopHtaccessAncestorWalk(dir) {
break
}
parent := filepath.Dir(dir)
if parent == dir {
break
}
dir = parent
}
for i, j := 0, len(dirs)-1; i < j; i, j = i+1, j-1 {
dirs[i], dirs[j] = dirs[j], dirs[i]
}
return dirs
}
func stopHtaccessAncestorWalk(dir string) bool {
// Stop at an account home: "<root>/<account>" is not inside an account,
// while "<root>/<account>/x" is.
if _, _, inside := accountRootOf(dir); inside {
return false
}
parent := filepath.Dir(dir)
return isAccountRoot(parent) && filepath.Base(dir) != "" && dir != parent
}
package checks
import (
"context"
"sync"
"github.com/pidginhost/csm/internal/phptaint"
)
// phpTaintWorker is the isolated analyzer this adapter routes through. It is
// nil until the daemon supervises one.
var (
phpTaintWorkerMu sync.RWMutex
phpTaintWorker PHPTaintAnalyzer
)
// PHPTaintAnalyzer is what the daemon supplies: an analyzer that runs in a
// process the caller can kill.
type PHPTaintAnalyzer interface {
Analyze(ctx context.Context, src []byte) phptaint.Report
}
// SetPHPTaintAnalyzer installs the supervised worker. Passing nil removes it,
// after which every candidate is reported as an unexamined coverage gap rather
// than analyzed in this process.
func SetPHPTaintAnalyzer(a PHPTaintAnalyzer) {
phpTaintWorkerMu.Lock()
defer phpTaintWorkerMu.Unlock()
phpTaintWorker = a
}
// phpTaintAnalyzerReady reports whether an isolated analyzer is available.
//
// The consumer does not dispatch without one. Analysis cannot safely happen in
// this process, so the alternative would be recording every candidate file as
// an unexamined gap -- turning "this feature is not active here" into a
// per-file coverage report on every scan, which tells an operator nothing they
// can act on. A gap is for content that should have been examined and was not.
func phpTaintAnalyzerReady() bool {
phpTaintWorkerMu.RLock()
defer phpTaintWorkerMu.RUnlock()
return phpTaintWorker != nil
}
// defaultPHPTaintAnalyze routes to the supervised worker, and reports a
// coverage gap when there is none.
//
// It deliberately does NOT fall back to calling phptaint.Analyze here. The
// parser can enter an infinite loop on attacker-controlled input; a loop is not
// a panic, so recover() cannot catch it, and the parser never checks context,
// so no deadline in this process can stop it. An in-process fallback would take
// a core from the daemon permanently the first time a crafted file was scanned
// -- which is the entire failure this indirection exists to prevent. Reporting
// a gap loses coverage on that file; the fallback would lose the daemon.
func defaultPHPTaintAnalyze(ctx context.Context, src []byte) phptaint.Report {
phpTaintWorkerMu.RLock()
worker := phpTaintWorker
phpTaintWorkerMu.RUnlock()
if worker == nil {
return phptaint.Report{
Status: phptaint.StatusWorkerFailure,
Reason: "no isolated php analyzer configured; content not examined",
}
}
return worker.Analyze(ctx, src)
}
// phpTaintAnalyze is indirected so tests can substitute the analyzer at the
// adapter boundary: forcing a panic or a specific status proves containment
// without needing a live worker or pathological PHP.
var phpTaintAnalyze = defaultPHPTaintAnalyze
// runPHPTaintAnalysis is the adapter recovery boundary around the analyzer
// call. The analyzer converts its own internal panics to StatusPanic, but a
// panic in adapter-side code around the call must have the same contained
// outcome instead of escaping through the owning check.
func runPHPTaintAnalysis(ctx context.Context, data []byte) (report phptaint.Report) {
defer func() {
if r := recover(); r != nil {
report = phptaint.Report{Status: phptaint.StatusPanic, Reason: "panic"}
}
}()
return phpTaintAnalyze(ctx, data)
}
package checks
import (
"bytes"
"context"
"crypto/sha256"
"encoding/binary"
"fmt"
"hash"
"os"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/phptaint"
)
// phpTaintDeepCursorCheck is the host-scope scan-cursor key for the scheduled
// PHP taint consumer of the shared deep-content walk. It is distinct from the
// JS consumer's key so the two advance independently: they admit different
// files, so a shared cursor would let one consumer's progress hide the other's
// unscanned remainder.
const phpTaintDeepCursorCheck = logicalOwnerPHPTaintDeep
// phpTaintDeepPerFileTimeout bounds one file's analysis as seen by this
// consumer. The supervised worker applies its own, shorter deadline and kills
// the process when it expires; this outer bound only covers the case where the
// worker layer itself becomes unresponsive.
const phpTaintDeepPerFileTimeout = 30 * time.Second
// Display bounds mirror the JS consumer's: message, details, and diagnostic
// example paths.
// maxPHPTaintGapPaths bounds the exact paths one run retains for carry-forward.
//
// The supervisor rejects non-candidates in-process, but an unavailable worker
// can still leave every candidate unexamined. Bound retained path strings for
// large scans. Past the bound the run stops enumerating and reports itself as
// unable to enumerate, which suppresses the purge wholesale rather than
// carrying forward an arbitrary prefix.
const maxPHPTaintGapPaths = 50_000
const (
phpTaintMessageMaxBytes = 512
phpTaintDetailsMaxBytes = 2048
phpTaintExampleMaxBytes = 256
)
// phpTaintGapCollector aggregates per-path PHP coverage gaps for one deep run:
// exact paths feed the carry-forward, counts and one example per status feed
// the php_taint_scan_incomplete diagnostic. A non-completed status is never
// counted as a clean file.
type phpTaintGapCollector struct {
paths map[string]struct{}
pathAliases map[string]struct{}
aliasesByPath map[string][]string
byStatus map[string]int
example map[string]string
recordCoverage func([]string)
resolveAliases func(string) ([]string, bool)
defeatInputs [][sha256.Size]byte
defeatOverflow hash.Hash
// unknown counts walk failures whose affected paths cannot be enumerated
// (an unreadable directory, a failed Lstat that may hide one). They are
// kept apart from paths because carry-forward needs exact paths, but they
// must still reach the operator: without this the loss is recorded only in
// a boolean that suppresses the purge, and a host running the PHP consumer
// without the YARA one is told nothing at all.
unknown int
unknownExample string
// pathsTruncated records that the exact-path set hit its bound, so the
// carry-forward can no longer be trusted to cover every gapped path.
pathsTruncated bool
}
func newPHPTaintGapCollector() *phpTaintGapCollector {
return &phpTaintGapCollector{
paths: map[string]struct{}{},
pathAliases: map[string]struct{}{},
aliasesByPath: map[string][]string{},
byStatus: map[string]int{},
example: map[string]string{},
}
}
func (g *phpTaintGapCollector) record(path, status string) {
g.recordSnapshot(path, status, "")
}
func (g *phpTaintGapCollector) recordSnapshot(path, status, contentSHA256 string) {
if isPHPTaintAnalyzerDefeatStatus(status) {
identity := sha256.Sum256([]byte(path + "\x00" + status + "\x00" + contentSHA256))
if len(g.defeatInputs) < maxPHPTaintGapPaths {
g.defeatInputs = append(g.defeatInputs, identity)
} else {
// Beyond the memory bound, retain all evidence in a streaming
// digest. Order changes may re-alert in this extreme case, but
// a dismissal must never hide failures beyond the retained set.
if g.defeatOverflow == nil {
g.defeatOverflow = sha256.New()
}
_, _ = g.defeatOverflow.Write(identity[:])
}
}
aliases, retained := g.aliasesByPath[path]
if !retained {
if len(g.paths) < maxPHPTaintGapPaths {
stable := true
if g.resolveAliases != nil {
aliases, stable = g.resolveAliases(path)
} else {
aliases = []string{coverageLexicalPath(path)}
}
if !stable {
g.recordUnknownRange(fmt.Sprintf("%s changed while its path identity was captured", path))
g.byStatus[status]++
if _, ok := g.example[status]; !ok {
g.example[status] = sanitizeJSTaintDisplay(path, phpTaintExampleMaxBytes)
}
return
}
g.paths[path] = struct{}{}
g.aliasesByPath[path] = aliases
for _, alias := range aliases {
g.pathAliases[alias] = struct{}{}
}
} else {
// Stop retaining paths, but never stop counting: the count is what
// tells an operator how much of the host went unexamined.
g.pathsTruncated = true
}
}
if len(aliases) > 0 && g.recordCoverage != nil {
g.recordCoverage(aliases)
}
g.byStatus[status]++
if _, ok := g.example[status]; !ok {
g.example[status] = sanitizeJSTaintDisplay(path, phpTaintExampleMaxBytes)
}
}
// pathsIncomplete reports that this run could not enumerate every gapped path,
// either because retention hit its bound or because the walk lost an unknown
// range. Its carry-forward set is therefore not authoritative and the purge
// must be suppressed for the whole owner.
func (g *phpTaintGapCollector) pathsIncomplete() bool {
return g.pathsTruncated || g.unknown > 0
}
// recordUnknownRange notes coverage lost over a range this walk cannot
// enumerate. It deliberately does not add to paths: claiming specific paths
// would be false, and the unknown range already forces a partial run, which
// suppresses the purge for every prior finding.
func (g *phpTaintGapCollector) recordUnknownRange(detail string) {
g.unknown++
if g.unknownExample == "" {
g.unknownExample = sanitizeJSTaintDisplay(detail, phpTaintExampleMaxBytes)
}
}
func (g *phpTaintGapCollector) empty() bool { return len(g.byStatus) == 0 && g.unknown == 0 }
func (g *phpTaintGapCollector) hasPath(path string) bool {
if _, ok := g.paths[path]; ok {
return true
}
for _, alias := range coveragePathAliases(path) {
if _, ok := g.pathAliases[alias]; ok {
return true
}
}
return false
}
// isPHPTaintAnalyzerDefeatStatus identifies per-file hard failures.
// StatusPanic covers defects anywhere in the analyzer stack, including the
// known upstream lexer defect. StatusTimeout means the isolated worker had to
// be killed after analysis stopped making progress. Neither may be attributed
// solely to the parser, but both need to remain visible instead of being buried
// under routine coverage gaps.
func isPHPTaintAnalyzerDefeatStatus(status string) bool {
return status == phptaint.StatusPanic.String() || status == phptaint.StatusTimeout.String()
}
// findings splits routine coverage gaps from content that crashed or stalled
// the analyzer. Collector-wide range-loss facts are attached exactly once.
func (g *phpTaintGapCollector) findings() []alert.Finding {
var out []alert.Finding
routine := make(map[string]int, len(g.byStatus))
defeats := make(map[string]int, 2)
for status, n := range g.byStatus {
if isPHPTaintAnalyzerDefeatStatus(status) {
defeats[status] = n
} else {
routine[status] = n
}
}
hasRoutineFinding := len(routine) > 0 || g.unknown > 0
if hasRoutineFinding {
out = append(out, g.buildFinding(routine, true, false))
}
if len(defeats) > 0 {
out = append(out, g.buildFinding(defeats, !hasRoutineFinding, true))
}
return out
}
// finding reports every gap in one alert. Retained for callers that do not
// need the split.
func (g *phpTaintGapCollector) finding() alert.Finding {
return g.buildFinding(g.byStatus, true, false)
}
func (g *phpTaintGapCollector) buildFinding(byStatus map[string]int, includeRangeLoss, analyzerDefeat bool) alert.Finding {
total := 0
statuses := make([]string, 0, len(byStatus))
for status, n := range byStatus {
total += n
statuses = append(statuses, status)
}
sort.Strings(statuses)
parts := make([]string, 0, len(statuses))
for _, status := range statuses {
parts = append(parts, fmt.Sprintf("%s=%d (example: %s)", status, byStatus[status], g.example[status]))
}
if includeRangeLoss && g.unknown > 0 {
parts = append(parts, fmt.Sprintf("unreadable-range=%d (example: %s)", g.unknown, g.unknownExample))
}
if includeRangeLoss && g.pathsTruncated {
parts = append(parts, fmt.Sprintf("exact paths retained for only the first %d", maxPHPTaintGapPaths))
}
message := fmt.Sprintf("PHP taint deep scan could not analyze %d file(s)", total)
// Routine coverage loss is one host condition. Analyzer defeats are
// input-specific so dismissing one cannot hide different failing files.
dedupKey := "coverage_gap"
if analyzerDefeat {
message = fmt.Sprintf("PHP taint deep scan was defeated by %d file(s) that crashed or stalled the analyzer", total)
dedupKey = g.analyzerDefeatDedupKey()
} else if total == 0 {
message = fmt.Sprintf("PHP taint deep scan could not cover %d location(s)", g.unknown)
}
return alert.Finding{
Severity: alert.Warning,
Check: "php_taint_scan_incomplete",
Message: message,
Details: strings.Join(parts, "; "),
DedupKey: dedupKey,
}
}
func (g *phpTaintGapCollector) analyzerDefeatDedupKey() string {
// Hash every full input identity, not just sanitized display examples.
// Sorting keeps ordinary traversal-order changes out of the alert key.
sort.Slice(g.defeatInputs, func(i, j int) bool {
return bytes.Compare(g.defeatInputs[i][:], g.defeatInputs[j][:]) < 0
})
digest := sha256.New()
for _, identity := range g.defeatInputs {
_, _ = digest.Write(identity[:])
}
if g.defeatOverflow != nil {
_, _ = digest.Write([]byte("overflow:"))
_, _ = digest.Write(g.defeatOverflow.Sum(nil))
}
return fmt.Sprintf("analyzer_defeat:%x", digest.Sum(nil))
}
// carryForwardPHPTaintFindings keeps at most one prior state finding for each
// path the current full cycle could not analyze, so a file that goes from
// analyzed to unexaminable does not silently lose its existing finding.
func carryForwardPHPTaintFindings(prior []alert.Finding, gaps *phpTaintGapCollector) []alert.Finding {
byPath := make(map[string]alert.Finding)
for _, finding := range prior {
if finding.Check != "php_remote_taint" || !gaps.hasPath(finding.FilePath) {
continue
}
current, exists := byPath[finding.FilePath]
if !exists || finding.Timestamp.After(current.Timestamp) ||
(finding.Timestamp.Equal(current.Timestamp) && finding.Key() < current.Key()) {
byPath[finding.FilePath] = finding
}
}
paths := make([]string, 0, len(byPath))
for path := range byPath {
paths = append(paths, path)
}
sort.Strings(paths)
carried := make([]alert.Finding, 0, len(paths))
for _, path := range paths {
finding := byPath[path]
finding.ScanCarryForward = true
carried = append(carried, finding)
}
return carried
}
// analyzePHPTaintSnapshot runs the PHP consumer on one complete in-memory
// snapshot and converts the result into at most one finding. Only StatusAnalyzed
// and StatusNotCandidate mean the file was examined; every other status is
// recorded as a known-path coverage gap.
func analyzePHPTaintSnapshot(ctx context.Context, path, contentSHA256 string, data []byte, gaps *phpTaintGapCollector) []alert.Finding {
fileCtx, cancel := context.WithTimeout(ctx, phpTaintDeepPerFileTimeout)
report := runPHPTaintAnalysis(fileCtx, data)
cancel()
switch report.Status {
case phptaint.StatusAnalyzed:
if len(report.Results) == 0 {
return nil
}
return []alert.Finding{phpTaintDeepFinding(path, contentSHA256, report)}
case phptaint.StatusNotCandidate:
return nil
default:
gaps.recordSnapshot(path, report.Status.String(), contentSHA256)
return nil
}
}
// phpTaintDeepFinding renders the single finding for one analyzed file. Every
// display field is sanitized and bounded; FilePath keeps the exact live path
// for remediation while only its display copy is sanitized.
func phpTaintDeepFinding(path, contentSHA256 string, report phptaint.Report) alert.Finding {
flows := make([]string, 0, len(report.Results))
for _, res := range report.Results {
flows = append(flows, fmt.Sprintf("%s -> %s (%s, %s)", res.Source, res.Sink, phpTaintConfidence(res.Confidence), res.Basis))
}
details := "Remotely fetched content reaches a code-execution construct. Evidence: " + strings.Join(flows, "; ")
if extra := report.TotalResults - len(report.Results); extra > 0 {
details += fmt.Sprintf("; %d additional flow(s) beyond returned evidence", extra)
}
if report.EvidenceTruncated {
details += " [evidence truncated]"
}
if len(report.PrecisionLoss) > 0 {
details += "; reduced precision: " + strings.Join(report.PrecisionLoss, ", ")
}
severity := phpTaintSeverity(report.Results)
return alert.Finding{
Severity: severity,
Check: "php_remote_taint",
Message: "PHP remote-source code execution data flow: " + sanitizeJSTaintDisplay(path, phpTaintMessageMaxBytes),
Details: sanitizeJSTaintDisplay(details, phpTaintDetailsMaxBytes),
DedupKey: phpTaintDedupKey(path, severity, report.Results),
FilePath: path,
ContentSHA256: contentSHA256,
DetectLogic: ContentDetectionVersion(),
}
}
// phpTaintDedupKey pins a finding's identity to the file, its severity and
// the distinct source and sink endpoints of its flows. Details carry the
// evidence wording, basis and context, which change between releases; if
// they fed the key, each such change would re-key every stored finding, drop
// its dismissal and alert again. The content hash is left out too: a library
// file still flagged after an update is the same finding. Any change to the
// reported flows, file or severity still makes a new one.
func phpTaintDedupKey(path string, severity alert.Severity, results []phptaint.Result) string {
type endpoint struct{ source, sink string }
seen := make(map[endpoint]bool, len(results))
pairs := make([]endpoint, 0, len(results))
for _, res := range results {
e := endpoint{res.Source, res.Sink}
if !seen[e] {
seen[e] = true
pairs = append(pairs, e)
}
}
sort.Slice(pairs, func(i, j int) bool {
if pairs[i].source != pairs[j].source {
return pairs[i].source < pairs[j].source
}
return pairs[i].sink < pairs[j].sink
})
identity := make([]byte, 0, 128)
appendField := func(value string) {
identity = binary.BigEndian.AppendUint64(identity, uint64(len(value)))
identity = append(identity, value...)
}
appendField(path)
appendField(severity.String())
for _, e := range pairs {
appendField(e.source)
appendField(e.sink)
}
digest := sha256.Sum256(identity)
return fmt.Sprintf("php-taint:%x", digest[:12])
}
// phpTaintSeverity grades a finding by the strongest flow it contains.
//
// Confidence is what separates a fetch this analyzer PROVED was remote from
// one it merely could not rule out, and on real hosts that distinction is the
// whole signal. Measured over 1,104,790 files on a production cPanel server:
// 30,276 analyzed, 33 findings, every one of them ConfidenceLow and every one
// a third-party library that legitimately reads a file and evaluates it --
// template compilers, cache layers, SDK bootstrap code. Nothing graded High or
// Certain. Reporting all of them at one severity would bury a genuine remote
// code-execution flow among library noise on its first scan.
//
// Low is downgraded, not dropped: a real cross-function flow whose URL sits at
// the caller grades Low too, because the acquiring call cannot see it. The
// finding stays for review; it just does not page anyone.
func phpTaintSeverity(results []phptaint.Result) alert.Severity {
strongest := phptaint.ConfidenceLow
for _, res := range results {
if res.Confidence > strongest {
strongest = res.Confidence
}
}
switch strongest {
case phptaint.ConfidenceCertain:
return alert.Critical
case phptaint.ConfidenceHigh:
return alert.High
default:
return alert.Warning
}
}
func phpTaintConfidence(c phptaint.Confidence) string {
switch c {
case phptaint.ConfidenceCertain:
return "certain"
case phptaint.ConfidenceHigh:
return "high"
default:
return "low"
}
}
// phpTaintOversizePeekBytes bounds the prefix read to decide whether an
// oversize file could be PHP. An open tag in a PHP file is at the top; a
// larger peek would only buy false positives from binary content that happens
// to contain the byte sequence.
const phpTaintOversizePeekBytes = 64 << 10
type phpRegularFilePrefixReader interface {
ReadRegularFilePrefix(string, os.FileInfo, int64) ([]byte, error)
}
package checks
import (
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var pluginInventoryBatches = newScanBatchMonitor()
// PluginInventoryQueueStatus includes sites waiting for a worker and inventories
// still executing or committing their result. Concurrent refreshes share no cap.
func PluginInventoryQueueStatus(now time.Time) queuehealth.Status {
return pluginInventoryBatches.snapshot(now)
}
package checks
import (
"context"
"sync"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/store"
)
type wpInventoryCycle struct {
mu sync.Mutex
calls map[*store.DB]*wpInventoryCycleCall
}
type wpInventoryCycleCall struct {
done chan struct{}
fresh bool
}
// ensurePluginCacheFresh reuses one refresh result for all consumers in a
// scan, even if the first consumer has finished before the next one starts.
func ensurePluginCacheFresh(ctx context.Context, cfg *config.Config, db *store.DB) bool {
cycle := wpInstallCacheFrom(ctx)
if cycle == nil {
return ensurePluginCacheFreshShared(ctx, cfg, db)
}
inventory := &cycle.inventory
inventory.mu.Lock()
if call := inventory.calls[db]; call != nil {
inventory.mu.Unlock()
select {
case <-call.done:
return call.fresh
case <-ctx.Done():
return false
}
}
call := &wpInventoryCycleCall{done: make(chan struct{})}
if inventory.calls == nil {
inventory.calls = make(map[*store.DB]*wpInventoryCycleCall)
}
inventory.calls[db] = call
inventory.mu.Unlock()
defer close(call.done)
call.fresh = ensurePluginCacheFreshShared(ctx, cfg, db)
return call.fresh
}
package checks
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
neturl "net/url"
"os"
"os/exec"
"path/filepath"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// wpCLIFlags are the extra flags and env CSM adds to every wp-cli invocation.
//
// WP_CLI_PHP_ARGS disables PHP's display_errors/error_reporting for the
// bootstrap so stray Notices/Warnings/Deprecated messages from the site
// don't get emitted (they'd only be a problem on stderr, which we already
// discard, but suppressing them also avoids exit-255 on strict hosts that
// promote warnings to errors).
//
// --skip-plugins and --skip-themes make wp-cli enumerate plugins from the
// filesystem without loading them. That removes the biggest source of log
// noise: one broken plugin (e.g. a PHP Parse error in litespeed-cache on a
// site nobody updated for years) would otherwise crash the whole `wp plugin
// list` call with exit 255, or spew backtraces from plugins that call
// wp_redirect() during admin bootstrap. Skipping loads gives us the list
// plus update_version unchanged.
const wpCLIFlags = `WP_CLI_PHP_ARGS='-d display_errors=0 -d error_reporting=0' wp --skip-plugins --skip-themes `
var wpOrgHTTPClient = &http.Client{Timeout: 10 * time.Second}
type wpOrgResponse struct {
Slug string `json:"slug"`
Version string `json:"version"`
Tested string `json:"tested"`
Error string `json:"error"`
}
// parseWPOrgPluginResponse parses a JSON body from the WordPress.org plugin
// information API into a store.PluginInfo. It returns an error if the response
// contains an error field (e.g. "Plugin not found.") or if the JSON is invalid.
func parseWPOrgPluginResponse(body []byte) (store.PluginInfo, error) {
var resp wpOrgResponse
if err := json.Unmarshal(body, &resp); err != nil {
return store.PluginInfo{}, fmt.Errorf("wporg: invalid JSON: %w", err)
}
if resp.Error != "" {
return store.PluginInfo{}, fmt.Errorf("wporg: %s", resp.Error)
}
return store.PluginInfo{
LatestVersion: resp.Version,
TestedUpTo: resp.Tested,
LastChecked: time.Now().Unix(),
}, nil
}
// fetchWPOrgPluginInfo queries the WordPress.org plugin information API for the
// given slug and returns the parsed PluginInfo.
func fetchWPOrgPluginInfo(ctx context.Context, slug string) (store.PluginInfo, error) {
url := "https://api.wordpress.org/plugins/info/1.2/?action=plugin_information" +
"&request[slug]=" + neturl.QueryEscape(slug) +
"&request[fields][version]=1&request[fields][tested]=1"
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return store.PluginInfo{}, fmt.Errorf("wporg: building request: %w", err)
}
resp, err := wpOrgHTTPClient.Do(req)
if err != nil {
return store.PluginInfo{}, fmt.Errorf("wporg: HTTP request failed: %w", err)
}
defer resp.Body.Close()
body, err := io.ReadAll(resp.Body)
if err != nil {
return store.PluginInfo{}, fmt.Errorf("wporg: reading response body: %w", err)
}
if resp.StatusCode != http.StatusOK {
return store.PluginInfo{}, fmt.Errorf("wporg: unexpected status %d", resp.StatusCode)
}
return parseWPOrgPluginResponse(body)
}
// parseVersion splits a dotted version string like "6.4.2" into []int{6, 4, 2}.
// Non-numeric segments are treated as 0.
func parseVersion(v string) []int {
if v == "" {
return nil
}
parts := strings.Split(v, ".")
out := make([]int, len(parts))
for i, p := range parts {
n, err := strconv.Atoi(p)
if err != nil {
n = 0
}
out[i] = n
}
return out
}
// compareVersions returns whether there is a major version gap and
// how many minor versions behind the installed version is.
func compareVersions(installed, available string) (majorGap bool, minorBehind int) {
iv := parseVersion(installed)
av := parseVersion(available)
if len(iv) < 2 || len(av) < 2 {
return false, 0
}
if av[0] > iv[0] {
return true, 0
}
if av[0] == iv[0] && av[1] > iv[1] {
return false, av[1] - iv[1]
}
return false, 0
}
// pluginAlertSeverity returns "critical", "high", "warning", or "" for a version gap.
func pluginAlertSeverity(installed, available string) string {
majorGap, minorBehind := compareVersions(installed, available)
if majorGap {
return "critical"
}
if minorBehind >= 3 {
return "high"
}
// Check if there is any difference at all.
iv := parseVersion(installed)
av := parseVersion(available)
if len(iv) < 2 || len(av) < 2 {
return ""
}
// Compare all parsed components to detect if available is actually newer.
// If installed >= available at every component, the site is up to date
// (or ahead, e.g. custom/premium builds). Only warn when behind.
maxLen := len(iv)
if len(av) > maxLen {
maxLen = len(av)
}
for i := 0; i < maxLen; i++ {
var a, b int
if i < len(iv) {
a = iv[i]
}
if i < len(av) {
b = av[i]
}
if b > a {
return "warning" // available is newer at this component
}
if a > b {
return "" // installed is ahead - not outdated
}
}
return "" // identical
}
const pluginCheckWorkers = 5
const defaultPluginCheckIntervalMin = 1440
type pluginCacheRefreshCall struct {
done chan struct{}
}
var (
pluginCacheRefreshMu sync.Mutex
pluginCacheRefreshes = make(map[*store.DB]*pluginCacheRefreshCall)
)
func pluginCacheFresh(db *store.DB, cfg *config.Config) bool {
lastRefresh := db.GetPluginRefreshTime()
if lastRefresh.IsZero() {
return false
}
intervalMin := defaultPluginCheckIntervalMin
if cfg != nil && cfg.Thresholds.PluginCheckIntervalMin > 0 {
intervalMin = cfg.Thresholds.PluginCheckIntervalMin
}
elapsed := time.Since(lastRefresh)
return elapsed >= 0 && elapsed <= time.Duration(intervalMin)*time.Minute
}
// ensurePluginCacheFresh coalesces refreshes shared by the outdated and
// known-vulnerable detectors. The checks run concurrently, so list order in
// the runner cannot establish the inventory dependency. Waiters reuse the
// leader's failed result as well, avoiding a duplicate walk in the same scan.
func ensurePluginCacheFreshShared(ctx context.Context, cfg *config.Config, db *store.DB) bool {
if pluginCacheFresh(db, cfg) {
return true
}
pluginCacheRefreshMu.Lock()
if call := pluginCacheRefreshes[db]; call != nil {
pluginCacheRefreshMu.Unlock()
select {
case <-call.done:
return pluginCacheFresh(db, cfg)
case <-ctx.Done():
return false
}
}
if pluginCacheFresh(db, cfg) {
pluginCacheRefreshMu.Unlock()
return true
}
call := &pluginCacheRefreshCall{done: make(chan struct{})}
pluginCacheRefreshes[db] = call
pluginCacheRefreshMu.Unlock()
func() {
defer func() {
pluginCacheRefreshMu.Lock()
delete(pluginCacheRefreshes, db)
close(call.done)
pluginCacheRefreshMu.Unlock()
}()
refreshPluginCache(ctx, db)
}()
return pluginCacheFresh(db, cfg)
}
// CheckOutdatedPlugins scans all WordPress installations for plugins with
// available updates and emits findings based on severity of the version gap.
// Results are cached in bbolt with a configurable refresh interval (default 24h).
func CheckOutdatedPlugins(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
db := store.Global()
if db == nil {
return nil
}
ensurePluginCacheFresh(ctx, cfg, db)
return evaluatePluginCache(db)
}
// findAllWPInstalls lists the WordPress installs to inventory plugins for.
// Discovery is shared (wpinstalls.go), which adds the panel's document-root
// map: a root outside the home layout used to have no plugin coverage, so its
// vulnerable plugins were never found and never virtually patched.
func findAllWPInstalls(ctx context.Context) []string {
seen := make(map[string]bool)
var results []string
installs := wpInstalls(ctx, "vulnerable_plugins")
if checkMarkedIncomplete(ctx, "vulnerable_plugins") {
// The same inventory backs both checks, regardless of which one won
// the concurrent refresh. Neither may purge after a partial walk.
markCheckIncomplete(ctx, "outdated_plugins")
}
for _, install := range installs {
if seen[install.ConfigPath] {
continue
}
seen[install.ConfigPath] = true
results = append(results, install.ConfigPath)
}
return results
}
func canonicalWPInstallPath(wpConfig string) (string, error) {
clean := filepath.Clean(wpConfig)
dir := filepath.Dir(clean)
home := wpInstallAccountRoot(clean)
if home == "" {
return "", fmt.Errorf("WordPress document root is outside an account home: %s", dir)
}
rel, err := filepath.Rel(home, dir)
if err != nil || rel == "." || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return "", fmt.Errorf("WordPress document root is outside an account home: %s", dir)
}
parts := strings.Split(rel, string(filepath.Separator))
resolved := home
const maxSymlinks = 40
symlinks := 0
for len(parts) > 0 {
part := parts[0]
parts = parts[1:]
if part == "" || part == "." || part == ".." {
return "", fmt.Errorf("invalid WordPress document-root component in %s", dir)
}
candidate := filepath.Join(resolved, part)
info, err := osFS.Lstat(candidate)
if err != nil {
return "", fmt.Errorf("inspect WordPress document root %s: %w", candidate, err)
}
if info.Mode()&os.ModeSymlink == 0 {
if !info.IsDir() {
return "", fmt.Errorf("WordPress document root is not a directory: %s", candidate)
}
resolved = candidate
continue
}
symlinks++
if symlinks > maxSymlinks {
return "", fmt.Errorf("too many symlinks in WordPress document root: %s", dir)
}
target, err := osFS.Readlink(candidate)
if err != nil {
return "", fmt.Errorf("resolve WordPress document root %s: %w", candidate, err)
}
if target == "" {
return "", fmt.Errorf("resolve WordPress document root %s: empty symlink target", candidate)
}
if !filepath.IsAbs(target) {
target = filepath.Join(resolved, target)
}
target = filepath.Clean(target)
if !isPathWithinOrEqual(target, home) {
return "", fmt.Errorf("WordPress document root %s resolves outside its account", candidate)
}
targetRel, err := filepath.Rel(home, target)
if err != nil {
return "", fmt.Errorf("resolve WordPress document root %s: %w", candidate, err)
}
var targetParts []string
if targetRel != "." {
targetParts = strings.Split(targetRel, string(filepath.Separator))
}
parts = append(targetParts, parts...)
resolved = home
}
return filepath.Join(resolved, filepath.Base(clean)), nil
}
func wpInstallAccountRoot(path string) string {
if home := homeAccountRoot(path); home != "" {
return home
}
clean := filepath.Clean(path)
if !filepath.IsAbs(clean) {
return ""
}
parts := strings.Split(clean, string(filepath.Separator))
for i, part := range parts {
if !isCPanelHomeBase(part) || i+2 >= len(parts) || !validAccountName.MatchString(parts[i+1]) {
continue
}
return filepath.Join(string(filepath.Separator), filepath.Join(parts[1:i+2]...))
}
return ""
}
// wpCLIPluginEntry mirrors the JSON output of `wp plugin list --format=json`.
type wpCLIPluginEntry struct {
Name string `json:"name"`
Status string `json:"status"`
Version string `json:"version"`
UpdateVersion string `json:"update_version"`
}
// refreshPluginCache discovers all WP installs, runs wp-cli to inventory
// plugins for each site, enriches free plugins via the WordPress.org API,
// and stores everything in bbolt.
func refreshPluginCache(ctx context.Context, db *store.DB) {
if incompleteCollectorFrom(ctx) == nil {
ctx, _ = withIncompleteCheckCollector(ctx)
}
wpConfigs := findAllWPInstalls(ctx)
discoveryIncomplete := checkMarkedIncomplete(ctx, "vulnerable_plugins")
coverage := newWPVerificationBatch(ctx, db, "plugins", "wp_plugin_inventory", wpConfigs)
defer func() {
if err := coverage.finish(ctx, !discoveryIncomplete); err != nil {
fmt.Fprintln(os.Stderr, "plugincheck: could not save verification history")
}
}()
if ctx.Err() != nil {
return
}
if len(wpConfigs) == 0 {
if discoveryIncomplete {
return
}
for path := range db.AllSitePlugins() {
if err := db.DeleteSitePlugins(path); err != nil {
fmt.Fprintf(os.Stderr, "plugincheck: prune failed for %s: %v\n", path, err)
return
}
}
return
}
var mu sync.Mutex
var wg sync.WaitGroup
successCount := 0
var timeoutCount, execFailCount, parseFailCount, cleanupFailCount int
slugsSeen := make(map[string]bool)
discoveredPaths := make(map[string]bool)
batch := pluginInventoryBatches.begin(len(wpConfigs), pluginCheckWorkers)
defer batch.abandon(ctx)
jobs := make(chan int, len(wpConfigs))
for i := 0; i < pluginCheckWorkers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for index := range jobs {
wpConfig, work := wpConfigs[index], batch.tasks[index]
work.admit()
stop := false
// Site URL discovery and plugin enumeration each have a command timeout.
work.run(ctx, 2*cmdTimeout, func() {
if err := ctx.Err(); err != nil {
work.withdraw(err)
stop = true
return
}
wpPath := filepath.Dir(wpConfig)
if err := ctx.Err(); err != nil {
work.withdraw(err)
stop = true
return
}
mu.Lock()
discoveredPaths[wpPath] = true
mu.Unlock()
sitePlugins, err := inventoryWPSite(ctx, wpConfig)
work.progress()
if err != nil {
parentErr := ctx.Err()
if (!errors.Is(err, context.Canceled) || parentErr == nil) && (!commandRefused(err) || errors.Is(err, errWPInventoryNoOutput)) {
work.fail()
}
if parentErr != nil {
work.withdraw(parentErr)
stop = true
return
}
coverage.record(wpPath, wpVerificationFailure(err, nil))
mu.Lock()
switch {
case errors.Is(err, context.DeadlineExceeded):
timeoutCount++
case errors.Is(err, errWPInventoryParse):
parseFailCount++
default:
execFailCount++
}
mu.Unlock()
if err := db.DeleteSitePlugins(wpPath); err != nil {
fmt.Fprintf(os.Stderr, "plugincheck: stale cache cleanup failed for %s: %v\n", wpPath, err)
mu.Lock()
cleanupFailCount++
mu.Unlock()
}
return
}
mu.Lock()
for _, p := range sitePlugins.Plugins {
slugsSeen[p.Slug] = true
}
mu.Unlock()
coverage.record(wpPath, store.WPVerificationResult{State: "verified"})
if err := db.SetSitePlugins(wpPath, sitePlugins); err != nil {
coverage.record(wpPath, store.WPVerificationResult{State: "unverified", Reason: "CSM could not save the plugin inventory"})
work.fail()
fmt.Fprintf(os.Stderr, "plugincheck: store failed for %s: %v\n", wpPath, err)
if cleanupErr := db.DeleteSitePlugins(wpPath); cleanupErr != nil {
fmt.Fprintf(os.Stderr, "plugincheck: stale cache cleanup failed for %s: %v\n", wpPath, cleanupErr)
}
mu.Lock()
cleanupFailCount++
mu.Unlock()
return
}
mu.Lock()
successCount++
mu.Unlock()
})
if stop {
return
}
}
}()
}
for index := range wpConfigs {
if ctx.Err() != nil {
break
}
jobs <- index
}
close(jobs)
wg.Wait()
batch.abandon(ctx)
if ctx.Err() != nil {
return
}
// Enrich free plugins via WordPress.org API (one lookup per unique slug).
mu.Lock()
slugList := make([]string, 0, len(slugsSeen))
for slug := range slugsSeen {
slugList = append(slugList, slug)
}
mu.Unlock()
for _, slug := range slugList {
if ctx.Err() != nil {
return
}
// Skip if we have a recent cached entry (< 24h).
if cached, ok := db.GetPluginInfo(slug); ok {
if time.Since(time.Unix(cached.LastChecked, 0)) < 24*time.Hour {
continue
}
}
info, err := fetchWPOrgPluginInfo(ctx, slug)
if err != nil {
// Not found on .org = premium/custom plugin, skip silently.
continue
}
_ = db.SetPluginInfo(slug, info)
}
// A partial walk cannot prove that an absent path was removed.
if !discoveryIncomplete {
allCached := db.AllSitePlugins()
for path := range allCached {
if !discoveredPaths[path] {
if err := db.DeleteSitePlugins(path); err != nil {
fmt.Fprintf(os.Stderr, "plugincheck: prune failed for %s: %v\n", path, err)
cleanupFailCount++
}
}
}
}
// Only mark refresh as complete if the majority of sites refreshed
// successfully. A partial failure (e.g. one wp-cli timeout on a 100-site
// server) should not freeze ALL stale data for 24 hours. But if most
// sites failed (e.g. transient PHP issue), don't mark as fresh - allow
// retry next cycle.
mu.Lock()
sc := successCount
to, exf, pf, cf := timeoutCount, execFailCount, parseFailCount, cleanupFailCount
mu.Unlock()
failCount := len(wpConfigs) - sc
ts := time.Now().Format("2006-01-02 15:04:05")
if sc == 0 {
fmt.Fprintf(os.Stderr, "[%s] plugincheck: refresh failed, 0/%d sites succeeded%s, not updating timestamp\n",
ts, len(wpConfigs), failureBreakdown(to, exf, pf))
return
}
if failCount > sc {
fmt.Fprintf(os.Stderr, "[%s] plugincheck: refresh partial, %d/%d sites failed%s, not updating timestamp\n",
ts, failCount, len(wpConfigs), failureBreakdown(to, exf, pf))
return
}
if cf > 0 {
fmt.Fprintf(os.Stderr, "[%s] plugincheck: refresh incomplete, %d stale cache cleanup(s) failed, not updating timestamp\n",
ts, cf)
return
}
if discoveryIncomplete {
fmt.Fprintf(os.Stderr, "[%s] plugincheck: discovery incomplete, not updating timestamp\n", ts)
return
}
if failCount > 0 {
fmt.Fprintf(os.Stderr, "[%s] plugincheck: refreshed %d/%d sites%s\n",
ts, sc, len(wpConfigs), failureBreakdown(to, exf, pf))
}
if err := db.SetPluginRefreshTime(time.Now()); err != nil {
fmt.Fprintf(os.Stderr, "plugincheck: refresh timestamp failed: %v\n", err)
}
}
// failureBreakdown formats " (timeout=N exec_fail=N json_fail=N)" when any
// category is non-zero, or "" otherwise. Keeps the refresh log to one line
// instead of one line per broken site.
func failureBreakdown(timeout, execFail, parseFail int) string {
if timeout == 0 && execFail == 0 && parseFail == 0 {
return ""
}
return fmt.Sprintf(" (timeout=%d exec_fail=%d json_fail=%d)", timeout, execFail, parseFail)
}
// evaluatePluginCache reads the cached plugin inventory and emits one
// aggregated finding per site listing every outdated active plugin. The
// per-site rollup keeps the alert channel under control during a deep
// scan tier on hosts with many sites: the previous one-finding-per-
// outdated-plugin shape produced ~1000 findings on a 200-account host
// and saturated the 500-deep alert channel buffer, dropping real
// signal under "alert channel full, dropping deep finding:
// outdated_plugins".
//
// Aggregation rules:
// - Severity = max of constituents (critical > high > warning).
// - Message = "<count> outdated plugins on <domain> (<account>):
// <severity-label>" - searchable and self-describing.
// - Details lists each plugin slug, installed version, available
// version, and per-plugin severity, one per line, so an operator
// triaging the alert sees the same per-plugin breakdown as before.
func evaluatePluginCache(db *store.DB) []alert.Finding {
var findings []alert.Finding
allSites := db.AllSitePlugins()
for wpPath, site := range allSites {
var (
detailLines []string
worstSeverity alert.Severity
worstSevLabel string
outdatedTotal int
)
// Track whether worstSeverity has been set at all: alert.Severity's
// zero value is Warning, so a strict "newer rank > current rank"
// comparison would never overwrite the initial state on a site
// whose constituents are all Warning, leaving worstSevLabel empty.
worstSet := false
for _, p := range site.Plugins {
severity, sevLabel, available, ok := pluginOutdatedSeverity(p, db)
if !ok {
continue
}
outdatedTotal++
detailLines = append(detailLines, fmt.Sprintf("- %s (%s): %s -> %s [%s]",
p.Slug, p.Name, p.InstalledVersion, available, sevLabel))
if !worstSet || severityRank(severity) > severityRank(worstSeverity) {
worstSeverity = severity
worstSevLabel = sevLabel
worstSet = true
}
}
if outdatedTotal == 0 {
continue
}
findings = append(findings, alert.Finding{
Severity: worstSeverity,
Check: "outdated_plugins",
Message: fmt.Sprintf("%d outdated plugin%s on %s (%s): worst severity %s",
outdatedTotal, plural(outdatedTotal), site.Domain, site.Account, worstSevLabel),
Details: fmt.Sprintf("Path: %s\nOutdated plugins (%d):\n%s",
wpPath, outdatedTotal, strings.Join(detailLines, "\n")),
})
}
return findings
}
// errWPInventoryParse marks a wp-cli plugin-list response that ran but produced
// unparseable JSON, so callers can tell a parse failure from an exec failure.
var errWPInventoryParse = errors.New("wp-cli plugin list: invalid JSON")
// A refusal without command output cannot account for the inventory work.
var errWPInventoryNoOutput = errors.New("wp-cli plugin list: no output")
// inventoryWPSite runs wp-cli for a single site (as the site owner) and returns
// its current plugin inventory. Shared by the periodic cache refresh and the
// per-finding re-check so both see identical results. Read-only: it inventories,
// it does not change anything.
func inventoryWPSite(ctx context.Context, wpConfig string) (store.SitePlugins, error) {
return inventoryWPSiteWithDomain(ctx, wpConfig, true)
}
func inventoryWPSiteForVerify(ctx context.Context, wpConfig string) (store.SitePlugins, error) {
return inventoryWPSiteWithDomain(ctx, wpConfig, false)
}
func inventoryWPSiteWithDomain(ctx context.Context, wpConfig string, includeDomain bool) (store.SitePlugins, error) {
wpPath := filepath.Dir(wpConfig)
user := wpConfigUser(wpPath)
if !validWPCLIUser(user) {
return store.SitePlugins{}, fmt.Errorf("invalid WordPress site owner: %q", user)
}
domain := user
if includeDomain {
domain = extractWPDomain(ctx, wpPath, user)
}
// Run wp plugin list as the site owner on stdout-only so PHP
// notices/warnings on stderr can't corrupt the JSON we parse. Use --path
// instead of a shell cd to avoid shell injection via crafted directory
// names on shared hosting.
out, err := runWPCLIStdout(ctx, user,
wpCLIFlags+"plugin list --fields=name,status,version,update_version --format=json --path="+shellQuote(wpPath),
)
if out == nil {
// Output retains a failed command's stderr on ExitError. A refusal
// there still answers the check without contaminating the JSON input.
var exit *exec.ExitError
if !errors.As(err, &exit) || len(exit.Stderr) == 0 {
return store.SitePlugins{}, errors.Join(errWPInventoryNoOutput, err)
}
}
if err != nil {
return store.SitePlugins{}, &wpInventoryError{err: err, result: wpVerificationFailure(err, out)}
}
var entries []wpCLIPluginEntry
if err := json.Unmarshal(out, &entries); err != nil {
return store.SitePlugins{}, fmt.Errorf("%w: %v", errWPInventoryParse, err)
}
site := store.SitePlugins{Account: user, Domain: domain}
for _, e := range entries {
site.Plugins = append(site.Plugins, store.SitePluginEntry{
Slug: e.Name,
Name: e.Name,
Status: e.Status,
InstalledVersion: e.Version,
UpdateVersion: e.UpdateVersion,
})
}
return site, nil
}
func validWPCLIUser(user string) bool {
if user == "" || user == "unknown" || strings.HasPrefix(user, "-") {
return false
}
for _, r := range user {
if r >= 'a' && r <= 'z' {
continue
}
if r >= 'A' && r <= 'Z' {
continue
}
if r >= '0' && r <= '9' {
continue
}
switch r {
case '_', '-', '.', '@', '$':
continue
default:
return false
}
}
return true
}
func runWPCLIStdout(ctx context.Context, user, command string) ([]byte, error) {
if !validWPCLIUser(user) {
return nil, fmt.Errorf("invalid WordPress site owner: %q", user)
}
// runuser, not su: under the hardened systemd unit /var/log is read-only, and
// su's PAM stack runs pam_lastlog on every call (floods the journal with
// "/var/log/lastlog: Read-only file system" and would record CSM's internal
// scans as user logins). runuser's PAM session has no pam_lastlog. Same
// login-shell semantics otherwise.
return cmdExec.RunContextStdout(ctx, "runuser", "-l", "-s", "/bin/bash", "-c", command, "--", user)
}
// pluginOutdatedSeverity classifies one plugin entry. It returns ok=false for
// inactive plugins, plugins with no known available version (custom/premium),
// and plugins already current. The "available" version prefers wp-cli's
// update_version and falls back to the cached WordPress.org latest version.
func pluginOutdatedSeverity(p store.SitePluginEntry, db *store.DB) (severity alert.Severity, sevLabel, available string, ok bool) {
if p.Status != "active" && p.Status != "active-network" {
return 0, "", "", false
}
available = p.UpdateVersion
if available == "" && db != nil {
if info, found := db.GetPluginInfo(p.Slug); found {
available = info.LatestVersion
}
}
if available == "" {
return 0, "", "", false
}
sevLabel = pluginAlertSeverity(p.InstalledVersion, available)
if sevLabel == "" {
return 0, "", "", false
}
switch sevLabel {
case "critical":
severity = alert.Critical
case "high":
severity = alert.High
default:
severity = alert.Warning
}
return severity, sevLabel, available, true
}
// countOutdatedActivePlugins reports how many active plugins on a site still
// have an available update, using the same classification as the alert path.
func countOutdatedActivePlugins(site store.SitePlugins, db *store.DB) int {
n := 0
for _, p := range site.Plugins {
if _, _, _, ok := pluginOutdatedSeverity(p, db); ok {
n++
}
}
return n
}
// severityRank orders severities so an aggregate can pick the worst.
// Higher returned value means more severe.
func severityRank(s alert.Severity) int {
switch s {
case alert.Critical:
return 3
case alert.High:
return 2
case alert.Warning:
return 1
default:
return 0
}
}
func plural(n int) string {
if n == 1 {
return ""
}
return "s"
}
// extractWPDomain runs `wp option get siteurl` to discover the site's domain.
// Falls back to directory name heuristics if wp-cli fails.
func extractWPDomain(ctx context.Context, wpPath, user string) string {
// Stdout-only: some sites print "WARNING: MYSQL_OPT_RECONNECT deprecated"
// or similar on stderr during wp-cli boot. Mixing that into the value
// would produce a poisoned domain like "Warning: ... https://site.com".
out, err := runWPCLIStdout(ctx, user,
wpCLIFlags+"option get siteurl --path="+shellQuote(wpPath),
)
if err == nil {
url := strings.TrimSpace(string(out))
if url != "" {
// Strip protocol prefix for display.
url = strings.TrimPrefix(url, "https://")
url = strings.TrimPrefix(url, "http://")
return url
}
}
// Fallback: use directory name after public_html (addon domain)
// or account name (main domain).
parts := strings.Split(wpPath, "/")
for i, p := range parts {
if p == "public_html" && i+1 < len(parts) {
return parts[i+1]
}
}
return user
}
// shellQuote wraps a string in single quotes for safe shell argument passing.
// Any embedded single quotes are escaped as '\” (end quote, literal quote, start quote).
func shellQuote(s string) string {
return "'" + strings.ReplaceAll(s, "'", `'\''`) + "'"
}
package checks
import (
"strings"
"github.com/pidginhost/csm/internal/store"
)
// postfixAuditParams are the parameters auditPostfix reads. postconf reports
// the effective value, so a parameter left at its built-in default still
// comes back with one.
var postfixAuditParams = []string{
"smtpd_relay_restrictions",
"smtpd_recipient_restrictions",
"smtpd_tls_security_level",
"smtpd_use_tls",
"smtpd_tls_protocols",
"smtpd_tls_mandatory_protocols",
"smtpd_sasl_auth_enable",
"smtpd_tls_auth_only",
"disable_vrfy_command",
}
// postconfSettings returns the effective value of each requested Postfix
// parameter. Only requested keys are kept, so warnings postconf prints
// alongside the values are never mistaken for settings.
func postconfSettings(params ...string) (map[string]string, bool) {
out, err := auditRunCmd("postconf", params...)
if err != nil {
return nil, false
}
wanted := make(map[string]bool, len(params))
for _, p := range params {
wanted[p] = true
}
settings := make(map[string]string, len(params))
for _, line := range strings.Split(string(out), "\n") {
key, value, found := strings.Cut(line, "=")
if !found {
continue
}
key = strings.TrimSpace(key)
if !wanted[key] {
continue
}
settings[key] = strings.TrimSpace(value)
}
return settings, len(settings) > 0
}
// auditPostfix checks the settings that decide whether a Postfix host relays
// for strangers or hands out credentials in the clear.
func auditPostfix() []store.AuditResult {
settings, ok := postconfSettings(postfixAuditParams...)
if !ok {
return []store.AuditResult{{
Category: "mail", Name: "mail_postfix_config", Title: "Postfix Configuration",
Status: "warn", Message: "Cannot query postfix configuration",
Fix: "Make 'postconf' runnable so postfix mail hardening can be audited.",
}}
}
return []store.AuditResult{
auditPostfixRelay(settings),
auditPostfixTLS(settings),
auditPostfixTLSProtocols(settings),
auditPostfixAuthOnly(settings),
auditPostfixVrfy(settings),
}
}
func auditPostfixRelay(settings map[string]string) store.AuditResult {
result := store.AuditResult{Category: "mail", Name: "mail_postfix_relay", Title: "Postfix Relay Control"}
// smtpd_relay_restrictions is the modern home for this rule; releases
// before 2.10 enforced it from smtpd_recipient_restrictions instead.
for _, key := range []string{"smtpd_relay_restrictions", "smtpd_recipient_restrictions"} {
if postfixRejectsUnauthDestination(settings[key]) {
result.Status = "pass"
result.Message = "Relaying to unauthorised destinations is rejected"
return result
}
}
result.Status = "fail"
result.Message = "No restriction list rejects relaying to unauthorised destinations"
result.Fix = "Add 'reject_unauth_destination' to smtpd_relay_restrictions in main.cf."
return result
}
func postfixRejectsUnauthDestination(value string) bool {
for _, token := range splitPostfixList(value) {
if strings.EqualFold(token, "reject_unauth_destination") ||
strings.EqualFold(token, "defer_unauth_destination") {
return true
}
}
return false
}
func auditPostfixTLS(settings map[string]string) store.AuditResult {
result := store.AuditResult{Category: "mail", Name: "mail_postfix_tls", Title: "Postfix Inbound TLS"}
level := strings.ToLower(settings["smtpd_tls_security_level"])
if level == "may" || level == "encrypt" {
result.Status = "pass"
result.Message = "Inbound TLS is offered (smtpd_tls_security_level = " + level + ")"
return result
}
// smtpd_use_tls predates smtpd_tls_security_level and still enables TLS
// on configurations that were never migrated.
if isPostfixYes(settings["smtpd_use_tls"]) {
result.Status = "pass"
result.Message = "Inbound TLS is offered via the legacy smtpd_use_tls setting"
return result
}
result.Status = "warn"
result.Message = "Inbound TLS is not offered, so mail and credentials cross the network in the clear"
result.Fix = "Set 'smtpd_tls_security_level = may' in main.cf."
return result
}
func auditPostfixTLSProtocols(settings map[string]string) store.AuditResult {
result := store.AuditResult{Category: "mail", Name: "mail_postfix_tls_protocols", Title: "Postfix TLS Protocols"}
var weak []string
for _, key := range []string{"smtpd_tls_protocols", "smtpd_tls_mandatory_protocols"} {
if !postfixProtocolsExcludeSSL(settings[key]) {
weak = append(weak, key)
}
}
if len(weak) == 0 {
result.Status = "pass"
result.Message = "SSLv2 and SSLv3 are excluded from both protocol lists"
return result
}
result.Status = "fail"
result.Message = "SSLv2 or SSLv3 is still permitted by " + strings.Join(weak, " and ")
result.Fix = "Set '>=TLSv1.2' for smtpd_tls_protocols and smtpd_tls_mandatory_protocols in main.cf."
return result
}
// postfixProtocolsExcludeSSL reports whether a protocol list keeps SSLv2 and
// SSLv3 out. Postfix accepts either an exclusion list (!SSLv2, !SSLv3) or a
// minimum-version floor (>=TLSv1.2), which rules the SSL versions out on its
// own.
func postfixProtocolsExcludeSSL(value string) bool {
excluded := make(map[string]bool)
for _, token := range splitPostfixList(value) {
if floor, ok := strings.CutPrefix(token, ">="); ok {
if strings.HasPrefix(strings.ToLower(floor), "tlsv") {
return true
}
continue
}
if name, ok := strings.CutPrefix(token, "!"); ok {
excluded[strings.ToLower(name)] = true
}
}
return excluded["sslv2"] && excluded["sslv3"]
}
func auditPostfixAuthOnly(settings map[string]string) store.AuditResult {
result := store.AuditResult{Category: "mail", Name: "mail_postfix_auth_only", Title: "Postfix Authentication Over TLS"}
if !isPostfixYes(settings["smtpd_sasl_auth_enable"]) {
result.Status = "pass"
result.Message = "SMTP authentication is not offered, so no credentials can be exposed"
return result
}
if isPostfixYes(settings["smtpd_tls_auth_only"]) {
result.Status = "pass"
result.Message = "SMTP authentication is offered only over TLS"
return result
}
result.Status = "fail"
result.Message = "SMTP authentication is offered on unencrypted connections"
result.Fix = "Set 'smtpd_tls_auth_only = yes' in main.cf."
return result
}
func auditPostfixVrfy(settings map[string]string) store.AuditResult {
result := store.AuditResult{Category: "mail", Name: "mail_postfix_vrfy", Title: "Postfix VRFY Command"}
if isPostfixYes(settings["disable_vrfy_command"]) {
result.Status = "pass"
result.Message = "The VRFY command is disabled"
return result
}
result.Status = "warn"
result.Message = "The VRFY command is enabled, letting senders confirm which mailboxes exist"
result.Fix = "Set 'disable_vrfy_command = yes' in main.cf."
return result
}
func splitPostfixList(value string) []string {
return strings.FieldsFunc(value, func(r rune) bool {
return r == ',' || r == ' ' || r == '\t' || r == '\r' || r == '\n'
})
}
func isPostfixYes(value string) bool {
switch strings.ToLower(strings.TrimSpace(value)) {
case "yes", "true", "1":
return true
}
return false
}
package checks
import (
"context"
"fmt"
"math"
"path/filepath"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/processctx"
"github.com/pidginhost/csm/internal/state"
)
func processStatusUID(data []byte) (int, bool) {
for _, line := range strings.Split(string(data), "\n") {
if strings.HasPrefix(line, "Uid:\t") {
fields := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(fields) == 0 {
return 0, false
}
uid, err := strconv.Atoi(fields[0])
return uid, err == nil && uid >= 0
}
}
return 0, false
}
func processIdentityForUID(uid int) (string, string) {
if uid < 0 || uid > math.MaxUint32 {
return "", ""
}
// #nosec G115 -- uid is range-checked against math.MaxUint32 above.
user := LookupUser(uint32(uid))
if uid >= 1000 && user != "" && !strings.HasPrefix(user, "uid:") {
return user, user
}
return user, ""
}
// suspiciousExeNames flags processes whose exe basename contains any of
// these substrings. Shared between the periodic CheckSuspiciousProcesses
// and the live BPF exec backend (which cannot see cmdline patterns and
// relies on exe-name + exe-path matching).
var suspiciousExeNames = []string{"defunct", "gsocket", "gs-netcat", "gs-sftp"}
// suspiciousExePaths flags processes whose exe path contains any of these
// directory prefixes. Shared with the live BPF exec backend.
var suspiciousExePaths = []string{"/tmp/", "/dev/shm/", "/.config/"}
// suspiciousCmdlinePatterns is checked only by the periodic
// CheckSuspiciousProcesses; the BPF exec backend cannot read cmdline at
// the moment of exec.
var suspiciousCmdlinePatterns = []string{
"/bin/sh -i", "/bin/bash -i", "bash -i",
"/dev/tcp/", "semutmerah", "gsocket",
"reverse", "nc -e", "ncat -e",
}
// EvaluateExec returns findings for a single execve event observed by the
// BPF live backend. Inputs are the (UID, PID, comm, exe, parentComm)
// tuple the kernel hook collects. It stamps the detection time because these
// findings enter the realtime bus directly. The legacy periodic checks
// (CheckSuspiciousProcesses, CheckFakeKernelThreads) keep using cmdline-aware
// detection that this function cannot replicate.
func EvaluateExec(uid uint32, pid uint32, comm, exe, parentComm string) []alert.Finding {
var out []alert.Finding
pidInt := int(pid)
detectedAt := time.Now()
if uid != 0 && len(comm) >= 2 && comm[0] == '[' && comm[len(comm)-1] == ']' {
out = append(out, alert.Finding{
Severity: alert.Critical,
Check: "fake_kernel_thread",
Message: fmt.Sprintf("Non-root process masquerading as kernel thread: %s", comm),
Details: fmt.Sprintf("PID: %d, UID: %d, exe: %s, parent: %s", pid, uid, exe, parentComm),
PID: pidInt,
Timestamp: detectedAt,
})
}
if uid == 0 {
return out
}
exeName := filepath.Base(exe)
exeNameLower := strings.ToLower(exeName)
for _, s := range suspiciousExeNames {
if strings.Contains(exeNameLower, s) {
out = append(out, alert.Finding{
Severity: alert.Critical,
Check: "suspicious_process",
Message: fmt.Sprintf("Suspicious process name: %s", exeName),
Details: fmt.Sprintf("PID: %d, UID: %d, exe: %s, comm: %s, parent: %s", pid, uid, exe, comm, parentComm),
PID: pidInt,
Timestamp: detectedAt,
})
break
}
}
for _, p := range suspiciousExePaths {
if strings.Contains(exe, p) {
out = append(out, alert.Finding{
Severity: alert.High,
Check: "suspicious_process",
Message: fmt.Sprintf("Process running from suspicious path: %s", exe),
Details: fmt.Sprintf("PID: %d, UID: %d, comm: %s, parent: %s", pid, uid, comm, parentComm),
PID: pidInt,
Timestamp: detectedAt,
})
break
}
}
return out
}
func CheckFakeKernelThreads(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
procs, _ := osFS.Glob("/proc/[0-9]*/status")
for _, statusPath := range procs {
pid := filepath.Base(filepath.Dir(statusPath))
data, err := osFS.ReadFile(statusPath)
if err != nil {
continue
}
var name, uid string
for _, line := range strings.Split(string(data), "\n") {
if strings.HasPrefix(line, "Name:\t") {
name = strings.TrimPrefix(line, "Name:\t")
}
if strings.HasPrefix(line, "Uid:\t") {
fields := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(fields) > 0 {
uid = fields[0]
}
}
}
// Kernel threads run as root (uid 0). Non-root process with
// a name that looks like a kernel thread is suspicious.
if uid == "0" || uid == "" {
continue
}
// Read cmdline - real kernel threads have empty cmdline
cmdline, _ := osFS.ReadFile(filepath.Join("/proc", pid, "cmdline"))
cmdStr := strings.TrimRight(strings.ReplaceAll(string(cmdline), "\x00", " "), " ")
safeCmdStr := redactProcCommandLine(cmdline)
// Check if the process name contains brackets (faking kernel thread)
// or if cmdline starts with [
if strings.HasPrefix(cmdStr, "[") || strings.HasPrefix(name, "[") {
// This is a non-root process masquerading as a kernel thread
exe, _ := osFS.Readlink(filepath.Join("/proc", pid, "exe"))
uidInt, _ := strconv.Atoi(uid)
pidInt, _ := strconv.Atoi(pid)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "fake_kernel_thread",
Timestamp: time.Now(),
Message: fmt.Sprintf("Non-root process masquerading as kernel thread: [%s]", name),
Details: fmt.Sprintf("PID: %s, UID: %d, exe: %s, cmdline: %s", pid, uidInt, exe, safeCmdStr),
PID: pidInt,
})
}
}
return findings
}
func CheckSuspiciousProcesses(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
suspiciousNames := suspiciousExeNames
suspiciousCmdline := suspiciousCmdlinePatterns
suspiciousPaths := suspiciousExePaths
procs, _ := osFS.Glob("/proc/[0-9]*/exe")
for _, exePath := range procs {
pid := filepath.Base(filepath.Dir(exePath))
pidInt, _ := strconv.Atoi(pid)
statusData, _ := osFS.ReadFile(filepath.Join("/proc", pid, "status"))
var uid string
for _, line := range strings.Split(string(statusData), "\n") {
if strings.HasPrefix(line, "Uid:\t") {
fields := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(fields) > 0 {
uid = fields[0]
}
}
}
if uid == "0" {
continue // Skip root processes for this check
}
exe, _ := osFS.Readlink(exePath)
cmdline, _ := osFS.ReadFile(filepath.Join("/proc", pid, "cmdline"))
cmdStr := strings.TrimRight(strings.ReplaceAll(string(cmdline), "\x00", " "), " ")
safeCmdStr := redactProcCommandLine(cmdline)
// Check executable name
exeName := filepath.Base(exe)
for _, s := range suspiciousNames {
if strings.Contains(strings.ToLower(exeName), s) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "suspicious_process",
Message: fmt.Sprintf("Suspicious process name: %s", exeName),
Timestamp: time.Now(),
Details: fmt.Sprintf("PID: %s, UID: %s, exe: %s, cmdline: %s", pid, uid, exe, safeCmdStr),
PID: pidInt,
})
}
}
// Check cmdline for suspicious patterns
cmdLower := strings.ToLower(cmdStr)
for _, s := range suspiciousCmdline {
if strings.Contains(cmdLower, strings.ToLower(s)) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "suspicious_process",
Message: fmt.Sprintf("Suspicious cmdline pattern: %s", s),
Timestamp: time.Now(),
Details: fmt.Sprintf("PID: %s, UID: %s, exe: %s, cmdline: %s", pid, uid, exe, safeCmdStr),
PID: pidInt,
})
break
}
}
// Check executable path
for _, s := range suspiciousPaths {
if strings.Contains(exe, s) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "suspicious_process",
Message: fmt.Sprintf("Process running from suspicious path: %s", exe),
Timestamp: time.Now(),
Details: fmt.Sprintf("PID: %s, UID: %s, cmdline: %s", pid, uid, safeCmdStr),
PID: pidInt,
})
break
}
}
}
return findings
}
// CheckPHPProcesses inspects running lsphp processes to detect active
// webshell execution. Only reads /proc cmdline - zero disk I/O.
func CheckPHPProcesses(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
suspiciousPHPPaths := []string{
"/tmp/",
"/dev/shm/",
"/wp-content/uploads/",
"/.config/",
}
procs, _ := osFS.Glob("/proc/[0-9]*/cmdline")
for _, cmdPath := range procs {
pid := filepath.Base(filepath.Dir(cmdPath))
pidInt, _ := strconv.Atoi(pid)
cmdline, err := osFS.ReadFile(cmdPath)
if err != nil {
continue
}
cmdStr := strings.ReplaceAll(string(cmdline), "\x00", " ")
safeCmdStr := redactProcCommandLine(cmdline)
// Only check lsphp processes
if !strings.Contains(cmdStr, "lsphp") {
continue
}
for _, sus := range suspiciousPHPPaths {
if strings.Contains(cmdStr, sus) {
statusData, _ := osFS.ReadFile(filepath.Join("/proc", pid, "status"))
uid, uidOK := processStatusUID(statusData)
uidText := ""
var proc *processctx.ProcessContext
if uidOK {
uidText = strconv.Itoa(uid)
userName, account := processIdentityForUID(uid)
proc = &processctx.ProcessContext{
PID: pidInt,
UID: uid,
User: userName,
Account: account,
}
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "php_suspicious_execution",
Message: fmt.Sprintf("PHP executing from suspicious path: %s", sus),
Details: fmt.Sprintf("PID: %s, UID: %s, cmdline: %s", pid, uidText, safeCmdStr),
PID: pidInt,
Process: proc,
})
break
}
}
}
return findings
}
func redactProcCommandLine(raw []byte) string {
return strings.TrimSpace(alert.RedactCommandLine(string(raw)))
}
package checks
import (
"context"
"errors"
"io"
"os"
"path/filepath"
"golang.org/x/sys/unix"
)
// ---------------------------------------------------------------------------
// OS abstraction — filesystem read operations
// ---------------------------------------------------------------------------
// OS abstracts filesystem operations (read and write) used by check functions.
// Production code uses realOS{}; tests swap in a mockOS via SetOS().
type OS interface {
ReadFile(name string) ([]byte, error)
ReadDir(name string) ([]os.DirEntry, error)
Stat(name string) (os.FileInfo, error)
Lstat(name string) (os.FileInfo, error)
Readlink(name string) (string, error)
Open(name string) (*os.File, error)
WriteFile(name string, data []byte, perm os.FileMode) error
MkdirAll(path string, perm os.FileMode) error
Remove(name string) error
Glob(pattern string) ([]string, error)
}
type realOS struct{}
var errNonRegularFile = errors.New("not a regular file")
var errFileChanged = errors.New("file changed while reading")
// #nosec G304 -- filesystem abstraction; check functions pass trusted paths.
func (realOS) ReadFile(name string) ([]byte, error) { return os.ReadFile(name) }
func (realOS) ReadDir(name string) ([]os.DirEntry, error) { return os.ReadDir(name) }
func (realOS) Stat(name string) (os.FileInfo, error) { return os.Stat(name) }
func (realOS) Lstat(name string) (os.FileInfo, error) { return os.Lstat(name) }
func (realOS) Readlink(name string) (string, error) { return os.Readlink(name) }
// #nosec G304 -- filesystem abstraction; check functions pass trusted paths.
func (realOS) openRegularFile(name string, extraFlags int) (*os.File, error) {
fd, err := unix.Open(name, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NONBLOCK|extraFlags, 0)
if err != nil {
return nil, err
}
// #nosec G115 -- unix.Open returns a non-negative descriptor whenever err is nil.
file := os.NewFile(uintptr(fd), name)
info, err := file.Stat()
if err != nil {
_ = file.Close()
return nil, err
}
if !info.Mode().IsRegular() {
_ = file.Close()
return nil, errNonRegularFile
}
return file, nil
}
// ReadRegularFile opens without blocking on a path raced to a FIFO, then
// verifies the opened object before reading from the inode bound to the fd.
func (r realOS) ReadRegularFile(name string) ([]byte, error) {
file, err := r.openRegularFile(name, 0)
if err != nil {
return nil, err
}
defer file.Close()
return io.ReadAll(file)
}
// ReadRegularFilePrefix bounds consumers that need only an admission prefix.
// Refusing symlinks and matching the opened inode to the caller's snapshot
// stop a path replacement from redirecting a content-based admission decision.
func (r realOS) ReadRegularFilePrefix(name string, expected os.FileInfo, limit int64) ([]byte, error) {
file, err := r.openRegularFile(name, unix.O_NOFOLLOW)
if err != nil {
return nil, err
}
defer file.Close()
opened, err := file.Stat()
if err != nil {
return nil, err
}
if !sameFileSnapshot(expected, opened) {
return nil, errFileChanged
}
prefix, err := io.ReadAll(io.LimitReader(file, limit))
if err != nil {
return prefix, err
}
after, err := file.Stat()
if err != nil {
return prefix, err
}
if !sameFileSnapshot(opened, after) {
return prefix, errFileChanged
}
return prefix, nil
}
func sameFileSnapshot(expected, actual os.FileInfo) bool {
return expected != nil && actual != nil &&
os.SameFile(expected, actual) &&
expected.Mode() == actual.Mode() &&
expected.Size() == actual.Size() &&
expected.ModTime().Equal(actual.ModTime())
}
// #nosec G304 -- filesystem abstraction; check functions pass trusted paths.
func (realOS) Open(name string) (*os.File, error) { return os.Open(name) }
// #nosec G306 -- callers pass explicit perm; intent is operator-readable mode.
func (realOS) WriteFile(name string, data []byte, perm os.FileMode) error {
return os.WriteFile(name, data, perm)
}
func (realOS) MkdirAll(path string, perm os.FileMode) error { return os.MkdirAll(path, perm) }
func (realOS) Remove(name string) error { return os.Remove(name) }
func (realOS) Glob(pattern string) ([]string, error) { return filepath.Glob(pattern) }
// osFS is the package-level filesystem provider. All check functions use
// this instead of calling os.ReadFile / os.ReadDir / etc. directly.
var osFS OS = realOS{}
// SetOS replaces the filesystem provider. Used by tests to inject mocks.
func SetOS(o OS) { osFS = o }
// ---------------------------------------------------------------------------
// CmdRunner abstraction — external command execution
// ---------------------------------------------------------------------------
// CmdRunner abstracts external command execution used by check functions.
// Production code uses realCmd{}; tests swap in a mockCmdRunner via SetCmdRunner().
//
// RunContext returns stdout+stderr merged (CombinedOutput) and is fine for
// tools that only write to stdout. RunContextStdout returns stdout only and
// should be used when the command prints structured output (JSON, a URL, ...)
// on stdout and chatter (warnings, PHP notices, MySQL deprecations) on stderr
// -- mixing them there would corrupt the parse. RunContextStdout also surfaces
// context.DeadlineExceeded on timeout so callers can distinguish "no output"
// from "empty output".
type CmdRunner interface {
Run(name string, args ...string) ([]byte, error)
RunAllowNonZero(name string, args ...string) ([]byte, error)
RunContext(parent context.Context, name string, args ...string) ([]byte, error)
RunContextStdout(parent context.Context, name string, args ...string) ([]byte, error)
RunWithEnv(name string, args []string, extraEnv ...string) ([]byte, error)
LookPath(file string) (string, error)
}
type realCmd struct{}
func (realCmd) Run(name string, args ...string) ([]byte, error) {
return runCmdReal(name, args...)
}
func (realCmd) RunAllowNonZero(name string, args ...string) ([]byte, error) {
return runCmdAllowNonZeroReal(name, args...)
}
func (realCmd) RunContext(parent context.Context, name string, args ...string) ([]byte, error) {
return runCmdCombinedContextReal(parent, name, args...)
}
func (realCmd) RunContextStdout(parent context.Context, name string, args ...string) ([]byte, error) {
return runCmdStdoutContextReal(parent, name, args...)
}
func (realCmd) RunWithEnv(name string, args []string, extraEnv ...string) ([]byte, error) {
return runCmdWithEnvReal(name, args, extraEnv...)
}
func (realCmd) LookPath(file string) (string, error) {
return lookupSystemCommand(file)
}
// cmdExec is the package-level command runner. All check functions use
// this instead of calling runCmd / exec.Command directly.
var cmdExec CmdRunner = realCmd{}
// SetCmdRunner replaces the command runner. Used by tests to inject mocks.
func SetCmdRunner(r CmdRunner) { cmdExec = r }
package checks
import (
"encoding/json"
"os"
"syscall"
"time"
)
// QuarantineMeta stores original file metadata alongside quarantined files.
type QuarantineMeta struct {
OriginalPath string `json:"original_path"`
Owner int `json:"owner_uid"`
Group int `json:"group_gid"`
Mode string `json:"mode"`
Size int64 `json:"size"`
QuarantineAt time.Time `json:"quarantined_at"`
OriginalModTime time.Time `json:"original_mtime,omitzero"`
Reason string `json:"reason"`
// FindingID ties the quarantine to the finding that caused it, using the
// same identifier the audit log and the action log emit. Empty when the
// quarantine came from an operator command rather than a detection.
FindingID string `json:"finding_id,omitempty"`
MessageID string `json:"message_id,omitempty"`
SpoolDir string `json:"spool_dir,omitempty"`
RestoreAction string `json:"restore_action,omitempty"`
ExpectedCurrentSHA256 string `json:"expected_current_sha256,omitempty"`
}
// UnmarshalJSON accepts the timestamp spelling used by historical manual
// fixes. Missing original mtimes remain unknown; archive mtimes are not evidence
// of when the original was modified.
func (m *QuarantineMeta) UnmarshalJSON(data []byte) error {
type wire QuarantineMeta
var decoded struct {
wire
LegacyQuarantineAt time.Time `json:"quarantine_at"`
}
if err := json.Unmarshal(data, &decoded); err != nil {
return err
}
if decoded.QuarantineAt.IsZero() {
decoded.QuarantineAt = decoded.LegacyQuarantineAt
}
*m = QuarantineMeta(decoded.wire)
return nil
}
func quarantineMetadata(path string, info os.FileInfo, reason string) QuarantineMeta {
meta := QuarantineMeta{
OriginalPath: path,
Mode: info.Mode().String(),
Size: info.Size(),
QuarantineAt: time.Now().UTC(),
OriginalModTime: info.ModTime().UTC(),
Reason: reason,
}
if stat, ok := info.Sys().(*syscall.Stat_t); ok {
meta.Owner, meta.Group = int(stat.Uid), int(stat.Gid)
}
return meta
}
package checks
import (
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"path/filepath"
"strings"
"time"
)
// quarantineSafeNameMax keeps the generated name, plus the timestamp prefix
// callers add, under the 255-byte filename limit.
const quarantineSafeNameMax = 180
func newQuarantinePath(dir, original string) string {
// Repeated cleanups can occur within one second. Independent names keep
// a new recovery point from replacing the only copy of an earlier state.
return filepath.Join(dir, time.Now().UTC().Format("20060102-150405")+"_"+rand.Text()+"_"+quarantineSafeName(original))
}
// quarantineSafeName turns a source path into one flat quarantine filename.
// Short paths keep the familiar slash-to-underscore form. A path that would
// exceed the filename limit is shortened to a hash of the whole path plus
// its tail, so the file name survives, two long paths cannot collide, and
// the move no longer fails with ENAMETOOLONG (which left the malware in
// place while the finding reported a quarantine attempt).
func quarantineSafeName(path string) string {
flat := strings.ReplaceAll(path, "/", "_")
if len(flat) <= quarantineSafeNameMax {
return flat
}
sum := sha256.Sum256([]byte(path))
prefix := hex.EncodeToString(sum[:6])
tail := flat[len(flat)-(quarantineSafeNameMax-len(prefix)-1):]
return prefix + "_" + tail
}
//go:build linux
package checks
import (
"fmt"
"os"
"syscall"
)
// var so Linux tests can interleave a source-path mutation after the verified
// fd is open without racing the test process itself.
var quarantineCopyByFD = copyQuarantineFileByFD
// quarantineFileTOCTOUSafe moves a single regular file into quarantine in
// a way that defends against the classic detect-then-quarantine race: an
// attacker who controls the directory can swap a legitimate file in
// between Lstat and Rename, tricking CSM into moving the wrong file out
// of the user's home. The defence:
//
// 1. Open the path with O_RDONLY|O_NOFOLLOW. Symlinks are refused at
// the kernel level; the fd is bound to the inode that existed at
// open time.
// 2. Fstat the fd and verify it still matches the inode we detected
// earlier (sameFileIdentity). A late swap loses here.
// 3. Copy from that verified fd into a private, independent quarantine
// inode. A hardlink is never used: the account could add another name
// after the initial fstat and keep the quarantine inode writable.
// Persist the copy, metadata, and quarantine directory before unlinking.
// 4. Unlink the source path only if it still resolves to the inode
// we quarantined. If an attacker swapped in a replacement after
// step 2, leave that replacement alone.
//
// Returns nil on success. Errors describe what failed; callers should
// not retry blindly because a failure usually means the file moved.
func quarantineFileTOCTOUSafe(path, qPath string, originalInfo os.FileInfo, metadata []byte) error {
if originalInfo == nil {
return fmt.Errorf("quarantine: missing original stat")
}
if originalInfo.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("quarantine: refused symlink at %s", path)
}
// O_NOFOLLOW makes open() fail with ELOOP if path resolved to a
// symlink in the final component; combined with the earlier Lstat
// rejection above, this closes the symlink-swap variant.
// #nosec G304 -- path is the quarantine subject; O_NOFOLLOW plus
// fd identity verification below fail closed on symlink and inode swaps.
fd, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW|syscall.O_NONBLOCK, 0)
if err != nil {
return fmt.Errorf("quarantine: open %s: %w", path, fileResponseSourceError(err))
}
defer fd.Close()
// Re-stat the open fd and confirm the inode matches what we
// detected. The race window between Lstat and OpenFile is narrow
// but real - an attacker who hits it gets caught here.
cur, err := fd.Stat()
if err != nil {
return fmt.Errorf("quarantine: fstat %s: %w", path, err)
}
if !sameFileIdentity(cur, originalInfo) {
return refuseFileResponse(fmt.Errorf("quarantine: file at %s changed between detection and quarantine (TOCTOU)", path))
}
// Defence against inode reuse: on busy tmpfs / ext4 mounts the kernel
// can hand out the freed inode to whatever the attacker wrote next.
// A matching inode is necessary but not sufficient; also require the
// content shape (size + mtime) to match what the detector recorded.
if !sameContentShape(cur, originalInfo) {
return refuseFileResponse(fmt.Errorf("quarantine: file at %s changed between detection and quarantine (TOCTOU, inode reused)", path))
}
// Refuse to quarantine a non-regular file (block, char, socket,
// FIFO). The detector only flags regular files, so a non-regular
// shape at this point means someone is trying to move CSM at a
// device node or pipe.
if !cur.Mode().IsRegular() {
return fmt.Errorf("quarantine: refusing non-regular file at %s (mode=%v)", path, cur.Mode())
}
// Always create an independent root-owned copy. Checking st_nlink before a
// hardlink is not sufficient: the account can add another name after the
// check and retain write access to the inode placed in quarantine.
if err = quarantineCopyByFD(fd, qPath, metadata); err != nil {
return fmt.Errorf("quarantine: copy %s -> %s: %w", path, qPath, err)
}
if err = removeQuarantinedSource(path, qPath, cur); err != nil {
return err
}
remaining, err := fd.Stat()
if err != nil {
return &quarantineCompletedWarning{message: fmt.Sprintf("quarantine: copied %s to %s and removed that name, but could not count surviving hard links: %v", path, qPath, err)}
}
if links := fileLinkCount(remaining); links > 0 {
return &quarantineCompletedWarning{message: fmt.Sprintf("quarantine: copied %s to %s and removed that name, but at least %d other hard link(s) to the same content remain reachable elsewhere", path, qPath, links)}
}
return nil
}
// fileLinkCount returns the inode's link count, or 1 when the stat carries
// no platform data.
func fileLinkCount(info os.FileInfo) uint64 {
if st, ok := info.Sys().(*syscall.Stat_t); ok {
return uint64(st.Nlink) // #nosec G115 -- link count, never negative.
}
return 1
}
package checks
import (
"errors"
"fmt"
"io"
"os"
"path/filepath"
"github.com/pidginhost/csm/internal/quarantinefs"
)
type quarantineCompletedWarning struct {
message string
}
func (w *quarantineCompletedWarning) Error() string {
return w.message
}
func completedQuarantineWarning(err error) (string, bool) {
var warning *quarantineCompletedWarning
if !errors.As(err, &warning) {
return "", false
}
return warning.Error(), true
}
// sameContentShape verifies that two stats describe a file with the same
// size and modification time. Used as a defence-in-depth check after
// sameFileIdentity passes, because inode reuse on tmpfs / ext4 lets an
// attacker recreate a file under the same path with a fresh ino that
// happens to match the freed slot.
func sameContentShape(a, b os.FileInfo) bool {
if a == nil || b == nil {
return false
}
if a.Size() != b.Size() {
return false
}
return a.ModTime().Equal(b.ModTime())
}
// copyQuarantineFileByFD copies the already-open source into qPath. The copy
// is created by the daemon (root-owned, 0600), so unlike a hard link it is
// never writable through any name the account still holds.
func copyQuarantineFileByFD(src *os.File, qPath string, metadata []byte) error {
if _, err := src.Seek(0, io.SeekStart); err != nil {
return fmt.Errorf("seek source: %w", err)
}
return quarantinefs.Store(qPath, src, metadata, 0600)
}
var quarantineUnlinkSource = os.Remove
var quarantineSyncSourceDir = quarantinefs.SyncDir
// removeQuarantinedSource unlinks the detected name once the content sits in
// quarantine, only if the name still resolves to the inode that was captured.
// A source that vanished is done. A source that now resolves to something
// else was swapped in by an attacker racing the unlink: the replacement is
// left alone, the captured copy is kept as evidence, and the caller is told
// the remediation did not complete instead of a silent success that left the
// replacement live under the detected name.
func removeQuarantinedSource(path, qPath string, original os.FileInfo) error {
info, err := os.Lstat(path)
if err != nil {
if os.IsNotExist(err) {
if syncErr := quarantineSyncSourceDir(filepath.Dir(path)); syncErr != nil {
return fmt.Errorf("quarantine: source vanished but removal is not durable; recovery copy retained at %s: %w", qPath, syncErr)
}
return nil
}
return fmt.Errorf("quarantine: stat source before unlink %s; recovery copy retained at %s: %w", path, qPath, fileResponseSourceError(err))
}
if info.Mode()&os.ModeSymlink != 0 || !sameFileIdentity(info, original) {
return refuseFileResponse(fmt.Errorf("quarantine: source at %s was replaced before unlink; the detected content is kept at %s and the replacement was left in place", path, qPath))
}
if !sameContentShape(info, original) {
// A copy made while the source changed may mix old and new bytes.
// Keep the live source and discard this untrustworthy recovery copy.
cause := fmt.Errorf("quarantine: source changed before unlink %s", path)
if err := os.Remove(qPath); err != nil && !os.IsNotExist(err) {
return errors.Join(cause, err)
}
if err := os.Remove(qPath + ".meta"); err != nil && !os.IsNotExist(err) {
return errors.Join(cause, err)
}
if err := quarantinefs.SyncDir(filepath.Dir(qPath)); err != nil {
return errors.Join(cause, err)
}
return refuseFileResponse(cause)
}
if err := quarantineUnlinkSource(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("quarantine: unlink source %s; recovery copy retained at %s: %w", path, qPath, err)
}
if err := quarantineSyncSourceDir(filepath.Dir(path)); err != nil {
return fmt.Errorf("quarantine: source removal is not durable; recovery copy retained at %s, inspect original path %s before retrying: %w", qPath, path, err)
}
return nil
}
package checks
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/quarantinefs"
"github.com/pidginhost/csm/internal/safepath"
)
// quarantineTarget moves a file or directory into quarantine and records the
// action. It is the single point every quarantine passes through, which is why
// the action record is taken here rather than at each caller.
func quarantineTarget(path, qPath string, info os.FileInfo, metadata QuarantineMeta) (err error) {
before := actionlog.FromInfo(info)
existed := actionlog.Metadata(qPath).Exists
defer func() { recordQuarantineAction(path, qPath, metadata, before, existed, err) }()
data, err := json.MarshalIndent(metadata, "", " ")
if err != nil {
return fmt.Errorf("encoding quarantine metadata: %w", err)
}
if err := quarantinefs.EnsureDir(filepath.Dir(qPath), 0700); err != nil {
return err
}
return quarantineTargetFn(path, qPath, info, data)
}
// quarantineTargetFn performs the move. It is indirected so the action record
// can be tested against a failing move without a read-only filesystem.
var quarantineTargetFn = func(path, qPath string, info os.FileInfo, data []byte) error {
if info.IsDir() {
return quarantineDirectory(path, qPath, info, data)
}
return quarantineFileTOCTOUSafe(path, qPath, info, data)
}
// recordQuarantineAction writes the unified action record for one quarantine.
// The digest of the removed content is the evidence a reviewer needs: it
// distinguishes "this exact file left the account" from "something was moved".
func recordQuarantineAction(path, qPath string, metadata QuarantineMeta, before *actionlog.FileState, existed bool, err error) {
rec := actionlog.Record{
Op: "respond.quarantine_file",
Actor: actionlog.DefaultActor(),
Target: path,
Reason: metadata.Reason,
FindingID: metadata.FindingID,
Before: before,
After: actionlog.Metadata(path),
Result: actionlog.Applied,
}
// Hash the retained copy only after capture, never a tenant-controlled
// source before the security action. Oversized copies remain unhashed.
if !existed {
retained := actionlog.Stat(qPath)
if retained.Exists {
rec.Before.Digest = retained.Digest
rec.RecoveryPath = qPath
}
}
if err != nil {
rec.Result = actionlog.Failed
rec.Error = err.Error()
if _, completed := completedQuarantineWarning(err); completed {
rec.Result = actionlog.Applied
} else if errors.Is(err, errFileResponseRefused) {
rec.Result = actionlog.Refused
}
}
actionlog.Write(rec)
}
var storeQuarantineBackup = func(path string, content []byte, metadata QuarantineMeta, mode os.FileMode) error {
data, err := json.MarshalIndent(metadata, "", " ")
if err != nil {
return fmt.Errorf("encoding quarantine metadata: %w", err)
}
return quarantinefs.Store(path, bytes.NewReader(content), data, mode)
}
var syncQuarantineTree = quarantinefs.SyncTree
var syncQuarantineDirectory = (*safepath.Dir).Sync
var renameQuarantineDirectory = (*safepath.Dir).RenameTo
func quarantineDirectory(path, qPath string, expected os.FileInfo, metadata []byte) error {
source, err := safepath.OpenDir(filepath.Dir(path))
if err != nil {
return err
}
defer func() { _ = source.Close() }()
quarantine, err := safepath.OpenDir(filepath.Dir(qPath))
if err != nil {
return err
}
defer func() { _ = quarantine.Close() }()
sourceName, name := filepath.Base(path), filepath.Base(qPath)
dir, err := source.OpenFile(sourceName, os.O_RDONLY|unix.O_DIRECTORY, 0)
if err != nil {
return err
}
defer dir.Close()
info, err := dir.Stat()
if err != nil {
return err
}
if !os.SameFile(info, expected) {
return errors.New("quarantine directory changed before capture")
}
if err := syncQuarantineTree(path, info); err != nil {
return fmt.Errorf("syncing quarantine directory: %w", err)
}
if err := dir.Close(); err != nil {
return fmt.Errorf("closing quarantine directory before move: %w", err)
}
if err := quarantinefs.WriteExclusive(qPath+".meta", bytes.NewReader(metadata), 0600); err != nil {
return fmt.Errorf("writing quarantine directory metadata: %w", err)
}
if err := renameQuarantineDirectory(source, sourceName, quarantine, name); err != nil {
return errors.Join(fmt.Errorf("moving quarantine directory: %w", err), os.Remove(qPath+".meta"))
}
got, statErr := quarantine.Stat(name)
if statErr != nil || !os.SameFile(got, info) {
if rollbackErr := quarantine.RenameTo(name, source, sourceName); rollbackErr != nil {
return fmt.Errorf("quarantine directory changed during capture; displaced directory retained at %s: %w", qPath, rollbackErr)
}
if syncErr := errors.Join(source.Sync(), quarantine.Sync()); syncErr != nil {
return fmt.Errorf("quarantine directory changed; rollback sync failed, inspect %s and %s: %w", path, qPath, syncErr)
}
if removeErr := os.Remove(qPath + ".meta"); removeErr != nil {
return fmt.Errorf("quarantine directory changed; original name restored but metadata retained at %s: %w", qPath+".meta", removeErr)
}
return errors.New("quarantine directory changed during capture; original name restored")
}
if err := syncQuarantineDirectory(quarantine); err != nil {
return fmt.Errorf("quarantine directory moved to %s, but destination sync failed; inspect both %s and %s before retrying: %w", qPath, path, qPath, err)
}
if err := syncQuarantineDirectory(source); err != nil {
return fmt.Errorf("quarantine directory retained at %s, but source removal is not durable; inspect %s before retrying: %w", qPath, path, err)
}
return nil
}
package checks
import (
"container/list"
"net"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
// RDNSCacheConfig is the config block for NewRDNSCache. Resolve is the
// function used to perform the actual reverse lookup; production
// callers wrap net.LookupAddr. ResolveDeadline bounds each lookup;
// 0 disables the deadline. MaxSize caps the number of cached entries
// to keep memory bounded on hosts that see a wide spread of remote
// IPs (BPF SMTP-egress is the motivating case); the oldest entry by
// cachedAt is evicted before a new one is inserted past the cap.
// 0 falls back to rdnsCacheDefaultMaxSize.
type RDNSCacheConfig struct {
TTL time.Duration
Resolve func(ip net.IP) (string, error)
ResolveDeadline time.Duration
MaxSize int
// MaxConcurrent caps deadline-bound lookups, including returned results
// still owned by their callers. A goroutine blocked on a wedged resolver cannot be
// cancelled in Go, so without a cap a burst of distinct IPs under deadline
// saturation spawns one abandonable goroutine per IP. 0 falls back to
// rdnsCacheDefaultMaxConcurrent.
MaxConcurrent int
}
const (
rdnsCacheDefaultMaxSize = 10000
rdnsCacheDefaultMaxConcurrent = 64
)
// RDNSCache is a small TTL cache around reverse DNS lookups. Cached
// negative results (resolver error / NXDOMAIN) are kept until TTL too,
// so the detector does not hammer a slow resolver on a known-bad IP.
// Entries are capped at maxSize; the oldest-by-cachedAt entry is
// evicted on insert once the cap is reached.
type RDNSCache struct {
mu sync.Mutex
ttl time.Duration
deadln time.Duration
maxSize int
resolve func(ip net.IP) (string, error)
now func() time.Time
order *list.List
entries map[string]*list.Element
sem chan struct{}
stats *queuehealth.Tracker
}
type rdnsEntry struct {
key string
host string
cachedAt time.Time
}
// NewRDNSCache returns a ready cache.
func NewRDNSCache(cfg RDNSCacheConfig) *RDNSCache {
maxSize := cfg.MaxSize
if maxSize <= 0 {
maxSize = rdnsCacheDefaultMaxSize
}
maxConcurrent := cfg.MaxConcurrent
if maxConcurrent <= 0 {
maxConcurrent = rdnsCacheDefaultMaxConcurrent
}
c := &RDNSCache{
ttl: cfg.TTL,
deadln: cfg.ResolveDeadline,
maxSize: maxSize,
resolve: cfg.Resolve,
now: time.Now,
order: list.New(),
entries: map[string]*list.Element{},
sem: make(chan struct{}, maxConcurrent),
}
if cfg.ResolveDeadline > 0 {
c.stats = queuehealth.NewSharedCapacity(maxConcurrent, cfg.ResolveDeadline)
}
return c
}
// evictOldestLocked drops the oldest cached entry. Caller holds c.mu.
func (c *RDNSCache) evictOldestLocked() {
el := c.order.Front()
if el == nil {
return
}
entry := el.Value.(*rdnsEntry)
delete(c.entries, entry.key)
c.order.Remove(el)
}
// Lookup returns the cached hostname for ip, or "" on miss/error/deadline.
// Lookup blocks the caller for at most cfg.ResolveDeadline; cache hits
// return immediately.
func (c *RDNSCache) Lookup(ip net.IP) string {
if ip == nil {
return ""
}
key := ip.String()
c.mu.Lock()
now := c.now()
if el, ok := c.entries[key]; ok {
e := el.Value.(*rdnsEntry)
if now.Sub(e.cachedAt) <= c.ttl {
c.mu.Unlock()
return e.host
}
}
c.mu.Unlock()
host := c.runWithDeadline(ip)
c.mu.Lock()
now = c.now()
if el, present := c.entries[key]; present {
e := el.Value.(*rdnsEntry)
e.host = host
e.cachedAt = now
c.order.MoveToBack(el)
c.mu.Unlock()
return host
}
if c.order.Len() >= c.maxSize {
c.evictOldestLocked()
}
c.entries[key] = c.order.PushBack(&rdnsEntry{key: key, host: host, cachedAt: now})
c.mu.Unlock()
return host
}
func (c *RDNSCache) runWithDeadline(ip net.IP) string {
if c.deadln <= 0 {
host, err := c.resolve(ip)
if err != nil {
return ""
}
return host
}
// Cap in-flight resolve goroutines. A goroutine blocked on a wedged
// resolver keeps its slot until the syscall finally returns, so under
// deadline saturation further lookups fail fast (return "" like a
// deadline miss) instead of spawning more abandonable goroutines.
work := c.acquireResolve()
if work == nil {
return ""
}
defer work.release()
ch := make(chan rdnsResult, 1)
go c.resolveTracked(work, ip, ch)
timer := time.NewTimer(c.deadln)
defer timer.Stop()
select {
case r := <-ch:
if r.err != nil {
return ""
}
return r.host
case <-timer.C:
work.fail()
return ""
}
}
package checks
import (
"errors"
"net"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type rdnsWork struct {
cache *RDNSCache
ticket queuehealth.Ticket
remaining atomic.Int32
failOnce sync.Once
}
type rdnsResult struct {
host string
err error
}
func (c *RDNSCache) acquireResolve() *rdnsWork {
ticket := c.stats.Begin(time.Now())
select {
case c.sem <- struct{}{}:
work := &rdnsWork{cache: c, ticket: ticket}
work.remaining.Store(2)
return work
default:
ticket.Reject(time.Now())
return nil
}
}
func (w *rdnsWork) fail() {
w.failOnce.Do(func() { w.cache.stats.Lose(time.Now(), 1) })
}
func (w *rdnsWork) release() {
// The caller can leave while DNS still runs, or DNS can return before
// its caller receives the buffered result. Both own the admitted slot.
if w.remaining.Add(-1) == 0 {
w.ticket.Finish(time.Now())
<-w.cache.sem
}
}
func (c *RDNSCache) resolveTracked(work *rdnsWork, ip net.IP, done chan<- rdnsResult) {
work.ticket.Start(time.Now())
completed := false
defer func() {
if !completed {
work.fail()
}
work.release()
}()
host, err := c.resolve(ip)
var dnsErr *net.DNSError
if err != nil && (!errors.As(err, &dnsErr) || !dnsErr.IsNotFound) {
work.fail()
}
done <- rdnsResult{host: host, err: err}
completed = true
}
// QueueStatuses reads memory only. Without a deadline, Lookup is synchronous
// and does not use this bounded resolver pool.
func (c *RDNSCache) QueueStatuses(now time.Time) map[string]queuehealth.Status {
if c.stats == nil {
return nil
}
return map[string]queuehealth.Status{"resolves": c.stats.Snapshot(now)}
}
package checks
import "sort"
// CheckInfo describes a single named check emitted as an alert.Finding.Check.
// Category groups related checks for display in the settings UI. Internal is
// true for checks that exist for plumbing (self-tests, plumbing findings) and
// should not appear in user-facing dropdowns like alerts.email.disabled_checks.
type CheckInfo struct {
Name string
Category string
Internal bool
// Correlation says how cross-account correlation treats this check.
// Every entry must set it; the zero value fails TestEveryCheckIsClassified.
Correlation CorrelationClass
// CorrelationReason names the policy that excludes an Ignored check. One
// of the reason constants in correlation_policy.go.
CorrelationReason string
// CorrelationGap documents a known missing producer identity path for an
// eligible check. It never changes eligibility.
CorrelationGap string
// Response is the check's automatic IP response policy. The zero value
// neither blocks nor challenges; see response_policy.go.
Response ResponsePolicy
}
// Category labels are the groupings shown in the multi-select UI. Keep the
// order below in sync with checkCategoryOrder so categories render in a sane
// order rather than alphabetically (Auth first, Internal last).
const (
CategoryAuth = "Authentication & Login"
CategoryBruteForce = "Brute Force"
CategoryMalware = "Malware & Webshells"
CategoryWeb = "Web & Application"
CategoryDatabase = "Database Content"
CategoryEmail = "Email & Phishing"
CategoryPerformance = "Performance"
CategoryNetwork = "Network & Firewall"
CategorySystem = "System Integrity"
CategoryWAF = "WAF & ModSecurity"
CategoryCorrelation = "Correlation & Health"
CategoryInternal = "Internal"
)
var checkCategoryOrder = []string{
CategoryAuth,
CategoryBruteForce,
CategoryMalware,
CategoryWeb,
CategoryDatabase,
CategoryEmail,
CategoryPerformance,
CategoryNetwork,
CategorySystem,
CategoryWAF,
CategoryCorrelation,
CategoryInternal,
}
// checkRegistry is the authoritative list of every Check string the daemon
// may emit. Adding a new alert.Finding Check name anywhere in internal/checks,
// internal/daemon, or internal/webui without also adding it here will fail
// TestCheckRegistryCoversProductionCode.
var checkRegistry = []CheckInfo{
// --- Authentication & Login ------------------------------------------
// The admin-panel detector's tight path set makes false positives unlikely;
// a challenge would only delay containment of repeated attacks.
{Name: "admin_panel_bruteforce", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "api_auth_failure", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockWithCpanelLogins}},
// API clients have no browser to answer the challenge.
{Name: "api_auth_failure_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockWithCpanelLogins, NeverChallenge: true}},
{Name: "api_tokens", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "bulk_password_change", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAccountAggregate},
{Name: "cpanel_file_upload", Category: CategoryAuth, Correlation: CorrelationSecurityEvent},
{Name: "cpanel_file_upload_realtime", Category: CategoryAuth, Correlation: CorrelationSecurityEvent},
{Name: "cpanel_login", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "cpanel_login_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "cpanel_multi_ip_login", Category: CategoryAuth, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{Block: BlockWithCpanelLogins}},
{Name: "cpanel_password_purge", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "cpanel_password_purge_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
// FTP authentication cannot be gated by an HTTP challenge.
{Name: "ftp_auth_failure_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockWithCpanelLogins, NeverChallenge: true}},
{Name: "ftp_bruteforce", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways}},
{Name: "ftp_login", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "ftp_login_after_bruteforce", Category: CategoryAuth, Correlation: CorrelationSecurityEvent},
// PAM breadth and brute-force signals come from non-browser clients.
{Name: "credential_stuffing", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "pam_bruteforce", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "pam_login", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "password_hijack_confirmed", Category: CategoryAuth, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "root_password_change", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "shadow_change", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "ssh_keys", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "ssh_login_unknown_ip", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational, Response: ResponsePolicy{Block: BlockAlways}},
{Name: "sshd_config_change", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "uid0_account", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope, Response: ResponsePolicy{NeverChallenge: true}},
// Pre-auth attacks on a public webmail login page can meet the gate on
// their next request, unlike successful post-auth webmail audit events.
{Name: "webmail_bruteforce", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockWithCpanelLogins, ChallengeFirst: true}},
{Name: "webmail_login_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "whm_account_action", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "whm_login_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "whm_password_change", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "whm_password_change_noninfra", Category: CategoryAuth, Correlation: CorrelationSecurityEvent},
{Name: "whm_unauth_scripts_realtime", Category: CategoryAuth, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
// --- Brute Force -----------------------------------------------------
{Name: "http_request_flood", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways}},
// URL enumeration is browser-visible; http_scanner_action can still opt
// this check out of the challenge in favor of a direct block.
{Name: "http_scanner_profile", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, ChallengeFirst: true}},
// A claimed crawler may still pass reverse-DNS verification next cycle.
// Challenge it without timeout escalation while that verdict is pending;
// with challenge routing disabled, the existing block policy still applies.
{Name: "http_claimed_bot_unverified", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, ChallengeFirst: true}},
{Name: "http_ua_spoof", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways}},
{Name: "http_distributed_flood", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "http_asn_crawl", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
// Mail protocols cannot answer an HTTP challenge. Compromise severity
// decides blocking.
{Name: "mail_account_compromised", Category: CategoryBruteForce, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{Block: BlockAlways, CriticalOnly: true, NeverChallenge: true}},
{Name: "mail_account_spray", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "mail_bruteforce", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "mail_bruteforce_suspected", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
// A subnet summary does not authorize a single-IP block; the subnet-spray
// path blocks the subnet itself.
{Name: "mail_subnet_spray", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "smtp_account_spray", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
// SMTP authentication, connection probes and subnet sprays have no browser
// at the other end, so a challenge cannot stop their traffic.
{Name: "smtp_bruteforce", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "smtp_probe_abuse", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "smtp_subnet_spray", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{NeverChallenge: true}},
// These WordPress probes reach public HTTP endpoints before authentication;
// the next request from the same source can be sent to the browser gate.
{Name: "wp_login_bruteforce", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, ChallengeFirst: true}},
{Name: "wp_user_enumeration", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{ChallengeFirst: true}},
{Name: "xmlrpc_abuse", Category: CategoryBruteForce, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, ChallengeFirst: true}},
// --- Malware & Webshells --------------------------------------------
{Name: "backdoor_binary", Category: CategoryMalware, Correlation: CorrelationMalwareArtifact, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "cgi_backdoor_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "cgi_suspicious_location_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "cross_account_malware", Category: CategoryMalware, Correlation: CorrelationDerived, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "executable_in_config_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "executable_in_tmp_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "fake_kernel_thread", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "group_writable_php", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "js_keylogger_dataflow", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "js_taint_scan_incomplete", Category: CategoryMalware, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "php_remote_taint", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "php_taint_scan_incomplete", Category: CategoryMalware, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "new_executable_in_config", Category: CategoryMalware, Correlation: CorrelationMalwareArtifact},
{Name: "new_php_in_sensitive_dir", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "new_php_in_sensitive_dir_clean", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "new_php_in_uploads", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "new_php_in_uploads_clean", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
// Retired: emitted by the file index until a20c6f76 and never registered
// afterwards. Kept registered so the file_index runner can purge findings
// written by older versions; nothing emits them today.
{Name: "new_php_in_languages", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "new_php_in_upgrade", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "new_suspicious_php", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "new_webshell_file", Category: CategoryMalware, Correlation: CorrelationMalwareArtifact},
{Name: "nulled_plugin", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "obfuscated_php", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "obfuscated_php_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "php_dropper_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "php_in_image_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "php_in_sensitive_dir_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "php_in_uploads_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "self_deleting_dropper_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "self_deleting_dropper_overflow", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "php_shield_block", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "php_shield_eval", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "php_shield_webshell", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "php_suspicious_execution", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "signature_match_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "suid_binary", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "suspicious_file", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "suspicious_php_content", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "suspicious_process", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "webshell", Category: CategoryMalware, Correlation: CorrelationMalwareArtifact, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "webshell_content_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "webshell_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent},
{Name: "world_writable_php", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "yara_match_realtime", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "yara_match_scheduled", Category: CategoryMalware, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "yara_scan_incomplete", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "yara_realtime_scan_error", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "yara_worker_crashed", Category: CategoryMalware, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
// --- Web & Application ----------------------------------------------
{Name: "htaccess_auto_prepend", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_cgi_handler_abuse", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_errordocument_hijack", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_filesmatch_shield", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_handler_abuse", Category: CategoryWeb, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "htaccess_header_injection", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_injection", Category: CategoryWeb, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "htaccess_injection_realtime", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_php_in_uploads", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_security_disabled", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "web_exposed_backup_archive", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "web_exposed_config_leak", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "web_exposed_repo_metadata", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "web_exposed_db_dump", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "web_exposed_phpinfo", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "web_exposed_sample_sql", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "web_exposed_source_backup", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "htaccess_spam_redirect", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "htaccess_user_agent_cloak", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "open_basedir", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "outdated_plugins", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "vulnerable_plugins", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "vulnerable_timthumb", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "php_config_change", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "php_config_scan_incomplete", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "php_config_realtime", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "symlink_attack", Category: CategoryWeb, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "wp_core_integrity", Category: CategoryWeb, Correlation: CorrelationSecurityEvent},
{Name: "wp_core_unverified", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "wp_plugin_inventory_unverified", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
// --- Database Content -----------------------------------------------
{Name: "database_dump", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "db_malicious_event", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_malicious_function", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_malicious_procedure", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "admin_cross_account_overlap", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonAccountAggregate},
{Name: "credential_reuse", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "supply_chain_vuln", Category: CategoryWeb, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "db_magic_token_user", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_malicious_trigger", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_options_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "db_options_new_external_script", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_options_plugin_notice_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_post_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "db_content_scan_incomplete", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "db_unexpected_event", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "db_unexpected_function", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "db_unexpected_procedure", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "db_unexpected_trigger", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "drupal_admin_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "drupal_content_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "drupal_settings_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "joomla_admin_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "joomla_content_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "joomla_extensions_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "magento_admin_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "magento_content_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "magento_settings_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "opencart_admin_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "opencart_content_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "opencart_settings_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_phantom_post_author", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_post_volume_burst", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_hidden_link_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_hostname_keyed_option", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_doorway_sitemap_routes", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_spam_taxonomy", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_stored_code_execution", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_stored_cloak_logic", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_rogue_admin", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "db_siteurl_hijack", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "db_siteurl_foreign_host", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_siteurl_invalid", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "db_spam_cleaned", Category: CategoryDatabase, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "db_spam_found", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
{Name: "db_spam_injection", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "db_suspicious_admin_email", Category: CategoryDatabase, Correlation: CorrelationSecurityEvent},
// --- Email & Phishing -----------------------------------------------
{Name: "credential_log_realtime", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_auth_failure_realtime", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "email_cloud_relay_abuse", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{Block: BlockAlways}},
{Name: "email_av_degraded", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_encrypted_archive", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_scanner_panic", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "realtime_scanner_panic", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_hold_bypass", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_late_verdict", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_parse_error", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_queue_overflow", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_quarantine_error", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_scan_error", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_av_timeout", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_compromised_account", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{Block: BlockAlways}},
{Name: "email_credential_leak", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_dkim_failure", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "email_filter_blackhole", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_filter_exfil", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_filter_forwarder", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_filter_pipe", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_mail_filters", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_malware", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "email_phishing_content", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "email_php_relay_abuse", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_php_relay_account_volume_capped", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_action_dry_run", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "email_php_relay_action_failed", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "email_php_relay_action_skipped", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "email_php_relay_cpanel_limit_unreadable", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_disabled", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_inotify_overflow", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_inotify_overflow_recovered", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_msgindex_persist_failed", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_no_exim", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_overflow_scan_truncated", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_path2b_disabled", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_policies_reload", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_rate_limit_hit", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "email_php_relay_sweep_failed", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_php_relay_watcher_failed", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "email_defer_fail_governor", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "email_pipe_forwarder", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_rate_critical", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_rate_warning", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_spam_outbreak", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "email_spf_rejection", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "email_suspicious_forwarder", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_suspicious_geo", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "email_weak_password", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "email_password_audit_incomplete", Category: CategoryEmail, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "exim_frozen_realtime", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "mail_per_account", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, CorrelationGap: gapEnvelopeSender},
{Name: "mail_queue", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "mail_queue_unavailable", Category: CategoryEmail, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "phishing_credential_log", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "phishing_directory", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "phishing_iframe", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "phishing_kit_archive", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "phishing_kit_realtime", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "phishing_page", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "phishing_php", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "phishing_realtime", Category: CategoryEmail, Correlation: CorrelationSecurityEvent},
{Name: "phishing_redirector", Category: CategoryEmail, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
// --- Performance -----------------------------------------------------
{Name: "perf_error_logs", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_load", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_memory", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_mysql_config", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_php_handler", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_php_processes", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_redis_config", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_wp_config", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_wp_cron", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
{Name: "perf_wp_transients", Category: CategoryPerformance, Correlation: CorrelationIgnored, CorrelationReason: reasonPerformance},
// --- Network & Firewall ---------------------------------------------
{Name: "backdoor_port", Category: CategoryNetwork, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "backdoor_port_outbound", Category: CategoryNetwork, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "c2_connection", Category: CategoryNetwork, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "direct_smtp_egress", Category: CategoryNetwork, Correlation: CorrelationSecurityEvent},
{Name: "dns_connection", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "dns_zone_change", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "exfiltration_paste_site", Category: CategoryNetwork, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "firewall", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "firewall_ports", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "firewall_ipv6_unmanaged", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "bad_asn_outbound", Category: CategoryNetwork, Correlation: CorrelationSecurityEvent},
{Name: "infra_ips_unresolvable", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
// Reputation on the HTTP path gives a browser one verifier before blocking.
// Critical sightings come from browserless vectors and bypass the gate.
{Name: "ip_reputation", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, ChallengeFirst: true}},
{Name: "reputation_quota_exhausted", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "threat_feed_stale", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "ssl_cert_issued", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "user_outbound_connection", Category: CategoryNetwork, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
// --- System Integrity ------------------------------------------------
{Name: "af_alg_enforcement_corrected", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "af_alg_socket_use", Category: CategorySystem, Correlation: CorrelationSecurityEvent},
{Name: "account_scan_error", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "account_scan_truncated", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "bpf_unavailable", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "bpf_ringbuf_error", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "check_panic", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "full_scan_file_too_large", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "crond_change", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "crontab_change", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
{Name: "dpkg_integrity", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "kernel_module", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "mysql_superuser", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "rpm_integrity", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "sensitive_file_modified", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
{Name: "signature_update_rollback", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "signature_update_rescan_queued", Category: CategorySystem, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "suspicious_crontab", Category: CategorySystem, Correlation: CorrelationSecurityEvent, Response: ResponsePolicy{NeverChallenge: true}},
// --- WAF & ModSecurity ----------------------------------------------
{Name: "modsec_block_escalation", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways}},
{Name: "modsec_block_realtime", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "modsec_classifier_gap", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "modsec_csm_block_escalation", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "modsec_low_confidence_burst", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
{Name: "modsec_warning_realtime", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide},
// The WAF already denied repeated attacks; keep containment direct.
{Name: "waf_attack_blocked", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, NeverChallenge: true}},
{Name: "modsec_disabled_vhost", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
// Retired in favour of modsec_disabled_vhost. Kept registered so the
// waf_status runner can still purge findings written by older versions.
{Name: "waf_bypass", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "waf_detection_only", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "waf_rules", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "waf_rules_stale", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
{Name: "waf_status", Category: CategoryWAF, Correlation: CorrelationIgnored, CorrelationReason: reasonPosture},
// --- Correlation & Health -------------------------------------------
{Name: "account_scan", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "auto_block", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "auto_response", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "auto_response_paused", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "challenge_route", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonResponse},
{Name: "check_timeout", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "config_reload_error", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "config_reload_restart_required", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "yara_forge_rollback", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "coordinated_attack", Category: CategoryCorrelation, Correlation: CorrelationDerived, Response: ResponsePolicy{NeverChallenge: true}},
{Name: "csm_health", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "fanotify_kernel_overflow", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "fanotify_overflow", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "integrity", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonHostScope},
// Aggregate suspicion gets a browser verifier before a hard block.
{Name: "local_threat_score", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonAttackerSide, Response: ResponsePolicy{Block: BlockAlways, ChallengeFirst: true}},
{Name: "mail_auth_backend_degraded", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "mail_log_source_unavailable", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "protection_queue_degraded", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
{Name: "protection_queue_recovered", Category: CategoryCorrelation, Correlation: CorrelationIgnored, CorrelationReason: reasonSelfHealth},
// --- Internal (not shown in user-facing dropdowns) -------------------
{Name: "test_alert", Category: CategoryInternal, Internal: true, Correlation: CorrelationIgnored, CorrelationReason: reasonInformational},
}
// AllCheckNames returns every registered Check name, sorted alphabetically.
// Includes internal names; callers that render user-facing UI should use
// PublicCheckInfos instead.
func AllCheckNames() []string {
out := make([]string, 0, len(checkRegistry))
for _, c := range checkRegistry {
out = append(out, c.Name)
}
sort.Strings(out)
return out
}
// PublicCheckInfos returns all non-Internal checks grouped by category in
// the canonical category order (see checkCategoryOrder). Within a category
// names are sorted alphabetically. This is the list the settings UI shows
// for alerts.email.disabled_checks.
func PublicCheckInfos() []CheckInfo {
byCategory := make(map[string][]CheckInfo, len(checkCategoryOrder))
for _, c := range checkRegistry {
if c.Internal {
continue
}
byCategory[c.Category] = append(byCategory[c.Category], c)
}
var out []CheckInfo
for _, cat := range checkCategoryOrder {
items := byCategory[cat]
sort.Slice(items, func(i, j int) bool { return items[i].Name < items[j].Name })
out = append(out, items...)
}
return out
}
// LookupCheck returns the registry entry for name, if any.
func LookupCheck(name string) (CheckInfo, bool) {
for _, c := range checkRegistry {
if c.Name == name {
return c, true
}
}
return CheckInfo{}, false
}
package checks
import (
"context"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
"syscall"
"github.com/pidginhost/csm/internal/actionlog"
)
// quarantineMoveChecks are the findings whose manual fix, and whose
// full-scan quarantine, is a plain move of the named file into quarantine.
// The automatic responder's broader set lives in autoQuarantineChecks; the
// full-scan set must stay equal to this one (pinned by test).
var quarantineMoveChecks = map[string]bool{
"webshell": true,
"new_webshell_file": true,
"obfuscated_php": true,
"suspicious_php_content": true,
"new_php_in_languages": true,
"new_php_in_upgrade": true,
"phishing_page": true,
"phishing_directory": true,
}
// eximMsgIDRegex validates Exim message ID format. Exim 4.96 and older use
// 6-6-2 ids; Exim 4.97 and newer use 6-11-4 ids.
var eximMsgIDRegex = regexp.MustCompile(`^[0-9A-Za-z]{6}-(?:[0-9A-Za-z]{6}-[0-9A-Za-z]{2}|[0-9A-Za-z]{11}-[0-9A-Za-z]{4})$`)
// Allowed roots for each fix action. Declared as vars (not consts) so tests
// can redirect remediation under t.TempDir() without writing to real /home,
// /tmp, or /var/spool. Production must not mutate these at runtime.
// A nil list means "the platform's account roots" (plus, for quarantine,
// the temp trees in quarantineExtraRoots); see effectiveFixRoots.
var (
fixPermissionsAllowedRoots []string
fixQuarantineAllowedRoots []string
fixHtaccessAllowedRoots []string
eximSpoolDirs = []string{"/var/spool/exim/input", "/var/spool/exim4/input"}
)
// chmodFunc performs the permission change for fixPermissions. It is a var so
// tests can simulate failures (e.g. a read-only mount returning EROFS) without
// an actual read-only filesystem, and assert an already-compliant file is
// never chmodded.
var chmodFunc = os.Chmod
// RemediationResult describes the outcome of a fix action.
type RemediationResult struct {
Success bool `json:"success"`
Action string `json:"action"` // human-readable description of what was done
Description string `json:"description"` // what fix was applied
Error string `json:"error,omitempty"`
// Refused distinguishes an unchanged, ineligible target from an I/O failure.
Refused bool `json:"-"`
// RemediationStatus lets a caller that supports more than one successful
// disposition distinguish an in-place clean from whole-file quarantine.
// It is transport metadata, not part of the generic remediation API.
RemediationStatus string `json:"-"`
// Reverted marks a virtual patch that had to be written again because
// something removed or damaged CSM's earlier block -- typically a backup
// plugin rewriting the .htaccess it owns.
Reverted bool `json:"reverted,omitempty"`
}
// FixDescription returns a human-readable description of what the fix will do
// for a given check type and file path. Returns empty string if no fix is available.
func FixDescription(checkType, message string, filePath ...string) string {
path := selectFindingPath(message, filePath...)
if isHtaccessHardenedFinding(checkType) {
if path != "" {
return fmt.Sprintf("Remove malicious directives from %s", path)
}
return ""
}
if quarantineMoveChecks[checkType] {
if path != "" {
return fmt.Sprintf("Quarantine %s to /opt/csm/quarantine/", path)
}
return ""
}
switch checkType {
case "world_writable_php", "group_writable_php":
if path != "" {
return fmt.Sprintf("Set permissions to 644 on %s", path)
}
case "backdoor_binary", "new_executable_in_config":
if path != "" {
return fmt.Sprintf("Kill process and quarantine %s", path)
}
case "suspicious_crontab":
if path != "" {
return fmt.Sprintf("Quarantine and truncate crontab %s", path)
}
return "Quarantine and truncate crontab"
case "htaccess_injection", "htaccess_handler_abuse":
if path != "" {
return fmt.Sprintf("Remove malicious directives from %s", path)
}
case "email_phishing_content":
msgID := extractEximMsgID(message)
if msgID != "" {
return fmt.Sprintf("Quarantine Exim spool message %s", msgID)
}
}
return ""
}
// HasFix returns true if the check type has a known automated fix.
func HasFix(checkType string) bool {
if isHtaccessHardenedFinding(checkType) || quarantineMoveChecks[checkType] {
return true
}
fixableChecks := map[string]bool{
"world_writable_php": true,
"group_writable_php": true,
"backdoor_binary": true,
"new_executable_in_config": true,
"htaccess_injection": true,
"htaccess_handler_abuse": true,
"email_phishing_content": true,
"suspicious_crontab": true,
}
return fixableChecks[checkType]
}
// ApplyFix executes the remediation action for a finding.
func ApplyFix(ctx context.Context, checkType, message, details string, filePath ...string) RemediationResult {
if err := ctx.Err(); err != nil {
return RemediationResult{Error: err.Error()}
}
path := selectFindingPath(message, filePath...)
if isHtaccessHardenedFinding(checkType) {
// CleanHtaccessFile re-runs the full detector registry, so a single
// action removes every malicious directive the audit found.
return CleanHtaccessFile(path)
}
if quarantineMoveChecks[checkType] {
return fixQuarantine(path)
}
switch checkType {
case "world_writable_php", "group_writable_php":
return fixPermissions(path, checkType)
case "backdoor_binary", "new_executable_in_config":
return fixKillAndQuarantine(ctx, path, details)
case "htaccess_injection", "htaccess_handler_abuse":
return fixHtaccess(path, message)
case "email_phishing_content":
return fixQuarantineSpoolMessage(message)
case "suspicious_crontab":
return fixSuspiciousCrontab(path)
default:
return RemediationResult{Error: fmt.Sprintf("no automated fix available for check type '%s'", checkType)}
}
}
// fixPermissions sets file permissions to 0644. checkType selects which write
// bit is the dangerous one: world-writable (0002) for world_writable_php,
// group-writable (0020) for group_writable_php.
//
// If the file no longer carries that bit -- because an operator already fixed
// it by hand, or it changed since the scan -- the finding is treated as
// already resolved and no chmod is attempted. This is the path an operator
// hits when they manually correct perms and then click "Apply automated fix":
// rather than erroring, the finding clears.
func fixPermissions(path, checkType string) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
path, info, err := resolveExistingFixPath(path, effectiveFixRoots(fixPermissionsAllowedRoots))
if err != nil {
return RemediationResult{Error: err.Error()}
}
oldMode := info.Mode().Perm()
dangerBit, label := os.FileMode(0002), "world-writable"
if checkType == "group_writable_php" {
dangerBit, label = 0020, "group-writable"
}
if oldMode&dangerBit == 0 {
return RemediationResult{
Success: true,
Action: fmt.Sprintf("verified %s: no longer %s (mode %o)", path, label, oldMode),
Description: fmt.Sprintf("File is already not %s; no change needed", label),
}
}
// #nosec G302 -- Intentional: this is the remediation that sets the
// canonical "safe web content" mode on a user file after we flagged
// the file as having dangerous perms (e.g. 0777). 0644 is what the
// webserver needs to serve static content as the file owner.
if err := chmodFunc(path, 0644); err != nil {
if errors.Is(err, syscall.EROFS) {
return RemediationResult{Error: fmt.Sprintf(
"cannot fix %s: the file is on a read-only mount (e.g. a backup snapshot or bind mount), not the live site. Dismiss or suppress this finding instead.",
path)}
}
return RemediationResult{Error: fmt.Sprintf("chmod failed: %v", err)}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("chmod 644 %s", path),
Description: fmt.Sprintf("Changed permissions from %o to 644", oldMode),
}
}
// fixQuarantine moves a file or directory to quarantine.
func fixQuarantine(path string) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
path, info, err := resolveExistingFixPath(path, effectiveFixRoots(fixQuarantineAllowedRoots, quarantineExtraRoots...))
if err != nil {
return RemediationResult{Error: err.Error()}
}
return quarantineResolvedTarget(path, info)
}
// quarantineResolvedTarget quarantines the exact object admitted by the
// caller's boundary check. Regular-file quarantine reopens the path and
// verifies this identity before copying or unlinking it.
func quarantineResolvedTarget(path string, info os.FileInfo) RemediationResult {
qPath := newQuarantinePath(quarantineDir, path)
var quarantineWarning string
meta := quarantineMetadata(path, info, "Fixed via CSM Web UI")
if err := quarantineTarget(path, qPath, info, meta); err != nil {
var completed bool
quarantineWarning, completed = completedQuarantineWarning(err)
if !completed {
return RemediationResult{Error: err.Error()}
}
}
description := fmt.Sprintf("Moved to quarantine: %s", qPath)
if quarantineWarning != "" {
description += ". Warning: " + quarantineWarning
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("quarantined %s -> %s", path, qPath),
Description: description,
}
}
// fixKillAndQuarantine kills any process using the file, then quarantines it.
func fixKillAndQuarantine(ctx context.Context, path, details string) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
resolvedPath, target, err := resolveExistingFixPath(path, effectiveFixRoots(fixQuarantineAllowedRoots, quarantineExtraRoots...))
if err != nil {
return RemediationResult{Error: err.Error()}
}
path = resolvedPath
// Try to extract and kill PID from details
pid := extractPID(details)
killed := false
var signalErr error
if pidInt, ok := parseProcessPID(pid); ok {
pid = strconv.Itoa(pidInt)
signalErr = signalProcess(ctx, pidInt, syscall.SIGKILL, func() error {
uid := getProcessUID(pid)
if uid == "0" || uid == "" || !processUsesFileIdentity(pidInt, target) {
return errProcessNotEligible
}
return nil
})
recordKillAction(nil, pid, path, signalErr)
killed = signalErr == nil
if errors.Is(signalErr, errProcessNotEligible) || errors.Is(signalErr, os.ErrProcessDone) {
signalErr = nil
}
}
if err := ctx.Err(); err != nil && !killed {
return RemediationResult{Error: err.Error()}
}
// Quarantine the same object used for the process decision. If the path was
// replaced after validation, the pinned-identity quarantine refuses it.
result := quarantineResolvedTarget(path, target)
if signalErr != nil {
result.Success = false
if result.Error != "" {
result.Error += "; "
}
result.Error += "process was not stopped: " + signalErr.Error()
}
if killed {
if result.Success {
result.Action = fmt.Sprintf("killed PID %s and %s", pid, result.Action)
result.Description = "Process killed and file quarantined"
} else {
result.Action = fmt.Sprintf("killed PID %s; quarantine failed", pid)
result.Description = "Process killed, but the file was not quarantined"
}
}
return result
}
// fixHtaccess removes malicious directives from an .htaccess file while
// preserving comments and known-safe directives (e.g., Wordfence, LiteSpeed).
func fixHtaccess(path, message string) (result RemediationResult) {
audit := newCleanAction(path)
defer func() { audit.finish(result.Error) }()
if path == "" {
return RemediationResult{Error: "could not extract file path"}
}
if filepath.Base(path) != ".htaccess" {
return RemediationResult{Error: "automated .htaccess remediation only applies to .htaccess files"}
}
path, _, err := resolveExistingFixPath(path, effectiveFixRoots(fixHtaccessAllowedRoots))
if err != nil {
return RemediationResult{Error: err.Error()}
}
// Same pinned-inode read and atomic replace as CleanHtaccessFile: the
// directory belongs to the account, so nothing here may follow a path
// the owner can redirect between the check and the write.
target, err := openCleanTarget(path)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot open: %v", err)}
}
defer target.Close()
audit.rec.Result = actionlog.Failed
data, err := io.ReadAll(target.File)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot read: %v", err)}
}
audit.capture(target, data)
audit.rec.Result = actionlog.Refused
dangerous := []string{"auto_prepend_file", "auto_append_file", "eval(", "base64_decode",
"gzinflate", "str_rot13", "addhandler", "sethandler"}
safe := []string{
"wordfence-waf.php", "litespeed", "advanced-headers.php", "rsssl",
"application/x-httpd-php", "application/x-httpd-ea-php", "application/x-httpd-alt-php",
"-execcgi", "sethandler none", "sethandler default-handler",
"text/html", "text/css", "text/javascript", "application/javascript",
"image/", "font/", ".woff", ".woff2", ".ttf", ".eot", ".svg",
"wordfence",
}
var cleaned []string
removed := 0
var phpHandlerContexts []phpHandlerOverlay
// Iterate logical directives so a malicious mapping split across an Apache
// line continuation is removed as a unit (every physical line it spans).
for _, logical := range joinHtaccessContinuations(strings.Split(string(data), "\n")) {
trimmed := strings.TrimSpace(logical.text)
lineLower := strings.ToLower(trimmed)
if strings.HasPrefix(trimmed, "#") {
cleaned = append(cleaned, logical.lines...)
continue
}
if ctx, ok := openPHPHandlerContext(trimmed); ok {
phpHandlerContexts = append(phpHandlerContexts, ctx)
cleaned = append(cleaned, logical.lines...)
continue
}
if closesPHPHandlerContext(trimmed) {
if len(phpHandlerContexts) > 0 {
phpHandlerContexts = phpHandlerContexts[:len(phpHandlerContexts)-1]
}
cleaned = append(cleaned, logical.lines...)
continue
}
isDangerous := false
if phpHandlerRemapsNonPHPInContext(lineLower, phpHandlerContexts) {
isDangerous = true
}
for _, d := range dangerous {
if strings.Contains(lineLower, d) {
isSafe := false
for _, s := range safe {
if strings.Contains(lineLower, s) {
isSafe = true
break
}
}
if !isSafe {
isDangerous = true
break
}
}
}
if isDangerous {
removed++
} else {
cleaned = append(cleaned, logical.lines...)
}
}
if removed == 0 {
return RemediationResult{Error: "no malicious directives found to remove"}
}
backupPath := newQuarantinePath(htaccessBackupDirRoot, path)
meta := quarantineMetadata(path, target.Info, "Pre-clean .htaccess backup")
audit.rec.Result = actionlog.Failed
audit.rec.Reason = meta.Reason
if err := storeQuarantineBackup(backupPath, data, meta, 0600); err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot create durable backup: %v", err)}
}
if err := audit.replace(target, []byte(strings.Join(cleaned, "\n")), backupPath); err != nil {
return RemediationResult{Error: fmt.Sprintf("write failed; backup retained at %s: %v", backupPath, err)}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("removed %d malicious directive(s) from %s", removed, path),
Description: fmt.Sprintf("Cleaned .htaccess: removed %d line(s) (backup: %s)", removed, backupPath),
}
}
// extractFilePathFromMessage extracts a file path from a finding message.
// Handles patterns like "World-writable PHP file: /path/to/file"
// and "Webshell found: /path/to/file"
func extractFilePathFromMessage(message string) string {
// Look for /home/ or /tmp/ paths
for _, prefix := range accountRootPrefixes("/tmp/", "/dev/shm/", "/var/tmp/") {
idx := strings.Index(message, prefix)
if idx < 0 {
continue
}
rest := message[idx:]
// Path ends at space, comma, newline, or end
endIdx := len(rest)
for i, c := range rest {
if c == ' ' || c == ',' || c == '\n' || c == ')' {
endIdx = i
break
}
}
return rest[:endIdx]
}
return ""
}
func selectFindingPath(message string, filePath ...string) string {
if len(filePath) > 0 {
path := filePath[0]
if strings.TrimSpace(path) != "" {
return path
}
}
return extractFilePathFromMessage(message)
}
func resolveExistingFixPath(path string, allowedRoots []string) (string, os.FileInfo, error) {
cleanPath, err := sanitizeFixPath(path, allowedRoots)
if err != nil {
return "", nil, err
}
info, err := osFS.Lstat(cleanPath)
if err != nil {
return "", nil, fmt.Errorf("file not found: %w", err)
}
if info.Mode()&os.ModeSymlink != 0 {
return "", nil, refuseFileResponse(fmt.Errorf("symlinked paths are not eligible for automated remediation: %s", cleanPath))
}
resolved, err := filepath.EvalSymlinks(cleanPath)
if err != nil {
return "", nil, fmt.Errorf("cannot resolve path: %w", err)
}
resolved, err = sanitizeFixPath(resolved, allowedRoots)
if err != nil {
return "", nil, err
}
if accountRoot := homeAccountRoot(cleanPath); accountRoot != "" && !isPathWithinOrEqual(resolved, accountRoot) {
return "", nil, refuseFileResponse(fmt.Errorf("resolved path escapes account boundary: %s", resolved))
}
resolvedInfo, err := osFS.Lstat(resolved)
if err != nil {
return "", nil, fmt.Errorf("file not found: %w", err)
}
if resolvedInfo.Mode()&os.ModeSymlink != 0 {
return "", nil, refuseFileResponse(fmt.Errorf("symlinked paths are not eligible for automated remediation: %s", resolved))
}
return resolved, resolvedInfo, nil
}
func sanitizeFixPath(path string, allowedRoots []string) (string, error) {
if strings.TrimSpace(path) == "" {
return "", refuseFileResponse(errors.New("file path is required"))
}
path = filepath.Clean(path)
if !filepath.IsAbs(path) {
return "", refuseFileResponse(errors.New("file path must be absolute"))
}
for _, root := range allowedRoots {
if fixTargetDepthBelow(path, root) >= fixTargetMinDepth(root) {
return path, nil
}
}
return "", refuseFileResponse(fmt.Errorf("file path is outside the allowed remediation roots: %s", path))
}
// fixTargetDepthBelow returns how many path components path lies below
// root, or 0 when path is root itself or not under it.
func fixTargetDepthBelow(path, root string) int {
cleanRoot := filepath.Clean(root)
if !strings.HasPrefix(path, cleanRoot+string(filepath.Separator)) {
return 0
}
rel := strings.TrimPrefix(path, cleanRoot+string(filepath.Separator))
return strings.Count(rel, string(filepath.Separator)) + 1
}
// fixTargetMinDepth is how far below a remediation root a target must lie.
// A root itself is never a target, and under /home neither is an account's
// home directory: quarantining or chmod-ing either takes a whole tree away.
func fixTargetMinDepth(root string) int {
if isAccountRoot(root) {
return 2
}
return 1
}
func isPathWithinOrEqual(path, base string) bool {
cleanPath := filepath.Clean(path)
cleanBase := filepath.Clean(base)
return cleanPath == cleanBase || strings.HasPrefix(cleanPath, cleanBase+string(filepath.Separator))
}
func homeAccountRoot(path string) string {
root, account, ok := accountRootOf(path)
if !ok {
return ""
}
return filepath.Join(root, account)
}
// extractEximMsgID extracts an Exim message ID from a finding message.
// Matches the pattern "(message: XXXXXX-XXXXXX-XX)" used by emailscan.go.
func extractEximMsgID(message string) string {
prefix := "(message: "
idx := strings.Index(message, prefix)
if idx < 0 {
return ""
}
rest := message[idx+len(prefix):]
end := strings.Index(rest, ")")
if end < 0 {
return ""
}
return strings.TrimSpace(rest[:end])
}
// fixQuarantineSpoolMessage moves Exim spool files (-H header and -D body)
// for a message ID into quarantine.
func fixQuarantineSpoolMessage(message string) RemediationResult {
msgID := extractEximMsgID(message)
if msgID == "" {
return RemediationResult{Error: "could not extract Exim message ID from finding"}
}
// Validate Exim message ID format to prevent path traversal
if !eximMsgIDRegex.MatchString(msgID) {
return RemediationResult{Error: fmt.Sprintf("invalid Exim message ID format: %s", msgID)}
}
var spoolDir string
for _, dir := range eximSpoolDirs {
if _, err := osFS.Stat(filepath.Join(dir, msgID+"-H")); err == nil {
spoolDir = dir
break
}
}
if spoolDir == "" {
return RemediationResult{Error: fmt.Sprintf("spool message %s not found (already delivered or removed)", msgID)}
}
base := newQuarantinePath(quarantineDir, "exim_"+msgID)
moved := 0
for _, suffix := range []string{"-H", "-D"} {
src := filepath.Join(spoolDir, msgID+suffix)
info, err := os.Lstat(src)
if os.IsNotExist(err) {
continue
}
if err != nil {
return RemediationResult{Error: fmt.Sprintf("cannot inspect spool file after quarantining %d files: %v", moved, err)}
}
meta := quarantineMetadata(src, info, "Phishing email quarantined via CSM Web UI")
meta.MessageID, meta.SpoolDir = msgID, spoolDir
dst := base + suffix
if err := quarantineTarget(src, dst, info, meta); err != nil {
return RemediationResult{Error: fmt.Sprintf("spool quarantine stopped after %d files; inspect recovery copies under %s: %v", moved, quarantineDir, err)}
}
moved++
}
if moved == 0 {
return RemediationResult{Error: fmt.Sprintf("no spool files found for message %s", msgID)}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("quarantined spool message %s (%d files)", msgID, moved),
Description: fmt.Sprintf("Exim spool files moved to quarantine for message %s", msgID),
}
}
package checks
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/url"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/eximlog"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/netutil"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/threatintel"
)
const (
reputationCacheFile = "reputation_cache.json"
reputationEximMainlog = "/var/log/exim_mainlog"
reputationWHMAccessLog = "/usr/local/cpanel/logs/access_log"
cacheExpiry = 6 * time.Hour
errorCacheExpiry = 1 * time.Hour // cache transient API errors to avoid retrying same IP
abuseConfidenceThreshold = 50
maxQueriesPerCycle = 5 // max AbuseIPDB API calls per 10-min cycle (~720/day, fits free tier)
maxCacheEntries = 5000 // cap cache size
reputationHealthReminder = 24 * time.Hour
reputationQuotaStateKey = "_reputation_health:quota"
reputationFeedStateKey = "_reputation_health:feed_stale"
)
var reputationHealthStateKeys = map[string]string{
logicalOwnerReputationQuota: reputationQuotaStateKey,
logicalOwnerReputationFeedStale: reputationFeedStateKey,
}
// maxDailyAbuseQueries is the store-backed daily circuit-breaker below
// the 1000/day free-tier ceiling. The 100-slot cushion below 1000 leaves
// room for API-side accounting differences and fallback paths that cannot
// share the store counter. Declared as a var (not const) so tests can lower it
// without burning seconds on 900 bbolt transactions. Production callers
// must not modify this.
var maxDailyAbuseQueries = 900
// abuseIPDBEndpoint is the URL queried for IP reputation. Declared as a
// var (not const) so tests can point it at an httptest server. Production
// callers must not modify this.
var abuseIPDBEndpoint = "https://api.abuseipdb.com/api/v2/check"
// abuseIPDBClient is the HTTP client used for AbuseIPDB queries. Declared
// at package scope so tests can swap in a mock client (e.g., one whose
// transport routes all traffic to an httptest server).
var abuseIPDBClient = &http.Client{Timeout: 10 * time.Second}
type reputationCache struct {
Entries map[string]*reputationEntry `json:"entries"`
// dirty tracks entries touched since load. nil means the cache was
// assembled directly rather than hydrated by loadReputationCache, in
// which case a save persists every entry.
dirty map[string]bool
// removed tracks evictions since load so bbolt saves can delete entries
// pruned from the in-memory view.
removed map[string]bool
}
type reputationEntry struct {
Score int `json:"score"`
Category string `json:"category"`
CheckedAt time.Time `json:"checked_at"`
}
// set records an entry and marks it changed so a bbolt-backed save can
// persist just this cycle's writes instead of the whole map.
func (c *reputationCache) set(ip string, e *reputationEntry) {
c.Entries[ip] = e
if c.dirty != nil {
c.dirty[ip] = true
delete(c.removed, ip)
}
}
// remove evicts an entry. Deleting any pending dirty mark keeps a later
// save from writing what eviction just removed; recording the removal lets
// bbolt delete a prior stored value for the same IP.
func (c *reputationCache) remove(ip string) {
delete(c.Entries, ip)
if c.dirty != nil {
delete(c.dirty, ip)
if c.removed != nil {
c.removed[ip] = true
}
}
}
// changedEntries returns what a bbolt-backed save must persist. Without
// change tracking every entry counts as changed. With tracking, only
// entries touched since load and still present are returned: re-putting
// the full map here used to resurrect every entry the TTL/cap prune had
// just deleted, so the bucket grew without bound.
func (c *reputationCache) changedEntries() map[string]store.ReputationEntry {
out := make(map[string]store.ReputationEntry)
if c.dirty == nil {
for ip, e := range c.Entries {
out[ip] = store.ReputationEntry{Score: e.Score, Category: e.Category, CheckedAt: e.CheckedAt}
}
return out
}
for ip := range c.dirty {
e, ok := c.Entries[ip]
if !ok {
continue
}
out[ip] = store.ReputationEntry{Score: e.Score, Category: e.Category, CheckedAt: e.CheckedAt}
}
return out
}
// CheckIPReputation looks up non-infra IPs against threat intelligence.
// Four-tier approach:
// 1. Skip if already blocked
// 2. Check local threat DB (permanent blocklist + free feeds)
// 3. Check AbuseIPDB cache
// 4. Query AbuseIPDB for truly unknown IPs (max 5/cycle, ~720/day)
func CheckIPReputation(ctx context.Context, cfg *config.Config, scanState *state.Store) []alert.Finding {
sdb := store.Global()
now := time.Now()
quotaExhausted := !abuseQuotaReady(sdb, now)
var findings []alert.Finding
supplementalAgg := newSupplementalThreatAggregator(cfg)
ips := collectRecentIPs(cfg)
if len(ips) == 0 {
return reputationHealthFindings(ctx, cfg, sdb, scanState, now, quotaExhausted)
}
authenticated := collectAuthenticatedIPs(cfg)
alreadyBlocked := loadAllBlockedIPs(cfg.StatePath)
threatDB := GetThreatDB()
cache := loadReputationCache(cfg.StatePath)
client := abuseIPDBClient
utcDay := now.UTC().Format("2006-01-02")
// Two-pass design so the slow AbuseIPDB HTTP queries can run in
// parallel:
//
// Pass 1 (serial): walk every IP and resolve via tier 1/2/3 plus
// the supplemental aggregator. Collect the IPs that genuinely
// need a tier-4 HTTP lookup into pendingQueries.
//
// Pass 2 (parallel, up to maxQueriesPerCycle workers): fan out
// queryAbuseIPDB and collect results.
//
// Pass 3 (serial): apply results back into the cache and emit
// findings.
//
// Pre-cache, all five HTTP queries ran in a serial loop, so a
// cycle paid ~5x worst-case AbuseIPDB latency. A busy production
// host saw ip_reputation averaging ~3.6 s per run because of
// this; the fan-out brings that down to ~max(single-call latency).
type pendingQuery struct {
ip string
source string
}
var pendingQueries []pendingQuery
for ip, source := range ips {
// A source that authenticated successfully holds valid credentials and is
// a real customer on a recycled dynamic/CGNAT address, not a drive-by
// scanner. Reputation auto-block keys on mere passive access, so without
// this skip a legitimate customer whose ISP-recycled IP appears in a
// public feed gets a 24h block the instant they open webmail. Genuine
// attacks from the same IP are still caught by the brute-force,
// compromise, and takeover detectors, which do not honour this exemption.
if authenticated[ip] {
continue
}
// Tier 1: Skip if already blocked
if alreadyBlocked[ip] {
continue
}
// Tier 2: Check local threat DB
if threatDB != nil {
if dbSource, found := threatDB.Lookup(ip); found {
findings = append(findings, alert.Finding{
Severity: reputationSightingSeverity(source),
Check: "ip_reputation",
Message: fmt.Sprintf(reputationMessagePrefix+"%s (source: %s)", ip, dbSource),
Details: fmt.Sprintf("Detected via: %s\nMatched in local threat intelligence database", source),
Timestamp: time.Now(),
SourceIP: ip,
})
continue
}
}
// Tier 3: Check AbuseIPDB cache. Treat entries with CheckedAt in
// the future (legacy data written by a prior buggy error-caching
// formula) as expired so they get re-queried or aged out.
if entry, ok := cache.Entries[ip]; ok {
age := time.Since(entry.CheckedAt)
if age >= 0 && age < cacheExpiry {
if entry.Score >= abuseConfidenceThreshold {
appendReputationFinding(&findings, ip, source, "AbuseIPDB", entry.Score, entry.Category)
} else if score, src, ok := supplementalThreatScore(ctx, supplementalAgg, ip); ok && score >= abuseConfidenceThreshold {
appendReputationFinding(&findings, ip, source, src, score, strings.ToLower(src)+" history")
}
continue
}
}
// Tier 4 candidate; defer the HTTP call to pass 2 unless the
// quota / config gates already preclude querying.
if cfg.Reputation.AbuseIPDBKey == "" || quotaExhausted || len(pendingQueries) >= maxQueriesPerCycle {
if score, src, ok := supplementalThreatScore(ctx, supplementalAgg, ip); ok && score >= abuseConfidenceThreshold {
appendReputationFinding(&findings, ip, source, src, score, strings.ToLower(src)+" history")
}
continue
}
pendingQueries = append(pendingQueries, pendingQuery{ip: ip, source: source})
}
// Pass 2: reserve daily quota slots up front and fan out the HTTP
// calls. The pre-reservation matches the prior "count the attempt
// before the call so a crash or network hang still consumes a slot"
// guarantee, while keeping near-cap cycles from spending more slots
// than the store can reserve.
type queryResult struct {
score int
category string
err error
}
var refusedQueries []pendingQuery
if len(pendingQueries) > 0 && sdb != nil {
reserved := sdb.ReserveAbuseQuerySlots(utcDay, len(pendingQueries), maxDailyAbuseQueries)
refusedQueries = pendingQueries[reserved:]
pendingQueries = pendingQueries[:reserved]
}
batch := &reputationQueryBatch{}
queryWork := make([]*reputationQueryWork, len(pendingQueries))
for i := range queryWork {
queryWork[i] = reputationQueries.begin(batch)
}
defer func() {
for _, work := range queryWork {
work.finish(false)
}
}()
// Reserved work already exists while the refused tail is scored. Publish
// it first so a stalled or abandoned fallback cannot hide those queries.
for _, q := range refusedQueries {
if supplemental, src, ok := supplementalThreatScore(ctx, supplementalAgg, q.ip); ok && supplemental >= abuseConfidenceThreshold {
appendReputationFinding(&findings, q.ip, q.source, src, supplemental, strings.ToLower(src)+" history")
}
}
results := make(map[string]queryResult, len(pendingQueries))
quotaErrorObserved := false
if len(pendingQueries) > 0 {
var mu sync.Mutex
var wg sync.WaitGroup
workers := len(pendingQueries)
if workers > maxQueriesPerCycle {
workers = maxQueriesPerCycle
}
jobs := make(chan int, len(pendingQueries))
for i := range pendingQueries {
jobs <- i
}
close(jobs)
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for index := range jobs {
q, work := pendingQueries[index], queryWork[index]
work.start(client.Timeout)
func() {
completed := false
defer func() { work.returned(completed) }()
score, category, err := queryAbuseIPDB(client, q.ip, cfg.Reputation.AbuseIPDBKey, work)
if err != nil && !abuseQuotaError(err) {
work.fail()
}
mu.Lock()
results[q.ip] = queryResult{score: score, category: category, err: err}
mu.Unlock()
completed = true
}()
}
}()
}
wg.Wait()
}
// Pass 3: apply each tier-4 result back into cache + findings.
// Serial so cache writes and quota-exhaustion handling stay
// consistent regardless of which worker observed which HTTP error.
var cacheWork []*reputationQueryWork
for index, q := range pendingQueries {
work := queryWork[index]
work.phase(true)
res, ok := results[q.ip]
if !ok {
work.finish(false)
continue
}
if res.err != nil {
if abuseQuotaError(res.err) {
quotaErrorObserved = true
resetAt := nextUTCMidnight(time.Now())
fmt.Fprintf(os.Stderr, "abuseipdb: quota exhausted (%v), pausing lookups until %s\n",
res.err, resetAt.Format(time.RFC3339))
// Persisted backoff is the load-bearing signal for the
// next cycle's classifier; the cycle-local quotaExhausted
// flag has no further reader past this loop.
if sdb != nil {
if err := sdb.SetAbuseQuotaExhaustedUntil(resetAt); err != nil {
work.fail()
fmt.Fprintf(os.Stderr, "abuseipdb: persisting quota backoff: %v\n", err)
}
}
if supplemental, src, ok := supplementalResultThreatScore(ctx, cfg, supplementalAgg, q.ip, work); ok && supplemental >= abuseConfidenceThreshold {
appendReputationFinding(&findings, q.ip, q.source, src, supplemental, strings.ToLower(src)+" history")
}
work.finish(true)
continue
}
cache.set(q.ip, &reputationEntry{
Score: -1,
Category: fmt.Sprintf("error: %v", res.err),
// CheckedAt is shifted into the past so time.Since returns
// ~(cacheExpiry-errorCacheExpiry) immediately; the Tier-3
// freshness check then flips false after a further
// errorCacheExpiry, giving a real ~1h TTL on error entries.
CheckedAt: time.Now().Add(-(cacheExpiry - errorCacheExpiry)),
})
if supplemental, src, ok := supplementalResultThreatScore(ctx, cfg, supplementalAgg, q.ip, work); ok && supplemental >= abuseConfidenceThreshold {
appendReputationFinding(&findings, q.ip, q.source, src, supplemental, strings.ToLower(src)+" history")
}
work.phase(false)
cacheWork = append(cacheWork, work)
continue
}
cache.set(q.ip, &reputationEntry{
Score: res.score,
Category: res.category,
CheckedAt: time.Now(),
})
score := res.score
category := res.category
provider := "AbuseIPDB"
if supplemental, src, ok := supplementalResultThreatScore(ctx, cfg, supplementalAgg, q.ip, work); ok && supplemental > score {
score = supplemental
category = strings.ToLower(src) + " history"
provider = src
}
if score >= abuseConfidenceThreshold {
appendReputationFinding(&findings, q.ip, q.source, provider, score, category)
}
work.phase(false)
cacheWork = append(cacheWork, work)
}
// All results awaiting the shared cache commit retain their own owner.
for _, work := range cacheWork {
work.phase(true)
}
cleanCache(cache)
saved := saveReputationCache(cfg.StatePath, cache)
for _, work := range cacheWork {
work.finish(saved)
}
healthNow := time.Now()
quotaExhausted = quotaErrorObserved || !abuseQuotaReady(sdb, healthNow)
findings = append(findings, reputationHealthFindings(ctx, cfg, sdb, scanState, healthNow, quotaExhausted)...)
return findings
}
func newSupplementalThreatAggregator(cfg *config.Config) *threatintel.Aggregator {
if !cfg.Reputation.Upstream.Enabled {
threatintel.ClearUpstreamMetricsSource()
}
if !cfg.Reputation.Rspamd.Enabled && !cfg.Reputation.Upstream.Enabled {
return nil
}
agg := threatintel.NewAggregator()
if cfg.Reputation.Rspamd.Enabled {
agg.Register(threatintel.NewRspamdSource(
cfg.Reputation.Rspamd.URL,
cfg.Reputation.Rspamd.Token,
cfg.Reputation.Rspamd.TokenEnv,
))
}
if cfg.Reputation.Upstream.Enabled {
upstream := threatintel.NewUpstreamSource(threatintel.UpstreamConfig{
URL: cfg.Reputation.Upstream.URL,
Token: cfg.Reputation.Upstream.Token,
TokenEnv: cfg.Reputation.Upstream.TokenEnv,
CacheTTL: time.Duration(cfg.Reputation.Upstream.CacheTTLMin) * time.Minute,
Timeout: time.Duration(cfg.Reputation.Upstream.TimeoutSec) * time.Second,
})
threatintel.RegisterUpstreamMetrics(metrics.Default(), upstream)
agg.Register(upstream)
}
return agg
}
func supplementalResultThreatScore(ctx context.Context, cfg *config.Config, agg *threatintel.Aggregator, ip string, work *reputationQueryWork) (int, string, bool) {
if agg == nil {
return 0, "", false
}
// The aggregator runs its sources serially. Keep their configured HTTP
// budget separate from cache writes and other local result handling.
var budget time.Duration
if cfg.Reputation.Rspamd.Enabled {
budget += 5 * time.Second
}
if cfg.Reputation.Upstream.Enabled {
timeout := cfg.Reputation.Upstream.TimeoutSec
if timeout == 0 {
timeout = 5
}
budget += time.Duration(timeout) * time.Second
}
work.consuming(ctx, budget)
defer work.phase(true)
return supplementalThreatScore(ctx, agg, ip)
}
// supplementalThreatScore queries the aggregator for ip and returns the
// aggregated score, the name of the highest-scoring individual source
// (capitalised for operator-facing messages), and whether a usable score
// was found. Returns ("", 0, false) when agg is nil or no source scored.
func supplementalThreatScore(ctx context.Context, agg *threatintel.Aggregator, ip string) (int, string, bool) {
if agg == nil {
return 0, "", false
}
res, err := agg.Score(ctx, ip)
if err != nil || res.AggregatedScore == 0 {
return 0, "", false
}
// Identify the source with the highest individual score so callers can
// label findings accurately (e.g. "Rspamd" vs "Upstream").
dominant := "supplemental"
max := 0
for name, s := range res.Sources {
if s > max {
max = s
dominant = name
}
}
return res.AggregatedScore, capitalizeProvider(dominant), true
}
// capitalizeProvider title-cases known source names for operator-facing
// messages ("rspamd" -> "Rspamd", "upstream" -> "Upstream").
func capitalizeProvider(name string) string {
if len(name) == 0 {
return name
}
return strings.ToUpper(name[:1]) + name[1:]
}
// reputationSightingSeverity grades a threat-intel sighting by what the IP
// was doing. Auth-surface contact (SSH, mail credential attacks) is an
// active threat and stays Critical; a passive web sighting of a listed IP
// is ambient scanner noise, downgraded to High so thousands of drive-by
// scanners do not drown compromise-class Criticals. Auto-block eligibility
// is keyed on the check name and is not affected by this severity. Unknown
// future surfaces fail closed to Critical.
func reputationSightingSeverity(detectedVia string) alert.Severity {
switch detectedVia {
case "HTTP request", "cPanel/WHM access":
return alert.High
default:
return alert.Critical
}
}
// reputationMessagePrefix opens every ip_reputation message; the address
// follows it up to " (". ReputationMessageSourceIP depends on that form.
const reputationMessagePrefix = "Known malicious IP accessing server: "
func appendReputationFinding(findings *[]alert.Finding, ip, detectedVia, provider string, score int, category string) {
*findings = append(*findings, alert.Finding{
Severity: reputationSightingSeverity(detectedVia),
Check: "ip_reputation",
Message: fmt.Sprintf(reputationMessagePrefix+"%s (%s score: %d/100)", ip, provider, score),
Details: fmt.Sprintf("Detected via: %s\nCategory: %s\nThis IP is reported in threat intelligence databases", detectedVia, category),
Timestamp: time.Now(),
SourceIP: ip,
})
}
// reputationHealthFindings surfaces degraded reputation coverage: an
// exhausted AbuseIPDB quota (tier-4 lookups paused) and threat feeds that
// have not refreshed in over a week (stale tier-2 data). Feeds that never
// loaded stay silent - a fresh install's first download may still be
// pending, which is not operator-actionable.
func reputationHealthFindings(ctx context.Context, cfg *config.Config, sdb *store.DB, scanState *state.Store, now time.Time, quotaExhausted bool) []alert.Finding {
var out []alert.Finding
disabled := disabledLogicalOwners(cfg)
quotaActive := cfg.Reputation.AbuseIPDBKey != "" && quotaExhausted
if _, off := disabled[logicalOwnerReputationQuota]; off {
clearReputationHealthState(scanState, reputationQuotaStateKey)
} else if quotaActive {
detail := "Daily AbuseIPDB query budget reached; uncached AbuseIPDB lookups resume at 00:00 UTC. Local threat DB, rspamd, upstream, and cached scores still apply when configured."
if sdb != nil {
if until := sdb.AbuseQuotaExhaustedUntil(); !until.IsZero() && now.Before(until) {
detail = fmt.Sprintf("API returned a quota error; uncached AbuseIPDB lookups are paused until %s. Local threat DB, rspamd, upstream, and cached scores still apply when configured.", until.UTC().Format(time.RFC3339))
}
}
finding := alert.Finding{
Severity: alert.Warning,
Check: "reputation_quota_exhausted",
Message: "AbuseIPDB quota exhausted: uncached AbuseIPDB lookups are paused",
Details: detail,
Timestamp: now,
}
appendReputationHealthFinding(ctx, &out, scanState, logicalOwnerReputationQuota, reputationQuotaStateKey, now, finding)
} else {
clearReputationHealthState(scanState, reputationQuotaStateKey)
}
lastRefresh := time.Time{}
if db := GetThreatDB(); db != nil {
lastRefresh = db.LastFeedRefresh()
}
feedActive := feedRefreshStale(lastRefresh, now)
if _, off := disabled[logicalOwnerReputationFeedStale]; off {
clearReputationHealthState(scanState, reputationFeedStateKey)
} else if feedActive {
finding := alert.Finding{
Severity: alert.Warning,
Check: "threat_feed_stale",
Message: "Threat intelligence feeds stale: reputation matching runs on old data",
Details: fmt.Sprintf("Last successful feed update: %s. Downloads retry on each deep-scan cycle; check outbound connectivity to the feed mirrors.", lastRefresh.UTC().Format(time.RFC3339)),
Timestamp: now,
}
appendReputationHealthFinding(ctx, &out, scanState, logicalOwnerReputationFeedStale, reputationFeedStateKey, now, finding)
} else {
clearReputationHealthState(scanState, reputationFeedStateKey)
}
return out
}
func appendReputationHealthFinding(ctx context.Context, out *[]alert.Finding, scanState *state.Store, owner, stateKey string, now time.Time, finding alert.Finding) {
if scanState != nil {
emit, err := scanState.ClaimRawTimestamp(stateKey, now, reputationHealthReminder)
if err != nil {
fmt.Fprintf(os.Stderr, "reputation: persisting %s reminder: %v\n", finding.Check, err)
}
if !emit {
markCheckIncomplete(ctx, owner)
return
}
}
*out = append(*out, finding)
}
func clearReputationHealthState(scanState *state.Store, stateKey string) {
if scanState == nil {
return
}
if err := scanState.DeleteRawAndSave(stateKey); err != nil {
fmt.Fprintf(os.Stderr, "reputation: clearing health reminder state %s: %v\n", stateKey, err)
}
}
func feedRefreshStale(lastRefresh, now time.Time) bool {
if lastRefresh.IsZero() {
return false
}
lastRefresh = lastRefresh.Round(0)
now = now.Round(0)
return !lastRefresh.After(now) && now.Sub(lastRefresh) > 7*24*time.Hour
}
// nextUTCMidnight returns 00:00 UTC on the day after now — the point at
// which AbuseIPDB's daily quota resets.
func nextUTCMidnight(now time.Time) time.Time {
u := now.UTC()
return time.Date(u.Year(), u.Month(), u.Day()+1, 0, 0, 0, 0, time.UTC)
}
// abuseQuotaReady reports whether we may call AbuseIPDB right now. It
// combines the persisted backoff (set when the API returns 429/402) with
// the daily query counter (stops before we approach the free-tier cap).
// Returns true when no bbolt store is available (fallback mode).
func abuseQuotaReady(sdb *store.DB, now time.Time) bool {
if sdb == nil {
return true
}
if until := sdb.AbuseQuotaExhaustedUntil(); !until.IsZero() && now.Before(until) {
return false
}
if sdb.AbuseQueryCount(now.UTC().Format("2006-01-02")) >= maxDailyAbuseQueries {
return false
}
return true
}
// collectRecentIPs gathers non-infra IPs from multiple log sources.
// Returns map of IP → source description (e.g. "SSH login", "Dovecot IMAP auth failure").
func collectRecentIPs(cfg *config.Config) map[string]string {
ips := make(map[string]string)
info := platform.Detect()
// SSH logins. Path differs by OS family (secure vs auth.log); the old
// hardcoded /var/log/secure made this loop dead on Debian/Ubuntu.
for _, line := range tailFile(info.AuthLogPath(), 50) {
if !strings.Contains(line, "Accepted") {
continue
}
if ip := extractIPAfterKeyword(line, "from"); ip != "" {
addIfNotInfra(ips, ip, "SSH login", cfg)
}
}
// Web server access logs, platform-detected (Apache/Nginx/LiteSpeed).
for _, path := range info.AccessLogPaths {
if path == "" || isWHMAccessLog(path) {
continue
}
lines := tailFile(path, 100)
if len(lines) == 0 {
continue
}
for _, line := range lines {
if ip := firstField(line); ip != "" {
addIfNotInfra(ips, ip, "HTTP request", cfg)
}
}
}
// Dovecot - IMAP/POP3 auth failures.
if mailLog := reputationMailLogPath(cfg, info); mailLog != "" {
for _, line := range tailFile(mailLog, 50) {
if strings.Contains(line, "auth failed") || strings.Contains(line, "Aborted login") {
if ip := extractIPAfterKeyword(line, "rip="); ip != "" {
addIfNotInfra(ips, ip, "Dovecot IMAP/POP3 auth failure", cfg)
}
}
}
}
if info.IsCPanel() {
// cPanel/WHM access log.
for _, line := range tailFile(reputationWHMAccessLog, 100) {
if ip := firstField(line); ip != "" {
addIfNotInfra(ips, ip, "cPanel/WHM access", cfg)
}
}
}
if shouldCollectEximMainlog(info) {
for _, line := range tailFile(reputationEximMainlog, 50) {
if strings.Contains(line, "authenticator failed") || strings.Contains(line, "rejected RCPT") {
if ip := eximlog.ClientIP(line); ip != "" {
addIfNotInfra(ips, ip, "SMTP auth failure", cfg)
}
}
}
}
return ips
}
// collectAuthenticatedIPs returns IPs that successfully authenticated to a
// mailbox in the recent log window. A source that holds valid mail credentials
// is a real customer, not a drive-by scanner: a threat-feed match on such an IP
// is almost always a recycled dynamic/CGNAT address rather than the attacker the
// feed once listed. Romanian and other residential ISPs rotate these pools
// aggressively, so an IP an attacker used last week routinely lands on a paying
// customer this week. Webmail authenticates through dovecot, so this also covers
// Horde/Roundcube users.
//
// Successful SSH logins are deliberately excluded: a clean SSH auth from a
// feed-listed IP is itself a red flag (a cracked or attacker-controlled host),
// and collectRecentIPs already surfaces those IPs precisely so their reputation
// is checked. Authenticated mail attackers (compromised accounts) remain covered
// by the brute-force, compromise, and account-takeover detectors, which do not
// honour this exemption. The mail window is wider than collectRecentIPs uses so
// a customer's success is not pushed out of view by a burst of attacker failures
// occupying the tail.
func collectAuthenticatedIPs(cfg *config.Config) map[string]bool {
authed := make(map[string]bool)
info := platform.Detect()
mailLog := reputationMailLogPath(cfg, info)
if mailLog == "" {
return authed
}
for _, line := range tailFile(mailLog, 200) {
if !strings.Contains(line, "-login: Logged in") {
continue
}
if ip := extractIPAfterKeyword(line, "rip="); ip != "" {
authed[ip] = true
}
}
return authed
}
func reputationMailLogPath(cfg *config.Config, info platform.Info) string {
if cfg == nil {
return info.MailLogPath()
}
if cfg.MailLogs.Source == "journal" {
return ""
}
if cfg.MailLogs.File != "" {
return cfg.MailLogs.File
}
return info.MailLogPath()
}
func isWHMAccessLog(path string) bool {
return filepath.Clean(path) == reputationWHMAccessLog
}
func shouldCollectEximMainlog(info platform.Info) bool {
if info.IsCPanel() {
return true
}
if _, err := osFS.Stat(reputationEximMainlog); err != nil {
return !os.IsNotExist(err)
}
return true
}
func addIfNotInfra(ips map[string]string, ip, source string, cfg *config.Config) {
if ip == "127.0.0.1" || ip == "::1" || ip == "" {
return
}
if isInfraIP(ip, cfg.InfraIPs) {
return
}
// One address can appear on several surfaces in the same scan. Keep the
// strongest sighting so an earlier passive web request cannot hide a later
// authentication attack and downgrade the resulting finding.
current, exists := ips[ip]
if !exists || reputationSightingSeverity(source) > reputationSightingSeverity(current) {
ips[ip] = source
}
}
func firstField(line string) string {
fields := strings.Fields(line)
if len(fields) == 0 {
return ""
}
ip := fields[0]
// Validate it looks like an IP (v4 or v6)
if strings.Count(ip, ".") == 3 || strings.Contains(ip, ":") {
return ip
}
return ""
}
func extractIPAfterKeyword(line, keyword string) string {
idx := strings.Index(line, keyword)
if idx < 0 {
return ""
}
rest := line[idx+len(keyword):]
rest = strings.TrimLeft(rest, " =")
fields := strings.Fields(rest)
if len(fields) == 0 {
return ""
}
if ip, ok := netutil.ParseIPToken(fields[0]); ok {
return ip
}
return ""
}
// loadAllBlockedIPs returns all IPs currently blocked in CSM.
// It reads firewall state through osFS so tests can inject a filesystem,
// then merges the legacy blocked_ips.json file.
func loadAllBlockedIPs(statePath string) map[string]bool {
blocked := make(map[string]bool)
// Read the authoritative firewall engine state. The engine persists
// every block to firewall/state.json; the parallel bbolt fw:blocked
// bucket is written only at migration, so reading it would return a
// frozen snapshot that misses live blocks.
fwPath := filepath.Join(statePath, "firewall", "state.json")
if fwData, err := osFS.ReadFile(fwPath); err == nil {
var fwState struct {
Blocked []struct {
IP string `json:"ip"`
ExpiresAt time.Time `json:"expires_at"`
} `json:"blocked"`
}
if uerr := json.Unmarshal(fwData, &fwState); uerr != nil {
fmt.Fprintf(os.Stderr, "reputation: %s is corrupt, alert suppression degraded: %v\n", fwPath, uerr)
} else {
now := time.Now()
for _, entry := range fwState.Blocked {
if entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt) {
blocked[entry.IP] = true
}
}
}
}
// Also read from blocked_ips.json (legacy CSM auto-block)
type blockedEntry struct {
IP string `json:"ip"`
ExpiresAt time.Time `json:"expires_at"`
}
type blockFile struct {
IPs []blockedEntry `json:"ips"`
Pending []struct {
IP string `json:"ip"`
} `json:"pending,omitempty"`
}
legacyPath := filepath.Join(statePath, "blocked_ips.json")
data, err := osFS.ReadFile(legacyPath)
if err == nil {
var bf blockFile
if uerr := json.Unmarshal(data, &bf); uerr != nil {
fmt.Fprintf(os.Stderr, "reputation: %s is corrupt, alert suppression degraded: %v\n", legacyPath, uerr)
} else {
now := time.Now()
for _, entry := range bf.IPs {
if now.Before(entry.ExpiresAt) {
blocked[entry.IP] = true
}
}
for _, entry := range bf.Pending {
blocked[entry.IP] = true
}
}
}
return blocked
}
type abuseIPDBResponse struct {
Data struct {
AbuseConfidenceScore int `json:"abuseConfidenceScore"`
UsageType string `json:"usageType"`
ISP string `json:"isp"`
TotalReports int `json:"totalReports"`
} `json:"data"`
Errors []struct {
Detail string `json:"detail"`
Status int `json:"status"`
} `json:"errors"`
}
// queryAbuseIPDB returns (score, category, error).
// Returns specific errors for rate limiting (429) and quota exhaustion (402).
func queryAbuseIPDB(client *http.Client, ip, apiKey string, work *reputationQueryWork) (score int, category string, err error) {
req, err := http.NewRequest("GET", abuseIPDBEndpoint+"?ipAddress="+url.QueryEscape(ip)+"&maxAgeInDays=90", nil)
if err != nil {
return 0, "", err
}
req.Header.Set("Key", apiKey)
req.Header.Set("Accept", "application/json")
resp, err := client.Do(req)
if err != nil {
return 0, "", err
}
defer func() {
work.cleanup(err)
_ = resp.Body.Close()
}()
if resp.StatusCode == 429 {
return 0, "", fmt.Errorf("429 rate limited")
}
if resp.StatusCode == 402 {
return 0, "", fmt.Errorf("402 quota exceeded")
}
if resp.StatusCode != 200 {
return 0, "", fmt.Errorf("HTTP %d", resp.StatusCode)
}
var result abuseIPDBResponse
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
return 0, "", err
}
if len(result.Errors) > 0 {
return 0, "", fmt.Errorf("API error: %s", result.Errors[0].Detail)
}
category = result.Data.UsageType
if result.Data.ISP != "" {
category += " (" + result.Data.ISP + ")"
}
if result.Data.TotalReports > 0 {
category += fmt.Sprintf(", %d reports", result.Data.TotalReports)
}
return result.Data.AbuseConfidenceScore, category, nil
}
func abuseQuotaError(err error) bool {
return strings.Contains(err.Error(), "429") || strings.Contains(err.Error(), "402")
}
// cleanCache removes expired entries and caps at maxCacheEntries.
func cleanCache(cache *reputationCache) {
// When using bbolt store, delegate cleanup to store methods, then
// mirror the prune on the in-memory view: entries deleted only in
// bbolt would linger in the map and could be pushed back by a save.
if sdb := store.Global(); sdb != nil {
sdb.CleanExpiredReputation(cacheExpiry)
sdb.EnforceReputationCap(maxCacheEntries)
}
pruneCacheEntries(cache)
}
// pruneCacheEntries drops expired entries from the in-memory map and
// enforces the size cap, evicting oldest first.
func pruneCacheEntries(cache *reputationCache) {
now := time.Now()
// Remove expired entries - use same expiry as cache freshness check
for ip, entry := range cache.Entries {
if now.Sub(entry.CheckedAt) > cacheExpiry {
cache.remove(ip)
}
}
// Cap at max entries - remove oldest if over limit
if len(cache.Entries) > maxCacheEntries {
type aged struct {
ip string
age time.Duration
}
entries := make([]aged, 0, len(cache.Entries))
for ip, entry := range cache.Entries {
entries = append(entries, aged{ip, now.Sub(entry.CheckedAt)})
}
// Sort by age descending (oldest first)
sort.Slice(entries, func(i, j int) bool {
return entries[i].age > entries[j].age
})
// Remove oldest until under limit
for i := 0; i < len(entries)-maxCacheEntries; i++ {
cache.remove(entries[i].ip)
}
}
}
func loadReputationCache(statePath string) *reputationCache {
cache := &reputationCache{
Entries: make(map[string]*reputationEntry),
dirty: make(map[string]bool),
removed: make(map[string]bool),
}
// Try bbolt store first - after migration the flat file is renamed to .bak.
if sdb := store.Global(); sdb != nil {
for ip, entry := range sdb.AllReputation() {
cache.Entries[ip] = &reputationEntry{
Score: entry.Score,
Category: entry.Category,
CheckedAt: entry.CheckedAt,
}
}
return cache
}
// Fallback: flat-file JSON (pre-migration).
data, err := osFS.ReadFile(filepath.Join(statePath, reputationCacheFile))
if err == nil {
_ = json.Unmarshal(data, cache)
if cache.Entries == nil {
cache.Entries = make(map[string]*reputationEntry)
}
}
return cache
}
func saveReputationCache(statePath string, cache *reputationCache) bool {
if sdb := store.Global(); sdb != nil {
changed := cache.changedEntries()
if len(changed) == 0 && len(cache.removed) == 0 {
return true
}
if err := sdb.ApplyReputationChanges(changed, cache.removed); err != nil {
// Keep the pending marks so a later save can retry the flush.
return false
}
clear(cache.dirty)
clear(cache.removed)
return true
}
// Fallback: flat-file JSON.
data, err := json.MarshalIndent(cache, "", " ")
if err != nil {
return false
}
tmpPath := filepath.Join(statePath, reputationCacheFile+".tmp")
if err := os.WriteFile(tmpPath, data, 0600); err != nil {
return false
}
return os.Rename(tmpPath, filepath.Join(statePath, reputationCacheFile)) == nil
}
package checks
import (
"context"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var reputationQueries = newReputationQueue()
type reputationQueueMonitor struct {
mu sync.Mutex
pending map[*reputationQueryWork]struct{}
losses *queuehealth.Tracker
}
type reputationQueryBatch struct {
consumer *reputationQueryWork // guarded by the queue mutex
}
type reputationQueryWork struct {
batch *reputationQueryBatch
queue *reputationQueueMonitor
at, deadline time.Time
running, queryStarted, queryDone, deliveryDone, failed bool
}
func newReputationQueue() *reputationQueueMonitor {
return &reputationQueueMonitor{pending: make(map[*reputationQueryWork]struct{}), losses: queuehealth.New(0, time.Minute)}
}
func (q *reputationQueueMonitor) begin(batch *reputationQueryBatch) *reputationQueryWork {
now := time.Now()
w := &reputationQueryWork{queue: q, batch: batch, at: now, deadline: now.Add(time.Minute)}
q.mu.Lock()
q.pending[w] = struct{}{}
q.mu.Unlock()
return w
}
func (w *reputationQueryWork) moveLocked(running bool, budget time.Duration) {
w.running = running
w.at = time.Now()
w.deadline = w.at.Add(budget)
}
func (w *reputationQueryWork) start(budget time.Duration) {
w.queue.mu.Lock()
w.queryStarted = true
w.moveLocked(true, budget)
w.queue.mu.Unlock()
}
func (w *reputationQueryWork) failLocked() {
if !w.failed {
w.failed = true
w.queue.losses.Lose(time.Now(), 1)
}
}
func (w *reputationQueryWork) fail() {
w.queue.mu.Lock()
w.failLocked()
w.queue.mu.Unlock()
}
func (w *reputationQueryWork) cleanup(err error) {
// Protocol calls outside the worker pool have no queued delivery owner.
if w == nil {
return
}
failed := err != nil && !abuseQuotaError(err)
w.queue.mu.Lock()
if failed {
w.failLocked()
}
w.moveLocked(true, time.Minute)
w.queue.mu.Unlock()
}
func (w *reputationQueryWork) returned(completed bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
if !completed {
w.failLocked()
}
w.queryDone = true
w.moveLocked(false, time.Minute)
w.releaseLocked()
}
func (w *reputationQueryWork) phase(running bool) {
w.queue.mu.Lock()
w.moveLocked(running, time.Minute)
if running {
w.batch.consumer = w
} else if w.batch.consumer == w {
w.batch.consumer = nil
}
w.queue.mu.Unlock()
}
func (w *reputationQueryWork) consuming(ctx context.Context, budget time.Duration) {
deadline := time.Now().Add(budget)
if parent, ok := ctx.Deadline(); ok && parent.Before(deadline) {
deadline = parent
}
w.queue.mu.Lock()
w.moveLocked(true, budget)
w.deadline = deadline
w.batch.consumer = w
w.queue.mu.Unlock()
}
func (w *reputationQueryWork) finish(success bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
if w.deliveryDone {
return
}
w.deliveryDone = true
if w.batch.consumer == w {
w.batch.consumer = nil
}
if !success {
w.failLocked()
}
// Dispatch can end before a worker starts. Already running HTTP keeps its
// separate owner until response cleanup and result publication finish.
if !w.queryStarted {
w.queryDone = true
}
w.releaseLocked()
}
func (w *reputationQueryWork) releaseLocked() {
if w.queryDone && w.deliveryDone {
delete(w.queue.pending, w)
}
}
// ReputationQueueStatus reads memory only, including queries outliving a check.
func ReputationQueueStatus(now time.Time) queuehealth.Status {
q := reputationQueries
q.mu.Lock()
defer q.mu.Unlock()
s := q.losses.Snapshot(now)
s.CapacityUnavailable = true
var waitingLate, runningLate bool
for w := range q.pending {
if w.running {
s.InFlight++
s.ProcessingSeconds = max(s.ProcessingSeconds, now.Sub(w.at).Seconds())
runningLate = runningLate || !now.Before(w.deadline)
} else {
s.Depth++
s.LagSeconds = max(s.LagSeconds, now.Sub(w.at).Seconds())
// A result waiting behind its own batch consumer follows that
// consumer's deadline. Another batch cannot lend it progress.
waitingLate = waitingLate || ((!w.queryDone || w.batch.consumer == nil) && !now.Before(w.deadline))
}
}
switch {
case waitingLate:
s.Reason = "backlog_lag"
case runningLate:
s.Reason = "processing_lag"
}
if s.Reason != "" {
s.Status = "degraded"
}
return s
}
package checks
import "sort"
// rollingCandidatesAfter selects up to `limit` candidates from the path-sorted,
// de-duplicated `sorted` slice whose path sorts strictly AFTER `lastPath`,
// wrapping once to the start of the list if the end is reached. It returns the
// selected paths (in scan order), the new cursor value (the last selected path,
// to persist as last_path), and whether a wrap occurred during this selection.
//
// Contract:
// - `sorted` MUST be ascending sorted + de-duplicated (caller guarantees).
// - lastPath == "" starts from the beginning.
// - Never returns more than len(sorted) items (one full cycle max), even if
// limit exceeds len(sorted).
// - Wrap happens at most once; selection stops when it would revisit the
// first item of this selection (no infinite loop, no double-scan in a cycle).
// - limit <= 0 returns (nil, lastPath, false) -- no rolling this cycle.
// - Empty `sorted` returns (nil, lastPath, false).
// - newLast is the last path actually selected; if nothing was selected it is
// lastPath unchanged.
// - Robust to add/remove between cycles: selection is by VALUE comparison
// against lastPath (sort.SearchStrings for the first index strictly greater
// than lastPath), never an integer offset -- a deleted lastPath or inserted
// earlier file does not corrupt progress.
//
// Documented end-of-list rule for limit >= len(sorted):
//
// When starting from "" (or any cursor where the pre-wrap tail plus the
// wrapped head together would cover every element), at most len(sorted) items
// are returned. If the initial cursor is at the very beginning ("") and the
// list fits inside limit, the traversal reaches the end without issuing a
// wrap, so wrapped=false. If the cursor is mid-list and the remaining tail
// plus the needed head from the wrap together account for all elements,
// wrapped=true and the selection still caps at len(sorted) with no repeats.
func rollingCandidatesAfter(sorted []string, lastPath string, limit int) (selected []string, newLast string, wrapped bool) {
if limit <= 0 || len(sorted) == 0 {
return nil, lastPath, false
}
// Cap to one full cycle so we never return duplicates.
maxCount := limit
if maxCount > len(sorted) {
maxCount = len(sorted)
}
// Find the first index strictly greater than lastPath.
// sort.SearchStrings returns the smallest i where sorted[i] >= lastPath.
// If sorted[i] == lastPath we advance by 1 to get strictly greater.
start := sort.SearchStrings(sorted, lastPath)
if start < len(sorted) && sorted[start] == lastPath {
start++
}
selected = make([]string, 0, maxCount)
// Phase 1: collect from start to end of sorted.
i := start
for len(selected) < maxCount && i < len(sorted) {
selected = append(selected, sorted[i])
i++
}
// Phase 2: if we still need more and reached the end, wrap to the beginning.
// Stop before `start` so we never revisit a Phase-1 item (no repeats). This
// also covers a cursor at/after the last element (start >= len(sorted) with
// start > 0, e.g. lastPath == max or a single-element list): Phase 1 collects
// nothing and the wrap re-covers from the head.
if len(selected) < maxCount && i >= len(sorted) && start > 0 {
wrapped = true
for j := 0; len(selected) < maxCount && j < start; j++ {
selected = append(selected, sorted[j])
}
}
if len(selected) == 0 {
return nil, lastPath, false
}
return selected, selected[len(selected)-1], wrapped
}
package checks
import (
"context"
"fmt"
"sort"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// autoResponseActions counts every auto-response action fired, by
// action class. Registered lazily on first observation.
var (
autoResponseActions *metrics.CounterVec
autoResponseActionsOnce sync.Once
)
func observeAutoResponse(action string, n int) {
if n <= 0 {
return
}
autoResponseActionsOnce.Do(func() {
autoResponseActions = metrics.NewCounterVec(
"csm_auto_response_actions_total",
"Auto-response actions fired, by action. Labels: action (kill|quarantine|block). Incremented once per finding the auto-response subsystem produced in each tier run; a batch of four IPs blocked in one cycle contributes 4 to action=block.",
[]string{"action"},
)
metrics.MustRegister("csm_auto_response_actions_total", autoResponseActions)
})
autoResponseActions.With(action).Add(float64(n))
}
func observeFileResponseActions(action string, findings []alert.Finding) {
n := 0
for _, finding := range findings {
if finding.Check == "auto_response" {
n++
}
}
observeAutoResponse(action, n)
}
// checkDuration is the per-check latency histogram for /metrics.
// Labelled by check name and tier so scrapers can spot a single check
// regressing without scanning logs. Buckets span the observed range
// from ~millisecond process-list passes up to the heavy-check timeout
// ceiling.
var (
checkDuration *metrics.HistogramVec
checkDurationOnce sync.Once
)
var checkDurationBuckets = []float64{0.01, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30, 60, 120, 180, 300, 600, 900}
func observeCheckDuration(name, tier string, d time.Duration) {
checkDurationOnce.Do(func() {
checkDuration = metrics.NewHistogramVec(
"csm_check_duration_seconds",
"Wall-clock time for each security check to complete. Label `name` is a check runner name; label `tier` is critical|deep|all. Use p95 across name to spot a single check regressing, and sum across name to track per-cycle pressure.",
[]string{"name", "tier"},
checkDurationBuckets,
)
metrics.MustRegister("csm_check_duration_seconds", checkDuration)
})
checkDuration.With(name, tier).Observe(d.Seconds())
}
// CheckFunc is the signature for all check functions.
// The context is cancelled when the check times out so goroutines can exit.
type CheckFunc func(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding
// splitDisabledChecks partitions checks by cfg.DisabledChecks. Finding names
// are the public vocabulary used by the settings UI and docs; runner names
// remain accepted for existing operator configs.
func splitDisabledChecks(cfg *config.Config, checks []namedCheck) (enabled, disabled []namedCheck) {
if cfg == nil || len(cfg.DisabledChecks) == 0 {
return checks, nil
}
disabledSet := make(map[string]struct{}, len(cfg.DisabledChecks))
knownRunners := make(map[string]struct{}, len(checks))
for _, nc := range checks {
knownRunners[nc.name] = struct{}{}
}
for _, name := range cfg.DisabledChecks {
name = strings.TrimSpace(name)
if name == "" {
continue
}
if _, ok := knownRunners[name]; ok {
disabledSet[name] = struct{}{}
continue
}
for _, runner := range runnerNamesForFinding(name) {
if _, ok := knownRunners[runner]; ok {
disabledSet[runner] = struct{}{}
}
}
}
ownerOff := disabledLogicalOwners(cfg)
for name := range disabledSet {
for _, owner := range physicalCheckLogicalOwners[name] {
if _, off := ownerOff[owner]; !off {
// An enabled hosted owner keeps the wrapper runnable; the
// check itself skips its disabled consumer.
delete(disabledSet, name)
break
}
}
}
if len(disabledSet) == 0 {
return checks, nil
}
enabled = make([]namedCheck, 0, len(checks))
disabledChecks := make([]namedCheck, 0, len(disabledSet))
for _, nc := range checks {
if _, skip := disabledSet[nc.name]; skip {
disabledChecks = append(disabledChecks, nc)
continue
}
enabled = append(enabled, nc)
}
return enabled, disabledChecks
}
func runnerNamesForFinding(finding string) []string {
return findingNameToRunnerNames[finding]
}
// DisabledCheckNames returns the sorted public finding-name vocabulary accepted
// by top-level disabled_checks for scheduled check execution. Runner IDs are
// also accepted by splitDisabledChecks for existing configs, but are not
// exposed in the UI.
func DisabledCheckNames() []string {
seen := make(map[string]struct{}, len(findingNameToRunnerNames))
for finding := range findingNameToRunnerNames {
seen[finding] = struct{}{}
}
for _, aliases := range logicalOwnerDisableAliases {
for _, alias := range aliases {
seen[alias] = struct{}{}
}
}
out := make([]string, 0, len(seen))
for finding := range seen {
info, ok := LookupCheck(finding)
if !ok || info.Internal {
continue
}
out = append(out, finding)
}
sort.Strings(out)
return out
}
// DisabledCheckConfigNames returns every value top-level disabled_checks
// accepts: this is exactly the set splitDisabledChecks honors -- every emitted
// finding name (including internal ones) plus every compatibility runner ID.
// It is broader than DisabledCheckNames (the UI vocabulary) so POST-side
// validation never rejects a value an existing operator config relies on.
func DisabledCheckConfigNames() []string {
seen := make(map[string]struct{}, len(findingNameToRunnerNames)+len(runnerFindingNames))
for finding := range findingNameToRunnerNames {
seen[finding] = struct{}{}
}
for runner := range runnerFindingNames {
seen[runner] = struct{}{}
}
// Logical-owner IDs and their public finding aliases are valid disable
// values; the coverage diagnostic stays rejected because it appears in no
// accepted vocabulary.
for _, aliases := range logicalOwnerDisableAliases {
for _, alias := range aliases {
seen[alias] = struct{}{}
}
}
out := make([]string, 0, len(seen))
for name := range seen {
out = append(out, name)
}
sort.Strings(out)
return out
}
var findingNameToRunnerNames = buildFindingNameToRunnerNames()
func buildFindingNameToRunnerNames() map[string][]string {
out := map[string][]string{}
for runner, findings := range runnerFindingNames {
for _, finding := range findings {
out[finding] = append(out[finding], runner)
}
}
return out
}
var runnerFindingNames = map[string][]string{
"admin_overlap": {"admin_cross_account_overlap"},
"credential_reuse": {"credential_reuse"},
"supply_chain": {"supply_chain_vuln"},
"af_alg_enforcement": {"af_alg_enforcement_corrected"},
"af_alg_socket_use": {"af_alg_socket_use"},
"api_auth_failures": {"api_auth_failure"},
"api_tokens": {"api_tokens"},
"cpanel_filemanager": {"cpanel_file_upload"},
"cpanel_logins": {"cpanel_login", "cpanel_multi_ip_login", "cpanel_password_purge"},
"crontabs": {"crond_change", "crontab_change", "suspicious_crontab"},
"database_dumps": {"database_dump"},
"db_content": {"db_content_scan_incomplete", "db_doorway_sitemap_routes", "db_hidden_link_injection", "db_hostname_keyed_option", "db_options_injection", "db_options_new_external_script", "db_options_plugin_notice_injection", "db_phantom_post_author", "db_post_injection", "db_post_volume_burst", "db_spam_taxonomy", "db_stored_cloak_logic", "db_stored_code_execution", "db_rogue_admin", "db_siteurl_foreign_host", "db_siteurl_hijack", "db_siteurl_invalid", "db_spam_cleaned", "db_spam_found", "db_spam_injection", "db_suspicious_admin_email"},
"db_content_drupal": {"drupal_admin_injection", "drupal_content_injection", "drupal_settings_injection"},
"db_content_joomla": {"joomla_admin_injection", "joomla_content_injection", "joomla_extensions_injection"},
"db_content_magento": {"magento_admin_injection", "magento_content_injection", "magento_settings_injection"},
"db_content_opencart": {"opencart_admin_injection", "opencart_content_injection", "opencart_settings_injection"},
"db_objects": {"db_magic_token_user", "db_malicious_event", "db_malicious_function", "db_malicious_procedure", "db_malicious_trigger", "db_unexpected_event", "db_unexpected_function", "db_unexpected_procedure", "db_unexpected_trigger"},
"dns_connections": {"dns_connection"},
"dns_zones": {"dns_zone_change"},
"email_content": {"email_phishing_content"},
"email_forwarder_audit": {"email_pipe_forwarder", "email_suspicious_forwarder"},
"email_mail_filters": {"email_filter_blackhole", "email_filter_exfil", "email_filter_forwarder", "email_filter_pipe", "email_mail_filters"},
"email_weak_password": {"email_weak_password", "email_password_audit_incomplete"},
"exfiltration_paste": {"exfiltration_paste_site"},
"fake_kernel_threads": {"fake_kernel_thread"},
// new_php_in_languages and new_php_in_upgrade were emitted until a20c6f76;
// they stay here so a completed scan clears rows written by older versions.
"file_index": {"new_executable_in_config", "new_php_in_languages", "new_php_in_sensitive_dir", "new_php_in_sensitive_dir_clean", "new_php_in_upgrade", "new_php_in_uploads", "new_php_in_uploads_clean", "new_suspicious_php", "new_webshell_file", "obfuscated_php", "suspicious_php_content"},
"filesystem": {"backdoor_binary", "suid_binary", "suspicious_file"},
"firewall": {"firewall", "firewall_ports", "firewall_ipv6_unmanaged"},
"ftp_logins": {"ftp_bruteforce", "ftp_login", "ftp_login_after_bruteforce"},
"group_writable_php": {"group_writable_php"},
"health": {"csm_health"},
"htaccess": append([]string{"htaccess_handler_abuse", "htaccess_injection"}, htaccessDetectorNames()...),
"exposed_files": {"web_exposed_config_leak", "web_exposed_db_dump", "web_exposed_backup_archive", "web_exposed_source_backup", "web_exposed_phpinfo", "web_exposed_sample_sql"},
"ip_reputation": {"ip_reputation"},
"kernel_modules": {"kernel_module"},
"local_threat_score": {"local_threat_score"},
"mail_per_account": {"mail_per_account"},
"mail_queue": {"mail_queue", "mail_queue_unavailable"},
"modsec_audit": {"waf_attack_blocked"},
"mysql_users": {"mysql_superuser"},
"nulled_plugins": {"nulled_plugin"},
"open_basedir": {"open_basedir"},
"outbound_connections": {"backdoor_port", "backdoor_port_outbound", "c2_connection"},
"outdated_plugins": {"outdated_plugins"},
"vulnerable_plugins": {"vulnerable_plugins"},
"vulnerable_timthumb": {"vulnerable_timthumb"},
"perf_error_logs": {"perf_error_logs"},
"perf_load": {"perf_load"},
"perf_memory": {"perf_memory"},
"perf_mysql_config": {"perf_mysql_config"},
"perf_php_handler": {"perf_php_handler"},
"perf_php_processes": {"perf_php_processes"},
"perf_redis_config": {"perf_redis_config"},
"perf_wp_config": {"perf_wp_config"},
"perf_wp_cron": {"perf_wp_cron"},
"perf_wp_transients": {"perf_wp_transients"},
"phishing": {"phishing_credential_log", "phishing_directory", "phishing_iframe", "phishing_kit_archive", "phishing_page", "phishing_php", "phishing_redirector"},
"yara_deep": {"yara_match_scheduled", "yara_scan_incomplete"},
"php_config_changes": {"php_config_change", "php_config_scan_incomplete"},
"php_content": {"obfuscated_php", "suspicious_php_content"},
"php_processes": {"php_suspicious_execution"},
"rpm_integrity": {"dpkg_integrity", "rpm_integrity"},
"shadow_changes": {"bulk_password_change", "root_password_change", "shadow_change"},
"ssh_keys": {"ssh_keys"},
"ssh_logins": {"ssh_login_unknown_ip"},
"sshd_config": {"sshd_config_change"},
"ssl_certs": {"ssl_cert_issued"},
"suspicious_processes": {"suspicious_process"},
"symlink_attacks": {"symlink_attack"},
"uid0_accounts": {"uid0_account"},
"user_outbound": {"user_outbound_connection", "direct_smtp_egress", "bad_asn_outbound"},
"waf_status": {"modsec_disabled_vhost", "waf_bypass", "waf_detection_only", "waf_rules", "waf_rules_stale", "waf_status"},
"webmail_logins": {"webmail_bruteforce"},
"webshells": {"webshell", "world_writable_php"},
"whm_access": {"whm_account_action", "whm_password_change"},
"wp_bruteforce": {"wp_login_bruteforce", "wp_user_enumeration", "xmlrpc_abuse", "http_request_flood", "http_scanner_profile", "http_claimed_bot_unverified", "http_ua_spoof", "http_distributed_flood", "http_asn_crawl"},
"wp_core": {"wp_core_integrity"},
"wp_plugin_inventory": {"wp_plugin_inventory_unverified"},
}
const (
logicalOwnerJSTaintDeep = "js_taint_deep"
logicalOwnerPHPTaintDeep = "php_taint_deep"
logicalOwnerReputationQuota = "reputation_quota_health"
logicalOwnerReputationFeedStale = "reputation_feed_health"
logicalOwnerWPCoreVerification = "wp_core_verification"
)
// logicalOwnerFindingNames maps a logical finding owner hosted inside another
// physical check to the finding names it owns. A logical owner is not a
// runnable check, so it must never be a runnerFindingNames key; the hosting
// check reports each owner's completion independently by calling
// markCheckIncomplete with the owner's exact name.
var logicalOwnerFindingNames = map[string][]string{
logicalOwnerJSTaintDeep: {"js_keylogger_dataflow", "js_taint_scan_incomplete"},
logicalOwnerPHPTaintDeep: {"php_remote_taint", "php_taint_scan_incomplete"},
logicalOwnerReputationQuota: {"reputation_quota_exhausted"},
logicalOwnerReputationFeedStale: {"threat_feed_stale"},
logicalOwnerWPCoreVerification: {"wp_core_unverified"},
}
// logicalOwnerDisableAliases maps a logical owner to the disabled_checks
// values that disable it. A physical host alias is included only when
// disabling that host is also meant to disable the hosted consumer.
var logicalOwnerDisableAliases = map[string][]string{
logicalOwnerJSTaintDeep: {logicalOwnerJSTaintDeep, "js_keylogger_dataflow"},
logicalOwnerPHPTaintDeep: {logicalOwnerPHPTaintDeep, "php_remote_taint"},
logicalOwnerReputationQuota: {logicalOwnerReputationQuota, "reputation_quota_exhausted", "ip_reputation"},
logicalOwnerReputationFeedStale: {logicalOwnerReputationFeedStale, "threat_feed_stale", "ip_reputation"},
logicalOwnerWPCoreVerification: {logicalOwnerWPCoreVerification, "wp_core_unverified", "wp_core"},
}
// physicalCheckLogicalOwners maps a runnable check to the logical owners it
// hosts. splitDisabledChecks keeps the physical wrapper runnable while any
// hosted owner is enabled, and the runner purges each owner by its own
// completion mark rather than the wrapper's.
var physicalCheckLogicalOwners = map[string][]string{
"yara_deep": {logicalOwnerJSTaintDeep, logicalOwnerPHPTaintDeep},
"ip_reputation": {logicalOwnerReputationQuota, logicalOwnerReputationFeedStale},
"wp_core": {logicalOwnerWPCoreVerification},
}
// disabledLogicalOwners returns the logical owners disabled by cfg.
func disabledLogicalOwners(cfg *config.Config) map[string]struct{} {
out := map[string]struct{}{}
if cfg == nil || len(cfg.DisabledChecks) == 0 {
return out
}
disabled := make(map[string]struct{}, len(cfg.DisabledChecks))
for _, name := range cfg.DisabledChecks {
if name = strings.TrimSpace(name); name != "" {
disabled[name] = struct{}{}
// A public finding alias also disables owners that explicitly
// inherit disablement from its physical check.
for _, runner := range runnerNamesForFinding(name) {
disabled[runner] = struct{}{}
}
}
}
for owner, aliases := range logicalOwnerDisableAliases {
for _, alias := range aliases {
if _, ok := disabled[alias]; ok {
out[owner] = struct{}{}
break
}
}
}
return out
}
// jsTaintDeepConsumerDisabled reports whether the scheduled JS taint consumer
// hosted by yara_deep is disabled.
func jsTaintDeepConsumerDisabled(cfg *config.Config) bool {
_, off := disabledLogicalOwners(cfg)[logicalOwnerJSTaintDeep]
return off
}
// phpTaintDeepConsumerDisabled reports whether the PHP taint consumer is
// disabled; the physical wrapper may still run for the other consumers.
func phpTaintDeepConsumerDisabled(cfg *config.Config) bool {
_, off := disabledLogicalOwners(cfg)[logicalOwnerPHPTaintDeep]
return off
}
// yaraDeepConsumerDisabled reports whether the YARA consumer itself is
// disabled; the physical wrapper may still run for the JS consumer.
func yaraDeepConsumerDisabled(cfg *config.Config) bool {
if cfg == nil {
return false
}
for _, name := range cfg.DisabledChecks {
name = strings.TrimSpace(name)
if name == "yara_deep" {
return true
}
for _, alias := range runnerFindingNames["yara_deep"] {
if name == alias {
return true
}
}
}
return false
}
// namedCheck pairs a check function with its name for timeout reporting.
type namedCheck struct {
name string
fn CheckFunc
}
// ForceAll forces all checks to run regardless of throttle (used by baseline).
var ForceAll bool
// Tier identifies which set of checks to run.
type Tier string
const (
TierCritical Tier = "critical" // Fast checks - processes, auth, network (~5 seconds)
TierDeep Tier = "deep" // Filesystem scans - webshells, htaccess, WP core (~90 seconds)
TierAll Tier = "all" // Both tiers
)
const checkTimeout = 5 * time.Minute
// heavyCheckTimeout applies to host-wide work over account web roots or their
// databases. On busy shared servers these legitimately run longer than the
// default 5-minute budget, so they get a wider window to avoid noisy
// check_timeout warnings while leaving fast checks aggressive.
const heavyCheckTimeout = 15 * time.Minute
// checkTimeoutDrainGrace bounds how long the runner waits for a check to hand
// back the findings it gathered before its deadline. A check that honors ctx
// returns within one file of the cancellation, and the slice is already paid
// for; a check wedged in a parser gets abandoned rather than holding the scan.
// A var so tests can shrink it without stalling on a deliberately stuck check.
var checkTimeoutDrainGrace = 5 * time.Second
// heavyChecks names the deep-tier checks that traverse every account's web
// roots or databases. Keep this list short and explicit; only checks that
// observably blow past 5 minutes on production hosts belong here.
var heavyChecks = map[string]bool{
"filesystem": true,
"webshells": true,
"htaccess": true,
"exposed_files": true,
"php_content": true,
"file_index": true,
"phishing": true,
"yara_deep": true,
"php_config_changes": true,
"db_content": true,
}
// timeoutFor returns the per-check execution budget. Heavy host-wide scans
// get heavyCheckTimeout, everything else gets checkTimeout.
// Indirected through timeoutForFunc so tests can shrink budgets
// without mutating the const.
func timeoutFor(name string) time.Duration {
return timeoutForFunc(name)
}
var timeoutForFunc = func(name string) time.Duration {
if heavyChecks[name] {
return heavyCheckTimeout
}
return checkTimeout
}
func criticalChecks() []namedCheck {
return []namedCheck{
{"fake_kernel_threads", CheckFakeKernelThreads},
{"suspicious_processes", CheckSuspiciousProcesses},
{"php_processes", CheckPHPProcesses},
{"shadow_changes", CheckShadowChanges},
{"uid0_accounts", CheckUID0Accounts},
{"ssh_keys", CheckSSHKeys},
{"sshd_config", CheckSSHDConfig},
{"ssh_logins", CheckSSHLogins},
{"api_tokens", CheckAPITokens},
{"crontabs", CheckCrontabs},
{"outbound_connections", CheckOutboundConnections},
{"user_outbound", CheckOutboundUserConnections},
{"dns_connections", CheckDNSConnections},
{"whm_access", CheckWHMAccess},
{"cpanel_logins", CheckCpanelLogins},
{"cpanel_filemanager", CheckCpanelFileManager},
{"firewall", CheckFirewall},
{"mail_queue", CheckMailQueue},
{"mail_per_account", CheckMailPerAccount},
{"kernel_modules", CheckKernelModules},
{"af_alg_socket_use", CheckAFAlgSocketUsage},
{"af_alg_enforcement", CheckAFAlgEnforcement},
{"mysql_users", CheckMySQLUsers},
{"database_dumps", CheckDatabaseDumps},
{"exfiltration_paste", CheckOutboundPasteSites},
{"wp_bruteforce", CheckWPBruteForce},
{"ftp_logins", CheckFTPLogins},
{"webmail_logins", CheckWebmailLogins},
{"api_auth_failures", CheckAPIAuthFailures},
{"ip_reputation", CheckIPReputation},
{"local_threat_score", CheckLocalThreatScore},
{"modsec_audit", CheckModSecAuditLog},
{"health", CheckHealth},
{"perf_load", CheckLoadAverage},
{"perf_php_processes", CheckPHPProcessLoad},
{"perf_memory", CheckSwapAndOOM},
}
}
func deepChecks() []namedCheck {
return []namedCheck{
{"filesystem", CheckFilesystem},
{"webshells", CheckWebshells},
{"htaccess", CheckHtaccess},
{"exposed_files", CheckExposedFiles},
{"wp_core", CheckWPCore},
{"file_index", CheckFileIndex},
{"php_content", CheckPHPContent},
{"yara_deep", CheckYARADeep},
{"phishing", CheckPhishing},
{"nulled_plugins", CheckNulledPlugins},
{"rpm_integrity", CheckRPMIntegrity},
{"group_writable_php", CheckGroupWritablePHP},
{"open_basedir", CheckOpenBasedir},
{"symlink_attacks", CheckSymlinkAttacks},
{"php_config_changes", CheckPHPConfigChanges},
{"dns_zones", CheckDNSZoneChanges},
{"ssl_certs", CheckSSLCertIssuance},
{"waf_status", CheckWAFStatus},
{"db_content", CheckDatabaseContent},
{"db_content_drupal", CheckDrupalContent},
{"db_content_joomla", CheckJoomlaContent},
{"db_content_magento", CheckMagentoContent},
{"db_content_opencart", CheckOpenCartContent},
{"db_objects", CheckDatabaseObjects},
{"admin_overlap", CheckAdminEmailOverlap},
{"credential_reuse", CheckCredentialReuse},
{"email_content", CheckOutboundEmailContent},
{"outdated_plugins", CheckOutdatedPlugins},
{"wp_plugin_inventory", CheckWPPluginVerification},
{"vulnerable_plugins", CheckVulnerablePlugins},
{"vulnerable_timthumb", CheckVulnerableTimThumb},
{"supply_chain", CheckSupplyChain},
{"email_weak_password", CheckEmailPasswords},
{"email_forwarder_audit", CheckForwarders},
{"email_mail_filters", CheckMailFilters},
{"perf_php_handler", CheckPHPHandler},
{"perf_mysql_config", CheckMySQLConfig},
{"perf_redis_config", CheckRedisConfig},
{"perf_error_logs", CheckErrorLogBloat},
{"perf_wp_config", CheckWPConfig},
{"perf_wp_transients", CheckWPTransientBloat},
{"perf_wp_cron", CheckWPCron},
}
}
func reducedDeepChecks() []namedCheck {
return []namedCheck{
{"yara_deep", CheckYARADeep},
// The file monitor sees close-after-write only. A shell written under
// a temporary name and renamed into place, or written through a bind
// mount outside the mount mark, never produces an event, so the
// budgeted rolling YAML content scan keeps running beside the
// rolling YARA scan instead of being dropped as "covered".
{"php_content", CheckPHPContent},
// Same blind spot, and the file index is the only check that closes
// it: the mask carries no FAN_MOVED_TO on any kernel and drops
// FAN_CREATE on EL8, so a renamed-in file is never reported. The
// index also owns the baseline the new-file diff runs against, which
// stops being refreshed for as long as the monitor stays attached.
{"file_index", CheckFileIndex},
// The rest of the rename-blind content scans, for the same reason.
{"webshells", CheckWebshells},
{"htaccess", CheckHtaccess},
{"phishing", CheckPhishing},
// Not merely rename-blind: a setuid bit is set by chmod, which raises
// no close-write event, so no realtime path reports one at all.
{"filesystem", CheckFilesystem},
// Confirmed by probing the vhost rather than by reading the file, so
// no file event stands in for it. Throttled, because that probe is a
// live request to a customer site and the findings are posture.
{"exposed_files", CheckExposedFiles},
{"wp_core", CheckWPCore},
{"nulled_plugins", CheckNulledPlugins},
{"rpm_integrity", CheckRPMIntegrity},
{"group_writable_php", CheckGroupWritablePHP},
{"open_basedir", CheckOpenBasedir},
{"symlink_attacks", CheckSymlinkAttacks},
{"php_config_changes", CheckPHPConfigChanges},
{"dns_zones", CheckDNSZoneChanges},
{"ssl_certs", CheckSSLCertIssuance},
{"waf_status", CheckWAFStatus},
{"db_content", CheckDatabaseContent},
{"db_content_drupal", CheckDrupalContent},
{"db_content_joomla", CheckJoomlaContent},
{"db_content_magento", CheckMagentoContent},
{"db_content_opencart", CheckOpenCartContent},
{"db_objects", CheckDatabaseObjects},
{"admin_overlap", CheckAdminEmailOverlap},
{"credential_reuse", CheckCredentialReuse},
{"email_content", CheckOutboundEmailContent},
{"outdated_plugins", CheckOutdatedPlugins},
{"wp_plugin_inventory", CheckWPPluginVerification},
{"vulnerable_plugins", CheckVulnerablePlugins},
{"vulnerable_timthumb", CheckVulnerableTimThumb},
{"supply_chain", CheckSupplyChain},
{"email_weak_password", CheckEmailPasswords},
{"email_forwarder_audit", CheckForwarders},
{"email_mail_filters", CheckMailFilters},
{"perf_php_handler", CheckPHPHandler},
{"perf_mysql_config", CheckMySQLConfig},
{"perf_redis_config", CheckRedisConfig},
{"perf_error_logs", CheckErrorLogBloat},
{"perf_wp_config", CheckWPConfig},
{"perf_wp_transients", CheckWPTransientBloat},
{"perf_wp_cron", CheckWPCron},
}
}
// PerfCheckNamesForTier returns the perf_* check names registered in the given tier.
// Used by the daemon to perform an atomic purge-and-merge when storing findings.
func PerfCheckNamesForTier(tier Tier) []string {
var names []string
for _, nc := range checksForTier(tier) {
if strings.HasPrefix(nc.name, "perf_") {
names = append(names, nc.name)
}
}
return names
}
// checkThrottleMin maps a check name to its minimum interval in minutes
// between executions. The runner reserves this before invoking the check
// function and stamps the slot only after the whole scan completes, so a
// timed-out or interrupted check keeps its slot and retries next cycle.
// Throttled checks that get skipped in a given cycle are NOT added to the
// per-scan purge list, so their previously-emitted findings stay in the
// latest set instead of being wiped every cycle. Without this gating in the
// runner, a deep scan that ran while the throttle window was still open
// would purge stale findings and merge nothing, hiding real issues until
// the next non-throttled cycle (or daemon restart).
var checkThrottleMin = map[string]int{
// Every candidate is confirmed with a live request to the customer's
// vhost, and the findings are posture that does not change between
// cycles, so this runs a few times a day rather than on each deep cycle.
"exposed_files": 360,
"perf_php_handler": 60,
"perf_mysql_config": 60,
"perf_redis_config": 60,
"perf_error_logs": 60,
"perf_wp_config": 60,
"perf_wp_transients": 60,
"perf_wp_cron": 60,
}
// LatestPurgeCheckNamesForTier returns every emitted finding name owned by a
// tier. The daemon uses this to replace a tier's current scan output without
// retaining stale findings from prior runs.
func LatestPurgeCheckNamesForTier(tier Tier) []string {
return withLogicalOwnerPurgeNames(checksForTier(tier))
}
// LatestPurgeCheckNamesForReducedDeep returns the emitted finding names owned
// by the reduced deep set used while fanotify covers filesystem events.
func LatestPurgeCheckNamesForReducedDeep() []string {
return withLogicalOwnerPurgeNames(reducedDeepChecks())
}
// withLogicalOwnerPurgeNames extends the physical ownership expansion with the
// finding names of every logical owner hosted by a check in the set. Only the
// tier-level "everything this tier owns" views use it; the runner's per-cycle
// purge tracks each logical owner's completion individually instead.
func withLogicalOwnerPurgeNames(toScan []namedCheck) []string {
names := latestPurgeCheckNamesForChecks(toScan)
seen := make(map[string]struct{}, len(names))
for _, name := range names {
seen[name] = struct{}{}
}
for _, nc := range toScan {
for _, owner := range physicalCheckLogicalOwners[nc.name] {
for _, name := range logicalOwnerFindingNames[owner] {
if _, ok := seen[name]; !ok {
seen[name] = struct{}{}
names = append(names, name)
}
}
}
}
sort.Strings(names)
return names
}
// perRunFindingNames lists finding names that describe a single run's scan
// coverage rather than discovered state. They are purged whenever their
// owning check ran and returned, even when that run marked itself
// incomplete: a rolling scan on a large host never completes in one run, so
// gating these on completion would let every run's status finding accumulate
// forever. Discovered-state names (yara_match_scheduled) stay
// completion-gated so mid-cycle windows never wipe earlier windows' finds.
// An empty entry lets a stateful logical owner preserve its last finding
// without inventing a per-run status finding.
var perRunFindingNames = map[string][]string{
"email_weak_password": {"email_password_audit_incomplete"},
"yara_deep": {"yara_scan_incomplete"},
"db_content": {"db_content_scan_incomplete"},
logicalOwnerJSTaintDeep: {"js_taint_scan_incomplete"},
logicalOwnerPHPTaintDeep: {"php_taint_scan_incomplete"},
logicalOwnerReputationQuota: {},
logicalOwnerReputationFeedStale: {},
logicalOwnerWPCoreVerification: {},
"php_config_changes": {"php_config_scan_incomplete"},
}
var latestVolatileCheckNames = []string{
"account_scan_truncated",
"auto_block",
"auto_response",
"auto_response_paused",
"challenge_route",
"check_panic",
"check_timeout",
}
// Keep merge completion and health publication in the same order. The
// store lock alone cannot prevent an older caller publishing after a newer
// merge once both have returned from the store.
var latestScanMergeMu sync.Mutex
// StoreLatestScanFindings replaces the latest findings owned by a scan, then
// rebuilds derived correlation findings from the merged current set. One-shot
// auto-response actions stay in history and alerts, not the active findings
// view.
func StoreLatestScanFindings(st *state.Store, purgeChecks []string, findings []alert.Finding) {
StoreLatestScanFindingsWithGaps(st, purgeChecks, findings, nil)
}
// StoreLatestScanFindingsWithGaps preserves the latest state for files a
// completed scan could not examine while replacing its covered state. gapPaths
// contains the lexical and resolved aliases captured when each gap occurred;
// the state store applies that frozen set under the same lock as the purge.
func StoreLatestScanFindingsWithGaps(st *state.Store, purgeChecks []string, findings []alert.Finding, gapPaths map[string]map[string]bool) {
StoreLatestScanFindingsWithCoverage(st, purgeChecks, findings, &state.ScanCoverage{PreservePaths: gapPaths})
}
// StoreLatestScanFindingsWithCoverage retires only completed checks or scopes
// and preserves current findings for file gaps in the same atomic operation.
func StoreLatestScanFindingsWithCoverage(st *state.Store, purgeChecks []string, findings []alert.Finding, coverage *state.ScanCoverage) {
if st == nil {
return
}
if len(purgeChecks) == 0 && len(findings) == 0 && (coverage == nil || len(coverage.CompletedScopes) == 0) {
return
}
// Cold detection runs commands; correlation under latestMu must only
// read cached platform roots.
platform.Detect()
latestScanMergeMu.Lock()
var healthChecks []string
for _, check := range purgeChecks {
switch check {
case "php_taint_scan_incomplete", "js_taint_scan_incomplete", "yara_scan_incomplete",
"email_password_audit_incomplete":
healthChecks = append(healthChecks, check)
case "db_content_scan_incomplete":
// Only the host summary is a per-run condition. A partial run
// cannot resolve the existing per-install multisite limit.
summary := alert.Finding{Check: check, DedupKey: dbContentHostCoverageDedupKey}
present := false
for _, f := range findings {
if f.Key() == summary.Key() {
present = true
break
}
}
if !present {
st.RearmDismissedFindings([]string{summary.Key()})
}
}
}
st.RearmAbsentDedupFindings(healthChecks, findings)
now := time.Now()
st.PurgeAndMergeFindingsDerivedWithCoverage(
latestPurgeWithVolatile(purgeChecks),
latestPersistentFindings(findings),
coverage,
DerivedCorrelationChecks(),
func(merged []alert.Finding) []alert.Finding {
res := defaultCorrelator.Correlate(merged, now)
for i := range res.Derived {
if res.Derived[i].Timestamp.IsZero() {
res.Derived[i].Timestamp = now
}
}
return res.Derived
},
)
// Adding derived findings can evict source rows at the active-set cap.
// Count the final set, not the intermediate input to derivation.
unattributed := defaultCorrelator.Correlate(st.LatestFindings(), now).Unattributed
reporter := defaultUnattributedReporter
warnings := reporter.record(unattributed, true)
latestScanMergeMu.Unlock()
reporter.warnCounts(warnings)
}
func latestPurgeWithVolatile(purgeChecks []string) []string {
out := make([]string, 0, len(purgeChecks)+len(latestVolatileCheckNames))
out = append(out, purgeChecks...)
out = append(out, latestVolatileCheckNames...)
return out
}
func latestPersistentFindings(findings []alert.Finding) []alert.Finding {
out := make([]alert.Finding, 0, len(findings))
for _, f := range findings {
if isLatestVolatileFinding(f.Check) || isLatestDerivedFinding(f.Check) {
continue
}
out = append(out, f)
}
return out
}
func isLatestVolatileFinding(check string) bool {
for _, name := range latestVolatileCheckNames {
if check == name {
return true
}
}
return false
}
func isLatestDerivedFinding(check string) bool {
return IsDerivedCorrelationCheck(check)
}
func checksForTier(tier Tier) []namedCheck {
switch tier {
case TierCritical:
return criticalChecks()
case TierDeep:
return deepChecks()
case TierAll:
return append(criticalChecks(), deepChecks()...)
default:
return nil
}
}
func latestPurgeCheckNamesForChecks(toScan []namedCheck) []string {
ran := make(map[string]struct{}, len(toScan))
for _, nc := range toScan {
ran[nc.name] = struct{}{}
}
seen := make(map[string]struct{})
for _, nc := range toScan {
seen[nc.name] = struct{}{}
for _, name := range runnerFindingNames[nc.name] {
// A finding name owned by several checks (php_content and
// file_index both emit obfuscated_php) is purged only once all
// its owners ran this cycle; otherwise one check completing
// wipes the other's live findings.
if !allFindingOwnersRan(name, ran) {
continue
}
seen[name] = struct{}{}
}
}
names := make([]string, 0, len(seen))
for name := range seen {
names = append(names, name)
}
sort.Strings(names)
return names
}
// allFindingOwnersRan reports whether every check that emits finding name
// is in ran.
func allFindingOwnersRan(name string, ran map[string]struct{}) bool {
for check, names := range runnerFindingNames {
if _, ok := ran[check]; ok {
continue
}
for _, n := range names {
if n == name {
return false
}
}
}
return true
}
// RunTier runs only the specified tier of checks. The second return value
// is the per-scan purge name list (emitted finding aliases owned by the
// checks that actually executed this cycle); pass it to
// StoreLatestScanFindings so throttled-out checks keep their prior
// findings.
//
// Passes the requested dry-run state into runParallel via a scoped
// parameter rather than the previous package-level toggle, so concurrent
// periodic scanners no longer race with a manual `csm check` invocation.
func RunTier(cfg *config.Config, store *state.Store, tier Tier) ([]alert.Finding, []string) {
return RunTierWithContext(context.Background(), cfg, store, tier)
}
// RunTierWithContext is RunTier with a caller-owned parent context. Daemon
// periodic scans pass their shutdown context here so an interrupted scan does
// not stall process exit.
func RunTierWithContext(ctx context.Context, cfg *config.Config, store *state.Store, tier Tier) ([]alert.Finding, []string) {
return runParallelWithContext(ctx, cfg, store, checksForTier(tier), string(tier), false)
}
// RunTierDryRun is the dry-run variant of RunTier: auto-response actions
// are skipped. Used by `csm check*` socket commands and the legacy CLI.
func RunTierDryRun(cfg *config.Config, store *state.Store, tier Tier) ([]alert.Finding, []string) {
return RunTierDryRunWithContext(context.Background(), cfg, store, tier)
}
// RunTierDryRunWithContext is RunTierDryRun with a caller-owned parent
// context. Control-socket scans use it to collect coverage gaps without
// enabling auto-response actions.
func RunTierDryRunWithContext(ctx context.Context, cfg *config.Config, store *state.Store, tier Tier) ([]alert.Finding, []string) {
return runParallelWithContext(ctx, cfg, store, checksForTier(tier), string(tier), true)
}
// RunReducedDeep runs only the deep checks that fanotify can't replace.
// Used by the daemon when fanotify is active.
//
// Filesystem and content scans remain scheduled because fanotify misses
// renames, permission changes and files planted before daemon startup.
//
// The second return value is the per-scan purge name list scoped to the
// checks that actually executed this cycle.
func RunReducedDeep(cfg *config.Config, store *state.Store) ([]alert.Finding, []string) {
return RunReducedDeepWithContext(context.Background(), cfg, store)
}
// RunReducedDeepWithContext is RunReducedDeep with a caller-owned parent
// context for daemon shutdown cancellation.
func RunReducedDeepWithContext(ctx context.Context, cfg *config.Config, store *state.Store) ([]alert.Finding, []string) {
return runParallelWithContext(ctx, cfg, store, reducedDeepChecks(), string(TierDeep), false)
}
// RunAll runs critical checks always. Deep checks run if throttle allows or
// ForceAll is set. The second return value is the per-scan purge name list
// scoped to the checks that actually executed this cycle.
func RunAll(cfg *config.Config, store *state.Store) ([]alert.Finding, []string) {
return runAll(cfg, store, false)
}
// RunAllDryRun is the dry-run variant of RunAll for `csm baseline`.
func RunAllDryRun(cfg *config.Config, store *state.Store) ([]alert.Finding, []string) {
return runAll(cfg, store, true)
}
func runAll(cfg *config.Config, store *state.Store, dryRun bool) ([]alert.Finding, []string) {
toRun := criticalChecks()
if ForceAll || store.ShouldRunThrottled("deep_scan", cfg.Thresholds.DeepScanIntervalMin) {
toRun = append(toRun, deepChecks()...)
}
return runParallelWithContext(context.Background(), cfg, store, toRun, string(TierAll), dryRun)
}
// runParallel executes the supplied checks concurrently. It returns the
// emitted findings plus the per-scan purge name list. Throttled checks whose
// window has not elapsed stay out of the purge list so the previous cycle's
// findings persist; the same applies to checks that hit their per-check
// timeout, since a timed-out run produced no results to merge back.
// Disabled checks do not run, but their names stay in the purge list so
// disabling a check clears any findings it previously owned.
func runParallel(cfg *config.Config, store *state.Store, checks []namedCheck, tier string, dryRun bool) ([]alert.Finding, []string) {
return runParallelWithContext(context.Background(), cfg, store, checks, tier, dryRun)
}
// scansInFlight counts check runs in progress in this process: tier and
// reduced deep scans from any caller, account scans and scan jobs.
var scansInFlight atomic.Int64
// ScanInProgress reports whether any check run is in progress.
func ScanInProgress() bool {
return scansInFlight.Load() > 0
}
func runParallelWithContext(parent context.Context, cfg *config.Config, store *state.Store, checks []namedCheck, tier string, dryRun bool) ([]alert.Finding, []string) {
scansInFlight.Add(1)
defer scansInFlight.Add(-1)
if parent == nil {
parent = context.Background()
}
coverageGaps := coverageGapsFrom(parent)
// Clear a reused handle before work starts. An interrupted run must never
// expose the preceding run's path set as its own.
coverageGaps.replace(nil, nil, nil)
enabledChecks, disabledChecks := splitDisabledChecks(cfg, checks)
// Logical owners hosted by checks in this set: a disabled owner purges
// its names every cycle like a disabled check, while an enabled owner is
// purged only on its own completion mark, independent of its host's.
ownerOff := disabledLogicalOwners(cfg)
hostedOwners := make(map[string][]string)
disabledOwnerSet := make(map[string]struct{})
for _, nc := range checks {
for _, owner := range physicalCheckLogicalOwners[nc.name] {
if _, off := ownerOff[owner]; off {
disabledOwnerSet[owner] = struct{}{}
} else {
hostedOwners[nc.name] = append(hostedOwners[nc.name], owner)
}
}
}
for owner := range disabledOwnerSet {
if stateKey, ok := reputationHealthStateKeys[owner]; ok {
clearReputationHealthState(store, stateKey)
}
}
// Hide the caller's sink from nested runners and late check goroutines. The
// private path collector below is copied out only after this runner's workers
// completed inside their budgets.
scanCtx := context.WithValue(parent, coverageGapsContextKey{}, (*CoverageGaps)(nil))
scanCtx, truncations := withAccountScanTruncationCollector(scanCtx)
scanCtx, coveragePaths := withCoveragePathCollector(scanCtx)
scanCtx, incompleteChecks := withIncompleteCheckCollector(scanCtx)
// One WordPress discovery per cycle, shared by every check that needs it.
scanCtx = withWPInstallCache(scanCtx)
var mu sync.Mutex
var findings []alert.Finding
var wg sync.WaitGroup
// completedChecks collects the checks whose function returned within
// budget (guarded by mu; workers run concurrently). Only completed
// checks join the purge list: a timed-out check produced no results,
// so purging its names would wipe every finding from earlier cycles
// while merging nothing back.
completedChecks := make([]namedCheck, 0, len(enabledChecks))
completedOwners := make([]string, 0)
completedThrottled := make([]string, 0)
recoveredPanicKeys := make([]string, 0)
// incompleteRan collects checks that returned within budget but marked
// themselves incomplete; their per-run status finding names still purge.
incompleteRan := make([]string, 0)
coverageGapPaths := make(map[string]map[string]bool)
completedScopeNames := make(map[string]map[string]bool)
addCoverageGapPaths := func(owner string) {
paths := coveragePaths.gapPaths(owner)
if len(paths) == 0 {
return
}
findingNames := []string{owner}
findingNames = append(findingNames, runnerFindingNames[owner]...)
findingNames = append(findingNames, logicalOwnerFindingNames[owner]...)
for _, findingName := range findingNames {
if coverageGapPaths[findingName] == nil {
coverageGapPaths[findingName] = make(map[string]bool, len(paths))
}
for path := range paths {
coverageGapPaths[findingName][path] = true
}
}
}
// Concurrency comes from the host-wide budget, so a periodic tier and an
// operator-triggered account scan cannot each run their own full set of
// checks on the same cores.
budget := scanBudgetFrom(scanCtx)
dispatches := checkDispatches.begin(len(enabledChecks), budget)
checkDispatches.observe(scanCtx, dispatches)
for i, nc := range enabledChecks {
wg.Add(1)
c := nc
task := dispatches[i]
// Check functions run against user filesystem content (unparsed PHP,
// crafted archives, foreign encodings), so contain both a panic in the
// runner and one inside the check execution. The inner recovery reports
// check_panic immediately with a stack trace.
obs.SafeGo("check-runner", task.wrap(func() {
defer wg.Done()
if !task.admit(scanCtx) {
task.withdraw(scanCtx)
return
}
if scanCtx.Err() != nil {
task.withdraw(scanCtx)
return
}
throttleReserved := false
if min, ok := checkThrottleMin[c.name]; ok && store != nil {
// Reserve atomically right before execution. The success
// stamp still happens only after the scan completes, so
// timeouts and interrupted scans can retry without allowing
// concurrent scans to start the same expensive check in the
// meantime.
if !store.ReserveThrottle(c.name, min) {
return
}
throttleReserved = true
}
// Run with cancellable context so timed-out checks stop
budget := timeoutFor(c.name)
ctx, cancel := context.WithTimeout(withCheckDispatch(scanCtx, task), budget)
start := time.Now()
execution := executeCheckAsync(ctx, "check-exec", func() []alert.Finding {
return c.fn(ctx, cfg, store)
})
defer execution.finishCaller()
select {
case outcome := <-execution.done:
execution.received()
if outcome.panicErr != "" {
cancel()
if throttleReserved {
store.ReleaseThrottle(c.name)
}
observeCheckDuration(c.name, tier, time.Since(start))
mu.Lock()
// The stack trace carries goroutine ids and addresses that
// differ on every run; a check that panics each cycle is
// one ongoing condition.
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "check_panic",
Message: fmt.Sprintf("Check '%s' stopped after an internal panic", c.name),
Details: outcome.panicErr,
DedupKey: "check:" + c.name,
Timestamp: time.Now(),
})
mu.Unlock()
return
}
results := outcome.findings
if ctx.Err() != nil {
execution.withdraw(ctx.Err())
cancel()
if throttleReserved {
store.ReleaseThrottle(c.name)
}
observeCheckDuration(c.name, tier, time.Since(start))
if scanCtx.Err() != nil {
return
}
mu.Lock()
// A heavy scan returns what it found before the budget ran
// out. The check stays out of completedChecks either way, so
// these findings merge into the latest set without
// authorizing a purge of anything the run never reached.
findings = append(findings, results...)
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "check_timeout",
Message: fmt.Sprintf("Check '%s' timed out after %s", c.name, budget),
Timestamp: time.Now(),
})
mu.Unlock()
return
}
cancel()
observeCheckDuration(c.name, tier, time.Since(start))
// An internal interval skip did not run an audit. Like a runner
// throttle skip, it cannot replace per-run coverage summaries
// or acknowledge recovery from an earlier failure.
if incompleteChecks.wasSkipped(c.name) {
if throttleReserved {
store.ReleaseThrottle(c.name)
}
return
}
mu.Lock()
recoveredPanicKeys = append(recoveredPanicKeys, (alert.Finding{Check: "check_panic", DedupKey: "check:" + c.name}).Key())
if scopes := coveragePaths.completedScopes(c.name); len(scopes) > 0 {
names := append([]string{c.name}, runnerFindingNames[c.name]...)
for _, name := range names {
completedScopeNames[name] = scopes
}
}
if !incompleteChecks.contains(c.name) {
completedChecks = append(completedChecks, c)
addCoverageGapPaths(c.name)
} else {
incompleteRan = append(incompleteRan, c.name)
}
// A hosted logical owner completes or stays partial on its own
// mark, honored only because the physical check returned in
// budget without a panic; a YARA-side failure must never purge
// JS findings the run did not re-emit, and vice versa.
for _, owner := range hostedOwners[c.name] {
if !incompleteChecks.contains(owner) {
completedOwners = append(completedOwners, owner)
addCoverageGapPaths(owner)
} else {
incompleteRan = append(incompleteRan, owner)
}
}
if len(results) > 0 {
findings = append(findings, results...)
}
if throttleReserved && incompleteChecks.contains(c.name) {
store.ReleaseThrottle(c.name)
} else if throttleReserved {
completedThrottled = append(completedThrottled, c.name)
}
mu.Unlock()
case <-ctx.Done():
execution.withdraw(ctx.Err())
// The deadline and the check's own return race: a heavy scan
// that honors ctx hands back its partial findings just after
// the deadline fires, and select picks whichever is ready
// first. Wait a bounded moment for them instead of discarding
// detections the scan already paid for.
var drained []alert.Finding
grace := time.NewTimer(checkTimeoutDrainGrace)
select {
case outcome := <-execution.done:
// A panic carries no findings, and the deadline is what the
// operator needs to see either way.
if outcome.panicErr == "" {
drained = outcome.findings
}
case <-grace.C:
case <-scanCtx.Done():
// Shutdown has no partial findings to publish and must
// not wait for checks that ignore cancellation.
}
grace.Stop()
cancel()
if throttleReserved {
store.ReleaseThrottle(c.name)
}
observeCheckDuration(c.name, tier, time.Since(start))
// Distinguish a real per-check timeout from a daemon shutdown.
// On shutdown the parent scan context is cancelled, so abort
// quietly rather than flooding the scan with bogus check_timeout
// findings for work that was deliberately interrupted.
if scanCtx.Err() != nil {
return
}
mu.Lock()
// The check stays out of completedChecks, so these merge into
// the latest set without authorizing a purge of the range the
// run never reached.
findings = append(findings, drained...)
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "check_timeout",
Message: fmt.Sprintf("Check '%s' timed out after %s", c.name, budget),
Timestamp: time.Now(),
})
mu.Unlock()
}
}))
}
wg.Wait()
if scanCtx.Err() != nil {
for _, name := range completedThrottled {
store.ReleaseThrottle(name)
}
// The scan did not complete, so keep the caller from replacing the
// last completed scan state with partial findings or an empty purge.
return nil, nil
}
for _, name := range completedThrottled {
store.MarkThrottledRan(name)
}
if store != nil {
store.RearmFindings(recoveredPanicKeys)
}
purgeChecks := make([]namedCheck, 0, len(disabledChecks)+len(completedChecks))
purgeChecks = append(purgeChecks, disabledChecks...)
purgeChecks = append(purgeChecks, completedChecks...)
now := time.Now()
findings = append(findings, truncations.findings(now)...)
for i := range findings {
if findings[i].Timestamp.IsZero() {
findings[i].Timestamp = now
}
}
// Cross-account correlation
if len(findings) > 0 {
platform.Detect()
}
correlated := CorrelateBatchFindings(findings)
for i := range correlated.Derived {
if correlated.Derived[i].Timestamp.IsZero() {
correlated.Derived[i].Timestamp = now
}
}
findings = append(findings, correlated.Derived...)
ReportUnattributedCorrelation(correlated.Unattributed)
// Auto-response: skip when the caller requested a dry run
// (check/baseline commands).
if !dryRun {
// Challenge routing and hard-blocking are resolved together, in order,
// so an eligible IP lands on the challenge list before the block stage
// checks membership. The block actions are appended in their own slot
// below (after the kill/quarantine/htaccess actions) to keep the
// emitted ordering stable; no intervening stage acts on action findings.
challengeActions, blockActions := ChallengeThenBlock(cfg, findings)
for i := range challengeActions {
if challengeActions[i].Timestamp.IsZero() {
challengeActions[i].Timestamp = now
}
}
findings = append(findings, challengeActions...)
killActions := AutoKillProcesses(parent, cfg, findings)
for i := range killActions {
if killActions[i].Timestamp.IsZero() {
killActions[i].Timestamp = now
}
}
observeAutoResponse("kill", len(killActions))
findings = append(findings, killActions...)
quarantineActions := AutoQuarantineFiles(cfg, findings)
for i := range quarantineActions {
if quarantineActions[i].Timestamp.IsZero() {
quarantineActions[i].Timestamp = now
}
}
observeFileResponseActions("quarantine", quarantineActions)
findings = append(findings, quarantineActions...)
htaccessActions := AutoCleanHtaccess(cfg, findings)
for i := range htaccessActions {
if htaccessActions[i].Timestamp.IsZero() {
htaccessActions[i].Timestamp = now
}
}
observeFileResponseActions("htaccess_clean", htaccessActions)
findings = append(findings, htaccessActions...)
vpatchActions := AutoVirtualPatchExposedFiles(cfg, findings)
for i := range vpatchActions {
if vpatchActions[i].Timestamp.IsZero() {
vpatchActions[i].Timestamp = now
}
}
observeAutoResponse("virtual_patch_exposed", len(vpatchActions))
findings = append(findings, vpatchActions...)
for i := range blockActions {
if blockActions[i].Timestamp.IsZero() {
blockActions[i].Timestamp = now
}
}
observeAutoResponse("block", len(blockActions))
findings = append(findings, blockActions...)
}
purgeNames := latestPurgeCheckNamesForChecks(purgeChecks)
for _, owner := range completedOwners {
purgeNames = append(purgeNames, logicalOwnerFindingNames[owner]...)
}
for owner := range disabledOwnerSet {
purgeNames = append(purgeNames, logicalOwnerFindingNames[owner]...)
}
// Selected checks that never completed (including timeouts, panics and
// throttle skips) did not examine their prior findings. Reuse the known
// completion sets to protect that state without authorizing any purge.
uncompletedOwners := make(map[string]bool)
for _, check := range enabledChecks {
uncompletedOwners[check.name] = true
for _, owner := range hostedOwners[check.name] {
uncompletedOwners[owner] = true
}
}
for _, check := range completedChecks {
delete(uncompletedOwners, check.name)
}
for _, owner := range completedOwners {
delete(uncompletedOwners, owner)
}
incompleteFindingNames := make(map[string]bool)
for owner := range uncompletedOwners {
names := append([]string{owner}, runnerFindingNames[owner]...)
names = append(names, logicalOwnerFindingNames[owner]...)
for _, name := range names {
incompleteFindingNames[name] = true
}
}
// Only checks that returned replace their per-run coverage summaries.
for _, owner := range incompleteRan {
for _, name := range perRunFindingNames[owner] {
delete(incompleteFindingNames, name)
}
}
coverageGaps.replace(coverageGapPaths, completedScopeNames, incompleteFindingNames)
return findings, mergePerRunPurgeNames(purgeNames, incompleteRan)
}
// mergePerRunPurgeNames preserves the purge list's sorted, duplicate-free
// contract. incompleteRan is populated by concurrent workers, so appending it
// directly made the returned order depend on goroutine completion order.
func mergePerRunPurgeNames(purgeNames, incompleteRan []string) []string {
seen := make(map[string]struct{}, len(purgeNames)+len(incompleteRan))
for _, name := range purgeNames {
seen[name] = struct{}{}
}
for _, owner := range incompleteRan {
for _, name := range perRunFindingNames[owner] {
seen[name] = struct{}{}
}
}
merged := make([]string, 0, len(seen))
for name := range seen {
merged = append(merged, name)
}
sort.Strings(merged)
return merged
}
package checks
import (
"context"
"errors"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type scanBatchMonitor struct {
mu sync.Mutex
batches map[*scanBatch]struct{}
losses *queuehealth.Tracker
}
type scanBatch struct {
monitor *scanBatchMonitor
tasks []*scanBatchTask
pending map[*scanBatchTask]struct{}
parallel int
progress time.Time
}
// Mutable batch and task state is guarded by the monitor mutex.
type scanBatchTask struct {
batch *scanBatch
started, deadline time.Time
failed bool
}
func newScanBatchMonitor() *scanBatchMonitor {
return &scanBatchMonitor{batches: make(map[*scanBatch]struct{}), losses: queuehealth.New(0, time.Minute)}
}
func (m *scanBatchMonitor) begin(count, parallel int) *scanBatch {
b := &scanBatch{monitor: m, tasks: make([]*scanBatchTask, count), pending: make(map[*scanBatchTask]struct{}, count), parallel: parallel, progress: time.Now()}
for i := range b.tasks {
b.tasks[i] = &scanBatchTask{batch: b}
b.pending[b.tasks[i]] = struct{}{}
}
if count > 0 {
m.mu.Lock()
m.batches[b] = struct{}{}
m.mu.Unlock()
}
return b
}
func (t *scanBatchTask) admit() {
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
t.started = time.Now()
t.deadline = t.started.Add(time.Minute)
t.batch.progress = t.started
}
func (t *scanBatchTask) executing(ctx context.Context, budget time.Duration) {
deadline := time.Now().Add(budget)
if parent, ok := ctx.Deadline(); ok && parent.Before(deadline) {
deadline = parent
}
m := t.batch.monitor
m.mu.Lock()
t.deadline = deadline
m.mu.Unlock()
}
func (t *scanBatchTask) progress() {
m := t.batch.monitor
m.mu.Lock()
t.deadline = time.Now().Add(time.Minute)
m.mu.Unlock()
}
func (t *scanBatchTask) failLocked() {
if !t.failed {
t.failed = true
t.batch.monitor.losses.Lose(time.Now(), 1)
}
}
func (t *scanBatchTask) fail() {
m := t.batch.monitor
m.mu.Lock()
t.failLocked()
m.mu.Unlock()
}
func (t *scanBatchTask) withdraw(err error) {
if errors.Is(err, context.DeadlineExceeded) {
t.fail()
}
}
func (t *scanBatchTask) finish() {
m := t.batch.monitor
m.mu.Lock()
defer m.mu.Unlock()
delete(t.batch.pending, t)
t.batch.progress = time.Now()
if len(t.batch.pending) == 0 {
delete(m.batches, t.batch)
}
}
// The callback reports failures and deadline withdrawal where work is
// abandoned. A later deadline cannot undo a successfully completed callback.
func (t *scanBatchTask) run(ctx context.Context, budget time.Duration, fn func()) {
completed := false
defer func() {
if !completed {
t.fail()
}
t.finish()
}()
t.executing(ctx, budget)
fn()
completed = true
}
// abandon runs after dispatch has stopped. Running workers retain their own
// tasks, including when a caller has canceled but an operation ignores it.
func (b *scanBatch) abandon(ctx context.Context) {
err := ctx.Err()
m := b.monitor
m.mu.Lock()
defer m.mu.Unlock()
for task := range b.pending {
if !task.started.IsZero() {
continue
}
if !errors.Is(err, context.Canceled) {
task.failLocked()
}
delete(b.pending, task)
}
if len(b.pending) == 0 {
delete(m.batches, b)
}
}
func (m *scanBatchMonitor) snapshot(now time.Time) queuehealth.Status {
m.mu.Lock()
defer m.mu.Unlock()
s := m.losses.Snapshot(now)
s.CapacityUnavailable = true
s.LagBasis = "consumer_progress"
var waitingLate, runningLate bool
for b := range m.batches {
waiting, running := 0, 0
for task := range b.pending {
if task.started.IsZero() {
waiting++
} else {
running++
s.ProcessingSeconds = max(s.ProcessingSeconds, now.Sub(task.started).Seconds())
runningLate = runningLate || !now.Before(task.deadline)
}
}
s.Depth += waiting
s.InFlight += running
if waiting > 0 {
lag := now.Sub(b.progress)
s.LagSeconds = max(s.LagSeconds, lag.Seconds())
// Filling a finite batch is expected while every worker is busy.
// Its own deadlines bound that work; a free slot needs progress.
waitingLate = waitingLate || (running < b.parallel && lag >= time.Minute)
}
}
switch {
case runningLate:
s.Reason = "processing_lag"
case waitingLate:
s.Reason = "backlog_lag"
}
if s.Reason != "" {
s.Status = "degraded"
}
return s
}
package checks
import (
"context"
"runtime"
"sync/atomic"
)
const (
// minScanParallelism keeps a single-core host making progress across the
// check set instead of serialising a scan behind one slow filesystem walk.
minScanParallelism = 2
// maxScanParallelism is the ceiling every scan path used to apply
// unconditionally. Checks walk account trees and hash file content, so
// past this point they starve each other on I/O rather than finishing
// sooner, and the Web UI stops answering.
maxScanParallelism = 5
)
// hostScanBudget bounds how many CPU-heavy checks run at once across the whole
// daemon. The tier runner and the per-account scan used to carry separate
// fixed limits, so an operator-triggered account scan during a periodic tier
// ran nine checks together on a four-core host.
var hostScanBudget = newScanBudget(scanParallelismFor(runtime.NumCPU()))
type scanBudgetKey struct{}
// withScanBudget scopes a scan to its own budget. Production scans carry none
// and share the host budget; tests use this so one test's checks cannot hold
// slots another test is waiting for.
func withScanBudget(ctx context.Context, slots int) context.Context {
return context.WithValue(ctx, scanBudgetKey{}, newScanBudget(slots))
}
// scanBudgetFrom returns the budget this scan draws from.
func scanBudgetFrom(ctx context.Context) *scanBudget {
if ctx != nil {
if budget, ok := ctx.Value(scanBudgetKey{}).(*scanBudget); ok && budget != nil {
return budget
}
}
return hostScanBudget
}
// scanParallelismFor sizes the budget from the host's core count, bounded at
// both ends.
func scanParallelismFor(cpus int) int {
if cpus < minScanParallelism {
return minScanParallelism
}
if cpus > maxScanParallelism {
return maxScanParallelism
}
return cpus
}
// scanBudget is a counting semaphore. Callers acquire one slot per check and
// must not hold one while acquiring another: nothing in the scan paths nests,
// and a nested acquisition would deadlock against a full budget.
type scanBudget struct {
slots chan struct{}
}
func newScanBudget(slots int) *scanBudget {
if slots < 1 {
slots = 1
}
return &scanBudget{slots: make(chan struct{}, slots)}
}
// acquire blocks until a slot is free or ctx is done. The runner and its
// asynchronous execution share ownership so cancellation cannot free a slot
// while a check that ignores its context is still consuming resources.
func (b *scanBudget) acquire(ctx context.Context) *scanSlot {
select {
case b.slots <- struct{}{}:
slot := &scanSlot{budget: b}
slot.owners.Store(1)
return slot
case <-ctx.Done():
return nil
}
}
func (b *scanBudget) size() int { return cap(b.slots) }
func (b *scanBudget) hasCapacity() bool { return len(b.slots) < b.size() }
type scanSlot struct {
budget *scanBudget
owners atomic.Int32
}
func (s *scanSlot) retain() {
if s != nil {
s.owners.Add(1)
}
}
func (s *scanSlot) release() {
if s != nil && s.owners.Add(-1) == 0 {
<-s.budget.slots
}
}
package checks
import (
"context"
"github.com/pidginhost/csm/internal/config"
)
const (
fullScanMaxFileMBDefault = 16
fullScanMaxFileMBMax = 4096
)
// scanForceContent reports whether the current scan should bypass the
// clean-file content cache (phpcontentcache.json). True only when ctx carries
// AccountScanOptions with ForceContent=true, i.e. an explicit full-scan audit.
// Normal scheduled scans and any context without options return false.
func scanForceContent(ctx context.Context) bool {
opts, ok := ScanOptionsFromContext(ctx)
return ok && opts.ForceContent
}
// scanForceFileIndex reports whether the current scan is a file-index audit
// run. True only when ctx carries AccountScanOptions with ForceFileIndex=true.
// In audit mode CheckFileIndex enumerates only the in-scope account, bypasses
// the directory mtime cache, and writes none of the three live state files
// (fileindex.current, fileindex.previous, dircache.json).
func scanForceFileIndex(ctx context.Context) bool {
opts, ok := ScanOptionsFromContext(ctx)
return ok && opts.ForceFileIndex
}
// scanRespectsIgnores reports whether the current scan should honour
// cfg.Suppressions.IgnorePaths. When ctx carries AccountScanOptions with
// RespectIgnores=false (i.e. an explicit full-scan / audit request), the caller
// wants full coverage and ignore_paths is bypassed. Normal scheduled scans and
// any call without options carry RespectIgnores=true (the safe default).
func scanRespectsIgnores(ctx context.Context, _ *config.Config) bool {
if opts, ok := ScanOptionsFromContext(ctx); ok {
return opts.RespectIgnores
}
return true
}
// scanMaxFileBytes returns the per-file byte limit for a full-scan context.
// When ctx carries AccountScanOptions, the options MaxFileBytes value is
// returned so an oversized file can be skipped with a warning. Returns 0 for
// all other contexts so normal scheduled scans are completely unaffected.
func scanMaxFileBytes(ctx context.Context) int64 {
opts, ok := ScanOptionsFromContext(ctx)
if !ok {
return 0
}
return opts.MaxFileBytes
}
// accountScanMaxFiles returns the effective file cap for the current scan. When
// ctx carries AccountScanOptions (i.e. the caller is RunAccountScanWithOptions),
// the options MaxFiles value wins so a full scan (MaxFiles=0) is uncapped.
// Otherwise it falls back to the config-derived value so normal scheduled scans
// are unaffected.
func accountScanMaxFiles(ctx context.Context, cfg *config.Config) int {
if opts, ok := ScanOptionsFromContext(ctx); ok {
return opts.MaxFiles
}
return effectiveAccountScanMaxFiles(cfg)
}
// AccountScanOptions controls how RunAccountScanWithOptions enumerates and
// content-scans an account. The zero value is NOT the default: MaxFiles 0 means
// uncapped. Callers use DefaultAccountScanOptions for normal behaviour.
type AccountScanOptions struct {
MaxFiles int // 0 = uncapped path ranking
ForceContent bool // true = bypass clean-file content caches
ForceFileIndex bool // true = bypass file-index dir mtime cache, do not write live index
RespectIgnores bool // false = also scan suppressions.ignore_paths
MaxFileBytes int64 // 0 = use each check's existing per-file limit
}
// DefaultAccountScanOptions returns the options that reproduce today's
// RunAccountScan behaviour. All callers that want the existing cap and cache
// semantics should use this rather than constructing AccountScanOptions directly.
func DefaultAccountScanOptions(cfg *config.Config) AccountScanOptions {
return AccountScanOptions{
MaxFiles: effectiveAccountScanMaxFiles(cfg),
ForceContent: false,
ForceFileIndex: false,
RespectIgnores: true,
MaxFileBytes: 0,
}
}
type scanOptionsKey struct{}
// ContextWithScanOptions attaches opts to ctx so that helpers called during a
// scan can read the active options without threading a parameter through every
// call site. Use ScanOptionsFromContext to retrieve them.
func ContextWithScanOptions(ctx context.Context, opts AccountScanOptions) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, scanOptionsKey{}, opts)
}
// ScanOptionsFromContext retrieves the AccountScanOptions stored by
// ContextWithScanOptions. ok is false when the context carries no options,
// which callers should treat as "use defaults".
func ScanOptionsFromContext(ctx context.Context) (AccountScanOptions, bool) {
if ctx == nil {
return AccountScanOptions{}, false
}
opts, ok := ctx.Value(scanOptionsKey{}).(AccountScanOptions)
return opts, ok
}
// FullScanOptions builds the canonical option set for an uncapped full-scan
// audit job: no file cap, force content + file-index (bypass the clean-file and
// directory-mtime caches), and the configured per-file byte ceiling.
// respectIgnores is the only caller-chosen knob. Centralised here so the control
// handler and the WebUI enqueue path cannot drift on this security-relevant set.
func FullScanOptions(cfg *config.Config, respectIgnores bool) AccountScanOptions {
return AccountScanOptions{
MaxFiles: 0,
ForceContent: true,
ForceFileIndex: true,
RespectIgnores: respectIgnores,
MaxFileBytes: FullScanMaxFileBytes(cfg),
}
}
// FullScanMaxFileBytes converts cfg.Thresholds.FullScanMaxFileMB to a byte
// limit for per-file content reads during a full-scan job. A configured value
// of 0, negative, or above the validated maximum falls back to 16 MiB so the
// caller never gets an unconstrained or overflowed limit from a bad field.
func FullScanMaxFileBytes(cfg *config.Config) int64 {
mb := fullScanMaxFileMBDefault
if cfg != nil && cfg.Thresholds.FullScanMaxFileMB > 0 && cfg.Thresholds.FullScanMaxFileMB <= fullScanMaxFileMBMax {
mb = cfg.Thresholds.FullScanMaxFileMB
}
return int64(mb) * 1024 * 1024
}
package checks
import (
"bytes"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"os"
"sync"
"time"
"github.com/pidginhost/csm/internal/state"
)
// selfWriteTTL bounds how long a CSM-performed write to a sensitive path
// suppresses the sensitive-file detectors. Short enough that an independent
// tamper layered on later is still caught by the next scan.
const selfWriteTTL = 15 * time.Minute
const durableSelfWriteVersion = 1
var (
selfWriteMu sync.Mutex
selfWrites = map[string]selfWriteRecord{}
selfWriteNow = time.Now // overridable in tests
// selfWriteStore persists self-write content and file identity across
// restarts. Nil in unit tests and one-shot CLI runs, where the in-memory
// ledger is enough.
selfWriteStore *state.Store
)
// selfWriteKey namespaces the durable record. The leading underscore marks it
// as housekeeping so the state sweeper never evicts it; the record is cleared
// explicitly by forgetSelfWrites, or superseded by the next write to the path.
func selfWriteKey(path string) string { return "_selfwrite:" + path }
// SetSelfWriteStore gives the self-write ledger somewhere durable to record
// what CSM wrote. Without it, a daemon restart -- or a crontab the cPanel
// wrapper reformats after CSM hands it over -- makes CSM's own write look like
// a third-party change to the sensitive-file detectors.
func SetSelfWriteStore(st *state.Store) {
selfWriteMu.Lock()
defer selfWriteMu.Unlock()
selfWriteStore = st
}
type selfWriteRecord struct {
hash string
identity *selfWriteFileIdentity
expires time.Time
}
type selfWriteFileIdentity struct {
Device uint64 `json:"device"`
Inode uint64 `json:"inode"`
ChangeSec int64 `json:"change_sec"`
ChangeNsec int64 `json:"change_nsec"`
}
type durableSelfWriteRecord struct {
Version int `json:"version"`
Hash string `json:"hash"`
Identity selfWriteFileIdentity `json:"identity"`
}
// RecordSelfWrite registers that CSM remediation just wrote content to a
// sensitive watched file. A current file is also recorded durably with its
// identity; a provisional record made before the write stays memory-only.
func RecordSelfWrite(path string, content []byte) {
sum := sha256.Sum256(content)
hash := hex.EncodeToString(sum[:])
identity, hasIdentity := currentSelfWriteIdentity(path, content)
selfWriteMu.Lock()
defer selfWriteMu.Unlock()
now := selfWriteNow()
pruneExpiredSelfWritesLocked(now)
rec := selfWriteRecord{
hash: hash,
expires: now.Add(selfWriteTTL),
}
if hasIdentity {
rec.identity = &identity
}
selfWrites[path] = rec
if selfWriteStore != nil && hasIdentity {
// Marshal of a fixed struct of scalars cannot fail; if it somehow did,
// losing the durable record costs a redundant finding, not safety, so
// remediation must not die here.
raw, err := json.Marshal(durableSelfWriteRecord{
Version: durableSelfWriteVersion,
Hash: hash,
Identity: identity,
})
if err != nil {
fmt.Fprintf(os.Stderr, "self-write: encode %s: %v\n", path, err)
return
}
if err := selfWriteStore.SetRawAndSave(selfWriteKey(path), string(raw)); err != nil {
fmt.Fprintf(os.Stderr, "self-write: persist %s: %v\n", path, err)
}
}
}
func forgetSelfWrites(paths ...string) {
selfWriteMu.Lock()
defer selfWriteMu.Unlock()
for _, path := range paths {
delete(selfWrites, path)
if selfWriteStore != nil {
if err := selfWriteStore.DeleteRawAndSave(selfWriteKey(path)); err != nil {
fmt.Fprintf(os.Stderr, "self-write: forget %s: %v\n", path, err)
}
}
}
}
// isExpectedSelfWrite reports whether content and file identity still match a
// CSM self-write. Provisional in-memory records expire; identity-bound durable
// records do not.
func isExpectedSelfWrite(path string, content []byte) bool {
selfWriteMu.Lock()
defer selfWriteMu.Unlock()
now := selfWriteNow()
pruneExpiredSelfWritesLocked(now)
sum := sha256.Sum256(content)
got := hex.EncodeToString(sum[:])
if rec, ok := selfWrites[path]; ok {
if got != rec.hash {
return false
}
// A provisional record is installed immediately before CSM invokes
// crontab so a synchronous filesystem event can be suppressed. Once
// the write finishes, RecordSelfWrite replaces it with file identity.
if rec.identity == nil {
return len(content) > 0
}
identity, ok := currentSelfWriteIdentity(path, content)
return ok && identity == *rec.identity
}
if selfWriteStore == nil {
return false
}
raw, ok := selfWriteStore.GetRaw(selfWriteKey(path))
if !ok {
return false
}
var want durableSelfWriteRecord
if err := json.Unmarshal([]byte(raw), &want); err != nil ||
want.Version != durableSelfWriteVersion ||
want.Hash == "" ||
got != want.Hash {
return false
}
identity, ok := currentSelfWriteIdentity(path, content)
return ok && identity == want.Identity
}
func pruneExpiredSelfWritesLocked(now time.Time) {
for path, rec := range selfWrites {
if now.After(rec.expires) {
delete(selfWrites, path)
}
}
}
func currentSelfWriteIdentity(path string, content []byte) (selfWriteFileIdentity, bool) {
f, err := osFS.Open(path)
if err != nil {
return selfWriteFileIdentity{}, false
}
defer func() { _ = f.Close() }()
got, err := io.ReadAll(io.LimitReader(f, int64(len(content))+1))
if err != nil || !bytes.Equal(got, content) {
return selfWriteFileIdentity{}, false
}
info, err := f.Stat()
if err != nil || !info.Mode().IsRegular() {
return selfWriteFileIdentity{}, false
}
return selfWriteIdentityFromFileInfo(info)
}
//go:build linux
package checks
import (
"os"
"syscall"
)
func selfWriteIdentityFromFileInfo(info os.FileInfo) (selfWriteFileIdentity, bool) {
st, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return selfWriteFileIdentity{}, false
}
return selfWriteFileIdentity{
Device: uint64(st.Dev),
Inode: st.Ino,
ChangeSec: st.Ctim.Sec,
ChangeNsec: st.Ctim.Nsec,
}, true
}
//go:build linux
package checks
import (
"os"
"syscall"
)
func sensitiveFileOwnership(info os.FileInfo) (uid, gid uint32, ok bool) {
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return 0, 0, false
}
return stat.Uid, stat.Gid, true
}
package checks
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// sensitiveWatchset is the static set of system-configuration paths CSM
// raises a finding on when any of them is opened for write. The set is
// intentionally narrow and not operator-configurable: an attacker who
// learns that a path is excluded gets a free landing pad.
//
// Glob entries expand at runtime; non-glob entries appear once.
var sensitiveWatchset = []string{
"/etc/shadow",
"/etc/gshadow",
"/etc/passwd",
"/etc/group",
"/etc/sudoers",
"/etc/sudoers.d/*",
"/etc/ssh/sshd_config",
"/etc/ssh/sshd_config.d/*",
"/etc/cron.d/*",
"/etc/cron.hourly/*",
"/etc/cron.daily/*",
"/etc/cron.weekly/*",
"/etc/cron.monthly/*",
"/var/spool/cron/*",
// Debian cron keeps user crontabs one level deeper than cronie.
"/var/spool/cron/crontabs/*",
}
const sensitiveFileBaselineKey = "_sensitive_file_hash:__baseline_complete"
// ExpandWatchset returns the absolute paths in the watchset, with globs
// expanded against the given filesystem root. Non-existent paths drop
// silently; the next refresh picks them up once they are created. root
// is "/" in production and a t.TempDir in tests.
func ExpandWatchset(root string) []string {
var out []string
for _, pat := range sensitiveWatchset {
full := filepath.Join(root, pat)
if strings.ContainsAny(pat, "*?[") {
matches, _ := filepath.Glob(full)
out = append(out, matches...)
continue
}
out = append(out, full)
}
return out
}
// classifySensitive returns a stable kind label for a watchset path so
// findings can vary their severity and message.
func classifySensitive(path string) string {
switch filepath.Base(path) {
case "shadow", "gshadow", "passwd", "group":
return "auth"
case "sshd_config":
return "sshd"
case "sudoers":
return "sudo"
}
dir := filepath.Dir(path)
if strings.Contains(dir, "/cron") || strings.Contains(dir, "/spool/cron") {
return "cron"
}
if strings.Contains(dir, "/sudoers.d") {
return "sudo"
}
if strings.Contains(dir, "/sshd_config.d") {
return "sshd"
}
return ""
}
// EvaluateSensitiveFileWrite returns a populated alert.Finding and true when
// the BPF live backend observed a write to a watchset path. It reads current
// content for content-bound self-write suppression.
// Returns false for paths classifySensitive does not recognise -- the BPF
// program already filters via its dev+inode map, but this guards against
// stale map entries pointing at unrelated files.
func EvaluateSensitiveFileWrite(path string, uid, pid uint32, comm string) (alert.Finding, bool) {
content, err := osFS.ReadFile(path)
return EvaluateSensitiveFileWriteSnapshot(path, uid, pid, comm, content, err == nil)
}
// EvaluateSensitiveFileWriteSnapshot evaluates a live write against bytes
// captured with the file state that will be recorded for refresh deduplication.
// Keeping those two observations together prevents a concurrent rename from
// making the finding describe or suppress different content.
func EvaluateSensitiveFileWriteSnapshot(path string, uid, pid uint32, comm string, content []byte, contentKnown bool) (alert.Finding, bool) {
kind := classifySensitive(path)
if kind == "" {
return alert.Finding{}, false
}
// Suppress writes CSM itself just performed (e.g. installing a per-user
// wp-cron). Content-bound: a tamper layered on top changes the hash and
// is still reported.
if contentKnown && isExpectedSelfWrite(path, content) {
return alert.Finding{}, false
}
// Durable complement to the TTL-bounded self-write ledger: a user crontab
// holding only CSM-installed WP-Cron jobs is CSM's own managed block, not
// attacker persistence, even after a restart cleared the ledger.
if suppressedAsManagedWPCron(path, content) {
return alert.Finding{}, false
}
sev := alert.High
if uid != 0 {
sev = alert.Critical
}
now := time.Now()
f := alert.Finding{
Severity: sev,
Check: "sensitive_file_modified",
Message: fmt.Sprintf("Write to sensitive system file: %s (uid=%d)", path, uid),
Details: fmt.Sprintf("Class: %s, PID: %d, Comm: %s, User: %s", kind, pid, comm, LookupUser(uid)),
FilePath: path,
Timestamp: now,
}
return rescoreSensitive(f, kind, content, pid, now), true
}
// EvaluateSensitiveFileAppearance returns a finding when a path no previous
// watchset refresh had seen shows up -- a genuinely new cron drop-in, sudoers
// fragment, or user crontab.
func EvaluateSensitiveFileAppearance(path string) (alert.Finding, bool) {
content, err := osFS.ReadFile(path)
return evaluateSensitiveWatchsetChange(path, "New sensitive system file appeared", content, err == nil)
}
// evaluateSensitiveWatchsetChange builds the finding both refresh-diff outcomes
// share. Callers pass the bytes that produced the compared digest so a later
// rewrite cannot change self-write suppression or cron-content scoring.
func evaluateSensitiveWatchsetChange(path, summary string, content []byte, contentKnown bool) (alert.Finding, bool) {
kind := classifySensitive(path)
if kind == "" {
return alert.Finding{}, false
}
if contentKnown {
if isExpectedSelfWrite(path, content) {
return alert.Finding{}, false
}
}
if suppressedAsManagedWPCron(path, content) {
return alert.Finding{}, false
}
now := time.Now()
f := alert.Finding{
Severity: alert.High,
Check: "sensitive_file_modified",
Message: fmt.Sprintf("%s: %s", summary, path),
Details: fmt.Sprintf("Class: %s", kind),
FilePath: path,
Timestamp: now,
}
var scoreContent []byte
if kind == "cron" {
scoreContent = content
}
return rescoreSensitive(f, kind, scoreContent, 0, now), true
}
// SensitiveFileState is the stable identity of one watchset path. Regular-file
// inode churn is deliberately excluded, while security metadata and symlink
// targets remain visible.
type SensitiveFileState struct {
ContentDigest string
PathIdentity string
}
// NextSensitiveDigests builds the state snapshot for a refresh cycle and
// returns the exact readable regular-file content behind each digest. A path
// whose content cannot be read keeps its previous digest so a transient read
// error does not surface as a content change. Non-regular objects retain type
// identity without being read. Paths absent from paths drop out. prev is only
// read and is never mutated.
func NextSensitiveDigests(prev map[string]SensitiveFileState, paths []string) (map[string]SensitiveFileState, map[string][]byte) {
next := make(map[string]SensitiveFileState, len(paths))
contents := make(map[string][]byte, len(paths))
for _, path := range paths {
state := prev[path]
identity, regularContent, identityKnown := sensitivePathIdentity(path)
if identityKnown {
state.PathIdentity = identity
}
if !regularContent {
next[path] = state
continue
}
data, err := readSensitiveRegularFile(path)
if err == nil {
sum := sha256.Sum256(data)
state.ContentDigest = hex.EncodeToString(sum[:])
contents[path] = data
}
next[path] = state
}
return next, contents
}
type sensitiveRegularFileReader interface {
ReadRegularFile(string) ([]byte, error)
}
// Production uses a nonblocking, fd-verified read. The fallback keeps custom
// test providers source-compatible after the preceding type check.
func readSensitiveRegularFile(path string) ([]byte, error) {
if reader, ok := osFS.(sensitiveRegularFileReader); ok {
return reader.ReadRegularFile(path)
}
return osFS.ReadFile(path)
}
func sensitivePathIdentity(path string) (identity string, regularContent, ok bool) {
info, err := osFS.Lstat(path)
if err != nil {
return "", false, false
}
if info.Mode()&os.ModeSymlink != 0 {
target, err := osFS.Readlink(path)
if err != nil {
return "", false, false
}
targetInfo, err := osFS.Stat(path)
if err != nil {
return "symlink\x00" + target + "\x00unresolved", false, true
}
resolvedIdentity, regular := sensitiveResolvedIdentity(targetInfo)
return "symlink\x00" + target + "\x00" + resolvedIdentity, regular, true
}
identity, regular := sensitiveResolvedIdentity(info)
return identity, regular, true
}
func sensitiveResolvedIdentity(info os.FileInfo) (string, bool) {
modeType := info.Mode() & os.ModeType
identity := fmt.Sprintf("type\x00%x\x00perm\x00%o", modeType, info.Mode().Perm())
if uid, gid, ok := sensitiveFileOwnership(info); ok {
identity += fmt.Sprintf("\x00owner\x00%d\x00%d", uid, gid)
}
if modeType == 0 {
return identity, true
}
if fileIdentity, ok := selfWriteIdentityFromFileInfo(info); ok {
identity += fmt.Sprintf("\x00%d\x00%d", fileIdentity.Device, fileIdentity.Inode)
}
return identity, false
}
// DiffSensitiveWatchset compares two refresh snapshots of the watchset and
// returns the findings the newer one warrants. Both maps are keyed by
// absolute path. contents holds the bytes used for cur's content digests.
//
// Identity is the path, never the inode. Keying on dev+inode reported every
// atomic rewrite (write temp, rename over) as a brand-new file, which is how
// /etc/passwd came to "appear" 16 times in a month on a live host. A path the
// previous snapshot knew about can only have changed, not appeared.
//
// Paths that vanished produce nothing: the caller unwatches the inode, and a
// deletion is not evidence of the tampering this watchset exists to catch.
//
// liveReported holds the exact states whose live-hook findings were delivered
// since the last refresh. A matching state is adopted without a duplicate.
// An appearance is still reported: the hook describes a write, not the fact
// that the path did not exist before.
func DiffSensitiveWatchset(prev, cur map[string]SensitiveFileState, contents map[string][]byte, liveReported map[string]SensitiveFileState) []alert.Finding {
paths := make([]string, 0, len(cur))
for path := range cur {
paths = append(paths, path)
}
sort.Strings(paths)
var findings []alert.Finding
for _, path := range paths {
curState := cur[path]
prevState, known := prev[path]
content, contentKnown := contents[path]
switch {
case !known:
if f, emit := evaluateSensitiveWatchsetChange(path, "New sensitive system file appeared", content, contentKnown); emit {
findings = append(findings, f)
}
case !sensitiveFileStateChanged(prevState, curState):
// Unknown evidence and equivalent regular-file rewrites do not
// establish a change.
case sensitiveLiveReportMatches(liveReported[path], curState):
// The live hook already delivered a finding for this exact state.
default:
summary := "Content changed on sensitive system file"
if prevState.PathIdentity != "" && curState.PathIdentity != "" && prevState.PathIdentity != curState.PathIdentity {
summary = "Path identity changed on sensitive system file"
}
if f, emit := evaluateSensitiveWatchsetChange(path, summary, content, contentKnown); emit {
findings = append(findings, f)
}
}
}
return findings
}
func sensitiveFileStateChanged(prev, cur SensitiveFileState) bool {
contentChanged := prev.ContentDigest != "" && cur.ContentDigest != "" && prev.ContentDigest != cur.ContentDigest
pathChanged := prev.PathIdentity != "" && cur.PathIdentity != "" && prev.PathIdentity != cur.PathIdentity
return contentChanged || pathChanged
}
func sensitiveLiveReportMatches(reported, cur SensitiveFileState) bool {
return reported.ContentDigest != "" && reported.PathIdentity != "" && reported == cur
}
// CheckSensitiveFiles is the periodic safety-net that runs when the BPF
// live monitor is unavailable or disabled. It content-hashes every watchset
// path and emits a finding when a hash differs from the previous run. The
// first run records baselines without emitting findings.
//
// CheckShadowChanges in auth.go does richer per-user diff and infra-IP
// suppression for /etc/shadow specifically; this catch-all complements
// that for sshd_config, sudoers, cron drop-ins, etc. Both run in parallel;
// audit-log dedup handles the (rare) overlap.
func CheckSensitiveFiles(_ context.Context, _ *config.Config, store *state.Store) []alert.Finding {
if store == nil {
return nil
}
var findings []alert.Finding
_, baselineComplete := store.GetRaw(sensitiveFileBaselineKey)
for _, path := range ExpandWatchset("/") {
data, err := osFS.ReadFile(path)
if err != nil {
continue
}
sum := sha256.Sum256(data)
hashHex := hex.EncodeToString(sum[:])
key := "_sensitive_file_hash:" + path
prev, ok := store.GetRaw(key)
if !ok {
store.SetRaw(key, hashHex)
if baselineComplete {
if f, emit := EvaluateSensitiveFileAppearance(path); emit {
findings = append(findings, f)
}
}
continue
}
if prev == hashHex {
continue
}
store.SetRaw(key, hashHex)
// A content change CSM itself made (e.g. installing a wp-cron) updates
// the stored baseline above but raises no finding.
if isExpectedSelfWrite(path, data) {
continue
}
if suppressedAsManagedWPCron(path, data) {
continue
}
kind := classifySensitive(path)
var contentForScore []byte
if kind == "cron" {
contentForScore = data
}
hashChange := alert.Finding{
Severity: alert.High,
Check: "sensitive_file_modified",
Message: fmt.Sprintf("Periodic check: content hash changed for %s", path),
Details: fmt.Sprintf("Previous: %s, Current: %s", prev, hashHex),
FilePath: path,
Timestamp: time.Now(),
}
findings = append(findings, rescoreSensitive(hashChange, kind, contentForScore, 0, time.Now()))
}
if !baselineComplete {
store.SetRaw(sensitiveFileBaselineKey, "1")
}
return findings
}
package checks
import (
"bytes"
"fmt"
"os"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// pkgManagerLogs is the ordered set of package-manager log files whose
// recent mtime acts as evidence of a legitimate root-driven file system
// change. RPM-family hosts touch dnf.rpm.log / yum.log; Debian-family
// hosts touch dpkg.log; minimal installs add history.log for unattended-
// upgrades. The variable is package-private (no operator override) so an
// attacker who learns CSM is here cannot point the daemon at an empty
// path -- the trade-off is that mtime spoofing requires root, which is
// already game-over for this detector class.
var pkgManagerLogs = []string{
"/var/log/dnf.rpm.log",
"/var/log/dnf.log",
"/var/log/yum.log",
"/var/log/dpkg.log",
"/var/log/apt/history.log",
}
// AncestryProvenance names the trusted system component that owns the process
// tree rooted at pid -- a package transaction, or the control panel's own
// maintenance -- and returns "" when nothing in the chain proves one. The
// string is recorded in the finding, so an operator can tell which evidence
// produced a demotion. Nil on hosts where the daemon has not wired it, in
// which case the pkg-window and cron-content layers still apply.
var AncestryProvenance func(pid uint32) string
// pkgManagerWindow returns true when any pkgManagerLogs file was modified
// within window. Reads file mtime only; does not parse log contents.
func pkgManagerWindow(now time.Time, window time.Duration) bool {
cutoff := now.Add(-window)
for _, p := range pkgManagerLogs {
fi, err := os.Stat(p)
if err != nil {
continue
}
if fi.ModTime().After(cutoff) {
return true
}
}
return false
}
// cronDangerTokens are byte fragments whose presence in a cron drop-in is
// inconsistent with vendor-shipped automation and consistent with
// post-exploitation persistence. The list is intentionally narrow: false
// positives here cancel the demote, which is the safe failure mode.
var cronDangerTokens = [][]byte{
[]byte("| sh"),
[]byte("|sh "),
[]byte("| bash"),
[]byte("|bash "),
[]byte("; sh "),
[]byte(";sh "),
[]byte("; bash"),
[]byte(";bash "),
[]byte("base64 -d"),
[]byte("base64 --decode"),
[]byte("base64_decode"),
[]byte("eval("),
[]byte("eval $("),
[]byte("eval \""),
[]byte("/tmp/"),
[]byte("/var/tmp/"),
[]byte("/dev/shm/"),
[]byte("python -c"),
[]byte("python3 -c"),
[]byte("perl -e"),
[]byte("ruby -e"),
[]byte("nc -e"),
[]byte("ncat -e"),
[]byte("bash -i"),
[]byte("\\x"),
[]byte("curl "),
[]byte("wget "),
}
// cronHasDangerTokens also uses the crontab detector so known persistence,
// including encoded payloads, cannot be downgraded by a smaller token list.
// The extra shell fragments are case-sensitive. The curl/wget tokens are
// broad on purpose: a vendor cron that needs to fetch is rare enough that
// flagging is the right default.
func cronHasDangerTokens(content []byte) bool {
for _, tok := range cronDangerTokens {
if bytes.Contains(content, tok) {
return true
}
}
return len(MatchCrontabPatternsDeep(string(content), nil)) > 0
}
// rescoreSensitive returns f with severity adjusted per provenance signals:
// - package-manager activity inside pkgWindowDefault demotes High to Warning
// - AncestryProvenance(pid) naming a trusted component demotes High to Warning
// - cron class with cronHasDangerTokens(content) vetoes any demote
//
// content and pid are optional (nil / 0). class is "" for non-classified
// findings. now is injected for deterministic testing.
func rescoreSensitive(f alert.Finding, class string, content []byte, pid uint32, now time.Time) alert.Finding {
if f.Severity != alert.High {
return f
}
veto := class == "cron" && len(content) > 0 && cronHasDangerTokens(content)
if veto {
return f
}
var reason string
if pkgManagerWindow(now, pkgWindowDefault) {
reason = "package manager active within window"
} else if pid != 0 && AncestryProvenance != nil {
reason = AncestryProvenance(pid)
}
if reason == "" {
return f
}
f.Severity = alert.Warning
if f.Details == "" {
f.Details = fmt.Sprintf("Demoted: %s", reason)
} else {
f.Details = fmt.Sprintf("%s [demoted: %s]", f.Details, reason)
}
return f
}
// PkgManagerRecentlyActive reports whether any package-manager log was
// modified within the provenance demotion window. Exported for the fanotify
// /tmp-executable demotion, which gates on the same evidence as
// rescoreSensitive.
func PkgManagerRecentlyActive(now time.Time) bool {
return pkgManagerWindow(now, pkgWindowDefault)
}
// pkgWindowDefault is the slack we give for a legitimate root-driven file
// system change after a package transaction. dnf scriptlets observed up to
// a few seconds between the rpm log entry and post-install file drops; 2
// minutes covers slower scriptlets without inviting a multi-minute window
// for an attacker who happened to time a transaction.
const pkgWindowDefault = 2 * time.Minute
package checks
import (
"path/filepath"
"regexp"
"strconv"
"strings"
)
// csmManagedWPCronJobRe matches a single crontab job line exactly as
// wpCronJobLine emits it. The command shape is fixed -- run wp-cron.php under a
// per-docroot flock -- and carries no general-purpose payload, so a job line
// matching this pattern cannot be repurposed for attacker persistence. Any
// foreign command (reverse shell, curl|bash, miner) fails the match, which is
// what keeps crontabIsExclusivelyCSMWPCron from suppressing a tampered crontab.
var csmManagedWPCronJobRe = regexp.MustCompile(
`^([0-9*/,-]+) \* \* \* \* cd ('(?:[^']|'\\'')*') && flock -n "\$HOME/\.csm-wpcron-([0-9a-f]{8})\.lock" ('(?:[^']|'\\'')*') -d max_execution_time=300 wp-cron\.php >/dev/null 2>&1$`)
var (
cpanelPHPBinRe = regexp.MustCompile(`^/opt/cpanel/ea-php[0-9]{2}/root/usr/bin/php$`)
cloudLinuxPHPBinRe = regexp.MustCompile(`^/opt/alt/php[0-9]{2}/usr/bin/php$`)
)
// safeCrontabShells is the set of SHELL values a fully CSM-managed crontab may
// carry. SHELL is honored by crond to exec every job, so an arbitrary value is
// a code-execution vector; only known shells (cPanel's jailshell and the
// standard system shells) are accepted. cPanel prepends the jailshell line to
// every user crontab it touches.
var safeCrontabShells = map[string]bool{
"/usr/local/cpanel/bin/jailshell": true,
"/bin/bash": true,
"/bin/sh": true,
"/usr/bin/bash": true,
"/usr/bin/sh": true,
}
// crontabIsExclusivelyCSMWPCron reports whether every meaningful line in a user
// crontab is either an inert header (blank, comment, or a vetted environment
// assignment) or a CSM-installed WP-Cron job line. Such a crontab is fully
// CSM-managed; flagging it as a sensitive-file change is a false positive that
// the in-memory, TTL-bounded self-write ledger misses after a daemon restart or
// a crontab reformat by crond/cPanel.
//
// Safety: this is a content-structure recognizer, not a path allowlist. A
// single foreign cron entry, an unrecognized environment assignment (PATH,
// BASH_ENV, ...), or an unsafe SHELL value makes it return false, so attacker
// persistence layered into a crontab is still surfaced. At least one CSM job
// line is required so an all-headers crontab is not mistaken for ours.
func crontabIsExclusivelyCSMWPCron(owner string, content []byte) bool {
sawCSMJob := false
pendingMarker := ""
for _, raw := range strings.Split(string(content), "\n") {
line := strings.TrimSpace(raw)
if line == "" {
pendingMarker = ""
continue
}
if strings.HasPrefix(line, "#") {
// Comments are never executed by crond. The CSM marker is retained
// only to pin the following managed job to the docroot CSM wrote.
pendingMarker = ""
if strings.HasPrefix(line, wpCronJobMarker) {
pendingMarker = strings.TrimSpace(strings.TrimPrefix(line, wpCronJobMarker))
}
continue
}
if name, val, ok := splitCrontabEnv(line); ok {
if !safeCrontabEnvAssignment(name, val) {
return false
}
pendingMarker = ""
continue
}
if docroot, ok := csmManagedWPCronJob(line, owner); ok && pendingMarker == docroot {
sawCSMJob = true
pendingMarker = ""
continue
}
return false
}
return sawCSMJob
}
func csmManagedWPCronJob(line, owner string) (string, bool) {
m := csmManagedWPCronJobRe.FindStringSubmatch(line)
if m == nil {
return "", false
}
minute, lockHex := m[1], m[3]
docroot, ok := unquoteShellSingle(m[2])
if !ok || !safeManagedWPCronDocroot(owner, docroot) {
return "", false
}
phpBin, ok := unquoteShellSingle(m[4])
if !ok || !safeManagedWPCronPHPBin(phpBin) {
return "", false
}
lockID, err := strconv.ParseUint(lockHex, 16, 32)
if err != nil || uint32(lockID) != wpCronLockID(docroot) {
return "", false
}
if !wpCronMinuteMatchesOwnerDocroot(minute, owner, docroot) {
return "", false
}
return docroot, true
}
func unquoteShellSingle(q string) (string, bool) {
if len(q) < 2 || q[0] != '\'' || q[len(q)-1] != '\'' {
return "", false
}
body := q[1 : len(q)-1]
var out strings.Builder
for len(body) > 0 {
i := strings.IndexByte(body, '\'')
if i < 0 {
out.WriteString(body)
return out.String(), true
}
out.WriteString(body[:i])
if !strings.HasPrefix(body[i:], `'\''`) {
return "", false
}
out.WriteByte('\'')
body = body[i+4:]
}
return out.String(), true
}
func safeManagedWPCronDocroot(owner, docroot string) bool {
if owner == "" || !safeWPCronDocroot(docroot) {
return false
}
_, account, ok := accountRootOf(docroot)
return ok && account == owner
}
func safeManagedWPCronPHPBin(path string) bool {
if !safeCronCommandString(path) || !filepath.IsAbs(path) || filepath.Clean(path) != path {
return false
}
switch path {
case "/usr/local/bin/php", "/usr/bin/php", "/bin/php":
return true
default:
return cpanelPHPBinRe.MatchString(path) || cloudLinuxPHPBinRe.MatchString(path)
}
}
func wpCronMinuteMatchesOwnerDocroot(minute, owner, docroot string) bool {
for interval := 1; interval <= wpCronMaxIntervalMin; interval++ {
if wpCronMinuteField(wpCronStaggerOffset(owner, docroot, interval), interval) == minute {
return true
}
}
return false
}
// splitCrontabEnv parses a crontab environment line of the form NAME=value.
// crond treats a line as an assignment (not a command) when the text left of
// the first '=' is a single bare identifier; a job line such as
// "* * * * * FOO=bar cmd" has spaces before '=' and is not an assignment. The
// value's surrounding quotes are stripped to match how cPanel writes them.
func splitCrontabEnv(line string) (name, value string, ok bool) {
eq := strings.IndexByte(line, '=')
if eq <= 0 {
return "", "", false
}
name = strings.TrimSpace(line[:eq])
if strings.ContainsAny(name, " \t") {
return "", "", false
}
for i, r := range name {
switch {
case r == '_':
case r >= 'A' && r <= 'Z':
case r >= 'a' && r <= 'z':
case r >= '0' && r <= '9' && i > 0:
default:
return "", "", false
}
}
value = strings.TrimSpace(line[eq+1:])
if value == "" {
return name, value, true
}
switch value[0] {
case '\'', '"':
if len(value) < 2 || value[len(value)-1] != value[0] {
return "", "", false
}
value = value[1 : len(value)-1]
}
return name, value, true
}
// safeCrontabEnvAssignment reports whether a crontab environment assignment is
// inert enough to appear in a crontab still considered fully CSM-managed.
// MAILTO is stored verbatim by crond and never executed, so any value is safe.
// SHELL and HOME influence job execution, so their values are constrained.
// Every other name (PATH, BASH_ENV, LD_*, ...) is rejected.
func safeCrontabEnvAssignment(name, value string) bool {
switch strings.ToUpper(name) {
case "MAILTO":
return true
case "SHELL":
return safeCrontabShells[value]
case "HOME":
return safeCrontabHome(value)
default:
return false
}
}
func safeCrontabHome(value string) bool {
if !safeCronCommandString(value) || !filepath.IsAbs(value) || filepath.Clean(value) != value {
return false
}
return value == "/root" || underAccountRoot(value) || strings.HasPrefix(value, "/root/")
}
// suppressedAsManagedWPCron reports whether a sensitive-file finding for a user
// crontab should be suppressed because the crontab is exclusively a
// CSM-installed WP-Cron block. Scoped to the platform's user crontab spool,
// the only place CSM installs WP-Cron jobs; system drop-ins under /etc/cron.d
// are never suppressed here.
func suppressedAsManagedWPCron(path string, content []byte) bool {
if len(content) == 0 {
return false
}
owner, ok := cronSpoolOwner(path)
if !ok {
return false
}
return crontabIsExclusivelyCSMWPCron(owner, content)
}
func cronSpoolOwner(path string) (string, bool) {
clean := filepath.ToSlash(filepath.Clean(path))
if filepath.ToSlash(filepath.Dir(clean)) != filepath.ToSlash(filepath.Clean(cronSpoolDir())) {
return "", false
}
owner := filepath.Base(clean)
if owner == "." || owner == "/" || owner == "" {
return "", false
}
if owner == "root" || !validCPUser.MatchString(owner) {
return "", false
}
return owner, true
}
package checks
import (
"os"
"path/filepath"
"strings"
)
// Recognisers used by the realtime fanotify restore/probe dedup logic.
// The scheduled file index deliberately does not call these as scan skips:
// upload PHP is indexed first and then classified by content.
// LooksLikeCpanelRestoreStaging recognises files inside cPanel's
// pkgacct/restorepkg staging tree. cPanel extracts the user backup
// as root into /home/cpanelpkgrestore.TMP.work.<id>/ for inspection,
// then re-extracts it under the user identity into /home/<account>/.
// Both extractions raise events; the user-context one carries the
// real signal, so the staging-side alert is a duplicate.
//
// The recogniser requires the marker to sit directly under /home (the
// only place cPanel ever creates it) plus a non-empty alphanumeric id
// of >=2 chars. A user account at /home/<user>/ cannot create
// siblings of itself, so this gate cannot be spoofed by a non-root
// attacker.
func LooksLikeCpanelRestoreStaging(path string) bool {
const homeRoot = "/home"
const marker = "/cpanelpkgrestore.TMP.work."
idx := strings.Index(path, marker)
if idx < 0 {
return false
}
if idx != len(homeRoot) {
return false
}
if !strings.HasPrefix(path, homeRoot) {
return false
}
rest := path[idx+len(marker):]
if rest == "" {
return false
}
end := strings.IndexByte(rest, '/')
var token string
if end < 0 {
token = rest
} else {
token = rest[:end]
}
if len(token) < 2 {
return false
}
for i := 0; i < len(token); i++ {
c := token[i]
switch {
case c >= '0' && c <= '9':
case c >= 'a' && c <= 'z':
case c >= 'A' && c <= 'Z':
default:
return false
}
}
return true
}
// LooksLikeWPOptimizeProbeByPath recognises WP-Optimize's per-server
// probe files using path structure alone. WP-Optimize writes tiny
// <?php files to /wp-content/uploads/wpo/.../test.php to test whether
// the host honours certain Apache/Nginx directives.
//
// Path-only gates (no content read):
//
// 1. Path lies under /wp-content/uploads/wpo/.
// 2. The basename is exactly "test.php" (the literal filename
// WP-Optimize uses for these probes; an attacker dropping
// /uploads/wpo/webshell.php fails this gate and continues to
// the standard alert).
// 3. The wp-optimize plugin directory is actually present in this
// site's wp-content/plugins/ tree (filesystem stat).
//
// The realtime path additionally applies a content shape gate before
// suppressing the duplicate alert, so the path predicate here stays
// narrow to the literal probe filename and installed plugin directory.
func LooksLikeWPOptimizeProbeByPath(path string) bool {
const marker = "/wp-content/uploads/wpo/"
if !strings.Contains(path, marker) {
return false
}
if filepath.Base(path) != "test.php" {
return false
}
uploadsIdx := strings.Index(path, "/wp-content/uploads/")
if uploadsIdx < 0 {
return false
}
pluginDir := path[:uploadsIdx] + "/wp-content/plugins/wp-optimize"
st, err := os.Stat(pluginDir)
if err != nil || !st.IsDir() {
return false
}
return true
}
package checks
import (
"fmt"
"github.com/pidginhost/csm/internal/alert"
)
// AttributeSocketOwner attaches the hosting account identified by a kernel
// socket or connection-event UID. A missing process snapshot does not remove
// this evidence. Root, service and unresolved UIDs remain unattributed.
// The message carries the account too, so dispatch and audit identities keep
// different accounts contacting the same destination separate.
func AttributeSocketOwner(f *alert.Finding, uid uint32) {
if uid == 0 {
return
}
if owner := HostingAccountForUser(LookupUser(uid)); owner != "" {
f.TenantID = owner
f.Message += fmt.Sprintf(" (account %s)", owner)
}
}
package checks
import (
"regexp"
"strings"
)
// SEO-spam context analysis for WordPress post content.
//
// Word-boundary keyword matching (see countSpamMatches) eliminated the
// substring false positives where "specialist" triggered "cialis" and
// "pharmaceutical" triggered "pharma". It did not eliminate a second
// class of false positive: legitimate prose mentions of an industry or
// product category ("our advisor covers consumer goods, energy,
// pharma" or "Industria alimentara si Pharma"). Word-boundary matching
// cannot distinguish the prose mention from the cloaked black-hat SEO
// link the lalimanro attack injected on the same site.
//
// This file classifies a keyword HIT as SPAM only when the surrounding
// HTML shows an attacker signal: CSS cloaking (off-screen absolute
// positioning, display:none, visibility:hidden, text-indent, micro
// height, font-size:0), an injection fingerprint (short hex HTML
// comment bracketing content), or an external anchor whose URL path
// itself contains the keyword ("/buy/pharma/" style commercial paths).
//
// Bare keyword mentions with none of those signals do not fire, so
// legitimate industry-vertical prose is silent.
//
// The proximity window is bounded to ±spamContextWindow bytes around
// the keyword match. Cloaking at the top of a long post unrelated to
// a keyword mention at the bottom does not spuriously associate.
// spamContextWindow bounds how far (in bytes) around a keyword match
// the analyzer looks for cloaking signals. 400 covers a typical
// cloaked-div attack (small <div> with style + <a> + keyword) without
// reaching into unrelated content.
const spamContextWindow = 400
// cssCloakPatterns are regexes that each, when matched, identify a
// CSS property value indicative of content cloaking. The list is
// conservative: each entry corresponds to a technique widely used in
// real SEO spam campaigns and rarely to nothing else at production
// scale. `(?i)` makes matching case-insensitive; `\s*` around colons
// and values tolerates the whitespace variants attackers use to evade
// naive string scanners ("display : none" vs "display:none").
var cssCloakPatterns = []*regexp.Regexp{
// display:none and visibility:hidden — classic hide.
regexp.MustCompile(`(?i)\bdisplay\s*:\s*none\b`),
regexp.MustCompile(`(?i)\bvisibility\s*:\s*hidden\b`),
// text-indent with a negative value — pushes text off-screen.
regexp.MustCompile(`(?i)\btext-indent\s*:\s*-\s*\d+`),
// Micro height with a positive integer of 0 or 1, paired with
// overflow:hidden to suppress contents. We require the pair
// because height:1 alone is occasionally legitimate.
regexp.MustCompile(`(?i)\bheight\s*:\s*[01](\s*px)?\b[^"'}]*\boverflow\s*:\s*hidden\b`),
regexp.MustCompile(`(?i)\boverflow\s*:\s*hidden\b[^"'}]*\bheight\s*:\s*[01](\s*px)?\b`),
// font-size:0 — classic invisible-text technique.
regexp.MustCompile(`(?i)\bfont-size\s*:\s*0\b`),
// position:absolute paired with a negative coordinate in the same
// style attribute (non-greedy [^"'}]* stays inside one attribute).
// Legitimate uses of position:absolute (menus, tooltips) have
// non-negative coordinates; the pairing is the signal.
regexp.MustCompile(`(?i)\bposition\s*:\s*absolute\b[^"'}]*\b(left|top|right|bottom)\s*:\s*-\s*\d+`),
regexp.MustCompile(`(?i)\b(left|top|right|bottom)\s*:\s*-\s*\d+[^"'}]*\bposition\s*:\s*absolute\b`),
// Off-screen via large negative margin on a block element.
regexp.MustCompile(`(?i)\bmargin(-left|-top)?\s*:\s*-\s*\d{4,}`),
}
// injectionFingerprintRe matches the short hex HTML comment (4-8 hex
// chars) attackers use to tag their own injections across a campaign.
// The length bracket is deliberate: shorter matches are too ambiguous,
// longer matches overlap with legitimate WordPress UUIDs.
var injectionFingerprintRe = regexp.MustCompile(`<!--\s*[a-f0-9]{4,8}\s*-->`)
// anchorHrefRe extracts the href attribute value from each anchor tag
// in a fragment. Double- and single-quoted values are both accepted.
var anchorHrefRe = regexp.MustCompile(`(?i)<a\b[^>]+href\s*=\s*["']([^"']+)["']`)
// externalSchemeRe recognises absolute or protocol-relative URLs. A
// relative URL ("/services/pharma/") is same-origin navigation and not
// an external spam link.
var externalSchemeRe = regexp.MustCompile(`^(?i)(https?:)?//`)
// contentHasSpamContext reports whether any occurrence of the keyword
// in content is accompanied by an SEO-spam signal within
// spamContextWindow bytes. Returning false is a direct statement that
// every keyword hit is a bare mention without cloaking context — the
// caller should suppress the finding in that case.
func contentHasSpamContext(content string, pattern dbSpamPattern) bool {
matches := spamKeywordMatchIndexes(pattern, content)
for _, m := range matches {
if hitHasSpamContext(content, m[0], m[1], pattern.keyword) {
return true
}
}
return false
}
// countCloakedSpamMatches returns the number of rows in contents whose
// text contains the spam keyword AND shows an accompanying SEO-spam
// context signal. It is the aggregator used by checkWPPosts to decide
// whether to emit a db_spam_injection finding: bare prose mentions of
// a keyword do not count; only cloaked/SEO-style injections do.
//
// Each qualifying row is counted exactly once regardless of how many
// keyword hits it contains — the finding is per-post, not per-hit.
func countCloakedSpamMatches(pattern dbSpamPattern, contents []string) int {
n := 0
for _, c := range contents {
if contentHasSpamContext(c, pattern) {
n++
}
}
return n
}
// hitHasSpamContext examines the ±spamContextWindow byte region around
// a single keyword hit for cloaking, injection-fingerprint, or
// external-spam-link signals. Extracted as a helper so tests can target
// one hit at a time.
func hitHasSpamContext(content string, start, end int, keyword string) bool {
ws := start - spamContextWindow
if ws < 0 {
ws = 0
}
we := end + spamContextWindow
if we > len(content) {
we = len(content)
}
window := content[ws:we]
if windowHasCSSCloaking(window) {
return true
}
if injectionFingerprintRe.MatchString(window) {
return true
}
if windowHasExternalSpamAnchor(window, keyword) {
return true
}
return false
}
// positionAbsoluteRe and negativeCoordRe are evaluated together in
// windowHasCSSCloaking for the "off-screen absolute positioning"
// signal. Keeping them as two independent regexes (rather than one
// paired regex with [^"'}]* between them) closes an evasion where the
// attacker splits the cloak across two CSS rules — e.g.
// `<style>.a{position:absolute}.b{left:-9999px}</style>` — which the
// paired form stops matching at the first `}`. Both signals must
// appear somewhere in the proximity window; the window itself is
// bounded (see spamContextWindow) so the association is not unbounded.
var (
positionAbsoluteRe = regexp.MustCompile(`(?i)\bposition\s*:\s*absolute\b`)
negativeCoordRe = regexp.MustCompile(`(?i)\b(left|top|right|bottom|margin-left|margin-top)\s*:\s*-\s*\d{2,}`)
)
// windowHasCSSCloaking returns true if the window contains any CSS
// declaration from cssCloakPatterns, OR if it contains both
// position:absolute and a negative coordinate somewhere in the window
// (independent-signal form, for rule-split evasions).
func windowHasCSSCloaking(window string) bool {
for _, re := range cssCloakPatterns {
if re.MatchString(window) {
return true
}
}
if positionAbsoluteRe.MatchString(window) && negativeCoordRe.MatchString(window) {
return true
}
return false
}
// windowHasExternalSpamAnchor returns true if the window contains an
// <a href> whose destination is an external URL (absolute or
// protocol-relative) AND whose URL path contains the spam keyword as
// a bounded segment. Same-origin relative links ("/services/pharma/")
// are skipped — they are internal navigation, not spam.
func windowHasExternalSpamAnchor(window, keyword string) bool {
anchors := anchorHrefRe.FindAllStringSubmatch(window, -1)
if len(anchors) == 0 {
return false
}
keywordLower := strings.ToLower(keyword)
for _, a := range anchors {
href := a[1]
if !externalSchemeRe.MatchString(href) {
continue
}
if !hrefPathContainsKeyword(href, keywordLower) {
continue
}
return true
}
return false
}
// hrefPathContainsKeyword checks whether the URL path portion of href
// contains the keyword bounded by non-alphanumeric characters. The
// boundary guarantees "/buy/pharma/cheap" matches "pharma" but
// "/pharmaceutical/" does not.
func hrefPathContainsKeyword(href, keywordLower string) bool {
// Strip scheme://host prefix. externalSchemeRe already confirmed
// the URL starts with //host or scheme://host.
rest := href
if i := strings.Index(rest, "://"); i >= 0 {
rest = rest[i+3:]
} else {
rest = strings.TrimPrefix(rest, "//")
}
// Everything after the first '/' is the path+query.
var path string
if i := strings.IndexByte(rest, '/'); i >= 0 {
path = rest[i:]
} else {
path = ""
}
lower := strings.ToLower(path)
idx := strings.Index(lower, keywordLower)
if idx < 0 {
return false
}
// Check bounded: preceding char (if any) and following char (if
// any) must not be alphanumeric. This disambiguates
// "/buy/pharma/" (bounded by '/') from "/pharmacist/" (bounded by
// 'c' after "pharma").
if idx > 0 {
c := lower[idx-1]
if isURLWordChar(c) {
return false
}
}
after := idx + len(keywordLower)
if after < len(lower) {
c := lower[after]
if isURLWordChar(c) {
return false
}
}
return true
}
// isURLWordChar reports whether a byte is an ASCII alphanumeric used
// for URL-path word-boundary analysis. Underscore is treated as a word
// character by convention; hyphen is NOT (so "buy-cheap-pharma" still
// matches "pharma" bounded by the trailing hyphen/slash).
func isURLWordChar(b byte) bool {
switch {
case b >= 'a' && b <= 'z':
return true
case b >= 'A' && b <= 'Z':
return true
case b >= '0' && b <= '9':
return true
case b == '_':
return true
}
return false
}
package checks
import (
"encoding/json"
"fmt"
"os"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
const sshLoginFollowKey = "_ssh_login_follow"
// sshLoginLookback bounds how far back a first run or a large read gap
// reports accepted logins. Anything older is history, not a new login, and
// ssh_login_unknown_ip is a critical, always-block finding.
const sshLoginLookback = time.Hour
var sshLoginMu sync.Mutex
type sshLoginFollow struct {
Follow followState `json:"follow"`
}
func loadSSHLoginFollow(store *state.Store) followState {
raw, ok := store.GetRaw(sshLoginFollowKey)
if !ok || raw == "" {
return followState{}
}
var decoded sshLoginFollow
if err := json.Unmarshal([]byte(raw), &decoded); err != nil || invalidFollowState(decoded.Follow) {
return followState{}
}
return decoded.Follow
}
func saveSSHLoginFollow(store *state.Store, st followState) {
b, err := json.Marshal(sshLoginFollow{Follow: st})
if err != nil {
return
}
if err := store.SetRawAndSave(sshLoginFollowKey, string(b)); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving SSH login follow state: %v\n", err)
}
}
// checkSSHLoginsFollow reads the auth log forward from the stored offset, so
// a login is seen however much brute-force noise follows it before the next
// cycle, and skips logins older than sshLoginLookback (first-run catch-up).
func checkSSHLoginsFollow(cfg *config.Config, store *state.Store) []alert.Finding {
sshLoginMu.Lock()
defer sshLoginMu.Unlock()
st := loadSSHLoginFollow(store)
lines, next, _, err := readNewSyslogLines(authLogPath(), st)
if err != nil {
return nil // leave stored state untouched
}
now := time.Now()
cutoff := now.Add(-sshLoginLookback)
var findings []alert.Finding
for _, line := range lines {
if !strings.Contains(line, "Accepted") {
continue
}
if at, ok := syslogLineTime(line, now); ok && at.Before(cutoff) {
continue
}
if f, ok := SSHAcceptedLoginFinding(line, cfg); ok {
findings = append(findings, f)
}
}
saveSSHLoginFollow(store, next)
return findings
}
package checks
import (
"context"
"encoding/json"
"fmt"
"path/filepath"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// Supply-chain dependency scanning.
//
// This is the scanner half of the supply-chain check: it parses
// composer.lock / package-lock.json dependency trees under customer
// document roots and matches the resolved versions against a local
// advisory database. The advisory database itself is operational data,
// not shipped in the binary -- an operator or a sync job writes
// <state>/advisories/supply-chain.json (format documented in
// docs/supply-chain-advisories.md). With no advisory file present the
// check is dormant: it parses nothing it cannot match and emits nothing.
// This mirrors the YARA-forge mirror posture (machinery in CSM, signed
// data delivered out of band).
// supplyChainAdvisoryRelPath is where CSM looks for the advisory DB,
// relative to the configured state directory.
const supplyChainAdvisoryRelPath = "advisories/supply-chain.json"
// supplyChainPkg is one resolved dependency from a lockfile.
type supplyChainPkg struct {
Ecosystem string // "composer" | "npm"
Name string
Version string
}
// supplyChainAdvisory is the OSV-subset advisory shape CSM matches
// against. A version is vulnerable when it falls inside any range:
// version >= introduced AND (fixed == "" OR version < fixed).
type supplyChainAdvisory struct {
Ecosystem string `json:"ecosystem"`
Package string `json:"package"`
Ranges []supplyChainAdvisoryRange `json:"ranges"`
ID string `json:"id"`
Severity string `json:"severity"`
Summary string `json:"summary"`
}
type supplyChainAdvisoryRange struct {
Introduced string `json:"introduced"`
Fixed string `json:"fixed"`
}
type supplyChainAdvisoryFile struct {
Advisories []supplyChainAdvisory `json:"advisories"`
}
// CheckSupplyChain scans customer dependency lockfiles for versions with
// known advisories. Dormant unless an advisory database is present at
// <state>/advisories/supply-chain.json.
func CheckSupplyChain(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if cfg == nil {
return nil
}
advisories := loadSupplyChainAdvisories(cfg.StatePath)
if len(advisories) == 0 {
return nil // no data -> nothing to match against
}
index := indexAdvisories(advisories)
var findings []alert.Finding
for _, lock := range discoverLockfiles(ctx) {
if ctx != nil && ctx.Err() != nil {
return findings
}
data, err := osFS.ReadFile(lock.path)
if err != nil {
continue
}
pkgs := lock.parse(data)
for _, p := range pkgs {
for _, adv := range index[advisoryKey(p.Ecosystem, p.Name)] {
if versionVulnerable(p.Version, adv.Ranges) {
findings = append(findings, supplyChainFinding(lock.account, lock.path, p, adv))
}
}
}
}
return findings
}
func loadSupplyChainAdvisories(statePath string) []supplyChainAdvisory {
if statePath == "" {
return nil
}
data, err := osFS.ReadFile(filepath.Join(statePath, supplyChainAdvisoryRelPath))
if err != nil {
return nil
}
var f supplyChainAdvisoryFile
if json.Unmarshal(data, &f) != nil {
return nil
}
return f.Advisories
}
func advisoryKey(ecosystem, pkg string) string {
return strings.ToLower(ecosystem) + "\x00" + strings.ToLower(pkg)
}
func indexAdvisories(advisories []supplyChainAdvisory) map[string][]supplyChainAdvisory {
out := map[string][]supplyChainAdvisory{}
for _, a := range advisories {
if a.Package == "" || a.Ecosystem == "" {
continue
}
k := advisoryKey(a.Ecosystem, a.Package)
out[k] = append(out[k], a)
}
return out
}
type lockfile struct {
path string
account string
parse func([]byte) []supplyChainPkg
}
type packageLockDependency struct {
Version string `json:"version"`
Dependencies map[string]packageLockDependency `json:"dependencies"`
}
// discoverLockfiles globs composer.lock and package-lock.json at the
// common project depths under customer home directories. Bounded by the
// glob shape (no recursive walk) so a deep node_modules tree cannot turn
// the scan into an unbounded crawl.
func discoverLockfiles(ctx context.Context) []lockfile {
patterns := []struct {
glob string
parse func([]byte) []supplyChainPkg
}{
// Relative to each account root; see accountHomeGlob.
{"*/public_html/composer.lock", parseComposerLock},
{"*/composer.lock", parseComposerLock},
{"*/public_html/package-lock.json", parsePackageLock},
{"*/package-lock.json", parsePackageLock},
}
var out []lockfile
seen := map[string]struct{}{}
for _, p := range patterns {
if ctx != nil && ctx.Err() != nil {
return out
}
matches, _ := accountHomeGlob(p.glob)
for _, m := range matches {
if _, dup := seen[m]; dup {
continue
}
seen[m] = struct{}{}
out = append(out, lockfile{path: m, account: extractUser(m), parse: p.parse})
}
}
return out
}
func parseComposerLock(data []byte) []supplyChainPkg {
var doc struct {
Packages []struct{ Name, Version string } `json:"packages"`
PackagesDev []struct{ Name, Version string } `json:"packages-dev"`
}
if json.Unmarshal(data, &doc) != nil {
return nil
}
var out []supplyChainPkg
for _, set := range [][]struct{ Name, Version string }{doc.Packages, doc.PackagesDev} {
for _, p := range set {
if p.Name == "" || p.Version == "" {
continue
}
out = append(out, supplyChainPkg{Ecosystem: "composer", Name: p.Name, Version: p.Version})
}
}
return out
}
func parsePackageLock(data []byte) []supplyChainPkg {
var doc struct {
Packages map[string]struct {
Version string `json:"version"`
} `json:"packages"`
Dependencies map[string]packageLockDependency `json:"dependencies"`
}
if json.Unmarshal(data, &doc) != nil {
return nil
}
var out []supplyChainPkg
// npm v2/v3: keyed by "node_modules/<name>" (the root "" entry is the
// project itself and has no node_modules prefix).
paths := make([]string, 0, len(doc.Packages))
for path := range doc.Packages {
paths = append(paths, path)
}
sort.Strings(paths)
for _, path := range paths {
v := doc.Packages[path]
name := npmNameFromPackagesKey(path)
if name == "" || v.Version == "" {
continue
}
out = append(out, supplyChainPkg{Ecosystem: "npm", Name: name, Version: v.Version})
}
// npm v1: dependency tree rooted at the top-level dependencies map.
if len(doc.Packages) == 0 {
return appendPackageLockV1Dependencies(out, doc.Dependencies)
}
return out
}
func appendPackageLockV1Dependencies(out []supplyChainPkg, deps map[string]packageLockDependency) []supplyChainPkg {
stack := []map[string]packageLockDependency{deps}
for len(stack) > 0 {
cur := stack[len(stack)-1]
stack = stack[:len(stack)-1]
names := make([]string, 0, len(cur))
for name := range cur {
names = append(names, name)
}
sort.Strings(names)
for _, name := range names {
dep := cur[name]
if name != "" && dep.Version != "" {
out = append(out, supplyChainPkg{Ecosystem: "npm", Name: name, Version: dep.Version})
}
if len(dep.Dependencies) > 0 {
stack = append(stack, dep.Dependencies)
}
}
}
return out
}
// npmNameFromPackagesKey extracts the package name from a v2/v3
// package-lock "packages" key. The key is the path "node_modules/<name>"
// (or nested "node_modules/a/node_modules/b"); the name is whatever
// follows the last "node_modules/". The root project key "" yields "".
func npmNameFromPackagesKey(key string) string {
if key == "" {
return ""
}
parts := strings.Split(key, "/")
for i := len(parts) - 2; i >= 0; i-- {
if parts[i] != "node_modules" {
continue
}
if parts[i+1] == "" {
return ""
}
if strings.HasPrefix(parts[i+1], "@") {
if i+2 >= len(parts) || parts[i+2] == "" {
return ""
}
return parts[i+1] + "/" + parts[i+2]
}
return parts[i+1]
}
return ""
}
// versionVulnerable reports whether version falls inside any advisory
// range: version >= introduced AND (fixed == "" OR version < fixed).
// An empty/"0" introduced means "from the beginning".
func versionVulnerable(version string, ranges []supplyChainAdvisoryRange) bool {
for _, r := range ranges {
introOK := r.Introduced == "" || r.Introduced == "0" || semverCompare(version, r.Introduced) >= 0
fixedOK := r.Fixed == "" || semverCompare(version, r.Fixed) < 0
if introOK && fixedOK {
return true
}
}
return false
}
// semverCompare compares two dotted versions numerically segment by
// segment, tolerating a leading "v" and ignoring any pre-release/build
// suffix after the first "-" or "+". Returns -1, 0, or 1. Non-numeric
// segments compare as 0 so a malformed version never panics.
func semverCompare(a, b string) int {
as := semverSegments(a)
bs := semverSegments(b)
n := len(as)
if len(bs) > n {
n = len(bs)
}
for i := 0; i < n; i++ {
var av, bv int
if i < len(as) {
av = as[i]
}
if i < len(bs) {
bv = bs[i]
}
if av < bv {
return -1
}
if av > bv {
return 1
}
}
return 0
}
func semverSegments(v string) []int {
v = strings.TrimSpace(v)
v = strings.TrimPrefix(v, "v")
v = strings.TrimPrefix(v, "V")
if i := strings.IndexAny(v, "-+"); i >= 0 {
v = v[:i]
}
parts := strings.Split(v, ".")
out := make([]int, 0, len(parts))
for _, p := range parts {
n, err := strconv.Atoi(strings.TrimSpace(p))
if err != nil {
n = 0
}
out = append(out, n)
}
return out
}
func supplyChainFinding(account, lockPath string, p supplyChainPkg, adv supplyChainAdvisory) alert.Finding {
sev := alert.High
switch strings.ToLower(adv.Severity) {
case "critical":
sev = alert.Critical
case "low", "medium", "moderate", "":
sev = alert.Warning
}
id := adv.ID
if id == "" {
id = "advisory"
}
fixed := "no fixed version published"
for _, r := range adv.Ranges {
if r.Fixed != "" {
fixed = "fixed in " + r.Fixed
break
}
}
return alert.Finding{
Severity: sev,
Check: "supply_chain_vuln",
Message: fmt.Sprintf("Vulnerable %s dependency %s %s (%s) on account %s",
p.Ecosystem, p.Name, p.Version, id, account),
Details: fmt.Sprintf("Account: %s\nLockfile: %s\nEcosystem: %s\nPackage: %s\nVersion: %s\nAdvisory: %s\nSeverity: %s\nFix: %s\n%s",
account, lockPath, p.Ecosystem, p.Name, p.Version, id, adv.Severity, fixed, adv.Summary),
FilePath: lockPath,
Timestamp: time.Now(),
}
}
package checks
import (
"strconv"
"strings"
"time"
)
// syslogLineTime returns the time a syslog line was written, from either the
// traditional BSD prefix ("Sep 2 04:12:33", no year, local time) or an
// RFC 3339 prefix as rsyslog writes with high-precision timestamps. The BSD
// form uses the most recent plausible year; a result more than a day in the
// future belongs to an earlier year (a January read of December lines). ok is
// false when the line carries neither.
func syslogLineTime(line string, now time.Time) (time.Time, bool) {
fields := strings.Fields(line)
if len(fields) == 0 {
return time.Time{}, false
}
if t, err := time.Parse(time.RFC3339Nano, fields[0]); err == nil {
return t, true
}
if len(fields) < 3 || !isSyslogTimestampPrefix(fields) {
return time.Time{}, false
}
stamp := fields[0] + " " + fields[1] + " " + fields[2]
parseYear := func(year int) (time.Time, error) {
return time.ParseInLocation("2006 Jan _2 15:04:05", strconv.Itoa(year)+" "+stamp, now.Location())
}
// Eight years covers the largest gap between Gregorian leap years. This
// also handles a leap-day record read more than one year later without
// treating an unparseable timestamp as a current event.
for yearsAgo := 0; yearsAgo <= 8; yearsAgo++ {
t, err := parseYear(now.Year() - yearsAgo)
if err == nil && !t.After(now.Add(24*time.Hour)) {
return t, true
}
}
return time.Time{}, false
}
package checks
import (
"bufio"
"context"
"fmt"
"io"
"path/filepath"
"slices"
"sort"
"strconv"
"strings"
"syscall"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/mysqlclient"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// CheckKernelModules compares loaded kernel modules against baseline.
// All modules present at baseline time are considered known.
// Only modules loaded AFTER baseline trigger alerts.
func CheckKernelModules(ctx context.Context, _ *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
modules := loadModuleList()
if len(modules) == 0 {
return nil
}
// Check if baseline exists for kernel modules
_, baselineExists := store.GetRaw("_kmod_baseline_set")
if !baselineExists {
// First run - store all current modules as baseline
for _, mod := range modules {
store.SetRaw("_kmod:"+mod, "baseline")
}
store.SetRaw("_kmod_baseline_set", "true")
return nil
}
// Check for modules not seen at baseline
for _, mod := range modules {
key := "_kmod:" + mod
_, known := store.GetRaw(key)
if !known {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "kernel_module",
Message: fmt.Sprintf("New kernel module loaded after baseline: %s", mod),
Details: "This module was not present when CSM baseline was set. Verify it is legitimate.",
})
// Store it so we don't re-alert
store.SetRaw(key, "new")
}
}
return findings
}
func loadModuleList() []string {
f, err := osFS.Open("/proc/modules")
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
var modules []string
scanner := bufio.NewScanner(f)
for scanner.Scan() {
fields := strings.Fields(scanner.Text())
if len(fields) >= 1 {
modules = append(modules, fields[0])
}
}
return modules
}
// CheckRPMIntegrity verifies critical system binaries haven't been modified.
// Only checks a small set of security-critical packages. Dispatches to
// rpm -V on RHEL-family systems and debsums/dpkg --verify on Debian family.
func CheckRPMIntegrity(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
info := platform.Detect()
switch {
case info.IsRHELFamily():
return checkRPMPackageIntegrity(rpmCriticalPackages)
case info.IsDebianFamily():
return checkDebianPackageIntegrity(debianCriticalPackages)
}
return nil
}
var rpmCriticalPackages = []string{
"openssh-server",
"shadow-utils",
"sudo",
"coreutils",
"util-linux",
"passwd",
}
var debianCriticalPackages = []string{
"openssh-server",
"passwd",
"sudo",
"coreutils",
"util-linux",
"login",
}
func checkRPMPackageIntegrity(packages []string) []alert.Finding {
var findings []alert.Finding
for _, pkg := range packages {
// rpm -V exits non-zero when it finds problems; treat that as
// "findings present" rather than command failure.
out, err := runCmdAllowNonZero("rpm", "-V", pkg)
if err != nil || out == nil {
continue
}
output := strings.TrimSpace(string(out))
if output == "" {
continue
}
// Parse rpm -V output: each line starts with flags
// S=size, 5=md5, T=mtime, etc. We care about S, 5, and M (mode)
for _, line := range strings.Split(output, "\n") {
if len(line) < 9 {
continue
}
flags := line[:9]
file := strings.TrimSpace(line[9:])
// Skip config files (c) and documentation (d)
if strings.Contains(line, " c ") || strings.Contains(line, " d ") {
continue
}
// Check for size (S) or checksum (5) changes. Config and doc files
// were already skipped above; report any tampered executable or
// shared library regardless of directory.
if strings.Contains(flags, "S") || strings.Contains(flags, "5") {
if looksExecutableOrLibrary(file) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "rpm_integrity",
Message: fmt.Sprintf("Modified system binary or library: %s (package: %s)", file, pkg),
Details: fmt.Sprintf("RPM verification flags: %s", flags),
})
}
}
}
}
return findings
}
// checkDebianPackageIntegrity verifies Debian/Ubuntu packages using debsums
// (from the debsums package) when available, falling back to dpkg --verify
// (built into dpkg and always present).
func checkDebianPackageIntegrity(packages []string) []alert.Finding {
// Prefer debsums: it's the Debian equivalent of `rpm -V` and reports
// changed files vs. the md5sum shipped by the package.
if _, err := cmdExec.LookPath("debsums"); err == nil {
return checkDebsums(packages)
}
return checkDpkgVerify(packages)
}
func checkDebsums(packages []string) []alert.Finding {
var findings []alert.Finding
for _, pkg := range packages {
// debsums -c exits 2 when it finds mismatches; treat as findings.
out, err := runCmdAllowNonZero("debsums", "-c", pkg)
if err != nil || out == nil {
continue
}
for _, file := range strings.Split(strings.TrimSpace(string(out)), "\n") {
file = strings.TrimSpace(file)
if file == "" {
continue
}
if !looksExecutableOrLibrary(file) {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "dpkg_integrity",
Message: fmt.Sprintf("Modified system binary or library: %s (package: %s)", file, pkg),
Details: "debsums reported md5 mismatch against the package manifest.",
})
}
}
return findings
}
func checkDpkgVerify(packages []string) []alert.Finding {
var findings []alert.Finding
for _, pkg := range packages {
// dpkg --verify exits 1 when it finds mismatches; treat as findings.
out, err := runCmdAllowNonZero("dpkg", "--verify", pkg)
if err != nil || out == nil {
continue
}
// dpkg --verify prints lines like:
// ??5?????? /usr/bin/passwd
// where position 2 == '5' means md5 mismatch.
for _, line := range strings.Split(strings.TrimSpace(string(out)), "\n") {
if len(line) < 10 {
continue
}
flags := line[:9]
file := strings.TrimSpace(line[9:])
// Skip config files (marked with 'c' after flags).
if strings.Contains(line, " c ") {
continue
}
if !strings.Contains(flags, "5") && !strings.Contains(flags, "S") {
continue
}
if !looksExecutableOrLibrary(file) {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "dpkg_integrity",
Message: fmt.Sprintf("Modified system binary or library: %s (package: %s)", file, pkg),
Details: fmt.Sprintf("dpkg --verify flags: %s", flags),
})
}
}
return findings
}
// looksExecutableOrLibrary reports whether the installed file at path is an
// executable or a shared library. A package-integrity mismatch on one of these
// is the threat we care about (a trojaned binary or .so), so it is reported
// wherever it lives. Judging by file type instead of a directory allowlist
// means an attacker cannot dodge the check by tampering a packaged binary that
// sits outside /usr/bin (e.g. under /usr/lib64, /usr/local, or /opt), while
// changed package-manager state files (manifests, caches, databases) -- which
// are not executable -- do not generate noise.
func looksExecutableOrLibrary(path string) bool {
info, err := osFS.Stat(path)
if err != nil || !info.Mode().IsRegular() {
return false
}
if info.Mode()&0o111 != 0 {
return true
}
// Shared libraries are commonly mode 0644; identify them by ELF magic so a
// trojaned .so is reported regardless of its permission bits.
f, err := osFS.Open(path)
if err != nil {
return false
}
defer func() { _ = f.Close() }()
var magic [4]byte
n, _ := io.ReadFull(f, magic[:])
return n == 4 && string(magic[:]) == "\x7fELF"
}
// CheckMySQLUsers queries for MySQL users with elevated privileges
// that aren't standard cPanel-managed users.
func CheckMySQLUsers(ctx context.Context, _ *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
rows, err := mysqlclient.RootQuery(ctx,
"SELECT user, host FROM mysql.user WHERE Super_priv='Y' AND user NOT IN ('root','mysql.session','mysql.sys','mysql.infoschema','debian-sys-maint')")
if err != nil {
return nil
}
if slices.Contains(rows, "mysql\tlocalhost") {
// Only MariaDB's stock socket account is exempt. mysql.user hides
// alternative auth plugins, so inspect the full definition without
// fetching password hashes. MySQL lacks global_priv; on any lookup
// failure the account remains in the audit for operator review.
stock, err := mysqlclient.RootQuery(ctx, stockMariaDBAccountQuery)
if err == nil && len(stock) == 1 && stock[0] == "1" {
rows = slices.DeleteFunc(rows, func(row string) bool { return row == "mysql\tlocalhost" })
}
}
sort.Strings(rows)
output := strings.TrimSpace(strings.Join(rows, "\n"))
out := []byte(output)
// Track known MySQL superusers
hash := hashBytes(out)
key := "_mysql_super_users"
prev, exists := store.GetRaw(key)
switch {
case !exists:
if output == "" {
break
}
// First run establishes the baseline. The query already excludes the
// standard cPanel/system superusers, so every account here is a
// non-standard privileged account. Baselining them silently would let
// a rogue superuser planted before CSM was installed (the
// CSM-installed-after-breach case) pass as "known" forever. Surface the
// pre-existing set for operator review instead.
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "mysql_superuser",
Message: "Non-standard MySQL superuser accounts present at baseline",
Details: fmt.Sprintf("Review these privileged accounts:\n%s", output),
})
case prev != hash:
details := "Current superusers: none"
if output != "" {
details = fmt.Sprintf("Current superusers:\n%s", output)
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "mysql_superuser",
Message: "MySQL superuser accounts changed",
Details: details,
})
}
store.SetRaw(key, hash)
return findings
}
// Match the authentication definition installed by mariadb-install-db:
// only the mysql OS user can log in, and password authentication is disabled.
// Extra alternatives or socket identity mappings must remain auditable.
const stockMariaDBAccountQuery = `SELECT 1 FROM mysql.global_priv
WHERE User='mysql' AND Host='localhost'
AND BINARY JSON_VALUE(Priv, '$.plugin') = 'mysql_native_password'
AND BINARY JSON_VALUE(Priv, '$.authentication_string') = 'invalid'
AND BINARY JSON_COMPACT(JSON_EXTRACT(Priv, '$.auth_or')) = '[{},{"plugin":"unix_socket"}]'`
// CheckGroupWritablePHP scans for PHP files that are group-writable
// where the group is the web server (nobody/www-data). This allows
// webshells to persist by the web server modifying PHP files.
func CheckGroupWritablePHP(ctx context.Context, _ *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
// Get web server group GIDs
webGroupGIDs := getWebServerGIDs()
if len(webGroupGIDs) == 0 {
return nil
}
homeDirs, _ := GetScanHomeDirs(ctx)
for _, homeEntry := range homeDirs {
if !homeEntry.IsDir() {
continue
}
docRoot := filepath.Join(scanHomeDirPath(homeEntry), "public_html")
scanGroupWritablePHP(docRoot, 4, webGroupGIDs, &findings)
}
return findings
}
func scanGroupWritablePHP(dir string, maxDepth int, webGIDs map[uint32]bool, findings *[]alert.Finding) {
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, entry := range entries {
name := entry.Name()
fullPath := dir + "/" + name
if entry.IsDir() {
// Skip known large/safe dirs
if name == "cache" || name == "node_modules" || name == "vendor" {
continue
}
scanGroupWritablePHP(fullPath, maxDepth-1, webGIDs, findings)
continue
}
if !isExecutablePHPName(strings.ToLower(name)) {
continue
}
info, err := entry.Info()
if err != nil {
continue
}
// Check group-write bit
if info.Mode()&0020 == 0 {
continue
}
// Check if group is web server
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
continue
}
if webGIDs[stat.Gid] {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "group_writable_php",
Message: fmt.Sprintf("Web-server group-writable PHP: %s", fullPath),
Details: fmt.Sprintf("Mode: %s, GID: %d", info.Mode(), stat.Gid),
})
}
}
}
// webServerGroupNames are group names that belong to a web server on any
// platform; the detected platform's own users are added at lookup time.
var webServerGroupNames = map[string]bool{
"nobody": true, "www-data": true, "apache": true, "www": true, "nginx": true,
}
func getWebServerGIDs() map[uint32]bool {
gids := make(map[uint32]bool)
data, err := osFS.ReadFile("/etc/group")
if err != nil {
return gids
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Split(line, ":")
if len(fields) < 3 {
continue
}
name := fields[0]
if webServerGroupNames[name] || slices.Contains(webServerUsers(), name) {
gid, parseErr := strconv.ParseUint(fields[2], 10, 32)
if parseErr != nil {
continue
}
gids[uint32(gid)] = true
}
}
return gids
}
package checks
import (
"io"
"os"
"github.com/pidginhost/csm/internal/jstaint"
"github.com/pidginhost/csm/internal/phptaint"
)
// jsTaintOversizePeekBytes bounds the prefix read that decides whether an
// oversize file could be JavaScript. It matches the PHP peek so both gates
// judge the same window of the same file.
const jsTaintOversizePeekBytes = 64 << 10
// taintSourcePrefixLooks reads a bounded prefix of a file too large to analyze
// and asks pred whether the content could be source of that language.
//
// Every failure answers yes. An unreadable file, a file swapped under us, a
// FIFO that would block: each is a file this scan could not examine, which is
// exactly what the coverage report exists to name. The failure direction that
// matters is the other one -- silently deciding a file was uninteresting and
// dropping it from the report.
func taintSourcePrefixLooks(path string, expected os.FileInfo, limit int64, pred func([]byte) bool) bool {
if reader, ok := osFS.(phpRegularFilePrefixReader); ok {
prefix, err := reader.ReadRegularFilePrefix(path, expected, limit)
if err != nil {
return true
}
return pred(prefix)
}
f, err := osFS.Open(path)
if err != nil {
return true
}
defer func() { _ = f.Close() }()
opened, err := f.Stat()
if err != nil || !opened.Mode().IsRegular() || !sameFileSnapshot(expected, opened) {
return true
}
prefix, err := io.ReadAll(io.LimitReader(f, limit))
if err != nil {
return true
}
after, err := f.Stat()
if err != nil || !sameFileSnapshot(opened, after) {
return true
}
return pred(prefix)
}
// jsFileMayBeJS reports whether a file too large to analyze nonetheless looks
// like JavaScript source, judged only by its leading bytes.
func jsFileMayBeJS(path string, expected os.FileInfo) bool {
return taintSourcePrefixLooks(path, expected, jsTaintOversizePeekBytes, jstaint.MayBeJSSource)
}
// phpFileMayBePHP reports whether a file too large to analyze nonetheless
// looks like PHP source, judged only by its leading bytes.
func phpFileMayBePHP(path string, expected os.FileInfo) bool {
return taintSourcePrefixLooks(path, expected, phpTaintOversizePeekBytes, phptaint.MayBePHPSource)
}
package checks
import (
"bufio"
"fmt"
"io"
"net"
"net/http"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/store"
)
// ThreatDB is a local IP reputation database built from:
// 1. CSM's own block history (permanent)
// 2. Public threat intelligence feeds (updated daily)
// 3. AbuseIPDB as fallback for unknown IPs
type ThreatDB struct {
mu sync.RWMutex
badIPs map[string]string // ip -> source/reason
badIPExpiry map[string]time.Time // ip -> temp-entry expiry; absent = permanent
badNets []*net.IPNet // flat CIDR list for Lookup, rebuilt from feedNets
feedIPs map[string]map[string]struct{} // feed name -> IPs, so overlapping feeds retain ownership
feedNets map[string][]*net.IPNet // feed name -> CIDRs, so a failed feed keeps coverage
whitelist map[string]bool // operator-managed (persisted) IPs to never flag
whitelistMeta map[string]*whitelistEntry // expiry metadata
// configWhitelist mirrors reputation.whitelist from csm.yaml. It is
// replaced wholesale on config reload and never persisted, so the
// file stays the source of truth for these entries.
configWhitelist map[string]bool
lastUpdate time.Time
dbPath string
// Stats for WebUI
PermanentCount int
FeedIPCount int
FeedNetCount int
LastFeedUpdate time.Time
LastUpdated time.Time // tracks when feeds were last successfully loaded
}
var (
globalThreatDB *ThreatDB
threatDBOnce sync.Once
)
// Minimum expected entries per feed - alerts if feed returns less (corrupted/down)
var feedMinEntries = map[string]int{
"spamhaus-drop": 50,
"spamhaus-edrop": 10,
"blocklist-de": 1000,
"cins-army": 5000,
}
// Free public threat intelligence feeds
var threatFeeds = []struct {
name string
url string
}{
{"spamhaus-drop", "https://www.spamhaus.org/drop/drop.txt"},
{"spamhaus-edrop", "https://www.spamhaus.org/drop/edrop.txt"},
{"blocklist-de", "https://lists.blocklist.de/lists/all.txt"},
{"cins-army", "https://cinsscore.com/list/ci-badguys.txt"},
}
// InitThreatDB initializes the global threat database.
func InitThreatDB(statePath string, whitelistIPs []string) *ThreatDB {
threatDBOnce.Do(func() {
db := &ThreatDB{
badIPs: make(map[string]string),
badIPExpiry: make(map[string]time.Time),
whitelist: make(map[string]bool),
configWhitelist: configWhitelistSet(whitelistIPs),
dbPath: filepath.Join(statePath, "threat_db"),
}
_ = os.MkdirAll(db.dbPath, 0700)
db.loadPermanentBlocklist()
db.loadPersistedWhitelist()
db.loadFeedCache()
globalThreatDB = db
})
return globalThreatDB
}
// GetThreatDB returns the global threat database.
func GetThreatDB() *ThreatDB {
return globalThreatDB
}
// SetGlobalThreatDBForTest installs a freshly-constructed threat DB rooted at
// statePath, bypassing the once-guard, and returns a function that restores the
// previous global. For tests only: lets a test exercise the threat-DB path
// without permanently polluting the global for order-dependent tests.
func SetGlobalThreatDBForTest(statePath string) func() {
prev := globalThreatDB
db := &ThreatDB{
badIPs: make(map[string]string),
badIPExpiry: make(map[string]time.Time),
whitelist: make(map[string]bool),
configWhitelist: make(map[string]bool),
dbPath: filepath.Join(statePath, "threat_db"),
}
_ = os.MkdirAll(db.dbPath, 0700)
globalThreatDB = db
return func() { globalThreatDB = prev }
}
func configWhitelistSet(ips []string) map[string]bool {
set := make(map[string]bool, len(ips))
for _, raw := range ips {
if ip := net.ParseIP(strings.TrimSpace(raw)); ip != nil {
set[ip.String()] = true
}
}
return set
}
// SetConfigWhitelist replaces the entries that come from reputation.whitelist
// in csm.yaml. Called on config reload; the hot-reload path reported success
// for that field while lookups kept honouring the startup list. Operator-
// managed entries added at runtime are untouched.
func (db *ThreatDB) SetConfigWhitelist(ips []string) {
db.mu.Lock()
db.configWhitelist = configWhitelistSet(ips)
db.mu.Unlock()
}
// IsConfigWhitelisted reports whether ip is managed by reputation.whitelist.
// Runtime removal must not claim to remove these entries because the config
// remains authoritative and restores them on reload.
func (db *ThreatDB) IsConfigWhitelisted(ip string) bool {
db.mu.RLock()
defer db.mu.RUnlock()
return db.configWhitelist[ip]
}
// ThreatMatch describes why an IP is flagged: which source named it, and
// whether that evidence is permanent (re-flags the address on every future
// sighting) or lapses with the block that recorded it.
type ThreatMatch struct {
Source string
Permanent bool
ExpiresAt time.Time // zero unless the entry lapses
}
// Lookup checks if an IP is in the local threat database.
// Returns (source, true) if found, ("", false) if unknown.
// Whitelisted IPs always return false.
func (db *ThreatDB) Lookup(ip string) (string, bool) {
match, ok := db.LookupMatch(ip)
return match.Source, ok
}
// LookupMatch is Lookup plus the lifetime of the matched evidence, so the
// Web UI can explain why an unblocked IP still scores as malicious.
func (db *ThreatDB) LookupMatch(ip string) (ThreatMatch, bool) {
db.mu.RLock()
defer db.mu.RUnlock()
// Never flag whitelisted IPs
if db.whitelist[ip] || db.configWhitelist[ip] {
return ThreatMatch{}, false
}
// Check exact IP match. Lapsed temp entries no longer count as
// evidence: honouring them here is what turned every temporary
// auto-block into a forever re-flag loop. The leftover row is
// removed by the periodic prune; until then, fall through to the
// feed data and CIDR ranges.
if source, ok := db.badIPs[ip]; ok {
exp, hasExp := db.badIPExpiry[ip]
if !hasExp {
// Feed-owned entries are refreshed from upstream, so only local
// rows (operator or legacy) count as permanent local evidence.
return ThreatMatch{Source: source, Permanent: !isFeedSourceName(source)}, true
}
if time.Now().Before(exp) {
return ThreatMatch{Source: source, ExpiresAt: exp}, true
}
if feed, ok := db.feedSourceLocked(ip); ok {
return ThreatMatch{Source: feed}, true
}
}
// Check CIDR ranges (supports both IPv4 and IPv6)
parsed := net.ParseIP(ip)
if parsed != nil {
for _, cidr := range db.badNets {
if cidr.Contains(parsed) {
return ThreatMatch{Source: "threat-feed-cidr"}, true
}
}
}
return ThreatMatch{}, false
}
// AddPermanent adds an IP to the permanent local blocklist.
// Called for deliberate operator blocks - persists across restarts and
// never expires. Upgrades an existing temp entry to permanent.
func (db *ThreatDB) AddPermanent(ip, reason string) {
db.mu.Lock()
_, exists := db.badIPs[ip]
_, wasTemp := db.badIPExpiry[ip]
_, wasFeed := db.feedSourceLocked(ip)
db.badIPs[ip] = reason
delete(db.badIPExpiry, ip)
db.mu.Unlock()
if sdb := store.Global(); sdb != nil {
_ = sdb.AddPermanentBlock(ip, reason)
return
}
// Only append to the flat-file fallback if this is a new local IP. A temp
// entry or feed-owned entry upgrading to operator evidence must still
// append, or the operator block would disappear after restart.
if exists && !wasTemp && !wasFeed {
return
}
// Fallback: flat-file permanent.txt.
f, err := os.OpenFile(filepath.Join(db.dbPath, "permanent.txt"), os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600)
if err != nil {
return
}
defer func() { _ = f.Close() }()
fmt.Fprintf(f, "%s # %s [%s]\n", ip, reason, time.Now().Format("2006-01-02"))
}
// AddTemporary records an auto-blocked IP for the lifetime of its firewall
// block. Unlike AddPermanent the entry lapses with the block: a permanent
// record turned every temporary auto-block into a forever "known malicious
// IP" that re-flagged (and re-blocked) the address on each later access.
// ttl <= 0 is ignored because auto-block evidence must never become a
// never-expiring threat row.
func (db *ThreatDB) AddTemporary(ip, reason string, ttl time.Duration) {
db.addExpiring(ip, reason, ttl, false, func(expiresAt time.Time) {
if sdb := store.Global(); sdb != nil {
_ = sdb.AddTempBlock(ip, reason, expiresAt)
}
})
}
// AddOperatorTemporary records a timed operator block (the Web UI 24h block)
// for the lifetime of its firewall block. The evidence is operator-sourced
// but lapses with the block, so a mistaken 24h block of a customer address
// does not leave it permanently malicious.
func (db *ThreatDB) AddOperatorTemporary(ip, reason string, ttl time.Duration) {
db.addExpiring(ip, reason, ttl, true, func(expiresAt time.Time) {
if sdb := store.Global(); sdb != nil {
_ = sdb.AddOperatorTempBlock(ip, reason, expiresAt)
}
})
}
func (db *ThreatDB) addExpiring(ip, reason string, ttl time.Duration, operator bool, persist func(time.Time)) {
if ttl <= 0 {
return
}
expiresAt := time.Now().Add(ttl)
db.mu.Lock()
if source, exists := db.badIPs[ip]; exists {
cur, isTemp := db.badIPExpiry[ip]
if !isTemp && (!operator || !isFeedSourceName(source)) {
// Keep permanent local evidence. Operator decisions must also
// survive feed withdrawal; feedIPs retains independent coverage.
db.mu.Unlock()
return
}
if time.Now().Before(cur) && cur.After(expiresAt) {
// Keep the longer live window; re-blocks extend, never truncate.
db.mu.Unlock()
return
}
}
db.badIPs[ip] = reason
db.badIPExpiry[ip] = expiresAt
db.mu.Unlock()
persist(expiresAt)
// Flat-file fallback (pre-migration) deliberately does not persist:
// permanent.txt has no expiry column, so a line there would recreate
// the forever-row this method exists to avoid. The firewall keeps the
// block itself across restarts.
}
// RemoveTemporary removes evidence that only lives as long as a firewall
// block: auto-block rows and timed operator blocks. It leaves permanent
// evidence untouched and restores feed ownership when the IP is
// independently present in a threat feed.
func (db *ThreatDB) RemoveTemporary(ip string) {
db.mu.Lock()
if _, isTemp := db.badIPExpiry[ip]; isTemp {
delete(db.badIPExpiry, ip)
if feed, ok := db.feedSourceLocked(ip); ok {
db.badIPs[ip] = feed
} else {
delete(db.badIPs, ip)
}
}
db.mu.Unlock()
}
// RemovePermanent removes an IP from the permanent blocklist and in-memory DB.
func (db *ThreatDB) RemovePermanent(ip string) {
db.mu.Lock()
delete(db.badIPExpiry, ip)
if feed, ok := db.feedSourceLocked(ip); ok {
db.badIPs[ip] = feed
} else {
delete(db.badIPs, ip)
}
db.mu.Unlock()
if sdb := store.Global(); sdb != nil {
_ = sdb.RemovePermanentBlock(ip)
return
}
// Fallback: rewrite permanent.txt without this IP.
path := filepath.Join(db.dbPath, "permanent.txt")
data, err := osFS.ReadFile(path)
if err != nil {
return
}
var kept []string
for _, line := range strings.Split(string(data), "\n") {
trimmed := strings.TrimSpace(line)
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
kept = append(kept, line)
continue
}
fields := strings.Fields(trimmed)
if len(fields) > 0 && fields[0] == ip {
continue // skip this IP
}
kept = append(kept, line)
}
tmpPath := path + ".tmp"
_ = os.WriteFile(tmpPath, []byte(strings.Join(kept, "\n")+"\n"), 0600)
_ = os.Rename(tmpPath, path)
}
// whitelistEntry tracks an IP with optional expiry.
type whitelistEntry struct {
ExpiresAt time.Time // zero = permanent
}
// AddWhitelist adds an IP to the permanent whitelist.
func (db *ThreatDB) AddWhitelist(ip string) {
db.addWhitelistEntry(ip, time.Time{})
}
// TempWhitelist adds an IP to the whitelist with a TTL.
func (db *ThreatDB) TempWhitelist(ip string, ttl time.Duration) {
db.addWhitelistEntry(ip, time.Now().Add(ttl))
}
func (db *ThreatDB) addWhitelistEntry(ip string, expiresAt time.Time) {
db.mu.Lock()
db.whitelist[ip] = true
if db.whitelistMeta == nil {
db.whitelistMeta = make(map[string]*whitelistEntry)
}
db.whitelistMeta[ip] = &whitelistEntry{ExpiresAt: expiresAt}
delete(db.badIPs, ip)
delete(db.badIPExpiry, ip)
db.mu.Unlock()
if sdb := store.Global(); sdb != nil {
permanent := expiresAt.IsZero()
_ = sdb.AddWhitelistEntry(ip, expiresAt, permanent)
return
}
db.saveWhitelistFile()
}
// RemoveWhitelist removes an operator-managed whitelist entry. Config-managed
// entries are replaced only by SetConfigWhitelist.
func (db *ThreatDB) RemoveWhitelist(ip string) {
db.mu.Lock()
delete(db.whitelist, ip)
delete(db.whitelistMeta, ip)
db.mu.Unlock()
if sdb := store.Global(); sdb != nil {
_ = sdb.RemoveWhitelistEntry(ip)
return
}
db.saveWhitelistFile()
}
// PruneExpiredWhitelist removes expired temporary whitelist entries.
// Called periodically from the daemon heartbeat.
func (db *ThreatDB) PruneExpiredWhitelist() int {
now := time.Now()
pruned := 0
db.mu.Lock()
for ip, entry := range db.whitelistMeta {
if !entry.ExpiresAt.IsZero() && now.After(entry.ExpiresAt) {
delete(db.whitelist, ip)
delete(db.whitelistMeta, ip)
pruned++
}
}
db.mu.Unlock()
if pruned > 0 {
if sdb := store.Global(); sdb != nil {
sdb.PruneExpiredWhitelist()
} else {
db.saveWhitelistFile()
}
fmt.Fprintf(os.Stderr, "[%s] Pruned %d expired whitelist entries\n",
time.Now().Format("2006-01-02 15:04:05"), pruned)
}
return pruned
}
// PruneExpiredThreats removes threat entries whose temp lifetime has
// lapsed, in memory and in the persistent store (which also drops legacy
// no-source auto-block rows written before expiry tagging existed).
// Called periodically from the daemon heartbeat; this is by-design
// lifecycle cleanup, not data retention, so it does not sit behind the
// opt-in retention sweeps.
func (db *ThreatDB) PruneExpiredThreats() int {
now := time.Now()
pruned := 0
db.mu.Lock()
for ip, exp := range db.badIPExpiry {
if now.Before(exp) {
continue
}
delete(db.badIPExpiry, ip)
// A feed may list the IP independently; hand the entry back to
// the feed instead of dropping it until the next feed rebuild.
if feed, ok := db.feedSourceLocked(ip); ok {
db.badIPs[ip] = feed
} else {
delete(db.badIPs, ip)
}
pruned++
}
db.mu.Unlock()
removed := pruned
if sdb := store.Global(); sdb != nil {
// The store count is authoritative: it also covers rows never
// loaded into memory (expired rows skipped at startup and legacy
// auto-block rows).
removed = sdb.PruneExpiredThreats()
}
if removed > 0 {
fmt.Fprintf(os.Stderr, "[%s] Pruned %d expired threat entries\n",
time.Now().Format("2006-01-02 15:04:05"), removed)
}
return removed
}
// isFeedSourceName reports whether a match source names a threat feed
// rather than local evidence.
func isFeedSourceName(source string) bool {
if source == "threat-feed-cidr" {
return true
}
for _, feed := range threatFeeds {
if feed.name == source {
return true
}
}
return false
}
// feedSourceLocked reports which feed lists ip, if any. Caller holds db.mu.
func (db *ThreatDB) feedSourceLocked(ip string) (string, bool) {
for _, feed := range threatFeeds {
if _, ok := db.feedIPs[feed.name][ip]; ok {
return feed.name, true
}
}
return "", false
}
// WhitelistInfo returns all whitelisted IPs with their expiry info.
type WhitelistIP struct {
IP string `json:"ip"`
ExpiresAt *time.Time `json:"expires_at,omitempty"` // nil = permanent
Permanent bool `json:"permanent"`
Configured bool `json:"configured,omitempty"`
}
func (db *ThreatDB) WhitelistedIPs() []WhitelistIP {
db.mu.RLock()
defer db.mu.RUnlock()
var ips []string
for ip := range db.whitelist {
ips = append(ips, ip)
}
for ip := range db.configWhitelist {
if !db.whitelist[ip] {
ips = append(ips, ip)
}
}
sort.Strings(ips)
result := make([]WhitelistIP, len(ips))
for i, ip := range ips {
w := WhitelistIP{IP: ip, Configured: db.configWhitelist[ip]}
if db.whitelist[ip] {
entry := db.whitelistMeta[ip]
w.Permanent = entry == nil || entry.ExpiresAt.IsZero()
if entry != nil && !entry.ExpiresAt.IsZero() {
t := entry.ExpiresAt
w.ExpiresAt = &t
}
}
result[i] = w
}
return result
}
func (db *ThreatDB) saveWhitelistFile() {
path := filepath.Join(db.dbPath, "whitelist.txt")
db.mu.RLock()
var lines []string
for ip := range db.whitelist {
entry := db.whitelistMeta[ip]
if entry != nil && !entry.ExpiresAt.IsZero() {
lines = append(lines, fmt.Sprintf("%s expires=%s", ip, entry.ExpiresAt.Format(time.RFC3339)))
} else {
lines = append(lines, fmt.Sprintf("%s permanent", ip))
}
}
db.mu.RUnlock()
sort.Strings(lines)
tmpPath := path + ".tmp"
_ = os.WriteFile(tmpPath, []byte(strings.Join(lines, "\n")+"\n"), 0600)
_ = os.Rename(tmpPath, path)
}
// loadPersistedWhitelist loads IPs from the bbolt store (if available)
// or from the flat-file whitelist.txt.
func (db *ThreatDB) loadPersistedWhitelist() {
if db.whitelistMeta == nil {
db.whitelistMeta = make(map[string]*whitelistEntry)
}
if sdb := store.Global(); sdb != nil {
entries := sdb.ListWhitelist()
now := time.Now()
for _, e := range entries {
// Skip expired entries
if !e.Permanent && !e.ExpiresAt.IsZero() && now.After(e.ExpiresAt) {
continue
}
db.whitelist[e.IP] = true
db.whitelistMeta[e.IP] = &whitelistEntry{ExpiresAt: e.ExpiresAt}
delete(db.badIPs, e.IP)
}
return
}
// Fallback: flat-file whitelist.txt.
path := filepath.Join(db.dbPath, "whitelist.txt")
f, err := osFS.Open(path)
if err != nil {
return
}
defer func() { _ = f.Close() }()
now := time.Now()
needsRewrite := false
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
fields := strings.Fields(line)
if len(fields) == 0 {
continue
}
ip := fields[0]
if net.ParseIP(ip) == nil {
continue
}
entry := &whitelistEntry{}
// Parse "expires=2026-03-28T19:00:00Z" if present
for _, f := range fields[1:] {
if strings.HasPrefix(f, "expires=") {
if t, err := time.Parse(time.RFC3339, f[8:]); err == nil {
entry.ExpiresAt = t
}
}
}
// Skip expired entries
if !entry.ExpiresAt.IsZero() && now.After(entry.ExpiresAt) {
needsRewrite = true
continue
}
db.whitelist[ip] = true
db.whitelistMeta[ip] = entry
delete(db.badIPs, ip)
}
if needsRewrite {
// Synchronous: the load path runs once at startup, so the cost
// is negligible, and a fire-and-forget goroutine would race the
// daemon's shutdown (potentially leaving a `.tmp` file behind or
// writing a half-serialized whitelist.txt if the process is
// killed before the rewrite lands).
db.saveWhitelistFile()
}
}
// Count returns the total number of entries in the database.
func (db *ThreatDB) Count() int {
db.mu.RLock()
defer db.mu.RUnlock()
return len(db.badIPs) + len(db.badNets)
}
// Stats returns statistics for the WebUI dashboard.
func (db *ThreatDB) Stats() map[string]interface{} {
db.mu.RLock()
defer db.mu.RUnlock()
stats := map[string]interface{}{
"permanent_ips": db.PermanentCount,
"feed_ips": db.FeedIPCount,
"feed_cidrs": db.FeedNetCount,
"total": len(db.badIPs) + len(db.badNets),
"whitelist": db.whitelistCountLocked(),
}
if !db.LastFeedUpdate.IsZero() {
stats["last_update"] = db.LastFeedUpdate
}
return stats
}
func (db *ThreatDB) whitelistCountLocked() int {
count := len(db.whitelist)
for ip := range db.configWhitelist {
if !db.whitelist[ip] {
count++
}
}
return count
}
// LastFeedRefresh returns when feeds last loaded successfully, preferring
// the in-memory timestamp over the persisted one. Zero means never.
func (db *ThreatDB) LastFeedRefresh() time.Time {
db.mu.RLock()
defer db.mu.RUnlock()
if !db.LastUpdated.IsZero() {
return db.LastUpdated
}
return db.lastUpdate
}
// FeedsStale returns true if threat feeds have not been updated in over 7 days.
func (db *ThreatDB) FeedsStale() bool {
db.mu.RLock()
defer db.mu.RUnlock()
lastRefresh := db.LastUpdated
if lastRefresh.IsZero() {
lastRefresh = db.lastUpdate
}
return lastRefresh.IsZero() || feedRefreshStale(lastRefresh, time.Now())
}
// UpdateFeeds downloads fresh threat intelligence feeds.
// Downloads outside the lock, then swaps data under lock to avoid blocking
// lookups. Each feed is swapped independently: a feed that fails to download
// (or fails validation) keeps its previously loaded IPs and CIDRs, so a
// transient outage never wipes that feed's coverage. lastUpdate only advances
// when at least one feed succeeded, otherwise the next cycle would skip the
// retry and serve zero feed data for the whole 20h window.
func (db *ThreatDB) UpdateFeeds() error {
db.mu.RLock()
lastUpdate := db.lastUpdate
db.mu.RUnlock()
// Only update once per day. Strip monotonic readings so suspend time is
// counted, and retry immediately if the wall clock moved behind the last
// refresh instead of suppressing downloads until it catches up.
now := time.Now().Round(0)
lastUpdate = lastUpdate.Round(0)
if !lastUpdate.IsZero() && !lastUpdate.After(now) && now.Sub(lastUpdate) < 20*time.Hour {
return nil
}
client := &http.Client{Timeout: 30 * time.Second}
// Download all feeds OUTSIDE the lock
type feedResult struct {
ips map[string]struct{}
nets []*net.IPNet
}
fresh := make(map[string]feedResult)
for _, feed := range threatFeeds {
ips, nets, err := downloadFeed(client, feed.url, feed.name)
if err != nil {
fmt.Fprintf(os.Stderr, "threatdb: error downloading %s: %v\n", feed.name, err)
continue // keep previous data for this feed
}
// Validate feed - reject partial downloads to avoid losing good data
minExpected := feedMinEntries[feed.name]
if minExpected > 0 && len(ips)+len(nets) < minExpected {
fmt.Fprintf(os.Stderr, "threatdb: WARNING %s returned only %d entries (expected >%d), keeping cached version\n",
feed.name, len(ips)+len(nets), minExpected)
continue // keep previous cached data for this feed
}
ipSet := make(map[string]struct{}, len(ips))
for _, ip := range ips {
ipSet[ip] = struct{}{}
}
fresh[feed.name] = feedResult{ips: ipSet, nets: nets}
// Cache to disk. IPs and CIDRs share one file: CIDR lines are
// distinguishable by "/", so the nets survive a daemon restart
// and legacy IP-only cache files keep loading unchanged.
lines := make([]string, 0, len(ipSet)+len(nets))
for ip := range ipSet {
lines = append(lines, ip)
}
for _, n := range nets {
lines = append(lines, n.String())
}
saveLines(filepath.Join(db.dbPath, feed.name+".txt"), lines)
}
if len(fresh) == 0 {
// Leave lastUpdate untouched so the next cycle retries instead of
// sitting on stale (or empty) data until the skip window expires.
return fmt.Errorf("threatdb: all %d feeds failed, keeping previous data", len(threatFeeds))
}
feedNames := make(map[string]bool, len(threatFeeds))
for _, feed := range threatFeeds {
feedNames[feed.name] = true
}
// Swap data UNDER the lock - fast operation
db.mu.Lock()
if db.feedIPs == nil {
db.feedIPs = make(map[string]map[string]struct{})
for ip, source := range db.badIPs {
if feedNames[source] {
if db.feedIPs[source] == nil {
db.feedIPs[source] = make(map[string]struct{})
}
db.feedIPs[source][ip] = struct{}{}
}
}
}
if db.feedNets == nil {
db.feedNets = make(map[string][]*net.IPNet)
}
for name, result := range fresh {
// Replace only this feed's previous entries; failed feeds keep
// their per-feed ownership and are merged back into badIPs below.
db.feedIPs[name] = result.ips
db.feedNets[name] = result.nets
}
totalIPs, totalNets := db.rebuildFeedLookup(feedNames)
now = time.Now()
db.lastUpdate = now
db.FeedIPCount = totalIPs
db.FeedNetCount = totalNets
db.LastFeedUpdate = now
db.LastUpdated = now
db.mu.Unlock()
// Save timestamp
_ = os.WriteFile(filepath.Join(db.dbPath, "last_update"),
[]byte(now.Format(time.RFC3339)), 0600)
fmt.Fprintf(os.Stderr, "threatdb: updated %d IPs + %d CIDR ranges (%d/%d feeds succeeded)\n",
totalIPs, totalNets, len(fresh), len(threatFeeds))
return nil
}
// loadPermanentBlocklist loads the permanent blocklist from bbolt store
// (if available) or from the flat-file permanent.txt.
func (db *ThreatDB) loadPermanentBlocklist() {
if sdb := store.Global(); sdb != nil {
blocks := sdb.AllPermanentBlocks()
now := time.Now()
live := 0
for _, b := range blocks {
// Expired rows stay out of memory, including legacy no-source
// auto-block rows (store classifies them in Expired): loading
// them would resume the permablock loop on upgraded hosts. The
// periodic prune deletes them from disk.
if b.Expired(now) {
continue
}
db.badIPs[b.IP] = "permanent-blocklist"
if !b.ExpiresAt.IsZero() {
db.badIPExpiry[b.IP] = b.ExpiresAt
}
live++
}
db.PermanentCount = live
return
}
// Fallback: flat-file permanent.txt.
path := filepath.Join(db.dbPath, "permanent.txt")
f, err := osFS.Open(path)
if err != nil {
return
}
defer func() { _ = f.Close() }()
seen := make(map[string]bool)
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
ip := strings.Fields(line)[0]
// Support both IPv4 and IPv6
if net.ParseIP(ip) != nil && !seen[ip] {
db.badIPs[ip] = "permanent-blocklist"
seen[ip] = true
}
}
db.PermanentCount = len(seen)
// Compact the file if it has duplicates (rewrite with unique entries)
if db.PermanentCount > 0 {
compactPermanentFile(path, seen)
}
}
// compactPermanentFile rewrites the permanent blocklist with unique entries only.
func compactPermanentFile(path string, uniqueIPs map[string]bool) {
// Read all lines to preserve comments/reasons
data, err := osFS.ReadFile(path)
if err != nil {
return
}
lines := strings.Split(string(data), "\n")
if len(lines) <= len(uniqueIPs)+5 {
return // not worth compacting - minimal duplicates
}
// Rewrite with deduplication
seen := make(map[string]bool)
var unique []string
for _, line := range lines {
trimmed := strings.TrimSpace(line)
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
unique = append(unique, line)
continue
}
ip := strings.Fields(trimmed)[0]
if !seen[ip] {
seen[ip] = true
unique = append(unique, line)
}
}
tmpPath := path + ".tmp"
_ = os.WriteFile(tmpPath, []byte(strings.Join(unique, "\n")+"\n"), 0600)
_ = os.Rename(tmpPath, path)
}
// loadFeedCache loads cached feed data from disk.
func (db *ThreatDB) loadFeedCache() {
data, err := osFS.ReadFile(filepath.Join(db.dbPath, "last_update"))
if err == nil {
if t, err := time.Parse(time.RFC3339, strings.TrimSpace(string(data))); err == nil {
db.lastUpdate = t
db.LastFeedUpdate = t
db.LastUpdated = t
}
}
if db.feedNets == nil {
db.feedNets = make(map[string][]*net.IPNet)
}
if db.feedIPs == nil {
db.feedIPs = make(map[string]map[string]struct{})
}
feedNames := make(map[string]bool, len(threatFeeds))
incomplete := false
for _, feed := range threatFeeds {
feedNames[feed.name] = true
db.feedIPs[feed.name] = make(map[string]struct{})
db.feedNets[feed.name] = nil
cachePath := filepath.Join(db.dbPath, feed.name+".txt")
_, statErr := osFS.Stat(cachePath)
cacheExists := statErr == nil
for _, line := range loadLines(cachePath) {
// Cache files mix plain IPs and CIDR lines ("/" marks a
// CIDR); legacy IP-only files parse the same way.
if strings.Contains(line, "/") {
if _, cidr, err := net.ParseCIDR(line); err == nil {
db.feedNets[feed.name] = append(db.feedNets[feed.name], cidr)
}
continue
}
db.feedIPs[feed.name][line] = struct{}{}
}
// The download path refuses a feed below its floor; a cache below it
// is a truncated write, not a smaller feed. Serve nothing from it and
// drop the update marker so the next cycle downloads instead of
// honouring the 20-hour skip. A feed with no cache at all (never
// fetched, or every download refused) is not truncated and must not
// force a refresh of the others.
if minExpected := feedMinEntries[feed.name]; cacheExists && minExpected > 0 {
if n := len(db.feedIPs[feed.name]) + len(db.feedNets[feed.name]); n < minExpected {
if n > 0 {
fmt.Fprintf(os.Stderr, "threatdb: WARNING cached %s holds only %d entries (expected >%d); ignoring it until refreshed\n", feed.name, n, minExpected)
}
db.feedIPs[feed.name] = make(map[string]struct{})
db.feedNets[feed.name] = nil
incomplete = true
}
}
}
if incomplete {
db.lastUpdate = time.Time{}
}
db.FeedIPCount, db.FeedNetCount = db.rebuildFeedLookup(feedNames)
// Warn on startup if feeds are stale
if db.LastUpdated.IsZero() && db.FeedIPCount == 0 && db.FeedNetCount == 0 {
fmt.Fprintf(os.Stderr, "threatdb: WARNING no threat feed data loaded, feeds have never been fetched\n")
} else if !db.LastUpdated.IsZero() && time.Since(db.LastUpdated) > 7*24*time.Hour {
fmt.Fprintf(os.Stderr, "threatdb: WARNING threat feeds are stale (last updated %s, %d days ago)\n",
db.LastUpdated.Format("2006-01-02"), int(time.Since(db.LastUpdated).Hours()/24))
}
}
// rebuildFeedLookup rebuilds the flat lookup maps from per-feed data. The
// caller must hold db.mu, or be in the startup load path before publication.
func (db *ThreatDB) rebuildFeedLookup(feedNames map[string]bool) (int, int) {
preserved := make(map[string]string, len(db.badIPs))
for ip, source := range db.badIPs {
if feedNames[source] {
continue
}
preserved[ip] = source
}
db.badIPs = preserved
totalIPs := 0
for _, feed := range threatFeeds {
for ip := range db.feedIPs[feed.name] {
totalIPs++
if _, exists := db.badIPs[ip]; !exists {
db.badIPs[ip] = feed.name
}
}
}
var nets []*net.IPNet
for _, feed := range threatFeeds {
nets = append(nets, db.feedNets[feed.name]...)
}
db.badNets = nets
return totalIPs, len(nets)
}
func downloadFeed(client *http.Client, url, name string) ([]string, []*net.IPNet, error) {
resp, err := client.Get(url)
if err != nil {
return nil, nil, err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != 200 {
return nil, nil, fmt.Errorf("HTTP %d", resp.StatusCode)
}
limited := io.LimitReader(resp.Body, 10*1024*1024)
scanner := bufio.NewScanner(limited)
var ips []string
var nets []*net.IPNet
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, ";") {
continue
}
if idx := strings.IndexAny(line, ";#"); idx > 0 {
line = strings.TrimSpace(line[:idx])
}
if strings.Contains(line, "/") {
_, cidr, err := net.ParseCIDR(line)
if err == nil {
nets = append(nets, cidr)
}
continue
}
ip := strings.Fields(line)[0]
if net.ParseIP(ip) != nil {
ips = append(ips, ip)
}
}
return ips, nets, nil
}
func saveLines(path string, lines []string) {
sort.Strings(lines) // sorted for diffing
// Written whole and renamed into place: an in-place truncate left a
// crash mid-write serving a partial feed for the next 20 hours.
data := strings.Join(lines, "\n")
if len(lines) > 0 {
data += "\n"
}
if err := atomicio.AtomicWrite(path, 0o600, []byte(data)); err != nil {
fmt.Fprintf(os.Stderr, "threatdb: error saving %s: %v\n", path, err)
}
}
func loadLines(path string) []string {
f, err := osFS.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
var lines []string
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line != "" && !strings.HasPrefix(line, "#") {
lines = append(lines, line)
}
}
return lines
}
package checks
import (
"bytes"
"context"
"fmt"
"path/filepath"
"regexp"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
)
// timThumbFixedVersion is the last TimThumb release. The project was abandoned
// in 2014; anything below this carries the CVE-2011-4106 remote-code-execution
// bug that let attackers write PHP into the image cache -- the entry vector in
// the 2026-07-20 cross-account compromise.
const timThumbFixedVersion = "2.8.14"
// timThumbReadHead bounds how much of each candidate file is read: the version
// define and feature constants are always near the top.
const timThumbReadHead = 16 * 1024
// timThumbScanDepth bounds the per-docroot recursive descent. TimThumb ships
// bundled in themes, sometimes several directories deep (framework/scripts).
const timThumbScanDepth = 10
var timThumbVersionRE = timThumbDefineMatcher("VERSION", `['"]([0-9]+(?:\.[0-9]+)*)['"]`)
// timThumbFeatureRE holds the boolean feature constants that make a copy
// directly exploitable, compiled once rather than per scanned file.
var timThumbFeatureRE = map[string]*regexp.Regexp{
timThumbWebshot: timThumbFeatureMatcher(timThumbWebshot),
timThumbExternal: timThumbFeatureMatcher(timThumbExternal),
timThumbAllExternal: timThumbFeatureMatcher(timThumbAllExternal),
}
const (
timThumbWebshot = "WEBSHOT_ENABLED"
timThumbExternal = "ALLOW_EXTERNAL"
timThumbAllExternal = "ALLOW_ALL_EXTERNAL_SITES"
)
// timThumbFeatureMatcher builds the define('NAME', true) matcher for a feature
// constant.
func timThumbFeatureMatcher(name string) *regexp.Regexp {
return timThumbDefineMatcher(name, `(true)`)
}
// timThumbFeatureDisabledRE matches define('NAME', false). Absence of a define
// is deliberately NOT treated as disabled: TimThumb sets these defaults
// elsewhere, so a file that simply does not mention the constant tells us
// nothing about whether the feature is on.
var timThumbFeatureDisabledRE = map[string]*regexp.Regexp{
timThumbWebshot: timThumbDisabledMatcher(timThumbWebshot),
timThumbExternal: timThumbDisabledMatcher(timThumbExternal),
timThumbAllExternal: timThumbDisabledMatcher(timThumbAllExternal),
}
func timThumbDisabledMatcher(name string) *regexp.Regexp {
return timThumbDefineMatcher(name, `(false)`)
}
// timThumbFeatureExplicitlyDisabled reports a visible define(NAME, false).
// An unknown name compiles on demand rather than silently reporting false, so a
// new caller cannot get a wrong answer; the map itself is never mutated, which
// keeps concurrent scans race-free.
func timThumbFeatureExplicitlyDisabled(head []byte, name string) bool {
re, ok := timThumbFeatureDisabledRE[name]
if !ok {
re = timThumbDisabledMatcher(name)
}
if timThumbDefineValue(head, re) == "" {
return false
}
return !timThumbFeatureEnabled(head, name)
}
// timThumbCandidateName reports whether a filename is a TimThumb-style script by
// convention. Content confirmation happens in looksLikeTimThumb.
func timThumbCandidateName(nameLower string) bool {
return nameLower == "timthumb.php" || nameLower == "thumb.php"
}
// looksLikeTimThumb confirms a file is actually TimThumb rather than an
// unrelated theme thumbnail helper. It keys on constants unique to TimThumb so
// a generic thumb.php is never flagged.
func looksLikeTimThumb(head []byte) bool {
lower := bytes.ToLower(head)
if bytes.Contains(lower, []byte("timthumb")) {
return true
}
// Older or renamed copies without the header name still carry these
// TimThumb-specific configuration constants.
if bytes.Contains(lower, []byte("block_external_leechers")) &&
(bytes.Contains(lower, []byte("webshot")) || bytes.Contains(lower, []byte("allow_external"))) {
return true
}
return false
}
// parseTimThumbVersion extracts the value of TimThumb's VERSION define, or ""
// when it is absent.
func parseTimThumbVersion(head []byte) string {
return timThumbDefineValue(head, timThumbVersionRE)
}
func timThumbDefineMatcher(name, valuePattern string) *regexp.Regexp {
return regexp.MustCompile(`(?i)^define\s*\(\s*['"]` + regexp.QuoteMeta(name) +
`['"]\s*,\s*` + valuePattern + `\s*\)`)
}
// timThumbDefineValue returns a literal from an executable define() call.
// Comment and string text cannot prove a copy is patched or a feature is off,
// so walk only PHP identifiers outside those regions before applying the small
// call matcher.
func timThumbDefineValue(head []byte, matcher *regexp.Regexp) string {
code := stripPHPCommentsFromCode(string(head))
inPHP := false
for i := 0; i < len(code); i++ {
if !inPHP {
if i+1 < len(code) && code[i] == '<' && code[i+1] == '?' &&
(i+5 > len(code) || !strings.EqualFold(code[i:i+5], "<?xml")) {
inPHP = true
i++
}
continue
}
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label) - 1
continue
}
if isPHPQuote(code[i]) || code[i] == '`' {
i = skipPHPString(code, i)
continue
}
if i+1 < len(code) && code[i] == '?' && code[i+1] == '>' {
inPHP = false
i++
continue
}
if !isPHPIdentifierStart(code[i]) {
continue
}
end := i + 1
for end < len(code) && isPHPIdentifierPart(code[end]) {
end++
}
if strings.EqualFold(code[i:end], "define") {
if match := matcher.FindStringSubmatch(code[i:]); len(match) == 2 {
return match[1]
}
}
i = end - 1
}
return ""
}
// timThumbVersionLess reports whether dotted numeric version a is older than b.
func timThumbVersionLess(a, b string) bool {
as, bs := strings.Split(a, "."), strings.Split(b, ".")
for i := 0; i < len(as) || i < len(bs); i++ {
var av, bv int
if i < len(as) {
av, _ = strconv.Atoi(as[i])
}
if i < len(bs) {
bv, _ = strconv.Atoi(bs[i])
}
if av != bv {
return av < bv
}
}
return false
}
// assessTimThumb grades a confirmed TimThumb file. A version below the last
// patched release (or an unparseable one) carries the known RCE and is High; a
// patched-but-abandoned copy is a Warning to remove.
func assessTimThumb(head []byte) (alert.Severity, []string) {
var reasons []string
version := parseTimThumbVersion(head)
exploitable := false
switch {
case version == "":
reasons = append(reasons, "version could not be determined")
exploitable = true
case timThumbVersionLess(version, timThumbFixedVersion):
reasons = append(reasons, fmt.Sprintf("version %s predates the last patch %s (CVE-2011-4106 remote code execution)", version, timThumbFixedVersion))
exploitable = true
default:
reasons = append(reasons, fmt.Sprintf("version %s is deprecated and unmaintained", version))
}
if timThumbFeatureEnabled(head, timThumbWebshot) {
reasons = append(reasons, "WebShot feature is enabled (remote command execution)")
exploitable = true
}
if timThumbFeatureEnabled(head, timThumbExternal) {
reasons = append(reasons, "external image fetching is enabled (SSRF and cache-poisoning surface)")
exploitable = true
}
if exploitable {
return alert.High, reasons
}
return alert.Warning, reasons
}
// timThumbMitigated reports whether a copy has nothing left to act on: the final
// release, with every remote-fetch and screenshot feature explicitly off. The
// remote-fetch path is what CVE-2011-4106 abused to write PHP into the image
// cache, so with it disabled on the last version there is no fix to apply and no
// residual exposure -- only the fact that the project is abandoned, which
// re-reporting every scan does not help anyone act on.
//
// Anything else (older release, unreadable version, any risky feature on) is
// still reported. This is a judgement about the file's own contents, never about
// where it lives.
func timThumbMitigated(head []byte) bool {
if parseTimThumbVersion(head) != timThumbFixedVersion {
return false
}
for _, feature := range []string{timThumbWebshot, timThumbExternal, timThumbAllExternal} {
if !timThumbFeatureExplicitlyDisabled(head, feature) {
return false
}
}
return true
}
// timThumbFeatureEnabled reports whether a TimThumb boolean constant is defined
// true. Matches define('NAME', true) with optional whitespace. Unknown names
// fall back to compiling on demand so a new caller cannot silently read false.
func timThumbFeatureEnabled(head []byte, name string) bool {
re, ok := timThumbFeatureRE[name]
if !ok {
re = timThumbFeatureMatcher(name)
}
return timThumbDefineValue(head, re) != ""
}
// CheckVulnerableTimThumb scans web document roots for bundled TimThumb scripts
// and reports each confirmed instance. Detection-only: TimThumb is legitimate
// (if abandoned) code, so it is never auto-quarantined -- removing it would
// break the theme.
func CheckVulnerableTimThumb(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
var findings []alert.Finding
homeDirs, _ := GetScanHomeDirs(ctx)
for _, homeEntry := range homeDirs {
if ctx.Err() != nil {
return findings
}
if !homeEntry.IsDir() {
continue
}
homeDir := scanHomeDirPath(homeEntry)
docRoots := []string{filepath.Join(homeDir, "public_html")}
subDirs, _ := osFS.ReadDir(homeDir)
for _, sd := range subDirs {
if sd.IsDir() && sd.Name() != "public_html" && sd.Name() != "mail" &&
!strings.HasPrefix(sd.Name(), ".") && sd.Name() != "etc" &&
sd.Name() != "logs" && sd.Name() != "ssl" && sd.Name() != "tmp" {
docRoots = append(docRoots, filepath.Join(homeDir, sd.Name()))
}
}
for _, docRoot := range docRoots {
scanForTimThumb(ctx, docRoot, timThumbScanDepth, &findings)
if ctx.Err() != nil {
return findings
}
}
}
return findings
}
// scanForTimThumb recursively walks dir for TimThumb scripts and appends a
// finding per confirmed instance.
func scanForTimThumb(ctx context.Context, dir string, maxDepth int, findings *[]alert.Finding) {
if ctx.Err() != nil || maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
return
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
if entry.IsDir() {
scanForTimThumb(ctx, fullPath, maxDepth-1, findings)
continue
}
if !timThumbCandidateName(strings.ToLower(name)) {
continue
}
head := readFileHead(fullPath, timThumbReadHead)
if head == nil || !looksLikeTimThumb(head) {
continue
}
if timThumbMitigated(head) {
continue
}
severity, reasons := assessTimThumb(head)
*findings = append(*findings, alert.Finding{
Severity: severity,
Check: "vulnerable_timthumb",
Message: fmt.Sprintf("Vulnerable TimThumb image resizer: %s", fullPath),
Details: "Remove or replace it. " + strings.Join(reasons, "; ") + ".",
FilePath: fullPath,
})
}
}
package checks
import (
"fmt"
"os"
"strconv"
"strings"
"sync"
"time"
)
const uidCacheMissTTL = time.Minute
// uidCache caches uid -> username from /etc/passwd. The first Lookup of an
// unknown uid reads and parses the file; subsequent lookups return from the
// in-memory map. Process-lifetime: callers that need fresh data after a
// useradd should call Refresh().
type uidCache struct {
lastHomeRead time.Time
path string
mu sync.RWMutex
m map[uint32]string
homes map[string]string
}
var defaultUIDCache = newUIDCache("/etc/passwd")
func newUIDCache(path string) *uidCache {
return &uidCache{path: path, m: map[uint32]string{}, homes: map[string]string{}}
}
// LookupUser returns the username for uid, or "uid:<n>" if not resolvable.
// Safe for concurrent use; the underlying cache is shared across the daemon.
func LookupUser(uid uint32) string { return defaultUIDCache.Lookup(uid) }
// swapDefaultUIDCacheForTest replaces defaultUIDCache with a cache pointed at
// path and returns a function that restores the original. Test-only helper:
// existing tests stub /etc/passwd via osFS, but the cache reads the real file
// directly (so the daemon never burns syscalls per-event). This shim lets the
// tests stage a fixture file and have LookupUser read from it.
func swapDefaultUIDCacheForTest(path string) func() {
prev := defaultUIDCache
defaultUIDCache = newUIDCache(path)
return func() { defaultUIDCache = prev }
}
// SwapUIDCacheForTest is the exported form of swapDefaultUIDCacheForTest for
// packages whose producers resolve uids through LookupUser.
func SwapUIDCacheForTest(path string) func() { return swapDefaultUIDCacheForTest(path) }
func (c *uidCache) Lookup(uid uint32) string {
c.mu.RLock()
if name, ok := c.m[uid]; ok {
c.mu.RUnlock()
return name
}
c.mu.RUnlock()
c.mu.Lock()
defer c.mu.Unlock()
if name, ok := c.m[uid]; ok {
return name
}
c.parseLocked()
if name, ok := c.m[uid]; ok {
return name
}
miss := fmt.Sprintf("uid:%d", uid)
c.m[uid] = miss
return miss
}
// Refresh drops the cache. The next Lookup re-reads /etc/passwd.
func (c *uidCache) Refresh() {
c.mu.Lock()
c.m = map[uint32]string{}
c.homes = map[string]string{}
c.lastHomeRead = time.Time{}
c.mu.Unlock()
}
// HomeDir returns the home directory recorded for name, or "" when the user
// is unknown. Producers use it to tell a hosting account from a system user.
func (c *uidCache) HomeDir(name string) string {
c.mu.RLock()
home, ok := c.homes[name]
fresh := time.Since(c.lastHomeRead) < uidCacheMissTTL
c.mu.RUnlock()
if ok || fresh {
return home
}
c.mu.Lock()
defer c.mu.Unlock()
if home, ok = c.homes[name]; ok || time.Since(c.lastHomeRead) < uidCacheMissTTL {
return home
}
// Throttle misses across all names without storing attacker-chosen keys.
// Expiry lets newly provisioned accounts become attributable after a miss.
c.parseLocked()
return c.homes[name]
}
// parseLocked replaces the cache contents with a fresh scan. Caller holds the
// write lock.
func (c *uidCache) parseLocked() {
c.lastHomeRead = time.Now()
data, err := os.ReadFile(c.path)
if err != nil {
return
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.SplitN(line, ":", 7)
if len(fields) < 3 {
continue
}
uid64, err := strconv.ParseUint(fields[2], 10, 32)
if err != nil {
continue
}
c.m[uint32(uid64)] = fields[0]
if len(fields) >= 6 {
c.homes[fields[0]] = fields[5]
}
}
}
package checks
import (
"net"
"net/url"
"strings"
)
// URL reputation — attack-indicator based classifier for external
// <script src="..."> URLs embedded in WordPress content.
//
// PHILOSOPHY
//
// Earlier versions of this package classified script sources against a
// hardcoded allowlist of "known safe" domains (Google Tag Manager,
// Cloudflare CDN, HubSpot, Stripe, etc.). In practice the allowlist is
// unmaintainable: every new widget service (OneTrust, Issuu, regional
// video embeds, tax-form widgets, etc.) adds an entry and operators
// still see HIGH-severity findings for legitimate third-party embeds.
//
// This file takes the inverse approach: rather than asking "is this
// domain on my list of safe services?", it asks "does this URL show
// attacker-characteristic markers?". A <script src> fires a finding
// only when at least one attack indicator is present:
//
// - the host is a raw IP address (attackers dodge domain reputation);
// - the host TLD is on a well-known abused-TLD list (.tk, .ml, .ga,
// .cf, .gq free-abuse; .top, .icu, .click, .pw and similar cheap
// gTLDs per Spamhaus recurring bad-TLD reports);
// - the host is on the existing short known-bad-exfil list
// (Cloudflare Workers free tier, Pastebin raw, GitHub Gist raw —
// legitimate content rarely loads from these hosts);
// - the scheme is plaintext HTTP (an external JS loader without TLS
// is a plaintext MITM target regardless of the destination);
// - the host is empty, contains no dot, or otherwise fails basic
// FQDN validation.
//
// The knownSafeDomains list is retained as a FAST-PATH tie-breaker: if
// the host matches a well-known service we return "not malicious"
// immediately, skipping further analysis. It is an optimization, not
// the primary filter. Unknown hosts on unremarkable TLDs (e.g.
// onetrust.com, issuu.com, trilulilu.ro, formular230.ro) pass because
// they have zero attack indicators — no allowlist growth needed.
//
// TRADE-OFF
//
// An attacker who hosts payload on a compromised mainstream domain
// (.com/.org/.net HTTPS, normal-looking path) is not caught. This gap
// existed under the prior allowlist too — it is a fundamental limit of
// URL-only classification. Closing it requires threat-intelligence
// correlation or content-based JS analysis, both of which are out of
// scope here. The defensive value is in raising the bar for casual
// injection, not in defeating sophisticated adversaries.
// abusedTLDs are top-level domains whose registrations are cheap,
// unverified, and overwhelmingly abused for phishing, malware, and SEO
// spam per recurring Spamhaus and KnowBe4 reports.
//
// The entry bar is high: a TLD is included only if (a) it has shown up
// in top-10 abuse rankings across multiple reporting years, AND
// (b) legitimate business usage is rare. Mixed-use TLDs with
// significant legitimate traffic (.xyz, .online, .site, .live, .space)
// are intentionally excluded to keep false-positive rates low.
//
// Entries are stored without the leading dot; comparison strips the
// leading dot from the observed TLD before lookup.
var abusedTLDs = map[string]bool{
// Former Freenom TLDs — free registration, no verification.
// Near-100% abuse rate; Freenom itself was shut down in 2023 but
// the TLDs remain in DNS.
"tk": true,
"ml": true,
"ga": true,
"cf": true,
"gq": true,
// Cheap new gTLDs consistently in the Spamhaus top-abused list.
"top": true,
"icu": true,
"click": true,
"pw": true,
"loan": true,
"work": true,
"download": true,
// Spamhaus badlist recurring entries — legitimate usage is
// essentially nonexistent at meaningful volume.
"kim": true,
"gdn": true,
"stream": true,
"bid": true,
"racing": true,
"win": true,
"party": true,
"science": true,
"trade": true,
}
// knownBadExfilHosts are hosts where legitimate WordPress content
// essentially never loads JavaScript from, but attackers routinely do.
// The list is intentionally short — see knownSafeDomains for the
// inverse fast-path.
var knownBadExfilHosts = []string{
// Cloudflare Workers free-tier subdomain (e.g. x.workers.dev).
// Legitimate apps host on custom domains; free .workers.dev is a
// common payload drop.
".workers.dev",
// Pastebin and GitHub Gist raw endpoints — legitimate sites do not
// load JS from these paths.
"pastebin.com",
"gist.githubusercontent.com",
// bit.ly and similar URL shorteners in a <script src> are a very
// strong attack signal — no legitimate embed uses a shortener for
// a JS asset.
"bit.ly",
"tinyurl.com",
"is.gd",
"cutt.ly",
}
// scriptSrcStrongReason classifies a <script src> URL by structural
// attack indicators that are context-independent: a raw IP host, an
// abused TLD, a known-bad exfil host, or an empty/unparseable/no-TLD
// host. These markers are rare-to-nonexistent in legitimate content of
// any age and remain valid signals whether the script appears in a
// freshly-written wp_options value or in decade-old post_content.
//
// Callers that operate on storage which is expected to hold current
// configuration (wp_options) should prefer scriptSrcMaliciousReason,
// which layers the plaintext-HTTP indicator on top. The HTTP signal
// catches attacker convenience ("don't bother with TLS") in fresh
// configuration but produces false positives on legacy author content
// where pre-TLS embeds are normal.
//
// The function does not consult knownSafeDomains — callers should do
// that first as a fast path. Separating the two concerns keeps this
// function single-purpose and trivially testable.
func scriptSrcStrongReason(rawURL string) (bool, string) {
normalised := rawURL
if strings.HasPrefix(normalised, "//") {
normalised = "https:" + normalised
}
u, err := url.Parse(normalised)
if err != nil || u == nil {
return true, "unparseable URL"
}
host := strings.ToLower(u.Hostname())
if host == "" {
return true, "empty host"
}
if ip := net.ParseIP(host); ip != nil {
return true, "raw IP address host"
}
for _, bad := range knownBadExfilHosts {
// Match the host itself or a subdomain of it, never a longer name
// that merely ends in the same characters (orbit.ly is not bit.ly).
bare := strings.TrimPrefix(bad, ".")
if host == bare || strings.HasSuffix(host, "."+bare) {
return true, "known-bad exfil host: " + bad
}
}
lastDot := strings.LastIndexByte(host, '.')
if lastDot < 0 || lastDot == len(host)-1 {
return true, "host without valid TLD"
}
tld := host[lastDot+1:]
if abusedTLDs[tld] {
return true, "abused TLD: ." + tld
}
return false, ""
}
// scriptSrcMaliciousReason classifies a <script src> URL with the full
// indicator set: everything scriptSrcStrongReason flags, plus plaintext
// HTTP. This is the classifier for wp_options and similar configuration
// storage, where a plaintext HTTP external script loader is a signal on
// its own (a site's analytics configuration should be HTTPS in 2026).
//
// The function does not consult knownSafeDomains — callers should do
// that first as a fast path.
func scriptSrcMaliciousReason(rawURL string) (bool, string) {
// Structural markers take precedence: a raw-IP host over HTTP should
// report the IP, not the scheme, because the scheme is merely the
// delivery method while the IP is the identity of the attacker
// infrastructure.
if bad, reason := scriptSrcStrongReason(rawURL); bad {
return true, reason
}
// Plaintext HTTP for an external script is a strong indicator in
// configuration storage. Protocol-relative URLs (//host/path)
// inherit the page's scheme and are not flagged here.
normalised := rawURL
if strings.HasPrefix(normalised, "//") {
normalised = "https:" + normalised
}
u, err := url.Parse(normalised)
if err != nil || u == nil {
// scriptSrcStrongReason would have caught this; defensive only.
return true, "unparseable URL"
}
if strings.EqualFold(u.Scheme, "http") && !strings.HasPrefix(rawURL, "//") {
return true, "plaintext HTTP external script"
}
return false, ""
}
// isAttackerScriptURL is the caller-facing predicate for contexts where
// a plaintext-HTTP external script is a signal on its own (wp_options
// and similar configuration storage). It combines the known-safe fast
// path with the strict attack-indicator classifier.
//
// The order matters: the fast path is checked first because it lets us
// short-circuit common legitimate widgets (Google Tag Manager,
// Cloudflare CDN, Stripe) without parsing the URL. Only unknown hosts
// are subjected to the attack-indicator analysis.
func isAttackerScriptURL(rawURL string) bool {
if isSafeScriptDomain(rawURL) {
return false
}
bad, _ := scriptSrcMaliciousReason(rawURL)
return bad
}
// isAttackerScriptURLInPost is the caller-facing predicate for
// post_content classification. It uses scriptSrcStrongReason, which
// omits the plaintext-HTTP indicator: legacy author embeds from the
// pre-TLS era are legitimate content, not injection, and must not
// produce db_post_injection findings.
//
// Fresh attacker injections still flag because they almost always point
// at structural markers (raw IP hosts, abused TLDs, cheap exfil hosts)
// rather than at an unremarkable mainstream-TLD host — and an attacker
// who did somehow land on a plaintext-HTTP mainstream host URL would
// still be caught by other checks (obfuscated_php_realtime scanning the
// attacker's dropper, remote payload URLs in the served page, etc.).
func isAttackerScriptURLInPost(rawURL string) bool {
if isSafeScriptDomain(rawURL) {
return false
}
bad, _ := scriptSrcStrongReason(rawURL)
return bad
}
package checks
import "bytes"
// .user.ini cPanel-managed signature detection.
//
// cPanel's MultiPHP INI Editor writes .user.ini files with a fixed
// four-line header that begins:
//
// ; cPanel-generated php ini directives, do not edit
// ; Manual editing of this file may result in unexpected behavior.
// ; To make changes to this file, use the cPanel MultiPHP INI Editor ...
// ; For more information, read our documentation ...
//
// When this header is present the values in the file reflect operator
// choices made through cPanel's UI (max_execution_time=0 for a backup
// importer, display_errors=On for a staging account, etc.). These are
// not attacker signals. The severity of findings on values in a
// cPanel-managed file should be reduced to informational.
//
// When the header is absent we cannot tell whether the site owner
// hand-edited the file or an attacker planted it. In that case we
// preserve the original severity.
//
// Detection rule (minimal and precise):
//
// The signature must appear on the FIRST non-blank line of the file.
//
// Reasoning: cPanel always writes the header at the top of the file
// and rewrites the whole file on every edit through the UI. An
// attacker appending content below keeps the header position intact.
// An attacker prepending content above pushes the header down; the
// file is then no longer a regular cPanel-managed file and the
// attacker's values take precedence anyway — we must NOT accept this
// as "managed" because doing so would let attackers suppress findings
// by inserting one of their own lines and then the real cPanel header.
// cpanelUserIniSignature is the exact string cPanel writes on the first
// comment line of a managed .user.ini. Case-sensitive: alternate
// capitalizations are rejected to avoid accepting forgeries.
const cpanelUserIniSignature = "cPanel-generated php ini directives"
// cpanelUserIniMaxLeadingBlanks caps how many blank lines may precede
// the signature. cPanel itself writes the signature as the very first
// line, so any tolerance here is purely for admins who may have
// round-tripped the file through a text editor that adds a trailing
// newline or through an FTP client that accumulates CRLFs. A small cap
// also closes a forgery route: without a bound, an attacker who wants
// to suppress severity on their injected values could prepend a large
// run of blank lines followed by the genuine cPanel header string, and
// a first-non-blank-line-only check would classify the file as
// managed despite the attacker owning all the content above the
// header.
const cpanelUserIniMaxLeadingBlanks = 5
// isCpanelManagedUserIni reports whether data begins with the cPanel
// MultiPHP-managed .user.ini header. Empty/whitespace-only input
// returns false. A small number (cpanelUserIniMaxLeadingBlanks) of
// leading blank lines is tolerated; beyond that, the file is not
// considered cPanel-managed even if the signature appears later.
func isCpanelManagedUserIni(data []byte) bool {
if len(data) == 0 {
return false
}
blanks := 0
start := 0
check := func(line []byte) (matched bool, terminate bool) {
line = bytes.TrimRight(line, "\r")
line = bytes.TrimSpace(line)
if len(line) == 0 {
blanks++
if blanks > cpanelUserIniMaxLeadingBlanks {
return false, true
}
return false, false
}
return bytes.Contains(line, []byte(cpanelUserIniSignature)), true
}
for i := 0; i < len(data); i++ {
if data[i] != '\n' {
continue
}
matched, terminate := check(data[start:i])
if terminate {
return matched
}
start = i + 1
}
if start < len(data) {
matched, _ := check(data[start:])
return matched
}
return false
}
package checks
import (
"strings"
)
type indirectFuncKind uint8
const (
indirectDecoder indirectFuncKind = 1 << iota
indirectEval
indirectShell
)
var indirectFuncKinds = map[string]indirectFuncKind{
"base64_decode": indirectDecoder,
"gzinflate": indirectDecoder,
"gzuncompress": indirectDecoder,
"gzdecode": indirectDecoder,
"bzdecompress": indirectDecoder,
"str_rot13": indirectDecoder,
"rawurldecode": indirectDecoder,
"eval": indirectEval,
"assert": indirectEval,
"create_function": indirectEval,
"system": indirectShell,
"passthru": indirectShell,
"exec": indirectShell,
"shell_exec": indirectShell,
"popen": indirectShell,
"proc_open": indirectShell,
"pcntl_exec": indirectShell,
}
type indirectAssignment struct {
kind indirectFuncKind
pos int
}
type indirectCall struct {
variable string
kind indirectFuncKind
pos int
lineStart int
lineEnd int
line string
}
// detectVarFuncDangerousAssignment returns true when content contains a
// `$var = "dangerous_name"` assignment followed by a corresponding
// `$var(` invocation that forms an execution sink. Decoder-only
// callbacks are not enough; direct base64_decode() calls are common in
// legitimate plugins, and indirect callbacks need the same restraint.
func detectVarFuncDangerousAssignment(content string) bool {
code := stripPHPCommentsFromCode(content)
assignments := findIndirectAssignments(code)
if len(assignments) == 0 {
return false
}
calls := findIndirectCalls(code, assignments)
if len(calls) == 0 {
return false
}
for _, call := range calls {
switch {
case call.kind&indirectShell != 0:
if lineContainsRequestVar(call.line) {
return true
}
case call.kind&indirectEval != 0:
if lineContainsRequestVar(call.line) ||
lineContainsDirectDecoderCall(call.line) ||
lineContainsIndirectCallKind(calls, call.lineStart, call.lineEnd, indirectDecoder) {
return true
}
case call.kind&indirectDecoder != 0:
if lineContainsDirectEvalCall(call.line) ||
lineContainsIndirectCallKind(calls, call.lineStart, call.lineEnd, indirectEval) {
return true
}
}
}
return false
}
func findIndirectAssignments(code string) map[string][]indirectAssignment {
assignments := map[string][]indirectAssignment{}
for i := 0; i < len(code); i++ {
// Literal examples must neither create nor overwrite a callable binding.
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label) - 1
continue
}
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
if code[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(code, i)
if !ok {
continue
}
j := skipPHPWhitespace(code, next)
if j >= len(code) || code[j] != '=' || (j+1 < len(code) && (code[j+1] == '=' || code[j+1] == '>')) {
i = next - 1
continue
}
assignment := indirectAssignment{pos: j + 1}
valueStart := skipPHPWhitespace(code, j+1)
if valueStart < len(code) && isPHPQuote(code[valueStart]) {
value, valueEnd, valueOK := readPHPFunctionString(code, valueStart)
assignment.pos = valueEnd
if valueOK {
assignment.kind = indirectFuncKinds[value]
}
}
assignments[variable] = append(assignments[variable], assignment)
i = next - 1
}
return assignments
}
func findIndirectCalls(code string, assignments map[string][]indirectAssignment) []indirectCall {
var calls []indirectCall
// Strip strings before slicing lines: a call's line may start inside a
// multiline literal whose opening quote is on an earlier line.
codeNoStrings := stripPHPStringsFromCode(code)
for i := 0; i < len(code); i++ {
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
i = phpHeredocEnd(code, bodyStart, label) - 1
continue
}
if isPHPQuote(code[i]) {
i = skipPHPString(code, i)
continue
}
if code[i] != '$' {
continue
}
variable, next, ok := readPHPVariableName(code, i)
if !ok {
continue
}
j := skipPHPWhitespace(code, next)
if j >= len(code) || code[j] != '(' {
i = next - 1
continue
}
kind, ok := indirectKindAt(assignments[variable], i)
if !ok {
i = next - 1
continue
}
lineStart, lineEnd := phpLineBounds(code, i)
calls = append(calls, indirectCall{
variable: variable,
kind: kind,
pos: i,
lineStart: lineStart,
lineEnd: lineEnd,
line: codeNoStrings[lineStart:lineEnd],
})
i = next - 1
}
return calls
}
func indirectKindAt(assignments []indirectAssignment, callPos int) (indirectFuncKind, bool) {
var last indirectFuncKind
found := false
for _, assignment := range assignments {
if assignment.pos > callPos {
break
}
last = assignment.kind
found = true
}
return last, found && last != 0
}
func lineContainsDirectDecoderCall(line string) bool {
codeLine := strings.ToLower(stripPHPStringsFromCode(line))
for name, kind := range indirectFuncKinds {
if kind&indirectDecoder != 0 && containsStandaloneFunc(codeLine, name+"(") {
return true
}
}
return false
}
func lineContainsDirectEvalCall(line string) bool {
codeLine := strings.ToLower(stripPHPStringsFromCode(line))
for name, kind := range indirectFuncKinds {
if kind&indirectEval != 0 && containsStandaloneFunc(codeLine, name+"(") {
return true
}
}
return false
}
func lineContainsRequestVar(line string) bool {
return containsRequestSuperglobal(stripPHPStringsFromCode(line))
}
func lineContainsIndirectCallKind(calls []indirectCall, lineStart, lineEnd int, kind indirectFuncKind) bool {
for _, call := range calls {
if call.lineStart == lineStart && call.lineEnd == lineEnd && call.kind&kind != 0 {
return true
}
}
return false
}
func stripPHPCommentsFromCode(code string) string {
var b strings.Builder
b.Grow(len(code))
for i := 0; i < len(code); i++ {
// Heredoc/nowdoc bodies are string literals, not code. Copy them
// verbatim so a '#', '//', or '/*' inside the body is not mistaken
// for a comment (which would corrupt the surrounding code).
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
end := phpHeredocEnd(code, bodyStart, label)
b.WriteString(code[i:end])
i = end - 1
continue
}
if isPHPQuote(code[i]) {
i = copyPHPString(&b, code, i)
continue
}
if code[i] == '/' && i+1 < len(code) && code[i+1] == '*' {
b.WriteString(" ")
i += 2
for i < len(code) {
if code[i] == '*' && i+1 < len(code) && code[i+1] == '/' {
b.WriteString(" ")
i++
break
}
writeCommentReplacementByte(&b, code[i])
i++
}
continue
}
if code[i] == '/' && i+1 < len(code) && code[i+1] == '/' {
b.WriteString(" ")
i += 2
for i < len(code) {
if code[i] == '\n' || code[i] == '\r' {
b.WriteByte(code[i])
break
}
b.WriteByte(' ')
i++
}
continue
}
// "#[" opens a PHP 8 attribute, which can precede a statement on the
// same line; only other "#" forms are comments.
if code[i] == '#' && isPHPLineCommentStart(code, i) {
b.WriteByte(' ')
i++
for i < len(code) {
if code[i] == '\n' || code[i] == '\r' {
b.WriteByte(code[i])
break
}
b.WriteByte(' ')
i++
}
continue
}
b.WriteByte(code[i])
}
return b.String()
}
func stripPHPStringsFromCode(code string) string {
var b strings.Builder
b.Grow(len(code))
for i := 0; i < len(code); i++ {
// Blank heredoc/nowdoc bodies (and their opener/closing label) the same
// way single/double-quoted strings are blanked, so their contents are
// not analysed as code and a quote inside the body cannot desync the
// scanner and swallow real code that follows the heredoc.
if label, bodyStart, ok := phpHeredocOpen(code, i); ok {
end := phpHeredocEnd(code, bodyStart, label)
for k := i; k < end; k++ {
if code[k] == '\n' || code[k] == '\r' {
b.WriteByte(code[k])
} else {
b.WriteByte(' ')
}
}
i = end - 1
continue
}
if isPHPQuote(code[i]) {
i = replacePHPString(&b, code, i)
continue
}
b.WriteByte(code[i])
}
return b.String()
}
// phpHeredocOpen reports whether code[i:] opens a heredoc or nowdoc. On success
// it returns the label and the byte index where the body begins (just past the
// opening line's newline). It recognises `<<<LABEL`, `<<<"LABEL"` (heredoc) and
// `<<<'LABEL'` (nowdoc), with optional spaces/tabs after `<<<`.
func phpHeredocOpen(code string, i int) (label string, bodyStart int, ok bool) {
if i+3 > len(code) || code[i] != '<' || code[i+1] != '<' || code[i+2] != '<' {
return "", 0, false
}
if i > 0 && code[i-1] == '<' {
return "", 0, false
}
j := i + 3
for j < len(code) && (code[j] == ' ' || code[j] == '\t') {
j++
}
var quote byte
if j < len(code) && (code[j] == '\'' || code[j] == '"') {
quote = code[j]
j++
}
if j >= len(code) || !isPHPIdentifierStart(code[j]) {
return "", 0, false
}
start := j
j++
for j < len(code) && isPHPIdentifierPart(code[j]) {
j++
}
label = code[start:j]
if quote != 0 {
if j >= len(code) || code[j] != quote {
return "", 0, false
}
j++
}
// The opening line ends at the next newline; only trailing whitespace and
// an optional CR may sit between the label and that newline.
for j < len(code) && (code[j] == ' ' || code[j] == '\t' || code[j] == '\r') {
j++
}
if j >= len(code) || code[j] != '\n' {
return "", 0, false
}
return label, j + 1, true
}
// phpHeredocEnd returns the byte index just past the closing label of a heredoc
// whose body begins at bodyStart. PHP 7.3+ permits the closing label to be
// indented; the label must appear at the start of a line (after optional
// spaces/tabs) and be followed by a non-identifier byte. An unterminated
// heredoc consumes the rest of the input.
func phpHeredocEnd(code string, bodyStart int, label string) int {
i := bodyStart
for i < len(code) {
lineEnd := i
for lineEnd < len(code) && code[lineEnd] != '\n' {
lineEnd++
}
k := i
for k < lineEnd && (code[k] == ' ' || code[k] == '\t') {
k++
}
if k+len(label) <= lineEnd && code[k:k+len(label)] == label {
after := k + len(label)
if after >= len(code) || !isPHPIdentifierPart(code[after]) {
return after
}
}
if lineEnd >= len(code) {
break
}
i = lineEnd + 1
}
return len(code)
}
func copyPHPString(b *strings.Builder, code string, start int) int {
quote := code[start]
b.WriteByte(code[start])
for i := start + 1; i < len(code); i++ {
b.WriteByte(code[i])
if code[i] == '\\' && i+1 < len(code) {
i++
b.WriteByte(code[i])
continue
}
if code[i] == quote {
return i
}
}
return len(code) - 1
}
func replacePHPString(b *strings.Builder, code string, start int) int {
quote := code[start]
b.WriteByte(' ')
for i := start + 1; i < len(code); i++ {
if code[i] == '\n' || code[i] == '\r' {
b.WriteByte(code[i])
} else {
b.WriteByte(' ')
}
if code[i] == '\\' && i+1 < len(code) {
i++
if code[i] == '\n' || code[i] == '\r' {
b.WriteByte(code[i])
} else {
b.WriteByte(' ')
}
continue
}
if code[i] == quote {
return i
}
}
return len(code) - 1
}
func writeCommentReplacementByte(b *strings.Builder, c byte) {
if c == '\n' || c == '\r' {
b.WriteByte(c)
return
}
b.WriteByte(' ')
}
func readPHPVariableName(code string, dollar int) (string, int, bool) {
if dollar+1 >= len(code) || !isPHPIdentifierStart(code[dollar+1]) {
return "", dollar + 1, false
}
i := dollar + 2
for i < len(code) && isPHPIdentifierPart(code[i]) {
i++
}
return code[dollar+1 : i], i, true
}
func readPHPFunctionString(code string, start int) (string, int, bool) {
quote := code[start]
var b strings.Builder
for i := start + 1; i < len(code); i++ {
if code[i] == '\\' && i+1 < len(code) {
if code[i+1] == quote || code[i+1] == '\\' {
i++
b.WriteByte(code[i])
continue
}
b.WriteByte(code[i])
continue
}
if code[i] == quote {
value, ok := normalizeIndirectFunctionName(b.String())
return value, i + 1, ok
}
b.WriteByte(code[i])
}
return "", len(code), false
}
func normalizeIndirectFunctionName(value string) (string, bool) {
value = strings.TrimLeft(value, "\\")
if !isPHPIdentifier(value) {
return "", false
}
return strings.ToLower(value), true
}
func skipPHPString(code string, start int) int {
quote := code[start]
for i := start + 1; i < len(code); i++ {
if code[i] == '\\' && i+1 < len(code) {
i++
continue
}
if code[i] == quote {
return i
}
}
return len(code) - 1
}
func skipPHPWhitespace(code string, start int) int {
for start < len(code) {
switch code[start] {
case ' ', '\t', '\n', '\r', '\f', '\v':
start++
default:
return start
}
}
return start
}
func phpLineBounds(code string, pos int) (int, int) {
start := pos
for start > 0 && code[start-1] != '\n' && code[start-1] != '\r' {
start--
}
end := pos
for end < len(code) && code[end] != '\n' && code[end] != '\r' {
end++
}
return start, end
}
func isPHPQuote(c byte) bool {
return c == '\'' || c == '"'
}
func isPHPIdentifier(value string) bool {
if value == "" || !isPHPIdentifierStart(value[0]) {
return false
}
for i := 1; i < len(value); i++ {
if !isPHPIdentifierPart(value[i]) {
return false
}
}
return true
}
func isPHPIdentifierStart(c byte) bool {
return c == '_' || (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z')
}
func isPHPIdentifierPart(c byte) bool {
return isPHPIdentifierStart(c) || (c >= '0' && c <= '9')
}
package checks
import (
"bytes"
"context"
"crypto/sha256"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"syscall"
"unicode"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/jstaint"
"github.com/pidginhost/csm/internal/phptaint"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/yara"
)
// VerifyResult reports whether a finding's underlying condition still holds.
//
// Checked is false when the finding's check type has no cheap, reliable
// single-target re-check (the caller should tell the operator to dismiss after
// manual review or run a full account scan). When Checked is true, Resolved
// reports whether the condition is gone and the finding can be cleared.
type VerifyResult struct {
Checked bool `json:"checked"`
Resolved bool `json:"resolved"`
// Demote marks a finding whose flagged content is gone but whose
// remediation cannot be proven, because the file changed since detection.
// It is never cleared -- an attacker must not retire a finding by editing
// the file -- but it stops ranking beside live threats.
Demote bool `json:"demote"`
Detail string `json:"detail"`
}
// VerifyInput carries everything a finding verifier may need. ContentSHA256 and
// DetectLogic are populated only for content findings emitted with a fingerprint.
// Context is optional; long-running verifiers use Background when it is nil.
type VerifyInput struct {
Check, Message, Details, Path string
ContentSHA256, DetectLogic string
Context context.Context
exposureVhosts *exposureVhostIndex
}
// presenceVerifiableChecks are findings whose remediation removes or
// quarantines a single flagged file, so the honest, cheap re-check is "is the
// flagged path still there?". Content-family checks are handled by
// reverifyContentFinding only when this package can re-run the same classifier
// that produced them; realtime PHP heuristic findings stay presence-based.
var presenceVerifiableChecks = []string{
"webshell", "webshell_realtime",
"webshell_content_realtime", "obfuscated_php_realtime",
"php_dropper_realtime",
"new_webshell_file", "new_suspicious_php",
"new_php_in_sensitive_dir", "new_php_in_uploads",
"php_in_sensitive_dir_realtime", "php_in_uploads_realtime",
"suspicious_file",
"nulled_plugin", "symlink_attack",
"vulnerable_timthumb",
"backdoor_binary", "new_executable_in_config",
"executable_in_config_realtime", "executable_in_tmp_realtime",
"cgi_backdoor_realtime", "cgi_suspicious_location_realtime",
"phishing_page", "phishing_directory", "phishing_php",
"phishing_kit_archive", "phishing_kit_realtime", "phishing_iframe",
"phishing_redirector", "phishing_credential_log", "phishing_realtime",
"credential_log_realtime",
}
// htaccessVerifiableChecks re-audit the .htaccess and resolve when no malicious
// directive remains (or the file is gone).
var htaccessVerifiableChecks = append(
[]string{"htaccess_injection", "htaccess_injection_realtime", "htaccess_handler_abuse"},
htaccessDetectorNames()...,
)
// findingVerifiers maps a finding's Check to a read-only re-check. A check not
// present here has no automated re-check -- either an event finding (a brute
// force, a past login: history cannot be re-evaluated by reading current state)
// or a condition we cannot yet cheaply and safely confirm. CanVerify reports
// membership so the Web UI shows the "Re-check" action only where it can act.
var findingVerifiers = buildFindingVerifiers()
var (
contentSignatureScanner = signatures.Global
contentYARAScanner = yara.Active
)
func buildFindingVerifiers() map[string]func(VerifyInput) VerifyResult {
m := map[string]func(VerifyInput) VerifyResult{}
register := func(fn func(VerifyInput) VerifyResult, names ...string) {
for _, n := range names {
m[n] = fn
}
}
register(func(in VerifyInput) VerifyResult { return verifyWriteBit(in.Path, 0002, "world-writable") },
"world_writable_php")
register(func(in VerifyInput) VerifyResult { return verifyWriteBit(in.Path, 0020, "group-writable") },
"group_writable_php")
register(func(in VerifyInput) VerifyResult {
return verifyPathAbsent(in.Path, effectiveFixRoots(fixQuarantineAllowedRoots, quarantineExtraRoots...))
},
presenceVerifiableChecks...)
register(reverifyContentFinding, contentReverifiableChecks...)
register(func(in VerifyInput) VerifyResult { return verifyHtaccessClean(in.Path) },
htaccessVerifiableChecks...)
register(func(in VerifyInput) VerifyResult { return verifyEximSpoolAbsent(in.Message) },
"email_phishing_content")
register(func(in VerifyInput) VerifyResult { return verifyCrontabClear(in.Path) },
"suspicious_crontab")
register(verifyExposedFile, exposedVerifiableChecks...)
register(func(in VerifyInput) VerifyResult { return verifyOutdatedPlugins(in.Details) },
"outdated_plugins")
register(func(in VerifyInput) VerifyResult { return verifyWPCoreIntegrity(in.Details) },
"wp_core_integrity")
register(func(in VerifyInput) VerifyResult { return verifyUID0Account(in.Message) },
"uid0_account")
register(func(in VerifyInput) VerifyResult { return verifySuidCleared(in.Path) },
"suid_binary")
register(func(in VerifyInput) VerifyResult { return verifyRPMIntegrity(in.Message) },
"rpm_integrity")
register(func(in VerifyInput) VerifyResult { return verifyDpkgIntegrity(in.Message) },
"dpkg_integrity")
register(func(in VerifyInput) VerifyResult { return verifyDBOptionsInjection(in.Message, in.Details) },
"db_options_injection")
register(func(in VerifyInput) VerifyResult { return verifyDBSiteurlHijack(in.Message, in.Details) },
"db_siteurl_hijack", "db_siteurl_invalid")
register(func(in VerifyInput) VerifyResult { return verifyDBPostInjection(in.Message, in.Details) },
"db_post_injection")
register(func(in VerifyInput) VerifyResult { return verifyDBSpamInjection(in.Message, in.Details) },
"db_spam_injection")
register(func(in VerifyInput) VerifyResult { return verifyDrupalSettingsInjection(in.Message, in.Details) },
"drupal_settings_injection")
register(func(in VerifyInput) VerifyResult { return verifyDrupalContentInjection(in.Message, in.Details) },
"drupal_content_injection")
register(func(in VerifyInput) VerifyResult { return verifyJoomlaExtensionsInjection(in.Message, in.Details) },
"joomla_extensions_injection")
register(func(in VerifyInput) VerifyResult { return verifyJoomlaContentInjection(in.Message, in.Details) },
"joomla_content_injection")
register(func(in VerifyInput) VerifyResult { return verifyMagentoSettingsInjection(in.Message, in.Details) },
"magento_settings_injection")
register(func(in VerifyInput) VerifyResult { return verifyMagentoContentInjection(in.Message, in.Details) },
"magento_content_injection")
register(func(in VerifyInput) VerifyResult { return verifyOpenCartSettingsInjection(in.Message, in.Details) },
"opencart_settings_injection")
register(func(in VerifyInput) VerifyResult { return verifyOpenCartContentInjection(in.Message, in.Details) },
"opencart_content_injection")
register(func(in VerifyInput) VerifyResult { return verifyDBObject(in.Message, in.Details, false) },
"db_unexpected_trigger", "db_unexpected_event",
"db_unexpected_procedure", "db_unexpected_function")
register(func(in VerifyInput) VerifyResult { return verifyDBObject(in.Message, in.Details, true) },
"db_malicious_trigger", "db_malicious_event",
"db_malicious_procedure", "db_malicious_function")
register(func(in VerifyInput) VerifyResult { return verifyDBMagicTokenUser(in.Message, in.Details) },
"db_magic_token_user")
register(func(in VerifyInput) VerifyResult { return verifyDBRogueAdmin(in.Message, in.Details) },
"db_rogue_admin")
register(func(in VerifyInput) VerifyResult { return verifyDBSuspiciousAdminEmail(in.Message, in.Details) },
"db_suspicious_admin_email")
register(func(in VerifyInput) VerifyResult { return verifyDrupalAdminInjection(in.Message, in.Details) },
"drupal_admin_injection")
register(func(in VerifyInput) VerifyResult { return verifyJoomlaAdminInjection(in.Message, in.Details) },
"joomla_admin_injection")
register(func(in VerifyInput) VerifyResult { return verifyMagentoAdminInjection(in.Message, in.Details) },
"magento_admin_injection")
register(func(in VerifyInput) VerifyResult { return verifyOpenCartAdminInjection(in.Message, in.Details) },
"opencart_admin_injection")
return m
}
// VerifyFinding re-evaluates a finding by check type + message/details/path.
// Preserved signature for CLI/legacy callers; carries no content fingerprint.
func VerifyFinding(checkType, message, details string, filePath ...string) VerifyResult {
return VerifyFindingInput(VerifyInput{
Check: checkType, Message: message, Details: details,
Path: selectFindingPath(message, filePath...),
})
}
// VerifyFindingInput re-evaluates a finding from a full VerifyInput. Callers
// that verify content findings must provide the stored detection fingerprint.
func VerifyFindingInput(in VerifyInput) VerifyResult {
if in.Path == "" {
in.Path = selectFindingPath(in.Message)
}
if fn, ok := findingVerifiers[in.Check]; ok {
return fn(in)
}
return VerifyResult{Checked: false, Detail: fmt.Sprintf("no automated re-check available for '%s'", in.Check)}
}
// reverifyContentFinding re-runs the content classifier that originally
// produced the finding. It resolves only when the file is gone OR the file's
// bytes are byte-for-byte identical to detection time AND the current
// classifier no longer flags them -- a superseded-heuristic false positive.
// A file modified since detection is never auto-cleared to prevent partial
// cleans or evasion edits from being mistaken for a fix.
func reverifyContentFinding(in VerifyInput) VerifyResult {
if in.Path == "" {
return VerifyResult{Checked: false, Detail: "could not extract file path from finding"}
}
clean, info, exists, err := readOnlyFixPath(in.Path, effectiveFixRoots(fixQuarantineAllowedRoots, quarantineExtraRoots...))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("file no longer present (removed or quarantined): %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: "path is not a regular file; not auto-verifiable"}
}
matched, label, currentHash, err := contentStillMatches(in.Check, clean, info)
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if matched {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("still flagged by current detection logic: %s", label)}
}
switch {
case in.ContentSHA256 != "" && currentHash != "" && currentHash == in.ContentSHA256:
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf(
"identical content (sha256 unchanged) is no longer flagged by current detection logic (%s) -- superseded-heuristic false positive",
ContentDetectionVersion())}
case in.ContentSHA256 != "":
return demotionForChangedContent(in.Check, clean, info, currentHash,
"file modified since detection (sha256 mismatch); not auto-cleared")
default:
return demotionForChangedContent(in.Check, clean, info, currentHash,
"current detection logic no longer flags this file, but no detection-time fingerprint exists to rule out modification")
}
}
// demotionEntropyCeiling rejects a replacement that reads as packed or encoded.
// The inert-source gate below is stricter, but checking entropy first keeps a
// large encoded comment blob ineligible too.
const demotionEntropyCeiling = 5.0
// changedContentDemotionEligible reports whether a changed-content finding may
// be considered for demotion. Every check whose condition this package can
// re-run qualifies: what makes a demotion safe is the inert-replacement gate,
// which reads the bytes, not the name of the detector that flagged them.
func changedContentDemotionEligible(check string) bool {
return IsContentReverifiable(check)
}
// demotionForChangedContent decides whether a finding whose file changed since
// detection, and which current detection no longer flags, may drop out of the
// live queue. It never clears the finding.
func demotionForChangedContent(check, path string, info os.FileInfo, classifiedHash, reason string) VerifyResult {
if !changedContentDemotionEligible(check) {
return VerifyResult{Checked: true, Resolved: false,
Detail: reason + "; changed content is not eligible for automatic demotion -- review manually"}
}
// The classifier and the demotion gates must inspect the same bytes. A
// bounded second snapshot avoids unbounded allocation, preserves inode and
// metadata identity, and its digest closes the gap between the two reads.
snap, err := readContentSnapshotForReverifyBounded(path, info, contentFingerprintMaxBytes)
if err != nil {
return VerifyResult{Checked: true, Resolved: false,
Detail: reason + "; replacement exceeds the demotion read limit or changed during verification -- review manually"}
}
if classifiedHash == "" || snap.sha256 != classifiedHash {
return VerifyResult{Checked: true, Resolved: false,
Detail: reason + "; replacement changed during verification -- review manually"}
}
content := string(snap.data)
if hasObfuscatedExecutionSignal(content) {
return VerifyResult{Checked: true, Resolved: false,
Detail: reason + "; replacement content still carries an obfuscated-execution signal -- review manually"}
}
if shannonEntropy(content) >= demotionEntropyCeiling {
return VerifyResult{Checked: true, Resolved: false,
Detail: reason + "; replacement content reads as packed or encoded -- review manually"}
}
if !isInertPHPReplacement(content) {
return VerifyResult{Checked: true, Resolved: false,
Detail: reason + "; replacement still contains active PHP or web content -- review manually"}
}
return VerifyResult{Checked: true, Resolved: false, Demote: true,
Detail: reason + "; replacement is an inert PHP stub -- confirm remediation"}
}
// isInertPHPReplacement deliberately proves a very small safe shape instead of
// trying to blacklist every way PHP can execute. Current detection not matching
// is already a precondition here; another deny-list would make any omitted
// include, callback, side-effect function, or inline script a severity-demotion
// bypass. Only an empty file or a comment-only, non-closing PHP stub is inert.
// Rejecting closing tags also rejects inline HTML or JavaScript that would stay
// live after PHP execution stops.
func isInertPHPReplacement(content string) bool {
// Preserve the byte after the opener even when it is trailing whitespace:
// trimming it could turn a non-tag into the valid "<?php" at EOF.
trimmed := strings.TrimLeftFunc(content, unicode.IsSpace)
if trimmed == "" {
return true
}
if len(trimmed) < len("<?php") || !strings.EqualFold(trimmed[:len("<?php")], "<?php") {
return false
}
if len(trimmed) > len("<?php") && !isPHPOpenTagSpace(trimmed[len("<?php")]) {
return false
}
if strings.Contains(trimmed, "?>") {
return false
}
code := stripPHPCommentsFromCode(phpCodeOnly(trimmed))
return strings.TrimSpace(code) == ""
}
// contentStillMatches re-runs the appropriate classifier for the check type.
// For files small enough to fingerprint, the returned hash and classifier result
// come from the same opened content snapshot.
func contentStillMatches(check, path string, info os.FileInfo) (bool, string, string, error) {
switch check {
case "signature_match_realtime":
s := contentSignatureScanner()
if s == nil || s.RuleCount() == 0 {
return false, "", "", fmt.Errorf("signature scanner unavailable")
}
// Bounded like the deep scan: the sweep runs unattended over every
// stored content finding, and a flagged file an attacker has since
// grown must leave the finding unresolved, not be read whole.
snap, err := readContentSnapshotForReverifyBounded(path, info, FullScanMaxFileBytes(config.Active()))
if err != nil {
return false, "", "", fmt.Errorf("cannot read file: %w", err)
}
if hits := s.ScanContent(snap.data, strings.ToLower(filepath.Ext(path))); len(hits) > 0 {
return true, fmt.Sprintf("%d signature match(es)", len(hits)), snap.sha256, nil
}
return false, "", snap.sha256, nil
case "yara_match_realtime", "yara_match_scheduled":
y := contentYARAScanner()
if y == nil || y.RuleCount() == 0 {
return false, "", "", fmt.Errorf("YARA scanner unavailable")
}
snap, err := readContentSnapshotForReverifyBounded(path, info, FullScanMaxFileBytes(config.Active()))
if err != nil {
return false, "", "", fmt.Errorf("cannot read file: %w", err)
}
// A payload too large for one IPC frame is retried by path, so a big
// file is re-checked rather than being stuck forever: the finding could
// otherwise never be confirmed or cleared.
hits, scannedSHA, err := yara.ScanContentOrPathChecked(y, path, snap.data, len(snap.data))
if err != nil {
// Fail closed: a scan that could not complete (worker down, a
// transport error) must not auto-clear a still-infected finding as
// a superseded false positive.
return false, "", "", fmt.Errorf("YARA scan error: %v", err)
}
if scannedSHA != "" {
snap.sha256 = scannedSHA
}
if len(hits) > 0 {
return true, fmt.Sprintf("%d YARA rule match(es)", len(hits)), snap.sha256, nil
}
return false, "", snap.sha256, nil
case "js_keylogger_dataflow":
if info.Size() > jstaint.MaxSourceBytes {
return false, "", "", fmt.Errorf("JS taint analysis did not complete: %s", jstaint.StatusOversize)
}
snap, err := readContentSnapshotForReverifyBounded(path, info, jstaint.MaxSourceBytes)
if err != nil {
if errors.Is(err, errContentSnapshotTooLarge) {
return false, "", "", fmt.Errorf("JS taint analysis did not complete: %s", jstaint.StatusOversize)
}
return false, "", "", fmt.Errorf("cannot read file: %v", err)
}
rctx, cancel := context.WithTimeout(context.Background(), jsTaintReverifyTimeout)
defer cancel()
report := runJSTaintAnalysis(rctx, snap.data)
switch report.Status {
case jstaint.StatusAnalyzed, jstaint.StatusNotCandidate:
if len(report.Results) > 0 {
return true, fmt.Sprintf("%d keystroke exfiltration flow(s)", report.TotalResults), snap.sha256, nil
}
return false, "", snap.sha256, nil
default:
// Fail closed: a coverage-gap status (oversize, parse error,
// resource limit, cancellation, panic) must not auto-clear a
// still-infected finding as a superseded false positive.
return false, "", "", fmt.Errorf("JS taint analysis did not complete: %s", report.Status)
}
case "php_remote_taint":
if info.Size() > phptaint.MaxSourceBytes {
return false, "", "", fmt.Errorf("PHP taint analysis did not complete: %s", phptaint.StatusOversize)
}
snap, err := readContentSnapshotForReverifyBounded(path, info, phptaint.MaxSourceBytes)
if err != nil {
if errors.Is(err, errContentSnapshotTooLarge) {
return false, "", "", fmt.Errorf("PHP taint analysis did not complete: %s", phptaint.StatusOversize)
}
return false, "", "", fmt.Errorf("cannot read file: %v", err)
}
rctx, cancel := context.WithTimeout(context.Background(), phpTaintDeepPerFileTimeout)
defer cancel()
report := runPHPTaintAnalysis(rctx, snap.data)
switch report.Status {
case phptaint.StatusAnalyzed, phptaint.StatusNotCandidate:
if len(report.Results) > 0 {
return true, fmt.Sprintf("%d remote-source code execution flow(s)", report.TotalResults), snap.sha256, nil
}
return false, "", snap.sha256, nil
default:
// The isolated worker is mandatory here too. Any incomplete status
// leaves the finding unresolved instead of falling through to the
// unrelated PHP heuristic classifier and clearing it.
return false, "", "", fmt.Errorf("PHP taint analysis did not complete: %s", report.Status)
}
default: // PHP heuristic content family
res, currentHash, err := analyzePHPContentForReverify(path, info)
if err != nil {
return false, "", "", fmt.Errorf("cannot read file: %v", err)
}
if res.severity >= 0 {
return true, strings.Join(res.indicators, ", "), currentHash, nil
}
return false, "", currentHash, nil
}
}
func analyzePHPContentForReverify(path string, expected os.FileInfo) (phpAnalysisResult, string, error) {
if expected != nil && expected.Size() <= contentFingerprintMaxBytes {
snap, err := readContentSnapshotForReverify(path, expected)
if err != nil {
return phpAnalysisResult{}, "", err
}
res := analyzePHPContentReaderAt(path, bytes.NewReader(snap.data), int64(len(snap.data)))
if !res.readOK {
return phpAnalysisResult{}, "", fmt.Errorf("content read failed")
}
return res, snap.sha256, nil
}
f, info, err := openReadOnlyPreservingIdentity(path, expected)
if err != nil {
return phpAnalysisResult{}, "", err
}
defer func() { _ = f.Close() }()
res := analyzePHPContentReaderAt(path, f, info.Size())
if !res.readOK {
return phpAnalysisResult{}, "", fmt.Errorf("content read failed")
}
return res, "", nil
}
type reverifyContentSnapshot struct {
data []byte
sha256 string
}
var errContentSnapshotTooLarge = errors.New("content snapshot exceeds read limit")
func readContentSnapshotForReverify(path string, expected os.FileInfo) (reverifyContentSnapshot, error) {
return readContentSnapshotForReverifyBounded(path, expected, 0)
}
func readContentSnapshotForReverifyBounded(path string, expected os.FileInfo, maxBytes int64) (reverifyContentSnapshot, error) {
f, info, err := openReadOnlyPreservingIdentity(path, expected)
if err != nil {
return reverifyContentSnapshot{}, err
}
defer func() { _ = f.Close() }()
if maxBytes > 0 && info.Size() > maxBytes {
return reverifyContentSnapshot{}, errContentSnapshotTooLarge
}
var reader io.Reader = f
if maxBytes > 0 {
reader = io.LimitReader(f, maxBytes+1)
}
data, err := io.ReadAll(reader)
if err != nil {
return reverifyContentSnapshot{}, err
}
after, err := f.Stat()
if err != nil {
return reverifyContentSnapshot{}, err
}
if !sameFileIdentity(after, info) || !sameCleanContentShape(after, info) {
return reverifyContentSnapshot{}, fmt.Errorf("file changed during verification")
}
if int64(len(data)) != info.Size() {
return reverifyContentSnapshot{}, fmt.Errorf("file changed during verification")
}
if maxBytes > 0 && int64(len(data)) > maxBytes {
return reverifyContentSnapshot{}, errContentSnapshotTooLarge
}
var digest string
if info.Size() <= contentFingerprintMaxBytes {
sum := sha256.Sum256(data)
digest = fmt.Sprintf("%x", sum)
}
return reverifyContentSnapshot{data: data, sha256: digest}, nil
}
func openReadOnlyPreservingIdentity(path string, expected os.FileInfo) (*os.File, os.FileInfo, error) {
// #nosec G304 -- path was validated by readOnlyFixPath against remediation
// roots; O_NOFOLLOW plus sameFileIdentity fails closed on inode swap.
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return nil, nil, err
}
info, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, nil, err
}
if !sameFileIdentity(info, expected) || !sameCleanContentShape(info, expected) {
_ = f.Close()
return nil, nil, fmt.Errorf("file changed during verification")
}
return f, info, nil
}
// CanVerify reports whether VerifyFinding has an automated re-check for the
// given check type. The Web UI gates the per-finding "Re-check" action on this
// so it never shows a button that could only report "not auto-verifiable".
func CanVerify(checkType string) bool {
_, ok := findingVerifiers[checkType]
return ok
}
func verifyWriteBit(path string, bit os.FileMode, label string) VerifyResult {
if path == "" {
return VerifyResult{Checked: false, Detail: "could not extract file path from finding"}
}
clean, info, exists, err := readOnlyFixPath(path, effectiveFixRoots(fixPermissionsAllowedRoots))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("file no longer exists: %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: "path is not a regular file; not auto-verifiable"}
}
if info.Mode().Perm()&bit == 0 {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("file is no longer %s (mode %o)", label, info.Mode().Perm())}
}
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("file is still %s (mode %o)", label, info.Mode().Perm())}
}
func verifyPathAbsent(path string, roots []string) VerifyResult {
if path == "" {
return VerifyResult{Checked: false, Detail: "could not extract file path from finding"}
}
clean, exists, err := readOnlyPathPresence(path, roots)
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("file no longer present (removed or quarantined): %s", clean)}
}
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("file is still present: %s", clean)}
}
func verifyHtaccessClean(path string) VerifyResult {
if path == "" {
return VerifyResult{Checked: false, Detail: "could not extract file path from finding"}
}
clean, info, exists, err := readOnlyFixPath(path, effectiveFixRoots(fixHtaccessAllowedRoots))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if filepath.Base(clean) != ".htaccess" {
return VerifyResult{Checked: false, Detail: "not a .htaccess file; not auto-verifiable"}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf(".htaccess no longer exists: %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: ".htaccess path is not a regular file; not auto-verifiable"}
}
content, err := readFilePreservingIdentity(clean, info)
if err != nil {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("cannot read .htaccess: %v", err)}
}
findings, _ := AuditHtaccessContent(clean, content)
if len(findings) == 0 {
return VerifyResult{Checked: true, Resolved: true, Detail: "no malicious directives remain in .htaccess"}
}
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("%d malicious directive(s) still present in .htaccess", len(findings))}
}
func verifyEximSpoolAbsent(message string) VerifyResult {
msgID := extractEximMsgID(message)
if msgID == "" {
return VerifyResult{Checked: false, Detail: "could not extract Exim message ID from finding"}
}
if !eximMsgIDRegex.MatchString(msgID) {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("invalid Exim message ID format: %s", msgID)}
}
for _, dir := range eximSpoolDirs {
if _, err := osFS.Lstat(filepath.Join(dir, msgID+"-H")); err == nil {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("spool message %s is still queued", msgID)}
} else if !os.IsNotExist(err) {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("cannot stat spool message %s: %v", msgID, err)}
}
}
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("spool message %s no longer present (delivered or removed)", msgID)}
}
func verifyCrontabClear(path string) VerifyResult {
if path == "" {
return VerifyResult{Checked: false, Detail: "could not extract crontab path from finding"}
}
clean, info, exists, err := readOnlyFixPath(path, fixCrontabAllowedRoots)
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("crontab no longer exists: %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: "crontab path is not a regular file; not auto-verifiable"}
}
if info.Size() == 0 {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("crontab is empty: %s", clean)}
}
// A non-empty crontab may still be legitimate; re-scanning its content is
// out of scope for a single-finding re-check, so report it as still
// present rather than guessing.
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("crontab is still present and non-empty: %s", clean)}
}
func readOnlyFixPath(path string, allowedRoots []string) (string, os.FileInfo, bool, error) {
clean, err := sanitizeFixPath(path, allowedRoots)
if err != nil {
return "", nil, false, err
}
root := matchingAllowedRoot(clean, allowedRoots)
if root == "" {
return "", nil, false, fmt.Errorf("file path is outside the allowed remediation roots: %s", clean)
}
current := root
info, err := osFS.Lstat(current)
if err != nil {
if os.IsNotExist(err) {
return clean, nil, false, nil
}
return "", nil, false, fmt.Errorf("cannot stat: %v", err)
}
if info.Mode()&os.ModeSymlink != 0 {
return "", nil, false, fmt.Errorf("symlinked paths are not eligible for automated verification: %s", current)
}
if current == clean {
return clean, info, true, nil
}
rel, err := filepath.Rel(current, clean)
if err != nil {
return "", nil, false, fmt.Errorf("cannot resolve path under allowed root: %v", err)
}
parts := strings.Split(rel, string(filepath.Separator))
for i, part := range parts {
if part == "" || part == "." {
continue
}
if !info.IsDir() {
return "", nil, false, fmt.Errorf("path ancestor is not a directory; not auto-verifiable: %s", current)
}
current = filepath.Join(current, part)
info, err = osFS.Lstat(current)
if err != nil {
if os.IsNotExist(err) {
return clean, nil, false, nil
}
return "", nil, false, fmt.Errorf("cannot stat: %v", err)
}
if info.Mode()&os.ModeSymlink != 0 {
return "", nil, false, fmt.Errorf("symlinked paths are not eligible for automated verification: %s", current)
}
if i < len(parts)-1 && !info.IsDir() {
return "", nil, false, fmt.Errorf("path ancestor is not a directory; not auto-verifiable: %s", current)
}
}
return clean, info, true, nil
}
func readOnlyPathPresence(path string, allowedRoots []string) (string, bool, error) {
clean, err := sanitizeFixPath(path, allowedRoots)
if err != nil {
return "", false, err
}
root := matchingAllowedRoot(clean, allowedRoots)
if root == "" {
return "", false, fmt.Errorf("file path is outside the allowed remediation roots: %s", clean)
}
current := root
info, err := osFS.Lstat(current)
if err != nil {
if os.IsNotExist(err) {
return clean, false, nil
}
return "", false, fmt.Errorf("cannot stat: %v", err)
}
if info.Mode()&os.ModeSymlink != 0 {
return "", false, fmt.Errorf("symlinked paths are not eligible for automated verification: %s", current)
}
if current == clean {
return clean, true, nil
}
rel, err := filepath.Rel(current, clean)
if err != nil {
return "", false, fmt.Errorf("cannot resolve path under allowed root: %v", err)
}
parts := strings.Split(rel, string(filepath.Separator))
for i, part := range parts {
if part == "" || part == "." {
continue
}
if !info.IsDir() {
return "", false, fmt.Errorf("path ancestor is not a directory; not auto-verifiable: %s", current)
}
current = filepath.Join(current, part)
info, err = osFS.Lstat(current)
if err != nil {
if os.IsNotExist(err) {
return clean, false, nil
}
return "", false, fmt.Errorf("cannot stat: %v", err)
}
if i < len(parts)-1 {
if info.Mode()&os.ModeSymlink != 0 {
return "", false, fmt.Errorf("symlinked paths are not eligible for automated verification: %s", current)
}
if !info.IsDir() {
return "", false, fmt.Errorf("path ancestor is not a directory; not auto-verifiable: %s", current)
}
}
}
return clean, true, nil
}
func matchingAllowedRoot(path string, allowedRoots []string) string {
var best string
for _, root := range allowedRoots {
cleanRoot := filepath.Clean(strings.TrimSpace(root))
if isPathWithinOrEqual(path, cleanRoot) && len(cleanRoot) > len(best) {
best = cleanRoot
}
}
return best
}
package checks
import (
"context"
"fmt"
"regexp"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/mysqlclient"
)
// Per-finding Re-check for the database-content family.
//
// These verifiers re-evaluate a single database finding against the live
// database and clear it only on confirmed evidence -- the row is gone, or its
// current value no longer matches the detector that raised it. The cardinal
// rule is the same as the filesystem re-checks: NEVER false-resolve. Clearing a
// db_* finding auto-dismisses a live database compromise, so any ambiguity
// (connection failure, query error, a site we can no longer locate, an
// unparseable finding) returns Checked:false and leaves the finding in place
// for a full account scan.
//
// Reliability over the wp-config password: re-checks query as root against the
// finding's schema (mysqlclient.RootQuerySchema) rather than the per-account
// credentials in the CMS config file, which drift on cPanel password rotations.
// The config file is read only to re-discover the schema name and table prefix.
// dbVerifyTimeout bounds a single synchronous database re-check.
var dbVerifyTimeout = 30 * time.Second
// validAccountName guards an account name parsed out of finding text before it
// is interpolated into a /home/<account>/... glob.
var validAccountName = regexp.MustCompile(`^[A-Za-z][A-Za-z0-9_-]{0,31}$`)
// messageAccountRe pulls account tokens out of WordPress finding messages.
var messageAccountRe = regexp.MustCompile(`(^|[,(])\s*account:\s*([A-Za-z][A-Za-z0-9_-]{0,31})([,)]|$)`)
// detailField returns the value of a "Key: value" line in a finding's Details
// block, or "" when the key is absent. The key match is case-sensitive and the
// value is whitespace-trimmed.
func detailField(details, key string) string {
value, ok := detailFieldPresent(details, key)
if !ok {
return ""
}
return value
}
func detailFieldPresent(details, key string) (string, bool) {
want := key + ":"
for _, line := range strings.Split(details, "\n") {
line = strings.TrimSpace(line)
if strings.HasPrefix(line, want) {
return strings.TrimSpace(line[len(want):]), true
}
}
return "", false
}
// dbFindingAccount resolves the cPanel account a database finding belongs to.
// CMS, DB-object, and admin findings carry it as an "Account:" detail line; the
// WordPress content findings carry it only in the message "(account: <user>)".
func dbFindingAccount(message, details string) string {
if a, ok := detailFieldPresent(details, "Account"); ok {
return validFindingAccount(a)
}
matches := messageAccountRe.FindAllStringSubmatch(message, -1)
if len(matches) > 0 {
return matches[len(matches)-1][2]
}
return ""
}
func validFindingAccount(account string) string {
account = strings.TrimSpace(account)
if !validAccountName.MatchString(account) {
return ""
}
return account
}
// runDBVerifyQueryRoot runs a root-credential query against an explicit schema
// and returns the rows plus the error. Unlike runMySQLQueryRoot (which collapses
// any error into a nil slice), this preserves the error so a re-check can tell
// "query ran, zero rows" (the row is gone/clean -> safe to resolve) apart from
// "query failed" (DB down, connection refused -> must NOT resolve).
func runDBVerifyQueryRoot(schema, query string, args ...any) ([]string, error) {
if strings.TrimSpace(schema) == "" {
return nil, fmt.Errorf("empty schema")
}
ctx, cancel := context.WithTimeout(context.Background(), dbVerifyTimeout)
defer cancel()
return mysqlclient.RootQuerySchema(ctx, schema, query, args...)
}
// dbVerifyQueryError is the standard Checked:false result for a failed re-check
// query. The raw error is intentionally not surfaced to the operator (it can
// carry DSN/host internals); the generic guidance is enough to act on.
func dbVerifyQueryError() VerifyResult {
return VerifyResult{Checked: false, Detail: "could not query the database (try again, or run an account scan)"}
}
// dbVerifyNotLocatable is the standard Checked:false result when the re-check
// cannot re-discover the install (site removed, DB renamed, config unreadable).
func dbVerifyNotLocatable(kind string) VerifyResult {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("could not locate the %s to re-check (run an account scan)", kind)}
}
// findWPVerifyPrefixes re-discovers the WordPress install for account whose
// database matches dbName, returning the table prefixes a verifier should query.
// New findings carry an exact Table prefix detail. Older findings do not, so the
// fallback enumerates every prefix in that WordPress install, including active
// multisite secondary blogs. ok=false means the caller refuses to guess.
func findWPVerifyPrefixes(account, dbName, details string) (prefixes []string, ok bool) {
if !validAccountName.MatchString(account) || dbName == "" {
return nil, false
}
// Shared discovery (wpinstalls.go): a re-check that cannot re-locate the
// install it flagged leaves the finding unresolvable for the life of the
// site, so it must see exactly what the detector saw.
patterns := wpInstallConfigPaths(wpInstallsForAccount(context.Background(), "db_content", account))
targetPrefix := detailField(details, "Table prefix")
if targetPrefix != "" && !validTablePrefix.MatchString(targetPrefix) {
return nil, false
}
seen := map[string]bool{}
for _, path := range patterns {
creds := parseWPConfig(path)
if creds.dbName != dbName {
continue
}
p, ok := resolveTablePrefix(creds)
if !ok {
return nil, false
}
if targetPrefix != "" {
if wpVerifyPrefixMatchesBase(targetPrefix, p) {
return []string{targetPrefix}, true
}
continue
}
addPrefix(&prefixes, seen, p)
if creds.multisite {
multisitePrefixes, err := wpVerifyMultisitePrefixes(dbName, p)
if err != nil {
return nil, false
}
for _, prefix := range multisitePrefixes {
addPrefix(&prefixes, seen, prefix)
}
}
}
if len(prefixes) == 0 {
return nil, false
}
return prefixes, true
}
func addPrefix(prefixes *[]string, seen map[string]bool, prefix string) {
if seen[prefix] {
return
}
seen[prefix] = true
*prefixes = append(*prefixes, prefix)
}
func wpVerifyPrefixMatchesBase(prefix, base string) bool {
if prefix == base {
return true
}
if !strings.HasPrefix(prefix, base) || !strings.HasSuffix(prefix, "_") {
return false
}
blogID := strings.TrimSuffix(strings.TrimPrefix(prefix, base), "_")
return isAllDigits(blogID)
}
func wpVerifyMultisitePrefixes(dbName, basePrefix string) ([]string, error) {
rows, err := runDBVerifyQueryRoot(dbName, fmt.Sprintf(
"SELECT blog_id FROM `%sblogs` WHERE archived = 0 AND deleted = 0 AND spam = 0 AND blog_id != 1",
basePrefix,
))
if err != nil {
return nil, err
}
var prefixes []string
for _, row := range rows {
blogID := strings.TrimSpace(row)
if blogID == "" || blogID == "1" || !isAllDigits(blogID) {
continue
}
prefixes = append(prefixes, fmt.Sprintf("%s%s_", basePrefix, blogID))
}
return prefixes, nil
}
// dbInjectionCoreOptions are the wp_options names that must never carry a
// <script> tag (mirrors the Path-2 core-option check in checkWPOptions).
var dbInjectionCoreOptions = map[string]bool{
"siteurl": true, "home": true, "blogname": true,
"blogdescription": true, "admin_email": true,
}
// verifyDBOptionsInjection re-reads the flagged wp_options row and resolves the
// finding when the option is gone or no longer carries an injected external
// script (mirrors checkWPOptions' two detection paths).
func verifyDBOptionsInjection(message, details string) VerifyResult {
dbName := detailField(details, "Database")
optName := detailField(details, "Option")
if optName == "" {
return VerifyResult{Checked: false, Detail: "could not parse the option name from the finding"}
}
prefixes, ok := findWPVerifyPrefixes(dbFindingAccount(message, details), dbName, details)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
present := false
for _, prefix := range prefixes {
rows, err := runDBVerifyQueryRoot(dbName,
fmt.Sprintf("SELECT option_value FROM `%soptions` WHERE option_name = ?", prefix), optName)
if err != nil {
return dbVerifyQueryError()
}
for _, row := range rows {
present = true
if optionValueStillMalicious(optName, row) {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("option %q still contains injected content", optName)}
}
}
}
if !present {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("option %q is no longer present", optName)}
}
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("option %q no longer contains injected content", optName)}
}
func optionValueStillMalicious(optName, value string) bool {
if extractMaliciousScriptURL(value) != "" {
return true
}
if dbInjectionCoreOptions[optName] && strings.Contains(strings.ToLower(value), "<script") {
return true
}
return false
}
// verifyDBSiteurlHijack re-reads the flagged siteurl/home option and resolves
// when the option is gone or no longer matches either site-address check in
// checkWPOptions.
func verifyDBSiteurlHijack(message, details string) VerifyResult {
dbName := detailField(details, "Database")
optName := siteurlOptionFromDetails(details)
if optName == "" {
return VerifyResult{Checked: false, Detail: "could not parse the option name from the finding"}
}
prefixes, ok := findWPVerifyPrefixes(dbFindingAccount(message, details), dbName, details)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
present := false
for _, prefix := range prefixes {
rows, err := runDBVerifyQueryRoot(dbName,
fmt.Sprintf("SELECT option_value FROM `%soptions` WHERE option_name = ?", prefix), optName)
if err != nil {
return dbVerifyQueryError()
}
for _, row := range rows {
present = true
if siteurlValueStillMalicious(row) {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("option %q is still a poisoned site address", optName)}
}
}
}
if !present {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("option %q is no longer present", optName)}
}
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("option %q is no longer a poisoned site address", optName)}
}
func siteurlValueStillMalicious(value string) bool {
lower := strings.ToLower(value)
if strings.Contains(lower, "eval(") || strings.Contains(lower, "<script") {
return true
}
_, poisoned := siteURLPoisonReason(value)
return poisoned
}
// siteurlOptionFromDetails extracts the option name from a site-address finding
// detail block whose body line is "<optName> = <value>". The detector only ever
// emits siteurl or home; any other shape returns "".
func siteurlOptionFromDetails(details string) string {
for _, line := range strings.Split(details, "\n") {
line = strings.TrimSpace(line)
for _, opt := range []string{"siteurl", "home"} {
if strings.HasPrefix(line, opt+" =") {
return opt
}
}
}
return ""
}
// verifyDBPostInjection searches the published posts for the injected pattern
// again, the way the detector found it, and resolves the finding only when no
// post still matches (mirrors checkWPPosts' post-status, post-type, and
// external-script filters). The finding lists at most five example post IDs;
// re-reading only those resolved a finding whose injection sat in every other
// post once the examples were cleaned.
func verifyDBPostInjection(message, details string) VerifyResult {
dbName := detailField(details, "Database")
pattern := detailField(details, "Pattern")
if pattern == "" {
return VerifyResult{Checked: false, Detail: "could not parse the injected pattern from the finding"}
}
requiresExternalScript, ok := lookupMalwarePattern(pattern)
if !ok {
return VerifyResult{Checked: false, Detail: "finding pattern is not auto-verifiable"}
}
prefixes, ok := findWPVerifyPrefixes(dbFindingAccount(message, details), dbName, details)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
like := likeContains(pattern)
for _, prefix := range prefixes {
var afterID uint64
for {
rows, err := runDBVerifyQueryRoot(dbName,
fmt.Sprintf("SELECT ID, post_content, post_content_filtered FROM `%sposts` WHERE ID > ? AND post_status='publish' AND post_type NOT IN (%s) AND (post_content LIKE ? OR post_content_filtered LIKE ?) ORDER BY ID LIMIT %d",
prefix, nonScannablePostTypesSQLList(), dbVerifyPostSearchLimit),
afterID, like, like)
if err != nil {
return dbVerifyQueryError()
}
for _, row := range rows {
parts := strings.SplitN(row, "\t", 3)
if len(parts) < 2 {
return VerifyResult{Checked: false, Detail: "database returned an unexpected post row; finding was not resolved"}
}
id, err := strconv.ParseUint(strings.TrimSpace(parts[0]), 10, 64)
if err != nil || id <= afterID {
return VerifyResult{Checked: false, Detail: "database returned an invalid post ID; finding was not resolved"}
}
afterID = id
content := parts[1]
if len(parts) == 3 {
content += "\n" + parts[2]
}
if postContentMatchesPattern(pattern, requiresExternalScript, content) {
return VerifyResult{Checked: true, Resolved: false, Detail: "an affected post still contains the injected pattern"}
}
}
if len(rows) < dbVerifyPostSearchLimit {
break
}
}
}
return VerifyResult{Checked: true, Resolved: true, Detail: "no affected post still contains the injected pattern"}
}
func postContentMatchesPattern(pattern string, requiresExternalScript bool, content string) bool {
if !strings.Contains(strings.ToLower(content), strings.ToLower(pattern)) {
return false
}
if requiresExternalScript {
return hasMaliciousExternalScriptInPost(content)
}
return true
}
// lookupMalwarePattern reports whether pattern is a known dbMalwarePatterns
// entry and, if so, whether it requires the external-script post-filter.
func lookupMalwarePattern(pattern string) (requiresExternalScript bool, ok bool) {
for _, mp := range dbMalwarePatterns {
if mp.pattern == pattern {
return mp.requiresExternalScript, true
}
}
return false, false
}
// verifyDBSpamInjection re-runs the cloaked-spam scan for the finding's keyword
// against the site's published posts and resolves when no post still matches
// (mirrors checkWPPosts' three-layer spam filter).
func verifyDBSpamInjection(message, details string) VerifyResult {
dbName := detailField(details, "Database")
keyword := spamKeywordFromMessage(message)
if keyword == "" {
return VerifyResult{Checked: false, Detail: "could not parse the spam keyword from the finding"}
}
sp, ok := lookupSpamPattern(keyword)
if !ok {
return VerifyResult{Checked: false, Detail: "spam keyword is not auto-verifiable"}
}
prefixes, ok := findWPVerifyPrefixes(dbFindingAccount(message, details), dbName, details)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
var contents []string
for _, prefix := range prefixes {
rows, err := runDBVerifyQueryRoot(dbName,
fmt.Sprintf("SELECT ID, post_content FROM `%sposts` WHERE post_status='publish' AND post_type NOT IN (%s) AND post_content LIKE ? LIMIT 200",
prefix, nonScannablePostTypesSQLList()), sp.likeFragment)
if err != nil {
return dbVerifyQueryError()
}
for _, row := range rows {
parts := strings.SplitN(row, "\t", 2)
if len(parts) >= 2 {
contents = append(contents, parts[1])
}
}
}
if countCloakedSpamMatches(sp, contents) == 0 {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("no published post still contains cloaked spam keyword %q", keyword)}
}
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("published posts still contain cloaked spam keyword %q", keyword)}
}
var spamKeywordRe = regexp.MustCompile(`cloaked spam keyword '([^']+)'`)
func spamKeywordFromMessage(message string) string {
m := spamKeywordRe.FindStringSubmatch(message)
if m == nil {
return ""
}
return m[1]
}
func lookupSpamPattern(keyword string) (dbSpamPattern, bool) {
for _, sp := range dbSpamPatterns {
if sp.keyword == keyword {
return sp, true
}
}
return dbSpamPattern{}, false
}
// parsePostIDList parses a "1, 2, 3" comma list into a deduplicated slice of
// numeric IDs. Non-numeric tokens are dropped so a malformed finding can never
// inject into the IN clause.
func parsePostIDList(s string) []string {
if s == "" {
return nil
}
seen := map[string]bool{}
var out []string
for _, tok := range strings.Split(s, ",") {
tok = strings.TrimSpace(tok)
if tok == "" || !isAllDigits(tok) || seen[tok] {
continue
}
seen[tok] = true
out = append(out, tok)
}
return out
}
// dbVerifyPostSearchLimit bounds the rows a post re-check reads back; one
// surviving match is enough to keep the finding open.
const dbVerifyPostSearchLimit = 50
// likeContains builds a bound LIKE argument that matches rows containing s
// literally: LIKE wildcards and the escape character inside s are escaped so
// a pattern such as "base64_decode" cannot match "base64Xdecode".
func likeContains(s string) string {
r := strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`)
return "%" + r.Replace(s) + "%"
}
package checks
import (
"context"
"fmt"
"strings"
)
// Per-finding Re-check for the administrator-account database findings: the
// WordPress rogue-admin and disposable-email-admin findings, and the four
// per-CMS administrator findings (Drupal, Joomla, Magento, OpenCart). Each
// resolves only when the specific flagged account is no longer present (or, for
// the WordPress rogue admin, no longer holds administrator capability). A
// surviving account -- which for the per-CMS Warning findings includes the
// legitimate site administrator -- keeps the finding active for operator review.
// Any query failure or undiscoverable install returns Checked:false.
// findWPAdminVerifyPrefixes re-derives the base table prefixes for WordPress
// installs whose database matches dbName. The users/usermeta tables are
// network-wide in multisite, so this returns only each install's base prefix,
// never secondary blog prefixes.
func findWPAdminVerifyPrefixes(account, dbName, details string) ([]string, bool) {
if !validAccountName.MatchString(account) || dbName == "" {
return nil, false
}
targetPrefix := detailField(details, "Table prefix")
if targetPrefix != "" && !validTablePrefix.MatchString(targetPrefix) {
return nil, false
}
var prefixes []string
seen := map[string]bool{}
// Shared discovery (wpinstalls.go): an administrator finding raised on a
// nested or panel-mapped install must be re-checkable there too.
patterns := wpInstallConfigPaths(wpInstallsForAccount(context.Background(), "db_content", account))
for _, path := range patterns {
creds := parseWPConfig(path)
if creds.dbName != dbName {
continue
}
prefix, ok := resolveTablePrefix(creds)
if !ok {
return nil, false
}
if targetPrefix != "" {
if targetPrefix == prefix {
return []string{prefix}, true
}
continue
}
addPrefix(&prefixes, seen, prefix)
}
if len(prefixes) == 0 {
return nil, false
}
return prefixes, true
}
// wpAdminQueryTables backtick-quotes the users and usermeta table names for the
// given (already validTablePrefix-validated) prefix.
func wpAdminQueryTables(prefix string) (users, usermeta string, ok bool) {
u, err := QuoteIdent(prefix + "users")
if err != nil {
return "", "", false
}
m, err := QuoteIdent(prefix + "usermeta")
if err != nil {
return "", "", false
}
return u, m, true
}
// verifyDBRogueAdmin resolves when the flagged WordPress user no longer holds
// administrator capability (deleted or demoted). The detector's "created in the
// last 7 days" heuristic is intentionally NOT re-applied -- the account is older
// now; what matters for the re-check is whether that specific account is still
// an administrator.
func verifyDBRogueAdmin(message, details string) VerifyResult {
dbName := detailField(details, "Database")
userID := detailField(details, "User ID")
if userID == "" || !isAllDigits(userID) {
return VerifyResult{Checked: false, Detail: "could not parse the user ID from the finding"}
}
prefixes, ok := findWPAdminVerifyPrefixes(dbFindingAccount(message, details), dbName, details)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
for _, prefix := range prefixes {
users, usermeta, ok := wpAdminQueryTables(prefix)
if !ok {
return VerifyResult{Checked: false, Detail: "could not validate the WordPress tables from the finding"}
}
rows, err := runDBVerifyQueryRoot(dbName, fmt.Sprintf(
"SELECT u.ID FROM %s u JOIN %s m ON u.ID = m.user_id WHERE m.meta_key = ? AND m.meta_value LIKE '%%administrator%%' AND u.ID = ?",
users, usermeta), prefix+"capabilities", userID)
if err != nil {
return dbVerifyQueryError()
}
if len(rows) > 0 {
return VerifyResult{Checked: true, Resolved: false, Detail: "the flagged account is still a WordPress administrator"}
}
}
return VerifyResult{Checked: true, Resolved: true, Detail: "the flagged account is no longer a WordPress administrator"}
}
// verifyDBSuspiciousAdminEmail resolves when no administrator still uses the
// flagged disposable email address.
func verifyDBSuspiciousAdminEmail(message, details string) VerifyResult {
dbName := detailField(details, "Database")
email := strings.ToLower(strings.TrimSpace(detailField(details, "Email")))
if email == "" {
return VerifyResult{Checked: false, Detail: "could not parse the admin email from the finding"}
}
prefixes, ok := findWPAdminVerifyPrefixes(dbFindingAccount(message, details), dbName, details)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
for _, prefix := range prefixes {
users, usermeta, ok := wpAdminQueryTables(prefix)
if !ok {
return VerifyResult{Checked: false, Detail: "could not validate the WordPress tables from the finding"}
}
rows, err := runDBVerifyQueryRoot(dbName, fmt.Sprintf(
"SELECT u.ID FROM %s u JOIN %s m ON u.ID = m.user_id WHERE m.meta_key = ? AND m.meta_value LIKE '%%administrator%%' AND LOWER(u.user_email) = ?",
users, usermeta), prefix+"capabilities", email)
if err != nil {
return dbVerifyQueryError()
}
if len(rows) > 0 {
return VerifyResult{Checked: true, Resolved: false, Detail: "an administrator still uses the flagged email address"}
}
}
return VerifyResult{Checked: true, Resolved: true, Detail: "no administrator still uses the flagged email address"}
}
// dbAdminRowID extracts the numeric account id from a per-CMS admin finding's
// "Row:" detail line, whose first tab- (or space-) separated field is the id.
func dbAdminRowID(details string) (string, bool) {
row := detailField(details, "Row")
if row == "" {
return "", false
}
first := row
if i := strings.IndexAny(row, "\t "); i >= 0 {
first = row[:i]
}
first = strings.TrimSpace(first)
if first == "" || !isAllDigits(first) {
return "", false
}
return first, true
}
// dbAdminPresenceResult maps a presence query result to a VerifyResult: a query
// error is never resolved; zero rows means the flagged account is gone.
func dbAdminPresenceResult(err error, rows []string) VerifyResult {
if err != nil {
return dbVerifyQueryError()
}
if len(rows) == 0 {
return VerifyResult{Checked: true, Resolved: true, Detail: "the flagged administrator account is no longer present"}
}
return VerifyResult{Checked: true, Resolved: false, Detail: "the flagged administrator account is still present"}
}
func verifyDrupalAdminInjection(message, details string) VerifyResult {
uid, ok := dbAdminRowID(details)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the account id from the finding"}
}
schema, ok := discoverDrupalSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Drupal site")
}
rows, err := runDBVerifyQueryRoot(schema,
"SELECT u.uid FROM users_field_data u JOIN user__roles r ON u.uid = r.entity_id WHERE r.roles_target_id = ? AND u.default_langcode = 1 AND u.uid = ?",
drupalAdminRoleID, uid)
return dbAdminPresenceResult(err, rows)
}
func verifyJoomlaAdminInjection(message, details string) VerifyResult {
id, ok := dbAdminRowID(details)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the account id from the finding"}
}
schema, prefix, ok := discoverJoomlaSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Joomla site")
}
users, err := QuoteIdent(prefix + "users")
if err != nil {
return VerifyResult{Checked: false, Detail: "could not validate the Joomla tables from the finding"}
}
mapTable, err := QuoteIdent(prefix + "user_usergroup_map")
if err != nil {
return VerifyResult{Checked: false, Detail: "could not validate the Joomla tables from the finding"}
}
rows, err := runDBVerifyQueryRoot(schema, fmt.Sprintf(
"SELECT u.id FROM %s u JOIN %s m ON u.id = m.user_id WHERE m.group_id = ? AND u.id = ?",
users, mapTable), joomlaSuperUserGroupID, id)
return dbAdminPresenceResult(err, rows)
}
func verifyMagentoAdminInjection(message, details string) VerifyResult {
id, ok := dbAdminRowID(details)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the account id from the finding"}
}
schema, prefix, ok := discoverMagentoSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Magento site")
}
table, err := QuoteIdent(prefix + "admin_user")
if err != nil {
return VerifyResult{Checked: false, Detail: "could not validate the Magento table from the finding"}
}
rows, err := runDBVerifyQueryRoot(schema,
fmt.Sprintf("SELECT user_id FROM %s WHERE user_id = ?", table), id)
return dbAdminPresenceResult(err, rows)
}
func verifyOpenCartAdminInjection(message, details string) VerifyResult {
id, ok := dbAdminRowID(details)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the account id from the finding"}
}
schema, prefix, ok := discoverOpenCartSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("OpenCart site")
}
table, err := QuoteIdent(prefix + "user")
if err != nil {
return VerifyResult{Checked: false, Detail: "could not validate the OpenCart table from the finding"}
}
rows, err := runDBVerifyQueryRoot(schema,
fmt.Sprintf("SELECT user_id FROM %s WHERE user_id = ?", table), id)
return dbAdminPresenceResult(err, rows)
}
package checks
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
)
// Per-finding Re-check for the non-WordPress CMS database-content findings
// (Drupal, Joomla, Magento, OpenCart). Each verifier re-discovers the single
// canonical install under /home/<account>/public_html, re-reads the one flagged
// row as root against that install's schema, and re-runs the SAME malware
// classifier the detector used. It resolves only when the row is gone or no
// longer classifies as malware; any query failure or undiscoverable install
// returns Checked:false so a live injection is never auto-cleared.
// dbVerifyRowMalware re-queries a single value column for one entity and
// resolves the finding when no returned row still classifies as malware.
// inPostContext selects the looser external-script predicate the detector uses
// for author-written content (article/product bodies) versus config storage.
func dbVerifyRowMalware(schema, query string, inPostContext bool, args ...any) VerifyResult {
rows, err := runDBVerifyQueryRoot(schema, query, args...)
if err != nil {
return dbVerifyQueryError()
}
for _, row := range rows {
if _, _, ok := classifyMalwareRow(row, inPostContext); ok {
return VerifyResult{Checked: true, Resolved: false, Detail: "the flagged row still contains injected content"}
}
}
return VerifyResult{Checked: true, Resolved: true, Detail: "the flagged row no longer contains injected content"}
}
// --- Drupal ---------------------------------------------------------------
func discoverDrupalSchema(account string) (schema string, ok bool) {
if !validAccountName.MatchString(account) {
return "", false
}
publicHTML := filepath.Join(accountHomeDir(account), "public_html")
if matched, err := looksLikeDrupal8Plus(publicHTML); err != nil || !matched {
return "", false
}
creds, err := parseDrupalSettings(context.Background(), filepath.Join(publicHTML, "sites", "default", "settings.php"))
if err != nil || creds.dbName == "" {
return "", false
}
return creds.dbName, true
}
func verifyDrupalSettingsInjection(message, details string) VerifyResult {
name := detailField(details, "Config name")
if name == "" {
return VerifyResult{Checked: false, Detail: "could not parse the config name from the finding"}
}
schema, ok := discoverDrupalSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Drupal site")
}
return dbVerifyRowMalware(schema, "SELECT data FROM config WHERE name = ?", false, name)
}
func verifyDrupalContentInjection(message, details string) VerifyResult {
entityID := detailField(details, "Node entity_id")
if entityID == "" || !isAllDigits(entityID) {
return VerifyResult{Checked: false, Detail: "could not parse the node id from the finding"}
}
schema, ok := discoverDrupalSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Drupal site")
}
return dbVerifyRowMalware(schema, "SELECT body_value FROM node_revision__body WHERE entity_id = ?", true, entityID)
}
// --- Joomla ---------------------------------------------------------------
func discoverJoomlaSchema(account string) (schema, prefix string, ok bool) {
if !validAccountName.MatchString(account) {
return "", "", false
}
path := filepath.Join(accountHomeDir(account), "public_html", "configuration.php")
if matched, err := looksLikeJoomlaConfig(context.Background(), path); err != nil || !matched {
return "", "", false
}
creds, err := parseJConfig(context.Background(), path)
if err != nil || creds.dbName == "" {
return "", "", false
}
prefix = creds.dbPrefix
if prefix == "" {
prefix = "jos_"
}
if !validTablePrefix.MatchString(prefix) {
return "", "", false
}
return creds.dbName, prefix, true
}
func verifyJoomlaExtensionsInjection(message, details string) VerifyResult {
name := detailField(details, "Extension")
if name == "" {
return VerifyResult{Checked: false, Detail: "could not parse the extension name from the finding"}
}
schema, prefix, ok := discoverJoomlaSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Joomla site")
}
return dbVerifyRowMalware(schema,
fmt.Sprintf("SELECT params FROM `%sextensions` WHERE name = ?", prefix), false, name)
}
func verifyJoomlaContentInjection(message, details string) VerifyResult {
id := detailField(details, "Article ID")
if id == "" || !isAllDigits(id) {
return VerifyResult{Checked: false, Detail: "could not parse the article id from the finding"}
}
schema, prefix, ok := discoverJoomlaSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Joomla site")
}
return dbVerifyRowMalware(schema,
fmt.Sprintf("SELECT introtext FROM `%scontent` WHERE id = ?", prefix), true, id)
}
// --- Magento --------------------------------------------------------------
func discoverMagentoSchema(account string) (schema, prefix string, ok bool) {
if !validAccountName.MatchString(account) {
return "", "", false
}
base := filepath.Join(accountHomeDir(account), "public_html", "app", "etc")
creds, err := parseMagentoM2(context.Background(), filepath.Join(base, "env.php"))
if err == nil {
if creds.dbName == "" {
return "", "", false
}
return magentoVerifyFinalize(creds)
}
// An unsafe or unreadable current config must not select an older schema
// and let a re-check clear the finding against that unrelated database.
if !errors.Is(err, os.ErrNotExist) {
return "", "", false
}
if creds, err := parseMagentoM1(context.Background(), filepath.Join(base, "local.xml")); err == nil && creds.dbName != "" {
return magentoVerifyFinalize(creds)
}
return "", "", false
}
// magentoVerifyFinalize validates the parsed prefix. Magento's default table
// prefix is empty (no prefix); a non-empty prefix must pass validTablePrefix
// before it is interpolated into a table name.
func magentoVerifyFinalize(creds magentoCreds) (string, string, bool) {
if creds.dbPrefix != "" && !validTablePrefix.MatchString(creds.dbPrefix) {
return "", "", false
}
return creds.dbName, creds.dbPrefix, true
}
func verifyMagentoSettingsInjection(message, details string) VerifyResult {
cfgPath := detailField(details, "Config path")
if cfgPath == "" {
return VerifyResult{Checked: false, Detail: "could not parse the config path from the finding"}
}
schema, prefix, ok := discoverMagentoSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Magento site")
}
return dbVerifyRowMalware(schema,
fmt.Sprintf("SELECT value FROM `%score_config_data` WHERE path = ?", prefix), false, cfgPath)
}
// magentoContentTableCols maps a Magento content table to its (idColumn,
// valueColumn). Only the three tables the detector scans are accepted; anything
// else returns ok=false so a malformed finding can never name an arbitrary
// table in the re-query.
func magentoContentTableCols(table string) (idCol, valueCol string, ok bool) {
switch table {
case "catalog_product_entity_text":
return "entity_id", "value", true
case "cms_block":
return "block_id", "content", true
case "cms_page":
return "page_id", "content", true
}
return "", "", false
}
func verifyMagentoContentInjection(message, details string) VerifyResult {
table := detailField(details, "Table")
rowID := detailField(details, "Row id")
idCol, valueCol, ok := magentoContentTableCols(table)
if !ok || rowID == "" || !isAllDigits(rowID) {
return VerifyResult{Checked: false, Detail: "could not parse the affected row from the finding"}
}
schema, prefix, ok := discoverMagentoSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("Magento site")
}
return dbVerifyRowMalware(schema,
fmt.Sprintf("SELECT %s FROM `%s%s` WHERE %s = ?", valueCol, prefix, table, idCol), true, rowID)
}
// --- OpenCart -------------------------------------------------------------
func discoverOpenCartSchema(account string) (schema, prefix string, ok bool) {
if !validAccountName.MatchString(account) {
return "", "", false
}
path := filepath.Join(accountHomeDir(account), "public_html", "config.php")
if matched, err := looksLikeOpenCart(context.Background(), path); err != nil || !matched {
return "", "", false
}
creds, err := parseOpenCartConfig(context.Background(), path)
if err != nil || creds.dbName == "" {
return "", "", false
}
prefix = creds.dbPrefix
if prefix == "" {
prefix = "oc_"
}
if !validTablePrefix.MatchString(prefix) {
return "", "", false
}
return creds.dbName, prefix, true
}
func verifyOpenCartSettingsInjection(message, details string) VerifyResult {
key := detailField(details, "Setting key")
if key == "" {
return VerifyResult{Checked: false, Detail: "could not parse the setting key from the finding"}
}
schema, prefix, ok := discoverOpenCartSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("OpenCart site")
}
return dbVerifyRowMalware(schema,
fmt.Sprintf("SELECT value FROM `%ssetting` WHERE `key` = ?", prefix), false, key)
}
// openCartContentTableCols maps an OpenCart content table to its id column. Both
// description tables share the "description" value column. Unknown tables are
// rejected.
func openCartContentTableCols(table string) (idCol string, ok bool) {
switch table {
case "product_description":
return "product_id", true
case "information_description":
return "information_id", true
}
return "", false
}
func verifyOpenCartContentInjection(message, details string) VerifyResult {
table := detailField(details, "Table")
rowID := detailField(details, "Row id")
idCol, ok := openCartContentTableCols(table)
if !ok || rowID == "" || !isAllDigits(rowID) {
return VerifyResult{Checked: false, Detail: "could not parse the affected row from the finding"}
}
schema, prefix, ok := discoverOpenCartSchema(dbFindingAccount(message, details))
if !ok {
return dbVerifyNotLocatable("OpenCart site")
}
return dbVerifyRowMalware(schema,
fmt.Sprintf("SELECT description FROM `%s%s` WHERE %s = ? AND language_id = 1", prefix, table, idCol), true, rowID)
}
package checks
import (
"context"
"fmt"
"strings"
"github.com/pidginhost/csm/internal/mysqlclient"
)
// Per-finding Re-check for the database persistence-mechanism findings
// (triggers, events, procedures, functions) and the backdoor magic-token user
// finding. The object re-checks re-read the object's CURRENT body from
// INFORMATION_SCHEMA (not the truncated copy stored in the finding details) and
// re-run the same malware classifier the detector used, so a partially-cleaned
// object is never mistaken for a fixed one. All queries run as root and any
// failure returns Checked:false.
// runDBVerifyQueryRootGlobal runs a root query with no pinned default schema,
// for INFORMATION_SCHEMA lookups whose WHERE clause already scopes the schema.
// If the account database was dropped entirely, the object is gone with it and
// the query simply returns zero rows -- the correct "resolved" signal.
func runDBVerifyQueryRootGlobal(query string, args ...any) ([]string, error) {
ctx, cancel := context.WithTimeout(context.Background(), dbVerifyTimeout)
defer cancel()
return mysqlclient.RootQuery(ctx, query, args...)
}
// verifyDBObject re-checks an unexpected/malicious trigger, event, procedure, or
// function. For the unexpected tier any surviving object keeps the finding
// active. For the malicious tier the finding clears when the object is gone or
// its current body no longer matches a malware pattern.
func verifyDBObject(message, details string, malicious bool) VerifyResult {
schema := detailField(details, "Schema")
kind := detailField(details, "Kind")
name := detailField(details, "Name")
if schema == "" || name == "" || !IsDBObjectKind(kind) {
return VerifyResult{Checked: false, Detail: "could not parse the database object from the finding"}
}
body, present, err := fetchDBObjectBody(schema, dbObjectKind(kind), name)
if err != nil {
return dbVerifyQueryError()
}
if !present {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("%s %q no longer exists", kind, name)}
}
if !malicious {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("%s %q is still present", kind, name)}
}
if bodyHasMalwarePattern(body) {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("%s %q still matches a malware pattern", kind, name)}
}
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("%s %q no longer matches a malware pattern", kind, name)}
}
// fetchDBObjectBody returns the current definition body of one database object
// and whether it still exists. The object name and schema are passed as bound
// parameters so a malformed finding cannot inject into the lookup.
func fetchDBObjectBody(schema string, kind dbObjectKind, name string) (body string, present bool, err error) {
var query string
var args []any
switch kind {
case dbObjectTrigger:
query = "SELECT ACTION_STATEMENT FROM INFORMATION_SCHEMA.TRIGGERS WHERE TRIGGER_SCHEMA = ? AND TRIGGER_NAME = ?"
args = []any{schema, name}
case dbObjectEvent:
query = "SELECT EVENT_DEFINITION FROM INFORMATION_SCHEMA.EVENTS WHERE EVENT_SCHEMA = ? AND EVENT_NAME = ?"
args = []any{schema, name}
case dbObjectProcedure, dbObjectFunction:
routineType := "PROCEDURE"
if kind == dbObjectFunction {
routineType = "FUNCTION"
}
query = "SELECT ROUTINE_DEFINITION FROM INFORMATION_SCHEMA.ROUTINES WHERE ROUTINE_SCHEMA = ? AND ROUTINE_NAME = ? AND ROUTINE_TYPE = ?"
args = []any{schema, name, routineType}
default:
return "", false, fmt.Errorf("unknown object kind")
}
rows, err := runDBVerifyQueryRootGlobal(query, args...)
if err != nil {
return "", false, err
}
if len(rows) == 0 {
return "", false, nil
}
return rows[0], true, nil
}
// verifyDBMagicTokenUser re-queries the WordPress users table for any account
// whose flagged user row still matches the trigger's backdoor activation
// pattern. The token is validated as the high-entropy [A-Za-z0-9_-]{10,32}
// shape the detector emits before it is used in the LIKE.
func verifyDBMagicTokenUser(message, details string) VerifyResult {
token := detailField(details, "Token")
if token == "" || !validMagicToken(token) {
return VerifyResult{Checked: false, Detail: "could not parse a valid activation token from the finding"}
}
schema := strings.TrimSpace(detailField(details, "Schema"))
if schema == "" {
return VerifyResult{Checked: false, Detail: "could not parse the database schema from the finding"}
}
userID, ok := dbMagicTokenUserID(details)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the WordPress user ID from the finding"}
}
prefix, ok := dbMagicTokenUserTablePrefix(message, details)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the WordPress table prefix from the finding"}
}
prefix, ok = locateDBMagicTokenUserPrefix(dbFindingAccount(message, details), schema, prefix)
if !ok {
return dbVerifyNotLocatable("WordPress site")
}
table, err := QuoteIdent(prefix + "users")
if err != nil {
return VerifyResult{Checked: false, Detail: "could not validate the WordPress users table from the finding"}
}
rows, err := runDBVerifyQueryRoot(schema,
fmt.Sprintf("SELECT ID FROM %s WHERE ID = ? AND display_name LIKE ?", table),
userID, "%"+token+"%")
if err != nil {
return dbVerifyQueryError()
}
if len(rows) == 0 {
return VerifyResult{Checked: true, Resolved: true, Detail: "the flagged user no longer carries the backdoor activation token"}
}
return VerifyResult{Checked: true, Resolved: false, Detail: "the flagged user still carries the backdoor activation token"}
}
func dbMagicTokenUserID(details string) (string, bool) {
userID := strings.TrimSpace(detailField(details, "User ID"))
if userID == "" || !isAllDigits(userID) {
return "", false
}
return userID, true
}
func dbMagicTokenUserTablePrefix(message, details string) (string, bool) {
if prefix, ok := validDBMagicTokenTablePrefix(detailField(details, "Table prefix")); ok {
return prefix, true
}
const marker = " in "
idx := strings.LastIndex(message, marker)
if idx < 0 {
return "", false
}
ref := strings.TrimSpace(message[idx+len(marker):])
if !strings.HasSuffix(ref, "users") {
return "", false
}
dot := strings.LastIndex(ref, ".")
if dot < 0 {
return "", false
}
return validDBMagicTokenTablePrefix(strings.TrimSuffix(ref[dot+1:], "users"))
}
func validDBMagicTokenTablePrefix(prefix string) (string, bool) {
prefix = strings.TrimSpace(prefix)
if prefix == "" || !validTablePrefix.MatchString(prefix) {
return "", false
}
return prefix, true
}
func locateDBMagicTokenUserPrefix(account, schema, prefix string) (string, bool) {
prefixes, ok := findWPVerifyPrefixes(account, schema, "Table prefix: "+prefix)
if !ok || len(prefixes) != 1 {
return "", false
}
return prefixes[0], true
}
package checks
import (
"context"
"fmt"
"regexp"
"strings"
"time"
)
// pkgVerifyTimeout bounds the synchronous package-manager re-check a Re-check
// click runs.
var pkgVerifyTimeout = 30 * time.Second
// pkgNameRe is the safe character set for a package name passed as an argument
// to rpm/dpkg/debsums. Findings carry CSM-generated package names, but the
// re-check still validates before exec so a malformed name can never inject.
var pkgNameRe = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._+:~-]*$`)
// parsePkgIntegrityFinding extracts the file and package from an
// rpm_integrity / dpkg_integrity message:
// "Modified system binary or library: <file> (package: <pkg>)".
func parsePkgIntegrityFinding(message string) (file, pkg string, ok bool) {
rest, ok := strings.CutPrefix(message, "Modified system binary or library: ")
if !ok {
return "", "", false
}
const sep = " (package: "
i := strings.LastIndex(rest, sep)
if i < 0 || !strings.HasSuffix(rest, ")") {
return "", "", false
}
file = strings.TrimSpace(rest[:i])
pkg = strings.TrimSpace(rest[i+len(sep) : len(rest)-1])
if file == "" || pkg == "" {
return "", "", false
}
return file, pkg, true
}
type packageVerifyOutputState struct {
targetFlagged bool
sawReport bool
sawUnknown bool
}
// verifyManifestLineFlagsFile reports whether an rpm -V / dpkg --verify output
// line marks the target file as size/checksum-modified (mirrors the detector:
// skip config/doc and require S or 5). The existing finding already classified
// the file; changing its mode or removing it must not clear a reported mismatch.
func verifyManifestLineFlagsFile(line, file string) (targetFlagged, recognized bool) {
flags, got, ok := parseManifestVerifyLine(line)
if !ok {
return false, false
}
if !manifestVerifyLineHasPackagePath(line, got) {
return false, false
}
if strings.Contains(line, " c ") || strings.Contains(line, " d ") {
return false, true
}
if !strings.Contains(flags, "S") && !strings.Contains(flags, "5") {
return false, true
}
return got == file, true
}
func manifestVerifyLineHasPackagePath(line, got string) bool {
if strings.HasPrefix(got, "/") {
return true
}
fields := strings.Fields(strings.TrimSpace(line[9:]))
if len(fields) != 2 || (fields[0] != "c" && fields[0] != "d") {
return false
}
return strings.HasPrefix(fields[1], "/")
}
func parseManifestVerifyLine(line string) (flags, file string, ok bool) {
if len(line) < 10 {
return "", "", false
}
flags = line[:9]
for _, ch := range flags {
if !strings.ContainsRune(".?SM5DLUGTP", ch) {
return "", "", false
}
}
file = strings.TrimSpace(line[9:])
if file == "" {
return "", "", false
}
return flags, file, true
}
func manifestOutputState(out []byte, file string) packageVerifyOutputState {
var state packageVerifyOutputState
for _, line := range strings.Split(strings.TrimSpace(string(out)), "\n") {
if strings.TrimSpace(line) == "" {
continue
}
targetFlagged, recognized := verifyManifestLineFlagsFile(line, file)
if !recognized {
state.sawUnknown = true
continue
}
state.sawReport = true
if targetFlagged {
state.targetFlagged = true
}
}
return state
}
func debsumsOutputState(out []byte, file string) packageVerifyOutputState {
var state packageVerifyOutputState
for _, line := range strings.Split(strings.TrimSpace(string(out)), "\n") {
got := strings.TrimSpace(line)
if got == "" {
continue
}
if !strings.HasPrefix(got, "/") {
state.sawUnknown = true
continue
}
state.sawReport = true
if got == file {
state.targetFlagged = true
}
}
return state
}
func resolvePkgVerifyOutput(state packageVerifyOutputState, file, verifier string) VerifyResult {
if state.targetFlagged {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("%s is still reported as modified", file)}
}
if state.sawReport && !state.sawUnknown {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("%s is no longer reported as modified", file)}
}
return VerifyResult{Checked: false, Detail: fmt.Sprintf("could not parse %s output (try again, or run an account scan)", verifier)}
}
func pkgVerifyTimedOut(ctx context.Context, verifier string) (VerifyResult, bool) {
if ctx.Err() == nil {
return VerifyResult{}, false
}
return VerifyResult{Checked: false, Detail: fmt.Sprintf("%s timed out (try again, or run an account scan)", verifier)}, true
}
// verifyRPMIntegrity re-runs `rpm -V <pkg>` and resolves the finding when the
// flagged file is no longer reported as modified (or the whole package verifies
// clean). Read-only, bounded; any command failure returns Checked:false so a
// real tampered binary is never auto-cleared.
func verifyRPMIntegrity(message string) VerifyResult {
file, pkg, ok := parsePkgIntegrityFinding(message)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the package finding"}
}
if !pkgNameRe.MatchString(pkg) {
return VerifyResult{Checked: false, Detail: "package name in finding is not auto-verifiable"}
}
ctx, cancel := context.WithTimeout(context.Background(), pkgVerifyTimeout)
defer cancel()
// rpm -V exits non-zero when files are modified; treat output, not exit
// code, as the signal.
out, err := cmdExec.RunContext(ctx, "rpm", "-V", pkg)
if res, timedOut := pkgVerifyTimedOut(ctx, "rpm -V"); timedOut {
return res
}
if strings.TrimSpace(string(out)) == "" {
if err == nil {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("package %s verifies clean", pkg)}
}
return VerifyResult{Checked: false, Detail: "could not run rpm -V (try again, or run an account scan)"}
}
return resolvePkgVerifyOutput(manifestOutputState(out, file), file, "rpm -V")
}
// verifyDpkgIntegrity re-runs debsums (preferred) or `dpkg --verify` for the
// package and resolves the finding when the flagged file is no longer reported
// as modified. Same safety contract as verifyRPMIntegrity.
func verifyDpkgIntegrity(message string) VerifyResult {
file, pkg, ok := parsePkgIntegrityFinding(message)
if !ok {
return VerifyResult{Checked: false, Detail: "could not parse the package finding"}
}
if !pkgNameRe.MatchString(pkg) {
return VerifyResult{Checked: false, Detail: "package name in finding is not auto-verifiable"}
}
ctx, cancel := context.WithTimeout(context.Background(), pkgVerifyTimeout)
defer cancel()
if _, err := cmdExec.LookPath("debsums"); err == nil {
// debsums -c prints one modified file path per line (exit 2 on mismatch).
out, err := cmdExec.RunContext(ctx, "debsums", "-c", pkg)
if res, timedOut := pkgVerifyTimedOut(ctx, "debsums"); timedOut {
return res
}
if strings.TrimSpace(string(out)) == "" {
if err == nil {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("package %s verifies clean", pkg)}
}
return VerifyResult{Checked: false, Detail: "could not run debsums (try again, or run an account scan)"}
}
return resolvePkgVerifyOutput(debsumsOutputState(out, file), file, "debsums")
}
out, err := cmdExec.RunContext(ctx, "dpkg", "--verify", pkg)
if res, timedOut := pkgVerifyTimedOut(ctx, "dpkg --verify"); timedOut {
return res
}
if strings.TrimSpace(string(out)) == "" {
if err == nil {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("package %s verifies clean", pkg)}
}
return VerifyResult{Checked: false, Detail: "could not run dpkg --verify (try again, or run an account scan)"}
}
return resolvePkgVerifyOutput(manifestOutputState(out, file), file, "dpkg --verify")
}
package checks
import (
"fmt"
"os"
"strings"
)
// allowedUID0 are the system accounts that legitimately carry UID 0; the
// detector and the re-check share this list so they agree on what counts as
// "unauthorized".
var allowedUID0 = map[string]bool{
"root": true, "sync": true, "shutdown": true,
"halt": true, "operator": true,
}
// classifyUID0Line parses one /etc/passwd line and reports the account name and
// whether it is an unauthorized UID 0 account.
func classifyUID0Line(line string) (user string, unauthorized bool) {
fields := strings.Split(line, ":")
if len(fields) < 4 {
return "", false
}
user = fields[0]
return user, fields[2] == "0" && !allowedUID0[user]
}
// verifyUID0Account re-reads /etc/passwd and resolves the finding when the
// flagged account is gone, no longer UID 0, or now an allowed system account.
// Read-only; an unreadable /etc/passwd returns Checked:false rather than a
// false clear.
func verifyUID0Account(message string) VerifyResult {
rest, ok := strings.CutPrefix(message, "Unauthorized UID 0 account: ")
if !ok {
return VerifyResult{Checked: false, Detail: "could not determine the account from the finding"}
}
user := rest
if user == "" {
return VerifyResult{Checked: false, Detail: "could not determine the account from the finding"}
}
if strings.ContainsAny(user, ":\r\n") {
return VerifyResult{Checked: false, Detail: "account name in finding is not auto-verifiable"}
}
data, err := osFS.ReadFile("/etc/passwd")
if err != nil {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("cannot read /etc/passwd: %v", err)}
}
found := false
for _, line := range strings.Split(string(data), "\n") {
u, unauthorized := classifyUID0Line(line)
if u != user {
continue
}
found = true
if unauthorized {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("%s is still an unauthorized UID 0 account", user)}
}
}
if found {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("%s is no longer an unauthorized UID 0 account", user)}
}
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("account %s no longer exists", user)}
}
// verifySuidCleared re-stats a flagged SUID binary and resolves the finding
// when the file is gone or no longer carries the setuid bit. Bounded to the
// same roots scanForSUID covers (/home, /tmp, /dev/shm, /var/tmp).
func verifySuidCleared(path string) VerifyResult {
if path == "" {
return VerifyResult{Checked: false, Detail: "could not extract file path from finding"}
}
clean, info, exists, err := readOnlyFixPath(path, effectiveFixRoots(fixQuarantineAllowedRoots, quarantineExtraRoots...))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("SUID binary no longer present: %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: "path is not a regular file; not auto-verifiable"}
}
if info.Mode()&os.ModeSetuid != 0 {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("file is still setuid (mode %s)", info.Mode())}
}
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("file is no longer setuid (mode %s)", info.Mode())}
}
package checks
import (
"context"
"fmt"
"path/filepath"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/store"
)
// wpVerifyAllowedRoots bounds where a WordPress re-check may run wp-cli. It is a
// var so tests can redirect under t.TempDir(); nil means the platform's
// account roots.
var wpVerifyAllowedRoots []string
// wpVerifyTimeout bounds the synchronous wp-cli re-scan a Re-check click runs.
var wpVerifyTimeout = 30 * time.Second
// findingDetailPath extracts the "Path: <dir>" value emitted in a finding's
// Details (outdated_plugins and the WordPress checks record the install path
// there). Returns "" when no such line is present.
func findingDetailPath(details string) string {
for _, line := range strings.Split(details, "\n") {
line = strings.TrimSpace(line)
if rest, ok := strings.CutPrefix(line, "Path:"); ok {
return strings.TrimSpace(rest)
}
}
return ""
}
func wpChecksumLineHasExtraneousCoreFile(line string) bool {
return strings.Contains(line, "should not exist") && !strings.Contains(line, "error_log")
}
// The two per-file shapes wp-cli has used. The closing summary line
// ("WordPress installation doesn't verify against checksums.") matches
// neither: the current shape needs the colon and the legacy one the singular
// "checksum." directly after the file name.
const (
wpChecksumMismatchCurrent = "doesn't verify against checksum: "
wpChecksumMismatchLegacy = " doesn't verify against checksum."
)
// wpChecksumModifiedCoreFile returns the install-relative path of a core file
// that wp-cli reports as changed, or "" when the line is not such a report.
// Localised packages legitimately ship their own root readme and license,
// which carry no code, so those two are not reported.
func wpChecksumModifiedCoreFile(line string) string {
rel := wpChecksumModifiedFilePath(line)
if rel == "readme.html" || rel == "license.txt" {
return ""
}
return rel
}
// Keep recognized but intentionally unreported mismatches distinct from command
// failures when accounting for a completed installation scan.
func wpChecksumModifiedFilePath(line string) string {
var rel string
if idx := strings.Index(line, wpChecksumMismatchCurrent); idx >= 0 {
rel = strings.TrimSpace(line[idx+len(wpChecksumMismatchCurrent):])
} else if idx := strings.Index(line, wpChecksumMismatchLegacy); idx >= 0 {
if fields := strings.Fields(line[:idx]); len(fields) > 0 {
rel = fields[len(fields)-1]
}
}
return rel
}
// wpCoreModifiedSeverity grades a modified core file by what an attacker could
// do with it.
//
// Critical is what auto-response acts on, so it is reserved for files that can
// carry executable content to a visitor: anything a PHP handler runs, and the
// scripts and templates served into the browser. Everything else -- stylesheets,
// images, translations, fonts -- still gets a finding, but a mismatch there is
// far more often an asset optimiser or an install whose version.php no longer
// names the release its files came from than it is an appended backdoor.
func wpCoreModifiedSeverity(path, rel string) alert.Severity {
if contenttype.IsExecutablePHPName(strings.ToLower(rel)) {
return alert.Critical
}
switch strings.ToLower(filepath.Ext(rel)) {
case ".js", ".mjs", ".html", ".htm", ".htaccess":
return alert.Critical
case ".svg", ".xml", ".xhtml":
// Markup, not an image format: SVG carries <script> and event
// handlers and the browser runs them. Judge it by what it holds,
// because most core SVG mismatches are an optimiser's whitespace.
if wpCoreMarkupIsActive(path) {
return alert.Critical
}
}
return alert.High
}
// wpCoreMarkupActiveMarkers are the element-level constructs that make markup
// executable. Event attributes are matched separately: SVG defines dozens of
// them, so any list of names would really be a list of the ones an attacker
// has to avoid.
var wpCoreMarkupActiveMarkers = []string{
"<script", "javascript:", "<foreignobject", "<!entity", "<handler",
"<animate", "<set ", "<use ", "data:text/html",
}
// wpCoreMarkupEventAttr matches any on<name>= handler attribute.
var wpCoreMarkupEventAttr = regexp.MustCompile(`(?i)\bon[a-z]+\s*=`)
// wpCoreMarkupPeekBytes bounds the read. A shipped core asset is far smaller;
// anything larger is judged active without reading further, because the part
// that was not read is exactly where content would be hidden.
const wpCoreMarkupPeekBytes = 256 << 10
// wpCoreMarkupIsActive reports whether a markup file carries anything the
// browser would execute.
//
// Every uncertain answer is "active": an unreadable file, one larger than the
// peek, a path raced to something that is not a regular file. None of those is
// evidence of innocence, and the point of the grade is that a core asset which
// cannot be shown inert keeps the higher severity.
func wpCoreMarkupIsActive(path string) bool {
if path == "" {
return true
}
info, err := osFS.Lstat(path)
if err != nil || !info.Mode().IsRegular() || info.Size() > wpCoreMarkupPeekBytes {
return true
}
// The path is in a tenant-writable tree, so it can be raced to a FIFO
// between the report and this read. Use the non-blocking, O_NOFOLLOW,
// regular-file-verified reader rather than a plain open.
reader, ok := osFS.(phpRegularFilePrefixReader)
if !ok {
return true
}
body, err := reader.ReadRegularFilePrefix(path, info, wpCoreMarkupPeekBytes)
if err != nil {
return true
}
lower := strings.ToLower(string(body))
for _, marker := range wpCoreMarkupActiveMarkers {
if strings.Contains(lower, marker) {
return true
}
}
return wpCoreMarkupEventAttr.MatchString(lower)
}
// wpCoreFilePathWithin joins a wp-cli reported relative path onto the install
// only when it stays inside it; wp-cli output is not a path oracle for the
// remediation and Re-check code that consumes FilePath.
func wpCoreFilePathWithin(wpPath, rel string) string {
clean := filepath.Clean(rel)
if filepath.IsAbs(clean) || clean == "." || clean == ".." || strings.HasPrefix(clean, "../") {
return ""
}
return filepath.Join(wpPath, clean)
}
// verifyOutdatedPlugins re-inventories a single WordPress site with wp-cli (run
// as the site owner) and resolves the finding when no active plugin still has
// an available update. It is heavier than the file re-checks but read-only and
// bounded by wpVerifyTimeout. Any failure to re-scan returns Checked:false so a
// finding is never falsely cleared on a transient wp-cli error.
func verifyOutdatedPlugins(details string) VerifyResult {
wpPath := findingDetailPath(details)
if wpPath == "" {
return VerifyResult{Checked: false, Detail: "could not determine the WordPress path from the finding"}
}
clean, _, exists, err := readOnlyFixPath(wpPath, effectiveFixRoots(wpVerifyAllowedRoots))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("WordPress install no longer exists: %s", clean)}
}
wpConfig, info, exists, err := readOnlyFixPath(filepath.Join(clean, "wp-config.php"), effectiveFixRoots(wpVerifyAllowedRoots))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("WordPress install no longer present: %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: "wp-config.php path is not a regular file; not auto-verifiable"}
}
ctx, cancel := context.WithTimeout(context.Background(), wpVerifyTimeout)
defer cancel()
site, err := inventoryWPSiteForVerify(ctx, wpConfig)
if err != nil {
return VerifyResult{Checked: false, Detail: fmt.Sprintf("could not re-scan plugins (try again, or run an account scan): %v", err)}
}
if n := countOutdatedActivePlugins(site, store.Global()); n > 0 {
return VerifyResult{Checked: true, Resolved: false, Detail: fmt.Sprintf("%d active plugin(s) still outdated", n)}
}
return VerifyResult{Checked: true, Resolved: true, Detail: "no active plugins are outdated anymore"}
}
// verifyWPCoreIntegrity re-runs `wp core verify-checksums` for one install and
// resolves the finding when no extraneous core file ("should not exist")
// remains -- mirroring CheckWPCore, which only flags those lines. It is
// read-only and bounded by wpVerifyTimeout. To avoid ever clearing a real
// compromise it resolves only when verification is clean or the install is
// gone; any wp-cli error (including modified files with no remaining
// extra-file line) returns Checked:false.
func verifyWPCoreIntegrity(details string) VerifyResult {
wpPath := findingDetailPath(details)
if wpPath == "" {
return VerifyResult{Checked: false, Detail: "could not determine the WordPress path from the finding"}
}
clean, _, exists, err := readOnlyFixPath(wpPath, effectiveFixRoots(wpVerifyAllowedRoots))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("WordPress install no longer exists: %s", clean)}
}
_, info, exists, err := readOnlyFixPath(filepath.Join(clean, "wp-config.php"), effectiveFixRoots(wpVerifyAllowedRoots))
if err != nil {
return VerifyResult{Checked: false, Detail: err.Error()}
}
if !exists {
return VerifyResult{Checked: true, Resolved: true, Detail: fmt.Sprintf("WordPress install no longer present: %s", clean)}
}
if !info.Mode().IsRegular() {
return VerifyResult{Checked: false, Detail: "wp-config.php path is not a regular file; not auto-verifiable"}
}
ctx, cancel := context.WithTimeout(context.Background(), wpVerifyTimeout)
defer cancel()
// Mirrors CheckWPCore: run as root with --allow-root (not su as the user).
// Args are passed directly (no shell), so the sanitized path cannot inject.
out, err := cmdExec.RunContext(ctx, "wp", "core", "verify-checksums", "--path="+clean, "--allow-root")
if err == nil {
return VerifyResult{Checked: true, Resolved: true, Detail: "WordPress core checksums verify clean"}
}
if len(out) == 0 {
return VerifyResult{Checked: false, Detail: "could not run wp core verify-checksums (try again, or run an account scan)"}
}
for _, line := range strings.Split(string(out), "\n") {
if wpChecksumLineHasExtraneousCoreFile(line) {
return VerifyResult{Checked: true, Resolved: false, Detail: "WordPress core still has extraneous files"}
}
if wpChecksumModifiedCoreFile(line) != "" {
return VerifyResult{Checked: true, Resolved: false, Detail: "WordPress core still has modified files"}
}
}
// Non-zero exit with neither an extraneous nor a modified file named: a
// wp-cli error we cannot interpret. Do not resolve.
return VerifyResult{Checked: false, Detail: "could not confirm core integrity (verify-checksums reported other issues); re-run or use an account scan"}
}
package checks
import (
"context"
_ "embed"
"fmt"
"strings"
"time"
"gopkg.in/yaml.v3"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// Known-vulnerable WordPress plugin detector.
//
// CheckOutdatedPlugins grades severity by how far a plugin is behind the latest
// release, which buries an actively-exploited version that is only a few
// releases stale. This detector instead matches the installed version against a
// curated feed of known-vulnerable ranges and elevates confirmed hits (CISA-KEV
// / unauthenticated RCE, privesc, SQLi, secret disclosure) to High or Critical
// regardless of version gap. It shares the cached inventory refresh with
// outdated_plugins and only alerts (v1 never disables a plugin, so a real
// customer site is never taken down).
//go:embed embed/plugin_vulns.yaml
var pluginVulnFeedData []byte
// pluginVuln is one known-vulnerable version range for a plugin slug.
type pluginVuln struct {
Slug string `yaml:"slug"`
CVE string `yaml:"cve"`
Title string `yaml:"title"`
FixedIn string `yaml:"fixed_in"`
MinAffected string `yaml:"min_affected"` // optional lower bound
VirtualPatch bool `yaml:"virtual_patch"` // CSM ships a ModSecurity rule for this CVE
KEV bool `yaml:"kev"`
Severity string `yaml:"severity"`
Reference string `yaml:"reference"`
}
type pluginVulnFeed struct {
Plugins []pluginVuln `yaml:"plugins"`
}
// loadPluginVulnFeed parses the curated feed and fails closed on individual
// malformed entries: an entry with missing or invalid range fields is dropped
// (rather than matching everything) so a bad line cannot create false positives.
func loadPluginVulnFeed(data []byte) ([]pluginVuln, error) {
var feed pluginVulnFeed
if err := yaml.Unmarshal(data, &feed); err != nil {
return nil, err
}
out := make([]pluginVuln, 0, len(feed.Plugins))
for _, v := range feed.Plugins {
if strings.TrimSpace(v.Slug) == "" || strings.TrimSpace(v.CVE) == "" || strings.TrimSpace(v.FixedIn) == "" {
continue
}
if _, ok := parsePluginVersion(v.FixedIn); !ok {
continue
}
if strings.TrimSpace(v.MinAffected) != "" {
cmp, ok := comparePluginVersions(v.MinAffected, v.FixedIn)
if !ok || cmp >= 0 {
continue
}
}
out = append(out, v)
}
return out, nil
}
// parsePluginVersion returns normalized decimal components without converting
// them to machine integers. WordPress versions are not strict semver, so a
// suffix on a numeric component is tolerated and ignored. A component that
// does not start with a digit makes the version unusable; an unknown version
// must never underflow into a false vulnerability match.
func parsePluginVersion(v string) ([]string, bool) {
v = strings.TrimSpace(v)
if len(v) > 1 && (v[0] == 'v' || v[0] == 'V') && v[1] >= '0' && v[1] <= '9' {
v = v[1:]
}
if v == "" {
return nil, false
}
parts := strings.Split(v, ".")
out := make([]string, 0, len(parts))
for _, part := range parts {
i := 0
for i < len(part) && part[i] >= '0' && part[i] <= '9' {
i++
}
if i == 0 {
return nil, false
}
digits := strings.TrimLeft(part[:i], "0")
if digits == "" {
digits = "0"
}
out = append(out, digits)
if i != len(part) {
if part[i] != '-' && part[i] != '+' {
return nil, false
}
break
}
}
return out, true
}
func comparePluginVersions(a, b string) (int, bool) {
av, aok := parsePluginVersion(a)
bv, bok := parsePluginVersion(b)
if !aok || !bok {
return 0, false
}
n := len(av)
if len(bv) > n {
n = len(bv)
}
for i := 0; i < n; i++ {
x, y := "0", "0"
if i < len(av) {
x = av[i]
}
if i < len(bv) {
y = bv[i]
}
if len(x) < len(y) {
return -1, true
}
if len(x) > len(y) {
return 1, true
}
if x < y {
return -1, true
}
if x > y {
return 1, true
}
}
return 0, true
}
// versionLess reports whether version a is strictly older than b. Invalid
// versions are incomparable and therefore never considered older.
func versionLess(a, b string) bool {
cmp, ok := comparePluginVersions(a, b)
return ok && cmp < 0
}
// matchPluginVuln reports whether an installed version falls inside a
// vulnerable range: strictly below fixed_in and, when a lower bound is given,
// at or above min_affected. An install at or above fixed_in is patched and
// never matches.
func matchPluginVuln(installed string, v pluginVuln) bool {
if strings.TrimSpace(installed) == "" {
return false
}
if !versionLess(installed, v.FixedIn) {
return false
}
if strings.TrimSpace(v.MinAffected) != "" && versionLess(installed, v.MinAffected) {
return false
}
return true
}
// vulnPluginSeverity maps a matched vulnerability to a finding severity. A
// version-matched known CVE is actionable by definition, so it is never a
// Warning: KEV/actively-exploited or an explicit "critical" is Critical, an
// explicit "high" is High, and anything else defaults to Critical.
func vulnPluginSeverity(v pluginVuln) alert.Severity {
if v.KEV {
return alert.Critical
}
if strings.EqualFold(strings.TrimSpace(v.Severity), "high") {
return alert.High
}
return alert.Critical
}
func vulnAllowKey(slug, version string) string {
return strings.ToLower(strings.TrimSpace(slug) + "@" + strings.TrimSpace(version))
}
// evaluatePluginVulns matches the cached per-site plugin inventory against the
// feed and returns one finding per confirmed vulnerable install, skipping any
// slug@version the operator has explicitly accepted via the allowlist. Each
// match carries the facts the WAF-coverage correlation needs and the finding
// itself does not record.
func evaluatePluginVulns(sites map[string]store.SitePlugins, feed []pluginVuln, allow map[string]bool) []vulnMatch {
bySlug := make(map[string][]pluginVuln, len(feed))
for _, v := range feed {
key := strings.ToLower(strings.TrimSpace(v.Slug))
bySlug[key] = append(bySlug[key], v)
}
var matches []vulnMatch
for wpPath, site := range sites {
for _, p := range site.Plugins {
for _, v := range bySlug[strings.ToLower(strings.TrimSpace(p.Slug))] {
if !matchPluginVuln(p.InstalledVersion, v) {
continue
}
if allow[vulnAllowKey(p.Slug, p.InstalledVersion)] {
continue
}
matches = append(matches, vulnMatch{
finding: buildVulnPluginFinding(wpPath, site, p, v),
active: vulnPluginActive(p.Status),
vpCovered: v.VirtualPatch,
})
}
}
}
return matches
}
// vulnMatch is one confirmed vulnerable install: the finding the detector
// built, plus whether WordPress actually loads the plugin and whether CSM
// ships a ModSecurity virtual patch for the CVE.
//
// An inactive plugin still earns an inventory finding because some plugins
// expose directly callable files, but that alone does not establish that this
// CVE is reachable without WordPress loading it, so it is never described as
// left open by a missing request filter.
type vulnMatch struct {
finding alert.Finding
active bool
vpCovered bool
}
func buildVulnPluginFinding(wpPath string, site store.SitePlugins, p store.SitePluginEntry, v pluginVuln) alert.Finding {
activeNote := "inactive"
if vulnPluginActive(p.Status) {
activeNote = "active"
}
kevNote := ""
if v.KEV {
kevNote = " [CISA-KEV: actively exploited]"
}
details := fmt.Sprintf("%s (%s)%s. Installed %s %s is %s; fixed in %s. Remediate: update to %s or newer, or remove the plugin.",
v.Title, v.CVE, kevNote, p.Slug, p.InstalledVersion, activeNote, v.FixedIn, v.FixedIn)
if v.Reference != "" {
details += "\nReference: " + v.Reference
}
details += "\nPath: " + wpPath
return alert.Finding{
Severity: vulnPluginSeverity(v),
Check: "vulnerable_plugins",
Message: fmt.Sprintf("Known-vulnerable plugin %s %s (%s) on %s", p.Slug, p.InstalledVersion, v.CVE, site.Domain),
Details: details,
Domain: site.Domain,
TenantID: site.Account,
Timestamp: time.Now(),
}
}
func vulnPluginActive(status string) bool {
switch strings.ToLower(strings.TrimSpace(status)) {
case "active", "active-network", "must-use":
return true
default:
return false
}
}
// CheckVulnerablePlugins matches the shared WordPress plugin inventory against
// the curated known-vulnerable feed. It participates in the same serialized
// refresh as CheckOutdatedPlugins and only reports -- it never disables a
// plugin.
func CheckVulnerablePlugins(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
if cfg != nil && !cfg.VulnerablePluginScanningEnabled() {
return nil
}
if incompleteCollectorFrom(ctx) == nil {
ctx, _ = withIncompleteCheckCollector(ctx)
}
db := store.Global()
if db == nil {
return nil
}
fresh := ensurePluginCacheFresh(ctx, cfg, db)
if ctx.Err() != nil {
return nil
}
// Discovery gaps retain the last usable inventory; evaluating it keeps a
// broken account map from hiding known vulnerabilities until discovery recovers.
if !fresh && !checkMarkedIncomplete(ctx, "vulnerable_plugins") {
return nil
}
feed, err := loadPluginVulnFeed(pluginVulnFeedData)
if err != nil || len(feed) == 0 {
return nil
}
matches := evaluatePluginVulns(db.AllSitePlugins(), feed, vulnPluginAllowSet(cfg))
if len(matches) == 0 {
return nil
}
var candidates []alert.Finding
for _, m := range matches {
if m.active {
candidates = append(candidates, m.finding)
}
}
if len(candidates) > 0 {
annotateUnprotected(matches, vpCoverageForHost(candidates))
}
findings := make([]alert.Finding, 0, len(matches))
for _, m := range matches {
findings = append(findings, m.finding)
}
return findings
}
func vulnPluginAllowSet(cfg *config.Config) map[string]bool {
if cfg == nil || len(cfg.Detection.VulnerablePluginAllow) == 0 {
return nil
}
set := make(map[string]bool, len(cfg.Detection.VulnerablePluginAllow))
for _, e := range cfg.Detection.VulnerablePluginAllow {
e = strings.TrimSpace(e)
if e == "" {
continue
}
set[strings.ToLower(e)] = true
}
return set
}
package checks
import (
"fmt"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/platform"
)
// Correlation between a known-vulnerable plugin and the compensating control
// that is supposed to hold the line until it is patched.
//
// CSM ships ModSecurity virtual patches for several of the CVEs in the plugin
// feed, so a vulnerable install is normally filtered while the operator
// schedules an update. When ModSecurity is off for that traffic the patch never
// executes, and the same finding now describes a directly reachable
// vulnerability. Reported separately, the two facts read as routine
// housekeeping; joined, they are an emergency.
//
// This adds no new detection: both inputs are findings CSM already produces.
// vpCoverage is the state that decides whether CSM's shipped virtual patches
// filter a given site.
type vpCoverage struct {
// engineMode is the host-wide SecRuleEngine value ("on", "detectiononly",
// "off"), empty when it could not be read.
engineMode string
// disabled lists the accounts and vhosts with ModSecurity switched off.
disabled []modsecDisabledScope
// aliases maps an account's addon domain to its unambiguous cPanel-associated
// subdomain (and vice versa). A per-vhost disabled flag can be recorded
// against that servername while the plugin inventory knows the public name.
aliases map[string][]string
}
// aliasKey scopes an alias set to one account: unrelated accounts may name the
// same docroot path, and one account's disabled flag says nothing about
// another's traffic.
func aliasKey(user, domain string) string {
return strings.ToLower(strings.TrimSpace(user)) + "\x00" + strings.ToLower(strings.TrimSpace(domain))
}
// vhostAliasSets links only an unambiguous addon/subdomain pair: the docroot
// group must contain exactly those two records and the subdomain must be the
// exact cPanel association <addon-domain>.<main-domain>. A shared docroot is
// not itself proof that two hostnames are aliases: parked domains, a main
// domain, and an addon can all intentionally route different sites from
// /home/<user>/public_html.
func vhostAliasSets(userdataDomains string) map[string][]string {
vhosts, complete := parseUserdataDomainRootsChecked(userdataDomains)
if !complete {
return nil
}
byRoot := make(map[string][]vhost, len(vhosts))
for _, vh := range vhosts {
root := aliasKey(vh.user, vh.docroot)
byRoot[root] = append(byRoot[root], vh)
}
sets := make(map[string][]string, len(vhosts))
for _, group := range byRoot {
if len(group) != 2 {
continue
}
firstType := strings.ToLower(strings.TrimSpace(group[0].typ))
secondType := strings.ToLower(strings.TrimSpace(group[1].typ))
var addon, sub vhost
switch {
case firstType == "addon" && secondType == "sub":
addon, sub = group[0], group[1]
case firstType == "sub" && secondType == "addon":
addon, sub = group[1], group[0]
default:
continue
}
mainDomain := cleanDomlogDomain(addon.mainDomain)
if mainDomain == "" ||
!strings.EqualFold(cleanDomlogDomain(sub.mainDomain), mainDomain) ||
!strings.EqualFold(sub.domain, addon.domain+"."+mainDomain) {
continue
}
for _, vh := range group {
sets[aliasKey(vh.user, vh.domain)] = []string{group[0].domain, group[1].domain}
}
}
return sets
}
// siteNames is every domain the site answers to: the one the inventory
// recorded plus its docroot peers.
func (c vpCoverage) siteNames(account, domain string) map[string]bool {
names := map[string]bool{strings.ToLower(strings.TrimSpace(domain)): true}
for _, peer := range c.aliases[aliasKey(account, domain)] {
names[strings.ToLower(strings.TrimSpace(peer))] = true
}
delete(names, "")
return names
}
// inertReason explains why CSM's virtual patches do not filter this site, or
// returns an empty reason when they do.
func (c vpCoverage) inertReason(account, domain string) (reason, source string) {
switch strings.ToLower(strings.TrimSpace(c.engineMode)) {
case "off":
return "the ModSecurity engine is off host-wide", ""
case "detectiononly":
return "the ModSecurity engine runs in DetectionOnly mode host-wide", ""
}
if strings.TrimSpace(account) == "" {
return "", ""
}
names := c.siteNames(account, domain)
var accountWide, exact, alias *modsecDisabledScope
for _, s := range c.disabled {
if !strings.EqualFold(strings.TrimSpace(s.User), strings.TrimSpace(account)) {
continue
}
if s.Domain == "" {
if accountWide == nil {
scope := s
accountWide = &scope
}
continue
}
if !names[strings.ToLower(strings.TrimSpace(s.Domain))] {
continue
}
if strings.EqualFold(s.Domain, domain) {
if exact == nil {
scope := s
exact = &scope
}
continue
}
if alias == nil {
scope := s
alias = &scope
}
}
if accountWide != nil {
return "ModSecurity is disabled for account " + accountWide.User + " (all domains)", accountWide.Source
}
if exact != nil {
return "ModSecurity is disabled for " + exact.Domain, exact.Source
}
if alias != nil {
return fmt.Sprintf("ModSecurity is disabled for %s, the cPanel-associated subdomain of %s", alias.Domain, domain), alias.Source
}
return "", ""
}
// annotateUnprotected rewrites vulnerable-plugin findings whose traffic no
// longer passes through ModSecurity, so the alert itself carries the fact that
// nothing stands between the vulnerability and the internet.
//
// Only an active install is rewritten: an inactive plugin is a finding because
// its files sit in the docroot, not because WordPress will run the vulnerable
// code path, so a missing request filter says nothing about its reachability.
// What the annotation claims depends on whether CSM ships a virtual patch for
// the CVE -- naming a patch that was never written would misdescribe the gap.
func annotateUnprotected(matches []vulnMatch, cov vpCoverage) {
for i := range matches {
if !matches[i].active {
continue
}
f := &matches[i].finding
// The production caller passes freshly built findings, but keeping this
// helper idempotent prevents a retrying caller from changing the alert
// identity and appending the same operator guidance repeatedly.
if strings.Contains(f.Message, " -- unprotected: ") ||
strings.Contains(f.Details, "\n\nUnprotected: ") {
continue
}
reason, source := cov.inertReason(f.TenantID, f.Domain)
if reason == "" {
continue
}
f.Severity = alert.Critical
f.Message += " -- unprotected: " + reason
located := reason
if source != "" {
located += " (" + source + ")"
}
gap := "No ModSecurity rule filters this traffic and no modsec audit record\n" +
"is written for it, so an attempt to exploit this leaves no WAF evidence.\n"
if matches[i].vpCovered {
gap = "CSM ships a virtual patch for this CVE and it cannot run here, and no\n" +
"modsec audit record is written for this traffic either, so the\n" +
"vulnerability is directly reachable and an attempt to exploit it\n" +
"leaves no WAF evidence.\n"
}
f.Details += "\n\nUnprotected: " + located + ".\n" + gap +
"Patch the plugin now, or restore filtering for this scope first."
}
}
// vpCoverageForHost is the seam the wiring is tested through: the snapshot it
// returns is assembled from platform detection and host config that a unit
// test cannot stage.
var vpCoverageForHost = currentVPCoverage
// currentVPCoverage reads the live ModSecurity state. Only cPanel receives
// CSM's virtual patches and only cPanel expresses per-vhost ModSecurity
// state, so on every other platform the question has no answer and the
// findings are left untouched.
func currentVPCoverage(findings []alert.Finding) vpCoverage {
info := platform.Detect()
if !info.IsCPanel() {
return vpCoverage{}
}
cov := vpCoverage{engineMode: checkEngineMode(info)}
if data, err := osFS.ReadFile(userdataDomainsPath); err == nil {
cov.aliases = vhostAliasSets(string(data))
}
// CheckWAFStatus already walks every userdata record in this scan tier.
// Correlation needs only the vulnerable sites, so read their account,
// domain, and unambiguous associated-subdomain paths instead of repeating the
// host-wide O(number of vhosts) traversal.
cov.disabled = modsecDisabledScopesForFindings(info, findings, cov.aliases)
return cov
}
package checks
import (
"bufio"
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"os"
"path/filepath"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/modsec"
"github.com/pidginhost/csm/internal/netutil"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// wafRulesAssembleRetryDelay is the wait between the first negative
// probe and the re-probe on cPanel+LiteSpeed hosts where cPanel's
// nightly modsec_assemble briefly leaves both `whmapi1
// modsec_get_configs` and the vendor dir empty while it rewrites the
// tree in place. Observed windows are <10s; 30s gives margin without
// meaningfully delaying the surrounding deep-scan tier. Tests override
// this to keep the suite fast.
var wafRulesAssembleRetryDelay = 30 * time.Second
// CheckWAFStatus verifies that ModSecurity is loaded, the engine is in
// enforcement mode (not DetectionOnly), OWASP/Comodo rules are active,
// and rules are up to date.
func CheckWAFStatus(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
manageHost := cfg == nil || !cfg.ObserveMode()
info := platform.Detect()
// If there is no web server at all, WAF concerns don't apply to this host.
if info.WebServer == platform.WSNone {
return findings
}
modsecActive := modsecDetected(info)
if !modsecActive {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "waf_status",
Message: "ModSecurity WAF is not active",
Details: wafInstallHint(info),
})
return findings // no point checking further
}
// --- Engine mode check ---
engineMode := checkEngineMode(info)
if engineMode == "detectiononly" {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "waf_detection_only",
Message: "ModSecurity is in DetectionOnly mode - attacks are logged but NOT blocked",
Details: "SecRuleEngine is set to DetectionOnly. Change to 'On' for enforcement:\nWHM > Security Center > ModSecurity > Edit Global Directive",
})
}
// --- Rule vendor check ---
ruleDirs := modsecRuleDirs(info)
hasRules := probeWAFRules(info, ruleDirs)
// cPanel+LiteSpeed: cPanel's nightly modsec_assemble rewrites the
// vendor tree in place, so for ~6-10s both `whmapi1
// modsec_get_configs` and the vendor dir return empty. A production
// false positive at 01:10:27 fired 6s after the rewrite. Re-probe
// once after a short delay before alerting; on a host that really
// has no rules, the re-probe is still negative and we alert in the
// same scan, so this doesn't shift detection to the next deep tier.
if !hasRules && info.IsCPanel() && info.WebServer == platform.WSLiteSpeed {
select {
case <-time.After(wafRulesAssembleRetryDelay):
case <-ctx.Done():
return findings
}
hasRules = probeWAFRules(info, ruleDirs)
}
if !hasRules {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "waf_rules",
Message: "ModSecurity has no WAF rules loaded",
Details: wafRulesHint(info),
})
}
// --- Rule age check + auto-update ---
if hasRules {
staleAge := checkRuleAge(info, ruleDirs)
if staleAge > 0 {
// Attempt auto-update before alerting
updated := false
if info.IsCPanel() && manageHost {
updated = autoUpdateWAFRules()
}
if updated {
// Re-check age after update
staleAge = checkRuleAge(info, ruleDirs)
}
if staleAge > 0 {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "waf_rules_stale",
Message: fmt.Sprintf("ModSecurity rules last updated %d days ago - update recommended", staleAge),
Details: wafRulesStaleHint(info),
})
}
}
}
// --- Virtual patch deployment ---
// Only cPanel has the modsec user config dirs we write into.
if info.IsCPanel() && manageHost {
reloadCommand := ""
if cfg != nil {
reloadCommand = cfg.ModSec.ReloadCommand
}
// Hosts without a reload command are warned at daemon startup; a
// finding here would repeat every scan with nothing CSM can verify.
if err := deployAndReconcileModSec(ctx, reloadCommand); err != nil {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "waf_status",
Message: "CSM ModSecurity rule activation could not be confirmed",
Details: err.Error(),
})
}
}
// --- Disabled ModSecurity scopes ---
// whmapi1 reports disabled *rules*, never a disabled engine, so this
// reads the userdata and conf.d state that actually gates filtering.
findings = append(findings, modsecDisabledFindings(info)...)
return findings
}
// probeWAFRules checks whether any WAF rule source -- cPanel's whmapi1
// active-config list or the on-disk vendor/CRS directories -- currently
// reports rules. Used by CheckWAFStatus directly and again on retry
// for the cPanel+LiteSpeed modsec_assemble race.
func probeWAFRules(info platform.Info, ruleDirs []string) bool {
if info.IsCPanel() {
if activeRules, ok := cPanelActiveRules(); ok {
if activeRules.present {
return true
}
// During LiteSpeed's modsec_assemble window, the WHM config
// list and filesystem can become visible in either order. Keep
// the existing filesystem backstop so either source can end the
// retry without letting retired trees affect other cPanel hosts.
if info.WebServer != platform.WSLiteSpeed {
return false
}
return hasRuleArtifacts(ruleDirs)
}
if out, _ := runCmd("whmapi1", "modsec_get_vendors"); out != nil {
outStr := string(out)
if strings.Contains(outStr, "comodo") || strings.Contains(outStr, "owasp") ||
strings.Contains(outStr, "OWASP") || strings.Contains(outStr, "Comodo") {
return true
}
}
}
return hasRuleArtifacts(ruleDirs)
}
// modsecDetected returns true if a ModSecurity module is loaded for the
// detected web server. It first consults the platform layer, then falls
// back to scanning config files.
func modsecDetected(info platform.Info) bool {
// cPanel fast path
if info.IsCPanel() {
if out, _ := runCmd("whmapi1", "modsec_is_installed"); out != nil &&
strings.Contains(string(out), "installed: 1") {
return true
}
}
// Generic file-based probes per web server
for _, conf := range expandPathGlobs(modsecActivationCandidates(info)) {
data, err := osFS.ReadFile(conf)
if err != nil {
continue
}
if modsecEnabledInConfig(info, string(data)) {
return true
}
}
return false
}
// modsecActivationCandidates returns the config files that can enable the
// ModSecurity module for the detected web server.
func modsecActivationCandidates(info platform.Info) []string {
var paths []string
switch info.WebServer {
case platform.WSApache:
if info.ApacheConfigDir != "" {
paths = append(paths,
filepath.Join(info.ApacheConfigDir, "httpd.conf"),
filepath.Join(info.ApacheConfigDir, "apache2.conf"),
filepath.Join(info.ApacheConfigDir, "modsec2.conf"),
filepath.Join(info.ApacheConfigDir, "conf.d", "modsec2.conf"),
filepath.Join(info.ApacheConfigDir, "mods-enabled", "security2.conf"),
filepath.Join(info.ApacheConfigDir, "conf-enabled", "security2.conf"),
filepath.Join(info.ApacheConfigDir, "conf.d", "mod_security.conf"),
filepath.Join(info.ApacheConfigDir, "conf.modules.d", "10-mod_security.conf"),
filepath.Join(info.ApacheConfigDir, "conf.d", "*.conf"),
filepath.Join(info.ApacheConfigDir, "mods-enabled", "*.conf"),
filepath.Join(info.ApacheConfigDir, "conf-enabled", "*.conf"),
)
}
case platform.WSNginx:
if info.NginxConfigDir != "" {
paths = append(paths,
filepath.Join(info.NginxConfigDir, "nginx.conf"),
filepath.Join(info.NginxConfigDir, "conf.d", "*.conf"),
filepath.Join(info.NginxConfigDir, "sites-enabled", "*"),
)
}
case platform.WSLiteSpeed:
paths = append(paths,
"/usr/local/lsws/conf/httpd_config.xml",
)
}
return paths
}
// modsecConfigCandidates returns the set of config files worth scanning
// for ModSecurity directives on the detected web server.
func modsecConfigCandidates(info platform.Info) []string {
paths := append([]string(nil), modsecActivationCandidates(info)...)
switch info.WebServer {
case platform.WSNginx:
if info.NginxConfigDir != "" {
paths = append(paths,
filepath.Join(info.NginxConfigDir, "modules-enabled", "*.conf"),
filepath.Join(info.NginxConfigDir, "modsec", "main.conf"),
)
}
case platform.WSLiteSpeed:
paths = append(paths, "/usr/local/lsws/conf/modsec2.conf")
}
return paths
}
func modsecEnabledInConfig(info platform.Info, contents string) bool {
scanner := bufio.NewScanner(strings.NewReader(contents))
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" {
continue
}
lineLower := strings.ToLower(line)
switch info.WebServer {
case platform.WSNginx:
if strings.HasPrefix(lineLower, "#") {
continue
}
if strings.HasPrefix(lineLower, "modsecurity on") ||
strings.HasPrefix(lineLower, "modsecurity_rules ") ||
strings.HasPrefix(lineLower, "modsecurity_rules_file ") {
return true
}
case platform.WSLiteSpeed:
if strings.Contains(lineLower, "mod_security") || strings.Contains(lineLower, "modsecurity") {
return true
}
default:
if strings.HasPrefix(lineLower, "#") {
continue
}
if strings.Contains(lineLower, "security2_module") ||
strings.HasPrefix(lineLower, "secruleengine ") ||
strings.Contains(lineLower, "mod_security2") {
return true
}
}
}
return false
}
func expandPathGlobs(paths []string) []string {
var expanded []string
seen := make(map[string]struct{})
for _, candidate := range paths {
matches := []string{candidate}
if strings.ContainsAny(candidate, "*?[") {
if globbed, err := osFS.Glob(candidate); err == nil && len(globbed) > 0 {
matches = globbed
}
}
for _, match := range matches {
if _, ok := seen[match]; ok {
continue
}
seen[match] = struct{}{}
expanded = append(expanded, match)
}
}
return expanded
}
// modsecRuleDirs delegates to the canonical helper in internal/modsec.
// Kept as a package-local thin wrapper because the existing waf check tests
// reference this name directly.
func modsecRuleDirs(info platform.Info) []string {
return modsec.RuleDirs(info)
}
// wafInstallHint returns platform-specific install instructions.
func wafInstallHint(info platform.Info) string {
switch {
case info.IsCPanel():
return "No ModSecurity module detected. Install: WHM > Security Center > ModSecurity"
case info.WebServer == platform.WSNginx && info.IsDebianFamily():
return "No ModSecurity module detected for Nginx.\nInstall: apt install libnginx-mod-http-modsecurity modsecurity-crs"
case info.WebServer == platform.WSApache && info.IsDebianFamily():
return "No ModSecurity module detected for Apache.\nInstall: apt install libapache2-mod-security2 modsecurity-crs && a2enmod security2"
case info.WebServer == platform.WSApache && info.IsRHELFamily():
return "No ModSecurity module detected for Apache.\nInstall (requires EPEL): dnf install -y epel-release && dnf install -y mod_security mod_security_crs && systemctl restart httpd"
case info.WebServer == platform.WSNginx && info.IsRHELFamily():
return "No ModSecurity module detected for Nginx.\nInstall (requires EPEL): dnf install -y epel-release && dnf install -y nginx-mod-http-modsecurity && systemctl restart nginx"
}
return "No ModSecurity module detected. The server has no web application firewall protecting against SQL injection, XSS, and other web attacks."
}
// wafRulesHint returns platform-specific rules-install instructions.
func wafRulesHint(info platform.Info) string {
if info.IsCPanel() {
return "ModSecurity is installed but has no OWASP or Comodo rules. Add rules: WHM > Security Center > ModSecurity Vendors"
}
if info.IsDebianFamily() {
return "ModSecurity is installed but has no rules loaded. Install OWASP CRS: apt install modsecurity-crs"
}
if info.IsRHELFamily() {
return "ModSecurity is installed but has no rules loaded. Install OWASP CRS: dnf install --enablerepo=epel modsecurity-crs"
}
return "ModSecurity is installed but has no rules loaded."
}
// wafRulesStaleHint returns platform-specific advice for updating stale
// ModSecurity vendor rules.
func wafRulesStaleHint(info platform.Info) string {
if info.IsCPanel() {
return "At least one vendor with active configuration files has not refreshed its rules in over a month. Inactive vendor trees are not counted. Check: WHM > Security Center > ModSecurity Vendors > Update"
}
if info.IsDebianFamily() {
return "Vendor rules should be updated at least monthly. Update with: apt update && apt upgrade modsecurity-crs"
}
if info.IsRHELFamily() {
return "Vendor rules should be updated at least monthly. Update with: dnf upgrade modsecurity-crs"
}
return "Vendor rules should be updated at least monthly."
}
// checkEngineMode determines the host-wide SecRuleEngine setting. On cPanel
// the generated configuration is authoritative: modsec2.conf turns the
// engine on and then includes modsec2.cpanel.conf (WHM "Edit Global
// Directive") and modsec2.user.conf, and Apache applies the last directive
// it parses, so the mode is whatever the last file in that chain says. A
// distro package's file is never consulted on cPanel: an inactive file
// reported as the host setting would manufacture a false unprotected
// Critical. Returns "on", "detectiononly", "off", or "" if unknown.
func checkEngineMode(info platform.Info) string {
if info.IsCPanel() {
return cPanelEngineMode(info)
}
configPaths := modsecConfigCandidates(info)
// Also include the top-level modsecurity.conf installed by distro packages.
configPaths = append(configPaths,
"/etc/modsecurity/modsecurity.conf",
"/etc/nginx/modsec/modsecurity.conf",
)
for _, path := range configPaths {
if mode := engineModeInFile(path); mode != "" {
return mode
}
}
return ""
}
// cPanelModsecEngineChain lists, in include order, the cPanel files that can
// set SecRuleEngine at server scope under the Apache-compatible config
// directory (cPanel + LiteSpeed reads the same tree through loadApacheConf).
// The modsec include directory is spelled both ways across EA4 and the
// older /usr/local/apache layout, and CSM already writes virtual patches to
// both, so each link tries every known spelling and uses the first that
// exists.
func cPanelModsecEngineChain(configDir string) [][]string {
return [][]string{
{filepath.Join(configDir, "conf.d", "modsec2.conf")},
{
filepath.Join(configDir, "conf.d", "modsec", "modsec2.cpanel.conf"),
filepath.Join(configDir, "conf.d", "modsec2.cpanel.conf"),
filepath.Join(configDir, "modsec2.cpanel.conf"),
},
{
filepath.Join(configDir, "conf.d", "modsec", "modsec2.user.conf"),
filepath.Join(configDir, "conf.d", "modsec2.user.conf"),
},
}
}
// cPanelEngineMode walks the include chain and returns the last engine
// directive. A chain file that exists but cannot be read, or that this
// parser cannot interpret authoritatively, makes the effective mode
// unknown: guessing from the files around it could hide a DetectionOnly
// set through WHM or report a stale On.
func cPanelEngineMode(info platform.Info) string {
configDir := filepath.Clean(info.ApacheCompatibleConfigDir())
if configDir == "." || configDir == string(filepath.Separator) {
return ""
}
mode := ""
for _, spellings := range cPanelModsecEngineChain(configDir) {
for _, path := range spellings {
m, known, err := engineDirectiveInFile(path)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
continue
}
return ""
}
if !known {
return ""
}
if m != "" {
mode = m
}
// The first spelling that exists is the file Apache reads.
break
}
}
return mode
}
// engineModeInFile is the distro-package form of engineDirectiveInFile: a
// missing, unreadable or ambiguous file simply yields unknown.
func engineModeInFile(path string) string {
mode, known, err := engineDirectiveInFile(path)
if err != nil || !known {
return ""
}
return mode
}
// lastEngineDirective returns the last server-scope SecRuleEngine value
// Apache would apply from path, or "" when the file is missing or cannot be
// interpreted authoritatively.
func lastEngineDirective(path string) string {
mode, known := readLastEngineDirective(path)
if !known {
return ""
}
return mode
}
// readLastEngineDirective also reports whether the file was interpreted
// authoritatively. A missing optional include has no effect, but an
// unreadable or ambiguous later include must invalidate an earlier mode
// instead of silently preserving it.
func readLastEngineDirective(path string) (string, bool) {
mode, known, err := engineDirectiveInFile(path)
if err != nil {
return "", errors.Is(err, os.ErrNotExist)
}
return mode, known
}
// engineDirectiveInFile opens path and returns the last server-scope
// SecRuleEngine value Apache would apply from it, lower-cased. known is
// false when the file could not be interpreted authoritatively; err carries
// the open or read failure so callers can tell a missing file from an
// unreadable one.
func engineDirectiveInFile(path string) (mode string, known bool, err error) {
f, err := osFS.Open(path)
if err != nil {
return "", false, err
}
defer func() { _ = f.Close() }()
mode, known, err = parseEngineDirectives(f)
return mode, known, err
}
// parseEngineDirectives scans Apache configuration text for the last
// server-scope SecRuleEngine directive. The cPanel wrapper commonly places
// the directive inside a positive mod_security2 IfModule block. A directive
// in an unrelated IfModule, in a request or virtual-host container, or in a
// runtime-conditional container makes the effective mode unknown because
// this parser cannot know which other Apache modules or defines are live.
func parseEngineDirectives(r io.Reader) (mode string, known bool, err error) {
type container struct {
name string
active bool
activeKnown bool
serverScope bool
}
var stack []container
active, activeKnown, serverScope := true, true, true
valid := true
scanner := bufio.NewScanner(r)
for scanner.Scan() {
line := strings.TrimSpace(stripApacheComment(scanner.Text()))
if line == "" {
continue
}
tag, isContainer, tagValid := parseApacheContainerTag(line)
if isContainer {
if !tagValid {
valid = false
continue
}
if tag.closing {
if len(stack) == 0 || !strings.EqualFold(stack[len(stack)-1].name, tag.name) {
valid = false
continue
}
stack = stack[:len(stack)-1]
active, activeKnown, serverScope = true, true, true
if len(stack) > 0 {
active = stack[len(stack)-1].active
activeKnown = stack[len(stack)-1].activeKnown
serverScope = stack[len(stack)-1].serverScope
}
continue
}
next := container{name: tag.name, active: active, activeKnown: activeKnown, serverScope: serverScope}
switch {
case strings.EqualFold(tag.name, "IfModule"):
condition, condKnown := activeModSecurityIfModule(tag.label)
switch {
case activeKnown && !active:
// An inactive parent keeps every nested directive inactive.
case condKnown && !condition:
next.active = false
next.activeKnown = true
case condKnown:
next.activeKnown = activeKnown
case activeKnown:
next.active = false
next.activeKnown = false
default:
next.active = false
}
case apacheEngineScopedContainer(tag.name):
// A directive in a request or virtual-host context does not
// define the server-wide engine mode.
next.serverScope = false
case activeKnown && active:
// IfDefine, IfVersion, IfFile and expression containers keep
// server scope but depend on runtime state unavailable here.
next.active = false
next.activeKnown = false
}
stack = append(stack, next)
active, activeKnown, serverScope = next.active, next.activeKnown, next.serverScope
continue
}
fields, fieldsValid := parseApacheDirectiveFields(line)
if !fieldsValid {
valid = false
continue
}
if !serverScope || len(fields) == 0 || !strings.EqualFold(fields[0], "SecRuleEngine") {
continue
}
if !activeKnown {
valid = false
continue
}
if !active {
continue
}
if len(fields) != 2 {
valid = false
continue
}
switch value := strings.ToLower(fields[1]); value {
case "on", "off", "detectiononly":
mode = value
default:
valid = false
}
}
if scanErr := scanner.Err(); scanErr != nil {
return "", false, scanErr
}
if len(stack) != 0 || !valid {
return "", false, nil
}
return mode, true, nil
}
func apacheEngineScopedContainer(name string) bool {
switch strings.ToLower(name) {
case "virtualhost", "directory", "directorymatch", "files", "filesmatch",
"location", "locationmatch", "proxy", "limit", "limitexcept":
return true
default:
return false
}
}
func activeModSecurityIfModule(label string) (active, known bool) {
fields, valid := parseApacheDirectiveFields(label)
if !valid || len(fields) != 2 {
return false, false
}
module := strings.ToLower(fields[1])
negated := strings.HasPrefix(module, "!")
module = strings.TrimPrefix(module, "!")
switch module {
case "mod_security2.c", "security2_module":
return !negated, true
default:
return false, false
}
}
// checkRuleAge returns the age of the rules that protect the host, or 0 when
// they were refreshed within the last 30 days.
//
// cPanel retains inactive vendor trees indefinitely, so its active-config API
// narrows the scan to loaded vendor paths. The newest artifact in each loaded
// vendor tree is its refresh time; the least recently refreshed loaded vendor
// decides the reported age. On other platforms, and when the API is
// unavailable, retaining the former oldest-file behavior is conservative: a
// fresh candidate tree that is not loaded must not hide stale rules in use.
func checkRuleAge(info platform.Info, ruleDirs []string) int {
artifactMtime := oldestRuleArtifact
if info.IsCPanel() {
activeRules, ok := cPanelActiveRules()
if ok {
if !activeRules.present || len(activeRules.vendorDirs) == 0 {
return 0
}
ruleDirs = activeRules.vendorDirs
artifactMtime = oldestRulesetRefresh
}
}
mtime, found := artifactMtime(ruleDirs)
if !found {
return 0
}
age := int(time.Since(mtime).Hours() / 24)
if age > 30 {
return age
}
return 0
}
func hasRuleArtifacts(ruleDirs []string) bool {
_, found := newestRuleArtifact(ruleDirs)
return found
}
type cPanelModSecVendor struct {
Path string `json:"path"`
VendorID string `json:"vendor_id"`
}
type cPanelModSecVendorsResponse struct {
Data struct {
Vendors []cPanelModSecVendor `json:"vendors"`
} `json:"data"`
Metadata struct {
Result int `json:"result"`
} `json:"metadata"`
}
type cPanelModSecConfig struct {
Active int `json:"active"`
VendorID string `json:"vendor_id"`
}
type cPanelModSecConfigsResponse struct {
Data struct {
Configs []cPanelModSecConfig `json:"configs"`
} `json:"data"`
Metadata struct {
Result int `json:"result"`
} `json:"metadata"`
}
type cPanelActiveRuleState struct {
present bool
vendorDirs []string
}
// cPanelActiveRules maps configuration files that WHM reports as active back
// to their vendor trees. Vendor-wide enabled state is insufficient because
// cPanel permits individual files to override it in either direction. The
// boolean is false when WHM cannot provide an authoritative mapping, in which
// case callers retain conservative filesystem behavior.
func cPanelActiveRules() (cPanelActiveRuleState, bool) {
out, err := runCmd("whmapi1", "modsec_get_configs", "--output=json")
if err != nil || len(out) == 0 {
return cPanelActiveRuleState{}, false
}
var configsResponse cPanelModSecConfigsResponse
if unmarshalErr := json.Unmarshal(out, &configsResponse); unmarshalErr != nil || configsResponse.Metadata.Result != 1 {
return cPanelActiveRuleState{}, false
}
activeVendorIDs := make(map[string]struct{})
activeRules := cPanelActiveRuleState{}
for _, config := range configsResponse.Data.Configs {
if config.Active != 1 {
continue
}
activeRules.present = true
if config.VendorID != "" {
activeVendorIDs[config.VendorID] = struct{}{}
}
}
if len(activeVendorIDs) == 0 {
return activeRules, true
}
out, err = runCmd("whmapi1", "modsec_get_vendors", "--output=json")
if err != nil || len(out) == 0 {
return cPanelActiveRuleState{}, false
}
var vendorsResponse cPanelModSecVendorsResponse
if unmarshalErr := json.Unmarshal(out, &vendorsResponse); unmarshalErr != nil || vendorsResponse.Metadata.Result != 1 {
return cPanelActiveRuleState{}, false
}
seen := make(map[string]struct{})
for _, vendor := range vendorsResponse.Data.Vendors {
if _, active := activeVendorIDs[vendor.VendorID]; !active {
continue
}
if vendor.Path == "" || !filepath.IsAbs(vendor.Path) {
return cPanelActiveRuleState{}, false
}
path := filepath.Clean(vendor.Path)
if path == string(filepath.Separator) {
return cPanelActiveRuleState{}, false
}
if _, exists := seen[path]; exists {
delete(activeVendorIDs, vendor.VendorID)
continue
}
seen[path] = struct{}{}
delete(activeVendorIDs, vendor.VendorID)
activeRules.vendorDirs = append(activeRules.vendorDirs, path)
}
if len(activeVendorIDs) != 0 {
return cPanelActiveRuleState{}, false
}
return activeRules, true
}
// oldestRulesetRefresh treats each directory as one ruleset. Files within a
// ruleset are not all rewritten by every update, so its newest artifact is the
// useful refresh signal. Every loaded ruleset must stay current, making the
// oldest of those per-directory refresh times the value to check.
func oldestRulesetRefresh(ruleDirs []string) (time.Time, bool) {
var oldest time.Time
found := false
for _, dir := range ruleDirs {
newest, ok := newestRuleArtifact([]string{dir})
if !ok {
continue
}
if !found || newest.Before(oldest) {
oldest = newest
found = true
}
}
return oldest, found
}
func newestRuleArtifact(ruleDirs []string) (time.Time, bool) {
return ruleArtifactMtime(ruleDirs, func(candidate, current time.Time) bool {
return candidate.After(current)
})
}
func oldestRuleArtifact(ruleDirs []string) (time.Time, bool) {
return ruleArtifactMtime(ruleDirs, func(candidate, current time.Time) bool {
return candidate.Before(current)
})
}
func ruleArtifactMtime(ruleDirs []string, replace func(time.Time, time.Time) bool) (time.Time, bool) {
var selected time.Time
found := false
pending := append([]string(nil), ruleDirs...)
visited := make(map[string]struct{})
for len(pending) > 0 {
dir := pending[len(pending)-1]
pending = pending[:len(pending)-1]
dir = filepath.Clean(dir)
if _, seen := visited[dir]; seen {
continue
}
visited[dir] = struct{}{}
entries, err := osFS.ReadDir(dir)
if err != nil {
continue
}
for _, entry := range entries {
path := filepath.Join(dir, entry.Name())
if entry.IsDir() {
pending = append(pending, path)
continue
}
if !isRuleArtifact(entry.Name()) {
continue
}
info, err := entry.Info()
if err != nil {
continue
}
if !found || replace(info.ModTime(), selected) {
selected = info.ModTime()
found = true
}
}
}
return selected, found
}
// isRuleArtifact reports whether a filename looks like a ModSecurity rule
// or data artifact (.conf, .data, .rules) so unrelated files like README
// or LICENSE don't dominate the age calculation.
func isRuleArtifact(name string) bool {
name = strings.ToLower(name)
return strings.HasSuffix(name, ".conf") ||
strings.HasSuffix(name, ".data") ||
strings.HasSuffix(name, ".rules")
}
// Markers delimiting the CSM-managed section inside modsec2.user.conf.
// That file is shared with operator-maintained rules (Host-scoped
// ctl:ruleRemoveById exclusions and the like), so CSM may only ever
// rewrite the bytes between these two lines. vpLegacyMarker is the header
// comment of the rules file itself, which is all that pre-delimiter CSM
// versions wrote; it locates those deployments for upgrade.
const (
vpBeginMarker = "# BEGIN CSM Custom ModSecurity Rules (managed by CSM - do not edit inside this block)"
vpEndMarker = "# END CSM Custom ModSecurity Rules"
vpLegacyMarker = "# CSM Custom ModSecurity Rules"
vpOverridesIncludeMarker = "# CSM overrides - managed by CSM rule management"
)
// deployVirtualPatches ensures CSM's custom ModSec rules are installed.
// These provide virtual patches for known WordPress CVEs.
//
// The destination is shared with operator rules, so CSM only ever creates
// or rewrites its own marker-delimited section; every byte outside the
// section is preserved verbatim.
func deployVirtualPatches() {
srcPath := "/opt/csm/configs/csm_modsec_custom.conf"
srcData, err := osFS.ReadFile(srcPath)
if err != nil {
return // no custom rules to deploy
}
section := buildVPSection(srcData)
for _, dest := range vpDestPaths {
dir := filepath.Dir(dest)
if _, err := osFS.Stat(dir); os.IsNotExist(err) {
continue
}
existing, err := osFS.ReadFile(dest)
if err != nil && !os.IsNotExist(err) {
// Present but unreadable: rewriting blind could destroy
// operator rules, so leave this candidate alone.
continue
}
merged, upToDate := mergeVPSection(existing, section)
if upToDate {
return
}
// #nosec G306 -- WAF rule file read by Apache/nginx as a different user.
if err := osFS.WriteFile(dest, merged, 0644); err != nil {
continue
}
fmt.Fprintf(os.Stderr, "[%s] Virtual patches deployed to %s\n",
time.Now().Format("2006-01-02 15:04:05"), dest)
return
}
}
// MergeModSecUserConfSection merges CSM's ModSecurity rules payload into
// the current contents of a modsec2.user.conf, confining CSM to its
// marker-delimited section so operator-maintained rules outside the
// section survive every deploy. It is exported because three call sites
// write this file (the WAF check cycle here, `csm install`, and the
// daemon startup config deploy); routing them all through one merge
// guarantees no caller ever whole-file-overwrites operator rules.
//
// existing is the current file contents (nil for a missing file). merged
// is only meaningful when changed is true; changed=false means the file
// already carries the wanted section and must not be rewritten.
func MergeModSecUserConfSection(existing, srcData []byte) (merged []byte, changed bool) {
merged, upToDate := mergeVPSection(existing, buildVPSection(srcData))
return merged, !upToDate
}
// RemoveModSecUserConfSections removes only content owned by CSM from the
// shared ModSecurity user configuration. Operator bytes outside the exact CSM
// marker lines are preserved.
func RemoveModSecUserConfSections(existing []byte) (cleaned []byte, changed bool) {
cleaned = existing
for {
if begin, _, ok := markerLineBounds(cleaned, vpBeginMarker); ok {
end := vpSectionEnd(cleaned[begin:])
if end < 0 {
end = vpBlockEndBeforePreservedTail(cleaned, begin) - begin
}
cleaned = removeByteRange(cleaned, begin, begin+end)
changed = true
continue
}
if legacy, _, ok := markerLineBounds(cleaned, vpLegacyMarker); ok {
end := vpBlockEndBeforePreservedTail(cleaned, legacy)
cleaned = removeByteRange(cleaned, legacy, end)
changed = true
continue
}
break
}
for {
markerStart, markerEnd, ok := markerLineBounds(cleaned, vpOverridesIncludeMarker)
if !ok {
break
}
removeEnd := markerEnd
if includeEnd := nextLineEnd(cleaned, markerEnd); includeEnd > markerEnd {
line := bytes.TrimSpace(cleaned[markerEnd:includeEnd])
if bytes.HasPrefix(line, []byte("Include ")) && bytes.Contains(line, []byte("modsec2.csm-overrides.conf")) {
removeEnd = includeEnd
}
}
cleaned = removeByteRange(cleaned, markerStart, removeEnd)
changed = true
}
return cleaned, changed
}
func removeByteRange(data []byte, start, end int) []byte {
out := make([]byte, 0, len(data)-(end-start))
out = append(out, data[:start]...)
return append(out, data[end:]...)
}
func nextLineEnd(data []byte, start int) int {
if start >= len(data) {
return start
}
if offset := bytes.IndexByte(data[start:], '\n'); offset >= 0 {
return start + offset + 1
}
return len(data)
}
// buildVPSection wraps the rules payload in the begin/end marker lines.
// The result is deterministic for a given payload so later cycles can
// recognize an up-to-date section by byte comparison.
func buildVPSection(srcData []byte) []byte {
section := make([]byte, 0, len(vpBeginMarker)+len(srcData)+len(vpEndMarker)+3)
section = append(section, vpBeginMarker...)
section = append(section, '\n')
section = append(section, srcData...)
if len(srcData) > 0 && srcData[len(srcData)-1] != '\n' {
section = append(section, '\n')
}
section = append(section, vpEndMarker...)
section = append(section, '\n')
return section
}
// mergeVPSection computes the new content for a modsec user conf so that
// it carries exactly one copy of the CSM section while every byte outside
// the section stays untouched. upToDate reports that the file already
// holds the wanted section and no write is needed.
func mergeVPSection(existing, section []byte) (merged []byte, upToDate bool) {
if len(existing) == 0 {
return section, false
}
if begin, _, ok := markerLineBounds(existing, vpBeginMarker); ok {
if end := vpSectionEnd(existing[begin:]); end >= 0 {
if bytes.Equal(existing[begin:begin+end], section) {
return nil, true
}
merged = append(merged, existing[:begin]...)
merged = append(merged, section...)
merged = append(merged, existing[begin+end:]...)
return merged, false
}
// A begin marker without an end marker is a malformed CSM block.
// Replace from the exact begin line so the next cycle is delimited
// again instead of falling through to the legacy header inside it.
blockEnd := vpBlockEndBeforePreservedTail(existing, begin)
merged = append(merged, existing[:begin]...)
merged = append(merged, section...)
merged = append(merged, existing[blockEnd:]...)
return merged, false
}
if legacy, _, ok := markerLineBounds(existing, vpLegacyMarker); ok {
// Pre-delimiter CSM versions appended the raw rules file, so the
// CSM content starts at this exact header line. Installer and
// daemon deploys appended the overrides Include after the raw
// rules, so preserve that tail when it is present.
legacyEnd := vpBlockEndBeforePreservedTail(existing, legacy)
merged = append(merged, existing[:legacy]...)
merged = append(merged, section...)
merged = append(merged, existing[legacyEnd:]...)
return merged, false
}
// Operator-only file: append the section, separated by one blank
// line, keeping the existing bytes exactly as they are.
merged = append(merged, existing...)
if existing[len(existing)-1] != '\n' {
merged = append(merged, '\n')
}
merged = append(merged, '\n')
merged = append(merged, section...)
return merged, false
}
func vpBlockEndBeforePreservedTail(existing []byte, blockStart int) int {
blockEnd := len(existing)
if tail, _, ok := markerLineBounds(existing[blockStart:], vpOverridesIncludeMarker); ok && tail > 0 {
blockEnd = blockStart + tail
if existing[blockEnd-1] == '\n' {
blockEnd--
}
}
return blockEnd
}
// markerLineBounds returns the start offset and end offset (including the
// trailing newline when present) for an exact marker line. Matching the
// whole line keeps operator comments that merely mention marker text from
// being treated as CSM-owned content.
func markerLineBounds(data []byte, marker string) (start, end int, ok bool) {
markerBytes := []byte(marker)
for start < len(data) {
lineEnd := bytes.IndexByte(data[start:], '\n')
end = len(data)
next := len(data)
if lineEnd >= 0 {
end = start + lineEnd
next = end + 1
}
line := data[start:end]
if len(line) > 0 && line[len(line)-1] == '\r' {
line = line[:len(line)-1]
}
if bytes.Equal(line, markerBytes) {
return start, next, true
}
start = next
}
return 0, 0, false
}
// vpSectionEnd returns the offset just past the end-marker line (its
// trailing newline included when present), or -1 when there is no end
// marker. data must start at the section's begin line.
func vpSectionEnd(data []byte) int {
if _, end, ok := markerLineBounds(data, vpEndMarker); ok {
return end
}
return -1
}
// autoUpdateWAFRules triggers ModSecurity vendor rule updates via whmapi1.
// Returns true if an update was successfully triggered.
func autoUpdateWAFRules() bool {
// Get installed vendors
out, err := runCmd("whmapi1", "modsec_get_vendors")
if err != nil || out == nil {
return false
}
// Parse vendor IDs from output (look for "vendor_id:" lines)
var vendors []string
for _, line := range strings.Split(string(out), "\n") {
line = strings.TrimSpace(line)
if strings.HasPrefix(line, "vendor_id:") || strings.HasPrefix(line, "id:") {
parts := strings.SplitN(line, ":", 2)
if len(parts) == 2 {
vid := strings.TrimSpace(parts[1])
if vid != "" {
vendors = append(vendors, vid)
}
}
}
}
if len(vendors) == 0 {
return false
}
// Update each vendor
updated := false
for _, vid := range vendors {
out, err := runCmd("whmapi1", "modsec_update_vendor", fmt.Sprintf("vendor_id=%s", vid))
if err == nil && out != nil && strings.Contains(string(out), "result: 1") {
fmt.Fprintf(os.Stderr, "[%s] WAF auto-update: vendor %s updated successfully\n",
time.Now().Format("2006-01-02 15:04:05"), vid)
updated = true
}
}
return updated
}
// CheckModSecAuditLog parses the ModSecurity audit log for blocked attacks.
// High-volume attackers are reported for potential auto-blocking.
func CheckModSecAuditLog(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
var findings []alert.Finding
logPaths := modsecAuditLogPaths()
if len(logPaths) == 0 {
return nil
}
var lines []string
for _, path := range logPaths {
lines = tailFile(path, modsecAuditTailLines)
if len(lines) > 0 {
break
}
}
if len(lines) == 0 {
return nil
}
// Count blocked attacks per IP
blocked := countModSecDenials(lines)
for ip := range blocked {
if isInfraIP(ip, cfg.InfraIPs) || !wafAttackerIsReportable(ip) {
delete(blocked, ip)
}
}
// Alert on high-volume attackers (auto-block integration via check name)
for ip, count := range blocked {
if count >= 20 {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "waf_attack_blocked",
SourceIP: ip,
Message: fmt.Sprintf("WAF blocking high-volume attacker: %s (%d blocked requests)", ip, count),
Details: wafBlockAdvice(ip, count),
})
}
}
return findings
}
// wafAttackerIsReportable reports whether a ModSecurity denial count belongs
// to an address worth telling the operator about.
//
// The control panel proxies its own traffic over loopback and, on cPanel,
// through the machine's public address rather than 127.0.0.1, so denials
// attributed to either accumulate on any busy host. Reporting those as a
// high-volume attacker and advising a permanent block points the operator at
// their own machine -- and the firewall's local-address guard excludes
// loopback, so the block is accepted rather than refused.
//
// A lookup failure fails open: a real attacker must never be suppressed by a
// transient syscall error. Loopback is decided without the lookup.
func wafAttackerIsReportable(ip string) bool {
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
if parsed.IsLoopback() || parsed.IsUnspecified() {
return false
}
return !netutil.IsHostAddress(ip)
}
// wafBlockAdvice is the operator guidance attached to a WAF attacker finding.
func wafBlockAdvice(ip string, count int) string {
details := fmt.Sprintf("IP %s has been blocked %d times by ModSecurity.", ip, count)
// The firewall refuses link-local blocks, so advising one sends the
// operator to an action that cannot succeed.
if parsed := net.ParseIP(ip); parsed.IsLinkLocalUnicast() || parsed.IsLinkLocalMulticast() {
return details + " Review the source of this link-local traffic; CSM does not block link-local addresses."
}
return details + " Consider permanent block via CSM."
}
// modsecAuditLogPaths yields the audit log candidates; a seam for tests.
var modsecAuditLogPaths = func() []string { return platform.Detect().ModSecAuditLogPaths }
// modsecAuditTailLines is how much of the audit log one cycle inspects.
// The serial format spends ten or more lines per transaction, so 200 lines
// could never hold the 20 denials the threshold asks for.
const modsecAuditTailLines = 4000
// countModSecDenials counts denied transactions per client address.
//
// In the serial audit format one transaction spans lettered sections
// between "--<id>-A--" and "--<id>-Z--"; the client address is on the A
// header and the denial message in H, so the address is carried across the
// transaction and each denied transaction counts once. Lines outside a
// transaction (concurrent format summaries, error-log style lines) keep
// the per-line rule: a denial marker plus an address on the same line.
func countModSecDenials(lines []string) map[string]int {
blocked := make(map[string]int)
var (
txID string
txIP string
txDenied bool
wantA bool
)
flush := func() {
if txID != "" && txIP != "" && txDenied {
blocked[txIP]++
}
txID, txIP, txDenied, wantA = "", "", false, false
}
for _, line := range lines {
if id, section, ok := modsecSectionBoundary(line); ok {
switch section {
case 'A':
flush()
txID, wantA = id, true
case 'Z':
if id == txID {
flush()
}
}
continue
}
if wantA {
if strings.TrimSpace(line) == "" {
continue
}
wantA = false
txIP = modsecAuditClientIP(line)
continue
}
if txID != "" {
if modsecDenialLine(line) {
txDenied = true
}
continue
}
if modsecDenialLine(line) {
if ip := extractIPFromLog(line); ip != "" {
blocked[ip]++
}
}
}
flush()
return blocked
}
func modsecDenialLine(line string) bool {
return strings.Contains(line, "403") || strings.Contains(line, "Access denied") ||
strings.Contains(line, "MODSEC") || strings.Contains(line, "mod_security")
}
// modsecSectionBoundary parses a serial-format "--<id>-<letter>--" line.
func modsecSectionBoundary(line string) (id string, section byte, ok bool) {
line = strings.TrimSpace(line)
if len(line) < 7 || !strings.HasPrefix(line, "--") || !strings.HasSuffix(line, "--") {
return "", 0, false
}
body := line[2 : len(line)-2]
dash := strings.LastIndexByte(body, '-')
if dash <= 0 || dash != len(body)-2 {
return "", 0, false
}
section = body[dash+1]
if section < 'A' || section > 'Z' {
return "", 0, false
}
return body[:dash], section, true
}
// modsecAuditClientIP reads the client address from an A-section header:
// "[timestamp] uniqueid client-ip client-port server-ip server-port".
func modsecAuditClientIP(line string) string {
fields := strings.Fields(line)
if len(fields) < 5 || !strings.HasPrefix(fields[0], "[") {
return ""
}
clientIP := strings.Trim(fields[3], "[]")
ip := net.ParseIP(clientIP)
if ip == nil {
return ""
}
return ip.String()
}
package checks
import (
"fmt"
"path/filepath"
"sort"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/platform"
)
// modsecDisabledScope is one place where ModSecurity is switched off for
// customer traffic. Domain is empty when the scope covers every vhost of
// the account.
type modsecDisabledScope struct {
User string
Domain string
Source string
}
// modsecDisabledScopes reports every vhost or account with ModSecurity
// turned off. A disabled scope voids every CSM virtual patch for that
// traffic and writes nothing to the modsec audit log, so an attack there
// is both unblocked and invisible.
//
// cPanel expresses "off" through several independent mechanisms and an
// audit of any one of them badly understates the exposure, so all of them
// are walked:
//
// - secruleengineoff in /var/cpanel/userdata (per domain)
// - modsec.conf under conf.d/userdata (per account, and per domain)
//
// The conf.d tree is mirrored into std and ssl. Because virtually all
// traffic is HTTPS, a scope present only under ssl still leaves the site
// unfiltered, so both trees carry equal weight here.
func modsecDisabledScopes(info platform.Info) []modsecDisabledScope {
if !info.IsCPanel() {
return nil
}
var scopes []modsecDisabledScope
scopes = append(scopes, userdataDisabledScopes()...)
scopes = append(scopes, confTreeDisabledScopes(info)...)
return dedupeScopes(scopes)
}
// modsecDisabledScopesForFindings reads only the cPanel paths that can affect
// the supplied vulnerable-plugin findings. The normal WAF audit already does
// the exhaustive /var/cpanel/userdata/*/* walk; repeating it for correlation
// makes a deep scan do thousands of redundant reads on a large shared host.
func modsecDisabledScopesForFindings(info platform.Info, findings []alert.Finding, aliases map[string][]string) []modsecDisabledScope {
if !info.IsCPanel() {
return nil
}
domainsByAccount := make(map[string]map[string]bool)
for _, finding := range findings {
account := strings.TrimSpace(finding.TenantID)
if !validAccountName.MatchString(account) {
continue
}
if domainsByAccount[account] == nil {
domainsByAccount[account] = make(map[string]bool)
}
domain := cleanDomlogDomain(finding.Domain)
if domain == "" {
continue
}
domainsByAccount[account][domain] = true
for _, peer := range aliases[aliasKey(account, domain)] {
if peer = cleanDomlogDomain(peer); peer != "" {
domainsByAccount[account][peer] = true
}
}
}
var scopes []modsecDisabledScope
for account, domains := range domainsByAccount {
scopes = append(scopes, targetedUserdataDisabledScopes(account, domains)...)
scopes = append(scopes, targetedConfTreeDisabledScopes(info, account, domains)...)
}
return dedupeScopes(scopes)
}
func targetedUserdataDisabledScopes(account string, domains map[string]bool) []modsecDisabledScope {
var scopes []modsecDisabledScope
for domain := range domains {
for _, name := range []string{domain, domain + "_SSL"} {
path := filepath.Join("/var/cpanel/userdata", account, name)
data, err := osFS.ReadFile(path)
if err != nil || !userdataSecRuleEngineOff(string(data)) {
continue
}
scopes = append(scopes, modsecDisabledScope{User: account, Domain: domain, Source: path})
}
}
return scopes
}
func targetedConfTreeDisabledScopes(info platform.Info, account string, domains map[string]bool) []modsecDisabledScope {
configDir := info.ApacheCompatibleConfigDir()
if configDir == "" {
return nil
}
var scopes []modsecDisabledScope
for _, base := range userdataTreeBases(configDir) {
accountPath := filepath.Join(base, account, "modsec.conf")
if confSecRuleEngineOff(accountPath) {
scopes = append(scopes, modsecDisabledScope{User: account, Source: accountPath})
}
for domain := range domains {
path := filepath.Join(base, account, domain, "modsec.conf")
if confSecRuleEngineOff(path) {
scopes = append(scopes, modsecDisabledScope{User: account, Domain: domain, Source: path})
}
}
}
return scopes
}
// userdataDisabledScopes walks the per-domain userdata flag. cPanel keeps
// <domain>, <domain>_SSL and <domain>.cache copies of the same record, so
// the domain name is normalised and duplicates collapse later.
func userdataDisabledScopes() []modsecDisabledScope {
var scopes []modsecDisabledScope
for _, path := range globPaths("/var/cpanel/userdata/*/*") {
domain, ok := cpanelUserdataDomain(filepath.Base(path))
if !ok {
continue
}
data, err := osFS.ReadFile(path)
if err != nil || !userdataSecRuleEngineOff(string(data)) {
continue
}
scopes = append(scopes, modsecDisabledScope{
User: filepath.Base(filepath.Dir(path)),
Domain: domain,
Source: path,
})
}
return scopes
}
// cpanelUserdataDomain accepts only per-vhost records. The same directory
// contains main/scope metadata and generated JSON/cache files; rejecting them
// by shape before ReadFile avoids bogus scopes and large cache reads.
func cpanelUserdataDomain(name string) (string, bool) {
if strings.HasSuffix(name, ".cache") || strings.HasSuffix(name, ".json") {
return "", false
}
domain := cleanDomlogDomain(strings.TrimSuffix(name, "_SSL"))
return domain, domain != ""
}
// confTreeDisabledScopes walks the Apache userdata include tree, where a
// modsec.conf directly under the account directory disables every vhost
// the account owns and one a level deeper disables a single domain.
func confTreeDisabledScopes(info platform.Info) []modsecDisabledScope {
configDir := info.ApacheCompatibleConfigDir()
if configDir == "" {
return nil
}
var scopes []modsecDisabledScope
for _, base := range userdataTreeBases(configDir) {
for _, path := range globPaths(filepath.Join(base, "*", "modsec.conf")) {
if !confSecRuleEngineOff(path) {
continue
}
scopes = append(scopes, modsecDisabledScope{
User: filepath.Base(filepath.Dir(path)),
Source: path,
})
}
for _, path := range globPaths(filepath.Join(base, "*", "*", "modsec.conf")) {
if !confSecRuleEngineOff(path) {
continue
}
domainDir := filepath.Dir(path)
scopes = append(scopes, modsecDisabledScope{
User: filepath.Base(filepath.Dir(domainDir)),
Domain: filepath.Base(domainDir),
Source: path,
})
}
}
return scopes
}
// userdataTreeBases returns the per-vhost include roots to walk.
//
// cPanel reports its config dir as /usr/local/apache/conf, and the
// userdata tree hangs directly off it. The distro-style path for the same
// directory is /etc/apache2/conf.d/userdata, so "conf.d" belongs to that
// spelling of the path, not under the cPanel config dir. Both layouts are
// walked because either spelling can be the one the platform layer
// reports; duplicate hits collapse during deduplication.
func userdataTreeBases(configDir string) []string {
var bases []string
for _, prefix := range [][]string{{"userdata"}, {"conf.d", "userdata"}} {
for _, tree := range []string{"std", "ssl"} {
parts := append(append([]string{configDir}, prefix...), tree, "2_4")
bases = append(bases, filepath.Join(parts...))
}
}
return bases
}
func confSecRuleEngineOff(path string) bool {
data, err := osFS.ReadFile(path)
if err != nil {
return false
}
return secRuleEngineOffDirective(string(data))
}
// dedupeScopes collapses the several files that can describe one scope
// into a single entry, keeping the lowest-sorting source so the reported
// path is stable across runs.
func dedupeScopes(scopes []modsecDisabledScope) []modsecDisabledScope {
sort.Slice(scopes, func(i, j int) bool {
if scopes[i].User != scopes[j].User {
return scopes[i].User < scopes[j].User
}
if scopes[i].Domain != scopes[j].Domain {
return scopes[i].Domain < scopes[j].Domain
}
return scopes[i].Source < scopes[j].Source
})
var out []modsecDisabledScope
seen := make(map[string]bool, len(scopes))
for _, s := range scopes {
key := s.User + "\x00" + s.Domain
if seen[key] {
continue
}
seen[key] = true
out = append(out, s)
}
return out
}
// globPaths expands a pattern, discarding the error: a malformed pattern
// is a programming bug and an unreadable directory simply yields nothing.
func globPaths(pattern string) []string {
matches, err := osFS.Glob(pattern)
if err != nil {
return nil
}
return matches
}
// userdataSecRuleEngineOff reports whether a cPanel userdata file carries
// the per-domain "ModSecurity off" flag. cPanel writes secruleengineoff: 0
// for the enabled state, so presence of the key is not enough.
func userdataSecRuleEngineOff(contents string) bool {
for _, line := range strings.Split(contents, "\n") {
key, value, found := strings.Cut(strings.TrimSpace(line), ":")
if !found || strings.TrimSpace(key) != "secruleengineoff" {
continue
}
if strings.TrimSpace(value) == "1" {
return true
}
}
return false
}
// secRuleEngineOffDirective reports whether an Apache config fragment
// turns the engine off. Files that only carry SecRuleRemoveById lines are
// legitimate false-positive workarounds and leave the engine running.
func secRuleEngineOffDirective(contents string) bool {
for _, line := range strings.Split(contents, "\n") {
fields := strings.Fields(line)
if len(fields) < 2 || strings.HasPrefix(fields[0], "#") {
continue
}
if strings.EqualFold(fields[0], "SecRuleEngine") && strings.EqualFold(fields[1], "Off") {
return true
}
}
return false
}
// modsecDisabledDetailCap bounds how many scopes are named in the finding
// details. The count in the message always reflects the true total.
const modsecDisabledDetailCap = 20
// modsecDisabledFindings reports disabled ModSecurity scopes as a single
// aggregated finding. Hosts routinely accumulate hundreds of these over
// the years, so one finding per scope would drown every other alert.
func modsecDisabledFindings(info platform.Info) []alert.Finding {
scopes := modsecDisabledScopes(info)
if len(scopes) == 0 {
return nil
}
var accounts, domains int
var b strings.Builder
for i, s := range scopes {
if s.Domain == "" {
accounts++
} else {
domains++
}
if i < modsecDisabledDetailCap {
target := s.Domain
if target == "" {
target = "account " + s.User + " (all domains)"
}
b.WriteString(" " + target + " -- " + s.Source + "\n")
}
}
if len(scopes) > modsecDisabledDetailCap {
fmt.Fprintf(&b, " ... and %d more\n", len(scopes)-modsecDisabledDetailCap)
}
return []alert.Finding{{
Severity: alert.High,
Check: "modsec_disabled_vhost",
Message: fmt.Sprintf("ModSecurity is disabled for %d scopes (%d account-wide, %d single-domain)",
len(scopes), accounts, domains),
Details: "Every CSM virtual patch is inert on this traffic and no modsec audit\n" +
"record is written, so attacks there are both unblocked and invisible.\n\n" +
b.String() +
"\nRe-enable per domain with:\n" +
" /usr/local/cpanel/bin/modsecuritydomains enable --domain=<domain>\n" +
"That tool only rewrites the userdata flag. Account-wide and per-domain\n" +
"modsec.conf files must be edited directly in BOTH the std and ssl\n" +
"userdata trees, then rebuild the web server config. Keep any\n" +
"SecRuleRemoveById lines already present -- those are rule exclusions,\n" +
"not an engine switch.",
}}
}
package checks
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"os"
"strings"
"sync"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/modsec"
"github.com/pidginhost/csm/internal/store"
)
// ErrModSecReloadNotConfigured reports a CSM rule section that differs from
// the last one the web server reloaded, on a host with no reload command.
var ErrModSecReloadNotConfigured = errors.New("modsec.reload_command is not set, so CSM cannot confirm its updated ModSecurity rules are active; they apply at the next web server reload")
// modsecActiveSectionKey holds the hash of the CSM section on disk at the last
// successful reload. Comparing against it, rather than reloading when a write
// happens, also activates sections written by the installer or an upgrade.
const modsecActiveSectionKey = "modsec:active_section_sha256"
var modsecReloadRunner = modsec.Reload
// vpDestPaths are the ModSecurity user configuration files CSM writes its
// section into, in order of preference.
var vpDestPaths = []string{
"/etc/apache2/conf.d/modsec/modsec2.user.conf",
"/usr/local/apache/conf/modsec2.user.conf",
}
// ModSecReloadReconciler is owned by the daemon and shared by startup and all
// its scans. A CLI opening the store must not implicitly gain reload authority.
type ModSecReloadReconciler struct {
mu sync.Mutex
db *store.DB
active string
pending bool // a successful reload whose metadata write needs retrying
}
type modsecReloadContextKey struct{}
// WithModSecReload lets daemon scans share their startup reconciler. Contexts
// without one can still deploy rules, but leave activation to the daemon.
func WithModSecReload(ctx context.Context, r *ModSecReloadReconciler) context.Context {
return context.WithValue(ctx, modsecReloadContextKey{}, r)
}
func deployAndReconcileModSec(ctx context.Context, command string) error {
r, _ := ctx.Value(modsecReloadContextKey{}).(*ModSecReloadReconciler)
if r == nil {
deployVirtualPatches()
return nil
}
// Serialize the write as well as the reload: a concurrent scan must not
// truncate the configuration while the web server is reading it.
r.mu.Lock()
defer r.mu.Unlock()
deployVirtualPatches()
if strings.TrimSpace(command) == "" {
return nil // only startup warns about an unset command
}
if err := ctx.Err(); err != nil {
return err
}
return r.reconcile(command)
}
// Reconcile activates sections already written by startup or the installer.
// Only failed reloads are repeated; metadata failures retry just the write.
func (r *ModSecReloadReconciler) Reconcile(reloadCommand string) error {
r.mu.Lock()
defer r.mu.Unlock()
return r.reconcile(reloadCommand)
}
func (r *ModSecReloadReconciler) reconcile(reloadCommand string) error {
db := store.Global()
if db == nil {
return nil
}
r.db = db
digest, err := installedVPSectionDigest()
if err != nil {
return err
}
if digest == "" {
return nil
}
if r.active == "" {
active, err := db.ReadMetaString(modsecActiveSectionKey)
if err != nil {
return fmt.Errorf("read active CSM ModSecurity rules: %w", err)
}
r.active = active
}
if r.active == digest {
return r.persistActive()
}
if strings.TrimSpace(reloadCommand) == "" {
return ErrModSecReloadNotConfigured
}
if _, err := modsecReloadRunner(reloadCommand); err != nil {
return fmt.Errorf("web server reload for CSM ModSecurity rules: %w", err)
}
r.active, r.pending = digest, true
csmlog.Info("web server reloaded to activate CSM ModSecurity rules", "command", reloadCommand)
return r.persistActive()
}
func (r *ModSecReloadReconciler) persistActive() error {
if !r.pending {
return nil
}
if err := r.db.SetMetaString(modsecActiveSectionKey, r.active); err != nil {
return fmt.Errorf("web server reloaded, but recording active CSM ModSecurity rules failed: %w", err)
}
r.pending = false
return nil
}
// Deployment can fall back when the preferred file cannot be written. Track
// every installed section so an older preferred copy cannot hide that update.
func installedVPSectionDigest() (string, error) {
var sums []byte
var readErr error
for _, dest := range vpDestPaths {
data, err := osFS.ReadFile(dest)
if err != nil {
if !os.IsNotExist(err) {
readErr = errors.Join(readErr, fmt.Errorf("read CSM ModSecurity rules in %s: %w", dest, err))
}
continue
}
begin, _, ok := markerLineBounds(data, vpBeginMarker)
if !ok {
continue
}
end := vpSectionEnd(data[begin:])
if end < 0 {
readErr = errors.Join(readErr, fmt.Errorf("unterminated CSM ModSecurity section in %s", dest))
continue
}
sum := sha256.Sum256(data[begin : begin+end])
sums = append(sums, sum[:]...)
}
if len(sums) == 0 {
return "", readErr
}
// The usual single section is identified by its own hash.
if len(sums) == sha256.Size {
return hex.EncodeToString(sums), nil
}
sum := sha256.Sum256(sums)
return hex.EncodeToString(sum[:]), nil
}
package checks
import (
"bufio"
"context"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"sort"
"strings"
"sync"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
const wpChecksumWorkers = 5 // concurrent wp core verify-checksums
// htaccessMaxLineBytes bounds a single .htaccess line for the token scanner.
// Legitimate directives are far shorter; a line past this is itself an
// anomaly, so the scanner fails closed (flags the file) rather than silently
// truncating the rest.
const htaccessMaxLineBytes = 1 << 20 // 1 MiB
// CheckHtaccess scans for malicious .htaccess directives using pure Go ReadDir.
func CheckHtaccess(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
suspiciousPatterns := []string{
"auto_prepend_file",
"auto_append_file",
"eval(",
"base64_decode",
"gzinflate",
"str_rot13",
"php_value disable_functions",
"addhandler",
"addtype",
"sethandler",
}
safePatterns := []string{
"wordfence-waf.php",
"litespeed",
"advanced-headers.php",
"rsssl",
// Standard handler directives for PHP/static files are safe
"application/x-httpd-php",
"application/x-httpd-php5",
"application/x-httpd-ea-php",
"application/x-httpd-alt-php",
"text/html",
"text/css",
"text/javascript",
"application/javascript",
"image/",
"font/",
"proxy:unix",
// Security plugins that use handler directives to BLOCK execution
"-execcgi", // Options -ExecCGI disables CGI (Wordfence pattern)
"sethandler none", // Disables all handlers (security measure)
"sethandler default-handler", // Resets to default (security measure)
// Legitimate MIME type additions
"application/font",
"application/vnd",
".woff",
".woff2",
".ttf",
".eot",
".svg",
// Wordfence code execution protection
"wordfence",
}
// Scan each user's document roots
homeDirs := scanHomeDirsWithCoverage(ctx, "htaccess")
for _, homeEntry := range homeDirs {
if ctx.Err() != nil {
return findings
}
if !homeEntry.IsDir() {
continue
}
homeDir := scanHomeDirPath(homeEntry)
docRoot := filepath.Join(homeDir, "public_html")
scanHtaccess(ctx, docRoot, htaccessScanMaxDepth, suspiciousPatterns, safePatterns, cfg, &findings)
// Also check addon domains
subDirs, err := osFS.ReadDir(homeDir)
markScanReadError(ctx, "htaccess", err)
for _, sd := range subDirs {
if sd.IsDir() && sd.Name() != "public_html" && sd.Name() != "mail" &&
!strings.HasPrefix(sd.Name(), ".") && sd.Name() != "etc" &&
sd.Name() != "logs" && sd.Name() != "ssl" && sd.Name() != "tmp" {
scanHtaccess(ctx, filepath.Join(homeDir, sd.Name()), htaccessScanMaxDepth, suspiciousPatterns, safePatterns, cfg, &findings)
}
}
}
return findings
}
// htaccessScanMaxDepth is how deep below a document root the scheduled
// .htaccess scan walks. Five levels stopped short of every uploads tree
// (wp-content/uploads/YYYY/MM/<dir> is already five), which is where a
// dropper plants the handler-enabling .htaccess that makes its .jpg run.
// Matched to the rolling content scan's depth so both see the same tree.
const htaccessScanMaxDepth = rollingWalkMaxDepth
func scanHtaccess(ctx context.Context, dir string, maxDepth int, suspicious, safe []string, cfg *config.Config, findings *[]alert.Finding) {
if ctx.Err() != nil {
return
}
if maxDepth <= 0 {
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
markScanReadError(ctx, "htaccess", err)
return
}
for _, entry := range entries {
if ctx.Err() != nil {
return
}
name := entry.Name()
fullPath := filepath.Join(dir, name)
if entry.IsDir() {
scanHtaccess(ctx, fullPath, maxDepth-1, suspicious, safe, cfg, findings)
continue
}
if name != ".htaccess" {
continue
}
// Skip suppressed paths (bypassed for explicit full-scan / audit requests).
suppressed := false
if scanRespectsIgnores(ctx, cfg) {
for _, ignore := range cfg.Suppressions.IgnorePaths {
if matchGlob(fullPath, ignore) {
suppressed = true
break
}
}
}
if suppressed {
continue
}
checkHtaccessFile(ctx, fullPath, suspicious, safe, findings)
// Run the hardened detector registry alongside the generic
// token scanner so per-pattern findings emit with their own
// names (htaccess_php_in_uploads, htaccess_filesmatch_shield,
// etc.) rather than collapsing into the catch-all categories.
// The two scans can both fire on the same line; downstream
// dedup at alert.Dispatch handles same-key dups.
hardenedFindings, _, complete := auditHtaccessFile(fullPath)
if !complete {
markCheckIncomplete(ctx, "htaccess")
}
*findings = append(*findings, hardenedFindings...)
}
}
// phpExtension reports whether ext (leading dot, lowercase) is a stock
// PHP-executed extension that legitimately maps to a PHP handler.
func phpExtension(ext string) bool {
return isExecutablePHPName("x" + ext)
}
// phpHandlerRemapsNonPHP reports whether a single lowercased .htaccess line is
// an AddHandler/AddType/SetHandler directive routing a non-PHP file extension
// to a PHP execution handler. cPanel MultiPHP legitimately maps PHP extensions
// to a PHP handler (e.g. application/x-httpd-ea-php74___lsphp .php .php7), so we
// flag only when at least one mapped extension is not PHP-family.
func phpHandlerRemapsNonPHP(lineLower string) bool {
return phpHandlerRemapsNonPHPInContext(lineLower, nil)
}
func phpHandlerRemapsNonPHPInContext(lineLower string, contexts []phpHandlerOverlay) bool {
fields := apacheDirectiveFields(lineLower)
if len(fields) < 2 {
return false
}
directive := fields[0]
switch directive {
case "addhandler", "addtype", "sethandler", "forcetype":
default:
return false
}
handler := fields[1]
if !handlerIsPHP(handler) {
return false
}
exts := normalizedExts(fields[2:])
if len(exts) == 0 {
return directiveHandlerContextTargetsNonPHP(directive, contexts)
}
for _, ext := range exts {
if !phpExtension(ext) {
return true
}
}
return false
}
func directiveHandlerContextTargetsNonPHP(directive string, contexts []phpHandlerOverlay) bool {
if directive != "sethandler" && directive != "forcetype" {
return false
}
for _, ctx := range contexts {
if ctx.unrestricted {
// A name-selecting container hands the PHP handler files that
// no extension list describes.
return true
}
for ext := range ctx.exts {
if !phpExtension(ext) {
return true
}
}
for name := range ctx.names {
if !isExecutablePHPName(name) {
return true
}
}
}
return false
}
func checkHtaccessFile(ctx context.Context, path string, suspicious, safe []string, findings *[]alert.Finding) {
f, err := osFS.Open(path)
if err != nil {
markScanReadError(ctx, "htaccess", err)
return
}
defer func() { _ = f.Close() }()
// The per-line ceiling below bounds one token, not the file: a
// .htaccess of a million short lines still had to be held whole.
if info, statErr := f.Stat(); statErr == nil && info.Size() > htaccessMaxFileBytes {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "htaccess_injection",
Message: fmt.Sprintf(".htaccess too large to audit: %s", path),
Details: fmt.Sprintf("File: %s\nSize: %d bytes exceeds the %d byte ceiling; a real .htaccess is a few kilobytes. Inspect it by hand.", path, info.Size(), htaccessMaxFileBytes),
FilePath: path,
})
return
}
// Read entire file to check cross-line context (e.g., AddHandler + Options -ExecCGI)
var lines []string
scanner := bufio.NewScanner(f)
// A .htaccess line longer than the default 64 KB token would make
// Scan stop early and silently drop every line after it, letting an
// attacker hide a malicious directive behind one padded line. Raise
// the ceiling, and if a line still exceeds it, fail closed below.
scanner.Buffer(make([]byte, 0, 64*1024), htaccessMaxLineBytes)
for scanner.Scan() {
lines = append(lines, scanner.Text())
}
if err := scanner.Err(); err != nil {
markCheckIncomplete(ctx, "htaccess")
// Could not read the whole file (oversized line or I/O error), so
// the cross-line and per-line analysis below is incomplete. Flag
// the file for review rather than reporting a clean partial scan.
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "htaccess_injection",
Message: "Unparseable .htaccess: oversized or unreadable line blocks full analysis",
Details: fmt.Sprintf("File: %s\nError: %v", path, err),
FilePath: path,
})
return
}
// Build full file content for context checks
fullContentLower := strings.ToLower(strings.Join(lines, "\n"))
// If file contains handler directives paired with -ExecCGI, the whole
// block is a security measure (e.g., Wordfence execution protection)
hasExecCGIBlock := strings.Contains(fullContentLower, "-execcgi")
var phpHandlerContexts []phpHandlerOverlay
for _, logical := range joinHtaccessContinuations(lines) {
lineNum := logical.start
trimmed := strings.TrimSpace(logical.text)
lineLower := strings.ToLower(trimmed)
// Skip comments entirely - commented-out directives are not active
if strings.HasPrefix(trimmed, "#") {
continue
}
if ctx, ok := openPHPHandlerContext(trimmed); ok {
phpHandlerContexts = append(phpHandlerContexts, ctx)
continue
}
if closesPHPHandlerContext(trimmed) {
if len(phpHandlerContexts) > 0 {
phpHandlerContexts = phpHandlerContexts[:len(phpHandlerContexts)-1]
}
continue
}
// A PHP execution handler mapped onto a non-PHP extension is the
// handler-remap webshell technique: an uploaded .jpg then runs as
// PHP. The safe-pattern and AddType skips below would otherwise
// suppress it because the handler name itself is a normal PHP
// handler, so this override fires first and unconditionally.
remapsNonPHP := phpHandlerRemapsNonPHP(lineLower)
if !remapsNonPHP && len(phpHandlerContexts) > 0 {
remapsNonPHP = phpHandlerRemapsNonPHPInContext(lineLower, phpHandlerContexts)
}
if remapsNonPHP {
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "htaccess_injection",
Message: "PHP handler mapped to non-PHP extension (handler remap)",
Details: fmt.Sprintf("File: %s (line %d)\nContent: %s", path, lineNum+1, trimmed),
FilePath: path,
})
continue
}
for _, pattern := range suspicious {
if !strings.Contains(lineLower, strings.ToLower(pattern)) {
continue
}
// A prelude directive is judged by its target file alone. The
// line-wide safe list below would let a target such as
// ".../uploads/fonts/x.ttf" or ".../litespeed/x.php" exempt itself
// with a word the attacker chose.
if m := reAutoPrependTarget.FindStringSubmatch(trimmed); m != nil {
if !autoPrependTargetSuspicious(m[1], path) {
continue
}
} else {
// Check per-line safe patterns
isSafe := false
for _, sp := range safe {
if strings.Contains(lineLower, strings.ToLower(sp)) {
isSafe = true
break
}
}
if isSafe {
continue
}
}
patternLower := strings.ToLower(pattern)
// For handler directives, apply context-aware checks
if patternLower == "addhandler" || patternLower == "sethandler" {
// Skip if paired with -ExecCGI (Wordfence protection)
if hasExecCGIBlock {
continue
}
// Skip Drupal security handlers
if strings.Contains(lineLower, "drupal_security") {
continue
}
// Skip SetHandler none/default (disabling handlers = security measure)
if strings.Contains(lineLower, "sethandler none") ||
strings.Contains(lineLower, "sethandler default") {
continue
}
// Skip AddHandler for standard CGI extensions only (.cgi, .pl)
if strings.Contains(lineLower, "addhandler") {
// Only flag if mapping non-standard extensions
standardCGI := true
hasNonStandard := false
// Check each extension on the line
for _, ext := range []string{".haxor", ".cgix", ".phtml", ".php3",
".php5", ".suspected", ".bak.php", ".shtml", ".sh"} {
if strings.Contains(lineLower, ext) {
hasNonStandard = true
break
}
}
// If line only has .cgi and/or .pl, it's standard
if !hasNonStandard && standardCGI {
onlyStandard := true
parts := strings.Fields(lineLower)
for _, p := range parts {
if strings.HasPrefix(p, ".") && p != ".cgi" && p != ".pl" && p != ".py" &&
p != ".php" && p != ".jsp" && p != ".asp" {
// Has non-standard extension
onlyStandard = false
break
}
}
if onlyStandard {
continue
}
}
}
}
// Skip AddType for any MIME type (application/*, text/*, x-mapp-*, etc.)
if patternLower == "addtype" {
// AddType is only dangerous if it maps to a PHP/CGI handler
// Standard MIME type declarations are safe
if strings.Contains(lineLower, "application/") ||
strings.Contains(lineLower, "text/") ||
strings.Contains(lineLower, "image/") ||
strings.Contains(lineLower, "font/") ||
strings.Contains(lineLower, "x-mapp-") ||
strings.Contains(lineLower, "audio/") ||
strings.Contains(lineLower, "video/") {
continue
}
}
*findings = append(*findings, alert.Finding{
Severity: alert.High,
Check: "htaccess_injection",
Message: fmt.Sprintf("Suspicious .htaccess directive: %s", pattern),
Details: fmt.Sprintf("File: %s (line %d)\nContent: %s", path, lineNum+1, trimmed),
FilePath: path,
})
}
}
// Special check: AddHandler mapping non-standard extensions WITHOUT -ExecCGI
// (actual attack pattern - e.g., AddHandler cgi-script .haxor)
if !hasExecCGIBlock && strings.Contains(fullContentLower, "addhandler") {
for _, logical := range joinHtaccessContinuations(lines) {
lineNum := logical.start
line := logical.text
lineLower := strings.ToLower(line)
if !strings.Contains(lineLower, "addhandler") {
continue
}
// Flag if it maps unusual extensions like .haxor, .cgix, etc.
dangerousExts := []string{".haxor", ".cgix", ".suspected", ".bak.php"}
for _, ext := range dangerousExts {
if strings.Contains(lineLower, ext) {
*findings = append(*findings, alert.Finding{
Severity: alert.Critical,
Check: "htaccess_handler_abuse",
Message: fmt.Sprintf("Malicious handler mapping for %s extension", ext),
Details: fmt.Sprintf("File: %s (line %d)\nContent: %s", path, lineNum+1, strings.TrimSpace(line)),
FilePath: path,
})
}
}
}
}
}
// CheckWPCore runs wp core verify-checksums for each WordPress installation
// using a bounded worker pool for concurrency.
// Installations that pass verification have their core files cached in
// GlobalCMSCache so the real-time scanner can skip signature matches
// on known-clean CMS files.
func CheckWPCore(ctx context.Context, cfg *config.Config, _ *state.Store) (findings []alert.Finding) {
if ctx == nil {
ctx = context.Background()
}
if incompleteCollectorFrom(ctx) == nil {
ctx, _ = withIncompleteCheckCollector(ctx)
}
wpConfigs := wpCoreScanRoots(ctx)
coverage := newWPVerificationBatch(ctx, store.Global(), "core", logicalOwnerWPCoreVerification, wpConfigs)
defer func() {
err := coverage.finish(ctx, !checkMarkedIncomplete(ctx, "wp_core"))
if _, disabled := disabledLogicalOwners(cfg)[logicalOwnerWPCoreVerification]; !disabled {
findings = append(findings, wpVerificationFindings(ctx, coverage.db, "core", logicalOwnerWPCoreVerification, err)...)
}
}()
if len(wpConfigs) == 0 {
return nil
}
cache := GlobalCMSCache()
var mu sync.Mutex
var wg sync.WaitGroup
batch := wpCoreBatches.begin(len(wpConfigs), wpChecksumWorkers)
defer batch.abandon(ctx)
jobs := make(chan int, len(wpConfigs))
for i := 0; i < wpChecksumWorkers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
for index := range jobs {
wpConfig, work := wpConfigs[index], batch.tasks[index]
work.admit()
stop := false
work.run(ctx, cmdTimeout, func() {
if parentErr := ctx.Err(); parentErr != nil {
work.withdraw(parentErr)
stop = true
return
}
wpPath := filepath.Dir(wpConfig)
user := wpConfigUser(wpPath)
out, err := runCmdCombinedContext(ctx, "wp", "core", "verify-checksums",
"--path="+wpPath, "--allow-root")
work.progress()
if parentErr := ctx.Err(); parentErr != nil {
work.withdraw(parentErr)
stop = true
return
}
if err == nil {
coverage.record(wpPath, store.WPVerificationResult{State: "verified"})
// Verification passed - cache all core files
cacheWPCoreFiles(cache, wpPath)
return
}
coverage.record(wpPath, wpVerificationFailure(err, out))
// Partial integrity output cannot complete a command killed by a signal.
var commandExit *exec.ExitError
if errors.As(err, &commandExit) && commandExit.ExitCode() < 0 {
work.fail()
}
if out == nil {
work.fail()
return
}
outStr := string(out)
var extraneous []string
reported := false
for _, line := range strings.Split(outStr, "\n") {
reported = reported || wpChecksumModifiedFilePath(line) != "" || strings.Contains(line, "should not exist")
if wpChecksumLineHasExtraneousCoreFile(line) {
extraneous = append(extraneous, strings.TrimSpace(line))
continue
}
// A shipped core file whose bytes changed is where backdoors
// are appended; that is worse than an extra file, and the
// path lets Re-check and the operator go straight to it.
if rel := wpChecksumModifiedCoreFile(line); rel != "" {
mu.Lock()
findings = append(findings, alert.Finding{
Severity: wpCoreModifiedSeverity(wpCoreFilePathWithin(wpPath, rel), rel),
Check: "wp_core_integrity",
Message: fmt.Sprintf("WordPress core file modified for %s", user),
Details: fmt.Sprintf("Path: %s\nFile: %s\n%s", wpPath, rel, line),
FilePath: wpCoreFilePathWithin(wpPath, rel),
})
mu.Unlock()
}
}
if len(extraneous) > 0 {
collapsed := wpCoreExtraneousFinding(user, wpPath, extraneous)
// No single file to name, so the install's owner carries
// the identity correlation needs.
if owner, ok := installOwner(wpConfig); ok {
collapsed.TenantID = owner
}
mu.Lock()
findings = append(findings, collapsed)
mu.Unlock()
}
if wpCoreVerificationCompleted(err, out) {
coverage.record(wpPath, store.WPVerificationResult{State: "modified"})
}
// wp-cli that ran and refused this tree answered the check.
if !reported && !commandRefused(err) {
work.fail()
}
})
if stop {
return
}
}
}()
}
for index := range wpConfigs {
if ctx.Err() != nil {
break
}
jobs <- index
}
close(jobs)
wg.Wait()
batch.abandon(ctx)
fmt.Fprintf(os.Stderr, "CMS hash cache: %d verified core files cached\n", cache.Size())
return findings
}
// wpCoreExtraneousSampleLimit bounds how many wp-cli lines the collapsed
// extra-file finding quotes. Enough to recognise the shape of the damage
// without turning one broken install into a wall of text.
const wpCoreExtraneousSampleLimit = 15
// wpCoreExtraneousFinding collapses every "should not exist" line wp-cli
// reported for one install into a single finding. The lines describe one
// condition -- this core is not the release it claims to be -- and a core
// rebuilt from an older release reports every file the newer one shipped, so
// emitting them per file buries the rest of the scan. Identity is pinned to
// the install so the row survives the operator deleting the files one by one.
func wpCoreExtraneousFinding(user, wpPath string, lines []string) alert.Finding {
sorted := append([]string(nil), lines...)
sort.Strings(sorted)
var details strings.Builder
fmt.Fprintf(&details, "Path: %s\n", wpPath)
fmt.Fprintf(&details, "Core files reported as extraneous: %d\n", len(sorted))
for _, line := range firstN(sorted, wpCoreExtraneousSampleLimit) {
details.WriteString(line)
details.WriteString("\n")
}
if extra := len(sorted) - wpCoreExtraneousSampleLimit; extra > 0 {
fmt.Fprintf(&details, "... and %d more\n", extra)
}
return alert.Finding{
Severity: alert.High,
Check: "wp_core_integrity",
Message: fmt.Sprintf("WordPress core integrity failure for %s", user),
Details: details.String(),
DedupKey: "extraneous:" + wpPath,
}
}
// cacheWPCoreFiles hashes all PHP files in wp-includes/ and wp-admin/
// for a verified-clean WordPress installation and adds them to the cache.
func cacheWPCoreFiles(cache *CMSHashCache, wpPath string) {
coreDirs := []string{
filepath.Join(wpPath, "wp-includes"),
filepath.Join(wpPath, "wp-admin"),
}
// Also cache root-level WP core files
rootFiles := []string{
"wp-cron.php", "wp-login.php", "wp-settings.php",
"wp-load.php", "wp-blog-header.php", "wp-links-opml.php",
"wp-mail.php", "wp-signup.php", "wp-activate.php",
"wp-comments-post.php", "wp-trackback.php", "xmlrpc.php",
"index.php",
}
for _, name := range rootFiles {
path := filepath.Join(wpPath, name)
if hash := HashFile(path); hash != "" {
if info, err := osFS.Stat(path); err == nil {
cache.Add(hash, info.Size())
}
}
}
for _, dir := range coreDirs {
_ = filepath.Walk(dir, func(path string, info os.FileInfo, err error) error {
if err != nil {
return err
}
if info.IsDir() {
return nil
}
name := strings.ToLower(info.Name())
if strings.HasSuffix(name, ".php") || strings.HasSuffix(name, ".js") {
if hash := HashFile(path); hash != "" {
cache.Add(hash, info.Size())
}
}
return nil
})
}
}
func extractUser(path string) string {
parts := strings.Split(path, "/")
for i, p := range parts {
if p == "home" && i+1 < len(parts) {
return parts[i+1]
}
}
return "unknown"
}
// wpCoreScanRoots lists the WordPress installs to verify. Discovery is shared
// (wpinstalls.go): core files are tampered with in subdomain and nested
// installs as readily as in a primary document root.
func wpCoreScanRoots(ctx context.Context) []string {
installs := wpInstalls(ctx, "wp_core")
out := make([]string, 0, len(installs))
for _, in := range installs {
out = append(out, in.ConfigPath)
}
return out
}
package checks
import (
"bufio"
"bytes"
"context"
"errors"
"fmt"
"io"
"os"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
)
// CheckWHMAccess parses the cPanel access log for WHM (port 2087) logins
// and password change API calls from non-infra IPs.
// Only reads the tail of the log - lightweight.
func CheckWHMAccess(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
var findings []alert.Finding
lines := tailFile("/usr/local/cpanel/logs/access_log", 200)
for _, line := range lines {
if !isWHMAccessLogLine(line) {
continue
}
// Extract IP (first field)
fields := strings.Fields(line)
if len(fields) < 1 {
continue
}
ip := fields[0]
// Skip infra IPs
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
continue
}
// Check for password change actions
passwordActions := []string{
"passwd", "change_root_password", "chpasswd",
"force_password_change", "resetpass",
}
lineLower := strings.ToLower(line)
for _, action := range passwordActions {
if strings.Contains(lineLower, action) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "whm_password_change",
Message: fmt.Sprintf("WHM password change from non-infra IP: %s", ip),
Details: truncateString(line, 200),
})
break
}
}
// Check for account management from unknown IPs
accountActions := []string{
"createacct", "killacct", "suspendacct", "unsuspendacct",
}
for _, action := range accountActions {
if strings.Contains(lineLower, action) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "whm_account_action",
Message: fmt.Sprintf("WHM account action from non-infra IP: %s", ip),
Details: truncateString(line, 200),
})
break
}
}
}
return findings
}
func isWHMAccessLogLine(line string) bool {
served := lastAccessLogField(line)
switch {
case served == "2087":
return accessLogQuotedFieldCount(line) >= 5
case strings.HasSuffix(served, ":2087"):
return accessLogQuotedFieldCount(line) >= 4
default:
return false
}
}
func lastAccessLogField(line string) string {
line = strings.TrimSpace(line)
if line == "" {
return ""
}
if strings.HasSuffix(line, "\"") {
end := len(line) - 1
start := strings.LastIndex(line[:end], "\"")
if start < 0 {
return ""
}
return line[start+1 : end]
}
fields := strings.Fields(line)
if len(fields) == 0 {
return ""
}
return fields[len(fields)-1]
}
func accessLogQuotedFieldCount(line string) int {
return strings.Count(line, "\"") / 2
}
var authLogPath = func() string { return platform.Detect().AuthLogPath() }
// CheckSSHLogins parses the platform authentication log for SSH logins from
// non-infra IPs. With a state store it reads the log forward-only from where
// the previous cycle stopped; without one it falls back to a per-cycle tail.
func CheckSSHLogins(ctx context.Context, cfg *config.Config, store *state.Store) []alert.Finding {
if cfg == nil {
cfg = &config.Config{}
}
if store != nil {
return checkSSHLoginsFollow(cfg, store)
}
var findings []alert.Finding
for _, line := range tailFile(authLogPath(), 100) {
if !strings.Contains(line, "Accepted") {
continue
}
if f, ok := SSHAcceptedLoginFinding(line, cfg); ok {
findings = append(findings, f)
}
}
return findings
}
// SSHAcceptedLoginFinding parses an sshd "Accepted <method> for <user> from
// <ip> port <n>" line and reports it unless the address is infrastructure.
// The daemon's realtime log watcher calls it so a login seen live and the same
// line re-read by CheckSSHLogins carry one identity; without that the state
// store sees two findings and the operator gets one login reported twice.
func SSHAcceptedLoginFinding(line string, cfg *config.Config) (alert.Finding, bool) {
if cfg == nil {
cfg = &config.Config{}
}
if !strings.Contains(line, "Accepted") {
return alert.Finding{}, false
}
parts := strings.Fields(line)
ipIdx := -1
for i, p := range parts {
if p == "from" && i+1 < len(parts) {
ipIdx = i + 1
break
}
}
if ipIdx < 0 || ipIdx >= len(parts) {
return alert.Finding{}, false
}
ip := parts[ipIdx]
if isInfraIP(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
return alert.Finding{}, false
}
user := "unknown"
for i, p := range parts {
if p == "for" && i+1 < len(parts) {
user = parts[i+1]
break
}
}
tenant := user
if tenant == "unknown" {
tenant = ""
}
return alert.Finding{
Severity: alert.Critical,
Check: "ssh_login_unknown_ip",
DedupKey: loginRecordKey(line),
Message: fmt.Sprintf("SSH login from non-infra IP: %s (user: %s)", ip, user),
Details: truncateString(line, 200),
SourceIP: ip,
TenantID: tenant,
}, true
}
// tailFile reads the last N lines of a file efficiently.
func tailFile(path string, maxLines int) []string {
if maxLines <= 0 {
return nil
}
f, err := osFS.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
// Seek to end and read backwards to find last N lines
info, err := f.Stat()
if err != nil {
return nil
}
// For small files, just read all
if info.Size() < 1024*1024 {
return readAllLines(f, maxLines)
}
data, err := readTailWindow(f, info.Size(), maxLines, maxTailWindowBytes)
if err != nil {
return readAllLines(f, maxLines)
}
return readAllLines(bytes.NewReader(data), maxLines)
}
func readTailWindow(f *os.File, size int64, maxLines int, maxBytes int64) ([]byte, error) {
const chunkSize int64 = 256 * 1024
if maxBytes <= 0 {
return nil, nil
}
offset := size
newlines := 0
var totalRead int64
chunks := make([][]byte, 0, 4)
for offset > 0 && newlines <= maxLines && totalRead < maxBytes {
n := chunkSize
if offset < n {
n = offset
}
if remaining := maxBytes - totalRead; remaining < n {
n = remaining
}
offset -= n
chunk := make([]byte, n)
read, err := f.ReadAt(chunk, offset)
if err != nil && !errors.Is(err, io.EOF) {
return nil, err
}
chunk = chunk[:read]
totalRead += int64(read)
newlines += bytes.Count(chunk, []byte{'\n'})
chunks = append(chunks, chunk)
}
total := 0
for _, chunk := range chunks {
total += len(chunk)
}
data := make([]byte, 0, total)
for i := len(chunks) - 1; i >= 0; i-- {
data = append(data, chunks[i]...)
}
if offset > 0 {
if firstNewline := bytes.IndexByte(data, '\n'); firstNewline >= 0 {
data = data[firstNewline+1:]
} else {
return nil, nil
}
}
return data, nil
}
const (
// maxLogLineBytes is the per-line cap for periodic log tailers.
// Oversized records are skipped after the reader advances past the
// terminator so a crafted long line cannot poison the next record.
maxLogLineBytes = 256 * 1024
// maxTailWindowBytes bounds the backward seek window before line
// parsing starts. Without this, a huge unterminated final record makes
// the tail reader cache the whole file while looking for maxLines.
maxTailWindowBytes int64 = 32 * 1024 * 1024
)
func readAllLines(r io.Reader, maxLines int) []string {
if maxLines <= 0 {
return nil
}
br := bufio.NewReaderSize(r, 64*1024)
var lines []string
for {
line, truncated, err := readBoundedLineLog(br, maxLogLineBytes)
if len(line) > 0 && !truncated {
lines = append(lines, trimLogLineEnding(line))
}
if err != nil {
break
}
}
if len(lines) > maxLines {
return lines[len(lines)-maxLines:]
}
return lines
}
func trimLogLineEnding(line string) string {
line = strings.TrimSuffix(line, "\n")
return strings.TrimSuffix(line, "\r")
}
// readBoundedLineLog reads up to and including the next '\n'. If the
// line exceeds maxBytes the returned data is truncated to maxBytes and
// the reader is advanced past the line's terminating newline so framing
// stays intact. Returns the same error semantics as
// bufio.Reader.ReadString.
func readBoundedLineLog(r *bufio.Reader, maxBytes int) (string, bool, error) {
var b strings.Builder
truncated := false
for {
chunk, err := r.ReadSlice('\n')
if len(chunk) > 0 {
switch {
case truncated:
// drain remainder so the next line is well-framed
case b.Len()+len(chunk) <= maxBytes:
b.Write(chunk)
default:
if room := maxBytes - b.Len(); room > 0 {
b.Write(chunk[:room])
}
truncated = true
}
}
if errors.Is(err, bufio.ErrBufferFull) {
continue
}
return b.String(), truncated, err
}
}
func truncateString(s string, maxLen int) string {
if len(s) <= maxLen {
return s
}
return s[:maxLen] + "..."
}
package checks
import (
"context"
"errors"
"fmt"
"os/exec"
"path/filepath"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
type wpVerificationBatch struct {
db *store.DB
kind, owner string
at time.Time
paths map[string]string
mu sync.Mutex
results map[string]store.WPVerificationResult
}
func newWPVerificationBatch(ctx context.Context, db *store.DB, kind, owner string, configs []string) *wpVerificationBatch {
at := time.Now()
if cycle := wpInstallCacheFrom(ctx); cycle != nil {
at = cycle.started
}
b := &wpVerificationBatch{db: db, kind: kind, owner: owner, at: at, paths: make(map[string]string), results: make(map[string]store.WPVerificationResult)}
for _, path := range configs {
b.paths[filepath.Dir(path)] = wpConfigUser(filepath.Dir(path))
}
return b
}
func (b *wpVerificationBatch) record(path string, result store.WPVerificationResult) {
b.mu.Lock()
b.results[path] = result
b.mu.Unlock()
}
// finish runs after workers join, including on partial discovery or shutdown.
// Attempted sites advance; unattempted sites retain their previous evidence.
func (b *wpVerificationBatch) finish(ctx context.Context, complete bool) error {
if b.db == nil {
return nil
}
err := b.db.UpdateWPVerification(b.kind, b.at, AccountFromContext(ctx), b.paths, b.results, complete && ctx.Err() == nil)
if err != nil {
markCheckIncomplete(ctx, b.owner)
}
return err
}
func wpVerificationFindings(ctx context.Context, db *store.DB, kind, owner string, updateErr error) []alert.Finding {
if db == nil {
return nil
}
check, label := "wp_core_unverified", "core verification"
if kind == "plugins" {
check, label = "wp_plugin_inventory_unverified", "plugin inventory"
}
rows, err := db.WPVerification(kind)
if err != nil || updateErr != nil {
markCheckIncomplete(ctx, owner)
return []alert.Finding{{Check: check, Severity: alert.Warning, Message: "WordPress " + label + " history unavailable", Details: "CSM could not read or save verification history. Check the state database; previous coverage findings are retained."}}
}
paths := make([]string, 0, len(rows))
for path := range rows {
paths = append(paths, path)
}
sort.Strings(paths)
scope := AccountFromContext(ctx)
byReason := make(map[string][]string)
for _, path := range paths {
row := rows[path]
if scope != "" && row.Account != scope {
continue
}
if row.State != "unverified" || row.Failures < 2 {
continue
}
byReason[row.Reason] = append(byReason[row.Reason], path)
}
reasons := make([]string, 0, len(byReason))
for reason := range byReason {
reasons = append(reasons, reason)
}
sort.Strings(reasons)
var findings []alert.Finding
for _, reason := range reasons {
group := byReason[reason]
// Only a host scan represents the whole cause. Account scans retain
// installation identities so a scoped result cannot replace or dedup
// against a host summary (or another account's summary).
if scope == "" && len(group) > wpVerificationCollapseCap {
findings = append(findings, wpVerificationCollapsedFinding(check, label, reason, group, rows))
continue
}
for _, path := range group {
row := rows[path]
findings = append(findings, alert.Finding{
Check: check, Severity: alert.Warning,
Message: "WordPress " + label + " repeatedly failed: " + strconv.QuoteToASCII(path),
Details: fmt.Sprintf("Installation: %s\nLast attempt: %s\nReason: %s\nVerification could not complete in consecutive scan cycles. Check wp-cli and this installation locally; raw command output is not stored.", strconv.QuoteToASCII(path), row.AttemptAt.UTC().Format(time.RFC3339), reason),
FilePath: path, TenantID: row.Account, DedupKey: path,
})
}
}
return findings
}
// wpVerificationCollapseCap is how many installations may share one cause
// before the finding stops naming them one alert at a time. A missing wp-cli
// or an unreachable checksum service fails every installation on the host in
// the same cycle, and one alert per installation would spend the whole
// alerts.max_per_hour budget on a single operational fault.
const wpVerificationCollapseCap = 10
// wpVerificationSampleLimit bounds how many installations a collapsed finding
// quotes. Enough to recognise the affected accounts without a wall of paths.
const wpVerificationSampleLimit = 10
// wpVerificationCollapsedFinding reports one cause that stopped verification
// across many installations. The cause is what the operator acts on, so the
// identity is the cause and not any single installation.
func wpVerificationCollapsedFinding(check, label, reason string, group []string, rows map[string]store.WPVerificationRecord) alert.Finding {
var last time.Time
accounts := make(map[string]struct{}, len(group))
for _, path := range group {
row := rows[path]
accounts[row.Account] = struct{}{}
if row.AttemptAt.After(last) {
last = row.AttemptAt
}
}
var details strings.Builder
fmt.Fprintf(&details, "Reason: %s\n", reason)
fmt.Fprintf(&details, "Installations affected: %d\n", len(group))
fmt.Fprintf(&details, "Last attempt: %s\n", last.UTC().Format(time.RFC3339))
for _, path := range firstN(group, wpVerificationSampleLimit) {
fmt.Fprintf(&details, "- %s\n", strconv.QuoteToASCII(path))
}
if extra := len(group) - wpVerificationSampleLimit; extra > 0 {
fmt.Fprintf(&details, "... and %d more\n", extra)
}
details.WriteString("Verification could not complete in consecutive scan cycles. One cause stopped every installation listed, so fix it once; raw command output is not stored.")
finding := alert.Finding{
Check: check, Severity: alert.Warning,
Message: fmt.Sprintf("WordPress %s repeatedly failed for %d installations", label, len(group)),
Details: details.String(),
DedupKey: "reason:" + reason,
}
if len(accounts) == 1 {
for account := range accounts {
finding.TenantID = account
}
}
return finding
}
// Integrity warnings can precede a later operational error. Only wp-cli's
// expected checksum-mismatch summary establishes a completed negative check.
func wpCoreVerificationCompleted(err error, out []byte) bool {
if !commandRefused(err) {
return false
}
completed := false
for _, line := range strings.Split(strings.ToLower(string(out)), "\n") {
line = strings.TrimSpace(line)
if strings.Contains(line, "fatal error") || strings.Contains(line, "parse error") {
return false
}
if line == "error: wordpress installation doesn't verify against checksums." {
completed = true
} else if strings.HasPrefix(line, "error:") {
return false
}
}
return completed
}
type wpInventoryError struct {
err error
result store.WPVerificationResult
}
func (e *wpInventoryError) Error() string { return e.result.Reason }
func (e *wpInventoryError) Unwrap() error { return e.err }
// wpVerificationFailure classifies output into fixed reasons. Never copy raw
// PHP output or an exec error into persisted evidence: either may carry secrets.
func wpVerificationFailure(err error, out []byte) store.WPVerificationResult {
var inventory *wpInventoryError
if errors.As(err, &inventory) {
return inventory.result
}
result := store.WPVerificationResult{State: "unverified", Reason: "wp-cli could not complete the check"}
var exit *exec.ExitError
switch {
case errors.Is(err, context.DeadlineExceeded):
result.Reason = "wp-cli timed out"
case errors.Is(err, context.Canceled):
result.Reason = "wp-cli was interrupted"
case errors.Is(err, exec.ErrNotFound):
result.Reason = "wp-cli executable is unavailable"
case errors.Is(err, errWPInventoryParse):
result.Reason = "wp-cli returned invalid plugin inventory JSON"
case errors.Is(err, errWPInventoryNoOutput):
result.Reason = "wp-cli returned no plugin inventory output"
case errors.As(err, &exit) && exit.ExitCode() < 0:
result.Reason = "wp-cli was terminated by a signal"
default:
text := strings.ToLower(string(out))
if len(out) == 0 && exit != nil {
text = strings.ToLower(string(exit.Stderr))
}
switch {
case commandRefused(err) && strings.Contains(text, "error: this does not seem to be a wordpress installation."):
result.State, result.Reason = "not_wordpress", "wp-cli did not find a WordPress installation"
case strings.Contains(text, "fatal error"), strings.Contains(text, "parse error"):
result.Reason = "WordPress configuration or PHP initialization failed"
case strings.Contains(text, "checksum"):
result.Reason = "WordPress checksum data could not be obtained or verified"
case strings.Contains(text, "database connection"):
result.Reason = "WordPress could not connect to its database"
case strings.Contains(text, "permission denied"):
result.Reason = "wp-cli could not read installation files or execute a required command"
case strings.Contains(text, "not found"):
result.Reason = "wp-cli or a required installation file is unavailable"
case len(text) == 0:
result.Reason = "wp-cli returned no diagnostic output"
}
}
return result
}
// CheckWPPluginVerification reports the shared inventory independently of its
// outdated/known-vulnerable consumers. The existing plugin refresh interval and
// disabled_checks controls remain authoritative.
func CheckWPPluginVerification(ctx context.Context, cfg *config.Config, _ *state.Store) []alert.Finding {
if ctx == nil {
ctx = context.Background()
}
if incompleteCollectorFrom(ctx) == nil {
ctx, _ = withIncompleteCheckCollector(ctx)
}
disabled := make(map[string]bool)
if cfg != nil {
for _, name := range cfg.DisabledChecks {
disabled[strings.TrimSpace(name)] = true
}
}
if disabled["outdated_plugins"] && (disabled["vulnerable_plugins"] || (cfg != nil && !cfg.VulnerablePluginScanningEnabled())) {
return nil
}
db := store.Global()
if db == nil {
return nil
}
ensurePluginCacheFresh(ctx, cfg, db)
var updateErr error
if checkMarkedIncomplete(ctx, "wp_plugin_inventory") {
updateErr = errors.New("verification history update failed")
}
return wpVerificationFindings(ctx, db, "plugins", "wp_plugin_inventory", updateErr)
}
package checks
import (
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var wpCoreBatches = newScanBatchMonitor()
// WPCoreQueueStatus includes selected installations through checksum execution,
// result collection and verified-file caching. Concurrent scans share no cap.
func WPCoreQueueStatus(now time.Time) queuehealth.Status {
status := wpCoreBatches.snapshot(now)
status.DepthUnit = "installations"
return status
}
package checks
import (
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/config"
)
// MigrateWPCronCrontabs upgrades CSM-managed wp-cron crontab lines installed
// by older releases (synchronized */N schedule, no overlap lock) to the
// current staggered format. The perf_wp_cron finding never re-fires once
// DISABLE_WP_CRON is set, so already-fixed accounts can only be reached by
// walking the spool directly. Gated on the same fix_wp_cron opt-in as the
// install path; returns the number of crontabs rewritten.
func MigrateWPCronCrontabs(cfg *config.Config) int {
if cfg == nil || !cfg.AutoResponse.Enabled || !cfg.AutoResponse.FixWPCron {
return 0
}
opts := WPCronFixOptions{
IntervalMinutes: cfg.Performance.WPCronFix.IntervalMinutes,
PHPBin: cfg.Performance.WPCronFix.PHPBin,
}
upgraded := 0
seen := map[string]bool{}
for _, dir := range wpCronSpoolDirs {
entries, err := osFS.ReadDir(dir)
if err != nil {
continue
}
for _, e := range entries {
if e.IsDir() {
continue
}
owner := e.Name()
if seen[owner] || !validCPUser.MatchString(owner) || owner == "root" {
continue
}
data, err := osFS.ReadFile(filepath.Join(dir, owner))
if err != nil {
continue
}
docroots := wpCronManagedDocroots(string(data))
if len(docroots) == 0 {
continue
}
seen[owner] = true
for _, docroot := range docroots {
installed, err := installUserWPCron(owner, docroot, opts)
if err == nil && installed {
upgraded++
}
}
}
}
return upgraded
}
// wpCronManagedDocroots extracts docroots from CSM marker lines. The marker
// is attacker-writable in principle (the spool file belongs to the account),
// so only clean absolute paths may flow into a crontab rewrite.
func wpCronManagedDocroots(crontab string) []string {
var roots []string
for _, line := range strings.Split(crontab, "\n") {
trimmed := strings.TrimSpace(line)
if !strings.HasPrefix(trimmed, wpCronJobMarker) {
continue
}
docroot := strings.TrimSpace(strings.TrimPrefix(trimmed, wpCronJobMarker))
if !safeWPCronDocroot(docroot) {
continue
}
roots = append(roots, docroot)
}
return roots
}
package checks
import (
"path/filepath"
"regexp"
"strings"
)
// cPanel records the MultiPHP version chosen for each vhost in a fixed column
// of /etc/userdatadomains, as an "ea-phpNN" or "alt-phpNN" token.
// Anything outside that shape is not a version we can turn into a path.
var userdataPHPVersionRe = regexp.MustCompile(`^(ea|alt)-php([0-9]{2})$`)
const userdataPHPVersionField = 9
// parseVhostPHPVersion pulls the MultiPHP token out of a /etc/userdatadomains
// row. The field has a fixed position after the IPv6-dedicated flag. Older
// rows may stop before it and newer rows may carry trailing empty fields, so
// accept only the fixed PHP-version column instead of searching attacker-adjacent
// fields for a version-shaped token.
func parseVhostPHPVersion(fields []string) string {
if len(fields) <= userdataPHPVersionField {
return ""
}
for _, trailing := range fields[userdataPHPVersionField+1:] {
if strings.TrimSpace(trailing) != "" {
return ""
}
}
tok := strings.TrimSpace(fields[userdataPHPVersionField])
if userdataPHPVersionRe.MatchString(tok) {
return tok
}
return ""
}
// phpBinForVersion maps a cPanel MultiPHP version token to its interpreter.
// EasyApache and CloudLinux alt-php lay their trees out differently, so the
// two shapes are built separately rather than by string substitution.
// Returns empty for anything that is not a well-formed version token, which is
// what keeps a malformed or attacker-influenced map out of a crontab line.
func phpBinForVersion(version string) string {
m := userdataPHPVersionRe.FindStringSubmatch(strings.TrimSpace(version))
if m == nil {
return ""
}
if m[1] == "alt" {
return "/opt/alt/php" + m[2] + "/usr/bin/php"
}
return "/opt/cpanel/ea-php" + m[2] + "/root/usr/bin/php"
}
// resolveDocrootPHPBin returns the PHP interpreter the owner's docroot is
// pinned to, or empty when the docroot is unknown, ambiguous, or unusable.
//
// WP-Cron has to run under the same interpreter as the site: a docroot pinned
// to an old MultiPHP version fatal-errors when driven by a newer system
// default, which silently kills scheduled tasks on that site.
//
// The result is deliberately restricted to the two known-good path shapes.
// safeManagedWPCronPHPBin accepts exactly those, so a resolved path never
// makes CSM report its own crontab as an unexpected change.
func resolveDocrootPHPBin(owner, docroot string) string {
content, err := osFS.ReadFile(userdataDomainsPath)
if err != nil {
return ""
}
want := filepath.Clean(docroot)
vhosts, _ := parseUserdataDomainRootsChecked(string(content))
// Match the most specific docroot this account owns that serves the target.
// A WordPress install in a subdirectory is not its own vhost, so the map has
// no entry for it, but cPanel serves it under the enclosing vhost and hence
// that vhost's PHP version. An exact entry always wins, because a subdomain
// docroot can be nested inside the main one.
version, bestLen := "", -1
for _, vh := range vhosts {
if vh.user != owner || !wpCronDocrootCovers(vh.docroot, want) {
continue
}
switch {
case len(vh.docroot) > bestLen:
version, bestLen = vh.phpVersion, len(vh.docroot)
case len(vh.docroot) == bestLen && vh.phpVersion != version:
// Two vhosts claim the same docroot with different versions;
// picking either would pin the wrong interpreter.
version = ""
}
}
if version == "" {
return ""
}
bin := phpBinForVersion(version)
if bin == "" || !safeManagedWPCronPHPBin(bin) {
return ""
}
return bin
}
// wpCronDocrootCovers reports whether vhostRoot serves docroot: either the same
// path or an ancestor of it. Comparison is path-segment aware so
// /home/a/public_html never covers /home/a/public_html_old, and a root shallower
// than /home/<user>/<dir> is rejected so inheritance cannot cross accounts.
func wpCronDocrootCovers(vhostRoot, docroot string) bool {
vhostRoot = filepath.Clean(vhostRoot)
if strings.Count(strings.TrimSuffix(vhostRoot, "/"), "/") < 3 {
return false
}
if vhostRoot == docroot {
return true
}
return strings.HasPrefix(docroot, vhostRoot+"/")
}
// resolveWPCronPHPBin distinguishes an operator override or unambiguous vhost
// mapping from fallback detection. Callers upgrading an existing managed line
// use that provenance to avoid replacing a known-good interpreter when the
// domain map is temporarily unavailable.
func resolveWPCronPHPBin(owner, docroot string, opts WPCronFixOptions) (string, bool) {
if opts.PHPBin != "" {
return opts.PHPBin, true
}
if bin := resolveDocrootPHPBin(owner, docroot); bin != "" {
return bin, true
}
return detectPHPBin(), false
}
// wpCronPHPBin picks the interpreter for a managed cron line. An operator who
// sets php_bin has overridden the choice deliberately, so that wins; otherwise
// the vhost's own version wins; detection is the last resort so a host without
// a usable domain map still gets an installable line.
func wpCronPHPBin(owner, docroot string, opts WPCronFixOptions) string {
bin, _ := resolveWPCronPHPBin(owner, docroot, opts)
return bin
}
package checks
import (
"bytes"
"fmt"
"hash/fnv"
"os"
"os/user"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"syscall"
)
// WPCronFixOptions carries operator-tunable parameters for the WP-Cron
// remediation. Both the Web UI handler and the daemon auto-response resolve
// these from config before calling the fix, so the remediation core itself
// stays free of config coupling.
type WPCronFixOptions struct {
// IntervalMinutes is how often the installed system cron runs wp-cron.php.
// Clamped to [1,60]; a non-positive value falls back to the 15-minute default.
IntervalMinutes int
// PHPBin is the interpreter the cron line invokes. Empty uses an unambiguous
// cPanel vhost version, then LookPath("php"), then /usr/local/bin/php.
PHPBin string
}
const (
wpCronDefaultIntervalMin = 15
wpCronMaxIntervalMin = 60
// wpCronEditMarker tags the line CSM inserts so the customer can see why
// WP-Cron was disabled and so re-running the fix stays idempotent.
wpCronEditMarker = "// CSM: WP-Cron disabled, served by system cron instead"
// wpCronJobMarker prefixes the managed crontab block for a given docroot.
wpCronJobMarker = "# CSM WP-Cron "
wpCronStopMarker = "stop editing"
)
// wpCronDefineRe matches a define of DISABLE_WP_CRON set to a truthy value,
// matching the detector's view of "already disabled".
var wpCronDefineRe = regexp.MustCompile(`(?i)define\s*\(\s*['"]DISABLE_WP_CRON['"]\s*,\s*['"]?(true|1)['"]?\s*\)`)
// validCPUser guards the username passed to `crontab -u`. cPanel usernames are
// lowercase alnum starting with a letter; rejecting anything else keeps a
// surprising file owner from reaching the crontab argument vector.
var validCPUser = regexp.MustCompile(`^[a-z][a-z0-9_-]{0,31}$`)
var wpCronHeredocStartRe = regexp.MustCompile(`<<<['"]?([A-Za-z_][A-Za-z0-9_]*)['"]?`)
// Crontab installs are read-modify-write; serialize each account in-process.
var wpCronCrontabLocks sync.Map
// wp-config.php edits are read-modify-write too. Deep and periodic scans can
// overlap, so serialize each config path before resolving, reading, and writing.
var wpCronConfigLocks sync.Map
// wpCronOwnerName resolves the account that owns a wp-config.php. It is a var
// so tests can inject a deterministic owner regardless of who runs `go test`.
var wpCronOwnerName = fileOwnerName
// FixDisableWPCron disables WP-Cron in a wp-config.php and installs a real
// per-user system cron that runs wp-cron.php on a fixed interval. It scopes
// writes to the default per-account roots (/home).
func FixDisableWPCron(path string, opts WPCronFixOptions) RemediationResult {
return FixDisableWPCronInRoots(path, fixPerfAllowedRoots, opts)
}
// FixDisableWPCronInRoots is FixDisableWPCron with caller-supplied roots so the
// Web UI can honor configured account_roots and tests can write under t.TempDir().
func FixDisableWPCronInRoots(path string, allowedRoots []string, opts WPCronFixOptions) RemediationResult {
if path == "" {
return RemediationResult{Error: "could not extract file path from finding"}
}
lockPath, err := sanitizeFixPath(path, allowedRoots)
if err != nil {
return RemediationResult{Error: err.Error()}
}
lock := wpCronConfigLock(lockPath)
lock.Lock()
defer lock.Unlock()
resolved, info, err := resolveExistingFixPath(lockPath, allowedRoots)
if err != nil {
return RemediationResult{Error: err.Error()}
}
if info.IsDir() {
return RemediationResult{Error: "refusing to edit a directory"}
}
if filepath.Base(resolved) != "wp-config.php" {
return RemediationResult{Error: fmt.Sprintf("automated WP-Cron fix only applies to wp-config.php (got %s)", filepath.Base(resolved))}
}
data, err := readFilePreservingIdentity(resolved, info)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("read failed: %v", err)}
}
isWordPress, validateErr := wpCronInstallIsValid(resolved, data)
if validateErr != nil {
return RemediationResult{Error: fmt.Sprintf("WordPress validation failed: %v", validateErr)}
}
if !isWordPress {
return RemediationResult{Error: "refusing WP-Cron fix because the directory is not a complete WordPress install"}
}
var actions []string
needsDefine := !wpCronHasActiveDisableDefine(data)
var rewritten []byte
if needsDefine {
var ok bool
rewritten, ok = insertDisableWPCron(data)
if !ok {
return RemediationResult{Error: "could not find a safe insertion point in wp-config.php (no \"stop editing\" marker or wp-settings.php require)"}
}
}
docroot := filepath.Dir(resolved)
owner, err := wpCronOwnerName(info)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("could not resolve account owner of wp-config.php: %v", err)}
}
cronInstalled, err := installUserWPCron(owner, docroot, opts)
if err != nil {
return RemediationResult{Error: fmt.Sprintf("system cron install failed: %v", err)}
}
if needsDefine {
if werr := writeFilePreservingOwner(resolved, rewritten, info); werr != nil {
if cronInstalled {
return RemediationResult{Error: fmt.Sprintf("system cron installed but wp-config.php update failed: %v", werr)}
}
return RemediationResult{Error: werr.Error()}
}
actions = append(actions, "disabled WP-Cron in wp-config.php")
}
if cronInstalled {
actions = append(actions, fmt.Sprintf("installed every-%d-minute system cron for %s", clampInterval(opts.IntervalMinutes), owner))
}
if len(actions) == 0 {
return RemediationResult{
Success: true,
Action: fmt.Sprintf("wp-cron already configured for %s", docroot),
Description: "WP-Cron already disabled and system cron already present; no change needed",
}
}
return RemediationResult{
Success: true,
Action: fmt.Sprintf("disable WP-Cron + install system cron for %s", docroot),
Description: strings.Join(actions, "; "),
}
}
// insertDisableWPCron returns wp-config.php bytes with the DISABLE_WP_CRON
// define inserted before the "stop editing" marker, or before the
// wp-settings.php require as a fallback. The second return is false when no
// safe insertion point exists, so the caller can refuse rather than append a
// define into an unfamiliar PHP file.
func insertDisableWPCron(data []byte) ([]byte, bool) {
lines := bytes.Split(data, []byte("\n"))
defineLine := []byte("define( 'DISABLE_WP_CRON', true ); " + wpCronEditMarker)
insertAt := wpCronInsertionLine(lines)
if insertAt < 0 {
return nil, false
}
out := make([][]byte, 0, len(lines)+1)
out = append(out, lines[:insertAt]...)
out = append(out, defineLine)
out = append(out, lines[insertAt:]...)
return bytes.Join(out, []byte("\n")), true
}
func wpCronHasActiveDisableDefine(data []byte) bool {
inBlockComment := false
heredocLabel := ""
for _, line := range strings.Split(string(data), "\n") {
code := wpCronActivePHPCode(line, &inBlockComment, &heredocLabel)
if wpCronDefineRe.MatchString(code) {
return true
}
}
return false
}
func wpCronInsertionLine(lines [][]byte) int {
inBlockComment := false
heredocLabel := ""
fallback := -1
for i, line := range lines {
safeAtLineStart := !inBlockComment && heredocLabel == ""
code := wpCronActivePHPCode(string(line), &inBlockComment, &heredocLabel)
if !safeAtLineStart {
continue
}
if bytes.Contains(bytes.ToLower(line), []byte(wpCronStopMarker)) {
return i
}
if fallback < 0 && strings.Contains(code, "wp-settings.php") {
fallback = i
}
}
return fallback
}
func wpCronActivePHPCode(line string, inBlockComment *bool, heredocLabel *string) string {
if *heredocLabel != "" {
if wpCronEndsHeredoc(line, *heredocLabel) {
*heredocLabel = ""
}
return ""
}
var out strings.Builder
quote := byte(0)
escaped := false
for i := 0; i < len(line); i++ {
if *inBlockComment {
if i+1 < len(line) && line[i] == '*' && line[i+1] == '/' {
*inBlockComment = false
i++
}
continue
}
c := line[i]
if quote != 0 {
out.WriteByte(c)
if escaped {
escaped = false
continue
}
if c == '\\' {
escaped = true
continue
}
if c == quote {
quote = 0
}
continue
}
if c == '\'' || c == '"' {
quote = c
out.WriteByte(c)
continue
}
if i+1 < len(line) && c == '/' && line[i+1] == '*' {
*inBlockComment = true
i++
continue
}
if i+1 < len(line) && c == '/' && line[i+1] == '/' {
break
}
if c == '#' {
break
}
out.WriteByte(c)
}
code := out.String()
if match := wpCronHeredocStartRe.FindStringSubmatch(code); len(match) == 2 {
*heredocLabel = match[1]
}
return code
}
func wpCronEndsHeredoc(line, label string) bool {
trimmed := strings.TrimSpace(line)
return trimmed == label || trimmed == label+";"
}
// installUserWPCron ensures the owner's crontab contains a CSM-managed line
// running wp-cron.php for docroot. It returns false (no error) when the line
// is already present. The crontab is rewritten via a spool file because the
// command runner has no stdin channel; `crontab -u <user> <file>` installs and
// validates it atomically.
func installUserWPCron(owner, docroot string, opts WPCronFixOptions) (bool, error) {
if !validCPUser.MatchString(owner) || owner == "root" {
return false, fmt.Errorf("refusing crontab edit for unexpected account name %q", owner)
}
if !safeWPCronDocroot(docroot) {
return false, fmt.Errorf("refusing crontab edit for unsafe WP-Cron docroot %q", docroot)
}
lock := wpCronCrontabLock(owner)
lock.Lock()
defer lock.Unlock()
// Resolve under the account lock so a caller that waited for another
// crontab edit cannot install a PHP mapping captured before that edit.
var authoritativePHPBin bool
opts.PHPBin, authoritativePHPBin = resolveWPCronPHPBin(owner, docroot, opts)
if !safeCronCommandString(opts.PHPBin) {
return false, fmt.Errorf("refusing crontab edit for unsafe WP-Cron php binary %q", opts.PHPBin)
}
existing := ""
if out, err := cmdExec.RunAllowNonZero("crontab", "-u", owner, "-l"); err == nil {
existing = string(out)
}
if !authoritativePHPBin {
if existingPHPBin := currentManagedWPCronPHPBin(existing, owner, docroot); existingPHPBin != "" {
opts.PHPBin = existingPHPBin
}
}
want := wpCronJobLine(owner, docroot, opts)
var buf bytes.Buffer
switch {
case wpCronUpgradeManagedLine(existing, docroot, want, &buf):
// Stale CSM-managed line rewritten in place (legacy synchronized
// schedule, changed interval, or changed php path).
case crontabHasWPCronJob(existing, docroot):
// Current managed line, or a customer-authored wp-cron entry CSM
// must not fight over.
return false, nil
default:
buf.WriteString(strings.TrimRight(existing, "\n"))
if buf.Len() > 0 {
buf.WriteByte('\n')
}
buf.WriteString(wpCronJobMarker + docroot + "\n")
buf.WriteString(want + "\n")
}
tmp, err := os.CreateTemp("", "csm-wpcron-*")
if err != nil {
return false, fmt.Errorf("create crontab spool: %v", err)
}
tmpPath := tmp.Name()
defer func() { _ = os.Remove(tmpPath) }()
if _, err := tmp.Write(buf.Bytes()); err != nil {
_ = tmp.Close()
return false, fmt.Errorf("write crontab spool: %v", err)
}
if err := tmp.Close(); err != nil {
return false, fmt.Errorf("close crontab spool: %v", err)
}
expected := append([]byte(nil), buf.Bytes()...)
preRecordCrontabSelfWrite(owner, expected)
if _, err := cmdExec.Run("crontab", "-u", owner, tmpPath); err != nil {
forgetSelfWrites(crontabSpoolPaths(owner)...)
return false, fmt.Errorf("crontab install: %v", err)
}
recordCrontabSelfWrite(owner, expected)
return true, nil
}
func preRecordCrontabSelfWrite(owner string, expected []byte) {
for _, p := range crontabSpoolPaths(owner) {
RecordSelfWrite(p, expected)
}
}
// recordCrontabSelfWrite registers the just-installed crontab with the
// self-write ledger so the sensitive-file detectors do not flag CSM's own
// change. The on-disk spool content (cron may normalize it) is what the
// detectors hash, so record the spool file rather than our staged buffer.
func recordCrontabSelfWrite(owner string, expected []byte) {
paths := crontabSpoolPaths(owner)
recorded := ""
for _, p := range paths {
data, err := osFS.ReadFile(p)
if err != nil {
continue
}
if crontabContentEqual(data, expected) || crontabExplainedBy(data, expected) {
RecordSelfWrite(p, data)
recorded = p
break
}
}
for _, p := range paths {
if p != recorded {
forgetSelfWrites(p)
}
}
}
// wpCronSpoolDirs lists where cron daemons keep per-user crontabs (cronie,
// then Debian-style cron). A var so tests can point it at a temp dir.
var wpCronSpoolDirs = []string{"/var/spool/cron", "/var/spool/cron/crontabs"}
func crontabSpoolPaths(owner string) []string {
paths := make([]string, 0, len(wpCronSpoolDirs))
for _, dir := range wpCronSpoolDirs {
paths = append(paths, filepath.Join(dir, owner))
}
return paths
}
// crontabExplainedBy reports whether every executable line on disk came from
// the buffer CSM installed. The cPanel crontab wrapper rewrites the spool after
// `crontab -u` -- it prepends its own SHELL= and drops comments -- so the bytes
// never match what was staged, and refusing to record a self-write then left
// CSM's own change looking like a third party's.
//
// This is not a blanket "trust whatever is on disk": a job line that is not in
// the staged buffer, or an environment assignment that could steer execution,
// means something other than the wrapper touched the file, and the change is
// reported instead of vouched for.
func crontabExplainedBy(got, want []byte) bool {
staged := make(map[string]int)
for _, raw := range strings.Split(string(normalizeCrontabLineEndings(want)), "\n") {
line := strings.TrimSpace(raw)
if line != "" && !strings.HasPrefix(line, "#") {
staged[line]++
}
}
if len(staged) == 0 {
return false
}
for _, raw := range strings.Split(string(normalizeCrontabLineEndings(got)), "\n") {
line := strings.TrimSpace(raw)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if staged[line] > 0 {
staged[line]--
continue
}
if name, val, ok := splitCrontabEnv(line); ok &&
strings.EqualFold(name, "SHELL") &&
val == "/usr/local/cpanel/bin/jailshell" {
continue
}
return false
}
for _, remaining := range staged {
if remaining != 0 {
return false
}
}
return true
}
func crontabContentEqual(got, want []byte) bool {
got = normalizeCrontabLineEndings(got)
want = normalizeCrontabLineEndings(want)
if bytes.Equal(got, want) {
return true
}
return bytes.HasSuffix(want, []byte("\n")) && bytes.Equal(got, bytes.TrimSuffix(want, []byte("\n")))
}
func normalizeCrontabLineEndings(data []byte) []byte {
return bytes.ReplaceAll(data, []byte("\r\n"), []byte("\n"))
}
// wpCronJobLine builds the crontab entry. CLI php is used (not an HTTP hit) so
// the job does not tie up a web worker, which is the load source the finding
// flags. max_execution_time caps a runaway cron pass.
//
// The minute field is staggered per account+docroot: a plain */N schedule is
// wall-clock aligned, so every managed site on the host fires in the same
// second and the load spikes once per interval. The offset hash is
// deterministic so reinstalls are idempotent. flock skips a run while the
// previous one still holds the lock; $HOME is expanded by crond, and the lock
// lives in the account home because /tmp is symlink-attackable.
func wpCronJobLine(owner, docroot string, opts WPCronFixOptions) string {
interval := clampInterval(opts.IntervalMinutes)
php := wpCronPHPBin(owner, docroot, opts)
return fmt.Sprintf(`%s * * * * cd %s && flock -n "$HOME/.csm-wpcron-%08x.lock" %s -d max_execution_time=300 wp-cron.php >/dev/null 2>&1`,
wpCronMinuteField(wpCronStaggerOffset(owner, docroot, interval), interval),
shellQuote(docroot), wpCronLockID(docroot), shellQuote(php))
}
// wpCronStaggerOffset spreads managed sites across the interval. Hashing
// owner+docroot (not just docroot) keeps multi-site accounts spread too.
func wpCronStaggerOffset(owner, docroot string, interval int) int {
h := fnv.New32a()
h.Write([]byte(owner))
h.Write([]byte{0})
h.Write([]byte(docroot))
return int(h.Sum32() % uint32(interval)) // #nosec G115 -- interval clamped to [1,60]
}
func wpCronLockID(docroot string) uint32 {
h := fnv.New32a()
h.Write([]byte(docroot))
return h.Sum32()
}
func wpCronMinuteField(offset, interval int) string {
switch {
case interval <= 1:
return "*"
case interval >= 60:
return strconv.Itoa(offset)
case 60%interval == 0:
return fmt.Sprintf("%d-59/%d", offset, interval)
default:
minutes := make([]int, 0, (60+interval-1)/interval)
for minute := 0; minute < 60; minute += interval {
minutes = append(minutes, (minute+offset)%60)
}
sort.Ints(minutes)
parts := make([]string, 0, len(minutes))
for _, minute := range minutes {
parts = append(parts, strconv.Itoa(minute))
}
return strings.Join(parts, ",")
}
}
// wpCronUpgradeManagedLine rewrites the job line under this docroot's CSM
// marker when it differs from want, writing the full crontab into buf. It
// returns false when there is no marker, the managed line already matches, or
// the line after the marker is not a wp-cron job for this docroot (a crontab
// the customer rearranged is left alone rather than guessed at).
func wpCronUpgradeManagedLine(existing, docroot, want string, buf *bytes.Buffer) bool {
lines := strings.Split(existing, "\n")
marker := wpCronJobMarker + docroot
for i, line := range lines {
if strings.TrimSpace(line) != marker {
continue
}
if i+1 >= len(lines) {
lines = append(lines, want)
writeCrontabLines(buf, lines)
return true
}
job := lines[i+1]
if strings.TrimSpace(job) == "" && i+1 == len(lines)-1 {
lines[i+1] = want
writeCrontabLines(buf, lines)
return true
}
if strings.TrimSpace(job) == want {
return false
}
if cmd := crontabCommand(job); cmd == "" || !commandRunsWPCronForDocroot(cmd, docroot) {
return false
}
lines[i+1] = want
writeCrontabLines(buf, lines)
return true
}
return false
}
func currentManagedWPCronPHPBin(existing, owner, docroot string) string {
lines := strings.Split(existing, "\n")
marker := wpCronJobMarker + docroot
for i, line := range lines {
if strings.TrimSpace(line) != marker || i+1 >= len(lines) {
continue
}
job := strings.TrimSpace(lines[i+1])
managedDocroot, ok := csmManagedWPCronJob(job, owner)
if !ok || managedDocroot != docroot {
return ""
}
match := csmManagedWPCronJobRe.FindStringSubmatch(job)
phpBin, ok := unquoteShellSingle(match[4])
if !ok {
return ""
}
return phpBin
}
return ""
}
func writeCrontabLines(buf *bytes.Buffer, lines []string) {
buf.WriteString(strings.TrimRight(strings.Join(lines, "\n"), "\n"))
buf.WriteByte('\n')
}
func safeWPCronDocroot(docroot string) bool {
return docroot != "" && filepath.IsAbs(docroot) && filepath.Clean(docroot) == docroot && safeCronCommandString(docroot)
}
func safeCronCommandString(s string) bool {
for i := 0; i < len(s); i++ {
if s[i] == '%' || s[i] < 0x20 || s[i] == 0x7f {
return false
}
}
return true
}
// crontabHasWPCronJob reports whether the crontab already runs wp-cron.php for
// docroot, regardless of interval or php path, so re-running the fix is a no-op.
func crontabHasWPCronJob(crontab, docroot string) bool {
docroot = filepath.Clean(docroot)
for _, line := range strings.Split(crontab, "\n") {
command := crontabCommand(line)
if command == "" {
continue
}
if commandRunsWPCronForDocroot(command, docroot) {
return true
}
}
return false
}
func wpCronCrontabLock(owner string) *sync.Mutex {
lock, _ := wpCronCrontabLocks.LoadOrStore(owner, &sync.Mutex{})
return lock.(*sync.Mutex)
}
func wpCronConfigLock(path string) *sync.Mutex {
lock, _ := wpCronConfigLocks.LoadOrStore(path, &sync.Mutex{})
return lock.(*sync.Mutex)
}
func crontabCommand(line string) string {
trimmed := strings.TrimSpace(line)
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
return ""
}
fields := strings.Fields(trimmed)
if len(fields) == 0 || strings.Contains(fields[0], "=") {
return ""
}
if strings.HasPrefix(fields[0], "@") {
if len(fields) < 2 {
return ""
}
return crontabCommandBeforeStdin(trimmed[len(fields[0]):])
}
if len(fields) < 6 {
return ""
}
rest := trimmed
for i := 0; i < 5; i++ {
rest = strings.TrimLeft(rest, " \t")
fieldEnd := strings.IndexAny(rest, " \t")
if fieldEnd < 0 {
return ""
}
rest = rest[fieldEnd:]
}
return crontabCommandBeforeStdin(rest)
}
func crontabCommandBeforeStdin(command string) string {
escaped := false
for i := 0; i < len(command); i++ {
switch {
case escaped:
escaped = false
case command[i] == '\\':
escaped = true
case command[i] == '%':
return strings.TrimSpace(command[:i])
}
}
return strings.TrimSpace(command)
}
func commandRunsWPCronForDocroot(command, docroot string) bool {
if !strings.Contains(command, "wp-cron.php") {
return false
}
words := shellWords(command)
wpCronPath := filepath.Clean(filepath.Join(docroot, "wp-cron.php"))
for _, word := range words {
if cleanShellPathWord(word) == wpCronPath {
return true
}
}
for i, word := range words {
if word != "cd" {
continue
}
j := i + 1
for j < len(words) && strings.HasPrefix(words[j], "-") {
j++
}
if j >= len(words) || filepath.Clean(words[j]) != docroot {
continue
}
for _, later := range words[j+1:] {
if filepath.Base(cleanShellPathWord(later)) == "wp-cron.php" {
return true
}
}
}
return false
}
func cleanShellPathWord(word string) string {
if i := strings.IndexAny(word, "?#"); i >= 0 {
word = word[:i]
}
return filepath.Clean(word)
}
func shellWords(command string) []string {
var words []string
var current strings.Builder
quote := byte(0)
escaped := false
flush := func() {
if current.Len() == 0 {
return
}
words = append(words, current.String())
current.Reset()
}
for i := 0; i < len(command); i++ {
c := command[i]
if quote != 0 {
if escaped {
current.WriteByte(c)
escaped = false
continue
}
if c == '\\' && quote == '"' {
escaped = true
continue
}
if c == quote {
quote = 0
continue
}
current.WriteByte(c)
continue
}
if escaped {
current.WriteByte(c)
escaped = false
continue
}
switch c {
case '\\':
escaped = true
case '\'', '"':
quote = c
case ' ', '\t', ';', '&', '|', '(', ')', '<', '>':
flush()
default:
current.WriteByte(c)
}
}
if escaped {
current.WriteByte('\\')
}
flush()
return words
}
func clampInterval(minutes int) int {
if minutes <= 0 {
return wpCronDefaultIntervalMin
}
if minutes > wpCronMaxIntervalMin {
return wpCronMaxIntervalMin
}
return minutes
}
func detectPHPBin() string {
if p, err := cmdExec.LookPath("php"); err == nil && p != "" {
return p
}
return "/usr/local/bin/php"
}
// fileOwnerName resolves the username that owns the wp-config.php so the cron
// runs as the account, not root. The owner is the source of truth for which
// account this WordPress install belongs to.
func fileOwnerName(info os.FileInfo) (string, error) {
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return "", fmt.Errorf("unsupported file info")
}
if stat.Uid == 0 {
// A customer wp-config.php should never be root-owned; installing a
// cron that runs wp-cron.php as root would be a privilege smell.
return "", fmt.Errorf("refusing to install a root-owned cron for wp-config.php")
}
uid := strconv.FormatUint(uint64(stat.Uid), 10)
u, err := user.LookupId(uid)
if err != nil {
return "", fmt.Errorf("uid %s: %v", uid, err)
}
return u.Username, nil
}
package checks
import (
"context"
"errors"
"io/fs"
"path/filepath"
"strings"
"sync"
"time"
)
// wpInstall is one discovered WordPress installation.
type wpInstall struct {
// ConfigPath is the wp-config.php file; DocRoot is the directory holding it.
ConfigPath string
DocRoot string
// Account owns the install, empty when the path is not attributable.
Account string
Served servedState
}
// wpSkipDirNames keep copies of a site out of discovery. A staging or backup
// tree holds a complete WordPress install, but scanning it means querying a
// database the site does not serve and fixing files nobody reaches.
var wpSkipDirNames = map[string]bool{
"cache": true, "backup": true, "backups": true, "staging": true, ".trash": true,
}
// wpDiscovery is one discovery result together with the coverage gap it
// produced. The gap travels with the result so a cached discovery can credit
// every later caller, not only the one that paid for the walk.
type wpDiscovery struct {
installs []wpInstall
panelDomains map[string][]string
// incomplete records that discovery could not see the whole host. Findings
// from an incomplete scan must survive the cycle's purge.
incomplete bool
}
// apply credits the coverage gap to the check that asked for the installs.
// Discovery is shared; completeness is not -- every consumer needs the gap
// recorded under its own name or its findings are purged as if the scan had
// been complete.
func (d wpDiscovery) apply(ctx context.Context, gapCheck string) {
if !d.incomplete || gapCheck == "" {
return
}
markCheckIncomplete(ctx, gapCheck)
}
// wpInstalls returns every WordPress installation visible to CSM, honouring the
// account scope carried by ctx. gapCheck names the check credited with any
// coverage gap discovery hits.
func wpInstalls(ctx context.Context, gapCheck string) []wpInstall {
installs, _ := wpInstallsWithDomains(ctx, gapCheck)
return installs
}
// wpInstallsWithDomains additionally returns the panel's account-to-domain map,
// which the database scan needs for its tenant-boundary checks. The map is nil
// when the panel's own map could not be parsed completely: a partial ownership
// map turns a legitimate domain into a foreign-host finding.
func wpInstallsWithDomains(ctx context.Context, gapCheck string) ([]wpInstall, map[string][]string) {
d := lookupWPInstalls(ctx, "")
d.apply(ctx, gapCheck)
return d.installs, d.panelDomains
}
// wpInstallsForAccount restricts discovery to one account. Fix, drop and
// re-check paths use it: they run outside a scan context but act on the
// account named by the finding they are resolving.
func wpInstallsForAccount(ctx context.Context, gapCheck, account string) []wpInstall {
if !validAccountName.MatchString(account) {
return nil
}
d := lookupWPInstalls(ctx, account)
d.apply(ctx, gapCheck)
return d.installs
}
type wpInstallCacheKey struct{}
// wpInstallCache memoises discovery for the length of one scan cycle. Nine
// WordPress consumers run per cycle and each used to walk every account home
// for itself.
type wpInstallCache struct {
inventory wpInventoryCycle
started time.Time
mu sync.Mutex
byAccount map[string]wpDiscovery
}
// withWPInstallCache attaches the memo to a scan context. Only the runner does
// this: fix, drop and re-check paths build their own context, so they always
// re-discover, which is what a caller that mutates the tree it just walked
// needs. Invalidation is structural -- no cycle context, no cache.
func withWPInstallCache(ctx context.Context) context.Context {
if ctx == nil {
ctx = context.Background()
}
return context.WithValue(ctx, wpInstallCacheKey{}, &wpInstallCache{
byAccount: make(map[string]wpDiscovery),
started: time.Now(),
})
}
func wpInstallCacheFrom(ctx context.Context) *wpInstallCache {
if ctx == nil {
return nil
}
cache, _ := ctx.Value(wpInstallCacheKey{}).(*wpInstallCache)
return cache
}
// lookupWPInstalls answers from the cycle memo when there is one. The cached
// value carries its coverage gap, which every caller re-applies under its own
// check name.
func lookupWPInstalls(ctx context.Context, account string) wpDiscovery {
cache := wpInstallCacheFrom(ctx)
if cache == nil {
return discoverWPInstalls(ctx, account)
}
key := account
if key == "" {
key = AccountFromContext(ctx)
}
// The lock is held across discovery on purpose: two consumers starting at
// once should cost one walk, not two.
cache.mu.Lock()
defer cache.mu.Unlock()
if cached, ok := cache.byAccount[key]; ok {
return cached
}
d := discoverWPInstalls(ctx, account)
cache.byAccount[key] = d
return d
}
// discoverWPInstalls merges cPanel's document-root map with a walk of the
// account home layout. Neither source is sufficient alone: the map reaches
// roots no walk can predict, and the walk reaches roots the panel has stopped
// serving but whose database is still live.
func discoverWPInstalls(ctx context.Context, account string) wpDiscovery {
scope := account
if scope == "" {
scope = AccountFromContext(ctx)
}
var d wpDiscovery
if scope != "" && !validAccountName.MatchString(scope) {
d.incomplete = true
return d
}
seen := make(map[string]bool)
mappedRoots := make(map[string]bool)
panelDomains := make(map[string][]string)
add := func(path, owner string, state servedState, missingIsIncomplete bool) {
if seen[path] {
return
}
info, err := osFS.Lstat(path)
if err != nil {
if missingIsIncomplete || !errors.Is(err, fs.ErrNotExist) {
d.incomplete = true
}
return
}
if !info.Mode().IsRegular() {
// A symlinked wp-config.php is not scannable here. wp-cli runs as
// root and follows it, so an account pointing its config at another
// account's would have that tenant's database inventoried and
// attributed to this one. The install is dropped and the gap
// recorded instead.
d.incomplete = true
return
}
seen[path] = true
if owner == "" {
_, owner, _ = accountRootOf(path)
}
d.installs = append(d.installs, wpInstall{
ConfigPath: path,
DocRoot: filepath.Dir(path),
Account: owner,
Served: state,
})
}
// cPanel publishes its actual domain-to-document-root map. It is
// authoritative for served roots and reaches layouts the home walk below
// cannot see, so it is consulted first.
vhostData, vhostErr := osFS.ReadFile(userdataDomainsPath)
domainMapComplete := false
switch {
case vhostErr == nil:
vhosts, complete := parseUserdataDomainRootsChecked(string(vhostData))
wildcardVhosts, wildcardComplete := parseWildcardUserdataDomainRootsChecked(string(vhostData))
vhosts = append(vhosts, wildcardVhosts...)
domainMapComplete = complete && wildcardComplete && len(vhosts) > 0
if !domainMapComplete {
d.incomplete = true
}
domainOwners := make(map[string]string, len(vhosts))
for _, vh := range vhosts {
root := filepath.Clean(vh.docroot)
if !docrootBelongsToCPanelUser(root, vh.user) {
d.incomplete = true
domainMapComplete = false
continue
}
wildcard := strings.HasPrefix(vh.domain, "*.")
domain := normalizeHost(strings.TrimPrefix(vh.domain, "*."))
if domain == "" {
d.incomplete = true
domainMapComplete = false
} else {
domainKey := domain
if wildcard {
domainKey = "*." + domain
}
owner, exists := domainOwners[domainKey]
if exists && owner != vh.user {
// The map is meant to have one authoritative owner per
// domain. An ambiguous owner cannot safely support a
// tenant-boundary check.
d.incomplete = true
domainMapComplete = false
} else if !exists {
domainOwners[domainKey] = vh.user
panelDomains[vh.user] = append(panelDomains[vh.user], domainKey)
}
}
if scope != "" && vh.user != scope {
continue
}
wpConfig, err := canonicalWPInstallPath(filepath.Join(root, "wp-config.php"))
if err != nil {
d.incomplete = true
continue
}
mappedRoots[wpConfig] = true
add(wpConfig, vh.user, servedByPanel, false)
}
case vhostMapFailureIsIncomplete(vhostErr):
d.incomplete = true
}
if domainMapComplete {
d.panelDomains = panelDomains
}
// The served map is not sufficient on its own. A document root the panel
// has stopped serving still holds a live database, and the compromise this
// scan was widened for sat in exactly such a root -- absent from the domain
// map, from /etc/userdomains, and from vhost userdata alike. Re-pointing the
// domain publishes it again, so the home layout is walked whatever the panel
// says. Anything the map did not name is not served today, but only when the
// map could be read at all.
homeState := notServed
if !domainMapComplete {
homeState = servedUnknown
}
globScope := scope
if globScope == "" {
globScope = "*"
}
if scope != "" {
// The primary document root is checked by name, not by glob. A fixer
// asked about one account must not lose that account's main install
// because the glob returned nothing, and a missing file here is an
// account without WordPress, not a coverage gap.
primary := filepath.Join(accountHomeDir(scope), "public_html", "wp-config.php")
if _, err := osFS.Lstat(primary); err == nil {
canonical, resolveErr := canonicalWPInstallPath(primary)
if resolveErr != nil {
d.incomplete = true
} else {
add(canonical, scope, homeState, false)
}
} else if !errors.Is(err, fs.ErrNotExist) {
d.incomplete = true
}
}
for _, pattern := range []string{
filepath.Join(globScope, "public_html", "wp-config.php"),
filepath.Join(globScope, "public_html", "*", "wp-config.php"),
filepath.Join(globScope, "*", "wp-config.php"),
} {
matches, err := accountHomeGlob(pattern)
if err != nil {
d.incomplete = true
}
for _, path := range matches {
// A symlinked document root (cPanel's www -> public_html) resolves
// to its target first, so one install is not discovered twice and
// an alias is not mistaken for a directory that is never a root.
path, err = canonicalWPInstallPath(path)
if err != nil {
d.incomplete = true
continue
}
if seen[path] || skipWPDiscoveryPath(path) {
continue
}
state := homeState
if mappedRoots[path] {
// Preserve the panel's declaration even if the first Lstat
// failed and the file appeared before the home walk.
state = servedByPanel
}
add(path, "", state, true)
}
}
return d
}
// skipWPDiscoveryPath rejects candidates that are not document roots: account
// data directories, dot-directories, and backup, cache or staging copies.
func skipWPDiscoveryPath(path string) bool {
dir := filepath.Base(filepath.Dir(path))
if nonDocRootDirs[dir] || strings.HasPrefix(dir, ".") {
return true
}
home := wpInstallAccountRoot(path)
if home == "" {
return true
}
rel, err := filepath.Rel(home, path)
if err != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
return true
}
// Skip names apply inside the account, not to the account name itself.
// cPanel permits accounts named backup, backups, cache, and staging.
for _, part := range strings.Split(rel, string(filepath.Separator)) {
if wpSkipDirNames[strings.ToLower(part)] {
return true
}
}
return false
}
// wpInstallConfigPaths projects installs to their wp-config.php paths for
// callers that work in paths rather than installs.
func wpInstallConfigPaths(installs []wpInstall) []string {
out := make([]string, 0, len(installs))
for _, in := range installs {
out = append(out, in.ConfigPath)
}
return out
}
package checks
import (
"container/heap"
"context"
"crypto/sha256"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/jstaint"
"github.com/pidginhost/csm/internal/phptaint"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/yara"
)
var activeYARABackend = yara.Active
var yaraAvailable = yara.Available
// yaraDeepScanMu serializes the persisted host-wide cursors. Manual baseline
// scans can overlap the daemon's scheduled scan; without serialization, the
// older window can overwrite a newer cursor after it returns.
var yaraDeepScanMu sync.Mutex
// yaraDeepNow is indirected so tests can drive the soft-deadline clock.
var yaraDeepNow = time.Now
// yaraDeepDeadlineMargin is how much of the runner budget the walk leaves
// unused so a partial run returns its findings before the context deadline.
// A check that overruns its budget gets every returned finding dropped by
// the runner, so stopping early is the only way partial coverage survives.
const yaraDeepDeadlineMargin = 45 * time.Second
// yaraDeepFullCycleStale bounds how long rolling coverage may go without
// completing a full pass before the check surfaces a warning.
const yaraDeepFullCycleStale = 30 * 24 * time.Hour
// yaraDeepCursorCheck is the host-scope scan-cursor key (account "") under
// which the rolling deep-scan records the YARA consumer's progress.
const yaraDeepCursorCheck = "yara_deep"
// deepScanConsumer is one consumer of the shared ordered deep-content walk.
// Each consumer resumes from its own persisted cursor and advances it
// independently, so a missing backend cannot stall another analyzer's
// coverage and one analyzer's gap cannot move another's resume point.
type deepScanConsumer struct {
name string
dispatch bool
resume string
lastScanned string
cur store.ScanCursorRecord
// staleCheck and label name this consumer in its own staleness finding.
// Coverage degradation must be reported under the consumer that actually
// degraded, never under a sibling's identity.
staleCheck string
label string
}
// wants reports whether this consumer still needs the given path this cycle.
func (c *deepScanConsumer) wants(path string) bool {
return c.dispatch && (c.resume == "" || path > c.resume)
}
// advance is monotonic-max: a leading consumer keeps its resume point while
// the lagging one catches up over the shared walk.
func (c *deepScanConsumer) advance(path string) {
if c.dispatch && path > c.lastScanned {
c.lastScanned = path
}
}
// resetDeepScanCursor clears a disabled consumer's persisted cursor so
// re-enabling it starts a full cycle instead of resuming a stale window.
func resetDeepScanCursor(db *store.DB, check string) {
if db == nil {
return
}
rec, ok, err := db.GetScanCursor("", check)
if err != nil || !ok {
return
}
if rec.LastPath == "" && rec.WrappedAt.IsZero() && rec.LastFullCycleTS.IsZero() {
return
}
var next store.ScanCursorRecord
next.Check = check
if err := db.PutScanCursor(next); err != nil {
fmt.Fprintf(os.Stderr, "%s: cursor reset: %v\n", check, err)
}
}
// scheduledYARADetails describes one deep-scan match. A loader is only half
// a backdoor: the payload it pulls in lives elsewhere and survives a clean-up
// of the matched file alone, so any non-executable file this content includes
// is named alongside the rule.
func scheduledYARADetails(ruleName string, content []byte) string {
return fmt.Sprintf("Scheduled deep scan matched YARA rule %s", ruleName) +
signatures.ReferencedPayloadDetail(content)
}
type yaraDeepScanEntry struct {
path string
sortKey string
info os.FileInfo
err error
inspected bool
}
type yaraDeepScanHeap []yaraDeepScanEntry
func (h yaraDeepScanHeap) Len() int { return len(h) }
func (h yaraDeepScanHeap) Less(i, j int) bool { return h[i].sortKey < h[j].sortKey }
func (h yaraDeepScanHeap) Swap(i, j int) { h[i], h[j] = h[j], h[i] }
func (h *yaraDeepScanHeap) Push(value any) { *h = append(*h, value.(yaraDeepScanEntry)) }
func (h *yaraDeepScanHeap) Pop() any {
old := *h
last := len(old) - 1
value := old[last]
old[last] = yaraDeepScanEntry{}
*h = old[:last]
return value
}
// CheckYARADeep is the shared rolling deep-content walk. It reads each file
// once and dispatches the same in-memory snapshot to three consumers with
// independent cursors and completion records: the YARA backend and the JS and
// PHP taint analyzers. A missing, disabled, or failed consumer neither prevents
// nor controls another consumer's progress.
func CheckYARADeep(ctx context.Context, cfg *config.Config, st *state.Store) []alert.Finding {
yaraDeepScanMu.Lock()
defer yaraDeepScanMu.Unlock()
if ctx.Err() != nil {
return nil
}
db := store.Global()
yaraConsumer := &deepScanConsumer{name: yaraDeepCursorCheck, staleCheck: "yara_scan_incomplete", label: "YARA"}
jsConsumer := &deepScanConsumer{name: jsTaintDeepCursorCheck, staleCheck: "js_taint_scan_incomplete", label: "JS taint"}
phpConsumer := &deepScanConsumer{name: phpTaintDeepCursorCheck, staleCheck: "php_taint_scan_incomplete", label: "PHP taint"}
consumers := []*deepScanConsumer{yaraConsumer, jsConsumer, phpConsumer}
yaraOff := yaraDeepConsumerDisabled(cfg)
jsOff := jsTaintDeepConsumerDisabled(cfg)
jsConsumer.dispatch = !jsOff
phpOff := phpTaintDeepConsumerDisabled(cfg)
phpReady := phpTaintAnalyzerReady()
// An absent isolated analyzer is not a coverage gap: it means the feature
// is not active on this host. Dispatching anyway would record every
// candidate as unexamined on every scan.
phpConsumer.dispatch = !phpOff && phpReady
if !phpOff && !phpReady {
// Keep prior state while the feature is inactive. Treating an inactive
// owner as completed would purge findings without examining their files.
markCheckIncomplete(ctx, logicalOwnerPHPTaintDeep)
}
// Disabled-consumer resets are persistent scan progress too. Commit them
// only on a normal return path so a hard-canceled shared walk writes no
// cursor state for any consumer.
resetDisabledCursors := func() {
if yaraOff {
resetDeepScanCursor(db, yaraDeepCursorCheck)
}
if jsOff {
resetDeepScanCursor(db, jsTaintDeepCursorCheck)
}
if phpOff {
resetDeepScanCursor(db, phpTaintDeepCursorCheck)
}
}
var findings []alert.Finding
backend := activeYARABackend()
yaraReady := backend != nil && backend.RuleCount() > 0
yaraConsumer.dispatch = !yaraOff && yaraReady
if !yaraOff && !yaraReady {
// A missing backend is a YARA coverage gap, never a JS one: mark only
// the YARA owner incomplete and leave its cursor unchanged so it can
// neither purge nor skip its own unscanned range, while the JS pass
// below still covers its cycle.
markCheckIncomplete(ctx, "yara_deep")
if yaraAvailable() {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "yara_scan_incomplete",
Message: "YARA deep scan could not start because no compiled rules are available",
})
}
}
if !yaraConsumer.dispatch && !jsConsumer.dispatch && !phpConsumer.dispatch {
if ctx.Err() != nil {
return nil
}
resetDisabledCursors()
return findings
}
maxBytes := int64(FullScanMaxFileBytes(cfg))
jsMaxBytes := int64(jstaint.MaxSourceBytes)
phpMaxBytes := int64(phptaint.MaxSourceBytes)
for _, c := range consumers {
if !c.dispatch || db == nil {
continue
}
rec, ok, err := db.GetScanCursor("", c.name)
if err != nil {
fmt.Fprintf(os.Stderr, "%s: cursor read: %v\n", c.name, err)
} else if ok {
c.cur = rec
}
c.resume = c.cur.LastPath
c.lastScanned = c.resume
}
// The shared walk starts at the earliest dispatchable resume point; a
// consumer whose cursor is ahead skips already-covered paths via wants()
// while the lagging consumer catches up on the same snapshot reads.
walkResume := ""
firstResume := true
for _, c := range consumers {
if !c.dispatch {
continue
}
if firstResume || c.resume < walkResume {
walkResume = c.resume
firstResume = false
}
}
var softDeadline time.Time
if deadline, ok := ctx.Deadline(); ok {
softDeadline = deadline.Add(-yaraDeepDeadlineMargin)
}
outOfTime := func() bool {
return !softDeadline.IsZero() && !yaraDeepNow().Before(softDeadline)
}
// The collector owns both the count and the per-status breakdown, so the
// finding can say what the gaps were instead of only how many.
var observedPath string
var observedInfo os.FileInfo
observePathIdentity := func(path string, info os.FileInfo) {
observedPath = path
observedInfo = info
}
resolveObservedAliases := func(path string) ([]string, bool) {
if path != observedPath || observedInfo == nil {
return coveragePathAliases(path), false
}
return stableCoveragePathAliases(path, observedInfo)
}
yaraGaps := newYARAGapCollector()
yaraGaps.resolveAliases = resolveObservedAliases
recordYARAPathGap := func(path, status, _ string) {
recordCoverageGapPaths(ctx, "yara_deep", yaraGaps.record(path, status))
}
recordYARAUnknownRangeGap := func(detail string) {
yaraGaps.recordUnknownRange(detail)
}
pathFailureStillAttributable := func(path string) (os.FileInfo, bool) {
info, err := osFS.Lstat(path)
// Every metadata error is ambiguous, including ENOENT: a delete can race
// this observation and replace the path with a directory immediately after
// it. Only an observed non-directory, non-symlink entry is narrow enough to
// preserve by exact path.
return info, err == nil && !info.IsDir() && info.Mode()&os.ModeSymlink == 0 &&
observedPath == path && observedInfo != nil && os.SameFile(observedInfo, info)
}
recordYARAScanGap := func(path string, scanErr error) {
// Inline IPC overflow retries by path. If that path changed into a
// directory or symlink before the retry failed, the unscanned range is no
// longer just the original file: children or a target subtree may now be
// hidden. The post-failure Lstat is inherently racy, so every metadata
// error takes the conservative unknown-range path.
info, attributable := pathFailureStillAttributable(path)
if !attributable {
recordYARAUnknownRangeGap(fmt.Sprintf("scanning %s failed after the path changed: %v", path, scanErr))
return
}
observePathIdentity(path, info)
recordYARAPathGap(path, "scan_error", fmt.Sprintf("scanning %s: %v", path, scanErr))
}
jsGaps := newJSTaintGapCollector()
phpGaps := newPHPTaintGapCollector()
jsGaps.resolveAliases = resolveObservedAliases
phpGaps.resolveAliases = resolveObservedAliases
jsGaps.recordCoverage = func(paths []string) {
recordCoverageGapPaths(ctx, logicalOwnerJSTaintDeep, paths)
}
phpGaps.recordCoverage = func(paths []string) {
recordCoverageGapPaths(ctx, logicalOwnerPHPTaintDeep, paths)
}
stoppedEarly := false
sep := string(filepath.Separator)
subtreePrefix := func(path string) string {
path = filepath.Clean(path)
if strings.HasSuffix(path, sep) {
return path
}
return path + sep
}
advanceAll := func(path string) {
for _, c := range consumers {
c.advance(path)
}
}
// subtreeCovered reports whether every path under dir sorts before the
// earliest resume point. Children all share the dir+sep prefix, so when
// resume does not itself start with that prefix, any child compares to
// resume exactly as the prefix does. Comparing the bare dir path instead
// would wrongly skip siblings like "ab/" when the cursor sits at "ab.zz"
// ('/' sorts after '.'). An exact prefix cursor records a subtree already
// accounted for this cycle, including an empty or unreadable directory.
subtreeCoveredAfter := func(resume, dir string) bool {
prefix := subtreePrefix(dir)
return resume != "" && (resume == prefix || (!strings.HasPrefix(resume, prefix) && prefix < resume))
}
subtreeCovered := func(dir string) bool {
return subtreeCoveredAfter(walkResume, dir)
}
consumerWantsSubtree := func(c *deepScanConsumer, dir string) bool {
return c.dispatch && !subtreeCoveredAfter(c.resume, dir)
}
var scanDir func(string)
scanDir = func(dir string) {
if ctx.Err() != nil || stoppedEarly {
return
}
if outOfTime() {
stoppedEarly = true
return
}
entries, err := osFS.ReadDir(dir)
if err != nil {
// The affected paths are unknowable, but a consumer that already
// covered this subtree must not inherit another consumer's gap.
yaraWants := consumerWantsSubtree(yaraConsumer, dir)
jsWants := consumerWantsSubtree(jsConsumer, dir)
phpWants := consumerWantsSubtree(phpConsumer, dir)
if yaraWants {
recordYARAUnknownRangeGap(fmt.Sprintf("reading %s: %v", dir, err))
}
if jsWants {
jsGaps.recordUnknownRange(fmt.Sprintf("reading %s: %v", dir, err))
}
if phpWants {
phpGaps.recordUnknownRange(fmt.Sprintf("reading %s: %v", dir, err))
}
advanceAll(subtreePrefix(dir))
return
}
// A directory's candidate paths start at dir+separator, not at the
// bare directory name. Bare-name order visits ab/ before ab.zz even
// though every child under ab/ sorts after ab.zz. The heap starts
// with each bare path as a lower bound, then requeues directories
// under their path+separator key after Lstat reveals their type.
ordered := make(yaraDeepScanHeap, len(entries))
for i, entry := range entries {
ordered[i].path = filepath.Join(dir, entry.Name())
ordered[i].sortKey = ordered[i].path
}
heap.Init(&ordered)
for ordered.Len() > 0 {
if ctx.Err() != nil || stoppedEarly {
return
}
if outOfTime() {
stoppedEarly = true
return
}
item := heap.Pop(&ordered).(yaraDeepScanEntry)
if !item.inspected {
item.info, item.err = osFS.Lstat(item.path)
item.inspected = true
}
path := item.path
if item.err == nil && item.info.Mode()&os.ModeSymlink == 0 && item.info.IsDir() {
if item.sortKey == path {
item.sortKey = subtreePrefix(path)
heap.Push(&ordered, item)
continue
}
if subtreeCovered(path) {
continue
}
if outOfTime() {
stoppedEarly = true
return
}
scanDir(path)
if ctx.Err() == nil && !stoppedEarly {
advanceAll(item.sortKey)
}
continue
}
yaraWants := yaraConsumer.wants(path)
jsWants := jsConsumer.wants(path)
phpWants := phpConsumer.wants(path)
if !yaraWants && !jsWants && !phpWants {
continue
}
if item.err != nil {
// A failed Lstat may hide a directory, so the gap range is
// unknowable for either path-based carry-forward.
if yaraWants {
recordYARAUnknownRangeGap(fmt.Sprintf("inspecting %s: %v", path, item.err))
}
if jsWants {
jsGaps.recordUnknownRange(fmt.Sprintf("inspecting %s: %v", path, item.err))
}
if phpWants {
phpGaps.recordUnknownRange(fmt.Sprintf("inspecting %s: %v", path, item.err))
}
advanceAll(path)
continue
}
info := item.info
if info.Mode()&os.ModeSymlink != 0 || !info.Mode().IsRegular() || info.Size() == 0 {
advanceAll(path)
continue
}
observePathIdentity(path, info)
if yaraWants && info.Size() > maxBytes {
// Oversize files are intentionally not opened, but they still
// count as covered progress for this rolling cycle. The gap
// advances only its own consumer: a soft-deadline stop right
// after this gate must not move the other consumer past a
// file it never received.
recordYARAPathGap(path, "oversize", fmt.Sprintf("%s exceeds the %d-byte scan limit", path, maxBytes))
yaraConsumer.advance(path)
yaraWants = false
}
if phpWants && info.Size() > phpMaxBytes {
// An oversize file is only a PHP coverage gap if it could have
// been PHP. The deep walk hands every readable file to this
// consumer and the analyzer's own pre-filter normally rejects
// the rest instantly -- but that pre-filter never runs on a
// file too large to send, so without this check every large
// error_log, JSON blob and media file on the host is reported
// as PHP content we failed to examine. Measured on a live
// host: 617 of 660 recorded gaps in one scan, which buries the
// handful that are real.
// Past the soft deadline the walk is stopping and must do
// no further I/O, so the peek is skipped and the gap is
// recorded unexamined -- the same answer an unreadable
// file gets.
if outOfTime() || phpFileMayBePHP(path, info) {
phpGaps.record(path, phptaint.StatusOversize.String())
}
phpConsumer.advance(path)
phpWants = false
}
if jsWants && info.Size() > jsMaxBytes {
// An oversize file is only a JS coverage gap if it could
// have been JavaScript. The deep walk hands every readable
// file to this consumer and the analyzer's own pre-filter
// normally rejects the rest instantly -- but that pre-filter
// never runs on a file too large to send, so without this
// check every large image, archive and database on the host
// is reported as JavaScript we failed to examine. This gate
// was missing while the PHP one beside it existed.
//
// This branch used to decide on metadata alone. Peeking
// costs an open, so as with PHP above it is skipped once the
// soft deadline has passed and the gap is recorded
// unexamined.
if outOfTime() || jsFileMayBeJS(path, info) {
jsGaps.record(path, jstaint.StatusOversize.String())
}
jsConsumer.advance(path)
jsWants = false
}
if !yaraWants && !jsWants && !phpWants {
continue
}
if outOfTime() {
stoppedEarly = true
return
}
// Advance before opening so permanently unreadable or unscannable
// files cannot wedge the cursor in place.
advanceAll(path)
file, err := osFS.Open(path)
if err != nil {
detail := fmt.Sprintf("opening %s: %v", path, err)
info, attributable := pathFailureStillAttributable(path)
if attributable {
observePathIdentity(path, info)
if yaraWants {
recordYARAPathGap(path, "open_error", detail)
}
if jsWants {
jsGaps.record(path, "read_error")
}
if phpWants {
phpGaps.record(path, "read_error")
}
} else {
if yaraWants {
recordYARAUnknownRangeGap(detail)
}
if jsWants {
jsGaps.recordUnknownRange(detail)
}
if phpWants {
phpGaps.recordUnknownRange(detail)
}
}
continue
}
openedInfo, statErr := file.Stat()
if statErr != nil {
_ = file.Close()
if yaraWants {
recordYARAUnknownRangeGap(fmt.Sprintf("inspecting opened path %s: %v", path, statErr))
}
if jsWants {
jsGaps.recordUnknownRange(fmt.Sprintf("inspecting opened path %s: %v", path, statErr))
}
if phpWants {
phpGaps.recordUnknownRange(fmt.Sprintf("inspecting opened path %s: %v", path, statErr))
}
continue
}
if !openedInfo.Mode().IsRegular() {
_ = file.Close()
if openedInfo.IsDir() {
// Lstat admitted a regular file, but Open bound a directory.
// Its children were never enumerated, so retaining only the
// exact path would purge findings from an unknown subtree.
if yaraWants {
recordYARAUnknownRangeGap(fmt.Sprintf("%s changed into a directory while it was being opened", path))
}
if jsWants {
jsGaps.recordUnknownRange(fmt.Sprintf("%s changed into a directory while it was being opened", path))
}
if phpWants {
phpGaps.recordUnknownRange(fmt.Sprintf("%s changed into a directory while it was being opened", path))
}
} else {
if yaraWants {
recordYARAPathGap(path, "changed_during_read", fmt.Sprintf("%s changed while it was being opened", path))
}
if jsWants {
jsGaps.record(path, "changed_during_read")
}
if phpWants {
phpGaps.record(path, "changed_during_read")
}
}
continue
}
observePathIdentity(path, openedInfo)
if yaraWants && openedInfo.Size() > maxBytes {
recordYARAPathGap(path, "changed_during_read", fmt.Sprintf("%s changed while it was being opened", path))
yaraWants = false
}
if phpWants && openedInfo.Size() > phpMaxBytes {
phpGaps.record(path, "changed_during_read")
phpWants = false
}
if jsWants && openedInfo.Size() > jsMaxBytes {
jsGaps.record(path, "changed_during_read")
jsWants = false
}
if !yaraWants && !jsWants && !phpWants {
_ = file.Close()
continue
}
readCap := int64(0)
if yaraWants {
readCap = maxBytes
}
if jsWants && jsMaxBytes > readCap {
readCap = jsMaxBytes
}
if phpWants && phpMaxBytes > readCap {
readCap = phpMaxBytes
}
data, readErr := io.ReadAll(io.LimitReader(file, readCap+1))
closeErr := file.Close()
if readErr != nil || closeErr != nil || int64(len(data)) > readCap {
if yaraWants {
recordYARAPathGap(path, "read_error", fmt.Sprintf("reading %s failed or exceeded the scan limit", path))
}
if jsWants {
jsGaps.record(path, "read_error")
}
if phpWants {
phpGaps.record(path, "read_error")
}
continue
}
if yaraWants && int64(len(data)) > maxBytes {
recordYARAPathGap(path, "read_error", fmt.Sprintf("reading %s failed or exceeded the scan limit", path))
yaraWants = false
}
if phpWants && int64(len(data)) > phpMaxBytes {
phpGaps.record(path, "changed_during_read")
phpWants = false
}
if jsWants && int64(len(data)) > jsMaxBytes {
jsGaps.record(path, "changed_during_read")
jsWants = false
}
fingerprint := sha256.Sum256(data)
contentSHA256 := fmt.Sprintf("%x", fingerprint)
if jsWants {
findings = append(findings, analyzeJSTaintSnapshot(ctx, path, contentSHA256, data, jsGaps)...)
}
if phpWants {
findings = append(findings, analyzePHPTaintSnapshot(ctx, path, contentSHA256, data, phpGaps)...)
phpConsumer.advance(path)
}
if !yaraWants {
continue
}
yaraSHA256 := contentSHA256
// Within the deep-scan size budget but possibly too large for one
// inline IPC frame once JSON base64 expands it; the helper retries
// by path so a payload cannot be hidden behind the inline ceiling.
// JS analysis is unaffected: it already consumed the snapshot above.
// #nosec G115 -- maxBytes is FullScanMaxFileBytes (int) widened to int64; the round-trip back to int is lossless.
matches, scannedSHA, scanErr := yara.ScanContentOrPathChecked(backend, path, data, int(maxBytes))
if scannedSHA != "" {
yaraSHA256 = scannedSHA
}
if scanErr != nil {
recordYARAScanGap(path, scanErr)
continue
}
for _, match := range matches {
finding := alert.Finding{
Severity: yaraMatchSeverity(match.Meta["severity"]),
Check: "yara_match_scheduled",
Message: fmt.Sprintf("YARA rule match [%s]: %s", match.RuleName, path),
Details: scheduledYARADetails(match.RuleName, data),
FilePath: path,
ContentSHA256: yaraSHA256,
DetectLogic: ContentDetectionVersion(),
}
findings = append(findings, finding)
}
}
}
roots := ResolveWebRoots(cfg)
normalizedRoots := roots[:0]
seenRoots := make(map[string]struct{}, len(roots))
for _, root := range roots {
root = filepath.Clean(root)
if _, exists := seenRoots[root]; exists {
continue
}
seenRoots[root] = struct{}{}
normalizedRoots = append(normalizedRoots, root)
}
roots = normalizedRoots
sort.Slice(roots, func(i, j int) bool { return subtreePrefix(roots[i]) < subtreePrefix(roots[j]) })
for _, root := range roots {
if ctx.Err() != nil || stoppedEarly {
break
}
if outOfTime() {
stoppedEarly = true
break
}
if subtreeCovered(root) {
continue
}
scanDir(root)
if ctx.Err() == nil && !stoppedEarly {
advanceAll(subtreePrefix(root))
}
}
if ctx.Err() != nil {
// The runner drops every finding a check returns after its budget
// expired, so there is nothing worth reporting; leave every cursor
// untouched and let the next run redo this window.
return nil
}
resetDisabledCursors()
now := yaraDeepNow().UTC()
if db != nil {
for _, c := range consumers {
if !c.dispatch {
continue
}
var next store.ScanCursorRecord
next.Check = c.name
if stoppedEarly {
next.LastPath = c.lastScanned
next.LastFullCycleTS = c.cur.LastFullCycleTS
next.WrappedAt = c.cur.WrappedAt
if next.WrappedAt.IsZero() {
next.WrappedAt = now
}
} else {
next.LastFullCycleTS = now
}
if err := db.PutScanCursor(next); err != nil {
fmt.Fprintf(os.Stderr, "%s: cursor write: %v\n", c.name, err)
}
}
}
// Rolling-cycle staleness is per consumer. All three keep their own cursor
// and their own WrappedAt, so the warning belongs to whichever cycle
// actually stalled. This was read from the YARA cursor alone, which left a
// stalled PHP or JS cycle silent, and silenced the warning entirely when
// yara_deep was disabled -- a supported configuration.
//
// Enumerated once so the shape is not rediscovered one field at a time.
// Persisted per-consumer state:
// cur.LastPath becomes resume and lastScanned. walkResume is the minimum
// active resume; wants and consumerWantsSubtree still gate each consumer,
// and advance plus the cursor write update only that consumer's progress.
// cur.WrappedAt starts that consumer's partial cycle and drives only its
// staleCheck/label warning here.
// cur.LastFullCycleTS is preserved or stamped in that consumer's record.
// Runtime per-consumer state:
// dispatch comes from that consumer's disable/readiness gates and controls
// its cursor, analysis, completion mark, status finding, and carry-forward.
// yaraWants/jsWants/phpWants and the three size limits govern only the
// matching consumer. The shared read cap is their maximum, so no sibling
// can reduce another's snapshot.
// incomplete/firstIncomplete belong only to YARA; jsGaps belongs only to
// JS; phpGaps belongs only to PHP. Each is
// reported and used for completion/carry-forward only in its owner block.
// Genuinely shared state:
// roots, path order, the context/soft deadline, opened snapshots, and
// stoppedEarly describe the one walk. After a consumer's resume point all
// active consumers have the same remaining ordered path space, so a walk
// that stops early cannot have completed one active consumer's cycle while
// leaving another active consumer's cycle unfinished.
// Known and deliberately out of scope: php_content_rolling.go writes
// WrappedAt for its per-account cursor and nothing reads it, so that scan
// has no staleness warning of its own.
if stoppedEarly {
for _, c := range consumers {
if !c.dispatch || c.cur.WrappedAt.IsZero() || now.Sub(c.cur.WrappedAt) <= yaraDeepFullCycleStale {
continue
}
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: c.staleCheck,
Message: fmt.Sprintf("Rolling %s deep scan has not completed a full pass since %s", c.label, c.cur.WrappedAt.Format("2006-01-02")),
Details: "Each run advances the cursor inside its time budget; a pass this stale means the budget is too small for the content volume.",
})
}
}
if yaraConsumer.dispatch {
// Same contract the PHP and JS consumers use. Only a partial window or
// an unknown-range loss suppresses the purge; a gap that names a file
// is handled by carrying that file's prior finding forward, so the
// scan can still retire everything it did examine. Reporting a bare
// count here is what made one permanently unreadable file look
// identical to a lost subtree, and froze every YARA finding on a host.
yaraPartial := stoppedEarly || yaraConsumer.resume != "" || yaraGaps.pathsIncomplete()
if yaraPartial {
markCheckIncomplete(ctx, "yara_deep")
}
if !yaraGaps.empty() {
findings = append(findings, yaraGaps.finding())
}
if !yaraPartial && st != nil {
// Keep the stable carried snapshots in-band so alert deduplication can
// retain its normal ongoing-finding re-alert cadence. Daemon callers
// also publish the path set for race-safe latest-state persistence.
findings = append(findings, carryForwardYARAFindings(st.LatestFindings(), yaraGaps)...)
}
}
if phpConsumer.dispatch {
phpPartial := stoppedEarly || phpConsumer.resume != "" || phpGaps.pathsIncomplete()
if phpPartial {
// Only a partial or unknown-range window suppresses the normal
// purge; a known-path gap is handled by the carry-forward below.
markCheckIncomplete(ctx, logicalOwnerPHPTaintDeep)
}
if !phpGaps.empty() {
findings = append(findings, phpGaps.findings()...)
}
if !phpPartial && st != nil {
// Path-specific carry-forward: this run is eligible to replace the
// PHP finding set, so a known-path coverage gap must re-emit that
// path's prior finding or the purge would clear it. A completed
// negative or missing path stays cleared.
findings = append(findings, carryForwardPHPTaintFindings(st.LatestFindings(), phpGaps)...)
}
}
if jsConsumer.dispatch {
jsPartial := stoppedEarly || jsConsumer.resume != "" || jsGaps.pathsIncomplete()
if jsPartial {
// Only a partial or unknown-range window suppresses the normal JS
// purge; a known-path gap is handled by the carry-forward below.
markCheckIncomplete(ctx, logicalOwnerJSTaintDeep)
}
if !jsGaps.empty() {
findings = append(findings, jsGaps.finding())
}
if !jsPartial && st != nil {
// Path-specific carry-forward: this run is eligible to replace
// the JS finding set, so a known-path coverage gap must re-emit
// that path's prior finding or the purge would clear it. A
// completed negative or missing path stays cleared.
findings = append(findings, carryForwardJSTaintFindings(st.LatestFindings(), jsGaps)...)
}
}
return findings
}
func yaraMatchSeverity(value string) alert.Severity {
switch strings.ToLower(value) {
case "warning", "low", "medium":
return alert.Warning
case "high":
return alert.High
default:
return alert.Critical
}
}
package checks
import (
"fmt"
"sort"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// maxYARAGapPaths bounds the exact-path set a single run retains. Past it the
// count keeps rising but the paths are no longer authoritative, which forces
// the run partial rather than letting a purge clear a path it can no longer
// name.
const maxYARAGapPaths = 5000
const yaraGapExampleMaxBytes = 200
// yaraGapCollector records what a YARA deep scan walked but could not examine.
// It mirrors the PHP and JS taint collectors deliberately: those two consumers
// already distinguish a gap that names a file from one that loses an unknown
// range, and YARA reporting only a bare count is what made a permanently
// unreadable file indistinguishable from a lost subtree.
type yaraGapCollector struct {
paths map[string]struct{}
pathAliases map[string]struct{}
aliasesByPath map[string][]string
resolveAliases func(string) ([]string, bool)
byStatus map[string]int
example map[string]string
unknown int
unknownExample string
pathsTruncated bool
}
func newYARAGapCollector() *yaraGapCollector {
return &yaraGapCollector{
paths: map[string]struct{}{},
pathAliases: map[string]struct{}{},
aliasesByPath: map[string][]string{},
byStatus: map[string]int{},
example: map[string]string{},
}
}
// record notes one file the scan could not examine, under a status naming why.
// It returns the aliases captured for an authoritative retained path so the
// atomic path-scoped purge uses exactly the same identity as carry-forward.
func (g *yaraGapCollector) record(path, status string) []string {
var aliases []string
if _, alreadyRetained := g.paths[path]; !alreadyRetained {
if len(g.paths) < maxYARAGapPaths {
stable := true
if g.resolveAliases != nil {
aliases, stable = g.resolveAliases(path)
} else {
// Must stay symmetric with hasPath, which queries by every
// alias. Recording only the lexical spelling meant a finding
// stored through a symlinked document root never matched its
// own gap, and was purged for a file the scan never read.
aliases = coveragePathAliases(path)
}
if !stable {
g.recordUnknownRange(fmt.Sprintf("%s changed while its path identity was captured", path))
g.byStatus[status]++
if _, ok := g.example[status]; !ok {
g.example[status] = sanitizeJSTaintDisplay(path, yaraGapExampleMaxBytes)
}
return nil
}
g.paths[path] = struct{}{}
g.aliasesByPath[path] = aliases
for _, alias := range aliases {
g.pathAliases[alias] = struct{}{}
}
} else {
// Keep counting; stop claiming the path set is complete.
g.pathsTruncated = true
}
} else {
aliases = g.aliasesByPath[path]
}
g.byStatus[status]++
if _, ok := g.example[status]; !ok {
g.example[status] = sanitizeJSTaintDisplay(path, yaraGapExampleMaxBytes)
}
return aliases
}
// recordUnknownRange notes coverage lost over a range this walk cannot
// enumerate -- an unreadable directory, a failed Lstat that may hide a subtree.
// It deliberately records no path: claiming one would be false, and an unknown
// range has to suppress the purge for the whole owner.
func (g *yaraGapCollector) recordUnknownRange(detail string) {
g.unknown++
if g.unknownExample == "" {
g.unknownExample = sanitizeJSTaintDisplay(detail, yaraGapExampleMaxBytes)
}
}
// pathsIncomplete reports that this run cannot enumerate every gapped path, so
// its carry-forward set is not authoritative and the purge must be suppressed.
func (g *yaraGapCollector) pathsIncomplete() bool {
return g.pathsTruncated || g.unknown > 0
}
func (g *yaraGapCollector) empty() bool { return len(g.byStatus) == 0 && g.unknown == 0 }
func (g *yaraGapCollector) hasPath(path string) bool {
if _, ok := g.paths[path]; ok {
return true
}
for _, alias := range coveragePathAliases(path) {
if _, ok := g.pathAliases[alias]; ok {
return true
}
}
return false
}
func (g *yaraGapCollector) finding() alert.Finding {
total := 0
statuses := make([]string, 0, len(g.byStatus))
for status, n := range g.byStatus {
total += n
statuses = append(statuses, status)
}
sort.Strings(statuses)
parts := make([]string, 0, len(statuses)+2)
for _, status := range statuses {
parts = append(parts, fmt.Sprintf("%s=%d (example: %s)", status, g.byStatus[status], g.example[status]))
}
if g.unknown > 0 {
parts = append(parts, fmt.Sprintf("unreadable-range=%d (example: %s)", g.unknown, g.unknownExample))
}
if g.pathsTruncated {
parts = append(parts, fmt.Sprintf("exact paths retained for only the first %d", maxYARAGapPaths))
}
message := fmt.Sprintf("YARA deep scan could not inspect %d file(s)", total)
if total == 0 {
message = fmt.Sprintf("YARA deep scan could not cover %d location(s)", g.unknown)
}
return alert.Finding{
Severity: alert.High,
Check: "yara_scan_incomplete",
Message: message,
Details: strings.Join(parts, "; "),
// One host-wide condition: counts and examples vary per cycle and
// must not re-alert while coverage stays degraded.
DedupKey: "coverage_gap",
}
}
// carryForwardYARAFindings keeps every distinct prior rule finding for paths
// this cycle could not examine. YARA can emit several rule matches for one
// file, so collapsing by path would silently retire all but one finding even
// though the scan formed no opinion about any of them. Duplicate snapshots of
// the same identity are collapsed by Key, with the newest snapshot winning.
func carryForwardYARAFindings(prior []alert.Finding, gaps *yaraGapCollector) []alert.Finding {
byKey := make(map[string]alert.Finding)
for _, finding := range prior {
if finding.Check != "yara_match_scheduled" || !gaps.hasPath(finding.FilePath) {
continue
}
key := finding.Key()
current, exists := byKey[key]
if !exists || finding.Timestamp.After(current.Timestamp) ||
(finding.Timestamp.Equal(current.Timestamp) && finding.FilePath < current.FilePath) {
byKey[key] = finding
}
}
keys := make([]string, 0, len(byKey))
for key := range byKey {
keys = append(keys, key)
}
sort.Strings(keys)
carried := make([]alert.Finding, 0, len(keys))
for _, key := range keys {
finding := byKey[key]
finding.ScanCarryForward = true
carried = append(carried, finding)
}
return carried
}
// Package cms is the single declaration of the content management systems
// CSM supports. Tests compare the taint analyzer's path constants and the
// database adapter owners against this table. Clean-corpus coverage is
// tracked separately.
package cms
import (
"fmt"
"strings"
)
// Kind identifies a supported CMS. Values are the canonical lowercase names
// used in the corpus manifest and in check owner names.
type Kind string
const (
WordPress Kind = "wordpress"
Joomla Kind = "joomla"
Drupal Kind = "drupal"
OpenCart Kind = "opencart"
Magento Kind = "magento"
)
// Descriptor records what the rest of the tree needs to know about a CMS.
type Descriptor struct {
Kind Kind
// DBContentCheck is the runner owner name of the adapter that scans this
// CMS's database ("db_content" for WordPress, "db_content_joomla", ...).
DBContentCheck string
// PathConstants are the lower-cased PHP constants the CMS defines at
// bootstrap that always hold a local filesystem path. They preserve the
// analyzer's provenance assumptions; they are not proof of a constant's
// runtime value.
PathConstants []string
}
var descriptors = []Descriptor{
{Kind: WordPress, DBContentCheck: "db_content", PathConstants: []string{"abspath"}},
{Kind: Joomla, DBContentCheck: "db_content_joomla", PathConstants: []string{
"jpath_root", "jpath_base", "jpath_site", "jpath_administrator", "jpath_api",
"jpath_cache", "jpath_cli", "jpath_component", "jpath_component_administrator",
"jpath_component_site", "jpath_configuration", "jpath_installation",
"jpath_libraries", "jpath_manifests", "jpath_plugins", "jpath_public", "jpath_themes",
}},
{Kind: Drupal, DBContentCheck: "db_content_drupal", PathConstants: []string{"drupal_root"}},
{Kind: OpenCart, DBContentCheck: "db_content_opencart", PathConstants: []string{
"dir_application", "dir_cache", "dir_catalog", "dir_config", "dir_download",
"dir_extension", "dir_image", "dir_language", "dir_logs", "dir_modification",
"dir_opencart", "dir_root", "dir_session", "dir_storage", "dir_system",
"dir_template", "dir_upload",
}},
{Kind: Magento, DBContentCheck: "db_content_magento", PathConstants: []string{"bp"}},
}
func (d Descriptor) clone() Descriptor {
out := d
out.PathConstants = append([]string(nil), d.PathConstants...)
return out
}
// All returns every supported CMS in declaration order. The result is a
// deep copy; mutating it does not change policy.
func All() []Descriptor {
out := make([]Descriptor, 0, len(descriptors))
for _, d := range descriptors {
out = append(out, d.clone())
}
return out
}
// Lookup returns the descriptor for k, or the zero descriptor and false.
func Lookup(k Kind) (Descriptor, bool) {
for _, d := range descriptors {
if d.Kind == k {
return d.clone(), true
}
}
return Descriptor{}, false
}
// Parse accepts only an exact declared kind. It does not trim or fold case,
// so a manifest or config value must be spelled canonically.
func Parse(s string) (Kind, bool) {
for _, d := range descriptors {
if string(d.Kind) == s {
return d.Kind, true
}
}
return "", false
}
// validateDescriptors reports the first shape violation in ds: kinds and
// owner names must be non-empty and unique, and path constants non-empty,
// lower case and unique within and across descriptors.
func validateDescriptors(ds []Descriptor) error {
seenKind := make(map[Kind]bool, len(ds))
seenOwner := make(map[string]bool, len(ds))
seenConst := make(map[string]Kind)
for _, d := range ds {
if d.Kind == "" {
return fmt.Errorf("descriptor with empty kind")
}
if seenKind[d.Kind] {
return fmt.Errorf("%s: kind declared twice", d.Kind)
}
seenKind[d.Kind] = true
if d.DBContentCheck == "" {
return fmt.Errorf("%s: empty DBContentCheck", d.Kind)
}
if seenOwner[d.DBContentCheck] {
return fmt.Errorf("%s: DBContentCheck %q reused", d.Kind, d.DBContentCheck)
}
seenOwner[d.DBContentCheck] = true
if len(d.PathConstants) == 0 {
return fmt.Errorf("%s: no path constants; local provenance cannot be analysed", d.Kind)
}
for _, c := range d.PathConstants {
if c == "" || c != strings.ToLower(c) {
return fmt.Errorf("%s: constant %q must be non-empty lower case", d.Kind, c)
}
if owner, dup := seenConst[c]; dup {
return fmt.Errorf("constant %q declared by both %s and %s", c, owner, d.Kind)
}
seenConst[c] = d.Kind
}
}
return nil
}
package config
// CanonicalCheckName maps persisted check names to the current finding names.
// Config exclusions, saved mutes and pending work must survive a producer rename.
func CanonicalCheckName(name string) string {
switch name {
case "ftp_login_realtime":
return "ftp_login"
case "ssh_login_realtime":
return "ssh_login_unknown_ip"
default:
return name
}
}
package config
import (
"net"
"os"
"path/filepath"
"strings"
"syscall"
"time"
)
// clamdSocketCandidates are the unix sockets a clamd is normally reachable on.
// The path is set by whoever packaged clamd, not by CSM: RHEL's clamd-scan,
// Debian's clamav-daemon and cPanel's bundled clamd all choose differently, and
// a host whose setting names the wrong one scans no mail at all while every
// health signal still reports the watcher as running.
//
// Every entry is a root-owned service directory. A world-writable location such
// as /tmp is deliberately absent: any account could bind a socket there, and
// CSM would then stream every attachment to it and believe the "OK" it answers.
var clamdSocketCandidates = []string{
"/var/run/clamd.scan/clamd.sock",
"/run/clamd.scan/clamd.sock",
"/var/run/clamav/clamd.ctl",
"/run/clamav/clamd.ctl",
"/var/run/clamav/clamd.sock",
"/run/clamav/clamd.sock",
"/usr/local/cpanel/3rdparty/var/clamav/clamd.sock",
}
// clamdDialTimeout bounds each probe. The candidate list is walked on a health
// path, so a hung socket must not hold the caller. A real clamd that is merely
// busy still answers PING promptly; it is a separate thread from scanning.
const clamdDialTimeout = 2 * time.Second
// ResolveClamdSocket returns the socket to talk to clamd on, and whether it had
// to be discovered because the configured one was not answering.
//
// The configured path always wins when clamd is answering on it. Only when it
// is not does CSM fall back to a well-known location, because refusing to scan
// mail is the worse failure: silence there looks exactly like clean mail.
//
// A discovered socket is only accepted when it is owned by root or by this
// process and sits in a directory no other account can write to, and when
// whatever is listening answers clamd's PING. Discovery must not be a way to
// point mail scanning at something an account controls.
func ResolveClamdSocket(configured string) (string, bool) {
if configured != "" && clamdSocketTrusted(configured) && clamdSocketAnswers(configured) {
return configured, false
}
for _, candidate := range clamdSocketCandidates {
if candidate == configured {
continue
}
if !clamdSocketTrusted(candidate) {
continue
}
if clamdSocketAnswers(candidate) {
return candidate, true
}
}
return configured, false
}
// clamdSocketTrusted reports whether only a privileged account could have put
// this socket here.
//
// What matters is the directory: whoever can write to it decides what the name
// resolves to, and CSM is about to stream every mail attachment to whatever is
// listening and believe the verdict it returns. The socket's own owner is not
// the test -- clamd is packaged to run as its own service user (Debian's
// clamav owns /run/clamav), so requiring root there would reject exactly the
// sockets this discovery exists to find.
//
// A sticky directory is refused rather than allowed: /tmp lets any account
// create the name first, and being unable to delete someone else's socket is
// no help when the attacker's is the one that got there.
func clamdSocketTrusted(path string) bool {
info, err := os.Lstat(path)
if err != nil || info.Mode()&os.ModeSymlink != 0 {
return false
}
return clamdDirTrusted(filepath.Dir(path))
}
// clamdMaxServiceUID is the ceiling for a packaged service account. Hosting
// accounts start well above it on every panel CSM supports.
const clamdMaxServiceUID = 500
func clamdDirTrusted(dir string) bool {
info, err := os.Lstat(dir)
if err != nil || !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
return false
}
if info.Mode().Perm()&0o022 != 0 {
return false
}
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return false
}
// Owned by root, by this process, or by the service account clamd runs
// as -- never by an account that also hosts websites.
uid := int(stat.Uid)
return uid == 0 || uid == os.Geteuid() || uid < clamdMaxServiceUID
}
// clamdSocketAnswers reports whether clamd itself is listening. Connecting
// proves only that something is there; PING proves it speaks the protocol CSM
// is about to hand mail to.
func clamdSocketAnswers(path string) bool {
conn, err := net.DialTimeout("unix", path, clamdDialTimeout)
if err != nil {
return false
}
defer func() { _ = conn.Close() }()
if setErr := conn.SetDeadline(time.Now().Add(clamdDialTimeout)); setErr != nil {
return false
}
if _, writeErr := conn.Write([]byte("zPING\x00")); writeErr != nil {
return false
}
buf := make([]byte, 16)
n, readErr := conn.Read(buf)
if readErr != nil || n == 0 {
return false
}
return strings.Contains(strings.ToUpper(string(buf[:n])), "PONG")
}
package config
import (
"errors"
"fmt"
"io/fs"
"os"
"path/filepath"
"sort"
"strings"
"syscall"
"gopkg.in/yaml.v3"
)
type confFragment struct {
path string
node *yaml.Node
}
// ValidateConfDir vets an operator-selected conf.d directory before any YAML
// fragments are loaded from it. The returned path is symlink-resolved so later
// reads do not depend on a mutable link name.
func ValidateConfDir(dir string) (string, error) {
if dir == "" {
return "", nil
}
if !filepath.IsAbs(dir) {
return "", fmt.Errorf("conf.d directory must be an absolute path, got %q", dir)
}
resolved, err := filepath.EvalSymlinks(dir)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return "", fmt.Errorf("conf.d directory does not exist: %s", dir)
}
return "", fmt.Errorf("conf.d directory symlink resolution: %w", err)
}
info, err := os.Stat(resolved)
if err != nil {
return "", fmt.Errorf("conf.d directory stat: %w", err)
}
if !info.IsDir() {
return "", fmt.Errorf("conf.d directory is not a directory: %s", resolved)
}
if trustErr := validateConfPathTrust("conf.d directory", resolved, info); trustErr != nil {
return "", trustErr
}
return resolved, nil
}
// LoadConfDir reads every *.yaml file in dir in lexicographic order and
// returns each as a parsed yaml.DocumentNode. A missing directory is not
// an error and returns an empty slice; an unreadable file or invalid YAML
// is fatal so operators see misconfigurations at startup.
func LoadConfDir(dir string) ([]*yaml.Node, error) {
frags, err := loadConfDirFragments(dir)
if err != nil {
return nil, err
}
out := make([]*yaml.Node, 0, len(frags))
for _, frag := range frags {
out = append(out, frag.node)
}
return out, nil
}
// ConfDirFragment is one conf.d drop-in fragment's filename and raw bytes,
// exported so the integrity hasher can cover the same fragment set the loader
// merges without duplicating the enumeration rules.
type ConfDirFragment struct {
Name string
Data []byte
}
// ConfDirFragmentDigestInput returns every non-empty trusted conf.d fragment in
// merge order (sorted .yaml/.yml, symlink-resolved, trust-validated) as
// name+content pairs for integrity hashing. An empty dir or no mergeable
// fragments yields nil so a config without conf.d hashes to the empty digest
// and its baseline is unaffected.
func ConfDirFragmentDigestInput(dir string) ([]ConfDirFragment, error) {
files, err := confDirFragmentFiles(dir)
if err != nil {
return nil, err
}
if len(files) == 0 {
return nil, nil
}
out := make([]ConfDirFragment, 0, len(files))
for _, ff := range files {
if _, ok, err := parseConfFragment(ff); err != nil {
return nil, err
} else if !ok {
continue
}
out = append(out, ConfDirFragment{Name: ff.name, Data: ff.data})
}
return out, nil
}
type confFragmentFile struct {
name string
path string
data []byte
}
// confDirFragmentFiles enumerates trusted conf.d fragment files and returns
// their raw bytes in merge order. Shared by loadConfDirFragments and the
// integrity hasher so both observe exactly the same fragment set.
func confDirFragmentFiles(dir string) ([]confFragmentFile, error) {
if dir == "" {
return nil, nil
}
resolved, err := filepath.EvalSymlinks(dir)
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return nil, nil
}
return nil, fmt.Errorf("conf.d directory symlink resolution: %w", err)
}
info, err := os.Stat(resolved)
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return nil, nil
}
return nil, fmt.Errorf("conf.d directory stat: %w", err)
}
if !info.IsDir() {
return nil, fmt.Errorf("conf.d directory is not a directory: %s", resolved)
}
if trustErr := validateConfPathTrust("conf.d directory", resolved, info); trustErr != nil {
return nil, trustErr
}
dir = resolved
entries, err := os.ReadDir(dir)
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return nil, nil
}
return nil, fmt.Errorf("reading %s: %w", dir, err)
}
names := make([]string, 0, len(entries))
for _, e := range entries {
if e.IsDir() {
continue
}
if !strings.HasSuffix(e.Name(), ".yaml") && !strings.HasSuffix(e.Name(), ".yml") {
continue
}
names = append(names, e.Name())
}
sort.Strings(names)
out := make([]confFragmentFile, 0, len(names))
for _, name := range names {
path := filepath.Join(dir, name)
data, err := readTrustedConfFragment(path)
if err != nil {
return nil, err
}
out = append(out, confFragmentFile{name: name, path: path, data: data})
}
return out, nil
}
func loadConfDirFragments(dir string) ([]confFragment, error) {
files, err := confDirFragmentFiles(dir)
if err != nil {
return nil, err
}
out := make([]confFragment, 0, len(files))
for _, ff := range files {
node, ok, err := parseConfFragment(ff)
if err != nil {
return nil, err
}
if !ok {
continue
}
out = append(out, confFragment{path: ff.path, node: node})
}
return out, nil
}
func parseConfFragment(ff confFragmentFile) (*yaml.Node, bool, error) {
var node yaml.Node
if err := yaml.Unmarshal(ff.data, &node); err != nil {
return nil, false, fmt.Errorf("parsing %s: %w", ff.path, err)
}
// Skip empty files (Unmarshal yields a zero-Content document).
if len(node.Content) == 0 {
return nil, false, nil
}
normalized, err := normalizeYAMLForMerge(&node)
if err != nil {
return nil, false, fmt.Errorf("parsing %s: %w", ff.path, err)
}
node = *normalized
if hasTopLevelKey(&node, "integrity") {
return nil, false, fmt.Errorf("conf.d fragment %s must not set daemon-managed integrity metadata", ff.path)
}
// The confd block decides which fragments the integrity digest covers, so
// a fragment that could set it would be able to exempt itself.
if hasTopLevelKey(&node, "confd") {
return nil, false, fmt.Errorf("conf.d fragment %s must not set the confd policy block; it belongs in the main config", ff.path)
}
return &node, true, nil
}
func hasTopLevelKey(root *yaml.Node, key string) bool {
cur := root
if cur.Kind == yaml.DocumentNode {
if len(cur.Content) == 0 {
return false
}
cur = cur.Content[0]
}
if cur.Kind != yaml.MappingNode {
return false
}
for i := 0; i+1 < len(cur.Content); i += 2 {
if cur.Content[i].Value == key {
return true
}
}
return false
}
func readTrustedConfFragment(path string) ([]byte, error) {
// #nosec G304 -- path is built from an operator-selected conf.d and a directory entry.
f, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("reading %s: %w", path, err)
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return nil, fmt.Errorf("stat %s: %w", path, err)
}
if !info.Mode().IsRegular() {
return nil, fmt.Errorf("conf.d fragment is not a regular file: %s", path)
}
if trustErr := validateConfPathTrust("conf.d fragment", path, info); trustErr != nil {
return nil, trustErr
}
data, err := readConfigBytesLimited(f)
if errors.Is(err, errConfigTooLarge) {
return nil, fmt.Errorf("conf.d fragment %s exceeds %d byte cap", path, MaxConfigBytes)
}
if err != nil {
return nil, fmt.Errorf("reading %s: %w", path, err)
}
return data, nil
}
func validateConfPathTrust(kind, path string, info os.FileInfo) error {
if mode := info.Mode().Perm(); mode&0022 != 0 {
return fmt.Errorf("%s %s has unsafe mode %04o (group or world writable); set 0750/0640 or stricter", kind, path, mode)
}
sys, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return nil
}
// #nosec G115 -- Linux uid_t is uint32; os.Geteuid returns the kernel
// effective UID and cannot overflow that type on supported hosts.
selfUID := uint32(os.Geteuid())
if sys.Uid != 0 && sys.Uid != selfUID {
return fmt.Errorf("%s %s owner uid=%d is neither root (0) nor process uid=%d; refusing to load untrusted YAML", kind, path, sys.Uid, selfUID)
}
return nil
}
package config
import (
"bytes"
"errors"
"fmt"
"io"
"net"
"net/url"
"os"
"path/filepath"
"regexp"
"strings"
"time"
"gopkg.in/yaml.v3"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/firewall"
)
// WebUIToken is one entry in WebUI.Tokens. Scope must be "admin" or "read".
type WebUIToken struct {
Name string `yaml:"name"`
Token string `yaml:"token"`
Scope string `yaml:"scope"`
}
const (
// DefaultBlockExpiry is how long an automatic temporary block lasts when
// auto_response.block_expiry is unset.
DefaultBlockExpiry = "24h"
// DefaultMaxBlocksPerHour is the safe hourly cap used when the operator
// leaves auto_response.max_blocks_per_hour unset or sets it to 0.
DefaultMaxBlocksPerHour = 50
// File response limits share one rolling window across automatic cleaners
// and quarantine. Zero selects these defaults, never an unlimited budget.
DefaultMaxFileActionsPerHour = 50
DefaultMaxFileActionsPerAccountPerHour = 10
DefaultMaxFileActionFailuresPerHour = 3
// MaxFileResponseLimit bounds retained hourly safety reservations.
MaxFileResponseLimit = 10000
// DefaultNetBlockThreshold is how many blocked addresses in one IPv4 /24
// or IPv6 /64 escalate to a subnet block when the key is unset.
DefaultNetBlockThreshold = 3
// DefaultNetBlockWindow is how far back blocked addresses count toward
// netblock_threshold when auto_response.netblock_window is unset. Hosts
// that rotate through a /24 keep one address blocked at a time, so only
// counting live blocks never reaches the threshold.
DefaultNetBlockWindow = "168h"
// DefaultPermBlockCount is how many temporary blocks inside the interval
// promote an address to a permanent block when the key is unset.
DefaultPermBlockCount = 4
// DefaultPermBlockInterval is the window used to count temporary blocks
// when auto_response.permblock_interval is unset.
DefaultPermBlockInterval = "24h"
// MinBlockEscalationCount is the lowest meaningful value for either
// counter. Below it there is no pattern to escalate from.
MinBlockEscalationCount = 2
)
const (
// DefaultExposedFileScanDepth is the number of directory levels below a
// document root inspected by the web-exposed-file detector.
DefaultExposedFileScanDepth = 2
// MaxExposedFileScanDepth prevents a configuration typo from turning the
// hourly detector into an effectively unbounded filesystem traversal.
MaxExposedFileScanDepth = 10
// DefaultDropperUnlinkTTLSec is how long the realtime monitor waits after
// a fresh PHP/executable lands under a web docroot before probing whether
// it self-deleted. Long enough that upgrade staging has settled, short
// enough that a dropper is reported minutes after it vanishes.
DefaultDropperUnlinkTTLSec = 300
// MinDropperUnlinkTTLSec / MaxDropperUnlinkTTLSec bound the operator
// knob: below 30s legitimate installers race the probe, above 1h the
// tracker holds state for little detection gain.
MinDropperUnlinkTTLSec = 30
MaxDropperUnlinkTTLSec = 3600
// DefaultHTTPScannerErrorPct is the fallback percentage gate for the
// URL scanner-profile detector when the operator leaves it unset.
DefaultHTTPScannerErrorPct = 90
// DefaultHTTPScannerMinDistinctPaths is the fallback breadth gate for
// distinct probe paths when the operator leaves it unset.
DefaultHTTPScannerMinDistinctPaths = 10
// HTTPScannerMaxDistinctPaths matches the detector's per-IP path cap.
// Higher configured breadth thresholds can never be reached.
HTTPScannerMaxDistinctPaths = 512
// SlowBruteMinThreshold is the smallest useful long-horizon failure
// threshold: the detector also requires this many distinct mailboxes.
SlowBruteMinThreshold = 3
// SlowBruteMaxThreshold matches the detector's bounded per-IP timestamp
// history. A larger value could never fire and would make the knob lie.
SlowBruteMaxThreshold = 1024
// SlowBruteMaxWindowMin bounds long-lived per-IP state while still allowing
// operators to cover paced attacks over a full week.
SlowBruteMaxWindowMin = 7 * 24 * 60
)
// DefaultHTTPScannerStatusCodes returns a fresh copy of the probe-error
// statuses used when the operator leaves the scanner status list unset.
func DefaultHTTPScannerStatusCodes() []int {
return []int{404, 403}
}
const (
// DefaultHTTPASNCrawlMinIPs is the minimum distinct source IPs from one
// ASN inside the window before http_asn_crawl fires. 0 disables the detector.
DefaultHTTPASNCrawlMinIPs = 25
// DefaultXMLRPCThreshold is the per-IP POST /xmlrpc.php count that trips
// xmlrpc_abuse in access-log based detectors. 0 disables the check; an
// absent key defaults to this. Raised from the legacy 30 so
// Jetpack/WooCommerce sites, which legitimately make many xmlrpc.php calls,
// are not hard-blocked.
DefaultXMLRPCThreshold = 100
// DefaultHTTPASNCrawlMinExpensive is the minimum uncacheable requests from
// the ASN inside the window before http_asn_crawl fires.
DefaultHTTPASNCrawlMinExpensive = 250
// DefaultHTTPASNCrawlMinSharePct is the minimum percentage of a single
// account's total uncacheable requests that must come from one ASN.
DefaultHTTPASNCrawlMinSharePct = 50
// DefaultHTTPASNCrawlHighAmpPct is the percentage threshold above which the
// ASN's share triggers a High-severity finding instead of Warning.
DefaultHTTPASNCrawlHighAmpPct = 50
// DefaultHTTPASNCrawlHighVolMult is the multiplier applied to MinExpensive
// for the High-severity volume threshold.
DefaultHTTPASNCrawlHighVolMult = 4
// DefaultHTTPASNCrawlMaxPrefix is the maximum number of distinct /24
// prefixes (IPv4) or /48 prefixes (IPv6) within the ASN before the finding
// is promoted to Critical (distributed saturation pattern).
DefaultHTTPASNCrawlMaxPrefix = 8
// DefaultHTTPASNCrawl16PrefPct is the percentage of IPs that must share a
// /16 prefix for the condensed-prefix heuristic to apply.
DefaultHTTPASNCrawl16PrefPct = 60
// DefaultHTTPASNCrawlMaxTrackedIPs caps how many source IPs per ASN the
// rolling window keeps in memory.
DefaultHTTPASNCrawlMaxTrackedIPs = 20000
// DefaultHTTPASNCrawlWindowMin is the rolling window length in minutes.
DefaultHTTPASNCrawlWindowMin = 60
// DefaultHTTPASNCrawlTempban is the ban duration applied when the detector
// fires and auto-response is enabled.
DefaultHTTPASNCrawlTempban = "24h"
)
// httpASNCrawlReverseProxySeed is the built-in safety list of reverse-proxy
// CDN ASNs (Cloudflare, Fastly, Akamai). Edge IPs from these never form a
// finding or tempban; real-client attribution requires web_server.trusted_proxies.
var httpASNCrawlReverseProxySeed = []uint{13335, 54113, 20940}
// MailLogsConfig controls how postfix/dovecot logs are read.
//
// source: auto - try file first; fall back to journal if file absent.
// source: file - require log file at the platform-default path.
// source: journal - read from systemd-journald (units must be set).
//
// Units is consulted for journal fallback: the daemon matches each
// systemd unit by name, appending ".service" for bare service names.
type MailLogsConfig struct {
Source string `yaml:"source"` // auto | file | journal
File string `yaml:"file,omitempty"` // override platform default
Units []string `yaml:"units,omitempty"` // for journal source
}
type Config struct {
ConfigFile string `yaml:"-"`
ConfigDir string `yaml:"-" hotreload:"restart"` // /etc/csm/conf.d (or operator override); empty means no drop-ins loaded
Hostname string `yaml:"hostname" hotreload:"restart"`
// Mode is the operator's posture for this host. "enforce" (default)
// leaves every subsystem under its own switch. "observe" declares that
// CSM must not change host state: the daemon skips the host integration
// files it otherwise deploys at startup, and a config that still enables
// a state-changing subsystem is refused by name instead of being
// silently rewritten. Rewriting would be persisted: the config
// re-signing path marshals the in-memory config back over csm.yaml.
Mode string `yaml:"mode" hotreload:"restart"`
Alerts struct {
Email struct {
Enabled bool `yaml:"enabled"`
To []string `yaml:"to"`
From string `yaml:"from"`
SMTP string `yaml:"smtp"`
DisabledChecks []string `yaml:"disabled_checks"`
} `yaml:"email"`
Webhook struct {
Enabled bool `yaml:"enabled"`
URL string `yaml:"url"`
Type string `yaml:"type"` // slack, discord, generic, phpanel
// HMACSecret is the shared secret used to sign each request when
// Type=="phpanel". Read from this field directly OR via HMACSecretEnv
// (env wins, for secret hygiene).
HMACSecret string `yaml:"hmac_secret,omitempty"`
HMACSecretEnv string `yaml:"hmac_secret_env,omitempty"`
// PerFinding documents the expected phpanel delivery shape. Phpanel
// webhooks always emit one signed POST per finding; other webhook
// types keep the existing digest delivery.
PerFinding bool `yaml:"per_finding,omitempty"`
} `yaml:"webhook"`
Heartbeat struct {
Enabled bool `yaml:"enabled"`
URL string `yaml:"url"`
} `yaml:"heartbeat"`
MaxPerHour int `yaml:"max_per_hour"`
// AuditLog ships every (deduplicated) finding to one or more
// SIEM-friendly destinations. Schema is stable: parsers can
// pin on the v=1 contract. Both sub-blocks default off; the
// alert pipeline behaves identically to before when neither
// is enabled.
AuditLog struct {
File struct {
Enabled bool `yaml:"enabled"`
Path string `yaml:"path"` // default: /var/log/csm/audit.jsonl
} `yaml:"file"`
Syslog struct {
Enabled bool `yaml:"enabled"`
Network string `yaml:"network"` // udp | tcp | unix | unixgram | tls
Address string `yaml:"address"` // host:port or filesystem path
Facility string `yaml:"facility"` // default: local0
TLSCAFile string `yaml:"tls_ca"` // optional CA cert for tls
} `yaml:"syslog"`
} `yaml:"audit_log"`
// BlockDigest emits a per-country roll-up of auto-blocked IPs so
// operators see when their customers' countries get blocked. The
// hotreload:"restart" tag overrides the safe Alerts parent: the
// collector and its ticker are built once at startup.
BlockDigest struct {
Enabled bool `yaml:"enabled"`
Countries []string `yaml:"countries"`
Interval string `yaml:"interval"`
Live bool `yaml:"live"`
SendOn string `yaml:"send_on"`
Channel string `yaml:"channel"`
MinBlock int `yaml:"min_block"`
} `yaml:"block_digest" hotreload:"restart"`
} `yaml:"alerts" hotreload:"safe"`
Integrity struct {
BinaryHash string `yaml:"binary_hash"`
ConfigHash string `yaml:"config_hash"`
// ConfdHash covers the conf.d drop-in fragments merged on top of the
// main config. Empty when no fragments exist, so configs without
// conf.d stay byte-identical to their pre-existing baseline. Without
// it an attacker with write access to conf.d could override any
// setting (auto_response, block_ips, dry_run) without tripping
// integrity verification.
ConfdHash string `yaml:"confd_hash"`
Immutable bool `yaml:"immutable"`
} `yaml:"integrity"`
// ConfD holds operator policy for the conf.d drop-in directory. It lives
// outside the integrity block on purpose: config_hash skips that block,
// so an exemption placed there could be added without tripping
// verification. Here it is covered by config_hash like any other key.
ConfD struct {
// IntegrityExempt names drop-in fragments (bare filenames) whose
// content is left out of integrity.confd_hash. A fragment its owning
// integration rewrites on its own schedule cannot be pinned by a
// static hash without turning every restart into a manual rehash.
// Every other fragment stays covered. Fragments cannot set this key.
IntegrityExempt []string `yaml:"integrity_exempt,omitempty"`
} `yaml:"confd,omitempty" hotreload:"safe"`
Thresholds struct {
MailQueueWarn int `yaml:"mail_queue_warn"`
MailQueueCrit int `yaml:"mail_queue_crit"`
StateExpiryHours int `yaml:"state_expiry_hours"`
DeepScanIntervalMin int `yaml:"deep_scan_interval_min"`
WPCoreCheckIntervalMin int `yaml:"wp_core_check_interval_min"`
WebshellScanIntervalMin int `yaml:"webshell_scan_interval_min"`
FilesystemScanIntervalMin int `yaml:"filesystem_scan_interval_min"`
// ExposedFileScanDepth bounds how many directory levels below each
// docroot the web-exposed-file detector descends (default 2, maximum
// 10). Dumps and backups almost always sit at or just under the web root.
ExposedFileScanDepth int `yaml:"exposed_file_scan_depth"`
MultiIPLoginThreshold int `yaml:"multi_ip_login_threshold"`
MultiIPLoginWindowMin int `yaml:"multi_ip_login_window_min"`
// CredStuffingDistinctAccounts is the number of distinct accounts a
// single source IP must fail against inside the multi_ip_login window
// to raise a credential_stuffing finding. This is the breadth signal
// (one source, many accounts) that the count-based pam_bruteforce
// detector does not catch. Default 5 when unset or <=0, matching the
// always-on posture of the sibling multi_ip_login_threshold.
CredStuffingDistinctAccounts int `yaml:"cred_stuffing_distinct_accounts"`
// PAMBruteforceThreshold is the number of PAM authentication
// failures (SSH, mail, FTP via pam_csm) from one source IP inside
// PAMBruteforceWindowMin minutes before pam_bruteforce fires and the
// address is auto-blocked. Defaults 5 failures in 10 minutes.
PAMBruteforceThreshold int `yaml:"pam_bruteforce_threshold"`
PAMBruteforceWindowMin int `yaml:"pam_bruteforce_window_min"`
PluginCheckIntervalMin int `yaml:"plugin_check_interval_min"`
BruteForceWindow int `yaml:"brute_force_window"`
// DomlogMaxFiles caps how many per-domain access logs the WP brute
// force check scans per cycle. Sites are ranked by recent mtime so
// the cap chops least-active domains. Default 500. Bump on hosts
// with many active domains so late-alphabet sites are not skipped.
DomlogMaxFiles int `yaml:"domlog_max_files"`
// PHPConfigWalkMaxDirs caps how many directories the PHP configuration
// scan walks below a single document root. Reaching it leaves the rest
// of that root unexamined, which is reported as an incomplete scan.
// Shared-hosting accounts routinely carry tens of thousands of
// directories under vendor and node_modules trees, so raise this rather
// than accept the gap. Default 50000.
PHPConfigWalkMaxDirs int `yaml:"php_config_walk_max_dirs"`
// PHPConfigWalkMaxEntries caps how many directory entries that same
// walk examines below one document root. A shallow tree holding very
// many files reaches this before the directory ceiling. Default 500000.
PHPConfigWalkMaxEntries int `yaml:"php_config_walk_max_entries"`
// AccountScanMaxFiles caps how many account and mail-domain paths
// account-scoped scanners iterate per cycle. This covers SSH keys,
// cPanel API tokens, Dovecot shadow files, CMS DB scans, forwarders,
// user crontabs, and account .config backdoor paths. Paths are
// ranked by mtime desc so the cap chops least-active accounts.
// Default 10000 covers typical cPanel hosts with no cap effect.
// Raise on very large multi-tenant hosts where late-mtime accounts
// get skipped.
AccountScanMaxFiles int `yaml:"account_scan_max_files"`
// CrontabBase64BlobMaxBytes caps a single base64 candidate before
// decoding in MatchCrontabPatternsDeep. Default 16384 encoded
// bytes (~12 KiB decoded) comfortably fits any realistic gsocket
// or `base64 -d|bash` payload while bounding work on adversarial
// input. Raise on hosts where csm_checks_crontab_base64_truncated_total
// shows recurring truncation. Must be a multiple of 4: standard
// base64 requires aligned input or the decode fails and the
// candidate is silently skipped.
CrontabBase64BlobMaxBytes int `yaml:"crontab_base64_blob_max_bytes"`
// DomlogTailLines is how many trailing lines the WP brute force
// check reads from each per-domain access log per cycle. Default
// 500 covers roughly 10 minutes of traffic on a busy site. Raise
// on hosts where slow-burn attacks against high-volume domains
// spread across more than 500 lines per scan interval, so the
// per-IP counter window is wide enough to trip a threshold.
DomlogTailLines int `yaml:"domlog_tail_lines"`
// DomlogMaxAgeMin is how many minutes back the WP brute force
// scanner accepts a per-domain access log as "fresh enough to
// scan." Logs whose mtime is older than this are skipped. Default
// 30 keeps the scan focused on currently-active vhosts. Raise on
// low-traffic hosts where a slow-burn dictionary attack against
// a quiet domain still needs to fall inside the window.
DomlogMaxAgeMin int `yaml:"domlog_max_age_min"`
// MailLogTailLines is how many trailing lines CheckMailPerAccount
// reads from /var/log/exim_mainlog per cycle. Default 500. Raise
// on busy mail hosts where a single account's spam burst spreads
// across more than 500 lines per cycle.
MailLogTailLines int `yaml:"mail_log_tail_lines"`
// SyslogMessagesTailLines is kept for direct CheckFTPLogins callers
// that do not pass a state store. The daemon's store-backed FTP detector
// follows /var/log/messages forward and ignores this setting.
SyslogMessagesTailLines int `yaml:"syslog_messages_tail_lines"`
// FTPFailWindowMin is the sliding-window length (minutes) over which the
// store-backed FTP detector accumulates per-IP pure-ftpd auth failures
// before alerting at ftpFailThreshold. 0 means use the built-in default
// (30). Valid range 1..1440.
FTPFailWindowMin int `yaml:"ftp_fail_window_min"`
SMTPBruteForceThreshold int `yaml:"smtp_bruteforce_threshold"`
SMTPBruteForceWindowMin int `yaml:"smtp_bruteforce_window_min"`
SMTPBruteForceSuppressMin int `yaml:"smtp_bruteforce_suppress_min"`
SMTPBruteForceSubnetThresh int `yaml:"smtp_bruteforce_subnet_threshold"`
SMTPAccountSprayThreshold int `yaml:"smtp_account_spray_threshold"`
SMTPBruteForceMaxTracked int `yaml:"smtp_bruteforce_max_tracked"`
// The slow pair catches attackers pacing below the fast window:
// SlowThreshold failures within SlowWindowMin minutes across at
// least three distinct mailboxes from one IP.
SMTPBruteForceSlowThreshold int `yaml:"smtp_bruteforce_slow_threshold"`
SMTPBruteForceSlowWindowMin int `yaml:"smtp_bruteforce_slow_window_min"`
// SMTP probe abuse counts raw inbound SMTP connect events per source
// IP (independent of AUTH outcome) so probe-and-disconnect scanners
// that never reach the AUTH stage are still caught. Threshold sized
// well above any legitimate MUA usage. Explicit 0 disables.
SMTPProbeThreshold int `yaml:"smtp_probe_threshold"`
SMTPProbeWindowMin int `yaml:"smtp_probe_window_min"`
SMTPProbeSuppressMin int `yaml:"smtp_probe_suppress_min"`
SMTPProbeMaxTracked int `yaml:"smtp_probe_max_tracked"`
MailBruteForceThreshold int `yaml:"mail_bruteforce_threshold"`
MailBruteForceWindowMin int `yaml:"mail_bruteforce_window_min"`
MailBruteForceSuppressMin int `yaml:"mail_bruteforce_suppress_min"`
MailBruteForceSubnetThresh int `yaml:"mail_bruteforce_subnet_threshold"`
MailAccountSprayThreshold int `yaml:"mail_account_spray_threshold"`
MailBruteForceMaxTracked int `yaml:"mail_bruteforce_max_tracked"`
// The slow pair mirrors the SMTP one for the dovecot-native path.
MailBruteForceSlowThreshold int `yaml:"mail_bruteforce_slow_threshold"`
MailBruteForceSlowWindowMin int `yaml:"mail_bruteforce_slow_window_min"`
// MailBruteAccountKey selects how the account is extracted from a
// dovecot/postfix log line for per-account brute-force scoring.
// - builtin:dovecot-user (default) - match `user=<...>`
// - builtin:postfix-sasl - match `sasl_username=<...>`
// - regex:<pattern> - capture group 1 is the account
MailBruteAccountKey string `yaml:"mail_brute_account_key,omitempty"`
// ModSecEscalationHits is the number of ModSecurity denies from a
// single source IP, inside ModSecEscalationWindowMin, that
// promote the IP from "logged" to "escalated" (Critical finding
// + firewall hand-off). Default 3. Lower it on hosts where
// low-and-slow scanners go below the floor for too long.
ModSecEscalationHits int `yaml:"modsec_escalation_hits"`
// ModSecEscalationWindowMin is the sliding-window size for the
// hit counter. Default 10. Bumping it (e.g. to 60-240) catches
// paced attackers that spread denies across hours without
// changing the trip count.
ModSecEscalationWindowMin int `yaml:"modsec_escalation_window_min"`
// ModSecLowConfidenceEscalationHits is the low-confidence-only
// backstop: how many low-confidence policy/anomaly denies (e.g.
// COMODO content-type/anomaly rules) from one IP within the
// escalation window force a firewall escalation even when the
// burst carries no attack signature. It closes the "only trip
// anomaly rules" bypass without banning a single customer's
// checkout retry. Default 30. Raise it on hosts with high-volume
// legitimate apps that bulk-trip policy rules.
ModSecLowConfidenceEscalationHits int `yaml:"modsec_low_confidence_escalation_hits"`
// HTTPFloodThreshold is the minimum requests per http_flood_window_min
// from one source IP that emits http_request_flood. 0 (default) disables
// the detector. Operators should sample baseline traffic before setting
// a nonzero value.
HTTPFloodThreshold int `yaml:"http_flood_threshold"`
// HTTPFloodWindowMin is the rate window in minutes for HTTPFloodThreshold
// counting. Default 5.
HTTPFloodWindowMin int `yaml:"http_flood_window_min"`
// HTTPUASpoofThreshold is the minimum per-IP per-window count
// of WPSpoofPingback, cache-confirmed negative ClaimedBot,
// ScriptingLang, Headless, or Empty UA requests that emits
// http_ua_spoof. KnownScanner still emits on count=1. Default 30.
HTTPUASpoofThreshold int `yaml:"http_ua_spoof_threshold"`
// HTTPDistributedMinIPs is the number of distinct source IPs that
// must each trip an HTTP-abuse threshold (wp-login / xmlrpc / user
// enumeration / request flood / UA spoof) against one vhost in a
// scan window before a single http_distributed_flood finding is
// emitted for that vhost. Only already-abusive IPs are counted, so
// a popular site's ordinary visitor spread does not trip it. 0
// disables the detector; the shipped sample sets 10.
HTTPDistributedMinIPs int `yaml:"http_distributed_min_ips"`
// HTTPScannerMinRequests is the minimum in-window requests from one
// source IP before the URL scanner-profile detector evaluates the IP.
// The volume gate keeps a visitor following a handful of dead links
// out of scope. 0 (default) disables the detector; the shipped
// sample suggests 30.
HTTPScannerMinRequests int `yaml:"http_scanner_min_requests"`
// HTTPScannerErrorPct is the minimum percentage of in-window
// requests answered with a probe-error status before
// http_scanner_profile fires. Default 90.
HTTPScannerErrorPct int `yaml:"http_scanner_error_pct"`
// HTTPScannerMinDistinctPaths is the minimum count of distinct
// error-status request paths (query strings stripped) before the
// detector fires. Repeated hits on one missing resource (a dead
// bookmark, a broken image) never look like URL enumeration no
// matter the volume. Default 10, maximum 512.
HTTPScannerMinDistinctPaths int `yaml:"http_scanner_min_distinct_paths"`
// HTTPScannerStatusCodes is the set of response statuses counted as
// probe errors. Default [404, 403]. 301 is deliberately excluded:
// http->https and www redirects make every legitimate visitor
// 301-heavy, and a site migration redirects entire domains. Add 301
// here only on hosts where that traffic shape is impossible.
HTTPScannerStatusCodes []int `yaml:"http_scanner_status_codes"`
// HTTPUAScriptingEnabled opts in to flagging scripting-language
// UA strings (curl, python-requests, wget, etc.) as spoof candidates.
// Off by default: many legitimate API integrations use these.
HTTPUAScriptingEnabled bool `yaml:"http_ua_scripting_enabled"`
// HTTPUAHeadlessEnabled opts in to flagging headless-browser UA
// strings (HeadlessChrome, PhantomJS, Playwright, etc.). Off by
// default: headless browsers are used by legitimate monitoring tools.
HTTPUAHeadlessEnabled bool `yaml:"http_ua_headless_enabled"`
// HTTPUAEmptyEnabled opts in to flagging requests with an empty or
// dash User-Agent. Off by default: some CDN health checks omit UA.
HTTPUAEmptyEnabled bool `yaml:"http_ua_empty_enabled"`
// FullScanMaxFileMB caps how large a single file may be (in MiB)
// before a full-scan job skips it and emits a full_scan_file_too_large
// warning. 0 falls back to the default (16 MiB). Raise only when a
// known-large legitimate file must be scanned; the cap exists to
// bound memory use during content decoding.
FullScanMaxFileMB int `yaml:"full_scan_max_file_mb"`
// ScanJobRetention is how many completed full-scan job records the
// job manager keeps before evicting the oldest. Default 20.
ScanJobRetention int `yaml:"scan_job_retention"`
// RollingCoverage enables Phase 3 rolling content-scan coverage: each
// periodic cycle also sweeps a bounded path-sorted slice of files past
// the mtime cap so dormant files are eventually content-scanned.
// Default true; set false to restore pure top-N-by-mtime behavior.
RollingCoverage bool `yaml:"rolling_coverage"`
// DropperDetection gates the realtime self-deleting-dropper detector:
// a PHP/executable file created under a web document root and
// unlinked before the TTL probe. Default true.
DropperDetection bool `yaml:"dropper_detection"`
// DropperUnlinkTTLSec is the tracking TTL in seconds for that
// detector. Default 300.
DropperUnlinkTTLSec int `yaml:"dropper_unlink_ttl_sec"`
// HTTPASNCrawlWindowMin is the rolling window in minutes for the
// single-ASN distributed crawl detector. Default 60.
HTTPASNCrawlWindowMin int `yaml:"http_asn_crawl_window_min"`
// XMLRPCThreshold is the per-IP POST /xmlrpc.php count that trips
// xmlrpc_abuse in access-log based detectors (a hard auto-block). 0
// disables the check; an absent key defaults to 100.
XMLRPCThreshold int `yaml:"xmlrpc_threshold"`
// HTTPASNCrawlMinIPs is the minimum distinct source IPs from one ASN
// inside the window before http_asn_crawl fires. 0 disables the
// detector; an absent key defaults to 25.
HTTPASNCrawlMinIPs int `yaml:"http_asn_crawl_min_ips"`
// HTTPASNCrawlMinExpensive is the minimum uncacheable requests from the
// ASN inside the window. Default 250.
HTTPASNCrawlMinExpensive int `yaml:"http_asn_crawl_min_expensive"`
// HTTPASNCrawlMinSharePct is the minimum percentage of one account's
// total uncacheable requests that must originate from the ASN. 1..100.
// Default 50.
HTTPASNCrawlMinSharePct int `yaml:"http_asn_crawl_min_share_pct"`
// HTTPASNCrawlHighAmpPct is the share percentage above which the finding
// is promoted to High severity. 1..100. Default 50.
HTTPASNCrawlHighAmpPct int `yaml:"http_asn_crawl_high_amp_pct"`
// HTTPASNCrawlHighVolumeMult multiplies MinExpensive to obtain the
// High-severity volume threshold. Default 4.
HTTPASNCrawlHighVolumeMult int `yaml:"http_asn_crawl_high_volume_mult"`
// HTTPASNCrawlSaturation is the lsphp process saturation count that
// triggers a Critical finding. 0 means use
// performance.php_process_warn_per_user.
HTTPASNCrawlSaturation int `yaml:"http_asn_crawl_saturation"`
// HTTPASNCrawlMaxPrefix is the maximum number of distinct /24 (IPv4) or
// /48 (IPv6) prefixes within the ASN before the finding is promoted to
// Critical (distributed saturation). Default 8.
HTTPASNCrawlMaxPrefix int `yaml:"http_asn_crawl_max_prefix"`
// HTTPASNCrawl16PrefPct is the percentage of IPs that must share a /16
// prefix for the condensed-prefix heuristic to apply. 1..100. Default 60.
HTTPASNCrawl16PrefPct int `yaml:"http_asn_crawl_16_pref_pct"`
// HTTPASNCrawlMaxTrackedIPs caps how many source IPs per ASN are kept in
// the rolling window. Default 20000.
HTTPASNCrawlMaxTrackedIPs int `yaml:"http_asn_crawl_max_tracked_ips"`
// HTTPASNCrawlAllowlistASNs is a list of ASNs that are never flagged by
// the detector. Ships empty; operator-supplied list only.
HTTPASNCrawlAllowlistASNs []uint `yaml:"http_asn_crawl_allowlist_asns"`
// HTTPASNCrawlReverseProxyASNs lists CDN/reverse-proxy ASNs whose edge
// IPs are never directly flagged by this detector.
// Ships with Cloudflare (13335), Fastly (54113), and Akamai (20940).
// Set to [] to clear the seed; absent key retains the built-in list.
HTTPASNCrawlReverseProxyASNs []uint `yaml:"http_asn_crawl_reverse_proxy_asns"`
} `yaml:"thresholds" hotreload:"safe"`
InfraIPs []string `yaml:"infra_ips" hotreload:"restart"`
StatePath string `yaml:"state_path" hotreload:"restart"`
// Detection groups operator-facing knobs that gate individual scanners.
Detection struct {
// DBObjectScanning toggles the MySQL persistence scanner
// (triggers/events/procedures/functions). Tri-state *bool
// matching the existing yara_worker_enabled pattern: nil =
// default-on, *true = explicit on, *false = explicit off.
// When off both the Critical (db_malicious_*) and Warning
// (db_unexpected_*) emit paths fall silent; the manual
// `csm db-clean drop-object` CLI keeps working so operators
// can act on objects discovered by other means.
DBObjectScanning *bool `yaml:"db_object_scanning"`
// DBObjectAllowlist suppresses the Warning tier
// (db_unexpected_*) for objects an operator has reviewed and
// accepted. Entries shaped <account>:<schema>:<type>:<name>.
// The Critical tier (db_malicious_*) ignores this list --
// pattern hits always fire.
DBObjectAllowlist []string `yaml:"db_object_allowlist"`
// VulnerablePluginScanning toggles the known-vulnerable WordPress
// plugin detector (matches installed versions against a curated
// CVE/KEV feed). Tri-state *bool: nil = default-on, *true = on,
// *false = off. Independent of outdated_plugins.
VulnerablePluginScanning *bool `yaml:"vulnerable_plugin_scanning"`
// VulnerablePluginAllow suppresses vulnerable_plugins findings for a
// specific slug@version the operator has reviewed and accepted (for
// example a back-ported/vendor-patched build the feed cannot see).
// Entries shaped <slug>@<version>, case-insensitive presence check.
VulnerablePluginAllow []string `yaml:"vulnerable_plugin_allow"`
// AdminOverlapMinAccounts is the threshold at which the
// cross-account admin email correlator emits a finding. Default
// 2 matches the most common compromise pattern on shared hosting
// -- a contractor account used across multiple customer cPanels
// is a single credential leak away from compromising every site
// they touch. Operators with deliberately shared internal admin
// emails (e.g. one ops team across many sites) can raise the
// threshold to silence the routine overlap.
AdminOverlapMinAccounts int `yaml:"admin_overlap_min_accounts"`
// AdminOverlapTrustedEmails suppresses cross-account admin
// overlap findings for exact, operator-reviewed email addresses.
AdminOverlapTrustedEmails []string `yaml:"admin_overlap_trusted_emails"`
// AdminOverlapTrustedDomains suppresses cross-account admin
// overlap findings for exact email domains used by trusted
// developer or reseller admin accounts.
AdminOverlapTrustedDomains []string `yaml:"admin_overlap_trusted_domains"`
// RescanOnSignatureUpdate fires a forced full-tree deep
// scan the next time the content of any file under
// cfg.Signatures.RulesDir changes. Tri-state *bool: nil = default-on,
// *true = explicit on, *false = explicit off. Off means the
// existing behaviour (deep-tier runs against the fanotify
// short-list when fanotify is active) is unchanged; new
// rules only catch files that change after the update.
RescanOnSignatureUpdate *bool `yaml:"rescan_on_signature_update"`
// AFAlgBackend selects the live AF_ALG (CVE-2026-31431, "Copy
// Fail") detection backend. Empty / "auto" picks BPF LSM if
// the binary was built with -tags bpf and the kernel supports
// it, otherwise the audit-log inotify listener. "bpf" forces
// BPF and disables the audit fallback (no live monitor if BPF
// is unavailable; the periodic critical-tier check still
// runs). "auditd" forces the audit listener even on BPF-
// capable kernels -- a kill switch when a BPF-tagged release
// misbehaves and the operator wants to revert without
// rebuilding. "none" disables the live monitor entirely.
AFAlgBackend string `yaml:"af_alg_backend"`
// ConnectionTrackerBackend selects the live outbound-connection
// tracker. Empty / "auto" tries BPF cgroup/connect4,6 first and
// falls back to the existing /proc/net/tcp polling. "bpf"
// requires BPF (no fallback). "legacy" pins polling. "none"
// disables the live tracker; the periodic check still runs.
ConnectionTrackerBackend string `yaml:"connection_tracker_backend"`
// ConnectionPollInterval is how often the legacy polling backend
// reads /proc/net/tcp(6). Ignored when the BPF backend is active.
// Empty / zero defaults to 30s.
ConnectionPollInterval time.Duration `yaml:"connection_poll_interval"`
// ExecMonitorBackend selects the live process-exec monitor.
// Empty / "auto" tries the sched_process_exec BPF tracepoint and
// falls back to the periodic /proc walk. "bpf" requires BPF
// (no fallback). "legacy" pins polling. "none" disables the live
// monitor; the periodic deep-tier checks still run.
ExecMonitorBackend string `yaml:"exec_monitor_backend"`
// ExecMonitorPollInterval is how often the legacy polling backend
// runs CheckSuspiciousProcesses + CheckFakeKernelThreads. Ignored
// when the BPF backend is active. Empty / zero defaults to 30m.
ExecMonitorPollInterval time.Duration `yaml:"exec_monitor_poll_interval"`
// SensitiveFilesBackend selects the live sensitive-file write
// monitor. Empty / "auto" tries the BPF LSM hook on /etc/shadow
// and friends, falling back to a periodic content-hash check.
// "bpf" requires BPF (no fallback). "legacy" pins polling.
// "none" disables the live monitor; the periodic check still runs.
SensitiveFilesBackend string `yaml:"sensitive_files_backend"`
// SensitiveFilesPollInterval is how often the BPF watchset map
// refreshes (to pick up newly-created files in glob directories
// and handle inode reuse) and how often the legacy polling
// backend runs the content-hash check. Empty / zero defaults to 5m.
SensitiveFilesPollInterval time.Duration `yaml:"sensitive_files_poll_interval"`
// DirectSMTPEgress flags non-MTA local processes opening
// outbound SMTP connections. Phase 3 of the BPF Incident
// Response Roadmap. Detection-only this phase; Phase 4 will
// add the auto-response action gated by DryRun.
DirectSMTPEgress struct {
Enabled bool `yaml:"enabled"`
Backend string `yaml:"backend"` // auto / bpf / legacy / none
// DryRun, when true (or absent for safety), reports findings
// but takes no detector-scoped action. Phase 3 emits findings
// regardless; the knob exists for the Phase 4 action.
DryRun *bool `yaml:"dry_run,omitempty"`
Ports []int `yaml:"ports,omitempty"`
} `yaml:"direct_smtp_egress" hotreload:"safe"`
// BadASNOutbound flags outbound connections whose destination IP
// resolves (via the GeoLite2-ASN database) to a bad or unexpected
// autonomous system. It is the third leg of the host-takeover
// chain correlator (alongside a new uid-0 account and a planted
// suid binary). Requires the GeoLite2-ASN database. Off by default;
// classification needs operator-supplied ASN lists.
BadASNOutbound struct {
Enabled bool `yaml:"enabled"`
// BlockedASNs are autonomous system numbers always treated as
// bad (e.g. known bulletproof hosters).
BlockedASNs []uint `yaml:"blocked_asns"`
// AllowedASNs, when non-empty, switches to allowlist mode: any
// destination ASN outside this set is treated as bad. Use on
// hosts whose legitimate egress is confined to a few providers.
AllowedASNs []uint `yaml:"allowed_asns"`
} `yaml:"bad_asn_outbound" hotreload:"safe"`
} `yaml:"detection" hotreload:"safe"`
Suppressions struct {
UPCPWindowStart string `yaml:"upcp_window_start"`
UPCPWindowEnd string `yaml:"upcp_window_end"`
KnownAPITokens []string `yaml:"known_api_tokens"`
IgnorePaths []string `yaml:"ignore_paths"`
SuppressWebmail bool `yaml:"suppress_webmail_alerts"` // don't alert on webmail logins
SuppressCpanelLogin bool `yaml:"suppress_cpanel_login_alerts"` // don't alert on cPanel direct logins
SuppressBlockedAlerts bool `yaml:"suppress_blocked_alerts"` // don't alert on IPs that were auto-blocked
TrustedCountries []string `yaml:"trusted_countries"` // ISO 3166-1 alpha-2 codes - suppress cPanel login alerts from these countries
} `yaml:"suppressions" hotreload:"safe"`
AutoResponse struct {
Enabled bool `yaml:"enabled"`
KillProcesses bool `yaml:"kill_processes"`
QuarantineFiles bool `yaml:"quarantine_files"`
MaxFileActionsPerHour int `yaml:"max_file_actions_per_hour"`
MaxFileActionsPerAccountPerHour int `yaml:"max_file_actions_per_account_per_hour"`
MaxFileActionFailuresPerHour int `yaml:"max_file_action_failures_per_hour"`
BlockIPs bool `yaml:"block_ips"`
BlockExpiry string `yaml:"block_expiry"` // e.g. "24h", "12h"
// HTTPASNCrawlTempban is the ban duration for http_asn_crawl findings
// when auto-response is enabled. Default "24h".
HTTPASNCrawlTempban string `yaml:"http_asn_crawl_tempban"`
EnforcePermissions bool `yaml:"enforce_permissions"` // auto-chmod 644 world/group-writable PHP files (default false)
FixWPCron bool `yaml:"fix_wp_cron"` // auto-disable WP-Cron + install per-user system cron on perf_wp_cron findings (default false)
BlockCpanelLogins bool `yaml:"block_cpanel_logins"` // block IPs on cPanel/webmail login alerts (default false)
// HTTPScannerAction selects the response for http_scanner_profile
// findings. "challenge" (default) routes the IP to the PoW
// challenge when the challenge subsystem is enabled and falls
// through to a firewall block when it is not; "block" always
// hard-blocks without offering a challenge.
HTTPScannerAction string `yaml:"http_scanner_action"`
NetBlock bool `yaml:"netblock"` // auto-block IPv4 /24 or IPv6 /64 at threshold
NetBlockThreshold int `yaml:"netblock_threshold"` // IPs from same IPv4 /24 or IPv6 /64 before subnet block (default 3)
NetBlockWindow string `yaml:"netblock_window"` // how far back blocked IPs count toward the threshold (default "168h")
// MaxBlocksPerHour caps per-IP auto-blocks per wall-clock hour.
// 0 uses DefaultMaxBlocksPerHour.
MaxBlocksPerHour int `yaml:"max_blocks_per_hour"`
PermBlock bool `yaml:"permblock"` // auto-promote to permanent after N temp blocks
PermBlockCount int `yaml:"permblock_count"` // temp blocks before permanent (default 4)
PermBlockInterval string `yaml:"permblock_interval"` // window for counting temp blocks (default "24h")
CleanDatabase bool `yaml:"clean_database"` // auto-clean malicious DB injections, revoke sessions, block attacker IPs (default false)
CleanHtaccess bool `yaml:"clean_htaccess"` // auto-clean .htaccess directives flagged by the hardened detectors (default false)
// VirtualPatchExposedFiles controls automatic .htaccess "Require all
// denied" rules for confirmed web_exposed_* findings. "off" (default,
// also the value for any unrecognised setting): detection only. "manual":
// deny rules are written only when the operator runs `csm virtual-patch`.
// "auto": deny rules are written on every scan except for warning-only
// sample SQL, gated by DryRun. Use VirtualPatchMode() for the safe default.
VirtualPatchExposedFiles string `yaml:"virtual_patch_exposed_files"`
DisableEnforceAFAlg bool `yaml:"disable_enforce_af_alg"` // suspend periodic AF_ALG enforcement; marker file + detection remain active (default false = enforce when marker present)
CopyFailKillProcess bool `yaml:"copy_fail_kill_process"` // SIGKILL processes caught opening AF_ALG sockets via the live listener (default false; alert-only)
// MailAuthRecovery optionally self-heals the mail auth backend
// (cPanel's cpdoveauthd). CSM always probes the socket, alerts on an
// outage, and pauses mail/SMTP brute-force auto-block while it is down
// (cPanel only); only the service restart is gated here and is off by
// default. A restart runs only after the backend is continuously down
// for DownGrace, so a brief blip during nightly maintenance never trips it.
MailAuthRecovery struct {
RestartEnabled bool `yaml:"restart_enabled"` // run a service restart after a sustained outage (default false)
DownGrace string `yaml:"down_grace"` // continuously-down duration before restarting (default "10m")
MaxRestartsPerHour int `yaml:"max_restarts_per_hour"` // hourly cap on restart attempts (default 3)
RestartCommand string `yaml:"restart_command"` // command to run (default cPanel restartsrv_dovecot)
} `yaml:"mail_auth_recovery" hotreload:"restart"`
// DryRun, when true (or absent - safety default), logs the intended
// action but does NOT touch nftables. Mirrors the PHPRelay.DryRun
// pattern: pointer-bool to distinguish "operator explicitly set false"
// from "operator omitted the key". Implicit nil means dry-run on, so
// flipping block_ips: true alone never causes a real block.
DryRun *bool `yaml:"dry_run,omitempty"`
// PHPRelay controls the auto-freeze behaviour that companion
// email PHP-relay detectors emit findings for. Freeze and DryRun
// are *bool so we can distinguish OMITTED from EXPLICIT FALSE in
// YAML. A plain bool zero-value is false, which would let an
// operator write `freeze: true` and (by forgetting `dry_run`)
// get LIVE freezes against their will. Pointer values: nil =
// "not set in YAML"; *true / *false = explicit. Use the
// FreezeEnabled() / DryRunEnabled() accessors on *Config to
// resolve the safe defaults rather than dereferencing directly.
PHPRelay struct {
Freeze *bool `yaml:"freeze"`
DryRun *bool `yaml:"dry_run"`
MaxActionsPerMinute int `yaml:"max_actions_per_minute"`
} `yaml:"php_relay"`
// VerdictCallback lets phpanel observe each block decision before it's
// applied. CSM POSTs the verdict to the configured URL with HMAC-SHA256
// signing (same scheme as the phpanel webhook); the response is
// advisory - phpanel can attach a tenant_id, return "allow" to keep
// the event audit-only, or omit a response entirely
// (CSM proceeds with its default verdict). NOT a per-tenant nftables
// enforcement: that's a separate, larger feature.
//
// Secret resolution happens at call time (the verdict.Client reads
// HMACSecretEnv per call), so operators can rotate via env without
// restarting the daemon.
VerdictCallback struct {
Enabled bool `yaml:"enabled"`
URL string `yaml:"url"`
HMACSecret string `yaml:"hmac_secret,omitempty"`
HMACSecretEnv string `yaml:"hmac_secret_env,omitempty"`
TimeoutSec int `yaml:"timeout_sec"`
// RequireResponseSignature controls whether the panel must sign
// its response body with the same HMAC scheme used on the
// request (X-CSM-Signature header). Default true: when a secret
// is configured, CSM rejects unsigned or forged responses to
// prevent an on-path attacker from downgrading block to allow.
// Set false only during a phpanel rollout that has not yet
// implemented response signing. In that mode, CSM still checks
// nonce or timestamp fields the panel echoes; responses that
// omit both keep the legacy advisory shape working.
RequireResponseSignature *bool `yaml:"require_response_signature,omitempty"`
// AllowUnsigned opts out of the default fail-closed posture and
// permits the verdict callback to fire without an HMAC secret,
// including advisory "allow" responses. Only set true while
// bootstrapping a new panel or during local testing; production
// deployments must keep this false so the daemon refuses to start
// when the secret env var is empty.
AllowUnsigned bool `yaml:"allow_unsigned,omitempty"`
} `yaml:"verdict_callback"`
} `yaml:"auto_response" hotreload:"safe"`
// BPFEnforcement is the optional in-kernel deny path for matched
// outbound connections. Phase 4 of the BPF Incident Response
// Roadmap. Defaults are all-safe: enforcement off, dry-run on.
// Operators flip live denial only after dry-run telemetry review.
BPFEnforcement struct {
Enabled bool `yaml:"enabled"`
// DryRun, when true (or absent for safety), logs intended
// denials but allows the connect. False = real deny.
DryRun *bool `yaml:"dry_run,omitempty"`
DirectSMTPEgress bool `yaml:"direct_smtp_egress"`
// VerdictCallback, when true, asks auto_response.verdict_callback
// for an advisory ALLOW override before recording a USERSPACE
// action (incident close, audit note). The in-kernel hook NEVER
// waits on this; it would add latency to every connect.
VerdictCallback bool `yaml:"verdict_callback"`
} `yaml:"bpf_enforcement" hotreload:"safe"`
Challenge struct {
Enabled bool `yaml:"enabled"` // enable challenge pages instead of hard block for some IPs
ListenAddr string `yaml:"listen_addr"` // bind address for the challenge listener (default: 127.0.0.1)
ListenPort int `yaml:"listen_port"` // port for challenge server (default: 8439)
// ListenAddr defaults to loopback because the production path
// keeps the listener private until an operator deliberately
// exposes it. The webserver integration redirects browsers to
// challenge.public_url, so installed direct mode needs a
// non-loopback listen address plus TLS material via
// challenge.tls_cert / tls_key, or the webui TLS fallback.
Secret string `yaml:"secret"` // HMAC secret for challenge tokens (auto-generated if empty)
Difficulty int `yaml:"difficulty"` // proof-of-work difficulty 0-5 (default: 2)
TrustedProxies []string `yaml:"trusted_proxies"` // IPs allowed to set X-Forwarded-For (empty = trust RemoteAddr only)
// TLSCert / TLSKey activate HTTPS on the challenge listener. Empty
// values keep loopback listeners on plain HTTP. Direct/public
// listeners fall back to webui.tls_cert / webui.tls_key so
// single-cert hosts can opt in without duplicating paths.
TLSCert string `yaml:"tls_cert"`
TLSKey string `yaml:"tls_key"`
// PublicURL is the external URL the webserver redirect target
// points at. Operators put it on an existing TLS-valid host so
// the integration does not need a new DNS record or cert.
// Example: https://server.example.com:8439/challenge with
// listen_addr=0.0.0.0, tls_cert/tls_key set to the host's
// cpanel cert.
// When empty the webserver-integration installer refuses to
// run; per-vhost reverse-proxy is no longer supported because
// LSWS proxy emulation does not honor it at server scope.
PublicURL string `yaml:"public_url"`
// CaptchaFallback shows a third-party CAPTCHA widget when JS is
// disabled. All fields default empty; the feature is off until
// the operator supplies provider + keys.
CaptchaFallback struct {
Provider string `yaml:"provider"` // "turnstile" | "hcaptcha" | "" (off)
SiteKey string `yaml:"site_key"` // public key embedded in the HTML widget
SecretKey string `yaml:"secret_key"` // verified server-side against the provider
Timeout time.Duration `yaml:"timeout"` // HTTP timeout for siteverify (default 10s)
} `yaml:"captcha_fallback"`
// VerifiedSession lets operators mint a signed cookie that
// bypasses the PoW for the cookie's TTL. The signing key is
// generated at daemon startup and rotates on restart.
VerifiedSession struct {
Enabled bool `yaml:"enabled"`
CookieName string `yaml:"cookie_name"` // default: csm_admin_session
TTL time.Duration `yaml:"ttl"` // default: 4h
AdminSecret string `yaml:"admin_secret"` // shared secret POST'd to /challenge/admin-token
} `yaml:"verified_session"`
// VerifiedCrawlers allows-passes traffic from search crawlers
// whose IP forward-confirms a reverse-DNS PTR matching one of
// the configured providers.
VerifiedCrawlers struct {
Enabled bool `yaml:"enabled"`
Providers []string `yaml:"providers"` // names: googlebot | bingbot
CacheTTL time.Duration `yaml:"cache_ttl"` // default: 15m
} `yaml:"verified_crawlers"`
// PortGate locks the challenge listener TCP port to specific
// source IPs via nftables. Enabled implies the daemon owns a
// dedicated `csm_chal` inet table whose chain drops all traffic
// to challenge.listen_port except: loopback, operator infra_ips,
// and IPs the IPList has just flagged. The set entry carries
// the same TTL as the challenge entry, so nftables expires the
// allow even if the daemon dies before calling Revoke. Auto-off
// when listen_addr is loopback (gate has no effect there).
PortGate struct {
Enabled bool `yaml:"enabled"`
} `yaml:"port_gate"`
} `yaml:"challenge" hotreload:"restart"`
PHPShield struct {
Enabled bool `yaml:"enabled"` // watch PHP Shield event log for alerts (default: false)
} `yaml:"php_shield" hotreload:"restart"`
Reputation struct {
AbuseIPDBKey string `yaml:"abuseipdb_key"`
Whitelist []string `yaml:"whitelist"` // IPs to never flag as malicious
// Rspamd queries the local rspamd controller for per-IP reject/junk
// counts. Disabled by default. URL must reach the controller HTTP port
// (default 11334). Token is the controller's admin password (rspamadm
// pw -e), supplied via env when possible. Token resolution happens at
// query time so operators can rotate via env without daemon restart.
Rspamd struct {
Enabled bool `yaml:"enabled"`
URL string `yaml:"url"`
Token string `yaml:"token,omitempty"`
TokenEnv string `yaml:"token_env,omitempty"`
} `yaml:"rspamd"`
// Upstream is an HTTP threat-intel source - typically a panel host that
// caches AbuseIPDB / proprietary scores on behalf of every agent in its
// fleet. Disabled by default. Token resolution happens at query time
// (see internal/threatintel/upstream_source.go) so operators can rotate
// the bearer via TokenEnv without restarting the daemon.
Upstream struct {
Enabled bool `yaml:"enabled"`
URL string `yaml:"url"`
Token string `yaml:"token,omitempty"` // discouraged - prefer TokenEnv
TokenEnv string `yaml:"token_env,omitempty"`
CacheTTLMin int `yaml:"cache_ttl_min"`
TimeoutSec int `yaml:"timeout_sec"`
} `yaml:"upstream"`
// BotVerifyEnabled controls async PTR+forward-A verification of
// claimed search-engine bot IPs. *bool so an explicit false
// survives SIGHUP reload without being overwritten by applyDefaults.
// Default true.
BotVerifyEnabled *bool `yaml:"bot_verify_enabled"`
// VerifiedBots extends the built-in good-bot allowlist with operator
// entries: claimed-UA substrings confirmed by forward-confirmed
// reverse DNS against the listed registrable-domain suffixes.
// Additive -- built-in bots (Googlebot etc.) always apply. Lets an
// operator stop a legitimate crawler (typically SEO/backlink bots)
// from tripping the HTTP scanner-profile detector without a code
// change. Suffixes are validated against shared-hosting domains so a
// bad entry cannot turn the allowlist into a bypass.
VerifiedBots []VerifiedBot `yaml:"verified_bots"`
// BotRanges controls the auto-updater that refreshes published
// AI-crawler IP ranges (OpenAI, Perplexity) used to verify those bots
// by address. Embedded snapshots work without it; the updater keeps
// them current via outbound HTTPS to the vendor endpoints. Default on.
BotRanges struct {
AutoUpdate *bool `yaml:"auto_update"` // nil = true
UpdateInterval string `yaml:"update_interval"` // default "24h"; min 1h
} `yaml:"bot_ranges" hotreload:"restart"`
// Report emits signed, minimized abuse reports for confirmed-abuse
// findings to a central abuse database or a private collector.
// Opt-in; default off. Keys/secrets resolve from *_env at startup.
Report struct {
Enabled bool `yaml:"enabled"`
Classes []string `yaml:"classes"` // bruteforce, php_relay, credential_stuffing, bad_asn_egress
SpoolPath string `yaml:"spool_path"` // bbolt file; default <state_dir>/abuse_reports.db
SpoolMax int `yaml:"spool_max"` // bounded queue size; default 10000
Targets []struct {
Name string `yaml:"name"`
URL string `yaml:"url"`
Transport string `yaml:"transport"` // ed25519 | hmac
NodeID string `yaml:"node_id"`
KeyID string `yaml:"key_id"`
KeyEnv string `yaml:"key_env"` // ed25519 hex private key, or hmac secret
TokenEnv string `yaml:"token_env"` // optional bearer for hmac collectors
} `yaml:"targets"`
} `yaml:"report" hotreload:"restart"`
// Central pulls a signed scored-set from the central abuse database and
// acts on it per Action. Opt-in; default off. Central data never hard
// blocks on its own (see Action / firebreaks).
Central struct {
Enabled bool `yaml:"enabled"`
SetURL string `yaml:"set_url"`
PubkeyEnv string `yaml:"pubkey_env"` // env holding the central Ed25519 public key (hex)
RefreshInterval string `yaml:"refresh_interval"` // e.g. "6h"; default 6h
Action string `yaml:"action"` // off | challenge | block_if_local_corroborated
BlockThreshold int `yaml:"block_threshold"` // node-local score threshold for block (default 80)
} `yaml:"central" hotreload:"restart"`
} `yaml:"reputation" hotreload:"safe"`
Signatures struct {
RulesDir string `yaml:"rules_dir"`
UpdateURL string `yaml:"update_url"`
AutoUpdate bool `yaml:"auto_update"` // auto-download rules daily (default: true if update_url set)
UpdateInterval string `yaml:"update_interval"` // how often to check (default: "24h")
SigningKey string `yaml:"signing_key"` // hex-encoded ed25519 public key for verifying rule updates
// AllowRuleCountDecrease permits a deliberately reduced, newer signed
// ruleset to bypass the default count-collapse guard.
AllowRuleCountDecrease bool `yaml:"allow_rule_count_decrease"`
YaraForge struct {
Enabled bool `yaml:"enabled"`
Tier string `yaml:"tier"` // "core", "extended", "full" (default: "core")
UpdateInterval string `yaml:"update_interval"` // default: "168h" (weekly)
DownloadURL string `yaml:"download_url"` // signed ZIP URL/template; supports {tier} and {version}
} `yaml:"yara_forge"`
// DisabledRules names rules to switch off: they are stripped from
// YARA-Forge downloads, skipped when the shipped .yml rules load,
// and stripped before the shipped .yar rules compile. `csm validate`
// warns about a name that matches no rule.
DisabledRules []string `yaml:"disabled_rules"`
// YaraWorkerEnabled is a tri-state: nil means "use system default"
// (default-on, per ROADMAP item 2 follow-up), *true means explicit on,
// *false means explicit off. Callers must nil-check before dereferencing;
// daemon.yaraWorkerOn() is the canonical accessor.
YaraWorkerEnabled *bool `yaml:"yara_worker_enabled"`
} `yaml:"signatures" hotreload:"restart"`
WebUI struct {
SessionLifetime string `yaml:"session_lifetime"`
SessionIdleTimeout string `yaml:"session_idle_timeout"`
Enabled bool `yaml:"enabled"`
Listen string `yaml:"listen"`
AuthToken string `yaml:"auth_token"`
MetricsToken string `yaml:"metrics_token" hotreload:"safe"` // optional Bearer token for /metrics; rotate via SIGHUP without restart
TLSCert string `yaml:"tls_cert"`
TLSKey string `yaml:"tls_key"`
UIDir string `yaml:"ui_dir"` // path to UI files on disk (default: /opt/csm/ui)
// AllowedOrigins lists extra browser origins ("https://host[:port]")
// whose API requests are accepted besides https://<hostname>:<port>.
// Loopback origins (SSH tunnels) are always accepted. Hot-reloadable.
AllowedOrigins []string `yaml:"allowed_origins,omitempty" hotreload:"safe"`
// Tokens is the multi-credential model added in v2.12.0. Each entry has
// a stable name (for audit), an opaque secret, and a scope that gates
// which endpoints accept it. Legacy AuthToken is preserved during the
// migration window so callers that read it directly keep working;
// applyDefaults populates Tokens from AuthToken when only the legacy
// field is set.
Tokens []WebUIToken `yaml:"tokens,omitempty"`
} `yaml:"webui" hotreload:"restart"`
EmailAV EmailAVConfig `yaml:"email_av" hotreload:"restart"`
EmailProtection struct {
PasswordCheckIntervalMin int `yaml:"password_check_interval_min"`
HighVolumeSenders []string `yaml:"high_volume_senders"`
RateWarnThreshold int `yaml:"rate_warn_threshold"`
RateCritThreshold int `yaml:"rate_crit_threshold"`
RateWindowMin int `yaml:"rate_window_min"`
KnownForwarders []string `yaml:"known_forwarders"`
// PHPRelay is the operator-tunable knob block for the email
// PHP-relay protection feature (Stage 1). Thresholds default to
// the values set in applyDefaults(), except documented zero
// sentinels such as AccountVolumePerHour auto-derive and
// FanoutDistinctRecipients disabling only that gate.
PHPRelay struct {
// The relay pipeline is wired once at startup; a reload can retune
// it but cannot start or stop it, so toggling it needs a restart.
Enabled bool `yaml:"enabled" hotreload:"restart"`
RateWindowMin int `yaml:"rate_window_min"`
HeaderScoreVolumeMin int `yaml:"header_score_volume_min"`
AbsoluteVolumePerHour int `yaml:"absolute_volume_per_hour"`
AccountVolumePerHour int `yaml:"account_volume_per_hour"`
ReputationFailuresPer24h int `yaml:"reputation_failures_per_24h"`
FanoutDistinctScripts int `yaml:"fanout_distinct_scripts"`
FanoutDistinctRecipients int `yaml:"fanout_distinct_recipients"`
FanoutWindowMin int `yaml:"fanout_window_min"`
BaselineSigma float64 `yaml:"baseline_sigma"`
BaselineObservationDays int `yaml:"baseline_observation_days"`
PoliciesDir string `yaml:"policies_dir"`
} `yaml:"php_relay"`
// CloudRelay scopes opt-out for the email_cloud_relay_abuse
// detector only. Use this when an operator legitimately runs a
// mailer on a public-cloud VM (Google Cloud, AWS, etc.) and the
// realtime/retro detectors keep false-firing on that mailbox.
// AllowUsers matches full mailboxes (case-insensitive). AllowDomains
// matches the domain part of the AUTH user (case-insensitive),
// covering every mailbox under that domain. Either match exits
// the detector before any window state is updated, so an
// allowlisted mailbox cannot prime the counter for another user.
// Leaving both empty preserves prior behavior. The shared
// EmailProtection.HighVolumeSenders list still applies as well.
CloudRelay struct {
AllowUsers []string `yaml:"allow_users"`
AllowDomains []string `yaml:"allow_domains"`
} `yaml:"cloud_relay"`
// ForwardGuard is the opt-in protection that holds spam/backscatter
// forward copies before they relay to an external provider. Default
// off; dry-run first. The MTA enforces; CSM only generates the rule
// and owns the quarantine, so it is never in the live mail path.
ForwardGuard ForwardGuardConfig `yaml:"forward_guard"`
} `yaml:"email_protection" hotreload:"safe"`
Firewall *firewall.FirewallConfig `yaml:"firewall" hotreload:"restart"`
GeoIP struct {
AccountID string `yaml:"account_id"`
LicenseKey string `yaml:"license_key"`
Editions []string `yaml:"editions"`
AutoUpdate *bool `yaml:"auto_update"` // nil = true when credentials set
UpdateInterval string `yaml:"update_interval"` // default "24h"
} `yaml:"geoip" hotreload:"restart"`
ModSecErrorLog string `yaml:"modsec_error_log" hotreload:"restart"`
ModSec struct {
RulesFile string `yaml:"rules_file"` // path to modsec2.user.conf
OverridesFile string `yaml:"overrides_file"` // path to csm-overrides.conf
ReloadCommand string `yaml:"reload_command"` // e.g. "systemctl restart lsws"
} `yaml:"modsec" hotreload:"restart"`
// WebServer overrides the auto-detected web server paths. Every field is
// optional: anything left blank or empty falls back to what
// platform.Detect() returned at startup. Intended for hosts with a
// custom layout (reverse proxy in front of a second daemon, non-standard
// package locations, chroot, etc.).
WebServer struct {
Type string `yaml:"type"` // "apache", "nginx", "litespeed" -- overrides auto-detect
ConfigDir string `yaml:"config_dir"` // e.g. /etc/apache2 or /etc/nginx
AccessLogs []string `yaml:"access_logs"` // candidate access-log paths, tried in order
ErrorLogs []string `yaml:"error_logs"` // candidate error-log paths (used for modsec denies)
ModSecAudits []string `yaml:"modsec_audit_logs"` // candidate ModSecurity audit-log paths
// TrustedProxies is the list of IP addresses or CIDR ranges whose
// X-Forwarded-For header is trusted for client IP extraction. When
// the connecting IP is not in this list, XFF is ignored and
// RemoteIP is used as-is.
TrustedProxies []string `yaml:"trusted_proxies"` // IP/CIDR sources allowed to supply X-Forwarded-For
// DomlogGlobs overrides the auto-detected per-vhost access-log glob
// patterns. When set, the platform default for the detected panel/OS
// is discarded and only these patterns are used. Leave empty to use
// the auto-detected globs.
DomlogGlobs []string `yaml:"domlog_globs"` // per-vhost access-log globs
} `yaml:"web_server" hotreload:"restart"`
// AccountRoots lets operators point the account-scan based checks at
// non-cPanel web root layouts. Each entry is a glob pattern expanded
// at check time. Validated directories also bound content remediation
// and quarantine restore. Examples:
//
// account_roots:
// - /var/www/*/public
// - /srv/http/*
// - /home/*/public_html # cPanel default (implicit when unset on cPanel)
//
// When unset, CSM uses the cPanel default of /home/*/public_html on
// cPanel hosts and an empty list on non-cPanel hosts (account-scan
// checks skip entirely). See docs/src/configuration.md for the full
// list of checks that consume this.
AccountRoots []string `yaml:"account_roots" hotreload:"restart"`
Performance struct {
Enabled *bool `yaml:"enabled"`
LoadHighMultiplier float64 `yaml:"load_high_multiplier"`
LoadCriticalMultiplier float64 `yaml:"load_critical_multiplier"`
PHPProcessWarnPerUser int `yaml:"php_process_warn_per_user"`
PHPProcessCriticalTotalMult int `yaml:"php_process_critical_total_multiplier"`
ErrorLogWarnSizeMB int `yaml:"error_log_warn_size_mb"`
MySQLJoinBufferMaxMB int `yaml:"mysql_join_buffer_max_mb"`
MySQLWaitTimeoutMax int `yaml:"mysql_wait_timeout_max"`
MySQLMaxConnectionsPerUser int `yaml:"mysql_max_connections_per_user"`
RedisBgsaveMinInterval int `yaml:"redis_bgsave_min_interval"`
RedisLargeDatasetGB int `yaml:"redis_large_dataset_gb"`
WPMemoryLimitMaxMB int `yaml:"wp_memory_limit_max_mb"`
WPTransientWarnMB int `yaml:"wp_transient_warn_mb"`
WPTransientCriticalMB int `yaml:"wp_transient_critical_mb"`
// WPCronFix tunes the WP-Cron remediation (manual fix from the Web UI
// and the daemon auto-response). Disabling WP-Cron without a real cron
// would stop scheduled tasks, so the fix also installs a per-user system
// cron that runs wp-cron.php on this interval.
WPCronFix struct {
IntervalMinutes int `yaml:"interval_minutes"` // system cron frequency; default 15, clamped to [1,60]
PHPBin string `yaml:"php_bin"` // cron interpreter override; empty => per-vhost, then detect
} `yaml:"wp_cron_fix"`
} `yaml:"performance" hotreload:"restart"`
Cloudflare struct {
Enabled bool `yaml:"enabled"`
RefreshHours int `yaml:"refresh_hours"`
} `yaml:"cloudflare" hotreload:"restart"`
C2Blocklist []string `yaml:"c2_blocklist" hotreload:"restart"`
BackdoorPorts []int `yaml:"backdoor_ports" hotreload:"restart"`
// DisabledChecks lists check names that should be skipped entirely by
// the runner (no execution, no finding, no email/webhook/audit). Use
// this when a whole category does not apply to a host (e.g. WAF/web
// checks on DNS-only cPanel servers). Distinct from
// alerts.email.disabled_checks, which only suppresses email but still
// runs the check and emits findings to other sinks.
DisabledChecks []string `yaml:"disabled_checks" hotreload:"safe"`
// Retention bounds bbolt growth in two independent ways:
// - Sweeps (opt-in via Enabled): a daily pass prunes per-bucket entries
// older than the configured TTL. Destructive, so off by default.
// - Compaction (automatic): the daemon reclaims freelist slack at startup
// when the file exceeds CompactMinSizeMB and is less than
// CompactFillRatio full. Non-destructive (it only rewrites the file to
// drop free pages), so it runs regardless of Enabled; set
// CompactMinSizeMB to 0 to disable it.
// All fields are hot-reload:"restart" because the retention goroutine and
// the startup compaction capture these on daemon start.
Retention struct {
Enabled bool `yaml:"enabled"` // opt-in (sweeps only)
FindingsDays int `yaml:"findings_days"` // default 90
HistoryDays int `yaml:"history_days"` // default 30
ReputationDays int `yaml:"reputation_days"` // default 180
SweepInterval string `yaml:"sweep_interval"` // default "24h"
CompactMinSizeMB int `yaml:"compact_min_size_mb"` // default 128; 0 disables auto-compaction
CompactFillRatio float64 `yaml:"compact_fill_ratio"` // default 0.5
} `yaml:"retention" hotreload:"restart"`
// Sentry ships panics and selected errors to a Sentry server for
// aggregation across hosts. Disabled by default; set enabled=true and
// provide a DSN from the Sentry project. Init is one-shot: changes
// require a daemon restart.
Sentry struct {
Enabled bool `yaml:"enabled"`
DSN string `yaml:"dsn"`
Environment string `yaml:"environment"` // e.g. "production", "staging"
SampleRate float64 `yaml:"sample_rate"` // 0 -> 1.0 (capture all errors)
Debug bool `yaml:"debug"` // SDK debug logs to stderr
} `yaml:"sentry" hotreload:"restart"`
// Debug exposes diagnostic endpoints. PprofListen, when set, binds a
// net/http/pprof server (heap/goroutine/CPU profiles) to the given
// host:port. It MUST be a loopback address (127.0.0.1 / ::1 / localhost):
// pprof leaks process internals and lets a caller trigger CPU/heap dumps,
// so it is never exposed off-box. Empty (default) disables it; reach it
// over an SSH tunnel. Restart-required because the listener starts once.
Debug struct {
PprofListen string `yaml:"pprof_listen"`
} `yaml:"debug" hotreload:"restart"`
// MailLogs selects the log source for the postfix/dovecot brute-force
// and relay detectors. Changing the source (file vs. journal) requires
// the daemon to re-attach its reader, so the field is tagged restart.
MailLogs MailLogsConfig `yaml:"mail_logs,omitempty" hotreload:"restart"`
// Updates controls the upstream release-availability poll surfaced
// in the Web UI top banner. The daemon never downloads or applies
// updates -- it only tells the operator that a newer version
// exists. Disable wholesale on air-gapped hosts.
Updates struct {
// CheckEnabled is a tri-state. nil means default-on; explicit
// false disables the poll entirely (no outbound HTTP, no
// package-manager probe). Use a pointer so the absence of the
// key in YAML is distinguishable from `check_enabled: false`.
CheckEnabled *bool `yaml:"check_enabled"`
// Interval is parsed by time.ParseDuration. Defaults to 24h;
// clamped to a 1h floor by updatecheck.New.
Interval string `yaml:"interval"`
// GitHubAPIURL overrides the default release endpoint. Tests
// and air-gapped mirrors use this; leave empty in production.
GitHubAPIURL string `yaml:"github_api_url,omitempty"`
// PackageName is the apt/dnf package name to query when the
// GitHub call fails. Defaults to "csm".
PackageName string `yaml:"package_name,omitempty"`
} `yaml:"updates" hotreload:"restart"`
// Incidents groups correlator-side knobs the operator can tune
// without code changes. Hot-reload "restart" because the daemon
// captures these on startup; flipping mid-run would race the
// retention loop.
Incidents struct {
// AutoClose resolves Open / Contained incidents whose UpdatedAt
// exceeds the per-kind idle threshold. Default-on with safe
// thresholds; mailbox takeover, mailbox bruteforce, credential
// spray, and web attack expire at 24h, web-account compromise at
// 7d, and host-level kinds never auto-close. Operators who want
// to monitor decisions without writing back can flip dry_run=true.
AutoClose struct {
// Enabled is tri-state: nil (default) means default-on; an
// explicit false in YAML disables. Pointer so absence in YAML
// is distinguishable from "enabled: false".
Enabled *bool `yaml:"enabled"`
DryRun bool `yaml:"dry_run"`
// ByKind maps incident kind -> idle threshold (parsed by
// time.ParseDuration). Kinds absent from the map are never
// auto-closed. Use a string-keyed map so the operator can
// add custom kinds without recompiling. Empty map falls back
// to safe defaults (mailbox_takeover=24h,
// mailbox_bruteforce=24h, credential_spray=24h,
// web_attack=24h, web_account_compromise=7d).
ByKind map[string]string `yaml:"by_kind,omitempty"`
} `yaml:"auto_close"`
// SpraySuppression collapses one source IP brute-forcing many
// distinct mailboxes/accounts into a single credential_spray
// super-incident. Default-OFF + dry_run=TRUE so the path ships
// dark; counters and audit log show what would have happened.
// Operators flip enabled=true and dry_run=false after watching
// the counters on their own infra.
SpraySuppression struct {
Enabled bool `yaml:"enabled"`
DryRun bool `yaml:"dry_run"`
DistinctMailboxes int `yaml:"distinct_mailboxes"`
SeverityEscalateAt int `yaml:"severity_escalate_at"`
PerCheck []string `yaml:"per_check"`
MaxTrackedIPs int `yaml:"max_tracked_ips"`
// BlockAtSeverity drives the firewall hand-off. Empty (default)
// means detection-only: the super-incident opens and counters
// move, but the IP is not blocked. "high" blocks as soon as the
// detector trips at DistinctMailboxes. "critical" waits for the
// severity escalation (SeverityEscalateAt distinct mailboxes)
// before blocking. Auto_response.dry_run and block_ips still
// gate the actual firewall call.
BlockAtSeverity string `yaml:"block_at_severity"`
} `yaml:"spray_suppression"`
// AutoBlock is the generic incident-driven firewall hand-off.
// Independent of SpraySuppression; applies to any non-spray
// incident kind that carries a remote_ip in its correlation
// key. Default-OFF so the path is dormant until an operator
// opts in. Block requests still respect auto_response.enabled
// and auto_response.block_ips at decision time.
AutoBlock struct {
Enabled bool `yaml:"enabled"`
BlockAtSeverity string `yaml:"block_at_severity"`
Kinds []string `yaml:"kinds"`
} `yaml:"auto_block"`
} `yaml:"incidents" hotreload:"restart"`
}
// UpdatesCheckEnabled reports the YAML-level state for the upstream
// release poll. Defaults to TRUE when omitted (most operators want
// the banner). Set `updates.check_enabled: false` to disable.
func (c *Config) UpdatesCheckEnabled() bool {
return c.Updates.CheckEnabled == nil || *c.Updates.CheckEnabled
}
// UpdatesInterval returns the parsed poll interval. Falls back to
// 24h on parse error or when unset; updatecheck applies the floor.
func (c *Config) UpdatesInterval() time.Duration {
if c.Updates.Interval == "" {
return 24 * time.Hour
}
d, err := time.ParseDuration(c.Updates.Interval)
if err != nil {
return 24 * time.Hour
}
return d
}
// BlockDigestInterval is the digest cadence. Empty, invalid, or non-positive
// values fall back to one hour.
func (c *Config) BlockDigestInterval() time.Duration {
if c.Alerts.BlockDigest.Interval == "" {
return time.Hour
}
d, err := time.ParseDuration(c.Alerts.BlockDigest.Interval)
if err != nil || d <= 0 {
return time.Hour
}
return d
}
// UpdatesPackageName returns the apt/dnf package name to query.
// Defaults to "csm".
func (c *Config) UpdatesPackageName() string {
if c.Updates.PackageName == "" {
return "csm"
}
return c.Updates.PackageName
}
// PHPRelayFreezeEnabled reports whether auto-freeze should run for the
// email PHP-relay detectors. Defaults to false when freeze was not set
// in YAML — the operator must opt in explicitly.
func (cfg *Config) PHPRelayFreezeEnabled() bool {
return cfg.AutoResponse.PHPRelay.Freeze != nil && *cfg.AutoResponse.PHPRelay.Freeze
}
// IncidentsAutoCloseEnabled reports whether the auto-close path should
// run. Defaults to TRUE when the YAML key is absent so a fresh
// installation drains stale incidents without explicit opt-in. An
// explicit `incidents.auto_close.enabled: false` disables.
func (cfg *Config) IncidentsAutoCloseEnabled() bool {
return cfg.Incidents.AutoClose.Enabled == nil || *cfg.Incidents.AutoClose.Enabled
}
// IncidentsAutoCloseThresholds returns the per-kind idle thresholds in
// parsed form. Built from the operator's by_kind YAML map, falling
// back to safe defaults (mailbox_takeover=24h, mailbox_bruteforce=24h,
// web_account_compromise=7d, credential_spray=24h, web_attack=24h) when the
// operator did not supply a map.
// Unparseable durations are skipped silently so a typo in one entry
// does not disable the rest.
func (cfg *Config) IncidentsAutoCloseThresholds() map[string]time.Duration {
out := defaultIncidentAutoCloseThresholds()
for kind, raw := range cfg.Incidents.AutoClose.ByKind {
if raw == "" {
delete(out, kind)
continue
}
d, err := time.ParseDuration(raw)
if err != nil || d <= 0 {
continue
}
out[kind] = d
}
return out
}
func defaultIncidentAutoCloseThresholds() map[string]time.Duration {
return map[string]time.Duration{
"mailbox_takeover": 24 * time.Hour,
"mailbox_bruteforce": 24 * time.Hour,
"credential_spray": 24 * time.Hour,
"web_account_compromise": 7 * 24 * time.Hour,
"web_attack": 24 * time.Hour,
}
}
// IncidentsAutoBlockKinds returns the configured kinds set in the shape
// the correlator expects. Empty result means "any non-spray kind".
func (cfg *Config) IncidentsAutoBlockKinds() map[string]bool {
src := cfg.Incidents.AutoBlock.Kinds
out := make(map[string]bool, len(src))
for _, k := range src {
if k == "" {
continue
}
out[k] = true
}
return out
}
// IncidentsSpraySuppressionPerCheck returns the configured per-check map
// in the shape the correlator expects. Falls back to a safe default
// when the operator did not supply a list.
func (cfg *Config) IncidentsSpraySuppressionPerCheck() map[string]bool {
src := cfg.Incidents.SpraySuppression.PerCheck
if len(src) == 0 {
// Each name must be a Finding.Check emitted in production code;
// the spray detector gates on Check equality, so an unmatched name
// silently drops that source from spray collapse. email_auth_failure_realtime
// is emitted by daemon/watcher.go; pam_bruteforce and credential_stuffing
// by daemon/pam_listener.go (the latter is the one-IP-many-accounts signal).
src = []string{
"email_auth_failure_realtime",
"pam_bruteforce",
"credential_stuffing",
}
}
out := make(map[string]bool, len(src))
for _, c := range src {
if c == "" {
continue
}
out[c] = true
}
return out
}
// PHPRelayDryRunEnabled reports the YAML-level dry-run state for the
// email PHP-relay auto-freeze. Defaults to TRUE when dry_run was not
// set, which is the safe shipped behaviour: an operator who enables
// freeze without thinking about dry-run gets a dry-run, not a live
// freeze. nil-or-explicit-true => true; explicit-false => false.
func (cfg *Config) PHPRelayDryRunEnabled() bool {
return cfg.AutoResponse.PHPRelay.DryRun == nil || *cfg.AutoResponse.PHPRelay.DryRun
}
// AutoResponseDryRunEnabled mirrors PHPRelayDryRunEnabled: nil-or-true means true.
// When dry_run is absent from YAML the operator gets safe dry-run behaviour;
// explicit false is required to enable live nftables blocking.
func (cfg *Config) AutoResponseDryRunEnabled() bool {
return cfg.AutoResponse.DryRun == nil || *cfg.AutoResponse.DryRun
}
// Virtual-patch modes for auto_response.virtual_patch_exposed_files.
const (
VirtualPatchOff = "off"
VirtualPatchManual = "manual"
VirtualPatchAuto = "auto"
)
// VirtualPatchMode resolves auto_response.virtual_patch_exposed_files to one of
// off/manual/auto. Any unrecognised or empty value is treated as off so a typo
// never silently enables automatic file blocking.
func (cfg *Config) VirtualPatchMode() string {
switch strings.ToLower(strings.TrimSpace(cfg.AutoResponse.VirtualPatchExposedFiles)) {
case VirtualPatchManual:
return VirtualPatchManual
case VirtualPatchAuto:
return VirtualPatchAuto
default:
return VirtualPatchOff
}
}
// VulnerablePluginScanningEnabled reports whether the known-vulnerable plugin
// detector runs. Tri-state *bool: nil (omitted) defaults to on.
func (cfg *Config) VulnerablePluginScanningEnabled() bool {
return cfg.Detection.VulnerablePluginScanning == nil || *cfg.Detection.VulnerablePluginScanning
}
// DirectSMTPEgressDryRunEnabled reports the YAML-level dry-run state
// for the direct SMTP egress detector. Defaults to TRUE when dry_run
// was omitted (safety default). Operators must explicitly set
// `dry_run: false` to flip the detector to active mode.
func (c *Config) DirectSMTPEgressDryRunEnabled() bool {
if c.Detection.DirectSMTPEgress.DryRun == nil {
return true
}
return *c.Detection.DirectSMTPEgress.DryRun
}
// BPFEnforcementDryRunEnabled reports the YAML-level dry-run state for
// BPF cgroup-deny enforcement. Defaults to TRUE when dry_run is omitted
// (safety default). Operators must explicitly set `dry_run: false` to
// flip the in-kernel program to live denial.
func (c *Config) BPFEnforcementDryRunEnabled() bool {
if c.AutoResponseDryRunEnabled() {
return true
}
if c.BPFEnforcement.DirectSMTPEgress && c.DirectSMTPEgressDryRunEnabled() {
return true
}
if c.BPFEnforcement.DryRun == nil {
return true
}
return *c.BPFEnforcement.DryRun
}
// BotVerifyEnabled reports whether async PTR+forward-A verification of
// claimed search-engine bots is on. Default true; explicit false is
// honored so operators on air-gapped networks can disable DNS calls.
func (c *Config) BotVerifyEnabled() bool {
if c.Reputation.BotVerifyEnabled == nil {
return true
}
return *c.Reputation.BotVerifyEnabled
}
// BotRangesAutoUpdate reports whether the AI-crawler IP-range auto-updater
// runs. Default true (embedded snapshots are used regardless).
func (c *Config) BotRangesAutoUpdate() bool {
if c.Reputation.BotRanges.AutoUpdate == nil {
return true
}
return *c.Reputation.BotRanges.AutoUpdate
}
type defaultPresence struct {
smtpProbeThreshold bool
smtpBruteSlowThreshold bool
mailBruteSlowThreshold bool
forwardGuard forwardGuardPresence
phpRelay phpRelayPresence
retention retentionPresence
blockDigestMinBlock bool
thresholdsRollingCoverage bool
thresholdsDropperDetection bool
httpASNCrawlMinIPs bool
httpASNCrawlReverseProxy bool
xmlrpcThreshold bool
integrityImmutable bool
suppressWebmail bool
autoResponse autoResponsePresence
// firewall records which firewall: keys the operator wrote, so a
// partial block keeps the defaults for everything unlisted. nil means
// the whole block was absent.
firewall map[string]bool
}
// autoResponsePresence distinguishes omitted values from explicit zero or
// empty values that validation must reject while the corresponding feature is
// enabled.
type autoResponsePresence struct {
blockExpiry bool
netBlockThreshold bool
netBlockWindow bool
permBlockCount bool
permBlockInterval bool
}
// forwardGuardPresence records which forward-guard fields were set explicitly,
// so an operator's literal false survives the safe default-true.
type forwardGuardPresence struct {
dryRun bool
retention bool
bounceBackscatter bool
spamFlagged bool
malware bool
badSenderIP bool
authFail bool
}
// phpRelayPresence records PHP-relay fields where a literal zero has runtime
// meaning and must not be overwritten by defaults.
type phpRelayPresence struct {
fanoutDistinctRecipients bool
}
// retentionPresence records fields where a literal zero has runtime meaning and
// must not be overwritten by defaults.
type retentionPresence struct {
compactMinSizeMB bool
}
// ForwardGuardConfig configures the email forward-guard. Enabled is the master
// switch (default off); DryRun (default true) accounts without holding. Each
// hold signal is individually toggleable; unset signals default on but only
// matter once Enabled.
type ForwardGuardConfig struct {
Enabled bool `yaml:"enabled"`
DryRun bool `yaml:"dry_run"`
HoldSignals ForwardHoldSignals `yaml:"hold_signals"`
SkipForwarders []string `yaml:"skip_forwarders"`
QuarantineRetentionDays int `yaml:"quarantine_retention_days"`
}
// ForwardHoldSignals toggles which layered signals may hold a forward copy.
type ForwardHoldSignals struct {
BounceBackscatter bool `yaml:"bounce_backscatter"`
SpamFlagged bool `yaml:"spam_flagged"`
Malware bool `yaml:"malware"`
BadSenderIP bool `yaml:"bad_sender_ip"`
AuthFail bool `yaml:"auth_fail"`
}
func applyDefaults(cfg *Config, presence defaultPresence) {
for i, name := range cfg.Alerts.Email.DisabledChecks {
cfg.Alerts.Email.DisabledChecks[i] = CanonicalCheckName(strings.TrimSpace(name))
}
// Defaults
if cfg.StatePath == "" {
cfg.StatePath = "/var/lib/csm/state"
}
if cfg.Mode == "" {
cfg.Mode = ModeEnforce
} else {
cfg.Mode = normalizeMode(cfg.Mode)
}
// Binary immutability defaults on; a config written before the key existed
// must not read as "disable protection". Explicit false is kept.
if !presence.integrityImmutable {
cfg.Integrity.Immutable = true
}
// Webmail login suppression defaults on, as both shipped templates and
// the documentation say; a config that omits the key must not start
// alerting on every webmail login. Explicit false is kept.
if !presence.suppressWebmail {
cfg.Suppressions.SuppressWebmail = true
}
if cfg.Alerts.AuditLog.File.Enabled && cfg.Alerts.AuditLog.File.Path == "" {
cfg.Alerts.AuditLog.File.Path = "/var/log/csm/audit.jsonl"
}
if cfg.Alerts.AuditLog.Syslog.Enabled {
if cfg.Alerts.AuditLog.Syslog.Network == "" {
cfg.Alerts.AuditLog.Syslog.Network = "udp"
}
if cfg.Alerts.AuditLog.Syslog.Facility == "" {
cfg.Alerts.AuditLog.Syslog.Facility = "local0"
}
}
if cfg.Alerts.Webhook.HMACSecretEnv != "" {
if v := os.Getenv(cfg.Alerts.Webhook.HMACSecretEnv); v != "" {
cfg.Alerts.Webhook.HMACSecret = v
}
}
{
bd := &cfg.Alerts.BlockDigest
if bd.Interval == "" {
bd.Interval = "1h"
}
if bd.SendOn == "" {
bd.SendOn = "any"
}
// A literal min_block: 0 means "send empty heartbeat digests" and must
// survive; only an absent key falls back to 1.
if bd.MinBlock == 0 && !presence.blockDigestMinBlock {
bd.MinBlock = 1
}
}
if cfg.Signatures.RulesDir == "" {
cfg.Signatures.RulesDir = "/opt/csm/rules"
}
if cfg.Signatures.YaraForge.Tier == "" {
cfg.Signatures.YaraForge.Tier = "core"
}
if cfg.Signatures.YaraForge.UpdateInterval == "" {
cfg.Signatures.YaraForge.UpdateInterval = "168h"
}
if cfg.Reputation.BotRanges.UpdateInterval == "" {
cfg.Reputation.BotRanges.UpdateInterval = "24h"
}
if cfg.WebUI.SessionLifetime == "" {
cfg.WebUI.SessionLifetime = DefaultBrowserSessionLifetime
}
if cfg.WebUI.SessionIdleTimeout == "" {
cfg.WebUI.SessionIdleTimeout = DefaultBrowserSessionIdleTimeout
}
if cfg.WebUI.Listen == "" {
cfg.WebUI.Listen = "0.0.0.0:9443"
}
if cfg.WebUI.AuthToken != "" && len(cfg.WebUI.Tokens) == 0 {
cfg.WebUI.Tokens = []WebUIToken{{
Name: "legacy-auth-token", Token: cfg.WebUI.AuthToken, Scope: "admin",
}}
}
if cfg.Thresholds.MailQueueWarn == 0 {
cfg.Thresholds.MailQueueWarn = 500
}
if cfg.Thresholds.MailQueueCrit == 0 {
cfg.Thresholds.MailQueueCrit = 2000
}
if cfg.Thresholds.StateExpiryHours == 0 {
cfg.Thresholds.StateExpiryHours = 24
}
if cfg.Thresholds.DeepScanIntervalMin == 0 {
cfg.Thresholds.DeepScanIntervalMin = 60
}
if cfg.Thresholds.WPCoreCheckIntervalMin == 0 {
cfg.Thresholds.WPCoreCheckIntervalMin = 60
}
if cfg.Thresholds.WebshellScanIntervalMin == 0 {
cfg.Thresholds.WebshellScanIntervalMin = 30
}
if cfg.Thresholds.FilesystemScanIntervalMin == 0 {
cfg.Thresholds.FilesystemScanIntervalMin = 30
}
if cfg.Thresholds.ExposedFileScanDepth == 0 {
cfg.Thresholds.ExposedFileScanDepth = DefaultExposedFileScanDepth
}
if cfg.Thresholds.PluginCheckIntervalMin == 0 {
cfg.Thresholds.PluginCheckIntervalMin = 1440
}
if cfg.Thresholds.BruteForceWindow == 0 {
cfg.Thresholds.BruteForceWindow = 5000
}
if cfg.Thresholds.PHPConfigWalkMaxDirs <= 0 {
cfg.Thresholds.PHPConfigWalkMaxDirs = 50000
}
if cfg.Thresholds.PHPConfigWalkMaxEntries <= 0 {
cfg.Thresholds.PHPConfigWalkMaxEntries = 500000
}
if cfg.Thresholds.DomlogMaxFiles == 0 {
cfg.Thresholds.DomlogMaxFiles = 500
}
if cfg.Thresholds.AccountScanMaxFiles == 0 {
cfg.Thresholds.AccountScanMaxFiles = 10000
}
if cfg.Thresholds.CrontabBase64BlobMaxBytes == 0 {
cfg.Thresholds.CrontabBase64BlobMaxBytes = 16384
}
if cfg.Thresholds.DomlogTailLines == 0 {
cfg.Thresholds.DomlogTailLines = 500
}
if cfg.Thresholds.DomlogMaxAgeMin == 0 {
cfg.Thresholds.DomlogMaxAgeMin = 30
}
if cfg.Thresholds.MailLogTailLines == 0 {
cfg.Thresholds.MailLogTailLines = 500
}
if cfg.Thresholds.SyslogMessagesTailLines == 0 {
cfg.Thresholds.SyslogMessagesTailLines = 200
}
if cfg.Thresholds.FTPFailWindowMin == 0 {
cfg.Thresholds.FTPFailWindowMin = 30
}
if cfg.Thresholds.CredStuffingDistinctAccounts == 0 {
cfg.Thresholds.CredStuffingDistinctAccounts = 5
}
if cfg.Thresholds.PAMBruteforceThreshold == 0 {
cfg.Thresholds.PAMBruteforceThreshold = 5
}
if cfg.Thresholds.PAMBruteforceWindowMin == 0 {
cfg.Thresholds.PAMBruteforceWindowMin = 10
}
if cfg.Thresholds.ModSecEscalationHits == 0 {
cfg.Thresholds.ModSecEscalationHits = 3
}
if cfg.Thresholds.ModSecEscalationWindowMin == 0 {
cfg.Thresholds.ModSecEscalationWindowMin = 10
}
if cfg.Thresholds.ModSecLowConfidenceEscalationHits == 0 {
cfg.Thresholds.ModSecLowConfidenceEscalationHits = 30
}
if cfg.Thresholds.HTTPFloodWindowMin <= 0 {
cfg.Thresholds.HTTPFloodWindowMin = 5
}
// HTTPFloodThreshold has no nonzero default: 0 means disabled and
// that is the shipped behavior.
if cfg.Thresholds.HTTPUASpoofThreshold <= 0 {
cfg.Thresholds.HTTPUASpoofThreshold = 30
}
if cfg.Thresholds.HTTPScannerErrorPct == 0 {
cfg.Thresholds.HTTPScannerErrorPct = DefaultHTTPScannerErrorPct
}
if cfg.Thresholds.HTTPScannerMinDistinctPaths == 0 {
cfg.Thresholds.HTTPScannerMinDistinctPaths = DefaultHTTPScannerMinDistinctPaths
}
if len(cfg.Thresholds.HTTPScannerStatusCodes) == 0 {
cfg.Thresholds.HTTPScannerStatusCodes = DefaultHTTPScannerStatusCodes()
}
if cfg.AutoResponse.HTTPScannerAction == "" {
cfg.AutoResponse.HTTPScannerAction = "challenge"
}
if cfg.Thresholds.SMTPBruteForceThreshold == 0 {
cfg.Thresholds.SMTPBruteForceThreshold = 5
}
if cfg.Thresholds.SMTPBruteForceWindowMin == 0 {
cfg.Thresholds.SMTPBruteForceWindowMin = 10
}
if cfg.Thresholds.SMTPBruteForceSuppressMin == 0 {
cfg.Thresholds.SMTPBruteForceSuppressMin = 60
}
if cfg.Thresholds.SMTPBruteForceSubnetThresh == 0 {
cfg.Thresholds.SMTPBruteForceSubnetThresh = 8
}
if cfg.Thresholds.SMTPAccountSprayThreshold == 0 {
cfg.Thresholds.SMTPAccountSprayThreshold = 12
}
if cfg.Thresholds.SMTPBruteForceMaxTracked == 0 {
cfg.Thresholds.SMTPBruteForceMaxTracked = 20000
}
// Explicit 0 disables the slow-brute signal; absence gets the default.
if cfg.Thresholds.SMTPBruteForceSlowThreshold == 0 && !presence.smtpBruteSlowThreshold {
cfg.Thresholds.SMTPBruteForceSlowThreshold = 40
}
if cfg.Thresholds.SMTPBruteForceSlowWindowMin == 0 {
cfg.Thresholds.SMTPBruteForceSlowWindowMin = 360
}
if cfg.Thresholds.SMTPProbeThreshold == 0 && !presence.smtpProbeThreshold {
cfg.Thresholds.SMTPProbeThreshold = 100
}
if cfg.Thresholds.SMTPProbeWindowMin == 0 {
cfg.Thresholds.SMTPProbeWindowMin = 5
}
if cfg.Thresholds.SMTPProbeSuppressMin == 0 {
cfg.Thresholds.SMTPProbeSuppressMin = 60
}
if cfg.Thresholds.SMTPProbeMaxTracked == 0 {
cfg.Thresholds.SMTPProbeMaxTracked = 20000
}
if cfg.Thresholds.MailBruteForceThreshold == 0 {
cfg.Thresholds.MailBruteForceThreshold = 5
}
if cfg.Thresholds.MailBruteForceWindowMin == 0 {
cfg.Thresholds.MailBruteForceWindowMin = 10
}
if cfg.Thresholds.MailBruteForceSuppressMin == 0 {
cfg.Thresholds.MailBruteForceSuppressMin = 60
}
if cfg.Thresholds.MailBruteForceSubnetThresh == 0 {
cfg.Thresholds.MailBruteForceSubnetThresh = 8
}
if cfg.Thresholds.MailAccountSprayThreshold == 0 {
cfg.Thresholds.MailAccountSprayThreshold = 12
}
if cfg.Thresholds.MailBruteForceMaxTracked == 0 {
cfg.Thresholds.MailBruteForceMaxTracked = 20000
}
// Explicit 0 disables the slow-brute signal; absence gets the default.
if cfg.Thresholds.MailBruteForceSlowThreshold == 0 && !presence.mailBruteSlowThreshold {
cfg.Thresholds.MailBruteForceSlowThreshold = 40
}
if cfg.Thresholds.MailBruteForceSlowWindowMin == 0 {
cfg.Thresholds.MailBruteForceSlowWindowMin = 360
}
if cfg.Alerts.MaxPerHour == 0 {
cfg.Alerts.MaxPerHour = 30
}
if cfg.Challenge.ListenAddr == "" {
cfg.Challenge.ListenAddr = "127.0.0.1"
}
if cfg.Challenge.ListenPort == 0 {
cfg.Challenge.ListenPort = 8439
}
if cfg.Challenge.Difficulty == 0 {
cfg.Challenge.Difficulty = 2
}
if cfg.Firewall == nil {
cfg.Firewall = firewall.DefaultConfig()
} else {
applyFirewallFieldDefaults(cfg.Firewall, presence.firewall)
}
if len(cfg.GeoIP.Editions) == 0 {
cfg.GeoIP.Editions = []string{"GeoLite2-City", "GeoLite2-ASN"}
}
if cfg.GeoIP.UpdateInterval == "" {
cfg.GeoIP.UpdateInterval = "24h"
}
EmailAVDefaults(&cfg.EmailAV)
if cfg.EmailProtection.PasswordCheckIntervalMin == 0 {
cfg.EmailProtection.PasswordCheckIntervalMin = 1440
}
if cfg.EmailProtection.RateWarnThreshold == 0 {
cfg.EmailProtection.RateWarnThreshold = 50
}
if cfg.EmailProtection.RateCritThreshold == 0 {
cfg.EmailProtection.RateCritThreshold = 100
}
if cfg.EmailProtection.RateWindowMin == 0 {
cfg.EmailProtection.RateWindowMin = 10
}
applyForwardGuardDefaults(&cfg.EmailProtection.ForwardGuard, presence.forwardGuard)
// EmailProtection.PHPRelay defaults. Freeze/DryRun are *bool and
// remain nil here -- accessors resolve the safe defaults
// (PHPRelayFreezeEnabled / PHPRelayDryRunEnabled) so we do NOT
// mutate them. AccountVolumePerHour stays at 0 by default to mark
// "auto-derive from cPanel maxemailsperhour" downstream.
if cfg.EmailProtection.PHPRelay.RateWindowMin == 0 {
cfg.EmailProtection.PHPRelay.RateWindowMin = 5
}
if cfg.EmailProtection.PHPRelay.HeaderScoreVolumeMin == 0 {
cfg.EmailProtection.PHPRelay.HeaderScoreVolumeMin = 5
}
if cfg.EmailProtection.PHPRelay.AbsoluteVolumePerHour == 0 {
cfg.EmailProtection.PHPRelay.AbsoluteVolumePerHour = 30
}
if cfg.EmailProtection.PHPRelay.ReputationFailuresPer24h == 0 {
cfg.EmailProtection.PHPRelay.ReputationFailuresPer24h = 3
}
if cfg.EmailProtection.PHPRelay.FanoutDistinctScripts == 0 {
cfg.EmailProtection.PHPRelay.FanoutDistinctScripts = 3
}
if cfg.EmailProtection.PHPRelay.FanoutDistinctRecipients == 0 && !presence.phpRelay.fanoutDistinctRecipients {
cfg.EmailProtection.PHPRelay.FanoutDistinctRecipients = 5
}
if cfg.EmailProtection.PHPRelay.FanoutWindowMin == 0 {
cfg.EmailProtection.PHPRelay.FanoutWindowMin = 5
}
if cfg.EmailProtection.PHPRelay.BaselineSigma == 0 {
cfg.EmailProtection.PHPRelay.BaselineSigma = 3.0
}
if cfg.EmailProtection.PHPRelay.BaselineObservationDays == 0 {
cfg.EmailProtection.PHPRelay.BaselineObservationDays = 7
}
if cfg.EmailProtection.PHPRelay.PoliciesDir == "" {
cfg.EmailProtection.PHPRelay.PoliciesDir = "/opt/csm/policies/php_relay"
}
if cfg.AutoResponse.PHPRelay.MaxActionsPerMinute == 0 {
cfg.AutoResponse.PHPRelay.MaxActionsPerMinute = 60
}
if cfg.AutoResponse.MaxFileActionsPerHour == 0 {
cfg.AutoResponse.MaxFileActionsPerHour = DefaultMaxFileActionsPerHour
}
if cfg.AutoResponse.MaxFileActionsPerAccountPerHour == 0 {
cfg.AutoResponse.MaxFileActionsPerAccountPerHour = DefaultMaxFileActionsPerAccountPerHour
}
if cfg.AutoResponse.MaxFileActionFailuresPerHour == 0 {
cfg.AutoResponse.MaxFileActionFailuresPerHour = DefaultMaxFileActionFailuresPerHour
}
if cfg.AutoResponse.MaxBlocksPerHour == 0 {
cfg.AutoResponse.MaxBlocksPerHour = DefaultMaxBlocksPerHour
}
// The block path keeps fallbacks for Config values assembled in code.
// Loaded configs resolve omitted values here so status surfaces report the
// effective settings and validation can reject explicit invalid values.
if cfg.AutoResponse.BlockExpiry == "" && !presence.autoResponse.blockExpiry {
cfg.AutoResponse.BlockExpiry = DefaultBlockExpiry
}
if cfg.AutoResponse.NetBlockThreshold == 0 && !presence.autoResponse.netBlockThreshold {
cfg.AutoResponse.NetBlockThreshold = DefaultNetBlockThreshold
}
if cfg.AutoResponse.NetBlockWindow == "" && !presence.autoResponse.netBlockWindow {
cfg.AutoResponse.NetBlockWindow = DefaultNetBlockWindow
}
if cfg.AutoResponse.PermBlockCount == 0 && !presence.autoResponse.permBlockCount {
cfg.AutoResponse.PermBlockCount = DefaultPermBlockCount
}
if cfg.AutoResponse.PermBlockInterval == "" && !presence.autoResponse.permBlockInterval {
cfg.AutoResponse.PermBlockInterval = DefaultPermBlockInterval
}
if cfg.AutoResponse.MailAuthRecovery.DownGrace == "" {
cfg.AutoResponse.MailAuthRecovery.DownGrace = "10m"
}
if cfg.AutoResponse.MailAuthRecovery.MaxRestartsPerHour == 0 {
cfg.AutoResponse.MailAuthRecovery.MaxRestartsPerHour = 3
}
if cfg.AutoResponse.MailAuthRecovery.RestartCommand == "" {
cfg.AutoResponse.MailAuthRecovery.RestartCommand = "/usr/local/cpanel/scripts/restartsrv_dovecot"
}
// Performance defaults.
// Enabled is a tri-state *bool: nil means "use system default (on)", true means
// explicitly enabled, false means explicitly disabled. We do NOT apply a default
// here so that callers can distinguish "operator left it unset" (nil) from
// "operator set it to true" (&true). All callers must nil-check before dereferencing;
// perfEnabled() in checks/performance.go treats nil as true.
if cfg.Performance.LoadHighMultiplier == 0 {
cfg.Performance.LoadHighMultiplier = 1.0
}
if cfg.Performance.LoadCriticalMultiplier == 0 {
cfg.Performance.LoadCriticalMultiplier = 2.0
}
if cfg.Performance.PHPProcessWarnPerUser == 0 {
cfg.Performance.PHPProcessWarnPerUser = 20
}
if cfg.Performance.PHPProcessCriticalTotalMult == 0 {
cfg.Performance.PHPProcessCriticalTotalMult = 5
}
if cfg.Performance.ErrorLogWarnSizeMB == 0 {
cfg.Performance.ErrorLogWarnSizeMB = 50
}
if cfg.Performance.MySQLJoinBufferMaxMB == 0 {
cfg.Performance.MySQLJoinBufferMaxMB = 64
}
if cfg.Performance.MySQLWaitTimeoutMax == 0 {
cfg.Performance.MySQLWaitTimeoutMax = 3600
}
if cfg.Performance.MySQLMaxConnectionsPerUser == 0 {
cfg.Performance.MySQLMaxConnectionsPerUser = 10
}
if cfg.Performance.RedisBgsaveMinInterval == 0 {
cfg.Performance.RedisBgsaveMinInterval = 900
}
if cfg.Performance.RedisLargeDatasetGB == 0 {
cfg.Performance.RedisLargeDatasetGB = 4
}
if cfg.Performance.WPMemoryLimitMaxMB == 0 {
cfg.Performance.WPMemoryLimitMaxMB = 512
}
if cfg.Performance.WPTransientWarnMB == 0 {
cfg.Performance.WPTransientWarnMB = 1
}
if cfg.Performance.WPTransientCriticalMB == 0 {
cfg.Performance.WPTransientCriticalMB = 10
}
if cfg.Performance.WPCronFix.IntervalMinutes == 0 {
cfg.Performance.WPCronFix.IntervalMinutes = 15
}
if cfg.Cloudflare.RefreshHours == 0 {
cfg.Cloudflare.RefreshHours = 6
}
// Retention: defaults apply whether or not the feature is enabled, so
// that flipping `enabled: true` without further tuning gives the
// documented behaviour.
if cfg.Retention.FindingsDays == 0 {
cfg.Retention.FindingsDays = 90
}
if cfg.Retention.HistoryDays == 0 {
cfg.Retention.HistoryDays = 30
}
if cfg.Retention.ReputationDays == 0 {
cfg.Retention.ReputationDays = 180
}
if cfg.Retention.SweepInterval == "" {
cfg.Retention.SweepInterval = "24h"
}
if cfg.Retention.CompactMinSizeMB == 0 && !presence.retention.compactMinSizeMB {
cfg.Retention.CompactMinSizeMB = 128
}
if cfg.Retention.CompactFillRatio == 0 {
cfg.Retention.CompactFillRatio = 0.5
}
if cfg.MailLogs.Source == "" {
cfg.MailLogs.Source = "auto"
}
if len(cfg.MailLogs.Units) == 0 {
cfg.MailLogs.Units = []string{"postfix", "dovecot"}
}
if cfg.Updates.Interval == "" {
cfg.Updates.Interval = "24h"
}
if cfg.Updates.PackageName == "" {
cfg.Updates.PackageName = "csm"
}
if cfg.Thresholds.MailBruteAccountKey == "" {
cfg.Thresholds.MailBruteAccountKey = "builtin:dovecot-user"
}
if cfg.Thresholds.FullScanMaxFileMB == 0 {
cfg.Thresholds.FullScanMaxFileMB = 16
}
if cfg.Thresholds.ScanJobRetention == 0 {
cfg.Thresholds.ScanJobRetention = 20
}
// A literal rolling_coverage: false means "disable rolling coverage" and
// must survive; only an absent key falls back to the safe default of true.
if !presence.thresholdsRollingCoverage {
cfg.Thresholds.RollingCoverage = true
}
// Same shape as rolling_coverage: a literal dropper_detection: false is an
// operator opt-out; only an absent key defaults to enabled.
if !presence.thresholdsDropperDetection {
cfg.Thresholds.DropperDetection = true
}
if cfg.Thresholds.DropperUnlinkTTLSec == 0 {
cfg.Thresholds.DropperUnlinkTTLSec = DefaultDropperUnlinkTTLSec
}
// http_asn_crawl: min_ips uses presence so an explicit 0 disables the
// detector; an absent key defaults to 25.
if !presence.httpASNCrawlMinIPs && cfg.Thresholds.HTTPASNCrawlMinIPs == 0 {
cfg.Thresholds.HTTPASNCrawlMinIPs = DefaultHTTPASNCrawlMinIPs
}
// xmlrpc_threshold uses presence so an explicit 0 disables the check; an
// absent key defaults to 100.
if !presence.xmlrpcThreshold && cfg.Thresholds.XMLRPCThreshold == 0 {
cfg.Thresholds.XMLRPCThreshold = DefaultXMLRPCThreshold
}
if cfg.Thresholds.HTTPASNCrawlMinExpensive == 0 {
cfg.Thresholds.HTTPASNCrawlMinExpensive = DefaultHTTPASNCrawlMinExpensive
}
if cfg.Thresholds.HTTPASNCrawlMinSharePct == 0 {
cfg.Thresholds.HTTPASNCrawlMinSharePct = DefaultHTTPASNCrawlMinSharePct
}
if cfg.Thresholds.HTTPASNCrawlHighAmpPct == 0 {
cfg.Thresholds.HTTPASNCrawlHighAmpPct = DefaultHTTPASNCrawlHighAmpPct
}
if cfg.Thresholds.HTTPASNCrawlHighVolumeMult == 0 {
cfg.Thresholds.HTTPASNCrawlHighVolumeMult = DefaultHTTPASNCrawlHighVolMult
}
if cfg.Thresholds.HTTPASNCrawlMaxPrefix == 0 {
cfg.Thresholds.HTTPASNCrawlMaxPrefix = DefaultHTTPASNCrawlMaxPrefix
}
if cfg.Thresholds.HTTPASNCrawl16PrefPct == 0 {
cfg.Thresholds.HTTPASNCrawl16PrefPct = DefaultHTTPASNCrawl16PrefPct
}
if cfg.Thresholds.HTTPASNCrawlMaxTrackedIPs == 0 {
cfg.Thresholds.HTTPASNCrawlMaxTrackedIPs = DefaultHTTPASNCrawlMaxTrackedIPs
}
if cfg.Thresholds.HTTPASNCrawlWindowMin == 0 {
cfg.Thresholds.HTTPASNCrawlWindowMin = DefaultHTTPASNCrawlWindowMin
}
// Reverse-proxy seed: absent key keeps the safety seed; an explicit empty
// list is honored as an operator override.
if !presence.httpASNCrawlReverseProxy && len(cfg.Thresholds.HTTPASNCrawlReverseProxyASNs) == 0 {
cfg.Thresholds.HTTPASNCrawlReverseProxyASNs = append([]uint(nil), httpASNCrawlReverseProxySeed...)
}
if cfg.AutoResponse.HTTPASNCrawlTempban == "" {
cfg.AutoResponse.HTTPASNCrawlTempban = DefaultHTTPASNCrawlTempban
}
if cfg.Reputation.Rspamd.URL == "" {
cfg.Reputation.Rspamd.URL = "http://127.0.0.1:11334"
}
// Token resolution happens at query time (see RspamdSource.Score).
if cfg.Reputation.Upstream.CacheTTLMin == 0 {
cfg.Reputation.Upstream.CacheTTLMin = 15
}
if cfg.Reputation.Upstream.TimeoutSec == 0 {
cfg.Reputation.Upstream.TimeoutSec = 5
}
// Token resolution happens at query time (UpstreamSource.resolveToken).
if cfg.AutoResponse.VerdictCallback.TimeoutSec == 0 {
cfg.AutoResponse.VerdictCallback.TimeoutSec = 2 // tight; the hook is on the block hot path
}
// Secret resolution happens at call time (verdict.Client reads env per call).
// Direct SMTP egress detector defaults. Backend "auto" lets the runtime
// pick BPF where available and fall back to legacy polling. Standard
// submission/relay ports cover the bulk of mass-mail abuse seen in the
// wild; operators can override via YAML to add e.g. 2525.
if cfg.Detection.DirectSMTPEgress.Backend == "" {
cfg.Detection.DirectSMTPEgress.Backend = "auto"
}
if len(cfg.Detection.DirectSMTPEgress.Ports) == 0 {
cfg.Detection.DirectSMTPEgress.Ports = []int{25, 465, 587}
}
}
// MaxConfigBytes caps the YAML config body size LoadBytes will parse.
// Real CSM configs (main + every drop-in fragment) top out near 64 KB
// even with verbose comments; 4 MB is several orders of magnitude
// above legitimate use and far below the size at which a malformed
// or attacker-supplied file would force the YAML parser to allocate
// gigabytes of intermediate state.
const MaxConfigBytes = 4 * 1024 * 1024
var errConfigTooLarge = errors.New("config input exceeds byte cap")
func readConfigBytesLimited(r io.Reader) ([]byte, error) {
data, err := io.ReadAll(io.LimitReader(r, MaxConfigBytes+1))
if err != nil {
return nil, err
}
if int64(len(data)) > MaxConfigBytes {
return nil, errConfigTooLarge
}
return data, nil
}
// LoadBytes decodes a YAML config body and applies all defaults,
// matching Load. ConfigFile is left empty; the caller sets it.
func LoadBytes(data []byte) (*Config, error) {
if len(data) > MaxConfigBytes {
return nil, fmt.Errorf("parsing config: input size %d exceeds %d byte cap", len(data), MaxConfigBytes)
}
presence, err := defaultPresenceFromYAML(data)
if err != nil {
return nil, fmt.Errorf("parsing config: %w", err)
}
cfg := &Config{}
dec := yaml.NewDecoder(bytes.NewReader(data))
dec.KnownFields(true)
if err := dec.Decode(cfg); err != nil && !errors.Is(err, io.EOF) {
return nil, fmt.Errorf("parsing config: %w", err)
}
applyDefaults(cfg, presence)
if err := validateMode(cfg); err != nil {
return nil, err
}
if err := validateWebUITokens(cfg); err != nil {
return nil, err
}
if err := validateMailLogs(cfg); err != nil {
return nil, err
}
if err := validateMailBruteAccountKey(cfg); err != nil {
return nil, err
}
if err := validateReputation(cfg); err != nil {
return nil, err
}
if err := validateVerifiedBotsConfig(cfg); err != nil {
return nil, err
}
if err := validateVerdictCallback(cfg); err != nil {
return nil, err
}
if err := validateDirectSMTPEgress(cfg); err != nil {
return nil, err
}
if err := validateForwardGuard(cfg); err != nil {
return nil, err
}
if err := validateHTTPASNCrawl(cfg); err != nil {
return nil, err
}
if err := validateFirewallConfig(cfg); err != nil {
return nil, err
}
return cfg, nil
}
func defaultPresenceFromYAML(data []byte) (defaultPresence, error) {
var presence defaultPresence
if len(bytes.TrimSpace(data)) == 0 {
return presence, nil
}
var raw struct {
Thresholds map[string]yaml.Node `yaml:"thresholds"`
EmailProtection struct {
ForwardGuard map[string]yaml.Node `yaml:"forward_guard"`
PHPRelay map[string]yaml.Node `yaml:"php_relay"`
} `yaml:"email_protection"`
Alerts struct {
BlockDigest map[string]yaml.Node `yaml:"block_digest"`
} `yaml:"alerts"`
Retention map[string]yaml.Node `yaml:"retention"`
Integrity map[string]yaml.Node `yaml:"integrity"`
Suppressions map[string]yaml.Node `yaml:"suppressions"`
AutoResponse map[string]yaml.Node `yaml:"auto_response"`
Firewall map[string]yaml.Node `yaml:"firewall"`
}
if err := yaml.Unmarshal(data, &raw); err != nil {
return presence, err
}
if raw.Firewall != nil {
presence.firewall = make(map[string]bool, len(raw.Firewall))
for key, node := range raw.Firewall {
// A null node means "no value", which is what an absent key means,
// so it must not count as present. The distinction matters for
// every field whose zero value is a real setting: `x: null` has to
// keep the shipped default while `x: 0` (or `[]`, or `false`)
// stays the operator's explicit choice.
presence.firewall[key] = !yamlNodeIsNull(&node)
}
}
_, presence.integrityImmutable = raw.Integrity["immutable"]
_, presence.suppressWebmail = raw.Suppressions["suppress_webmail_alerts"]
if node, ok := raw.AutoResponse["block_expiry"]; ok && !yamlNodeIsNull(&node) {
presence.autoResponse.blockExpiry = true
}
if node, ok := raw.AutoResponse["netblock_threshold"]; ok && !yamlNodeIsNull(&node) {
presence.autoResponse.netBlockThreshold = true
}
if node, ok := raw.AutoResponse["netblock_window"]; ok && !yamlNodeIsNull(&node) {
presence.autoResponse.netBlockWindow = true
}
if node, ok := raw.AutoResponse["permblock_count"]; ok && !yamlNodeIsNull(&node) {
presence.autoResponse.permBlockCount = true
}
if node, ok := raw.AutoResponse["permblock_interval"]; ok && !yamlNodeIsNull(&node) {
presence.autoResponse.permBlockInterval = true
}
_, presence.smtpProbeThreshold = raw.Thresholds["smtp_probe_threshold"]
if node, ok := raw.Thresholds["smtp_bruteforce_slow_threshold"]; ok && !yamlNodeIsNull(&node) {
presence.smtpBruteSlowThreshold = true
}
if node, ok := raw.Thresholds["mail_bruteforce_slow_threshold"]; ok && !yamlNodeIsNull(&node) {
presence.mailBruteSlowThreshold = true
}
_, presence.thresholdsRollingCoverage = raw.Thresholds["rolling_coverage"]
_, presence.thresholdsDropperDetection = raw.Thresholds["dropper_detection"]
_, presence.httpASNCrawlMinIPs = raw.Thresholds["http_asn_crawl_min_ips"]
_, presence.httpASNCrawlReverseProxy = raw.Thresholds["http_asn_crawl_reverse_proxy_asns"]
_, presence.xmlrpcThreshold = raw.Thresholds["xmlrpc_threshold"]
_, presence.phpRelay.fanoutDistinctRecipients = raw.EmailProtection.PHPRelay["fanout_distinct_recipients"]
if node, ok := raw.Alerts.BlockDigest["min_block"]; ok && node.Tag != "!!null" {
presence.blockDigestMinBlock = true
}
if node, ok := raw.Retention["compact_min_size_mb"]; ok && node.Tag != "!!null" {
presence.retention.compactMinSizeMB = true
}
fg := raw.EmailProtection.ForwardGuard
_, presence.forwardGuard.dryRun = fg["dry_run"]
_, presence.forwardGuard.retention = fg["quarantine_retention_days"]
if node, ok := fg["hold_signals"]; ok {
var sig map[string]yaml.Node
if err := node.Decode(&sig); err != nil {
return presence, err
}
_, presence.forwardGuard.bounceBackscatter = sig["bounce_backscatter"]
_, presence.forwardGuard.spamFlagged = sig["spam_flagged"]
_, presence.forwardGuard.malware = sig["malware"]
_, presence.forwardGuard.badSenderIP = sig["bad_sender_ip"]
_, presence.forwardGuard.authFail = sig["auth_fail"]
}
return presence, nil
}
func yamlNodeIsNull(node *yaml.Node) bool {
seen := make(map[*yaml.Node]struct{})
for node != nil && node.Kind == yaml.AliasNode {
if node.Alias == nil {
return false
}
if _, ok := seen[node.Alias]; ok {
return false
}
seen[node.Alias] = struct{}{}
node = node.Alias
}
return node != nil && node.Tag == "!!null"
}
// applyForwardGuardDefaults fills unset forward-guard fields. dry_run and every
// hold signal default on; an explicit false (tracked via presence) is kept.
func applyForwardGuardDefaults(fg *ForwardGuardConfig, p forwardGuardPresence) {
if !p.dryRun {
fg.DryRun = true
}
if !p.bounceBackscatter {
fg.HoldSignals.BounceBackscatter = true
}
if !p.spamFlagged {
fg.HoldSignals.SpamFlagged = true
}
if !p.malware {
fg.HoldSignals.Malware = true
}
if !p.badSenderIP {
fg.HoldSignals.BadSenderIP = true
}
if !p.authFail {
fg.HoldSignals.AuthFail = true
}
if !p.retention {
fg.QuarantineRetentionDays = 14
}
}
// applyFirewallFieldDefaults fills absent firewall: keys from DefaultConfig.
// The whole-pointer default used to be all-or-nothing: any partial firewall
// block ("enabled: true" alone) lost every unlisted default and produced a
// DROP-policy chain with empty accept lists - an instant lockout. Presence
// tracking keeps operator-explicit values, including empty lists, zeros,
// and false. Fields whose default is the zero value need no entry here.
func applyFirewallFieldDefaults(fc *firewall.FirewallConfig, present map[string]bool) {
def := firewall.DefaultConfig()
if !present["tcp_in"] {
fc.TCPIn = def.TCPIn
}
if !present["tcp_out"] {
fc.TCPOut = def.TCPOut
}
if !present["udp_in"] {
fc.UDPIn = def.UDPIn
}
if !present["udp_out"] {
fc.UDPOut = def.UDPOut
}
if !present["restricted_tcp"] {
fc.RestrictedTCP = def.RestrictedTCP
}
if !present["passive_ftp_start"] {
fc.PassiveFTPStart = def.PassiveFTPStart
}
if !present["passive_ftp_end"] {
fc.PassiveFTPEnd = def.PassiveFTPEnd
}
if !present["conn_rate_limit"] {
fc.ConnRateLimit = def.ConnRateLimit
}
if !present["syn_flood_protection"] {
fc.SYNFloodProtection = def.SYNFloodProtection
}
if !present["conn_limit"] {
fc.ConnLimit = def.ConnLimit
}
if !present["port_flood"] {
fc.PortFlood = def.PortFlood
}
if !present["udp_flood"] {
fc.UDPFlood = def.UDPFlood
}
if !present["udp_flood_rate"] {
fc.UDPFloodRate = def.UDPFloodRate
}
if !present["udp_flood_burst"] {
fc.UDPFloodBurst = def.UDPFloodBurst
}
if !present["drop_nolog"] {
fc.DropNoLog = def.DropNoLog
}
if !present["deny_ip_limit"] {
fc.DenyIPLimit = def.DenyIPLimit
}
if !present["deny_temp_ip_limit"] {
fc.DenyTempIPLimit = def.DenyTempIPLimit
}
if !present["smtp_ports"] {
fc.SMTPPorts = def.SMTPPorts
}
if !present["log_dropped"] {
fc.LogDropped = def.LogDropped
}
if !present["log_rate"] {
fc.LogRate = def.LogRate
}
}
// validateForwardGuard rejects nonsensical or unsafe forward-guard configs. The
// guard only matters when enabled, so an off guard never fails validation.
func validateForwardGuard(cfg *Config) error {
fg := cfg.EmailProtection.ForwardGuard
if !fg.Enabled {
return nil
}
if fg.QuarantineRetentionDays <= 0 {
return fmt.Errorf("email_protection.forward_guard.quarantine_retention_days must be > 0 when enabled")
}
// Enforce mode requires a signal exim can actually evaluate at routing time.
// spam_flagged/malware/auth_fail are accounted in dry-run but not yet
// enforceable, so enforcing with only those on would silently hold nothing.
if !fg.DryRun && !fg.HoldSignals.BounceBackscatter && !fg.HoldSignals.BadSenderIP {
return fmt.Errorf("email_protection.forward_guard: enforce mode (dry_run:false) requires bounce_backscatter or bad_sender_ip enabled; spam_flagged/malware/auth_fail are dry-run only until exim content scanning is enabled")
}
return nil
}
func validateReputation(cfg *Config) error {
up := cfg.Reputation.Upstream
if up.CacheTTLMin != 0 && (up.CacheTTLMin < 1 || up.CacheTTLMin > 1440) {
return fmt.Errorf("reputation.upstream.cache_ttl_min must be between 1 and 1440")
}
if up.TimeoutSec != 0 && (up.TimeoutSec < 1 || up.TimeoutSec > 60) {
return fmt.Errorf("reputation.upstream.timeout_sec must be between 1 and 60")
}
if !up.Enabled {
return nil
}
if strings.TrimSpace(up.URL) == "" {
return fmt.Errorf("reputation.upstream.enabled=true but url is empty")
}
parsed, err := url.Parse(up.URL)
if err != nil {
return fmt.Errorf("reputation.upstream.url: %w", err)
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return fmt.Errorf("reputation.upstream.url must use http or https")
}
if parsed.Host == "" {
return fmt.Errorf("reputation.upstream.url must include host")
}
if parsed.Scheme == "http" && !isLoopbackHost(parsed.Hostname()) {
return fmt.Errorf("reputation.upstream.url must use https for non-loopback hosts (bearer token would otherwise leak in plaintext)")
}
return nil
}
// isLoopbackHost keeps plain HTTP limited to same-host panel deployments.
func isLoopbackHost(host string) bool {
if host == "" {
return false
}
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
if ip == nil {
return false
}
return ip.IsLoopback()
}
func validateVerdictCallback(cfg *Config) error {
_, err := validateVerdictCallbackField(cfg)
return err
}
func validateVerdictCallbackField(cfg *Config) (string, error) {
vc := cfg.AutoResponse.VerdictCallback
if vc.TimeoutSec != 0 && (vc.TimeoutSec < 1 || vc.TimeoutSec > 30) {
return "auto_response.verdict_callback.timeout_sec", fmt.Errorf("auto_response.verdict_callback.timeout_sec must be between 1 and 30")
}
if !vc.Enabled {
return "", nil
}
rawURL := strings.TrimSpace(vc.URL)
if rawURL == "" {
return "auto_response.verdict_callback.url", fmt.Errorf("auto_response.verdict_callback.enabled=true but url is empty")
}
parsed, err := url.Parse(rawURL)
if err != nil {
return "auto_response.verdict_callback.url", fmt.Errorf("auto_response.verdict_callback.url: %w", err)
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return "auto_response.verdict_callback.url", fmt.Errorf("auto_response.verdict_callback.url must use http or https")
}
if parsed.Host == "" {
return "auto_response.verdict_callback.url", fmt.Errorf("auto_response.verdict_callback.url must include host")
}
if err := validateVerdictCallbackSecret(verdictCallbackForValidation{
HMACSecret: vc.HMACSecret,
HMACSecretEnv: vc.HMACSecretEnv,
AllowUnsigned: vc.AllowUnsigned,
}); err != nil {
return verdictCallbackSecretField(vc.HMACSecretEnv), err
}
return "", nil
}
// validateVerdictCallbackSecret enforces fail-closed posture on the
// outbound HMAC: when the callback is enabled, either hmac_secret or the
// hmac_secret_env-named env var must resolve to a non-empty value, OR
// the operator must explicitly set allow_unsigned: true to acknowledge
// that requests and responses will run without integrity protection.
//
// Without this check a misconfigured deployment (env var typoed, secret
// not yet rotated in) silently emits unsigned POSTs while the daemon
// keeps reporting healthy, and any on-path actor can forge or replay
// block decisions.
func validateVerdictCallbackSecret(vc verdictCallbackForValidation) error {
if vc.AllowUnsigned {
return nil
}
if strings.TrimSpace(vc.HMACSecret) != "" {
return nil
}
if vc.HMACSecretEnv != "" {
if strings.TrimSpace(os.Getenv(vc.HMACSecretEnv)) != "" {
return nil
}
return fmt.Errorf("auto_response.verdict_callback.enabled=true but env var %q is empty or unset; set the secret, or opt in with allow_unsigned: true", vc.HMACSecretEnv)
}
return fmt.Errorf("auto_response.verdict_callback.enabled=true requires hmac_secret or hmac_secret_env (or allow_unsigned: true to acknowledge unsigned requests and responses)")
}
// verdictCallbackForValidation isolates the fields validateVerdictCallbackSecret
// needs without re-spelling the anonymous struct literal in config.go.
type verdictCallbackForValidation struct {
HMACSecret string
HMACSecretEnv string
AllowUnsigned bool
}
func verdictCallbackSecretField(hmacSecretEnv string) string {
if hmacSecretEnv != "" {
return "auto_response.verdict_callback.hmac_secret_env"
}
return "auto_response.verdict_callback.hmac_secret"
}
func validateDirectSMTPEgress(cfg *Config) error {
d := cfg.Detection.DirectSMTPEgress
switch strings.ToLower(strings.TrimSpace(d.Backend)) {
case "", "auto", "bpf", "legacy", "none":
default:
return fmt.Errorf("detection.direct_smtp_egress.backend must be auto, bpf, legacy, or none")
}
for i, p := range d.Ports {
if p < 1 || p > 65535 {
return fmt.Errorf("detection.direct_smtp_egress.ports[%d] must be between 1 and 65535", i)
}
}
return nil
}
func validateBPFEnforcement(cfg *Config) error {
if !cfg.BPFEnforcement.Enabled {
return nil
}
switch strings.ToLower(strings.TrimSpace(cfg.Detection.ConnectionTrackerBackend)) {
case "", "auto", "bpf":
case "legacy", "none":
return fmt.Errorf("bpf_enforcement.enabled=true requires detection.connection_tracker_backend=auto or bpf")
default:
return fmt.Errorf("detection.connection_tracker_backend must be auto, bpf, legacy, or none")
}
gates := 0
if cfg.BPFEnforcement.DirectSMTPEgress {
if !cfg.Detection.DirectSMTPEgress.Enabled {
return fmt.Errorf("bpf_enforcement.direct_smtp_egress requires detection.direct_smtp_egress.enabled=true")
}
switch strings.ToLower(strings.TrimSpace(cfg.Detection.DirectSMTPEgress.Backend)) {
case "", "auto", "bpf":
case "legacy", "none":
return fmt.Errorf("bpf_enforcement.direct_smtp_egress requires detection.direct_smtp_egress.backend=auto or bpf")
default:
return fmt.Errorf("detection.direct_smtp_egress.backend must be auto, bpf, legacy, or none")
}
gates++
}
if gates == 0 {
return fmt.Errorf("bpf_enforcement.enabled=true requires at least one feature gate (direct_smtp_egress)")
}
return nil
}
func validateWebUITokens(cfg *Config) error {
if _, _, err := cfg.BrowserSessionDurations(); err != nil {
return err
}
seenNames := make(map[string]struct{}, len(cfg.WebUI.Tokens))
seenTokens := make(map[string]struct{}, len(cfg.WebUI.Tokens))
for i, tok := range cfg.WebUI.Tokens {
name := strings.TrimSpace(tok.Name)
if name == "" {
return fmt.Errorf("webui.tokens[%d]: empty name", i)
}
if tok.Scope != "admin" && tok.Scope != "read" {
return fmt.Errorf("webui.tokens[%d]: unknown scope %q (use admin or read)", i, tok.Scope)
}
if tok.Token == "" {
return fmt.Errorf("webui.tokens[%d]: empty token", i)
}
if _, ok := seenNames[name]; ok {
return fmt.Errorf("webui.tokens[%d]: duplicate name %q", i, tok.Name)
}
seenNames[name] = struct{}{}
if _, ok := seenTokens[tok.Token]; ok {
return fmt.Errorf("webui.tokens[%d]: duplicate token", i)
}
seenTokens[tok.Token] = struct{}{}
}
return nil
}
func validateMailLogsField(cfg *Config) (string, error) {
switch cfg.MailLogs.Source {
case "", "auto", "file", "journal":
default:
return "mail_logs.source", fmt.Errorf("mail_logs.source: must be auto, file, or journal (got %q)", cfg.MailLogs.Source)
}
for i, unit := range cfg.MailLogs.Units {
if strings.TrimSpace(unit) == "" {
return "mail_logs.units", fmt.Errorf("mail_logs.units[%d]: empty unit", i)
}
}
return "", nil
}
func validateMailLogs(cfg *Config) error {
_, err := validateMailLogsField(cfg)
return err
}
func validateMailBruteAccountKeyField(cfg *Config) (string, error) {
key := cfg.Thresholds.MailBruteAccountKey
switch {
case key == "", key == "builtin:dovecot-user", key == "builtin:postfix-sasl":
// ok
case strings.HasPrefix(key, "regex:"):
re, err := regexp.Compile(strings.TrimPrefix(key, "regex:"))
if err != nil {
return "thresholds.mail_brute_account_key", fmt.Errorf("mail_brute_account_key: invalid regex: %w", err)
}
if re.NumSubexp() < 1 {
return "thresholds.mail_brute_account_key", fmt.Errorf("mail_brute_account_key: regex must contain at least one capture group")
}
default:
return "thresholds.mail_brute_account_key", fmt.Errorf("mail_brute_account_key: %q must be builtin:* or regex:*", key)
}
return "", nil
}
func validateMailBruteAccountKey(cfg *Config) error {
_, err := validateMailBruteAccountKeyField(cfg)
return err
}
// validateHTTPASNCrawl checks http_asn_crawl thresholds and auto-response for
// values that would make the detector mis-behave at runtime. Called from
// LoadBytes after defaults are applied, so percent fields are already 1..100
// unless the operator explicitly set them to out-of-range values.
func validateHTTPASNCrawl(cfg *Config) error {
th := cfg.Thresholds
// Integer fields that must not be negative.
negChecks := []struct {
name string
val int
}{
{"thresholds.http_asn_crawl_window_min", th.HTTPASNCrawlWindowMin},
{"thresholds.http_asn_crawl_min_ips", th.HTTPASNCrawlMinIPs},
{"thresholds.http_asn_crawl_min_expensive", th.HTTPASNCrawlMinExpensive},
{"thresholds.http_asn_crawl_high_volume_mult", th.HTTPASNCrawlHighVolumeMult},
{"thresholds.http_asn_crawl_saturation", th.HTTPASNCrawlSaturation},
{"thresholds.http_asn_crawl_max_prefix", th.HTTPASNCrawlMaxPrefix},
{"thresholds.http_asn_crawl_max_tracked_ips", th.HTTPASNCrawlMaxTrackedIPs},
}
for _, c := range negChecks {
if c.val < 0 {
return fmt.Errorf("%s must not be negative (got %d)", c.name, c.val)
}
}
// Percent fields: after defaults are applied these are always >= 1.
// An operator-supplied value outside 1..100 is rejected.
pctChecks := []struct {
name string
val int
}{
{"thresholds.http_asn_crawl_min_share_pct", th.HTTPASNCrawlMinSharePct},
{"thresholds.http_asn_crawl_high_amp_pct", th.HTTPASNCrawlHighAmpPct},
{"thresholds.http_asn_crawl_16_pref_pct", th.HTTPASNCrawl16PrefPct},
}
for _, c := range pctChecks {
if c.val < 1 || c.val > 100 {
return fmt.Errorf("%s must be in 1..100 (got %d)", c.name, c.val)
}
}
if th.HTTPASNCrawlMinIPs > 0 && th.HTTPASNCrawlMaxTrackedIPs < th.HTTPASNCrawlMinIPs {
return fmt.Errorf("thresholds.http_asn_crawl_max_tracked_ips must be >= thresholds.http_asn_crawl_min_ips when http_asn_crawl is enabled")
}
// Every ASN in both lists must be in 1..4294967295.
asnLists := []struct {
name string
list []uint
}{
{"thresholds.http_asn_crawl_allowlist_asns", th.HTTPASNCrawlAllowlistASNs},
{"thresholds.http_asn_crawl_reverse_proxy_asns", th.HTTPASNCrawlReverseProxyASNs},
}
for _, al := range asnLists {
for _, asn := range al.list {
if asn < 1 || asn > 4294967295 {
return fmt.Errorf("%s contains invalid ASN %d (must be 1..4294967295)", al.name, asn)
}
}
}
// Tempban must be a positive duration.
if tb := cfg.AutoResponse.HTTPASNCrawlTempban; tb != "" {
d, err := time.ParseDuration(tb)
if err != nil {
return fmt.Errorf("auto_response.http_asn_crawl_tempban: unparseable duration %q: %w", tb, err)
}
if d <= 0 {
return fmt.Errorf("auto_response.http_asn_crawl_tempban: duration must be positive (got %q)", tb)
}
}
return nil
}
func Load(path string) (*Config, error) {
// #nosec G304 -- path is operator-supplied config file (CLI flag / env).
f, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("reading config %s: %w", path, err)
}
defer f.Close()
data, err := readConfigBytesLimited(f)
if errors.Is(err, errConfigTooLarge) {
return nil, fmt.Errorf("config %s exceeds %d byte cap", path, MaxConfigBytes)
}
if err != nil {
return nil, fmt.Errorf("reading config %s: %w", path, err)
}
cfg, err := LoadBytes(data)
if err != nil {
return nil, err
}
cfg.ConfigFile = path
return cfg, nil
}
// LoadWithDir loads the main config file and then merges every YAML fragment
// from confDir on top in lexicographic order. A missing confDir is not an
// error. Unknown fields in fragments are rejected (KnownFields=true).
func LoadWithDir(path, confDir string) (*Config, error) {
// #nosec G304 -- path is operator-supplied (CLI flag).
f, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("reading config %s: %w", path, err)
}
defer f.Close()
mainData, err := readConfigBytesLimited(f)
if errors.Is(err, errConfigTooLarge) {
return nil, fmt.Errorf("config %s exceeds %d byte cap", path, MaxConfigBytes)
}
if err != nil {
return nil, fmt.Errorf("reading config %s: %w", path, err)
}
cfg, err := loadBytesWithDir(mainData, confDir, path)
if err != nil {
return nil, err
}
cfg.ConfigFile = path
cfg.ConfigDir = confDir
return cfg, nil
}
// LoadBytesWithDir decodes an in-memory main config after merging the current
// confDir fragments on top. It is used to validate an edited main file before
// that file is committed to disk. ConfigFile is left empty.
func LoadBytesWithDir(data []byte, confDir string) (*Config, error) {
mergedBytes, err := MergeBytesWithDir(data, confDir)
if err != nil {
return nil, err
}
cfg, err := LoadBytes(mergedBytes)
if err != nil {
return nil, err
}
cfg.ConfigDir = confDir
return cfg, nil
}
// MergeBytesWithDir merges an in-memory main config with the current confDir
// fragments and returns YAML that still preserves whether fields were omitted.
func MergeBytesWithDir(data []byte, confDir string) ([]byte, error) {
if len(data) > MaxConfigBytes {
return nil, fmt.Errorf("parsing config: input size %d exceeds %d byte cap", len(data), MaxConfigBytes)
}
return mergeBytesWithDir(data, confDir, "config input")
}
func loadBytesWithDir(mainData []byte, confDir, source string) (*Config, error) {
mergedBytes, err := mergeBytesWithDir(mainData, confDir, source)
if err != nil {
return nil, err
}
return LoadBytes(mergedBytes)
}
func mergeBytesWithDir(mainData []byte, confDir, source string) ([]byte, error) {
var merged yaml.Node
if unmarshalErr := yaml.Unmarshal(mainData, &merged); unmarshalErr != nil {
return nil, fmt.Errorf("parsing %s: %w", source, unmarshalErr)
}
normalized, err := normalizeYAMLForMerge(&merged)
if err != nil {
return nil, fmt.Errorf("parsing %s: %w", source, err)
}
merged = *normalized
frags, err := loadConfDirFragments(confDir)
if err != nil {
return nil, err
}
for _, frag := range frags {
DeepMergeTracked(&merged, frag.node, func(keyPath, oldVal, newVal string) {
fmt.Fprintf(os.Stderr, "confd: %s overrides %s: %q -> %q\n",
frag.path,
keyPath,
redactConfigScalarForLog(keyPath, oldVal),
redactConfigScalarForLog(keyPath, newVal),
)
})
}
mergedBytes, err := yaml.Marshal(&merged)
if err != nil {
return nil, fmt.Errorf("marshaling merged config: %w", err)
}
return mergedBytes, nil
}
func Save(cfg *Config) error {
data, err := yaml.Marshal(cfg)
if err != nil {
return fmt.Errorf("marshaling config: %w", err)
}
path, err := saveTargetPath(cfg.ConfigFile)
if err != nil {
return err
}
// csm.yaml is the daemon's only config: a truncate-in-place write torn
// by a crash would block the next daemon start with a parse error.
return atomicio.AtomicWrite(path, 0600, data)
}
func saveTargetPath(path string) (string, error) {
info, err := os.Lstat(path)
if err != nil {
if os.IsNotExist(err) {
return path, nil
}
return "", fmt.Errorf("checking config path: %w", err)
}
if info.Mode()&os.ModeSymlink == 0 {
return path, nil
}
// FHS migration leaves the legacy config path as a symlink to the
// main config. Save must update the target instead of replacing the link.
resolved, err := filepath.EvalSymlinks(path)
if err != nil {
return "", fmt.Errorf("resolving config symlink: %w", err)
}
return resolved, nil
}
package config
import (
"fmt"
"net"
"net/netip"
"net/url"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/firewall"
)
// outboundDependency is one TCP destination the daemon dials because the
// operator enabled the feature that needs it. what reads as a clause the
// warning can quote, e.g. "alerts.webhook.url dials panel.example.com on
// TCP port 8443".
type outboundDependency struct {
what string
host string // hostname or literal IP; empty when the endpoints are vendor-fixed
port int
daemonRoot bool // the daemon's root UID can use smtp_block's per-UID accepts
}
// builtInHTTPSEndpoints names the outbound HTTPS destinations that are not
// configurable: threat feeds, AbuseIPDB, MaxMind, YARA Forge, AI-crawler
// range feeds and the release check. They share one warning because dropping
// 443 from tcp_out silences all of them at once.
const builtInHTTPSEndpoints = "built-in HTTPS endpoints (threat feeds, signature, GeoIP and bot-range updates, AbuseIPDB, release check)"
// outboundDependencies lists every TCP destination the daemon will dial with
// the features currently enabled. Loopback destinations are left out because
// the output chain accepts loopback ahead of any port rule.
func outboundDependencies(cfg *Config) []outboundDependency {
deps := []outboundDependency{{
what: builtInHTTPSEndpoints + " need TCP port 443 outbound",
port: 443,
daemonRoot: true,
}}
add := func(label, host string, port int) {
if egressLoopbackHost(host) {
return
}
deps = append(deps, outboundDependency{
what: fmt.Sprintf("%s dials %s on TCP port %d", label, host, port),
host: host,
port: port,
daemonRoot: true,
})
}
addURL := func(label, raw string) {
if host, port, ok := urlDialTarget(raw); ok {
add(label, host, port)
}
}
addHostPort := func(label, raw string) {
if host, port, ok := hostPortDialTarget(raw); ok {
add(label, host, port)
}
}
if cfg.Alerts.Email.Enabled && cfg.Alerts.Email.SMTP != "" {
addHostPort("alerts.email.smtp", cfg.Alerts.Email.SMTP)
}
if cfg.Alerts.Webhook.Enabled && cfg.Alerts.Webhook.URL != "" {
addURL("alerts.webhook.url", cfg.Alerts.Webhook.URL)
}
if cfg.Alerts.Heartbeat.Enabled && cfg.Alerts.Heartbeat.URL != "" {
addURL("alerts.heartbeat.url", cfg.Alerts.Heartbeat.URL)
}
syslog := cfg.Alerts.AuditLog.Syslog
if syslog.Enabled && (syslog.Network == "tcp" || syslog.Network == "tls") {
addHostPort("alerts.audit_log.syslog.address", syslog.Address)
}
if cfg.AutoResponse.VerdictCallback.Enabled && cfg.AutoResponse.VerdictCallback.URL != "" {
addURL("auto_response.verdict_callback.url", cfg.AutoResponse.VerdictCallback.URL)
}
if cfg.Reputation.Rspamd.Enabled && cfg.Reputation.Rspamd.URL != "" {
addURL("reputation.rspamd.url", cfg.Reputation.Rspamd.URL)
}
if cfg.Reputation.Upstream.Enabled && cfg.Reputation.Upstream.URL != "" {
addURL("reputation.upstream.url", cfg.Reputation.Upstream.URL)
}
if cfg.Reputation.Report.Enabled {
for i, target := range cfg.Reputation.Report.Targets {
addURL(fmt.Sprintf("reputation.report.targets[%d].url", i), target.URL)
}
}
if cfg.Reputation.Central.Enabled && cfg.Reputation.Central.SetURL != "" {
addURL("reputation.central.set_url", cfg.Reputation.Central.SetURL)
}
// The signature updater runs whenever an update URL is set.
if cfg.Signatures.UpdateURL != "" {
addURL("signatures.update_url", cfg.Signatures.UpdateURL)
}
if cfg.Signatures.YaraForge.Enabled && cfg.Signatures.YaraForge.DownloadURL != "" {
addURL("signatures.yara_forge.download_url", cfg.Signatures.YaraForge.DownloadURL)
}
if cfg.Sentry.Enabled && cfg.Sentry.DSN != "" {
addURL("sentry.dsn", cfg.Sentry.DSN)
}
if cfg.UpdatesCheckEnabled() && cfg.Updates.GitHubAPIURL != "" {
addURL("updates.github_api_url", cfg.Updates.GitHubAPIURL)
}
return deps
}
// urlDialTarget resolves the host and port a URL is dialed on the way the
// HTTP client does: an explicit port wins, otherwise the scheme default. A
// URL that the client cannot dial produces no egress requirement.
func urlDialTarget(raw string) (host string, port int, ok bool) {
u, err := url.Parse(raw)
if err != nil || u.Host == "" {
return "", 0, false
}
host = u.Hostname()
if host == "" {
return "", 0, false
}
scheme := strings.ToLower(u.Scheme)
if scheme != "http" && scheme != "https" {
return "", 0, false
}
if p := u.Port(); p != "" {
port, err = strconv.Atoi(p)
if err != nil || port < 1 || port > 65535 {
return "", 0, false
}
return host, port, true
}
switch scheme {
case "https":
return host, 443, true
case "http":
return host, 80, true
}
return "", 0, false
}
// hostPortDialTarget splits a host:port dial address. Addresses without a
// port cannot be dialed and produce no egress requirement.
func hostPortDialTarget(raw string) (host string, port int, ok bool) {
if raw != strings.TrimSpace(raw) {
return "", 0, false
}
host, portStr, err := net.SplitHostPort(raw)
if err != nil || host == "" {
return "", 0, false
}
// net.Dial accepts both numeric ports and service names. Resolve through
// the same service database so a valid address such as host:smtp cannot
// evade the egress-policy warning.
port, err = net.LookupPort("tcp", portStr)
if err != nil || port < 1 || port > 65535 {
return "", 0, false
}
return host, port, true
}
// firewallEgressResults warns when the enabled firewall's outbound policy
// would refuse a connection the daemon itself needs. The output chain is
// default-drop and ends in a TCP reset, so the failure reads as the far end
// being down while the host looks healthy locally. Like the inbound checks
// these are warnings, never errors: an operator may route through a proxy
// that validation cannot see.
//
// Coverage is limited to what the daemon dials from its own config plus the
// ports declared under firewall.required_tcp_out. A third-party agent's
// egress is only checked when it declares its ports there.
func firewallEgressResults(cfg *Config) []ValidationResult {
fw := cfg.Firewall
if fw == nil || !fw.Enabled {
return nil
}
// Mirror the engine: the output chain is accept-all unless some outbound
// list is set; IPv4 is accepted wholesale when only IPv6 lists are set;
// IPv6 is accepted wholesale unless it is managed.
tcp6Out := fw.TCP6Out
if len(tcp6Out) == 0 {
tcp6Out = fw.TCPOut
}
udp6Out := fw.UDP6Out
if len(udp6Out) == 0 {
udp6Out = fw.UDPOut
}
ipv4Filtered := len(fw.TCPOut) > 0 || len(fw.UDPOut) > 0
ipv6Filtered := fw.IPv6 && (ipv4Filtered || len(tcp6Out) > 0 || len(udp6Out) > 0)
if !ipv4Filtered && !ipv6Filtered {
return nil
}
// smtp_block installs per-UID accepts and then a drop for each mail port
// ahead of both the IPv4 bypass and the configured port rules. The daemon
// runs as root, which is always accepted; required_tcp_out describes a
// separate service whose UID this validator cannot prove is allowed.
smtpRestricted := func(port int) bool {
return fw.SMTPBlock && containsPort(fw.SMTPPorts, port)
}
blocked4 := func(dep outboundDependency) bool {
if smtpRestricted(dep.port) {
return !dep.daemonRoot
}
if outAllowReachesAnyDst(fw, dep.port, false) {
return false
}
return ipv4Filtered && !containsPort(fw.TCPOut, dep.port)
}
blocked6 := func(dep outboundDependency) bool {
if smtpRestricted(dep.port) {
return !dep.daemonRoot
}
if outAllowReachesAnyDst(fw, dep.port, true) {
return false
}
return !containsPort(tcp6Out, dep.port)
}
deps := outboundDependencies(cfg)
for _, port := range fw.RequiredTCPOut {
if port < 1 || port > 65535 {
continue // reported as an error by firewallValueResults
}
deps = append(deps, outboundDependency{
what: fmt.Sprintf("firewall.required_tcp_out declares TCP port %d", port),
port: port,
})
}
var results []ValidationResult
for _, dep := range deps {
wantV4, wantV6 := dialFamilies(dep.host)
isBlocked4 := wantV4 && blocked4(dep)
if isBlocked4 {
results = append(results, ValidationResult{"warn", "firewall.tcp_out",
egressBlockedMessage(dep, "tcp_out", smtpRestricted(dep.port)) +
outAllowScopedNote(fw, dep.port, false)})
}
// An inherited tcp6_out is the same list as tcp_out, so the IPv4
// warning above covers both families only when IPv4 is filtered too.
// With only udp6_out configured, the engine bypasses IPv4 wholesale
// but still drops IPv6 TCP, so that shape needs a tcp6_out warning.
isBlocked6 := wantV6 && ipv6Filtered && blocked6(dep)
if isBlocked6 && (len(fw.TCP6Out) > 0 || !isBlocked4) {
results = append(results, ValidationResult{"warn", "firewall.tcp6_out",
egressBlockedMessage(dep, "tcp6_out", smtpRestricted(dep.port)) +
outAllowScopedNote(fw, dep.port, true)})
}
}
return results
}
// outAllowReachesAnyDst reports whether a tcp_out_allow rule opens the port to
// every destination in the requested family. Only then can it prove the endpoint is
// reachable, because it cannot resolve a hostname to test a scoped prefix.
func outAllowReachesAnyDst(fw *firewall.FirewallConfig, port int, ipv6 bool) bool {
for _, r := range fw.TCPOutAllow {
network := outAllowNetworkForPort(fw, r, port, ipv6)
if network == nil {
continue
}
if ones, _ := network.Mask.Size(); ones == 0 {
return true
}
}
return false
}
// outAllowNetworkForPort mirrors the engine's emission checks and SMTP order
// so a skipped rule cannot hide a lockout or claim to cover a blocked port.
func outAllowNetworkForPort(fw *firewall.FirewallConfig, r firewall.OutAllowRule, port int, ipv6 bool) *net.IPNet {
if !validPort(r.PortStart) || !validPort(r.PortEnd) || port < r.PortStart || port > r.PortEnd {
return nil
}
if ipv6 && !fw.IPv6 || fw.SMTPBlock && containsPort(fw.SMTPPorts, port) {
return nil
}
network, err := firewall.ParseOutAllowDst(r.Dst)
if err != nil || (network.IP.To4() == nil) != ipv6 {
return nil
}
return network
}
// outAllowScopedNote annotates a lockout warning when a scoped tcp_out_allow
// rule covers the port. The warning stands -- silence here would be the silent
// lockout these checks exist to catch -- but the operator is told where to look.
func outAllowScopedNote(fw *firewall.FirewallConfig, port int, ipv6 bool) string {
var dsts []string
for _, r := range fw.TCPOutAllow {
network := outAllowNetworkForPort(fw, r, port, ipv6)
if network == nil {
continue
}
if ones, _ := network.Mask.Size(); ones == 0 {
continue
}
dsts = append(dsts, r.Dst)
}
if len(dsts) == 0 {
return ""
}
return fmt.Sprintf("; tcp_out_allow covers the port for %s only, so confirm the endpoint resolves into that range",
strings.Join(dsts, ", "))
}
func egressBlockedMessage(dep outboundDependency, policy string, smtpRestricted bool) string {
if smtpRestricted {
return fmt.Sprintf("%s but smtp_block restricts TCP port %d to allowed UIDs; required_tcp_out does not identify an allowed user", dep.what, dep.port)
}
if policy == "tcp6_out" {
return fmt.Sprintf("IPv6 is managed and tcp6_out does not allow it: %s", dep.what)
}
return fmt.Sprintf("%s but tcp_out does not allow it; once the firewall applies, connections to that port are refused", dep.what)
}
// dialFamilies reports which IP families a destination can be dialed over. A
// literal address pins one family; a hostname (or the vendor endpoints, which
// have no single host) may resolve to either.
func dialFamilies(host string) (ipv4, ipv6 bool) {
ip, err := netip.ParseAddr(host)
if err != nil {
return true, true
}
if ip.Unmap().Is4() {
return true, false
}
return false, true
}
// egressLoopbackHost recognizes scoped IPv6 literals as well as the common
// forms handled by isLoopbackHost. URL.Hostname and net.SplitHostPort retain
// a zone such as "%lo", which net.ParseIP cannot parse even though the
// kernel still routes ::1 over the output chain's accepted loopback device.
func egressLoopbackHost(host string) bool {
if isLoopbackHost(host) {
return true
}
if strings.EqualFold(strings.TrimSuffix(host, "."), "localhost") {
return true
}
ip, err := netip.ParseAddr(host)
return err == nil && ip.Unmap().IsLoopback()
}
package config
import "time"
// EmailAVConfig holds email antivirus scanning settings.
type EmailAVConfig struct {
Enabled bool `yaml:"enabled"`
ClamdSocket string `yaml:"clamd_socket"`
ScanTimeout string `yaml:"scan_timeout"`
MaxAttachmentSize int64 `yaml:"max_attachment_size"`
MaxArchiveDepth int `yaml:"max_archive_depth"`
MaxArchiveFiles int `yaml:"max_archive_files"`
MaxExtractionSize int64 `yaml:"max_extraction_size"`
QuarantineInfected bool `yaml:"quarantine_infected"`
ScanConcurrency int `yaml:"scan_concurrency"`
// FailMode controls behavior when scanning cannot complete.
// "open" (default): deliver mail when engines are down, scans time out, or quarantine fails.
// "tempfail": defer delivery (Exim retries later) when scanning cannot complete or infected mail cannot be quarantined.
FailMode string `yaml:"fail_mode"`
}
// ScanTimeoutDuration parses the ScanTimeout string as a time.Duration.
func (c *EmailAVConfig) ScanTimeoutDuration() time.Duration {
if c.ScanTimeout == "" {
return 30 * time.Second
}
d, err := time.ParseDuration(c.ScanTimeout)
if err != nil {
return 30 * time.Second
}
return d
}
// DefaultClamdSocket is the socket assumed when the operator sets none. It is
// RHEL's clamd-scan path; other packagings put it elsewhere, which is what
// ResolveClamdSocket exists to cope with.
const DefaultClamdSocket = "/var/run/clamd.scan/clamd.sock"
// EmailAVDefaults applies default values to an EmailAVConfig.
func EmailAVDefaults(c *EmailAVConfig) {
if c.ClamdSocket == "" {
c.ClamdSocket = DefaultClamdSocket
}
if c.ScanTimeout == "" {
c.ScanTimeout = "30s"
}
if c.MaxAttachmentSize == 0 {
c.MaxAttachmentSize = 25 * 1024 * 1024 // 25 MB
}
if c.MaxArchiveDepth == 0 {
c.MaxArchiveDepth = 1
}
if c.MaxArchiveFiles == 0 {
c.MaxArchiveFiles = 50
}
if c.MaxExtractionSize == 0 {
c.MaxExtractionSize = 100 * 1024 * 1024 // 100 MB
}
if c.ScanConcurrency == 0 {
c.ScanConcurrency = 4
}
}
package config
import (
"net"
"strings"
"github.com/pidginhost/csm/internal/firewall"
)
// EffectiveFirewallConfig returns the copied configuration handed to the
// firewall engine. Keeping these derived values in one place prevents status
// surfaces from describing the persisted sections instead of the live policy.
func EffectiveFirewallConfig(cfg *Config) *firewall.FirewallConfig {
if cfg == nil {
return nil
}
effective := firewall.FirewallConfig{}
if cfg.Firewall != nil {
effective = *cfg.Firewall
}
effective.InfraIPs = firewall.MergeInfraIPs(cfg.InfraIPs, effective.InfraIPs)
if !cfg.Challenge.Enabled || !cfg.Challenge.PortGate.Enabled ||
cfg.Challenge.ListenPort <= 0 || challengeListenAddrIsLoopback(cfg.Challenge.ListenAddr) {
return &effective
}
effective.TCPIn = appendUniquePort(effective.TCPIn, cfg.Challenge.ListenPort)
effective.RestrictedTCP = removePort(effective.RestrictedTCP, cfg.Challenge.ListenPort)
return &effective
}
func challengeListenAddrIsLoopback(addr string) bool {
addr = strings.TrimSpace(addr)
if addr == "" {
return true
}
host := addr
if h, _, err := net.SplitHostPort(addr); err == nil {
host = h
}
host = strings.Trim(host, "[]")
if host == "" {
return false
}
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func appendUniquePort(ports []int, port int) []int {
for _, existing := range ports {
if existing == port {
return ports
}
}
out := append([]int(nil), ports...)
return append(out, port)
}
func removePort(ports []int, port int) []int {
out := make([]int, 0, len(ports))
for _, existing := range ports {
if existing != port {
out = append(out, existing)
}
}
return out
}
package config
import (
"reflect"
"sync/atomic"
)
// Hot-reload policy (ROADMAP item 7).
//
// Each top-level field of Config carries an optional `hotreload`
// struct tag:
//
// - "safe": a SIGHUP reload swaps the field in place; readers
// on the next tick see the new value.
// - "restart": the field cannot be applied without a full daemon
// restart (fanotify watched roots, bbolt path, web UI
// listener). A SIGHUP that touches a restart field
// logs a warning and leaves the prior config live.
// - (none): treated as restart-required. Tagging every safe
// field explicitly is the strict default; this closes
// the door on accidental hot-swaps of a fresh field
// someone adds without considering the safety of live
// mutation.
const (
TagSafe = "safe"
TagRestart = "restart"
)
// active holds the current live config. Readers on hot paths (check
// tick handlers, alert dispatchers, metrics-auth) call Active() to
// pick up the latest snapshot after a SIGHUP. Writers (daemon
// startup, SIGHUP reload) call SetActive.
var active atomic.Pointer[Config]
// Active returns the current live Config pointer. Returns nil if
// SetActive has not been called; callers on hot paths are expected
// to nil-check once per call. Daemon startup calls SetActive before
// any tick runs, so a nil return in production is a bug.
func Active() *Config {
return active.Load()
}
// SetActive installs cfg as the current live config.
func SetActive(cfg *Config) {
active.Store(cfg)
}
// Change describes a single top-level field that differs between an
// old and a new Config.
type Change struct {
// Field is the YAML name (from the `yaml:"..."` struct tag), or
// the Go field name if no yaml tag is set.
Field string
// Tag is the hotreload classification: TagSafe, TagRestart, or
// "" for fields with no explicit tag (treated as TagRestart).
Tag string
}
// ReloadPolicy is the top-level hot-reload manifest exposed to tests and
// operator surfaces that need to explain whether a Settings section can be
// applied live or waits for a restart.
type ReloadPolicy struct {
Field string
Tag string
RestartRequired bool
}
// HotReloadManifest returns every operator-owned top-level Config field with
// its effective reload policy. Fields without an explicit supported tag are
// classified restart-required, matching Diff and RestartRequired.
func HotReloadManifest() []ReloadPolicy {
cfgType := reflect.TypeOf(Config{})
policies := make([]ReloadPolicy, 0, cfgType.NumField())
for i := 0; i < cfgType.NumField(); i++ {
field := cfgType.Field(i)
if !field.IsExported() || isReloadManifestIgnoredRoot(field) {
continue
}
name := yamlFieldName(field)
if name == "" || name == "-" {
continue
}
tag := field.Tag.Get("hotreload")
if tag != TagSafe && tag != TagRestart {
tag = TagRestart
}
policies = append(policies, ReloadPolicy{
Field: name,
Tag: tag,
RestartRequired: tag != TagSafe,
})
}
return policies
}
func isReloadManifestIgnoredRoot(field reflect.StructField) bool {
return field.Name == "ConfigFile" || field.Name == "ConfigDir" || field.Name == "Integrity"
}
// Diff reports which Config fields differ between old and new,
// classified by hotreload tag.
//
// The walk is recursive: if a top-level field is tagged, its tag
// applies to any change inside. If a nested field has its own tag,
// that tag wins over the parent (field-level overrides let a single
// safe field sit inside an otherwise restart-required parent, which
// is how webui.metrics_token can hot-reload even though the rest of
// WebUI needs a restart).
//
// Each Change carries the YAML path from root (e.g. "thresholds" for
// the top-level struct, "webui.metrics_token" for a nested leaf).
// The tag is the nearest tagged ancestor on that path; if nothing on
// the path is tagged, the Change's Tag is "" and the caller should
// treat that as TagRestart.
//
// Granularity rule: if a tagged ancestor classifies the whole
// subtree uniformly (parent tag applies, no nested overrides on
// changed leaves), the Change is reported at the parent level. That
// keeps the common case ("I changed three thresholds") as one
// "thresholds" Change. When a subtree contains a differently-tagged
// leaf, that leaf is reported separately with its own tag, and the
// parent (minus that leaf) is reported with the inherited tag.
func Diff(oldCfg, newCfg *Config) []Change {
if oldCfg == nil || newCfg == nil {
return nil
}
oldV := reflect.ValueOf(*oldCfg)
newV := reflect.ValueOf(*newCfg)
return diffStruct(oldV, newV, "", "")
}
// diffStruct walks two reflect.Values of the same struct type and
// returns Changes for every differing field. parentPath is the
// already-composed YAML dotted path down to this struct (empty at
// the root). parentTag is the effective hotreload tag inherited
// from the nearest tagged ancestor.
func diffStruct(oldV, newV reflect.Value, parentPath, parentTag string) []Change {
var changes []Change
t := oldV.Type()
for i := 0; i < t.NumField(); i++ {
field := t.Field(i)
if !field.IsExported() {
continue
}
// ConfigFile / ConfigDir / Integrity are daemon-managed process
// metadata, not operator policy fields.
if parentPath == "" && isReloadManifestIgnoredRoot(field) {
continue
}
oldField := oldV.Field(i).Interface()
newField := newV.Field(i).Interface()
if reflect.DeepEqual(oldField, newField) {
continue
}
// Effective tag for this field: its own explicit tag wins;
// otherwise inherit from the parent path.
tag := field.Tag.Get("hotreload")
if tag != TagSafe && tag != TagRestart {
tag = parentTag
}
name := yamlFieldName(field)
path := name
if parentPath != "" {
path = parentPath + "." + name
}
// If the field is itself a struct (not a pointer, slice,
// map), recurse so nested overrides can surface separately.
// Pointer-to-struct is treated as a leaf because the
// reflect.DeepEqual already told us the pointer target
// changed; re-walking it would produce duplicate noise.
if field.Type.Kind() == reflect.Struct {
nested := diffStruct(oldV.Field(i), newV.Field(i), path, tag)
// If every nested Change carries the same tag and there
// is no mixed classification, collapse to a single
// Change at this level. Operators rarely need the
// granularity "I changed thresholds.mail_queue_warn";
// the collapse keeps the common case clean.
if collapsed, ok := collapseIfUniform(nested, path, tag); ok {
changes = append(changes, collapsed)
} else {
changes = append(changes, nested...)
}
continue
}
changes = append(changes, Change{Field: path, Tag: tag})
}
return changes
}
// collapseIfUniform returns (Change{Field:path, Tag:parentTag}, true)
// when every nested change inherits parentTag (i.e. nothing nested
// overrode it). Returns (_, false) when the subtree contains a
// differently-tagged leaf, which means the caller must keep the
// granular changes.
func collapseIfUniform(nested []Change, path, parentTag string) (Change, bool) {
if len(nested) == 0 {
return Change{}, false
}
for _, c := range nested {
if c.Tag != parentTag {
return Change{}, false
}
}
return Change{Field: path, Tag: parentTag}, true
}
// RestartRequired returns true if any change in the diff carries a
// TagRestart classification (or no tag, which collapses to restart).
func RestartRequired(changes []Change) bool {
for _, c := range changes {
if c.Tag != TagSafe {
return true
}
}
return false
}
// yamlFieldName returns the yaml tag's primary name if set, else the
// Go field name. Strips any `,omitempty` / `,inline` suffix.
func yamlFieldName(f reflect.StructField) string {
tag := f.Tag.Get("yaml")
if tag == "" || tag == "-" {
return f.Name
}
for i := 0; i < len(tag); i++ {
if tag[i] == ',' {
return tag[:i]
}
}
return tag
}
package config
import (
"fmt"
"gopkg.in/yaml.v3"
)
// CollisionFn is invoked when DeepMergeTracked detects a scalar in the
// overlay overwriting a different scalar in the base. Identical-value
// rewrites are not reported because the operator cannot act on them.
// keyPath uses dotted YAML notation rooted at the document
// ("mail_logs.source"); top-level keys have no parent.
type CollisionFn func(keyPath, oldVal, newVal string)
// DeepMerge merges overlay into base in place and returns base.
// Both inputs must be DocumentNodes. Rules:
// - mapping ∩ mapping → key-by-key recurse
// - sequence ∩ sequence → append (base then overlay), with duplicate
// scalar entries removed from all-scalar lists
// - any other combination → overlay replaces base
//
// AliasNodes are treated as opaque scalars: an overlay alias replaces the
// base node; an alias inside base/overlay is not resolved before merging.
func DeepMerge(base, overlay *yaml.Node) *yaml.Node {
return DeepMergeTracked(base, overlay, nil)
}
// DeepMergeTracked is DeepMerge with an optional collision callback. The
// callback fires once per scalar-vs-scalar override across the document
// tree. Callers can pass nil for the previous DeepMerge behaviour.
func DeepMergeTracked(base, overlay *yaml.Node, onCollision CollisionFn) *yaml.Node {
if base == nil || overlay == nil {
return base
}
// An empty yaml.Unmarshal result has Kind==0; treat it as an empty document.
if base.Kind == 0 {
base.Kind = yaml.DocumentNode
}
if overlay.Kind == 0 {
overlay.Kind = yaml.DocumentNode
}
if base.Kind != yaml.DocumentNode || overlay.Kind != yaml.DocumentNode {
return base
}
if len(overlay.Content) == 0 {
return base
}
if len(base.Content) == 0 {
base.Content = overlay.Content
return base
}
mergeNodesAt(base.Content[0], overlay.Content[0], "", onCollision)
return base
}
const maxYAMLMergeExpansionNodes = 100_000
// normalizeYAMLForMerge resolves aliases and YAML merge keys before custom
// main+conf.d merging. Scalar nodes are preserved verbatim: decoding through
// interface{} would coerce date-like strings and other tagged values before
// Config gets its typed decode.
func normalizeYAMLForMerge(root *yaml.Node) (*yaml.Node, error) {
if root == nil || !containsAliasOrMerge(root) {
return root, nil
}
n := yamlMergeNormalizer{active: make(map[*yaml.Node]bool)}
// The initial scan already proved the tree needs normalization. Clone the
// whole tree in one pass so a deeply nested alias cannot make each ancestor
// rescan the same descendants.
return n.normalize(root)
}
func containsAliasOrMerge(node *yaml.Node) bool {
if node == nil {
return false
}
if node.Kind == yaml.AliasNode {
return true
}
if node.Kind == yaml.MappingNode {
for i := 0; i+1 < len(node.Content); i += 2 {
if isYAMLMergeKey(node.Content[i]) {
return true
}
}
}
for _, child := range node.Content {
if containsAliasOrMerge(child) {
return true
}
}
return false
}
type yamlMergeNormalizer struct {
active map[*yaml.Node]bool
created int
}
func (n *yamlMergeNormalizer) normalize(node *yaml.Node) (*yaml.Node, error) {
if node == nil {
return nil, nil
}
if node.Kind == yaml.AliasNode {
if node.Alias == nil {
return nil, fmt.Errorf("YAML alias %q has no anchor", node.Value)
}
if n.active[node.Alias] {
return nil, fmt.Errorf("YAML anchor %q contains itself", node.Value)
}
return n.normalize(node.Alias)
}
if n.active[node] {
return nil, fmt.Errorf("YAML anchor %q contains itself", node.Anchor)
}
if n.created >= maxYAMLMergeExpansionNodes {
return nil, fmt.Errorf("YAML alias expansion exceeds %d nodes", maxYAMLMergeExpansionNodes)
}
n.created++
n.active[node] = true
defer delete(n.active, node)
if node.Kind == yaml.MappingNode {
return n.normalizeMapping(node)
}
clone := *node
clone.Anchor = ""
clone.Alias = nil
clone.Content = make([]*yaml.Node, 0, len(node.Content))
for _, child := range node.Content {
normalized, err := n.normalize(child)
if err != nil {
return nil, err
}
clone.Content = append(clone.Content, normalized)
}
return &clone, nil
}
func (n *yamlMergeNormalizer) normalizeMapping(node *yaml.Node) (*yaml.Node, error) {
clone := *node
clone.Anchor = ""
clone.Alias = nil
clone.Content = make([]*yaml.Node, 0, len(node.Content))
explicit := make(map[string]bool, len(node.Content)/2)
var mergeValue *yaml.Node
for i := 0; i+1 < len(node.Content); i += 2 {
key := node.Content[i]
if isYAMLMergeKey(key) {
if mergeValue != nil {
return nil, fmt.Errorf("YAML mapping contains multiple merge keys")
}
mergeValue = node.Content[i+1]
continue
}
normalizedKey, err := n.normalize(key)
if err != nil {
return nil, err
}
normalizedValue, err := n.normalize(node.Content[i+1])
if err != nil {
return nil, err
}
clone.Content = append(clone.Content, normalizedKey, normalizedValue)
if id, ok := yamlScalarKeyID(normalizedKey); ok {
explicit[id] = true
}
}
if mergeValue == nil {
return &clone, nil
}
mergeMappings, err := n.normalizeMergeValue(mergeValue)
if err != nil {
return nil, err
}
for _, mapping := range mergeMappings {
if err := validateMergeMappingKeys(mapping); err != nil {
return nil, err
}
for i := 0; i+1 < len(mapping.Content); i += 2 {
id, _ := yamlScalarKeyID(mapping.Content[i])
if explicit[id] {
continue
}
explicit[id] = true
clone.Content = append(clone.Content, mapping.Content[i], mapping.Content[i+1])
}
}
return &clone, nil
}
func (n *yamlMergeNormalizer) normalizeMergeValue(node *yaml.Node) ([]*yaml.Node, error) {
normalized, err := n.normalize(node)
if err != nil {
return nil, err
}
switch normalized.Kind {
case yaml.MappingNode:
return []*yaml.Node{normalized}, nil
case yaml.SequenceNode:
mappings := make([]*yaml.Node, 0, len(normalized.Content))
for _, child := range normalized.Content {
if child.Kind != yaml.MappingNode {
return nil, fmt.Errorf("YAML map merge requires a map or sequence of maps")
}
mappings = append(mappings, child)
}
return mappings, nil
default:
return nil, fmt.Errorf("YAML map merge requires a map or sequence of maps")
}
}
func validateMergeMappingKeys(mapping *yaml.Node) error {
seen := make(map[string]bool, len(mapping.Content)/2)
for i := 0; i+1 < len(mapping.Content); i += 2 {
id, ok := yamlScalarKeyID(mapping.Content[i])
if !ok {
return fmt.Errorf("YAML map merge contains a non-scalar key")
}
if seen[id] {
return fmt.Errorf("YAML map merge contains duplicate key %q", mapping.Content[i].Value)
}
seen[id] = true
}
return nil
}
func yamlScalarKeyID(key *yaml.Node) (string, bool) {
if key == nil || key.Kind != yaml.ScalarNode {
return "", false
}
return key.ShortTag() + "\x00" + key.Value, true
}
func isYAMLMergeKey(key *yaml.Node) bool {
return key != nil && key.Kind == yaml.ScalarNode && key.Value == "<<" &&
(key.Tag == "" || key.Tag == "!" || key.ShortTag() == "!!merge")
}
func mergeNodesAt(b, o *yaml.Node, path string, onCollision CollisionFn) {
switch {
case b.Kind == yaml.MappingNode && o.Kind == yaml.MappingNode:
mergeMapAt(b, o, path, onCollision)
case b.Kind == yaml.SequenceNode && o.Kind == yaml.SequenceNode:
b.Content = dedupScalarSequence(append(b.Content, o.Content...))
default:
if onCollision != nil && b.Kind == yaml.ScalarNode && o.Kind == yaml.ScalarNode && b.Value != o.Value {
onCollision(path, b.Value, o.Value)
}
*b = *o
}
}
// dedupScalarSequence removes duplicate scalar entries (by value+tag),
// keeping the first occurrence and preserving order. It only acts when every
// element is a scalar: lists of maps (e.g. webui.tokens) keep every entry,
// where position and identity matter. Idempotent-by-content security lists
// (infra_ips, c2_blocklist, trusted_countries, disabled_checks) merged from a
// fragment that repeats a main-config entry would otherwise carry duplicates
// into validation and enforcement on every load.
func dedupScalarSequence(content []*yaml.Node) []*yaml.Node {
for _, n := range content {
if n.Kind != yaml.ScalarNode {
return content
}
}
seen := make(map[string]struct{}, len(content))
out := content[:0]
for _, n := range content {
key := n.Tag + "\x00" + n.Value
if _, dup := seen[key]; dup {
continue
}
seen[key] = struct{}{}
out = append(out, n)
}
return out
}
func mergeMapAt(b, o *yaml.Node, parent string, onCollision CollisionFn) {
for i := 0; i+1 < len(o.Content); i += 2 {
key := o.Content[i].Value
val := o.Content[i+1]
childPath := key
if parent != "" {
childPath = parent + "." + key
}
if idx := findKey(b, key); idx >= 0 {
mergeNodesAt(b.Content[idx+1], val, childPath, onCollision)
} else {
b.Content = append(b.Content, o.Content[i], val)
}
}
}
func findKey(m *yaml.Node, key string) int {
for i := 0; i+1 < len(m.Content); i += 2 {
if m.Content[i].Value == key {
return i
}
}
return -1
}
package config
import (
"fmt"
"strings"
)
// Operating modes for the Mode field.
const (
// ModeEnforce is the default: every subsystem is governed by its own
// switch and CSM keeps its host integration files current.
ModeEnforce = "enforce"
// ModeObserve declares that CSM must not change host state on this
// host. Detection, correlation, alerting and the audit sinks all run;
// automatic host remediation and integration updates do not. Two probes
// documented in the capability matrix still write: BPF capability
// discovery, and the kcarectl query that refreshes its own cache.
ModeObserve = "observe"
)
// ObserveMode reports whether this host runs in the observe posture.
func (cfg *Config) ObserveMode() bool {
return normalizeMode(cfg.Mode) == ModeObserve
}
func normalizeMode(mode string) string {
return strings.ToLower(strings.TrimSpace(mode))
}
// observeConflict names one config key whose value would let CSM change host
// state, and the value an observe-mode host has to use instead.
type observeConflict struct {
key string
want string
}
// observeConflicts lists switches for automatic changes to host files,
// processes or traffic. New state-changing subsystems belong here so
// validation can report every conflict in one pass.
func observeConflicts(cfg *Config) []observeConflict {
var out []observeConflict
add := func(on bool, key, want string) {
if on {
out = append(out, observeConflict{key: key, want: want})
}
}
add(cfg.AutoResponse.Enabled, "auto_response.enabled", "false")
add(cfg.AutoResponse.CopyFailKillProcess, "auto_response.copy_fail_kill_process", "false")
add(cfg.Firewall != nil && cfg.Firewall.Enabled, "firewall.enabled", "false")
add(cfg.PHPShield.Enabled, "php_shield.enabled", "false")
add(cfg.BPFEnforcement.Enabled, "bpf_enforcement.enabled", "false")
add(cfg.EmailProtection.ForwardGuard.Enabled, "email_protection.forward_guard.enabled", "false")
add(cfg.EmailAV.QuarantineInfected, "email_av.quarantine_infected", "false")
add(cfg.EmailAV.Enabled && cfg.EmailAV.FailMode == "tempfail", "email_av.fail_mode", "open")
add(cfg.AutoResponse.PHPRelay.Freeze != nil && *cfg.AutoResponse.PHPRelay.Freeze,
"auto_response.php_relay.freeze", "false")
add(cfg.AutoResponse.MailAuthRecovery.RestartEnabled,
"auto_response.mail_auth_recovery.restart_enabled", "false")
add(cfg.VirtualPatchMode() == VirtualPatchAuto,
"auto_response.virtual_patch_exposed_files", "off or manual")
return out
}
// validateMode rejects an unknown mode, and rejects an observe-mode config
// that still enables a subsystem which writes to the host. Every conflicting
// key is named in one error so an operator fixes the file in one pass.
func validateMode(cfg *Config) error {
switch normalizeMode(cfg.Mode) {
case ModeEnforce:
return nil
case ModeObserve:
default:
return fmt.Errorf("mode: %q is not a valid mode (use %q or %q)", cfg.Mode, ModeEnforce, ModeObserve)
}
conflicts := observeConflicts(cfg)
if len(conflicts) == 0 {
return nil
}
parts := make([]string, 0, len(conflicts))
for _, c := range conflicts {
parts = append(parts, fmt.Sprintf("%s (set %s)", c.key, c.want))
}
return fmt.Errorf("mode: observe forbids changing host state, but these keys still enable it: %s",
strings.Join(parts, ", "))
}
package config
import "slices"
const redactedValue = "***REDACTED***"
// RedactedValue is the placeholder Redact writes over a secret. A settings
// save that sends it back means "keep the stored secret".
const RedactedValue = redactedValue
var sensitiveScalarPaths = map[string]struct{}{
"alerts.webhook.hmac_secret": {},
"auto_response.verdict_callback.hmac_secret": {},
"challenge.captcha_fallback.secret_key": {},
"challenge.secret": {},
"challenge.verified_session.admin_secret": {},
"geoip.license_key": {},
"integrity.binary_hash": {},
"integrity.config_hash": {},
"integrity.confd_hash": {},
"reputation.abuseipdb_key": {},
"reputation.rspamd.token": {},
"reputation.upstream.token": {},
"sentry.dsn": {},
"webui.auth_token": {},
"webui.metrics_token": {},
}
// Redact returns a copy of the config with sensitive fields replaced.
// Empty fields are left empty (not replaced with the redaction marker).
// The original config is not modified.
func Redact(cfg *Config) *Config {
// Shallow copy the struct
c := *cfg
// Redact secrets (only if non-empty)
if c.WebUI.AuthToken != "" {
c.WebUI.AuthToken = redactedValue
}
if c.WebUI.MetricsToken != "" {
c.WebUI.MetricsToken = redactedValue
}
if len(c.WebUI.Tokens) > 0 {
c.WebUI.Tokens = append([]WebUIToken(nil), c.WebUI.Tokens...)
for i := range c.WebUI.Tokens {
if c.WebUI.Tokens[i].Token != "" {
c.WebUI.Tokens[i].Token = redactedValue
}
}
}
if c.Alerts.Webhook.HMACSecret != "" {
c.Alerts.Webhook.HMACSecret = redactedValue
}
if c.GeoIP.LicenseKey != "" {
c.GeoIP.LicenseKey = redactedValue
}
if c.Reputation.AbuseIPDBKey != "" {
c.Reputation.AbuseIPDBKey = redactedValue
}
if c.Reputation.Rspamd.Token != "" {
c.Reputation.Rspamd.Token = redactedValue
}
if c.Reputation.Upstream.Token != "" {
c.Reputation.Upstream.Token = redactedValue
}
if c.AutoResponse.VerdictCallback.HMACSecret != "" {
c.AutoResponse.VerdictCallback.HMACSecret = redactedValue
}
if c.Challenge.Secret != "" {
c.Challenge.Secret = redactedValue
}
if c.Challenge.CaptchaFallback.SecretKey != "" {
c.Challenge.CaptchaFallback.SecretKey = redactedValue
}
if c.Challenge.VerifiedSession.AdminSecret != "" {
c.Challenge.VerifiedSession.AdminSecret = redactedValue
}
if c.Integrity.BinaryHash != "" {
c.Integrity.BinaryHash = redactedValue
}
if c.Integrity.ConfigHash != "" {
c.Integrity.ConfigHash = redactedValue
}
if c.Integrity.ConfdHash != "" {
c.Integrity.ConfdHash = redactedValue
}
if c.Sentry.DSN != "" {
c.Sentry.DSN = redactedValue
}
// Credential-bearing URLs keep scheme and host only.
c.Alerts.Webhook.URL = RedactURL(c.Alerts.Webhook.URL)
c.Alerts.Heartbeat.URL = RedactURL(c.Alerts.Heartbeat.URL)
c.AutoResponse.VerdictCallback.URL = RedactURL(c.AutoResponse.VerdictCallback.URL)
c.Reputation.Rspamd.URL = RedactURL(c.Reputation.Rspamd.URL)
c.Reputation.Upstream.URL = RedactURL(c.Reputation.Upstream.URL)
if len(c.Reputation.Report.Targets) > 0 {
targets := slices.Clone(c.Reputation.Report.Targets)
for i := range targets {
targets[i].URL = RedactURL(targets[i].URL)
}
c.Reputation.Report.Targets = targets
}
// Deep-copy Firewall pointer so we don't share it with the original
if cfg.Firewall != nil {
fw := *cfg.Firewall
c.Firewall = &fw
}
return &c
}
func redactConfigScalarForLog(keyPath, value string) string {
if value == "" {
return value
}
if _, ok := sensitiveScalarPaths[keyPath]; ok {
return redactedValue
}
if _, ok := urlScalarPaths[keyPath]; ok {
return RedactURL(value)
}
return value
}
package config
import "net/url"
const redactedURLTail = "[REDACTED]"
// RedactURL keeps only the scheme and host of a URL for display. Webhook,
// heartbeat and callback URLs carry their credential in the path, query or
// userinfo (Slack and Discord webhooks, healthcheck pings, token query
// parameters), so everything past the host is replaced. A bare
// scheme://host[:port][/] is returned unchanged; an unparseable value is
// hidden entirely.
func RedactURL(raw string) string {
if raw == "" {
return raw
}
u, err := url.Parse(raw)
if err != nil || u.Scheme == "" || u.Host == "" {
return redactedURLTail
}
bare := u.User == nil && (u.Path == "" || u.Path == "/") && u.RawQuery == "" && u.Fragment == "" && u.RawPath == ""
if bare {
return raw
}
return u.Scheme + "://" + u.Host + "/" + redactedURLTail
}
// urlScalarPaths are config keys whose value is a credential-bearing URL;
// hot-reload diff logging shows them through RedactURL.
var urlScalarPaths = map[string]struct{}{
"alerts.webhook.url": {},
"alerts.heartbeat.url": {},
"auto_response.verdict_callback.url": {},
"reputation.rspamd.url": {},
"reputation.upstream.url": {},
}
package config
import (
"reflect"
"strings"
"time"
)
// Schema returns a JSON Schema (draft-07-style, partial) describing the
// Config struct via reflection over yaml: tags. Used by phpanel's config
// editor for client-side validation. Not a complete spec implementation -
// covers YAML field names, scalar/container types, and nested objects.
//
// IMPORTANT: This schema is structural only. Imperative validation rules
// enforced by Validate() (e.g., mail_logs.source must be auto/file/journal,
// webui.tokens[].scope must be admin or read) are NOT encoded here.
// Phpanel must still call `csm validate` for the authoritative check.
func Schema() map[string]interface{} {
return reflectStruct(reflect.TypeOf(Config{}))
}
func reflectStruct(t reflect.Type) map[string]interface{} {
if t.Kind() == reflect.Pointer {
t = t.Elem()
}
if t.Kind() != reflect.Struct {
return map[string]interface{}{"type": "object"}
}
props := map[string]interface{}{}
for i := 0; i < t.NumField(); i++ {
f := t.Field(i)
tag := f.Tag.Get("yaml")
if tag == "" || tag == "-" {
continue
}
name, _ := splitYAMLTag(tag)
if name == "" {
continue
}
props[name] = reflectField(f.Type)
}
return map[string]interface{}{
"type": "object",
"properties": props,
}
}
func reflectField(t reflect.Type) map[string]interface{} {
for t.Kind() == reflect.Pointer {
t = t.Elem()
}
if t == reflect.TypeOf(time.Duration(0)) {
return map[string]interface{}{"type": "string", "format": "duration"}
}
switch t.Kind() {
case reflect.String:
return map[string]interface{}{"type": "string"}
case reflect.Bool:
return map[string]interface{}{"type": "boolean"}
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64,
reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
return map[string]interface{}{"type": "integer"}
case reflect.Float32, reflect.Float64:
return map[string]interface{}{"type": "number"}
case reflect.Slice, reflect.Array:
return map[string]interface{}{"type": "array", "items": reflectField(t.Elem())}
case reflect.Map:
return map[string]interface{}{"type": "object", "additionalProperties": reflectField(t.Elem())}
case reflect.Struct:
return reflectStruct(t)
default:
return map[string]interface{}{}
}
}
func splitYAMLTag(tag string) (string, []string) {
parts := strings.Split(tag, ",")
return parts[0], parts[1:]
}
package config
import (
"errors"
"fmt"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/sshdconf"
"golang.org/x/text/language"
)
// ValidationResult represents a single validation finding.
type ValidationResult struct {
Level string // "error", "warn", "ok"
Field string // dotted path matching YAML keys
Message string
}
// String implements the Stringer interface for nice printing.
func (v ValidationResult) String() string {
return fmt.Sprintf("[%s] %s: %s", strings.ToUpper(v.Level), v.Field, v.Message)
}
// Validate checks the config for errors, warnings, and emits OK for valid sections.
func Validate(cfg *Config) []ValidationResult {
var results []ValidationResult
for index, pattern := range cfg.AccountRoots {
if err := platform.ValidateAccountRootPattern(pattern); err != nil {
results = append(results, ValidationResult{"error", fmt.Sprintf("account_roots[%d]", index), err.Error()})
}
}
if cfg.ObserveMode() {
results = append(results, ValidationResult{"ok", "mode", "observe (CSM does not change host state on this host)"})
} else {
results = append(results, ValidationResult{"ok", "mode", ModeEnforce})
}
// --- Hostname ---
if cfg.Hostname == "" || cfg.Hostname == "SET_HOSTNAME_HERE" {
results = append(results, ValidationResult{"error", "hostname", "hostname is not set"})
} else {
results = append(results, ValidationResult{"ok", "hostname", cfg.Hostname})
}
// --- Alerts ---
if !cfg.Alerts.Email.Enabled && !cfg.Alerts.Webhook.Enabled {
results = append(results, ValidationResult{"error", "alerts", "no alert method enabled (enable email or webhook)"})
}
// --- Email alerts ---
if cfg.Alerts.Email.Enabled {
if len(cfg.Alerts.Email.To) == 0 {
results = append(results, ValidationResult{"error", "alerts.email.to", "email alerts enabled but no recipients configured"})
} else {
valid := true
for _, to := range cfg.Alerts.Email.To {
if to == "SET_EMAIL_HERE" || !strings.Contains(to, "@") {
results = append(results, ValidationResult{"error", "alerts.email.to", fmt.Sprintf("invalid email recipient: %s", to)})
valid = false
}
}
if valid {
results = append(results, ValidationResult{"ok", "alerts.email.to", strings.Join(cfg.Alerts.Email.To, ", ")})
}
}
if cfg.Alerts.Email.From == "" {
results = append(results, ValidationResult{"error", "alerts.email.from", "email alerts enabled but no from address configured"})
}
if cfg.Alerts.Email.SMTP == "" {
results = append(results, ValidationResult{"error", "alerts.email.smtp", "email alerts enabled but no SMTP server configured"})
} else {
results = append(results, ValidationResult{"ok", "alerts.email.smtp", cfg.Alerts.Email.SMTP})
}
}
// --- Webhook ---
if cfg.Alerts.Webhook.Enabled {
if cfg.Alerts.Webhook.URL == "" {
results = append(results, ValidationResult{"error", "alerts.webhook.url", "webhook alerts enabled but no URL configured"})
} else {
results = append(results, ValidationResult{"ok", "alerts.webhook.url", RedactURL(cfg.Alerts.Webhook.URL)})
}
switch cfg.Alerts.Webhook.Type {
case "", "slack", "discord", "generic", "phpanel":
default:
results = append(results, ValidationResult{"error", "alerts.webhook.type", fmt.Sprintf("unknown webhook type %q", cfg.Alerts.Webhook.Type)})
}
if cfg.Alerts.Webhook.Type == "phpanel" {
secret := cfg.Alerts.Webhook.HMACSecret
if cfg.Alerts.Webhook.HMACSecretEnv != "" {
if v := os.Getenv(cfg.Alerts.Webhook.HMACSecretEnv); v != "" {
secret = v
}
}
if secret == "" {
field := "alerts.webhook.hmac_secret"
if cfg.Alerts.Webhook.HMACSecretEnv != "" {
field = "alerts.webhook.hmac_secret_env"
}
results = append(results, ValidationResult{"error", field, "phpanel webhook enabled but no HMAC secret configured"})
}
}
}
// --- Heartbeat ---
if cfg.Alerts.Heartbeat.Enabled {
if cfg.Alerts.Heartbeat.URL == "" {
results = append(results, ValidationResult{"error", "alerts.heartbeat.url", "heartbeat enabled but no URL configured"})
} else {
results = append(results, ValidationResult{"ok", "alerts.heartbeat.url", RedactURL(cfg.Alerts.Heartbeat.URL)})
}
}
// --- MaxPerHour ---
if cfg.Alerts.MaxPerHour <= 0 {
results = append(results, ValidationResult{"error", "alerts.max_per_hour", "max_per_hour must be > 0"})
}
// --- WebUI ---
results = append(results, shortTokenWarnings(cfg)...)
for _, origin := range cfg.WebUI.AllowedOrigins {
if err := validateBrowserOrigin(origin); err != nil {
results = append(results, ValidationResult{"error", "webui.allowed_origins", fmt.Sprintf("%q: %v", origin, err)})
}
}
if cfg.WebUI.Enabled {
if err := validateWebUITokens(cfg); err != nil {
results = append(results, ValidationResult{"error", "webui.tokens", err.Error()})
}
tokenCount, adminCount := webUITokenCounts(cfg)
if tokenCount == 0 {
results = append(results, ValidationResult{"error", "webui.tokens", "webui enabled but no auth token configured"})
} else {
results = append(results, ValidationResult{"ok", "webui", fmt.Sprintf("listening on %s", cfg.WebUI.Listen)})
if adminCount == 0 {
results = append(results, ValidationResult{"warn", "webui.tokens", "no admin-scope token configured; browser login and admin API calls are disabled"})
}
}
}
// --- Trusted countries ---
for _, cc := range cfg.Suppressions.TrustedCountries {
if len(cc) != 2 {
results = append(results, ValidationResult{"error", "suppressions.trusted_countries", fmt.Sprintf("invalid country code: %q (expected 2-letter ISO code)", cc)})
}
}
// Credentials only authorize a download; they say nothing about whether a
// database exists. Checking them instead of the database let a wrong key
// pass validation while no database was ever fetched, so the setting was
// inert and reported healthy. Check what the daemon actually reads.
if len(cfg.Suppressions.TrustedCountries) > 0 && !geoIPCityDatabasePresent(cfg.StatePath) {
remedy := "Provision that database, or set geoip.account_id and geoip.license_key and run csm update-geoip"
if cfg.GeoIP.AccountID != "" && cfg.GeoIP.LicenseKey != "" {
remedy = "Credentials are set but no database has been downloaded; run csm update-geoip and check it reports success"
}
results = append(results, ValidationResult{"warn", "suppressions.trusted_countries", "configured but no GeoLite2-City database is installed, so country lookups return nothing and no address is ever treated as trusted. " + remedy})
}
// --- Block digest ---
switch cfg.Alerts.BlockDigest.SendOn {
case "", "any", "customer":
default:
results = append(results, ValidationResult{"error", "alerts.block_digest.send_on", fmt.Sprintf("unknown send_on %q (want any|customer)", cfg.Alerts.BlockDigest.SendOn)})
}
switch cfg.Alerts.BlockDigest.Channel {
case "":
if cfg.Alerts.BlockDigest.Enabled && !cfg.Alerts.Email.Enabled && !cfg.Alerts.Webhook.Enabled {
results = append(results, ValidationResult{"error", "alerts.block_digest.channel", "empty channel requires email or webhook alerts to be enabled"})
}
case "email":
if cfg.Alerts.BlockDigest.Enabled && !cfg.Alerts.Email.Enabled {
results = append(results, ValidationResult{"error", "alerts.block_digest.channel", "channel email requires email alerts to be enabled"})
}
case "webhook":
if cfg.Alerts.BlockDigest.Enabled && !cfg.Alerts.Webhook.Enabled {
results = append(results, ValidationResult{"error", "alerts.block_digest.channel", "channel webhook requires webhook alerts to be enabled"})
}
default:
results = append(results, ValidationResult{"error", "alerts.block_digest.channel", fmt.Sprintf("unknown channel %q (want email|webhook)", cfg.Alerts.BlockDigest.Channel)})
}
if cfg.Alerts.BlockDigest.Interval != "" {
if d, err := time.ParseDuration(cfg.Alerts.BlockDigest.Interval); err != nil {
results = append(results, ValidationResult{"error", "alerts.block_digest.interval", fmt.Sprintf("unparseable duration: %s", cfg.Alerts.BlockDigest.Interval)})
} else if d <= 0 {
results = append(results, ValidationResult{"error", "alerts.block_digest.interval", "interval must be > 0"})
}
}
for _, cc := range cfg.Alerts.BlockDigest.Countries {
if len(cc) != 2 {
results = append(results, ValidationResult{"error", "alerts.block_digest.countries", fmt.Sprintf("invalid country code: %q (expected 2-letter ISO code)", cc)})
}
}
if cfg.Alerts.BlockDigest.MinBlock < 0 {
results = append(results, ValidationResult{"error", "alerts.block_digest.min_block", "min_block must be >= 0"})
}
// --- Duration fields ---
if cfg.AutoResponse.BlockExpiry != "" {
if d, err := time.ParseDuration(cfg.AutoResponse.BlockExpiry); err != nil {
results = append(results, ValidationResult{"error", "auto_response.block_expiry", fmt.Sprintf("unparseable duration: %s", cfg.AutoResponse.BlockExpiry)})
} else if d <= 0 {
results = append(results, ValidationResult{"error", "auto_response.block_expiry", "block_expiry must be a positive duration"})
}
} else if cfg.AutoResponse.Enabled && cfg.AutoResponse.BlockIPs {
results = append(results, ValidationResult{"error", "auto_response.block_expiry", fmt.Sprintf("block_expiry must be a positive duration when IP blocking is enabled; omit the key to use %s", DefaultBlockExpiry)})
}
if v := strings.ToLower(strings.TrimSpace(cfg.AutoResponse.VirtualPatchExposedFiles)); v != "" &&
v != VirtualPatchOff && v != VirtualPatchManual && v != VirtualPatchAuto {
results = append(results, ValidationResult{"error", "auto_response.virtual_patch_exposed_files", fmt.Sprintf("must be off, manual, or auto (got %q)", cfg.AutoResponse.VirtualPatchExposedFiles)})
}
if v := cfg.Thresholds.DropperUnlinkTTLSec; v != 0 && (v < MinDropperUnlinkTTLSec || v > MaxDropperUnlinkTTLSec) {
results = append(results, ValidationResult{"error", "thresholds.dropper_unlink_ttl_sec", fmt.Sprintf("must be between %d and %d seconds (got %d)", MinDropperUnlinkTTLSec, MaxDropperUnlinkTTLSec, v)})
}
if cfg.Signatures.UpdateInterval != "" {
if _, err := time.ParseDuration(cfg.Signatures.UpdateInterval); err != nil {
results = append(results, ValidationResult{"error", "signatures.update_interval", fmt.Sprintf("unparseable duration: %s", cfg.Signatures.UpdateInterval)})
}
}
if cfg.Signatures.YaraForge.UpdateInterval != "" {
if _, err := time.ParseDuration(cfg.Signatures.YaraForge.UpdateInterval); err != nil {
results = append(results, ValidationResult{"error", "signatures.yara_forge.update_interval", fmt.Sprintf("unparseable duration: %s", cfg.Signatures.YaraForge.UpdateInterval)})
}
}
if cfg.EmailAV.ScanTimeout != "" {
if _, err := time.ParseDuration(cfg.EmailAV.ScanTimeout); err != nil {
results = append(results, ValidationResult{"error", "email_av.scan_timeout", fmt.Sprintf("unparseable duration: %s", cfg.EmailAV.ScanTimeout)})
}
}
if cfg.GeoIP.UpdateInterval != "" {
if _, err := time.ParseDuration(cfg.GeoIP.UpdateInterval); err != nil {
results = append(results, ValidationResult{"error", "geoip.update_interval", fmt.Sprintf("unparseable duration: %s", cfg.GeoIP.UpdateInterval)})
}
}
if cfg.Reputation.BotRanges.UpdateInterval != "" {
if d, err := time.ParseDuration(cfg.Reputation.BotRanges.UpdateInterval); err != nil {
results = append(results, ValidationResult{"error", "reputation.bot_ranges.update_interval", fmt.Sprintf("unparseable duration: %s", cfg.Reputation.BotRanges.UpdateInterval)})
} else if d < time.Hour {
results = append(results, ValidationResult{"error", "reputation.bot_ranges.update_interval", "update_interval must be at least 1h"})
}
}
if cfg.AutoResponse.NetBlockWindow != "" {
if d, err := time.ParseDuration(cfg.AutoResponse.NetBlockWindow); err != nil {
results = append(results, ValidationResult{"error", "auto_response.netblock_window", fmt.Sprintf("unparseable duration: %s", cfg.AutoResponse.NetBlockWindow)})
} else if d <= 0 {
results = append(results, ValidationResult{"error", "auto_response.netblock_window", "netblock_window must be a positive duration"})
}
} else if cfg.AutoResponse.NetBlock {
results = append(results, ValidationResult{"error", "auto_response.netblock_window", fmt.Sprintf("netblock_window must be a positive duration when netblock is enabled; omit the key to use %s", DefaultNetBlockWindow)})
}
if cfg.AutoResponse.PermBlockInterval != "" {
if d, err := time.ParseDuration(cfg.AutoResponse.PermBlockInterval); err != nil {
results = append(results, ValidationResult{"error", "auto_response.permblock_interval", fmt.Sprintf("unparseable duration: %s", cfg.AutoResponse.PermBlockInterval)})
} else if d <= 0 {
results = append(results, ValidationResult{"error", "auto_response.permblock_interval", "permblock_interval must be a positive duration"})
}
} else if cfg.AutoResponse.PermBlock {
results = append(results, ValidationResult{"error", "auto_response.permblock_interval", fmt.Sprintf("permblock_interval must be a positive duration when permblock is enabled; omit the key to use %s", DefaultPermBlockInterval)})
}
if cfg.AutoResponse.MailAuthRecovery.DownGrace != "" {
if d, err := time.ParseDuration(cfg.AutoResponse.MailAuthRecovery.DownGrace); err != nil {
results = append(results, ValidationResult{"error", "auto_response.mail_auth_recovery.down_grace", fmt.Sprintf("unparseable duration: %s", cfg.AutoResponse.MailAuthRecovery.DownGrace)})
} else if d <= 0 {
results = append(results, ValidationResult{"error", "auto_response.mail_auth_recovery.down_grace", "down_grace must be > 0"})
}
}
// --- Retention ---
if cfg.Retention.Enabled {
if cfg.Retention.SweepInterval != "" {
if _, err := time.ParseDuration(cfg.Retention.SweepInterval); err != nil {
results = append(results, ValidationResult{"error", "retention.sweep_interval", fmt.Sprintf("unparseable duration: %s", cfg.Retention.SweepInterval)})
}
}
if cfg.Retention.FindingsDays < 0 {
results = append(results, ValidationResult{"error", "retention.findings_days", fmt.Sprintf("findings_days must be >= 0 (0 disables the sweep), got %d", cfg.Retention.FindingsDays)})
}
if cfg.Retention.HistoryDays < 0 {
results = append(results, ValidationResult{"error", "retention.history_days", fmt.Sprintf("history_days must be >= 0, got %d", cfg.Retention.HistoryDays)})
}
if cfg.Retention.ReputationDays < 0 {
results = append(results, ValidationResult{"error", "retention.reputation_days", fmt.Sprintf("reputation_days must be >= 0, got %d", cfg.Retention.ReputationDays)})
}
}
if cfg.Retention.CompactMinSizeMB < 0 {
results = append(results, ValidationResult{"error", "retention.compact_min_size_mb", fmt.Sprintf("compact_min_size_mb must be >= 0, got %d", cfg.Retention.CompactMinSizeMB)})
}
if cfg.Retention.CompactFillRatio < 0 || cfg.Retention.CompactFillRatio > 1 || (cfg.Retention.Enabled && cfg.Retention.CompactFillRatio == 0) {
results = append(results, ValidationResult{"error", "retention.compact_fill_ratio", fmt.Sprintf("compact_fill_ratio must be in (0, 1], got %v", cfg.Retention.CompactFillRatio)})
}
results = append(results, confdResults(cfg)...)
// --- Firewall ---
if cfg.Firewall != nil {
for _, e := range validateDOSExemptRanges(cfg.Firewall.DOSExemptRanges) {
results = append(results, ValidationResult{"error", "firewall.dos_exempt_ranges", e})
}
if cfg.Firewall.Enabled {
// 0 disables the connection meter: it is what the engine does with
// the value and what the web UI's own help text promises. Absent
// keys are filled from the shipped defaults before validation runs
// (see applyFirewallFieldDefaults), so a 0 reaching here is always
// the operator's explicit choice rather than an unset field.
if cfg.Firewall.ConnRateLimit < 0 {
results = append(results, ValidationResult{"error", "firewall.conn_rate_limit", fmt.Sprintf("conn_rate_limit must be >= 0 when firewall enabled (0 = disabled), got %d", cfg.Firewall.ConnRateLimit)})
}
if cfg.Firewall.ConnLimit < 0 {
results = append(results, ValidationResult{"error", "firewall.conn_limit", "conn_limit must be >= 0 when firewall enabled (0 = disabled)"})
}
if cfg.Firewall.ConnRateLimit >= 0 && cfg.Firewall.ConnLimit >= 0 {
results = append(results, ValidationResult{"ok", "firewall", fmt.Sprintf("enabled, conn_rate_limit=%s, conn_limit=%s",
limitSummary(cfg.Firewall.ConnRateLimit), limitSummary(cfg.Firewall.ConnLimit))})
}
}
results = append(results, firewallLockoutResults(cfg)...)
results = append(results, firewallEgressResults(cfg)...)
results = append(results, firewallValueResults(cfg.Firewall)...)
}
results = append(results, centralActionResults(cfg)...)
results = append(results, blockAtSeverityResults(cfg)...)
// --- Challenge ---
if cfg.Challenge.Difficulty < 0 || cfg.Challenge.Difficulty > 5 {
results = append(results, ValidationResult{"error", "challenge.difficulty", fmt.Sprintf("difficulty must be 0-5, got %d", cfg.Challenge.Difficulty)})
}
if cfg.Challenge.ListenPort < 0 || cfg.Challenge.ListenPort > 65535 {
results = append(results, ValidationResult{"error", "challenge.listen_port", fmt.Sprintf("listen_port must be 0-65535, got %d", cfg.Challenge.ListenPort)})
} else if cfg.Challenge.Enabled && cfg.Challenge.ListenPort == 0 {
results = append(results, ValidationResult{"error", "challenge.listen_port", fmt.Sprintf("listen_port must be 1-65535 when challenge.enabled, got %d", cfg.Challenge.ListenPort)})
}
// --- EmailAV ---
if cfg.EmailAV.Enabled && cfg.EmailAV.MaxAttachmentSize <= 0 {
results = append(results, ValidationResult{"error", "email_av.max_attachment_size", "max_attachment_size must be > 0 when email_av enabled"})
}
if cfg.EmailAV.FailMode != "" && cfg.EmailAV.FailMode != "open" && cfg.EmailAV.FailMode != "tempfail" {
results = append(results, ValidationResult{"error", "email_av.fail_mode",
fmt.Sprintf("invalid fail_mode %q: must be \"open\" or \"tempfail\"", cfg.EmailAV.FailMode)})
}
if cfg.Signatures.UpdateURL != "" && cfg.Signatures.SigningKey == "" {
results = append(results, ValidationResult{"error", "signatures.signing_key",
"signing_key is required when signatures.update_url is configured"})
}
if cfg.Signatures.YaraForge.Enabled && cfg.Signatures.SigningKey == "" {
results = append(results, ValidationResult{"error", "signatures.signing_key",
"signing_key is required when signatures.yara_forge.enabled is true"})
}
if cfg.Signatures.YaraForge.Enabled && cfg.Signatures.YaraForge.DownloadURL == "" {
results = append(results, ValidationResult{"error", "signatures.yara_forge.download_url",
"download_url is required because upstream YARA Forge releases do not publish CSM detached signatures"})
}
if cfg.Signatures.YaraForge.DownloadURL != "" {
if err := validateSignatureURL(cfg.Signatures.YaraForge.DownloadURL, true); err != nil {
results = append(results, ValidationResult{"error", "signatures.yara_forge.download_url", err.Error()})
}
}
if cfg.Signatures.UpdateURL != "" {
if err := validateSignatureURL(cfg.Signatures.UpdateURL, false); err != nil {
results = append(results, ValidationResult{"error", "signatures.update_url", err.Error()})
}
}
if err := validateDirectSMTPEgress(cfg); err != nil {
results = append(results, ValidationResult{"error", "detection.direct_smtp_egress", err.Error()})
}
if err := validateBPFEnforcement(cfg); err != nil {
results = append(results, ValidationResult{"error", "bpf_enforcement", err.Error()})
}
// --- EmailProtection ---
if cfg.EmailProtection.RateWarnThreshold > 0 && cfg.EmailProtection.RateWarnThreshold < 10 {
results = append(results, ValidationResult{"warn", "email_protection.rate_warn_threshold", "rate_warn_threshold < 10 may cause excessive alerts"})
}
if cfg.EmailProtection.RateCritThreshold > 0 && cfg.EmailProtection.RateCritThreshold <= cfg.EmailProtection.RateWarnThreshold {
results = append(results, ValidationResult{"error", "email_protection.rate_crit_threshold", "rate_crit_threshold must be > rate_warn_threshold"})
}
if cfg.EmailProtection.RateWindowMin > 0 && (cfg.EmailProtection.RateWindowMin < 5 || cfg.EmailProtection.RateWindowMin > 60) {
results = append(results, ValidationResult{"error", "email_protection.rate_window_min", "rate_window_min must be between 5 and 60"})
}
if cfg.EmailProtection.PasswordCheckIntervalMin > 0 && cfg.EmailProtection.PasswordCheckIntervalMin < 60 {
results = append(results, ValidationResult{"warn", "email_protection.password_check_interval_min", "password_check_interval_min < 60 may cause high CPU from doveadm"})
}
if cfg.Thresholds.FTPFailWindowMin != 0 &&
(cfg.Thresholds.FTPFailWindowMin < 1 || cfg.Thresholds.FTPFailWindowMin > 1440) {
results = append(results, ValidationResult{"error", "thresholds.ftp_fail_window_min",
fmt.Sprintf("ftp_fail_window_min must be between 1 and 1440, got %d", cfg.Thresholds.FTPFailWindowMin)})
}
// --- EmailProtection.PHPRelay bounds ---
// Bounds checks fire only when the operator has supplied a value
// (zero means "use the applyDefaults value" or, for AccountVolumePerHour,
// "auto-derive from cPanel maxemailsperhour"). PoliciesDir is NOT
// validated here -- filesystem state probes belong in ValidateDeep.
pr := cfg.EmailProtection.PHPRelay
if pr.RateWindowMin != 0 && (pr.RateWindowMin < 1 || pr.RateWindowMin > 60) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.rate_window_min", fmt.Sprintf("rate_window_min must be between 1 and 60, got %d", pr.RateWindowMin)})
}
if pr.HeaderScoreVolumeMin != 0 && (pr.HeaderScoreVolumeMin < 2 || pr.HeaderScoreVolumeMin > 100) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.header_score_volume_min", fmt.Sprintf("header_score_volume_min must be between 2 and 100, got %d", pr.HeaderScoreVolumeMin)})
}
if pr.AbsoluteVolumePerHour != 0 && (pr.AbsoluteVolumePerHour < 10 || pr.AbsoluteVolumePerHour > 1000) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.absolute_volume_per_hour", fmt.Sprintf("absolute_volume_per_hour must be between 10 and 1000, got %d", pr.AbsoluteVolumePerHour)})
}
// AccountVolumePerHour: 0 is the documented "auto-derive" sentinel;
// only reject explicitly out-of-range positive values.
if pr.AccountVolumePerHour < 0 || pr.AccountVolumePerHour > 5000 {
results = append(results, ValidationResult{"error", "email_protection.php_relay.account_volume_per_hour", fmt.Sprintf("account_volume_per_hour must be between 0 (auto-derive) and 5000, got %d", pr.AccountVolumePerHour)})
}
if pr.ReputationFailuresPer24h != 0 && (pr.ReputationFailuresPer24h < 1 || pr.ReputationFailuresPer24h > 50) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.reputation_failures_per_24h", fmt.Sprintf("reputation_failures_per_24h must be between 1 and 50, got %d", pr.ReputationFailuresPer24h)})
}
if pr.FanoutDistinctScripts != 0 && (pr.FanoutDistinctScripts < 2 || pr.FanoutDistinctScripts > 20) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.fanout_distinct_scripts", fmt.Sprintf("fanout_distinct_scripts must be between 2 and 20, got %d", pr.FanoutDistinctScripts)})
}
if pr.FanoutDistinctRecipients != 0 && (pr.FanoutDistinctRecipients < 1 || pr.FanoutDistinctRecipients > 100) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.fanout_distinct_recipients", fmt.Sprintf("fanout_distinct_recipients must be between 1 and 100, got %d", pr.FanoutDistinctRecipients)})
}
if pr.FanoutWindowMin != 0 && (pr.FanoutWindowMin < 1 || pr.FanoutWindowMin > 60) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.fanout_window_min", fmt.Sprintf("fanout_window_min must be between 1 and 60, got %d", pr.FanoutWindowMin)})
}
if pr.BaselineSigma != 0 && (pr.BaselineSigma < 2.0 || pr.BaselineSigma > 6.0) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.baseline_sigma", fmt.Sprintf("baseline_sigma must be between 2.0 and 6.0, got %v", pr.BaselineSigma)})
}
if pr.BaselineObservationDays != 0 && (pr.BaselineObservationDays < 1 || pr.BaselineObservationDays > 30) {
results = append(results, ValidationResult{"error", "email_protection.php_relay.baseline_observation_days", fmt.Sprintf("baseline_observation_days must be between 1 and 30, got %d", pr.BaselineObservationDays)})
}
// --- AutoResponse.PHPRelay bounds ---
if cfg.AutoResponse.PHPRelay.MaxActionsPerMinute != 0 && (cfg.AutoResponse.PHPRelay.MaxActionsPerMinute < 1 || cfg.AutoResponse.PHPRelay.MaxActionsPerMinute > 600) {
results = append(results, ValidationResult{"error", "auto_response.php_relay.max_actions_per_minute", fmt.Sprintf("max_actions_per_minute must be between 1 and 600, got %d", cfg.AutoResponse.PHPRelay.MaxActionsPerMinute)})
}
for key, value := range map[string]int{
"max_file_actions_per_hour": cfg.AutoResponse.MaxFileActionsPerHour,
"max_file_actions_per_account_per_hour": cfg.AutoResponse.MaxFileActionsPerAccountPerHour,
"max_file_action_failures_per_hour": cfg.AutoResponse.MaxFileActionFailuresPerHour,
} {
if value < 0 || value > MaxFileResponseLimit {
results = append(results, ValidationResult{"error", "auto_response." + key, fmt.Sprintf("must be between 0 and %d (0 uses the default)", MaxFileResponseLimit)})
}
}
if cfg.AutoResponse.MaxBlocksPerHour < 0 {
results = append(results, ValidationResult{"error", "auto_response.max_blocks_per_hour", fmt.Sprintf("max_blocks_per_hour must be >= 0 (0 uses default %d), got %d", DefaultMaxBlocksPerHour, cfg.AutoResponse.MaxBlocksPerHour)})
}
if cfg.AutoResponse.MailAuthRecovery.MaxRestartsPerHour < 0 {
results = append(results, ValidationResult{"error", "auto_response.mail_auth_recovery.max_restarts_per_hour", fmt.Sprintf("max_restarts_per_hour must be >= 0 (0 uses default 3), got %d", cfg.AutoResponse.MailAuthRecovery.MaxRestartsPerHour)})
}
if cfg.AutoResponse.MailAuthRecovery.RestartEnabled && strings.TrimSpace(cfg.AutoResponse.MailAuthRecovery.RestartCommand) == "" {
results = append(results, ValidationResult{"error", "auto_response.mail_auth_recovery.restart_command", "restart_command is required when restart_enabled is true"})
}
// --- SMTP brute-force thresholds ---
t := cfg.Thresholds
if t.ExposedFileScanDepth != 0 &&
(t.ExposedFileScanDepth < 1 || t.ExposedFileScanDepth > MaxExposedFileScanDepth) {
results = append(results, ValidationResult{
"error",
"thresholds.exposed_file_scan_depth",
fmt.Sprintf("exposed_file_scan_depth must be between 1 and %d", MaxExposedFileScanDepth),
})
}
if t.DomlogMaxFiles != 0 && (t.DomlogMaxFiles < 1 || t.DomlogMaxFiles > 100000) {
results = append(results, ValidationResult{"error", "thresholds.domlog_max_files", "domlog_max_files must be between 1 and 100000"})
}
if t.AccountScanMaxFiles != 0 && (t.AccountScanMaxFiles < 1 || t.AccountScanMaxFiles > 100000) {
results = append(results, ValidationResult{"error", "thresholds.account_scan_max_files", "account_scan_max_files must be between 1 and 100000"})
}
if t.FullScanMaxFileMB != 0 && (t.FullScanMaxFileMB < 1 || t.FullScanMaxFileMB > 4096) {
results = append(results, ValidationResult{"error", "thresholds.full_scan_max_file_mb", "full_scan_max_file_mb must be between 1 and 4096"})
}
if t.ScanJobRetention != 0 && (t.ScanJobRetention < 1 || t.ScanJobRetention > 1000) {
results = append(results, ValidationResult{"error", "thresholds.scan_job_retention", "scan_job_retention must be between 1 and 1000"})
}
if t.CrontabBase64BlobMaxBytes != 0 {
if t.CrontabBase64BlobMaxBytes < 1024 || t.CrontabBase64BlobMaxBytes > 1048576 {
results = append(results, ValidationResult{"error", "thresholds.crontab_base64_blob_max_bytes", "crontab_base64_blob_max_bytes must be between 1024 and 1048576"})
} else if t.CrontabBase64BlobMaxBytes%4 != 0 {
results = append(results, ValidationResult{"error", "thresholds.crontab_base64_blob_max_bytes", "crontab_base64_blob_max_bytes must be a multiple of 4 (standard base64 alignment)"})
}
}
if t.DomlogTailLines != 0 && (t.DomlogTailLines < 10 || t.DomlogTailLines > 100000) {
results = append(results, ValidationResult{"error", "thresholds.domlog_tail_lines", "domlog_tail_lines must be between 10 and 100000"})
}
if t.DomlogMaxAgeMin != 0 && (t.DomlogMaxAgeMin < 1 || t.DomlogMaxAgeMin > 1440) {
results = append(results, ValidationResult{"error", "thresholds.domlog_max_age_min", "domlog_max_age_min must be between 1 and 1440"})
}
if t.MailLogTailLines != 0 && (t.MailLogTailLines < 10 || t.MailLogTailLines > 100000) {
results = append(results, ValidationResult{"error", "thresholds.mail_log_tail_lines", "mail_log_tail_lines must be between 10 and 100000"})
}
if t.SyslogMessagesTailLines != 0 && (t.SyslogMessagesTailLines < 10 || t.SyslogMessagesTailLines > 100000) {
results = append(results, ValidationResult{"error", "thresholds.syslog_messages_tail_lines", "syslog_messages_tail_lines must be between 10 and 100000"})
}
if t.CredStuffingDistinctAccounts != 0 && (t.CredStuffingDistinctAccounts < 2 || t.CredStuffingDistinctAccounts > 200) {
results = append(results, ValidationResult{"error", "thresholds.cred_stuffing_distinct_accounts", "cred_stuffing_distinct_accounts must be between 2 and 200"})
}
if t.PAMBruteforceThreshold != 0 && (t.PAMBruteforceThreshold < 2 || t.PAMBruteforceThreshold > 1000) {
results = append(results, ValidationResult{"error", "thresholds.pam_bruteforce_threshold", "pam_bruteforce_threshold must be between 2 and 1000"})
}
if t.PAMBruteforceWindowMin != 0 && (t.PAMBruteforceWindowMin < 1 || t.PAMBruteforceWindowMin > 1440) {
results = append(results, ValidationResult{"error", "thresholds.pam_bruteforce_window_min", "pam_bruteforce_window_min must be between 1 and 1440"})
}
if t.HTTPScannerMinRequests < 0 {
results = append(results, ValidationResult{"error", "thresholds.http_scanner_min_requests", "http_scanner_min_requests must be >= 0 (0 disables the detector)"})
}
if t.HTTPScannerErrorPct != 0 && (t.HTTPScannerErrorPct < 1 || t.HTTPScannerErrorPct > 100) {
results = append(results, ValidationResult{"error", "thresholds.http_scanner_error_pct", "http_scanner_error_pct must be between 1 and 100"})
}
if t.HTTPScannerMinDistinctPaths != 0 && (t.HTTPScannerMinDistinctPaths < 1 || t.HTTPScannerMinDistinctPaths > HTTPScannerMaxDistinctPaths) {
results = append(results, ValidationResult{"error", "thresholds.http_scanner_min_distinct_paths", fmt.Sprintf("http_scanner_min_distinct_paths must be between 1 and %d", HTTPScannerMaxDistinctPaths)})
}
for _, code := range t.HTTPScannerStatusCodes {
if code < 100 || code > 599 {
results = append(results, ValidationResult{"error", "thresholds.http_scanner_status_codes", "http_scanner_status_codes entries must be HTTP status codes between 100 and 599"})
break
}
}
switch cfg.AutoResponse.HTTPScannerAction {
case "", "challenge", "block":
default:
results = append(results, ValidationResult{"error", "auto_response.http_scanner_action", fmt.Sprintf("http_scanner_action must be %q or %q, got %q", "challenge", "block", cfg.AutoResponse.HTTPScannerAction)})
}
if t.SMTPBruteForceThreshold != 0 && (t.SMTPBruteForceThreshold < 2 || t.SMTPBruteForceThreshold > 50) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_threshold", "smtp_bruteforce_threshold must be between 2 and 50"})
}
if t.SMTPBruteForceWindowMin != 0 && (t.SMTPBruteForceWindowMin < 1 || t.SMTPBruteForceWindowMin > 60) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_window_min", "smtp_bruteforce_window_min must be between 1 and 60"})
}
if t.SMTPBruteForceSuppressMin != 0 && (t.SMTPBruteForceSuppressMin < 1 || t.SMTPBruteForceSuppressMin > 1440) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_suppress_min", "smtp_bruteforce_suppress_min must be between 1 and 1440"})
}
if t.SMTPBruteForceSubnetThresh != 0 && (t.SMTPBruteForceSubnetThresh < 2 || t.SMTPBruteForceSubnetThresh > 64) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_subnet_threshold", "smtp_bruteforce_subnet_threshold must be between 2 and 64"})
}
if t.SMTPAccountSprayThreshold != 0 && (t.SMTPAccountSprayThreshold < 2 || t.SMTPAccountSprayThreshold > 200) {
results = append(results, ValidationResult{"error", "thresholds.smtp_account_spray_threshold", "smtp_account_spray_threshold must be between 2 and 200"})
}
if t.SMTPBruteForceMaxTracked != 0 && (t.SMTPBruteForceMaxTracked < 1000 || t.SMTPBruteForceMaxTracked > 200000) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_max_tracked", "smtp_bruteforce_max_tracked must be between 1000 and 200000"})
}
if t.SMTPBruteForceSlowThreshold != 0 && (t.SMTPBruteForceSlowThreshold < SlowBruteMinThreshold || t.SMTPBruteForceSlowThreshold > SlowBruteMaxThreshold) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_slow_threshold", fmt.Sprintf("smtp_bruteforce_slow_threshold must be 0 (disabled) or between %d and %d", SlowBruteMinThreshold, SlowBruteMaxThreshold)})
}
if t.SMTPBruteForceSlowWindowMin != 0 && (t.SMTPBruteForceSlowWindowMin < 1 || t.SMTPBruteForceSlowWindowMin > SlowBruteMaxWindowMin) {
results = append(results, ValidationResult{"error", "thresholds.smtp_bruteforce_slow_window_min", fmt.Sprintf("smtp_bruteforce_slow_window_min must be between 1 and %d", SlowBruteMaxWindowMin)})
}
if t.SMTPProbeThreshold != 0 && (t.SMTPProbeThreshold < 10 || t.SMTPProbeThreshold > 10000) {
results = append(results, ValidationResult{"error", "thresholds.smtp_probe_threshold", "smtp_probe_threshold must be between 10 and 10000"})
}
if t.SMTPProbeWindowMin != 0 && (t.SMTPProbeWindowMin < 1 || t.SMTPProbeWindowMin > 60) {
results = append(results, ValidationResult{"error", "thresholds.smtp_probe_window_min", "smtp_probe_window_min must be between 1 and 60"})
}
if t.SMTPProbeSuppressMin != 0 && (t.SMTPProbeSuppressMin < 1 || t.SMTPProbeSuppressMin > 1440) {
results = append(results, ValidationResult{"error", "thresholds.smtp_probe_suppress_min", "smtp_probe_suppress_min must be between 1 and 1440"})
}
if t.SMTPProbeMaxTracked != 0 && (t.SMTPProbeMaxTracked < 1000 || t.SMTPProbeMaxTracked > 200000) {
results = append(results, ValidationResult{"error", "thresholds.smtp_probe_max_tracked", "smtp_probe_max_tracked must be between 1000 and 200000"})
}
if t.MailBruteForceThreshold != 0 && (t.MailBruteForceThreshold < 2 || t.MailBruteForceThreshold > 50) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_threshold", "mail_bruteforce_threshold must be between 2 and 50"})
}
if t.MailBruteForceWindowMin != 0 && (t.MailBruteForceWindowMin < 1 || t.MailBruteForceWindowMin > 60) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_window_min", "mail_bruteforce_window_min must be between 1 and 60"})
}
if t.MailBruteForceSuppressMin != 0 && (t.MailBruteForceSuppressMin < 1 || t.MailBruteForceSuppressMin > 1440) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_suppress_min", "mail_bruteforce_suppress_min must be between 1 and 1440"})
}
if t.MailBruteForceSubnetThresh != 0 && (t.MailBruteForceSubnetThresh < 2 || t.MailBruteForceSubnetThresh > 64) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_subnet_threshold", "mail_bruteforce_subnet_threshold must be between 2 and 64"})
}
if t.MailAccountSprayThreshold != 0 && (t.MailAccountSprayThreshold < 2 || t.MailAccountSprayThreshold > 200) {
results = append(results, ValidationResult{"error", "thresholds.mail_account_spray_threshold", "mail_account_spray_threshold must be between 2 and 200"})
}
if t.MailBruteForceMaxTracked != 0 && (t.MailBruteForceMaxTracked < 1000 || t.MailBruteForceMaxTracked > 200000) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_max_tracked", "mail_bruteforce_max_tracked must be between 1000 and 200000"})
}
if t.MailBruteForceSlowThreshold != 0 && (t.MailBruteForceSlowThreshold < SlowBruteMinThreshold || t.MailBruteForceSlowThreshold > SlowBruteMaxThreshold) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_slow_threshold", fmt.Sprintf("mail_bruteforce_slow_threshold must be 0 (disabled) or between %d and %d", SlowBruteMinThreshold, SlowBruteMaxThreshold)})
}
if t.MailBruteForceSlowWindowMin != 0 && (t.MailBruteForceSlowWindowMin < 1 || t.MailBruteForceSlowWindowMin > SlowBruteMaxWindowMin) {
results = append(results, ValidationResult{"error", "thresholds.mail_bruteforce_slow_window_min", fmt.Sprintf("mail_bruteforce_slow_window_min must be between 1 and %d", SlowBruteMaxWindowMin)})
}
if field, err := validateMailBruteAccountKeyField(cfg); err != nil {
results = append(results, ValidationResult{"error", field, err.Error()})
}
// --- Mail log source ---
if field, err := validateMailLogsField(cfg); err != nil {
results = append(results, ValidationResult{"error", field, err.Error()})
}
// --- Reputation.Rspamd ---
if cfg.Reputation.Rspamd.Enabled {
secret := cfg.Reputation.Rspamd.Token
if cfg.Reputation.Rspamd.TokenEnv != "" {
if v := os.Getenv(cfg.Reputation.Rspamd.TokenEnv); v != "" {
secret = v
}
}
if secret == "" {
results = append(results, ValidationResult{"warn", "reputation.rspamd.token", "rspamd enabled but no token configured (rspamd controller history may require auth)"})
}
}
// --- Reputation.Upstream ---
if cfg.Reputation.Upstream.Enabled {
secret := cfg.Reputation.Upstream.Token
if cfg.Reputation.Upstream.TokenEnv != "" {
if v := os.Getenv(cfg.Reputation.Upstream.TokenEnv); v != "" {
secret = v
}
}
if secret == "" {
results = append(results, ValidationResult{"warn", "reputation.upstream.token", "upstream enabled but no token configured (panel endpoint may require auth)"})
}
}
// --- Reputation.VerifiedBots ---
results = append(results, validateVerifiedBots(cfg)...)
// --- AutoResponse.VerdictCallback ---
if field, err := validateVerdictCallbackField(cfg); err != nil {
results = append(results, ValidationResult{"error", field, err.Error()})
}
// --- Debug / pprof ---
if addr := strings.TrimSpace(cfg.Debug.PprofListen); addr != "" {
host, _, err := net.SplitHostPort(addr)
host = strings.TrimSpace(host)
loopback := err == nil && host != "" &&
(strings.EqualFold(host, "localhost") || (net.ParseIP(host) != nil && net.ParseIP(host).IsLoopback()))
if loopback {
results = append(results, ValidationResult{"ok", "debug.pprof_listen", addr})
} else {
results = append(results, ValidationResult{"error", "debug.pprof_listen",
fmt.Sprintf("must be a loopback host:port (127.0.0.1/::1/localhost); %q would expose pprof off-box and is ignored at runtime", addr)})
}
}
results = append(results, blockEscalationResults(cfg)...)
// --- Warnings ---
results = append(results, validateWarnings(cfg)...)
return results
}
// minWebUITokenLength is the shortest Web UI or metrics token that is not
// reported as guessable. The installer generates 64 hex characters.
const minWebUITokenLength = 32
// shortTokenWarnings reports tokens shorter than minWebUITokenLength. They
// are warnings, never errors: startup, `csm validate` and `csm doctor` print
// them, and an existing short token keeps working. Results name the token,
// never its value.
func shortTokenWarnings(cfg *Config) []ValidationResult {
advice := fmt.Sprintf("shorter than %d characters and guessable; replace it with a random value (the installer generates 64 hex characters)", minWebUITokenLength)
var out []ValidationResult
for _, tok := range cfg.WebUI.Tokens {
if tok.Token != "" && len(tok.Token) < minWebUITokenLength {
out = append(out, ValidationResult{"warn", "webui.tokens", fmt.Sprintf("token %q is %s", tok.Name, advice)})
}
}
if len(cfg.WebUI.Tokens) == 0 && cfg.WebUI.AuthToken != "" && len(cfg.WebUI.AuthToken) < minWebUITokenLength {
out = append(out, ValidationResult{"warn", "webui.auth_token", "auth_token is " + advice})
}
if cfg.WebUI.MetricsToken != "" && len(cfg.WebUI.MetricsToken) < minWebUITokenLength {
out = append(out, ValidationResult{"warn", "webui.metrics_token", "metrics_token is " + advice})
}
return out
}
func webUITokenCounts(cfg *Config) (tokens, admins int) {
if len(cfg.WebUI.Tokens) == 0 && cfg.WebUI.AuthToken != "" {
tokens++
admins++
}
for _, tok := range cfg.WebUI.Tokens {
if tok.Token == "" {
continue
}
tokens++
if tok.Scope == "admin" {
admins++
}
}
return tokens, admins
}
func blockEscalationResults(cfg *Config) []ValidationResult {
var results []ValidationResult
// Escalation counters below 2 describe no pattern: one address is not a
// subnet, and one temporary block is not a repeat offender. Omitted keys
// get defaults during Load, so lower values here were explicit.
if cfg.AutoResponse.NetBlock && cfg.AutoResponse.NetBlockThreshold < MinBlockEscalationCount {
results = append(results, ValidationResult{"error", "auto_response.netblock_threshold",
fmt.Sprintf("netblock_threshold must be at least %d when netblock is enabled; raise the threshold or disable netblock (got %d)", MinBlockEscalationCount, cfg.AutoResponse.NetBlockThreshold)})
}
if cfg.AutoResponse.PermBlock && cfg.AutoResponse.PermBlockCount < MinBlockEscalationCount {
results = append(results, ValidationResult{"error", "auto_response.permblock_count",
fmt.Sprintf("permblock_count must be at least %d when permblock is enabled; raise the count or disable permblock (got %d)", MinBlockEscalationCount, cfg.AutoResponse.PermBlockCount)})
}
return results
}
func hasEffectiveInfraIPs(cfg *Config) bool {
var firewallInfra []string
if cfg.Firewall != nil {
firewallInfra = cfg.Firewall.InfraIPs
}
return len(firewall.MergeInfraIPs(cfg.InfraIPs, firewallInfra)) > 0
}
// validateWarnings checks for non-fatal configuration issues.
func validateWarnings(cfg *Config) []ValidationResult {
var results []ValidationResult
// GeoIP credentials set but auto_update explicitly false
if cfg.GeoIP.AccountID != "" && cfg.GeoIP.LicenseKey != "" {
if cfg.GeoIP.AutoUpdate != nil && !*cfg.GeoIP.AutoUpdate {
results = append(results, ValidationResult{"warn", "geoip", "GeoIP credentials configured but auto_update is disabled"})
}
}
// Auto-response enabled but no actions
if cfg.AutoResponse.Enabled {
if !cfg.AutoResponse.KillProcesses && !cfg.AutoResponse.QuarantineFiles && !cfg.AutoResponse.BlockIPs {
results = append(results, ValidationResult{"warn", "auto_response", "auto_response enabled but no actions configured (kill/quarantine/block all false)"})
}
}
// block_ips wants to mutate nftables, but the firewall engine that
// would apply those rules is disabled or absent. Without this check
// the daemon happily logs "auto-blocked" actions that never reach
// the kernel, and operators only notice when attackers keep coming
// back.
if cfg.AutoResponse.Enabled && cfg.AutoResponse.BlockIPs {
if cfg.Firewall == nil || !cfg.Firewall.Enabled {
results = append(results, ValidationResult{"warn", "auto_response.block_ips", "auto-response wants to block IPs but firewall is disabled; blocks will be no-ops"})
}
}
// The firewall trims blank entries while merging the two config sections.
// Validation must use that same effective list or placeholders can hide a
// lockout warning even though they create no kernel accept rule.
hasInfra := hasEffectiveInfraIPs(cfg)
if !hasInfra {
results = append(results, ValidationResult{"warn", "infra_ips", "no infra_ips configured in either top-level or firewall section"})
}
// Firewall enabled but no infra IPs (lockout risk)
if cfg.Firewall != nil && cfg.Firewall.Enabled && !hasInfra {
results = append(results, ValidationResult{"warn", "firewall", "firewall enabled but no infra_ips configured - risk of lockout"})
}
return results
}
// confdResults rejects integrity_exempt entries that can never match a
// fragment. The digest compares bare filenames, so a path, a glob or a file
// the loader would not merge leaves the operator believing a fragment is
// exempt while every rewrite of it still fails the next restart.
func confdResults(cfg *Config) []ValidationResult {
var results []ValidationResult
for _, entry := range cfg.ConfD.IntegrityExempt {
if err := validateExemptFragmentName(entry); err != nil {
results = append(results, ValidationResult{"error", "confd.integrity_exempt", err.Error()})
}
}
return results
}
func validateExemptFragmentName(name string) error {
if name == "" {
return errors.New("entries must be conf.d fragment filenames, got an empty entry")
}
if strings.ContainsAny(name, `/\*?[`) || name == "." || name == ".." {
return fmt.Errorf("%q is not a bare fragment filename; list the file name only, without directories or wildcards", name)
}
if !strings.HasSuffix(name, ".yaml") && !strings.HasSuffix(name, ".yml") {
return fmt.Errorf("%q is not a .yaml or .yml fragment, so conf.d never loads it", name)
}
return nil
}
// firewallValueResults rejects firewall values the engine would otherwise
// accept and quietly reinterpret: a port outside 1-65535 cannot select the
// intended service, an inverted passive-FTP range opens nothing, and a
// port_flood proto that is not "udp" is treated as TCP no matter what the
// operator typed.
func firewallValueResults(fw *firewall.FirewallConfig) []ValidationResult {
var results []ValidationResult
for _, list := range []struct {
field string
ports []int
}{
{"firewall.tcp_in", fw.TCPIn},
{"firewall.tcp_out", fw.TCPOut},
{"firewall.udp_in", fw.UDPIn},
{"firewall.udp_out", fw.UDPOut},
{"firewall.tcp6_in", fw.TCP6In},
{"firewall.tcp6_out", fw.TCP6Out},
{"firewall.required_tcp_out", fw.RequiredTCPOut},
{"firewall.udp6_in", fw.UDP6In},
{"firewall.udp6_out", fw.UDP6Out},
{"firewall.restricted_tcp", fw.RestrictedTCP},
{"firewall.drop_nolog", fw.DropNoLog},
{"firewall.smtp_ports", fw.SMTPPorts},
} {
for _, port := range list.ports {
if !validPort(port) {
results = append(results, ValidationResult{"error", list.field,
fmt.Sprintf("port %d is out of range (1-65535)", port)})
}
}
}
// The engine only builds the range rule when both ends are set, so a
// half-configured range is silently inert rather than wrong.
if fw.PassiveFTPStart != 0 || fw.PassiveFTPEnd != 0 {
if !validPort(fw.PassiveFTPStart) {
results = append(results, ValidationResult{"error", "firewall.passive_ftp_start",
fmt.Sprintf("port %d is out of range (1-65535)", fw.PassiveFTPStart)})
}
if !validPort(fw.PassiveFTPEnd) {
results = append(results, ValidationResult{"error", "firewall.passive_ftp_end",
fmt.Sprintf("port %d is out of range (1-65535)", fw.PassiveFTPEnd)})
}
if validPort(fw.PassiveFTPStart) && validPort(fw.PassiveFTPEnd) && fw.PassiveFTPStart > fw.PassiveFTPEnd {
results = append(results, ValidationResult{"error", "firewall.passive_ftp_start",
fmt.Sprintf("passive FTP range starts at %d but ends at %d", fw.PassiveFTPStart, fw.PassiveFTPEnd)})
}
}
for _, code := range fw.CountryBlock {
if !validCountryCode(code) {
results = append(results, ValidationResult{"error", "firewall.country_block",
fmt.Sprintf("%q is not a two-letter ISO country code", code)})
}
}
for i, pf := range fw.PortFlood {
switch {
case !validPort(pf.Port):
results = append(results, ValidationResult{"error", "firewall.port_flood",
fmt.Sprintf("entry %d: port %d is out of range (1-65535)", i, pf.Port)})
// The case handling is deliberately asymmetric. The engine selects UDP
// with an exact `Proto == "udp"` match and treats every other value as
// TCP, so "TCP" resolves to the protocol the operator meant and is
// safe to accept, while "UDP" would silently become TCP and must be
// rejected. Do not "tidy" this into a single case-insensitive compare.
case pf.Proto != "" && !strings.EqualFold(pf.Proto, "tcp") && pf.Proto != "udp":
results = append(results, ValidationResult{"error", "firewall.port_flood",
fmt.Sprintf("entry %d: proto %q must be \"tcp\" or lowercase \"udp\"", i, pf.Proto)})
case pf.Hits <= 0:
results = append(results, ValidationResult{"error", "firewall.port_flood",
fmt.Sprintf("entry %d: hits must be > 0, got %d", i, pf.Hits)})
case pf.Seconds <= 0:
results = append(results, ValidationResult{"error", "firewall.port_flood",
fmt.Sprintf("entry %d: seconds must be > 0, got %d", i, pf.Seconds)})
}
}
results = append(results, outAllowResults(fw)...)
return results
}
// outAllowResults validates firewall.tcp_out_allow. A malformed entry emits no
// nftables rule at all, so every shape error is reported rather than left to
// fail open as "the fetch just does not work".
func outAllowResults(fw *firewall.FirewallConfig) []ValidationResult {
var results []ValidationResult
for i, r := range fw.TCPOutAllow {
prefix := fmt.Sprintf("entry %d (dst %q)", i, r.Dst)
network, err := firewall.ParseOutAllowDst(r.Dst)
if err != nil {
results = append(results, ValidationResult{"error", "firewall.tcp_out_allow",
fmt.Sprintf("%s: %v", prefix, err)})
continue
}
switch {
case !validPort(r.PortStart):
results = append(results, ValidationResult{"error", "firewall.tcp_out_allow",
fmt.Sprintf("%s: port_start %d is out of range (1-65535)", prefix, r.PortStart)})
continue
case !validPort(r.PortEnd):
results = append(results, ValidationResult{"error", "firewall.tcp_out_allow",
fmt.Sprintf("%s: port_end %d is out of range (1-65535)", prefix, r.PortEnd)})
continue
case r.PortStart > r.PortEnd:
results = append(results, ValidationResult{"error", "firewall.tcp_out_allow",
fmt.Sprintf("%s: range starts at %d but ends at %d, so it matches nothing",
prefix, r.PortStart, r.PortEnd)})
continue
}
// The output chain emits the smtp_block drop before these rules, so an
// overlap cannot actually bypass it. Rejected anyway: rule ordering
// must not be the only guard between a config key and outbound mail.
if fw.SMTPBlock {
for _, port := range fw.SMTPPorts {
if port >= r.PortStart && port <= r.PortEnd {
results = append(results, ValidationResult{"error", "firewall.tcp_out_allow",
fmt.Sprintf("%s: range %d-%d covers smtp port %d while smtp_block is on",
prefix, r.PortStart, r.PortEnd, port)})
break
}
}
}
if ones, bits := network.Mask.Size(); ones == 0 && bits > 0 {
results = append(results, ValidationResult{"warn", "firewall.tcp_out_allow",
fmt.Sprintf("%s: opens ports %d-%d (%d ports) to every destination; scope it to the hosts this server dials",
prefix, r.PortStart, r.PortEnd, r.PortEnd-r.PortStart+1)})
}
if network.IP.To4() == nil && !fw.IPv6 {
results = append(results, ValidationResult{"warn", "firewall.tcp_out_allow",
fmt.Sprintf("%s: IPv6 destination emits no rule while firewall.ipv6 is off", prefix)})
}
}
return results
}
// centralActions mirrors the action constants in internal/reporting, which
// this package cannot import: reporting depends on alert, and alert depends
// on config. TestCentralActionsMatchReportingConstants (external test package,
// so it may import both) fails if the two lists ever drift.
var centralActions = [...]string{"off", "challenge", "block_if_local_corroborated"}
// ValidCentralActions returns every central-intelligence action accepted by
// config validation.
func ValidCentralActions() []string {
return append([]string(nil), centralActions[:]...)
}
// centralActionResults rejects an unrecognised central action. The consumer
// parses it with a default branch that resolves to "challenge", so a typo
// silently turns a corroborated-block policy into a challenge policy and the
// only trace is one daemon log line at startup.
func centralActionResults(cfg *Config) []ValidationResult {
action := cfg.Reputation.Central.Action
if action == "" {
return nil
}
for _, valid := range centralActions {
if action == valid {
return nil
}
}
return []ValidationResult{{"error", "reputation.central.action",
fmt.Sprintf("invalid action %q: must be one of %s", action, strings.Join(centralActions[:], ", "))}}
}
// blockAtSeverityResults rejects an unrecognised incident block threshold.
// The correlator matches "high" and "critical" and ignores anything else, so
// a typo leaves the operator believing incident blocking is armed when the
// hand-off can never fire.
func blockAtSeverityResults(cfg *Config) []ValidationResult {
var results []ValidationResult
for _, entry := range []struct {
field string
value string
}{
{"incidents.spray_suppression.block_at_severity", cfg.Incidents.SpraySuppression.BlockAtSeverity},
{"incidents.auto_block.block_at_severity", cfg.Incidents.AutoBlock.BlockAtSeverity},
} {
if entry.value == "" {
continue
}
switch strings.ToLower(entry.value) {
case "high", "critical":
default:
results = append(results, ValidationResult{"error", entry.field,
fmt.Sprintf("invalid severity %q: must be \"high\" or \"critical\" (empty disables blocking)", entry.value)})
}
}
return results
}
func validPort(p int) bool { return p >= 1 && p <= 65535 }
func validCountryCode(code string) bool {
if len(code) != 2 {
return false
}
region, err := language.ParseRegion(code)
return err == nil && region.IsCountry()
}
// limitSummary renders a firewall limit for operator-facing output so a
// disabled protection reads as "disabled" instead of as a bare 0 that looks
// like a missing value.
func limitSummary(v int) string {
if v == 0 {
return "disabled"
}
return strconv.Itoa(v)
}
// firewallLockoutResults reports the ways an enabled firewall can cut the
// operator off from the host. The web UI runs the same checks before a save;
// running them here covers hand-edited csm.yaml, `csm doctor`, and daemon
// startup, which previously got no warning at all.
//
// These are warnings, never errors. Fronting the web UI with a reverse proxy
// or reaching it over a VPN are legitimate reasons to leave the port out of
// tcp_in, and validation must not refuse a deliberate configuration.
func firewallLockoutResults(cfg *Config) []ValidationResult {
fw := cfg.Firewall
if fw == nil || !fw.Enabled {
return nil
}
var results []ValidationResult
hasInfra := hasEffectiveInfraIPs(cfg)
port, portKnown := webUIListenPort(cfg.WebUI.Listen)
// Naming the port in restricted_tcp is how an operator asks for an
// infra-only service, and the infra accept rule matches before any port
// rule, so such a port stays reachable without appearing in tcp_in. The
// shipped defaults ship exactly this pair; warning about it would make
// every stock install report WARN. Without infra_ips there is nothing left
// to reach it from, so the warning still stands.
infraOnly := portKnown && hasInfra && containsPort(fw.RestrictedTCP, port)
if cfg.WebUI.Enabled && portKnown && !infraOnly && !containsPort(fw.TCPIn, port) {
results = append(results, ValidationResult{"warn", "firewall.tcp_in",
fmt.Sprintf("web UI listens on %d but tcp_in does not allow it; the next firewall apply drops new web UI connections", port)})
}
// A non-empty tcp6_in overrides tcp_in. When empty, IPv6 inherits tcp_in,
// so the IPv4 check above already covers the effective port policy.
if cfg.WebUI.Enabled && portKnown && !infraOnly && fw.IPv6 && len(fw.TCP6In) > 0 && !containsPort(fw.TCP6In, port) {
results = append(results, ValidationResult{"warn", "firewall.tcp6_in",
fmt.Sprintf("IPv6 is managed and tcp6_in does not allow web UI port %d", port)})
}
// restricted_tcp filters matching ports out of the public allow lists; it
// does not add reachability. Without infra_ips, only an overlap with an
// effective allow list changes a port from public to unreachable.
tcp6In := fw.TCP6In
if len(tcp6In) == 0 {
tcp6In = fw.TCPIn
}
restrictedAllowed := portsOverlap(fw.RestrictedTCP, fw.TCPIn) ||
(fw.IPv6 && portsOverlap(fw.RestrictedTCP, tcp6In))
if restrictedAllowed && !hasInfra {
msg := "restricted_tcp filters ports out of the public allow list, but no infra_ips are configured, so those ports are reachable from nowhere"
webUIAllowed := containsPort(fw.TCPIn, port) || (fw.IPv6 && containsPort(tcp6In, port))
if cfg.WebUI.Enabled && portKnown && webUIAllowed && containsPort(fw.RestrictedTCP, port) {
msg = fmt.Sprintf("web UI port %d is in restricted_tcp but no infra_ips are configured, so the web UI is reachable from nowhere", port)
}
results = append(results, ValidationResult{"warn", "firewall.restricted_tcp", msg})
}
return results
}
// sshdConfigPath is where the SSH lockout guard reads the listen ports from.
// It is a test hook, not an operator setting: sshd's own path is fixed.
var sshdConfigPath = sshdconf.DefaultPath
// SetSSHDConfigPath points the SSH lockout guard at another sshd config and
// returns a func restoring the previous path. Tests use it to stay hermetic;
// production always reads sshd's own path.
func SetSSHDConfigPath(path string) func() {
previous := sshdConfigPath
sshdConfigPath = path
return func() { sshdConfigPath = previous }
}
// probeSSHLockout warns when an enabled firewall would not accept the ports
// sshd actually listens on. The shipped tcp_in leaves 22 out, so a host that
// never moved sshd loses SSH on the next apply.
//
// It reads the host, which is why it is a deep probe rather than part of
// Validate: the same csm.yaml must validate identically on any machine.
func probeSSHLockout(cfg *Config) []ValidationResult {
fw := cfg.Firewall
if fw == nil || !fw.Enabled {
return nil
}
// No sshd config means no evidence of a listener. Falling back to the
// compiled default here would warn on every host that does not run sshd.
sshd := sshdconf.Parse(sshdconf.OSFS{}, sshdConfigPath)
if !sshd.Present() {
return nil
}
hasInfra := hasEffectiveInfraIPs(cfg)
var missingTCPIn, missingTCP6In []int
ipv4Ports, ipv6Ports := sshd.RemoteListenPorts()
for _, port := range ipv4Ports {
// Naming the port in restricted_tcp is how an operator asks for an
// infra-only listener. The infra accept rule matches before any port
// rule, so such a port stays reachable without appearing in tcp_in.
if hasInfra && containsPort(fw.RestrictedTCP, port) {
continue
}
if !containsPort(fw.TCPIn, port) {
missingTCPIn = append(missingTCPIn, port)
}
}
if fw.IPv6 {
for _, port := range ipv6Ports {
if hasInfra && containsPort(fw.RestrictedTCP, port) {
continue
}
if len(fw.TCP6In) == 0 {
if !containsPort(fw.TCPIn, port) && !containsPort(missingTCPIn, port) {
missingTCPIn = append(missingTCPIn, port)
}
} else if !containsPort(fw.TCP6In, port) {
missingTCP6In = append(missingTCP6In, port)
}
}
}
var results []ValidationResult
if len(missingTCPIn) > 0 {
subject, pronoun := portPhrase(missingTCPIn)
results = append(results, ValidationResult{"warn", "firewall.tcp_in",
fmt.Sprintf("sshd listens on %s but tcp_in does not allow %s; the next firewall apply drops new SSH connections", subject, pronoun)})
}
if len(missingTCP6In) > 0 {
subject, _ := portPhrase(missingTCP6In)
results = append(results, ValidationResult{"warn", "firewall.tcp6_in",
fmt.Sprintf("IPv6 is managed and tcp6_in does not allow sshd %s", subject)})
}
return results
}
// portPhrase renders a port list plus the pronoun that agrees with it.
func portPhrase(ports []int) (subject, pronoun string) {
sort.Ints(ports)
parts := make([]string, 0, len(ports))
for _, p := range ports {
parts = append(parts, strconv.Itoa(p))
}
if len(parts) == 1 {
return "port " + parts[0], "it"
}
return "ports " + strings.Join(parts, ", "), "them"
}
// webUIListenPort extracts the port from a host:port listen string, falling
// back to the shipped default when the field is unset.
func webUIListenPort(listen string) (int, bool) {
if listen == "" {
listen = "0.0.0.0:9443"
}
_, portStr, err := net.SplitHostPort(listen)
if err != nil {
return 0, false
}
port, err := strconv.Atoi(portStr)
if err != nil {
return 0, false
}
return port, true
}
func containsPort(ports []int, want int) bool {
for _, p := range ports {
if p == want {
return true
}
}
return false
}
func portsOverlap(a, b []int) bool {
for _, port := range a {
if containsPort(b, port) {
return true
}
}
return false
}
// ValidateDeep performs connectivity probes against configured services.
// It does NOT call Validate(); the caller should invoke both separately.
func ValidateDeep(cfg *Config) []ValidationResult {
var results []ValidationResult
// State directory
results = append(results, probeStatePath(cfg.StatePath)...)
results = append(results, probeStatePathSandbox(cfg.StatePath)...)
// Signature rules directory
if cfg.Signatures.RulesDir != "" {
results = append(results, probeRulesDir(cfg.Signatures.RulesDir)...)
}
// SMTP
if cfg.Alerts.Email.Enabled && cfg.Alerts.Email.SMTP != "" {
results = append(results, probeSMTP(cfg.Alerts.Email.SMTP)...)
}
// ClamAV socket
if cfg.EmailAV.Enabled && cfg.EmailAV.ClamdSocket != "" {
results = append(results, probeClamd(cfg.EmailAV.ClamdSocket)...)
}
// TLS cert/key (only when custom paths set)
if cfg.WebUI.TLSCert != "" {
if _, err := os.Stat(cfg.WebUI.TLSCert); err != nil {
results = append(results, ValidationResult{"error", "webui.tls_cert", fmt.Sprintf("file not found: %s", cfg.WebUI.TLSCert)})
} else {
results = append(results, ValidationResult{"ok", "webui.tls_cert", cfg.WebUI.TLSCert})
}
}
if cfg.WebUI.TLSKey != "" {
if _, err := os.Stat(cfg.WebUI.TLSKey); err != nil {
results = append(results, ValidationResult{"error", "webui.tls_key", fmt.Sprintf("file not found: %s", cfg.WebUI.TLSKey)})
} else {
results = append(results, ValidationResult{"ok", "webui.tls_key", cfg.WebUI.TLSKey})
}
}
// Webhook
if cfg.Alerts.Webhook.Enabled && cfg.Alerts.Webhook.URL != "" {
results = append(results, probeWebhook(cfg.Alerts.Webhook.URL)...)
}
// GeoIP database files
if cfg.GeoIP.AccountID != "" && cfg.GeoIP.LicenseKey != "" && len(cfg.GeoIP.Editions) > 0 {
results = append(results, probeGeoIPDBs(cfg.StatePath, cfg.GeoIP.Editions)...)
}
// SSH reachability under the configured firewall policy
results = append(results, probeSSHLockout(cfg)...)
return results
}
// ValidateDeepSection runs only the deep probes relevant to the named
// section, so a save to section X does not fail on an unrelated probe
// for section Y. Section names match the webui settings schema IDs.
//
// For any section without deep probes, returns nil.
func ValidateDeepSection(cfg *Config, section string) []ValidationResult {
switch section {
case "alerts":
var results []ValidationResult
if cfg.Alerts.Email.Enabled && cfg.Alerts.Email.SMTP != "" {
results = append(results, probeSMTP(cfg.Alerts.Email.SMTP)...)
}
if cfg.Alerts.Webhook.Enabled && cfg.Alerts.Webhook.URL != "" {
results = append(results, probeWebhook(cfg.Alerts.Webhook.URL)...)
}
return results
case "email_av":
if cfg.EmailAV.Enabled && cfg.EmailAV.ClamdSocket != "" {
return probeClamd(cfg.EmailAV.ClamdSocket)
}
case "geoip":
if cfg.GeoIP.AccountID != "" && cfg.GeoIP.LicenseKey != "" && len(cfg.GeoIP.Editions) > 0 {
return probeGeoIPDBs(cfg.StatePath, cfg.GeoIP.Editions)
}
case "firewall":
return probeSSHLockout(cfg)
case "challenge":
// probeListenPortAvailable is not yet implemented in this codebase.
// When added, invoke it here: return probeListenPortAvailable(cfg.Challenge.ListenPort).
return nil
}
return nil
}
// systemdUnitFile is the installed service unit the sandbox probe reads.
var systemdUnitFile = "/etc/systemd/system/csm.service"
// probeStatePathSandbox checks that the daemon, not just this CLI process,
// can write state_path: under ProtectSystem=strict only StateDirectory and
// ReadWritePaths grants are writable, and a state_path outside them makes
// the daemon crash-loop on its first write while validate passes.
func probeStatePathSandbox(statePath string) []ValidationResult {
unit, err := readSystemdUnitAndDropIns(systemdUnitFile)
if err != nil {
return nil
}
covered, known := unitCoversStatePath(unit, statePath)
if !known || covered {
return nil
}
return []ValidationResult{{"error", "state_path", fmt.Sprintf("%s is not writable under the service unit's ProtectSystem and ReadWritePaths settings (unit: %s)", statePath, systemdUnitFile)}}
}
// #nosec G304 -- path is the installed service unit's package-constant path,
// and the drop-ins come from globbing that unit's own .d directory.
func readSystemdUnitAndDropIns(path string) (string, error) {
data, err := os.ReadFile(path)
if err != nil {
return "", err
}
var unit strings.Builder
unit.Write(data)
matches, err := filepath.Glob(filepath.Join(path+".d", "*.conf"))
if err != nil {
return "", err
}
for _, dropIn := range matches {
data, err = os.ReadFile(dropIn)
if err != nil {
return "", err
}
unit.WriteByte('\n')
unit.Write(data)
}
return unit.String(), nil
}
// unitCoversStatePath parses a systemd unit and reports whether statePath
// is writable to the service: covered is true when ProtectSystem is not
// strict or the path sits under a StateDirectory or ReadWritePaths grant;
// known is false when the unit text carries no sandbox directives at all.
func unitCoversStatePath(unit, statePath string) (covered, known bool) {
statePath = filepath.Clean(statePath)
protectSystem := "false"
var stateGrants, pathGrants []string
section := ""
for _, raw := range systemdLogicalLines(unit) {
line := strings.TrimSpace(raw)
if strings.HasPrefix(line, "[") && strings.HasSuffix(line, "]") {
section = strings.ToLower(strings.TrimSpace(line[1 : len(line)-1]))
continue
}
if section != "service" {
continue
}
key, value, ok := strings.Cut(line, "=")
if !ok {
continue
}
key = strings.TrimSpace(key)
value = strings.TrimSpace(value)
switch key {
case "ProtectSystem":
known = true
protectSystem = strings.ToLower(strings.Trim(value, "\"'"))
case "StateDirectory":
known = true
if value == "" {
stateGrants = nil
continue
}
words, valid := splitSystemdWords(value)
if !valid {
continue
}
for _, name := range words {
name, _, _ = strings.Cut(name, ":")
if name != "" {
stateGrants = append(stateGrants, filepath.Join("/var/lib", name))
}
}
case "ReadWritePaths":
known = true
if value == "" {
pathGrants = nil
continue
}
words, valid := splitSystemdWords(value)
if !valid {
continue
}
for _, p := range words {
p = strings.TrimLeft(p, "-+!")
if p != "" {
pathGrants = append(pathGrants, filepath.Clean(p))
}
}
}
}
if !known {
return false, false
}
for _, g := range append(stateGrants, pathGrants...) {
if pathContains(g, statePath) {
return true, true
}
}
switch protectSystem {
case "strict":
return false, true
case "full":
return !pathUnderAny(statePath, "/usr", "/boot", "/efi", "/etc"), true
case "true", "yes":
return !pathUnderAny(statePath, "/usr", "/boot", "/efi"), true
default:
return true, true
}
}
func systemdLogicalLines(unit string) []string {
var lines []string
var logical strings.Builder
for _, raw := range strings.Split(unit, "\n") {
raw = strings.TrimSuffix(raw, "\r")
trailingSlashes := 0
for i := len(raw) - 1; i >= 0 && raw[i] == '\\'; i-- {
trailingSlashes++
}
continued := trailingSlashes%2 == 1
if continued {
raw = raw[:len(raw)-1]
}
logical.WriteString(raw)
if continued {
logical.WriteByte(' ')
continue
}
lines = append(lines, logical.String())
logical.Reset()
}
if logical.Len() > 0 {
lines = append(lines, logical.String())
}
return lines
}
func splitSystemdWords(value string) ([]string, bool) {
var words []string
var word strings.Builder
var quote byte
escaped := false
flush := func() {
if word.Len() > 0 {
words = append(words, word.String())
word.Reset()
}
}
for i := 0; i < len(value); i++ {
c := value[i]
if escaped {
word.WriteByte(c)
escaped = false
continue
}
if c == '\\' {
escaped = true
continue
}
if quote != 0 {
if c == quote {
quote = 0
} else {
word.WriteByte(c)
}
continue
}
if c == '\'' || c == '"' {
quote = c
continue
}
if c == ' ' || c == '\t' {
flush()
continue
}
word.WriteByte(c)
}
if escaped {
word.WriteByte('\\')
}
if quote != 0 {
return nil, false
}
flush()
return words, true
}
func pathContains(parent, child string) bool {
parent = filepath.Clean(parent)
child = filepath.Clean(child)
if parent == string(filepath.Separator) {
return filepath.IsAbs(child)
}
return child == parent || strings.HasPrefix(child, parent+string(filepath.Separator))
}
func pathUnderAny(path string, roots ...string) bool {
for _, root := range roots {
if pathContains(root, path) {
return true
}
}
return false
}
// probeStatePath checks that the state directory exists and is writable.
func probeStatePath(path string) []ValidationResult {
info, err := os.Stat(path)
if err != nil {
return []ValidationResult{{"error", "state_path", fmt.Sprintf("directory not found: %s", path)}}
}
if !info.IsDir() {
return []ValidationResult{{"error", "state_path", fmt.Sprintf("not a directory: %s", path)}}
}
probe := filepath.Join(path, ".csm-validate-probe")
// #nosec G304 -- filepath.Join under operator-configured statePath.
f, err := os.Create(probe)
if err != nil {
return []ValidationResult{{"error", "state_path", fmt.Sprintf("directory not writable: %s", path)}}
}
f.Close()
os.Remove(probe)
return []ValidationResult{{"ok", "state_path", path}}
}
// probeRulesDir checks that the rules directory exists and contains rule files.
func probeRulesDir(path string) []ValidationResult {
info, err := os.Stat(path)
if err != nil {
return []ValidationResult{{"error", "signatures.rules_dir", fmt.Sprintf("directory not found: %s", path)}}
}
if !info.IsDir() {
return []ValidationResult{{"error", "signatures.rules_dir", fmt.Sprintf("not a directory: %s", path)}}
}
// Check for rule files
for _, pattern := range []string{"*.yaml", "*.yml", "*.yar", "*.yara"} {
matches, _ := filepath.Glob(filepath.Join(path, pattern))
if len(matches) > 0 {
return []ValidationResult{{"ok", "signatures.rules_dir", fmt.Sprintf("%s (%d rule files)", path, len(matches))}}
}
}
return []ValidationResult{{"error", "signatures.rules_dir", fmt.Sprintf("no rule files (.yaml/.yml/.yar/.yara) found in %s", path)}}
}
// probeSMTP attempts a TCP dial to the SMTP server.
func probeSMTP(addr string) []ValidationResult {
conn, err := net.DialTimeout("tcp", addr, 3*time.Second)
if err != nil {
return []ValidationResult{{"error", "alerts.email.smtp", fmt.Sprintf("cannot connect to %s: %v", addr, err)}}
}
_ = conn.Close()
return []ValidationResult{{"ok", "alerts.email.smtp", fmt.Sprintf("connected to %s", addr)}}
}
// probeClamd attempts to connect to the ClamAV unix socket, falling back to the
// well-known locations so a host whose setting names the wrong path is told
// which one actually answers instead of only that mail is not being scanned.
func probeClamd(socket string) []ValidationResult {
resolved, discovered := ResolveClamdSocket(socket)
if !discovered && socket != "" && !clamdSocketTrusted(socket) {
return []ValidationResult{{"error", "email_av.clamd_socket", fmt.Sprintf(
"%s is in a directory other accounts can write to, so what answers there is not necessarily clamd; move the socket somewhere only root or the clamd service account can write",
socket)}}
}
conn, err := net.DialTimeout("unix", resolved, 3*time.Second)
if err != nil {
return []ValidationResult{{"error", "email_av.clamd_socket", fmt.Sprintf("cannot connect to %s: %v", resolved, err)}}
}
_ = conn.Close()
if discovered {
return []ValidationResult{{"warn", "email_av.clamd_socket", fmt.Sprintf(
"nothing is listening on the configured %s; using %s, which is answering. Set clamd_socket to it so the fallback is not needed",
socket, resolved)}}
}
return []ValidationResult{{"ok", "email_av.clamd_socket", fmt.Sprintf("connected to %s", resolved)}}
}
// probeWebhook performs an HTTP HEAD request to verify the webhook endpoint is reachable.
// DNS/TCP/TLS failures are errors; HTTP status codes (even 401/403/404/405) mean reachable.
func probeWebhook(rawURL string) []ValidationResult {
client := &http.Client{Timeout: 5 * time.Second}
resp, err := client.Head(rawURL)
if err != nil {
var urlErr *url.Error
if errors.As(err, &urlErr) {
err = urlErr.Err
}
return []ValidationResult{{"error", "alerts.webhook.url", fmt.Sprintf("cannot reach %s: %v", RedactURL(rawURL), err)}}
}
resp.Body.Close()
return []ValidationResult{{"ok", "alerts.webhook.url", fmt.Sprintf("reachable (HTTP %d)", resp.StatusCode)}}
}
// validateBrowserOrigin checks that raw has the shape of a browser Origin
// header value the web UI can match: https://host[:port] and nothing else.
func validateBrowserOrigin(raw string) error {
u, err := url.Parse(strings.TrimSpace(raw))
if err != nil {
return fmt.Errorf("parse: %w", err)
}
if !strings.EqualFold(u.Scheme, "https") {
return fmt.Errorf("must be an https origin")
}
if u.Hostname() == "" {
return fmt.Errorf("missing host")
}
if u.User != nil || u.Path != "" || u.RawQuery != "" || u.Fragment != "" {
return fmt.Errorf("must be a bare https://host[:port] origin")
}
return nil
}
func validateSignatureURL(raw string, allowTemplates bool) error {
candidate := strings.TrimSpace(raw)
if allowTemplates {
candidate = sampleSignatureURLTemplate(candidate)
}
u, err := url.Parse(candidate)
if err != nil {
return fmt.Errorf("parse: %w", err)
}
// Signed content and the rollback guard bound what a tampered download
// can do, but a plain-http mirror still hands an on-path attacker every
// old signed release to replay and every request to observe.
if strings.ToLower(u.Scheme) != "https" {
return fmt.Errorf("signatures URL must be an https URL")
}
host := u.Hostname()
if host == "" {
return fmt.Errorf("missing host in %q", raw)
}
return validateSignaturesHost(host)
}
func sampleSignatureURLTemplate(raw string) string {
raw = strings.ReplaceAll(raw, "{tier}", "core")
raw = strings.ReplaceAll(raw, "{version}", "v1.0.0")
return raw
}
// validateSignaturesHost rejects URL hosts that point at loopback,
// link-local, or RFC1918 / ULA ranges. Scoped IPv6 literals are valid URL
// hosts, so use netip instead of net.ParseIP to keep the zone intact.
func validateSignaturesHost(host string) error {
lower := strings.TrimSuffix(strings.ToLower(host), ".")
if lower == "localhost" || lower == "localhost.localdomain" {
return fmt.Errorf("signatures URL host %q is loopback; refuse for production downloads", host)
}
if addr, err := netip.ParseAddr(lower); err == nil {
addr = addr.Unmap()
if addr.IsLoopback() {
return fmt.Errorf("signatures URL host %s is loopback", host)
}
if addr.IsPrivate() {
return fmt.Errorf("signatures URL host %s is RFC1918 / ULA private", host)
}
if addr.IsLinkLocalUnicast() || addr.IsLinkLocalMulticast() {
return fmt.Errorf("signatures URL host %s is link-local", host)
}
if addr.IsUnspecified() {
return fmt.Errorf("signatures URL host %s is unspecified", host)
}
}
return nil
}
// probeGeoIPDBs checks that expected GeoIP database files exist on disk.
func probeGeoIPDBs(statePath string, editions []string) []ValidationResult {
var results []ValidationResult
allOK := true
for _, edition := range editions {
dbPath := filepath.Join(statePath, "geoip", edition+".mmdb")
if _, err := os.Stat(dbPath); err != nil {
results = append(results, ValidationResult{"error", "geoip", fmt.Sprintf("database not found: %s", dbPath)})
allOK = false
}
}
if allOK {
results = append(results, ValidationResult{"ok", "geoip", fmt.Sprintf("all %d edition databases present", len(editions))})
}
return results
}
// validateDOSExemptRanges checks each entry in the dos_exempt_ranges list.
// Per entry: whitespace is trimmed; empty entries are rejected; CIDR notation
// is parsed via net.ParseCIDR and default routes (/0) are rejected; bare IP
// addresses are accepted as /32 or /128 equivalents; anything else is an error.
// Each error string is prefixed with "firewall.dos_exempt_ranges[<i>]:".
func validateDOSExemptRanges(entries []string) []string {
var errs []string
for i, raw := range entries {
entry := strings.TrimSpace(raw)
if entry == "" {
errs = append(errs, fmt.Sprintf("firewall.dos_exempt_ranges[%d]: empty entry", i))
continue
}
if _, ipnet, err := net.ParseCIDR(entry); err == nil {
ones, _ := ipnet.Mask.Size()
if ones == 0 {
errs = append(errs, fmt.Sprintf("firewall.dos_exempt_ranges[%d]: %s is a default route and cannot be used as an exempt range", i, entry))
}
continue
}
if net.ParseIP(entry) != nil {
continue
}
errs = append(errs, fmt.Sprintf("firewall.dos_exempt_ranges[%d]: %q is not a valid CIDR or IP address", i, entry))
}
return errs
}
// validateFirewallConfig validates firewall fields that are checked at load
// time to prevent invalid configuration from reaching the daemon. It returns
// the first multi-error joined string so LoadBytes can return a single error.
func validateFirewallConfig(cfg *Config) error {
if cfg.Firewall == nil {
return nil
}
errs := validateDOSExemptRanges(cfg.Firewall.DOSExemptRanges)
if len(errs) == 0 {
return nil
}
return errors.New(strings.Join(errs, "; "))
}
// geoIPCityDatabasePresent reports whether the City database the daemon loads
// exists. The daemon opens only the state-path copy, so that is the location
// that decides whether a country lookup can resolve anything.
func geoIPCityDatabasePresent(statePath string) bool {
if statePath == "" {
return false
}
info, err := os.Stat(filepath.Join(statePath, "geoip", "GeoLite2-City.mmdb"))
return err == nil && info.Mode().IsRegular() && info.Size() > 0
}
package config
import (
"fmt"
"strings"
"golang.org/x/net/publicsuffix"
"github.com/pidginhost/csm/internal/netutil"
)
// VerifiedBot is one operator-configured good bot. A request whose UA contains
// any UASubstrings is treated as claiming this identity, and is trusted only
// if the source IP matches one configured verification method.
type VerifiedBot struct {
Name string `yaml:"name" json:"name"`
UASubstrings []string `yaml:"ua_substrings,omitempty" json:"ua_substrings,omitempty"`
RDNSSuffixes []string `yaml:"rdns_suffixes,omitempty" json:"rdns_suffixes,omitempty"`
// IPRanges are published CIDRs (or single IPs) for bots that verify by
// address rather than reverse DNS -- typically AI agents (PerplexityBot,
// GPTBot, ClaudeBot). Membership is checked synchronously; no rDNS lookup.
IPRanges []string `yaml:"ip_ranges,omitempty" json:"ip_ranges,omitempty"`
}
// verifiedBotMinUALen rejects UA substrings short enough to match unrelated
// traffic ("bot", "go"). Real crawler tokens are longer.
const verifiedBotMinUALen = 4
// browserUATokens are substrings that appear in ordinary browser UAs. Entries
// keyed only on these tokens would match real users, so validation rejects
// them unless the substring also carries a crawler-specific token.
var browserUATokens = map[string]bool{
"mozilla": true, "applewebkit": true, "webkit": true, "gecko": true,
"chrome": true, "safari": true, "firefox": true, "edge": true,
"opera": true, "msie": true, "trident": true, "windows": true,
"macintosh": true, "linux": true, "android": true, "iphone": true,
"ipad": true, "x11": true, "mobile": true,
}
var crawlerUATokens = []string{
"bot", "crawler", "spider", "externalhit", "inspectiontool", "lighthouse",
}
// sharedHostingSuffixes are domains where reverse DNS is assigned to whoever
// rents the address, so an attacker can obtain a PTR (and forward A) under
// them. Allowlisting such a suffix would let any tenant spoof a crawler, so
// they are rejected. Matched as a suffix so subdomains are caught too.
var sharedHostingSuffixes = []string{
"amazonaws.com", "googleusercontent.com", "appspot.com", "run.app",
"cloudfront.net", "azurewebsites.net", "cloudapp.azure.com",
"herokuapp.com", "workers.dev", "pages.dev", "netlify.app",
"vercel.app", "ondigitalocean.app", "digitaloceanspaces.com",
"github.io", "gitlab.io", "fastly.net", "akamaitechnologies.com",
"akamai.net", "cloudflare.net", "colocrossing.com", "contabo.net",
"hetzner.com", "your-server.de", "ovh.net", "ip-linodeusercontent.com",
"linode.com", "vultrusercontent.com",
}
// commonPublicSuffixes are multi-label public suffixes that are not
// registrable on their own; a crawler can never legitimately be the whole
// suffix. Bare single-label TLDs are caught by the label-count check.
var commonPublicSuffixes = map[string]bool{
"co.uk": true, "org.uk": true, "gov.uk": true, "ac.uk": true,
"com.au": true, "net.au": true, "org.au": true, "co.jp": true,
"com.br": true, "com.cn": true, "co.in": true, "co.za": true,
"com.tr": true, "com.mx": true, "co.kr": true, "com.sg": true,
}
func validateVerifiedBots(cfg *Config) []ValidationResult {
var results []ValidationResult
seen := map[string]bool{}
seenUA := map[string]string{}
for i, b := range cfg.Reputation.VerifiedBots {
field := fmt.Sprintf("reputation.verified_bots[%d]", i)
name := strings.ToLower(strings.TrimSpace(b.Name))
if name == "" {
results = append(results, ValidationResult{"error", field + ".name", "verified bot name is required"})
continue
}
if seen[name] {
results = append(results, ValidationResult{"error", field + ".name",
fmt.Sprintf("duplicate verified bot name %q", name)})
}
seen[name] = true
hasUA := false
for _, raw := range b.UASubstrings {
s := strings.ToLower(strings.TrimSpace(raw))
if s == "" {
continue
}
hasUA = true
if len(s) < verifiedBotMinUALen {
results = append(results, ValidationResult{"error", field + ".ua_substrings",
fmt.Sprintf("UA substring %q is too short (min %d chars)", s, verifiedBotMinUALen)})
}
if browserUASubstringFootgun(s) {
results = append(results, ValidationResult{"error", field + ".ua_substrings",
fmt.Sprintf("UA substring %q matches ordinary browsers and would allowlist real users", s)})
}
if prev, ok := seenUA[s]; ok && prev != name {
results = append(results, ValidationResult{"error", field + ".ua_substrings",
fmt.Sprintf("UA substring %q is already used by verified bot %q", s, prev)})
}
for prev, prevName := range seenUA {
if prevName == name || prev == s {
continue
}
if strings.Contains(prev, s) || strings.Contains(s, prev) {
results = append(results, ValidationResult{"error", field + ".ua_substrings",
fmt.Sprintf("UA substring %q overlaps verified bot %q substring %q", s, prevName, prev)})
break
}
}
if _, ok := seenUA[s]; !ok {
seenUA[s] = name
}
}
if !hasUA {
results = append(results, ValidationResult{"error", field + ".ua_substrings",
"at least one ua_substring is required"})
}
hasSuffix := false
for _, raw := range b.RDNSSuffixes {
s := strings.TrimPrefix(strings.ToLower(strings.TrimSpace(raw)), ".")
if s == "" {
continue
}
hasSuffix = true
if msg := verifiedBotSuffixError(s); msg != "" {
results = append(results, ValidationResult{"error", field + ".rdns_suffixes", msg})
}
}
hasRange := false
for _, raw := range b.IPRanges {
s := strings.TrimSpace(raw)
if s == "" {
continue
}
hasRange = true
if msg := verifiedBotIPRangeError(s); msg != "" {
results = append(results, ValidationResult{"error", field + ".ip_ranges", msg})
}
}
// A bot needs at least one way to be confirmed: an rDNS suffix (for
// crawlers with forward-confirmable reverse DNS) or an IP range (for
// AI agents that publish address ranges instead of rDNS).
if !hasSuffix && !hasRange {
results = append(results, ValidationResult{"error", field,
"at least one rdns_suffix or ip_range is required"})
}
}
return results
}
// verifiedBotIPRangeError validates an operator-supplied CIDR or single IP.
// It rejects ranges too broad to be a crawler fleet and non-public space, so
// the allowlist cannot be turned into a blanket detection bypass.
func verifiedBotIPRangeError(s string) string {
n := netutil.ParseCIDROrIP(s)
if n == nil {
return fmt.Sprintf("ip_range %q is not a valid CIDR or IP", s)
}
ones, bits := n.Mask.Size()
if bits == 32 && ones < 16 {
return fmt.Sprintf("ip_range %q is too broad (minimum prefix /16 for IPv4)", s)
}
if bits == 128 && ones < 32 {
return fmt.Sprintf("ip_range %q is too broad (minimum prefix /32 for IPv6)", s)
}
if !netutil.IsPublicIP(n.IP) {
return fmt.Sprintf("ip_range %q is not a public address range", s)
}
return ""
}
func browserUASubstringFootgun(s string) bool {
if !containsKnownUAToken(s, browserUATokens) {
return false
}
for _, token := range crawlerUATokens {
if strings.Contains(s, token) {
return false
}
}
return true
}
func containsKnownUAToken(s string, tokens map[string]bool) bool {
for token := range tokens {
for start := 0; start < len(s); {
idx := strings.Index(s[start:], token)
if idx == -1 {
break
}
idx += start
end := idx + len(token)
if uaTokenBoundary(s, idx-1) && uaTokenBoundary(s, end) {
return true
}
start = idx + 1
}
}
return false
}
func uaTokenBoundary(s string, idx int) bool {
if idx < 0 || idx >= len(s) {
return true
}
c := s[idx]
if c >= 'a' && c <= 'z' {
return false
}
if c >= '0' && c <= '9' {
return false
}
return true
}
func validateVerifiedBotsConfig(cfg *Config) error {
for _, r := range validateVerifiedBots(cfg) {
if r.Level == "error" {
return fmt.Errorf("%s: %s", r.Field, r.Message)
}
}
return nil
}
func verifiedBotSuffixError(s string) string {
if strings.ContainsAny(s, " /:") {
return fmt.Sprintf("rdns_suffix %q is not a domain", s)
}
if len(s) > 253 {
return fmt.Sprintf("rdns_suffix %q is too long", s)
}
labels := strings.Split(s, ".")
if len(labels) < 2 {
return fmt.Sprintf("rdns_suffix %q must be a registrable domain (e.g. seranking.com), not a bare TLD", s)
}
for _, l := range labels {
if l == "" {
return fmt.Sprintf("rdns_suffix %q has an empty label", s)
}
if len(l) > 63 {
return fmt.Sprintf("rdns_suffix %q has a label longer than 63 characters", s)
}
if strings.HasPrefix(l, "-") || strings.HasSuffix(l, "-") {
return fmt.Sprintf("rdns_suffix %q has a label starting or ending with hyphen", s)
}
for _, r := range l {
if (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') || r == '-' {
continue
}
return fmt.Sprintf("rdns_suffix %q has invalid domain characters", s)
}
}
if commonPublicSuffixes[s] {
return fmt.Sprintf("rdns_suffix %q is a public suffix, not a registrable domain", s)
}
if _, err := publicsuffix.EffectiveTLDPlusOne(s); err != nil {
return fmt.Sprintf("rdns_suffix %q is a public suffix, not a registrable domain", s)
}
for _, bad := range sharedHostingSuffixes {
if s == bad || strings.HasSuffix(s, "."+bad) {
return fmt.Sprintf("rdns_suffix %q is shared hosting where reverse DNS is attacker-controlled; not allowed", s)
}
}
return ""
}
package config
import (
"fmt"
"time"
)
const DefaultBrowserSessionLifetime = "24h"
const DefaultBrowserSessionIdleTimeout = "30m"
// BrowserSessionDurations also accepts an un-defaulted configuration, as used
// by embedded callers. Explicit zero durations never disable session expiry.
func (c *Config) BrowserSessionDurations() (time.Duration, time.Duration, error) {
lifetimeText, idleText := c.WebUI.SessionLifetime, c.WebUI.SessionIdleTimeout
if lifetimeText == "" {
lifetimeText = DefaultBrowserSessionLifetime
}
if idleText == "" {
idleText = DefaultBrowserSessionIdleTimeout
}
lifetime, err := time.ParseDuration(lifetimeText)
if err != nil || lifetime < time.Second || lifetime > 30*24*time.Hour {
return 0, 0, fmt.Errorf("webui.session_lifetime must be between 1s and 720h")
}
idle, err := time.ParseDuration(idleText)
if err != nil || idle < time.Second || idle > lifetime {
return 0, 0, fmt.Errorf("webui.session_idle_timeout must be between 1s and session_lifetime")
}
return lifetime, idle, nil
}
package config
import (
"bytes"
"fmt"
"reflect"
"sort"
"strings"
"gopkg.in/yaml.v3"
)
// YAMLChange describes a single scalar, list, or map replacement to apply
// to a YAML document. Path is the dotted YAML key path from the document
// root. Value is the Go value to serialise; nil means YAML null.
type YAMLChange struct {
Path []string
Value interface{}
}
// lineIndex maps 1-based line numbers to byte offsets of their first byte.
// lineIndex[i] is the byte offset of line i+1 (i.e. lineIndex[0] == 0 == start of line 1).
type lineIndex []int
func buildLineIndex(data []byte) lineIndex {
idx := lineIndex{0} // line 1 starts at offset 0
for i, b := range data {
if b == '\n' {
idx = append(idx, i+1) // line starts after the newline
}
}
return idx
}
// offset returns the byte offset for a 1-based line and 1-based column.
func (li lineIndex) offset(line, col int) int {
if line < 1 || line > len(li) {
// line beyond end of file -- treat as EOF
if len(li) == 0 {
return 0
}
return li[len(li)-1]
}
return li[line-1] + col - 1
}
// lineEnd returns the byte offset of the '\n' at the end of the given 1-based line,
// or the length of data if the last line has no trailing newline.
func (li lineIndex) lineEnd(line int, data []byte) int {
if line < len(li) {
// next line starts at li[line]; the '\n' is at li[line]-1
return li[line] - 1
}
// last line
return len(data)
}
// lineStart returns the byte offset of the start of 1-based line.
func (li lineIndex) lineStart(line int) int {
if line < 1 {
return 0
}
if line > len(li) {
return li[len(li)-1]
}
return li[line-1]
}
// maxLine returns the maximum Line value in the node subtree.
func maxLine(n *yaml.Node) int {
m := n.Line
for _, child := range n.Content {
if cl := maxLine(child); cl > m {
m = cl
}
}
return m
}
// findNode walks the yaml.Node document tree and returns the key node and value node
// for the last segment of path, plus the parent mapping node.
// Returns (keyNode, valueNode, parentMapping, error).
func findNode(root *yaml.Node, path []string) (keyN, valN, parentMap *yaml.Node, err error) {
// root should be a DocumentNode; its Content[0] is the real root mapping.
cur := root
if cur.Kind == yaml.DocumentNode {
if len(cur.Content) == 0 {
return nil, nil, nil, fmt.Errorf("empty document")
}
cur = cur.Content[0]
}
for i, seg := range path {
if cur.Kind != yaml.MappingNode {
return nil, nil, nil, fmt.Errorf("path %v: segment %q: expected mapping, got kind %d", path[:i+1], seg, cur.Kind)
}
found := false
for j := 0; j+1 < len(cur.Content); j += 2 {
k := cur.Content[j]
v := cur.Content[j+1]
if k.Value == seg {
if i == len(path)-1 {
return k, v, cur, nil
}
cur = v
found = true
break
}
}
if !found {
// segment not found; return the parent mapping so caller can insert
if i == len(path)-1 {
return nil, nil, cur, nil // key missing, parent is cur
}
return nil, nil, nil, fmt.Errorf("path %v: segment %q not found", path[:i+1], seg)
}
}
return nil, nil, nil, fmt.Errorf("empty path")
}
// renderValueInline renders a scalar value to its YAML inline form.
// nil -> "null", bool -> "true"/"false", numbers via fmt, strings quoted if needed.
func renderValueInline(v interface{}) (string, error) {
if v == nil {
return "null", nil
}
// Use yaml.Marshal on a single-value map to get the marshalled scalar,
// then extract just the value part.
type wrapper struct {
V interface{} `yaml:"v"`
}
b, err := yaml.Marshal(wrapper{V: v})
if err != nil {
return "", err
}
// b looks like "v: VALUE\n"
s := strings.TrimPrefix(string(b), "v: ")
s = strings.TrimSuffix(s, "\n")
if strings.ContainsAny(s, "\n\r") {
return "", fmt.Errorf("value of type %T cannot be rendered inline", v)
}
return s, nil
}
// renderKeyInline renders a mapping key string in a YAML-safe form.
// If the key needs quoting (contains special characters, starts with special
// indicators, etc.) yaml.Marshal will add the necessary quotes.
// Returns an error if the key cannot be represented as a single-line YAML key
// (e.g. the key itself contains literal newlines).
func renderKeyInline(key string) (string, error) {
// Marshal a scalar node tagged !!str -- yaml.v3 will add quotes when the
// bare value would be misinterpreted (e.g. "1", ":", " x").
node := &yaml.Node{
Kind: yaml.ScalarNode,
Value: key,
Tag: "!!str",
}
b, err := yaml.Marshal(node)
if err != nil {
return "", err
}
// yaml.Marshal of a scalar node produces "value\n"
rendered := strings.TrimSuffix(string(b), "\n")
// Block literal (|) and folded (>) scalars span multiple lines and cannot
// serve as a simple inline mapping key (key: value on one line).
if strings.ContainsAny(rendered, "\n\r") || strings.HasPrefix(rendered, "|") || strings.HasPrefix(rendered, ">") {
return "", fmt.Errorf("key %q requires multi-line YAML representation and cannot be used as a simple mapping key", key)
}
return rendered, nil
}
// renderValueBlock renders a sequence or mapping value to a block of lines,
// indented at the given column (1-based). Returns lines like:
//
// " - a\n - b\n"
func renderValueBlock(v interface{}, indent int) (string, error) {
raw, err := yaml.Marshal(v)
if err != nil {
return "", err
}
// raw is like "- a\n- b\n" for a sequence, or "key: val\n" for a mapping.
prefix := strings.Repeat(" ", indent-1)
lines := strings.Split(string(raw), "\n")
var sb strings.Builder
for _, line := range lines {
if line == "" {
continue
}
sb.WriteString(prefix)
sb.WriteString(line)
sb.WriteByte('\n')
}
return sb.String(), nil
}
// isBlockValue returns true when the node needs block rendering (sequence or mapping
// that is not on the same line as its key).
func isBlockValue(keyN, valN *yaml.Node) bool {
if valN.Kind == yaml.ScalarNode {
return false
}
return valN.Line > keyN.Line
}
// splice replaces data[start:end] with replacement.
func splice(data []byte, start, end int, replacement []byte) []byte {
var buf bytes.Buffer
buf.Write(data[:start])
buf.Write(replacement)
buf.Write(data[end:])
return buf.Bytes()
}
// edit holds a resolved splice operation.
type edit struct {
start int
end int
replacement []byte
}
// YAMLEdit applies changes to data and returns the new document bytes.
// For every path that already exists, only the value span is rewritten
// at the same indent; untouched bytes (including all comments and
// whitespace) remain byte-identical. For a path that does not exist,
// a new key:value block is appended to the parent mapping at the parent's
// indent. Applies later edits first so earlier offsets remain valid.
func YAMLEdit(data []byte, changes []YAMLChange) ([]byte, error) {
if len(changes) == 0 {
return data, nil
}
seen := make(map[string]struct{}, len(changes))
for _, ch := range changes {
key := strings.Join(ch.Path, "\x00")
if _, dup := seen[key]; dup {
return nil, fmt.Errorf("yamledit: duplicate path %v", ch.Path)
}
seen[key] = struct{}{}
}
var root yaml.Node
if err := yaml.Unmarshal(data, &root); err != nil {
return nil, fmt.Errorf("yamledit: parse: %w", err)
}
li := buildLineIndex(data)
var edits []edit
for _, ch := range changes {
if len(ch.Path) == 0 {
return nil, fmt.Errorf("yamledit: empty path")
}
keyN, valN, parentMap, err := findNode(&root, ch.Path)
if err != nil {
return nil, fmt.Errorf("yamledit: %w", err)
}
if valN == nil {
// Key does not exist -- insert into parentMap.
var ed edit
ed, err = buildInsertEdit(data, li, parentMap, ch.Path[len(ch.Path)-1], ch.Value)
if err != nil {
return nil, fmt.Errorf("yamledit: insert %v: %w", ch.Path, err)
}
edits = append(edits, ed)
continue
}
// Key exists -- replace value span.
var ed edit
ed, err = buildReplaceEdit(data, li, keyN, valN, ch.Value)
if err != nil {
return nil, fmt.Errorf("yamledit: replace %v: %w", ch.Path, err)
}
edits = append(edits, ed)
}
// Sort by start offset descending so we splice end-to-start.
sort.Slice(edits, func(i, j int) bool {
return edits[i].start > edits[j].start
})
result := data
for _, ed := range edits {
result = splice(result, ed.start, ed.end, ed.replacement)
}
// Validate that the output is still parseable YAML. This catches edge cases
// where unusual input formats (complex key notation, etc.) produce invalid output.
var check yaml.Node
if err := yaml.Unmarshal(result, &check); err != nil {
return nil, fmt.Errorf("yamledit: output is not valid YAML: %w", err)
}
return result, nil
}
// buildReplaceEdit computes the splice for replacing an existing value node.
func buildReplaceEdit(data []byte, li lineIndex, keyN, valN *yaml.Node, value interface{}) (edit, error) {
if isBlockValue(keyN, valN) {
// Block sequence or mapping: value occupies one or more complete lines
// starting at valN.Line. Replace from the start of valN.Line to the
// end of the last line in the subtree.
lastLine := maxLine(valN)
start := li.lineStart(valN.Line)
end := li.lineEnd(lastLine, data)
if end < len(data) && data[end] == '\n' {
end++ // include the trailing newline so we replace whole lines
}
rendered, err := renderValueBlock(value, valN.Column)
if err != nil {
return edit{}, fmt.Errorf("path %v: %w", keyN.Value, err)
}
return edit{start: start, end: end, replacement: []byte(rendered)}, nil
}
// Flow / scalar: try inline first. If the new value is a sequence or
// mapping that cannot fit on one line (e.g. replacing `foo: []` with a
// multi-item list) fall back to block rendering and expand the span to
// cover the whole line, so the result is `foo:\n - a\n - b\n`.
start := li.offset(valN.Line, valN.Column)
end := li.lineEnd(valN.Line, data)
rendered, err := renderValueInline(value)
if err == nil {
return edit{start: start, end: end, replacement: []byte(rendered)}, nil
}
if !needsBlockFallback(value) {
return edit{}, fmt.Errorf("path %v: %w", keyN.Value, err)
}
// Replace from the key's column (beginning of `key:`) to end of the key's
// line, emitting `key:\n<block>`. Using keyN.Column keeps the existing
// indent. A trailing `# comment` on the original line is kept attached
// to the key line so operator annotations survive a multi-select save.
keyStart := li.offset(keyN.Line, keyN.Column)
keyEnd := li.lineEnd(keyN.Line, data)
renderedKey, kerr := renderKeyInline(keyN.Value)
if kerr != nil {
return edit{}, fmt.Errorf("path %v: render key: %w", keyN.Value, kerr)
}
block, berr := renderValueBlock(value, keyN.Column+2)
if berr != nil {
return edit{}, fmt.Errorf("path %v: %w", keyN.Value, berr)
}
trailingComment := extractInlineComment(data[keyStart:keyEnd])
return edit{start: keyStart, end: keyEnd, replacement: []byte(renderedKey + ":" + trailingComment + "\n" + strings.TrimRight(block, "\n"))}, nil
}
// extractInlineComment scans a single YAML line for a trailing `#` comment
// (one preceded by whitespace, outside string/rune literals) and returns it
// with its leading whitespace, e.g. " # keep empty". Returns "" if there
// is no comment. The scanner is deliberately simple — it handles the common
// case of a flow value followed by whitespace + `#`; it does not attempt to
// parse every YAML quoting variant (block scalars etc. do not reach this
// path).
func extractInlineComment(line []byte) string {
inDQ := false
inSQ := false
for i := 0; i < len(line); i++ {
c := line[i]
if c == '"' && !inSQ {
// Count backslashes before this quote to detect escapes.
bs := 0
for j := i - 1; j >= 0 && line[j] == '\\'; j-- {
bs++
}
if bs%2 == 0 {
inDQ = !inDQ
}
continue
}
if c == '\'' && !inDQ {
inSQ = !inSQ
continue
}
if inDQ || inSQ {
continue
}
if c != '#' {
continue
}
// A `#` only starts a comment when preceded by whitespace (or at the
// start of the line). Otherwise it is a legal scalar character.
if i == 0 || line[i-1] == ' ' || line[i-1] == '\t' {
// Include the whitespace run before `#` so output keeps spacing.
start := i
for start > 0 && (line[start-1] == ' ' || line[start-1] == '\t') {
start--
}
return string(line[start:])
}
}
return ""
}
// needsBlockFallback reports whether value is a sequence or mapping type
// that may exceed one line and therefore warrants switching the replacement
// from inline to block rendering.
func needsBlockFallback(value interface{}) bool {
switch value.(type) {
case []string, []interface{}, map[string]interface{}, map[interface{}]interface{}:
return true
}
return false
}
// buildInsertEdit computes the splice for appending a new key to a mapping node.
func buildInsertEdit(data []byte, li lineIndex, parentMap *yaml.Node, key string, value interface{}) (edit, error) {
// Find the insertion point: end of the last content line of parentMap.
// If parentMap has no content (empty mapping), insert after the mapping's own line.
insertLine := parentMap.Line
if len(parentMap.Content) > 0 {
// last child is parentMap.Content[len-1]
last := parentMap.Content[len(parentMap.Content)-1]
insertLine = maxLine(last)
}
insertOff := li.lineEnd(insertLine, data)
needsNewline := false
if insertOff < len(data) && data[insertOff] == '\n' {
insertOff++ // insert after the newline
} else if insertOff > 0 && data[insertOff-1] != '\n' {
// Last line has no trailing newline; we must add one before the new key.
needsNewline = true
}
// Match the indent of existing siblings rather than inferring from the
// parent node's own column, which in a nested mapping does not reflect
// child indentation.
siblingCol := parentMap.Column
if len(parentMap.Content) >= 1 {
siblingCol = parentMap.Content[0].Column
}
prefix := strings.Repeat(" ", siblingCol-1)
renderedKey, err := renderKeyInline(key)
if err != nil {
return edit{}, fmt.Errorf("render key %q: %w", key, err)
}
var rendered string
if valueNeedsBlockInsert(value) {
block, err := renderValueBlock(value, siblingCol+2)
if err != nil {
return edit{}, fmt.Errorf("path %v: %w", key, err)
}
rendered = prefix + renderedKey + ":\n" + block
} else {
inline, err := renderValueInline(value)
if err != nil {
return edit{}, fmt.Errorf("path %v: %w", key, err)
}
rendered = prefix + renderedKey + ": " + inline + "\n"
}
if needsNewline {
rendered = "\n" + rendered
}
return edit{start: insertOff, end: insertOff, replacement: []byte(rendered)}, nil
}
func valueNeedsBlockInsert(value interface{}) bool {
if value == nil {
return false
}
v := reflect.ValueOf(value)
switch v.Kind() {
case reflect.Array, reflect.Map, reflect.Struct:
return true
case reflect.Slice:
return v.Type().Elem().Kind() != reflect.Uint8
default:
return false
}
}
// Package contenttype classifies file content shared by the malware scanners.
package contenttype
import (
"bytes"
"path/filepath"
"strings"
)
var compressedArchiveMagics = [][]byte{
{'P', 'K', 0x03, 0x04}, // ZIP local file header
{'P', 'K', 0x05, 0x06}, // ZIP empty archive
{'P', 'K', 0x07, 0x08}, // ZIP spanned data descriptor
{0x1f, 0x8b}, // gzip
{'B', 'Z', 'h'}, // bzip2
{0xfd, '7', 'z', 'X', 'Z', 0x00}, // xz
{'7', 'z', 0xbc, 0xaf, 0x27, 0x1c}, // 7z
{'R', 'a', 'r', '!', 0x1a, 0x07, 0x00}, // RAR 4.x
{'R', 'a', 'r', '!', 0x1a, 0x07, 0x01, 0x00}, // RAR 5.x
}
// IsCompressedArchive reports whether data starts with a supported compressed
// archive signature. Tar and PHAR are intentionally excluded because their raw
// content is executable or otherwise meaningful to the signature engines.
func IsCompressedArchive(data []byte) bool {
if len(data) < 4 {
return false
}
for _, magic := range compressedArchiveMagics {
if bytes.HasPrefix(data, magic) {
return true
}
}
return false
}
// archiveExtensions are the names under which a compressed container is
// stored on a web host: plain archives plus the zip-based document, package
// and extension formats. Lowercase, leading dot.
var archiveExtensions = map[string]struct{}{
".zip": {}, ".zipx": {}, ".jar": {}, ".war": {}, ".ear": {}, ".apk": {},
".docx": {}, ".xlsx": {}, ".pptx": {}, ".odt": {}, ".ods": {}, ".odp": {},
".epub": {}, ".whl": {}, ".egg": {}, ".crx": {}, ".xpi": {}, ".ipa": {},
".nupkg": {}, ".vsix": {},
".gz": {}, ".tgz": {}, ".svgz": {},
".bz2": {}, ".tbz": {}, ".tbz2": {},
".xz": {}, ".txz": {},
".7z": {},
".rar": {},
}
// IsArchiveExt reports whether ext (leading dot, any case) names a compressed
// container format.
func IsArchiveExt(ext string) bool {
_, ok := archiveExtensions[strings.ToLower(ext)]
return ok
}
// IsArchiveFile reports whether both the name and the leading bytes identify a
// compressed container. Bytes alone are not enough: PHP echoes whatever
// precedes its open tag and runs the rest, so any interpreted file can start
// with archive magic and still execute. Only a file that also carries an
// archive extension is left to the extraction-time scan.
func IsArchiveFile(name string, data []byte) bool {
return IsArchiveExt(filepath.Ext(name)) && IsCompressedArchive(data)
}
package contenttype
import (
"bytes"
"encoding/binary"
"strings"
)
// imageExts are the file extensions served as raster images. The set exists
// for path-based dispatch only. Detection never trusts it: every content
// decision goes through ImageContainer, so renaming a payload cannot hide it.
var imageExts = map[string]bool{
".png": true,
".jpg": true,
".jpeg": true,
".jpe": true,
".gif": true,
".webp": true,
".ico": true,
".cur": true,
".bmp": true,
".tif": true,
".tiff": true,
}
// IsImageExt reports whether ext (with the leading dot, any case) names a
// raster image format.
func IsImageExt(ext string) bool {
return imageExts[strings.ToLower(ext)]
}
// ImageContainer identifies the raster image format data begins with and
// returns its display name. A hostile file can carry any extension, so the
// verdict comes from the leading bytes alone.
//
// The signatures are deliberately longer than the shortest unique prefix.
// "BM" and the ICO lead-in are two and four bytes wide, which ordinary text
// reaches by chance, so the structural fields that follow them are checked
// too: an image container that also holds PHP is a strong detection signal
// and a weak magic test would spend it on plain files.
func ImageContainer(data []byte) (string, bool) {
switch {
case bytes.HasPrefix(data, []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a}):
return "PNG", true
case bytes.HasPrefix(data, []byte{0xff, 0xd8, 0xff}):
return "JPEG", true
case bytes.HasPrefix(data, []byte("GIF87a")), bytes.HasPrefix(data, []byte("GIF89a")):
return "GIF", true
case len(data) >= 12 && bytes.HasPrefix(data, []byte("RIFF")) && bytes.Equal(data[8:12], []byte("WEBP")):
return "WebP", true
case isICO(data):
return "ICO", true
case isBMP(data):
return "BMP", true
case bytes.HasPrefix(data, []byte{0x49, 0x49, 0x2a, 0x00}), bytes.HasPrefix(data, []byte{0x4d, 0x4d, 0x00, 0x2a}):
return "TIFF", true
}
return "", false
}
// isICO checks the ICONDIR header: two reserved zero bytes, a type of 1 (icon)
// or 2 (cursor), and at least one directory entry.
func isICO(data []byte) bool {
if len(data) < 6 {
return false
}
if data[0] != 0 || data[1] != 0 {
return false
}
imageType := binary.LittleEndian.Uint16(data[2:4])
if imageType != 1 && imageType != 2 {
return false
}
return binary.LittleEndian.Uint16(data[4:6]) > 0
}
// isBMP checks the BITMAPFILEHEADER: the "BM" tag, the two reserved words
// which the format requires to be zero, and a pixel-data offset that lands
// past the header.
func isBMP(data []byte) bool {
if len(data) < 14 || data[0] != 'B' || data[1] != 'M' {
return false
}
if binary.LittleEndian.Uint32(data[6:10]) != 0 {
return false
}
return binary.LittleEndian.Uint32(data[10:14]) >= 14
}
package contenttype
import "strings"
// executablePHPExtensions are the extensions a stock PHP-capable web server
// (Apache mod_php / PHP-FPM via EasyApache4, LiteSpeed LSAPI, Nginx + php-fpm)
// routes to the PHP interpreter by default. Any file with one of these names
// can execute PHP, so a content scan that skipped them would let a webshell
// hide behind a non-".php" name. ".phps" is deliberately excluded: the stock
// handler renders it as highlighted source, it does not execute. Lowercase,
// leading dot.
var executablePHPExtensions = []string{
".php", ".php2", ".php3", ".php4", ".php5", ".php6", ".php7", ".php8",
".phtml", ".pht",
}
// IsExecutablePHPExt reports whether ext (leading dot, any case) is one a
// stock PHP handler executes. Rule sets are written against ".php"; every
// extension here must be matched against the same rules or a payload hides
// behind the name.
func IsExecutablePHPExt(ext string) bool {
ext = strings.ToLower(ext)
for _, e := range executablePHPExtensions {
if ext == e {
return true
}
}
return false
}
// IsExecutablePHPName reports whether a (lowercased) filename has an extension
// that a stock PHP handler executes. Shared by the realtime fanotify path, the
// periodic content scanners and the rule engines so none of them drift apart.
// It is a coarse, default-deny gate for content analysis only; per-directory
// .htaccess handler remappings are layered on top by the checks package.
func IsExecutablePHPName(nameLower string) bool {
for _, ext := range executablePHPExtensions {
if strings.HasSuffix(nameLower, ext) {
return true
}
}
return false
}
// IsPHPSourceName reports whether a file should receive PHP content analysis.
// It deliberately includes .phps even though IsExecutablePHPName does not:
// stock handlers render .phps as source, but the bytes can still hold a staged
// payload that becomes executable after a rename.
func IsPHPSourceName(nameLower string) bool {
return IsExecutablePHPName(nameLower) || strings.HasSuffix(nameLower, ".phps")
}
package contenttype
import "regexp"
// phpOpenTag matches the two PHP openers a web server executes under a stock
// configuration. The bare "<?" short tag is off by default and is also the
// XML and SVG declaration opener, so it is not treated as PHP here.
var phpOpenTag = regexp.MustCompile(`(?i)<\?(?:php(?:[ \t\r\n]|$)|=)`)
// HasPHPOpenTag reports whether data contains a PHP opening tag anywhere. It
// establishes PHP context for content whose name says nothing useful: an
// image carrying a payload, or a rule helper handed a file of unknown type.
func HasPHPOpenTag(data []byte) bool {
return phpOpenTag.Match(data)
}
// Package control defines the wire protocol between the CSM daemon and
// its local command-line client. The daemon listens on a Unix socket and
// accepts one line-framed JSON request per connection, replying with one
// line-framed JSON response. All types defined here are stable enough to
// be imported by both sides.
package control
import (
"encoding/json"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/incident"
"github.com/pidginhost/csm/internal/store"
)
// DefaultSocketPath is the Unix socket the daemon binds and the client
// dials. It sits next to pam.sock under /var/run/csm so both follow the
// same directory permissions established by the daemon at startup.
const DefaultSocketPath = "/var/run/csm/control.sock"
// Command names. Use constants so typos fail at compile time.
const (
CmdTierRun = "tier.run"
CmdStatus = "status"
CmdHistoryRead = "history.read"
CmdRulesReload = "rules.reload"
CmdGeoIPReload = "geoip.reload"
CmdBotRangesReload = "botranges.reload"
// Phase 2 additions.
CmdBaseline = "baseline"
CmdFirewallStatus = "firewall.status"
CmdFirewallPorts = "firewall.ports"
CmdFirewallGrep = "firewall.grep"
CmdFirewallAudit = "firewall.audit"
CmdFirewallBlock = "firewall.block"
CmdFirewallUnblock = "firewall.unblock"
CmdFirewallAllow = "firewall.allow"
CmdFirewallRemoveAllow = "firewall.remove_allow"
CmdFirewallAllowPort = "firewall.allow_port"
CmdFirewallRemovePort = "firewall.remove_port"
CmdFirewallTempBan = "firewall.tempban"
CmdFirewallTempAllow = "firewall.tempallow"
CmdFirewallDenySubnet = "firewall.deny_subnet"
CmdFirewallRemoveSubnet = "firewall.remove_subnet"
CmdFirewallDenyFile = "firewall.deny_file"
CmdFirewallAllowFile = "firewall.allow_file"
CmdFirewallFlush = "firewall.flush"
CmdFirewallRestart = "firewall.restart"
CmdFirewallApplyConfirmed = "firewall.apply_confirmed"
CmdFirewallConfirm = "firewall.confirm"
// Durable firewall actions. Actions lists what recovery could not settle,
// and ActionResolve records the outcome an operator established by hand.
CmdFirewallActions = "firewall.actions"
CmdFirewallActionResolve = "firewall.action_resolve"
CmdFirewallRollbackStatus = "firewall.rollback_status"
CmdFirewallRollbackRevert = "firewall.rollback_revert"
CmdFirewallRollbackOK = "firewall.rollback_confirm"
// Backup / restore. Export runs through the daemon (live read of
// bbolt). Import does not -- it requires a stopped daemon, so the
// CLI opens the archive and target paths directly.
CmdStoreExport = "store.export"
// Audit-log backfill. Streams every finding from the history
// bucket newer than the supplied cutoff so SIEM operators can
// seed their pipeline with prior findings the first time they
// enable audit_log:.
CmdHistorySince = "history.since"
CmdPHPRelayStatus = "phprelay.status"
CmdPHPRelayIgnoreScript = "phprelay.ignore_script"
CmdPHPRelayUnignore = "phprelay.unignore"
CmdPHPRelayIgnoreList = "phprelay.ignore_list"
CmdPHPRelayDryRun = "phprelay.dry_run"
CmdPHPRelayThaw = "phprelay.thaw"
// Clears one address's accumulated local threat-scoring state. Kept
// separate from the firewall commands: it changes no block, allow or
// whitelist entry, only what local_threat_score reads.
CmdThreatForget = "threat.forget"
// Phase 2 incident correlation.
CmdIncidentsList = "incidents.list"
CmdIncidentsShow = "incidents.show"
CmdIncidentsStatus = "incidents.status"
CmdIncidentsBulkStatus = "incidents.bulk_status"
// Full-scan job commands (uncapped-full-scan feature).
CmdScanEnqueue = "scan.enqueue"
CmdScanStatus = "scan.status"
CmdScanReport = "scan.report"
CmdScanCancel = "scan.cancel"
)
// Request is the single JSON object the client sends per connection.
// Args is a raw message so each command defines its own typed payload
// without forcing the transport layer to know about every command.
type Request struct {
Cmd string `json:"cmd"`
Args json.RawMessage `json:"args,omitempty"`
}
// Response is the single JSON object the daemon returns per connection.
// Ok=false means the request did not run to completion; inspect Error.
// Ok=true means the handler ran and Result holds the typed payload.
type Response struct {
OK bool `json:"ok"`
Result json.RawMessage `json:"result,omitempty"`
Error string `json:"error,omitempty"`
}
// TierRunArgs carries parameters for CmdTierRun.
// Tier is one of "critical", "deep", or "all". Alerts=true feeds the
// daemon's normal alert pipeline; false means "scan and report counts
// only" (used by the CLI's dry-run `check-*` commands when they migrate).
type TierRunArgs struct {
Tier string `json:"tier"`
Alerts bool `json:"alerts"`
}
// TierRunResult summarises what the tier run produced. Findings is only
// populated when TierRunArgs.Alerts was false (the dry-run `csm check*`
// path) so live tier runs do not pay the marshalling cost for the
// no-op case.
type TierRunResult struct {
Findings int `json:"findings"`
NewFindings int `json:"new_findings"`
ElapsedMs int64 `json:"elapsed_ms"`
FindingList []alert.Finding `json:"finding_list,omitempty"`
}
// StatusResult mirrors what `csm status` historically printed plus the
// extended health.Snapshot fields used by `csm status --json`. The
// pre-existing fields stay in place; any new client should prefer
// reading Snapshot directly.
type StatusResult struct {
Version string `json:"version"`
UptimeSec int64 `json:"uptime_sec"`
LatestScanTime string `json:"latest_scan_time,omitempty"`
LatestFindings int `json:"latest_findings"`
HistoryCount int `json:"history_count"`
DroppedAlerts int64 `json:"dropped_alerts"`
// Snapshot is the full health view added in v2.12.0. Older clients
// that only knew the six legacy fields above will simply ignore it.
Snapshot *health.Snapshot `json:"snapshot,omitempty"`
}
// HistoryReadArgs carries parameters for CmdHistoryRead.
type HistoryReadArgs struct {
Limit int `json:"limit"`
Offset int `json:"offset"`
}
// HistoryReadResult holds a page of historical findings, newest-first.
type HistoryReadResult struct {
Findings []alert.Finding `json:"findings"`
Total int `json:"total"`
}
// BaselineArgs carries parameters for CmdBaseline.
// Confirm mirrors the CLI's --confirm flag: required when existing
// history would be wiped.
type BaselineArgs struct {
Confirm bool `json:"confirm"`
}
// BaselineResult reports what a fresh baseline wrote.
type BaselineResult struct {
Findings int `json:"findings"`
HistoryCleared int `json:"history_cleared"`
BinaryHash string `json:"binary_hash"`
ConfigHash string `json:"config_hash"`
// NeedsConfirm=true means the daemon refused because Confirm was
// false and HistoryCleared would have been non-zero. The other
// fields are populated so the CLI can print the same warning as
// today without a second round trip.
NeedsConfirm bool `json:"needs_confirm,omitempty"`
}
// FirewallIPArgs is shared by block / unblock / allow / remove-style
// commands that take a single IP and an optional reason. Timeout is
// consumed by tempban/tempallow and MUST parse via time.ParseDuration
// (e.g. "24h", "1h30m", "5m"); empty means permanent. Invalid strings
// are rejected by the handler with a parse error.
type FirewallIPArgs struct {
IP string `json:"ip"`
Reason string `json:"reason,omitempty"`
Timeout string `json:"timeout,omitempty"`
}
// FirewallPortArgs covers allow-port / remove-port.
type FirewallPortArgs struct {
IP string `json:"ip"`
Port int `json:"port"`
Proto string `json:"proto,omitempty"` // "tcp" or "udp"; empty = tcp
Reason string `json:"reason,omitempty"`
}
// FirewallSubnetArgs covers deny-subnet / remove-subnet.
type FirewallSubnetArgs struct {
CIDR string `json:"cidr"`
Reason string `json:"reason,omitempty"`
}
// FirewallFileArgs carries a batch of IPs for deny-file / allow-file.
// The client reads the file locally and sends the contents over the
// socket so the daemon does not need to open arbitrary paths.
type FirewallFileArgs struct {
IPs []string `json:"ips"`
Reason string `json:"reason,omitempty"`
}
// FirewallGrepArgs matches the CLI's positional pattern.
type FirewallGrepArgs struct {
Pattern string `json:"pattern"`
}
// FirewallAuditArgs matches the CLI's optional limit. Limit=0 means
// "use the handler default" (currently 50 lines, matching the old CLI).
type FirewallAuditArgs struct {
Limit int `json:"limit"`
}
// FirewallActionResolveArgs carries an operator decision about one durable
// firewall action. Outcome is "applied" or "rejected", in the operator's own
// terms; Note records the evidence they went on.
type FirewallActionResolveArgs struct {
ID string `json:"id"`
Outcome string `json:"outcome"`
Note string `json:"note,omitempty"`
}
// FirewallApplyConfirmedArgs mirrors the CLI's minutes argument.
type FirewallApplyConfirmedArgs struct {
Minutes int `json:"minutes"`
}
// FirewallAckResult is the minimal ack returned by mutating firewall
// commands that do not need to report state back (block, allow, etc).
// Message is a short human-readable string the CLI can print verbatim.
type FirewallAckResult struct {
Message string `json:"message"`
}
// ThreatForgetResult reports what clearing an address's scoring state
// actually removed. Found distinguishes "cleared a stale record" from
// "there was nothing to clear", which the operator cannot otherwise tell
// apart and which decides whether the alert will stop.
// If legacy spellings created multiple records for the same address, Events
// is their total event count and Score is the highest removed record's score.
type ThreatForgetResult struct {
IP string `json:"ip"`
Found bool `json:"found"`
Score int `json:"score"`
Events int `json:"events"`
Message string `json:"message"`
}
// FirewallRollbackStatus reports the pending tentative-apply state for
// the firewall settings section. Pending=false means nothing in flight;
// the other fields are zero in that case.
type FirewallRollbackStatus struct {
Pending bool `json:"pending"`
AppliedAtRFC3339 string `json:"applied_at,omitempty"`
ExpiresAtRFC3339 string `json:"expires_at,omitempty"`
SecondsRemaining int64 `json:"seconds_remaining,omitempty"`
AppliedBy string `json:"applied_by,omitempty"`
PrevHash string `json:"prev_hash,omitempty"`
NewHash string `json:"new_hash,omitempty"`
}
// FirewallStatusResult mirrors what `csm firewall status` prints.
// Fields match the CLI output one-to-one so the client can format
// identically without calling firewall.LoadState itself.
type FirewallStatusResult struct {
Enabled bool `json:"enabled"`
TCPIn []string `json:"tcp_in"`
TCPOut []string `json:"tcp_out"`
UDPIn []string `json:"udp_in"`
UDPOut []string `json:"udp_out"`
Restricted []string `json:"restricted"`
PassiveFTPStart int `json:"passive_ftp_start"`
PassiveFTPEnd int `json:"passive_ftp_end"`
TCPOutAllow []string `json:"tcp_out_allow,omitempty"`
InfraIPCount int `json:"infra_ip_count"`
BlockedCount int `json:"blocked_count"`
BlockedNetCount int `json:"blocked_net_count"`
AllowedCount int `json:"allowed_count"`
SYNFlood bool `json:"syn_flood"`
ConnRateLimit int `json:"conn_rate_limit"`
LogDropped bool `json:"log_dropped"`
LogRate int `json:"log_rate"`
RecentBlocked []FirewallBlockedEntry `json:"recent_blocked,omitempty"`
}
// FirewallBlockedEntry is one entry in FirewallStatusResult.RecentBlocked.
type FirewallBlockedEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
BlockedAt string `json:"blocked_at"`
ExpiresAt string `json:"expires_at,omitempty"`
}
// FirewallListResult is returned by ports / grep / audit: a freeform
// line-based payload the CLI prints verbatim.
type FirewallListResult struct {
Lines []string `json:"lines"`
}
// StoreExportArgs configures CmdStoreExport. DstPath must be an
// absolute path the daemon can write to; the CLI does not pre-create
// it.
type StoreExportArgs struct {
DstPath string `json:"dst_path"`
// Stage asks the daemon to write the archive under its own state
// directory (the one place its sandbox guarantees writable) instead of
// DstPath; the result's Path names the staged file for the CLI to move.
Stage bool `json:"stage,omitempty"`
}
// HistorySinceArgs configures CmdHistorySince. Since is RFC 3339; an
// empty value yields nothing (caller mistake, not "everything", to
// avoid accidental whole-DB dumps over the socket).
type HistorySinceArgs struct {
Since string `json:"since"`
}
// HistorySinceResult holds every finding newer than Since, oldest
// first so a downstream JSONL writer produces chronologically-sorted
// lines.
type HistorySinceResult struct {
Findings []alert.Finding `json:"findings"`
}
// StoreExportResult summarises a successful export.
type StoreExportResult struct {
Path string `json:"path"`
Bytes int64 `json:"bytes"`
ArchiveSHA256 string `json:"archive_sha256"`
BboltSHA256 string `json:"bbolt_sha256"`
}
// PHPRelayStatusRequest carries no parameters.
type PHPRelayStatusRequest struct{}
// PHPRelayStatusResponse summarises the running detector state.
type PHPRelayStatusResponse struct {
Enabled bool `json:"enabled"`
Platform string `json:"platform"`
DryRun bool `json:"dry_run"`
EffectiveAccountLimit int `json:"effective_account_limit"`
ScriptsTracked int `json:"scripts_tracked"`
IPsTracked int `json:"ips_tracked"`
AccountsTracked int `json:"accounts_tracked"`
MsgIDIndexSize int `json:"msgid_index_size"`
IgnoresActive int `json:"ignores_active"`
RecentFindings map[string]int `json:"recent_findings"` // path -> count last 1h
}
type PHPRelayIgnoreScriptRequest struct {
ScriptKey string `json:"script_key"`
ForHours int `json:"for_hours"` // 0 -> default 7d (168h)
Persist bool `json:"persist"`
Reason string `json:"reason"`
AddedBy string `json:"added_by"`
}
type PHPRelayIgnoreScriptResponse struct {
ExpiresAt time.Time `json:"expires_at"`
}
type PHPRelayUnignoreRequest struct {
ScriptKey string `json:"script_key"`
Persist bool `json:"persist"`
}
type PHPRelayIgnoreEntry struct {
ScriptKey string `json:"script_key"`
ExpiresAt time.Time `json:"expires_at"`
AddedBy string `json:"added_by"`
Reason string `json:"reason"`
}
type PHPRelayIgnoreListResponse struct {
Entries []PHPRelayIgnoreEntry `json:"entries"`
}
type PHPRelayDryRunRequest struct {
Mode string `json:"mode"` // "on" | "off" | "reset"
Persist bool `json:"persist"`
}
type PHPRelayDryRunResponse struct {
Effective bool `json:"effective"`
Source string `json:"source"` // "runtime" | "bbolt" | "csm.yaml"
}
type PHPRelayThawRequest struct {
MsgID string `json:"msg_id"`
By string `json:"by"`
}
type PHPRelayThawResponse struct {
Stderr string `json:"stderr,omitempty"`
}
// IncidentShowArgs targets a single incident by id.
type IncidentShowArgs struct {
ID string `json:"id"`
}
// IncidentListArgs pages CmdIncidentsList. All=true returns every
// matching incident and should be reserved for offline export-style use.
type IncidentListArgs struct {
Limit int `json:"limit,omitempty"`
Offset int `json:"offset,omitempty"`
Status string `json:"status,omitempty"` // all / active / open / contained / resolved / dismissed
All bool `json:"all,omitempty"`
}
// IncidentListResult is the bounded list envelope returned by
// CmdIncidentsList. Total is computed before paging so the CLI can show
// the next offset without a second command.
type IncidentListResult struct {
Items []incident.Incident `json:"items"`
Total int `json:"total"`
Offset int `json:"offset"`
Limit int `json:"limit"`
Status string `json:"status"`
}
// IncidentStatusArgs transitions an incident's status.
type IncidentStatusArgs struct {
ID string `json:"id"`
Status string `json:"status"` // open / contained / resolved / dismissed
Details string `json:"details,omitempty"`
}
// IncidentBulkStatusArgs previews or applies the same status transition
// to stale incidents matching the supplied exact filters.
type IncidentBulkStatusArgs struct {
Status string `json:"status,omitempty"` // active / open / contained
To string `json:"to,omitempty"` // resolved / dismissed
OlderThanSeconds int64 `json:"older_than_seconds,omitempty"`
LastSeenBefore time.Time `json:"last_seen_before,omitempty"`
Kind string `json:"kind,omitempty"`
Domain string `json:"domain,omitempty"`
Account string `json:"account,omitempty"`
Mailbox string `json:"mailbox,omitempty"`
Limit int `json:"limit,omitempty"`
Apply bool `json:"apply,omitempty"`
Confirm bool `json:"confirm,omitempty"`
Details string `json:"details,omitempty"`
}
// IncidentBulkStatusResult reports the full match count plus the
// bounded set of incidents that were previewed or changed.
type IncidentBulkStatusResult struct {
DryRun bool `json:"dry_run"`
Matched int `json:"matched"`
Updated int `json:"updated"`
Limit int `json:"limit"`
Status string `json:"status"`
To string `json:"to"`
OlderThanSeconds int64 `json:"older_than_seconds,omitempty"`
LastSeenBefore time.Time `json:"last_seen_before,omitempty"`
Items []incident.BulkStatusItem `json:"items"`
}
// ScanEnqueueRequest carries parameters for CmdScanEnqueue.
// Phase 1 accepts only Scope="account". Target is the cPanel username to scan.
// RespectIgnores gates whether cfg.Suppressions.IgnorePaths apply during the scan.
// Quarantine records the quarantine intent for use by a later phase; no live
// auto-response is wired in Phase 1.
type ScanEnqueueRequest struct {
Scope string `json:"scope"`
Target string `json:"target"`
RespectIgnores bool `json:"respect_ignores"`
Quarantine bool `json:"quarantine"`
}
// ValidScanAccountTarget reports whether target is a single account name, not
// a path. Full-scan handlers join this value under /home, so slashes, dot-only
// names, spaces, and control bytes must be rejected at the protocol boundary.
func ValidScanAccountTarget(target string) bool {
if target == "" || target == "." || target == ".." {
return false
}
for i := 0; i < len(target); i++ {
b := target[i]
if b >= 'a' && b <= 'z' ||
b >= 'A' && b <= 'Z' ||
b >= '0' && b <= '9' ||
b == '_' || b == '-' || b == '.' {
continue
}
return false
}
return true
}
// ScanEnqueueResponse carries the job ID and initial state ("queued").
type ScanEnqueueResponse struct {
JobID string `json:"job_id"`
State string `json:"state"`
}
// ScanStatusRequest carries an optional job ID for CmdScanStatus.
// An empty JobID requests the full job list; a concrete ID returns that job only.
type ScanStatusRequest struct {
JobID string `json:"job_id,omitempty"`
}
// ScanStatusResponse carries a single job or the full list depending on the request.
type ScanStatusResponse struct {
Job *store.ScanJobRecord `json:"job,omitempty"`
Jobs []store.ScanJobRecord `json:"jobs,omitempty"`
}
// ScanReportRequest carries parameters for CmdScanReport.
// Offset and Limit follow the usual slice semantics; Limit=0 returns all findings.
type ScanReportRequest struct {
JobID string `json:"job_id"`
Offset int `json:"offset,omitempty"`
Limit int `json:"limit,omitempty"`
}
// ScanReportResponse carries the job record, the requested page of findings, and
// the total finding count for the job (before paging).
type ScanReportResponse struct {
Job store.ScanJobRecord `json:"job"`
Findings []alert.Finding `json:"findings"`
Total int `json:"total"`
}
// ScanCancelRequest carries the job ID for CmdScanCancel.
type ScanCancelRequest struct {
JobID string `json:"job_id"`
}
// ScanCancelResponse reports the job ID and its state after cancellation.
type ScanCancelResponse struct {
JobID string `json:"job_id"`
State string `json:"state"`
}
package corpusgate
import (
"crypto/sha256"
"encoding/hex"
"fmt"
"net/url"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/cms"
)
// ManifestVersion is the only accepted manifest version. Version 2 made the
// per-source cms field mandatory and added the pending list.
const ManifestVersion = 2
// PendingCMS records a supported CMS with no pinned clean source yet and the
// roadmap item that blocks it. It is coverage metadata only: nothing at scan
// time reads it, and it confers no runtime trust.
type PendingCMS struct {
CMS string `json:"cms"`
Reason string `json:"reason"`
}
// Validate checks the manifest's own metadata: the version, every source's
// fields, and that every supported CMS is either sourced or explicitly
// pending, never both. The supported set always comes from internal/cms so a
// caller cannot validate against a smaller one. Archive authentication,
// extraction confinement and inventory checks stay in Prepare.
func (m Manifest) Validate() error {
if m.Version != ManifestVersion {
return fmt.Errorf("corpus manifest version %d is not supported; version %d with a cms field on every source is required", m.Version, ManifestVersion)
}
if len(m.Sources) == 0 {
return fmt.Errorf("corpus manifest has no source")
}
sourced := make(map[cms.Kind]bool)
seenID := make(map[string]bool, len(m.Sources))
for _, s := range m.Sources {
if err := s.validate(); err != nil {
return err
}
if seenID[s.ID] {
return fmt.Errorf("source %q: id declared twice", s.ID)
}
seenID[s.ID] = true
kind, ok := cms.Parse(s.CMS)
if !ok {
return fmt.Errorf("source %q: cms %q is not a supported kind", s.ID, s.CMS)
}
sourced[kind] = true
}
pending := make(map[cms.Kind]bool)
for _, p := range m.Pending {
kind, ok := cms.Parse(p.CMS)
if !ok {
return fmt.Errorf("pending cms %q is not a supported kind", p.CMS)
}
if pending[kind] {
return fmt.Errorf("pending cms %q is declared twice", p.CMS)
}
if strings.TrimSpace(p.Reason) == "" {
return fmt.Errorf("pending cms %q: reason is required", p.CMS)
}
if sourced[kind] {
return fmt.Errorf("cms %q is both sourced and pending", p.CMS)
}
pending[kind] = true
}
for _, d := range cms.All() {
if !sourced[d.Kind] && !pending[d.Kind] {
return fmt.Errorf("cms %q has neither a pinned source nor a pending entry", d.Kind)
}
}
return nil
}
func (s Source) validate() error {
if s.ID == "" || strings.ContainsAny(s.ID, "/\\\x00") || s.ID == "." || s.ID == ".." {
return fmt.Errorf("source %q: id must be a non-empty path-safe name", s.ID)
}
if s.Version == "" || strings.ContainsAny(s.Version, "/\\\x00") {
return fmt.Errorf("source %q: version must be a non-empty path-safe string", s.ID)
}
u, err := url.Parse(s.URL)
if err != nil || u.Scheme != "https" || u.Host == "" {
return fmt.Errorf("source %q: url must be https with a host", s.ID)
}
digest, err := hex.DecodeString(s.SHA256)
if err != nil || len(digest) != sha256.Size {
return fmt.Errorf("source %q: sha256 must be 64 hex characters", s.ID)
}
if s.License == "" {
return fmt.Errorf("source %q: license is required", s.ID)
}
if !filepath.IsLocal(s.LicenseFile) {
return fmt.Errorf("source %q: license_file must be a local path inside the archive", s.ID)
}
if s.Files < 1 || s.Files > maxSourceFiles {
return fmt.Errorf("source %q: files must be between 1 and %d", s.ID, maxSourceFiles)
}
if s.CMS == "" {
return fmt.Errorf("source %q: cms field is required in manifest version %d", s.ID, ManifestVersion)
}
return nil
}
// maxSourceFiles bounds a pinned archive's inventory.
const maxSourceFiles = 30000
// Package corpusgate provisions pinned clean applications for detector tests.
package corpusgate
import (
"archive/zip"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"sort"
"strings"
"time"
)
type Manifest struct {
Version int `json:"version"`
Sources []Source `json:"sources"`
// Pending lists supported CMSs with no pinned source yet. See Validate.
Pending []PendingCMS `json:"pending,omitempty"`
}
type Source struct {
ID string `json:"id"`
// CMS names the supported kind this archive belongs to; a plugin or
// theme names the CMS it runs on.
CMS string `json:"cms"`
Version string `json:"version"`
URL string `json:"url"`
SHA256 string `json:"sha256"`
License string `json:"license"`
LicenseFile string `json:"license_file"`
Files int `json:"files"`
}
type File struct {
Path string `json:"path"`
SHA256 string `json:"sha256"`
Bytes int64 `json:"bytes"`
}
const maxArchive = 512 << 20
const maxExpanded = 1 << 30
// Prepare refuses to reuse extracted trees: only authenticated archive caches
// survive runs, so removed vendor files cannot inflate the next inventory.
func Prepare(ctx context.Context, manifest Manifest, cache, destination string) ([]File, error) {
// Validate before any directory, cache or network side effect so a
// drifted manifest is refused without leaving artifacts behind.
if err := manifest.Validate(); err != nil {
return nil, err
}
if err := os.Mkdir(destination, 0700); err != nil {
return nil, fmt.Errorf("new corpus directory: %w", err)
}
if err := os.MkdirAll(cache, 0700); err != nil {
return nil, err
}
cacheRoot, err := os.OpenRoot(cache)
if err != nil {
return nil, err
}
defer func() { _ = cacheRoot.Close() }()
var inventory []File
for _, s := range manifest.Sources {
archive := filepath.Join(cache, s.ID+"-"+s.Version+".zip")
if _, err := os.Stat(archive); os.IsNotExist(err) {
if downloadErr := download(ctx, s.URL, archive); downloadErr != nil {
return nil, downloadErr
}
} else if err != nil {
return nil, err
}
f, err := cacheRoot.Open(filepath.Base(archive))
if err != nil {
return nil, err
}
h := sha256.New()
n, copyErr := io.Copy(h, io.LimitReader(f, maxArchive+1))
if copyErr != nil || n > maxArchive || hex.EncodeToString(h.Sum(nil)) != s.SHA256 {
_ = f.Close()
return nil, fmt.Errorf("archive checksum/size/read failure for %s", s.ID)
}
rows, err := extract(f, n, destination, s)
closeErr := f.Close()
if err == nil {
err = closeErr
}
if err != nil {
return nil, fmt.Errorf("extract %s: %w", s.ID, err)
}
inventory = append(inventory, rows...)
}
sort.Slice(inventory, func(i, j int) bool { return inventory[i].Path < inventory[j].Path })
return inventory, nil
}
func download(ctx context.Context, address, destination string) error {
client := &http.Client{Timeout: 3 * time.Minute, CheckRedirect: func(req *http.Request, via []*http.Request) error {
if req.URL.Scheme != "https" || len(via) >= 10 {
return fmt.Errorf("unsafe corpus redirect")
}
return nil
}}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, address, nil)
if err != nil {
return err
}
response, err := client.Do(req)
if err != nil {
return err
}
defer response.Body.Close()
if response.StatusCode != http.StatusOK {
return fmt.Errorf("corpus download HTTP %d", response.StatusCode)
}
file, err := os.CreateTemp(filepath.Dir(destination), ".corpus-download-*")
if err != nil {
return err
}
defer os.Remove(file.Name())
n, copyErr := io.Copy(file, io.LimitReader(response.Body, maxArchive+1))
closeErr := file.Close()
if copyErr != nil {
return copyErr
}
if closeErr != nil {
return closeErr
}
if n > maxArchive {
return fmt.Errorf("corpus archive too large")
}
return os.Rename(file.Name(), destination)
}
func extract(archive io.ReaderAt, size int64, destination string, source Source) ([]File, error) {
z, err := zip.NewReader(archive, size)
if err != nil {
return nil, err
}
dir := filepath.Join(destination, source.ID)
if err = os.Mkdir(dir, 0700); err != nil {
return nil, err
}
root, err := os.OpenRoot(dir)
if err != nil {
return nil, err
}
defer func() { _ = root.Close() }()
var rows []File
var expanded int64
license := false
for _, entry := range z.File {
name := strings.TrimSuffix(entry.Name, "/")
if !filepath.IsLocal(name) || strings.Contains(name, "\\") {
return nil, fmt.Errorf("unsafe archive name %q", name)
}
if entry.FileInfo().IsDir() {
if err := root.MkdirAll(name, 0700); err != nil {
return nil, err
}
continue
}
if !entry.Mode().IsRegular() {
return nil, fmt.Errorf("unsupported archive entry %q", name)
}
if len(rows) >= source.Files || entry.UncompressedSize64 > maxExpanded || expanded+int64(entry.UncompressedSize64) > maxExpanded {
return nil, fmt.Errorf("corpus exceeds pinned inventory or size limit")
}
if err := root.MkdirAll(filepath.Dir(name), 0700); err != nil {
return nil, err
}
out, err := root.OpenFile(name, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0600)
if err != nil {
return nil, err
}
in, err := entry.Open()
if err != nil {
_ = out.Close()
return nil, err
}
h := sha256.New()
n, copyErr := io.Copy(io.MultiWriter(out, h), io.LimitReader(in, maxExpanded-expanded+1))
readClose, writeClose := in.Close(), out.Close()
if copyErr != nil {
return nil, copyErr
}
if readClose != nil {
return nil, readClose
}
if writeClose != nil {
return nil, writeClose
}
expanded += n
if expanded > maxExpanded || n != int64(entry.UncompressedSize64) {
return nil, fmt.Errorf("invalid expanded size")
}
if name == source.LicenseFile && n > 0 {
license = true
}
rows = append(rows, File{Path: filepath.ToSlash(filepath.Join(source.ID, name)), SHA256: hex.EncodeToString(h.Sum(nil)), Bytes: n})
}
if len(rows) != source.Files || !license {
return nil, fmt.Errorf("incomplete corpus: files=%d want=%d license=%t", len(rows), source.Files, license)
}
return rows, nil
}
func WriteJSON(path string, value any) error {
data, err := json.MarshalIndent(value, "", " ")
if err != nil {
return err
}
// #nosec G703 -- The local command or test runner selects this artifact path; vendor content never supplies it.
return os.WriteFile(path, append(data, '\n'), 0600)
}
package corpusgate
import (
"fmt"
"os"
"path/filepath"
"sort"
"strings"
)
func Root(key string) (string, error) {
root := os.Getenv(key)
if root == "" && os.Getenv("CSM_CORPUS_REQUIRED") == "1" {
return "", fmt.Errorf("required corpus variable %s is missing", key)
}
return root, nil
}
type Report struct {
Engine string `json:"engine"`
Scanned int `json:"scanned"`
Hits map[string]int `json:"hits"`
Thresholds map[string]int `json:"thresholds"`
Statuses map[string]int `json:"statuses,omitempty"`
StatusThresholds map[string]int `json:"status_thresholds,omitempty"`
}
func (r Report) Validate() error {
if r.Scanned <= 0 {
return fmt.Errorf("%s scanned no files", r.Engine)
}
var bad []string
for rule, count := range r.Hits {
if count > r.Thresholds[rule] {
bad = append(bad, fmt.Sprintf("%s=%d (maximum %d)", rule, count, r.Thresholds[rule]))
}
}
for status, count := range r.Statuses {
if status != "analyzed" && status != "not_candidate" && count > r.StatusThresholds[status] {
bad = append(bad, fmt.Sprintf("status %s=%d (maximum %d)", status, count, r.StatusThresholds[status]))
}
}
sort.Strings(bad)
if len(bad) > 0 {
return fmt.Errorf("%s corpus regressions: %s", r.Engine, strings.Join(bad, ", "))
}
return nil
}
// Save writes before validation, retaining evidence from failing gates too.
func (r Report) Save() error {
dir := os.Getenv("CSM_CORPUS_REPORT_DIR")
if dir == "" {
if os.Getenv("CSM_CORPUS_REQUIRED") == "1" {
return fmt.Errorf("CSM_CORPUS_REPORT_DIR is required")
}
return nil
}
thresholds := make(map[string]int, len(r.Hits))
for rule := range r.Hits {
thresholds[rule] = r.Thresholds[rule]
}
for rule, limit := range r.Thresholds {
thresholds[rule] = limit
}
r.Thresholds = thresholds
return WriteJSON(filepath.Join(dir, r.Engine+".json"), r)
}
package daemon
import (
"context"
"log"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/reporting"
)
type abuseReportConsumer struct {
stop <-chan struct{}
interval time.Duration
persist func(reporting.Report) error
drain func(context.Context)
queue *queuehealth.Channel[reporting.Report]
mu sync.Mutex
closed bool
}
func newAbuseReportConsumer(stop <-chan struct{}, capacity int, interval time.Duration, persist func(reporting.Report) error, drain func(context.Context)) *abuseReportConsumer {
return &abuseReportConsumer{
stop: stop, interval: interval, persist: persist, drain: drain,
queue: queuehealth.NewChannel[reporting.Report](capacity, time.Minute),
}
}
func (c *abuseReportConsumer) QueueStatuses(now time.Time) map[string]queuehealth.Status {
return map[string]queuehealth.Status{"ingress": c.queue.Snapshot(now)}
}
func (c *abuseReportConsumer) enqueue(r reporting.Report) bool {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
c.queue.Lose(1)
return false
}
select {
case <-c.stop:
c.queue.Lose(1)
return false
default:
return c.queue.TrySend(r)
}
}
func (c *abuseReportConsumer) closeAdmission() {
c.mu.Lock()
defer c.mu.Unlock()
if !c.closed {
c.closed = true
c.queue.Close()
}
}
func (c *abuseReportConsumer) run() {
ctx, cancel := context.WithCancel(context.Background())
canceled := make(chan struct{})
go func() {
defer close(canceled)
select {
case <-c.stop:
cancel()
case <-ctx.Done():
}
}()
var loggedLoss uint64
defer func() {
cancel()
<-canceled
c.closeAdmission()
c.queue.DiscardPending()
c.logDropped(&loggedLoss)
}()
ticker := time.NewTicker(c.interval)
defer ticker.Stop()
for {
select {
case <-c.stop:
c.persistRemaining()
return
default:
}
select {
case <-c.stop:
c.persistRemaining()
return
case work := <-c.queue.Items():
c.process(work)
case <-ticker.C:
c.logDropped(&loggedLoss)
c.drain(ctx)
}
}
}
func (c *abuseReportConsumer) persistRemaining() {
// Stop admission before draining: a captured report hook must not append
// work after the final empty check or keep shutdown running indefinitely.
c.closeAdmission()
for work := range c.queue.Items() {
c.process(work)
}
}
func (c *abuseReportConsumer) process(work queuehealth.Work[reporting.Report]) {
work.Ticket.Start(time.Now())
settled := false
defer func() {
if !settled {
work.Ticket.Reject(time.Now())
}
}()
err := c.persist(work.Value)
if err != nil {
work.Ticket.Reject(time.Now())
} else {
work.Ticket.Finish(time.Now())
}
settled = true
if err != nil {
log.Printf("abuse-reporting: report persistence failed: %v", err)
}
}
func (c *abuseReportConsumer) logDropped(previous *uint64) {
total := c.queue.Snapshot(time.Now()).DroppedTotal
if n := total - *previous; n > 0 {
log.Printf("abuse-reporting: %d report(s) incomplete", n)
*previous = total
}
}
package daemon
import (
"crypto/ed25519"
"encoding/hex"
"log"
"net"
"os"
"path/filepath"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/reporting"
)
const (
abuseReportSpoolFile = "abuse_reports.db"
abuseReportSpoolDefault = 10000
abuseReportQueueDefault = abuseReportSpoolDefault
abuseReportDrainEvery = time.Minute
)
// abuseReportFirebreak returns the predicate that keeps protected addresses
// out of abuse reports. Var so tests that report documentation-range
// fixtures can lift it.
var abuseReportFirebreak = func(d *Daemon) func(string) bool { return d.centralFirebreak() }
// startAbuseReporting wires the abuse reporter from config: it sets
// alert.ReportHook so confirmed-abuse findings are gated, minimized, and
// queued for spooling, and returns the spool drain loop to run as a supervised
// goroutine.
// It returns nil when reporting is disabled or misconfigured (logged), leaving
// the alert path untouched.
func (d *Daemon) startAbuseReporting() func() {
alert.SetReportHook(nil)
rc := d.cfg.Reputation.Report
if !rc.Enabled {
return nil
}
targets := buildReportTargets(rc.Targets)
if len(targets) == 0 {
log.Printf("abuse-reporting: enabled but no usable targets configured; reporting stays off")
return nil
}
enabled := classSet(rc.Classes)
if len(enabled) == 0 {
log.Printf("abuse-reporting: enabled but no valid classes configured; reporting stays off")
return nil
}
spoolPath := rc.SpoolPath
if spoolPath == "" {
spoolPath = filepath.Join(d.cfg.StatePath, abuseReportSpoolFile)
}
max := rc.SpoolMax
if max <= 0 {
max = abuseReportSpoolDefault
}
spool, err := reporting.NewSpool(spoolPath, "reports", max)
if err != nil {
log.Printf("abuse-reporting: cannot open spool %s: %v; reporting stays off", spoolPath, err)
return nil
}
spooler := reporting.NewSpooler(spool, reporting.NewSender(nil, nil), targets, abuseReportDrainEvery)
// The same firebreak that guards central-intel actions keeps
// infrastructure, Cloudflare edges and verified crawlers out of the
// shared abuse set.
firebreak := abuseReportFirebreak(d)
gate := reporting.Gate{
Enabled: enabled,
Protected: func(ip net.IP) bool { return firebreak(ip.String()) },
}
stopCh := make(chan struct{})
doneCh := make(chan struct{})
d.abuseReportStop = stopCh
d.abuseReportDone = doneCh
consumer := newAbuseReportConsumer(stopCh, abuseReportQueueSize(max), abuseReportDrainEvery, spooler.Enqueue, spooler.DrainOnce)
d.registerQueueSource("abuse_reporting", abuseReportQueues{ingress: consumer, spool: spool})
alert.SetReportHook(func(f alert.Finding) {
if r, ok := gate.Consider(f); ok {
consumer.enqueue(r)
}
})
log.Printf("abuse-reporting: enabled for %d target(s), %d class(es)", len(targets), len(enabled))
return func() {
defer close(doneCh)
defer func() {
alert.SetReportHook(nil)
_ = spool.Close()
}()
consumer.run()
}
}
type abuseReportQueues struct {
ingress *abuseReportConsumer
spool *reporting.Spool
}
func (q abuseReportQueues) QueueStatuses(now time.Time) map[string]queuehealth.Status {
statuses := q.ingress.QueueStatuses(now)
for name, status := range q.spool.QueueStatuses(now) {
statuses[name] = status
}
return statuses
}
func (d *Daemon) stopAbuseReporting() {
if d.abuseReportStop == nil {
return
}
close(d.abuseReportStop)
<-d.abuseReportDone
d.abuseReportStop = nil
d.abuseReportDone = nil
}
func abuseReportQueueSize(spoolMax int) int {
if spoolMax > 0 && spoolMax < abuseReportQueueDefault {
return spoolMax
}
return abuseReportQueueDefault
}
// classSet parses configured class names into the set the gate accepts,
// skipping unknown values with a log line.
func classSet(names []string) map[reporting.Class]bool {
known := map[reporting.Class]bool{
reporting.ClassBruteforce: true,
reporting.ClassPHPRelay: true,
reporting.ClassCredentialStuffing: true,
reporting.ClassBadASNEgress: true,
}
out := make(map[reporting.Class]bool)
for _, n := range names {
c := reporting.Class(n)
if known[c] {
out[c] = true
} else {
log.Printf("abuse-reporting: ignoring unknown report class %q", n)
}
}
return out
}
// reportTargetConfig mirrors the per-target config shape; declared so the
// builder takes a concrete slice type from the anonymous struct in config.
type reportTargetConfig = struct {
Name string `yaml:"name"`
URL string `yaml:"url"`
Transport string `yaml:"transport"`
NodeID string `yaml:"node_id"`
KeyID string `yaml:"key_id"`
KeyEnv string `yaml:"key_env"`
TokenEnv string `yaml:"token_env"`
}
// buildReportTargets resolves configured targets into sender targets, reading
// key material from the environment. Invalid targets are skipped with a log
// line rather than failing startup.
func buildReportTargets(cfgTargets []reportTargetConfig) []reporting.Target {
var targets []reporting.Target
for _, ct := range cfgTargets {
if ct.Name == "" || ct.URL == "" || ct.NodeID == "" || ct.KeyID == "" {
log.Printf("abuse-reporting: skipping target with missing name/url/node_id/key_id")
continue
}
if err := reporting.ValidateTargetURL(ct.URL); err != nil {
log.Printf("abuse-reporting: target %q: URL must be HTTPS or loopback HTTP; skipping", ct.Name)
continue
}
t := reporting.Target{
Name: ct.Name,
URL: ct.URL,
NodeID: ct.NodeID,
KeyID: ct.KeyID,
}
secret := os.Getenv(ct.KeyEnv)
switch reporting.Transport(ct.Transport) {
case reporting.TransportEd25519:
raw, err := hex.DecodeString(secret)
if err != nil || len(raw) != ed25519.PrivateKeySize {
log.Printf("abuse-reporting: target %q: configured key_env must hold a 64-byte hex Ed25519 key; skipping", ct.Name)
continue
}
t.Transport = reporting.TransportEd25519
t.Ed25519Key = ed25519.PrivateKey(raw)
case reporting.TransportHMAC:
if secret == "" {
log.Printf("abuse-reporting: target %q: configured key_env is empty; skipping", ct.Name)
continue
}
t.Transport = reporting.TransportHMAC
t.HMACSecret = []byte(secret)
if ct.TokenEnv != "" {
t.BearerToken = os.Getenv(ct.TokenEnv)
}
default:
log.Printf("abuse-reporting: target %q: unknown transport %q; skipping", ct.Name, ct.Transport)
continue
}
targets = append(targets, t)
}
return targets
}
package daemon
import (
"path/filepath"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
)
// defaultLogDir is the directory the packaged unit creates for CSM's logs.
const defaultLogDir = "/var/log/csm"
// actionLogPath puts the action log beside the SIEM audit log, so an operator
// who moved that file gets both streams in one directory.
func actionLogPath(cfg *config.Config) string {
dir := defaultLogDir
if configured := cfg.Alerts.AuditLog.File.Path; configured != "" {
dir = filepath.Dir(configured)
}
return actionlog.DefaultPath(dir)
}
// installActionLog points the process-wide action recorder at the log file.
// Every subsystem that changes host state records through it, so this runs
// before the watchers and the check scheduler start.
func (d *Daemon) installActionLog() {
path := actionLogPath(d.cfg)
sink := actionlog.NewFileSink(func() string { return path }, func(err error) {
csmlog.Warn("action log write failed", "path", path, "err", err)
})
actionlog.SetSink(sink, d.cfg.Hostname)
}
//go:build linux
package daemon
import (
"bytes"
"context"
"errors"
"fmt"
"os"
"sync/atomic"
"syscall"
"time"
"unsafe"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
)
// auditLogPath is the file the kernel auditd writes events to. Var (not
// const) so tests can redirect it under t.TempDir().
var auditLogPath = "/var/log/audit/audit.log"
// AFAlgAuditListener tails /var/log/audit/audit.log via inotify, parses
// each new line for the csm_af_alg_socket auditd key, and emits a
// Critical alert.Finding (plus optional kill/quarantine reactions)
// within milliseconds of the syscall. Sub-second response on hosts where
// BPF LSM is not available.
//
// The listener seeks to end-of-file at startup — events that pre-date
// the daemon are intentionally not re-alerted; the periodic critical-tier
// CheckAFAlgSocketUsage handles backfill via its (timestamp, serial)
// cursor in state.Store.
type AFAlgAuditListener struct {
alertCh chan<- alert.Finding
cfg *config.Config
path string
inotifyFd int
file *os.File
pos int64
leftover []byte // partial line accumulator across reads
droppedOversize bool // mid-drop of a line that overflowed the cap
cursorAnchor []byte // bytes immediately before pos, used to detect copytruncate rewrites
// Re-open retry state. A rotation (or a broken inotify fd) sets
// reopenPending; each tick retries open() once its backoff has elapsed,
// so one failed re-open no longer blinds the listener forever.
reopenPending bool
reopenBackoff time.Duration
reopenNotBefore time.Time
eventCount atomic.Uint64 // observed by tests / metrics
}
// Re-open backoff bounds. Start small so a normal logrotate race (the
// replacement file lands a beat after the move) recovers within a tick or two,
// then grow to a cap so a genuinely absent audit log does not spin.
const (
afAlgReopenBackoffInitial = 1 * time.Second
afAlgReopenBackoffMax = 30 * time.Second
)
func nextAFAlgReopenBackoff(cur time.Duration) time.Duration {
if cur <= 0 {
return afAlgReopenBackoffInitial
}
next := cur * 2
if next > afAlgReopenBackoffMax {
return afAlgReopenBackoffMax
}
return next
}
// afAlgMaxLeftoverBytes caps the partial-line accumulator. A real auditd
// SYSCALL record is well under 1 KiB; 64 KiB leaves generous headroom while
// bounding memory if a record never terminates.
const afAlgMaxLeftoverBytes = 64 * 1024
// afAlgCursorAnchorBytes is the suffix length remembered at the current read
// cursor. On each tail tick the listener verifies those bytes still sit before
// pos; if copytruncate rewrote the file past the old cursor between ticks, the
// anchor no longer matches and the listener resets to offset 0 instead of
// skipping the fresh records.
const afAlgCursorAnchorBytes = 64
// Mode reports the live-monitor backend kind. Matches the BPF path's
// "bpf-lsm" return so the coordinator and operator-visible logs use a
// stable, machine-readable label.
func (l *AFAlgAuditListener) Mode() string { return "auditd-tail" }
// EventCount returns the number of csm_af_alg_socket events this
// listener has parsed since startup. Operational metric; not exported
// to Prometheus here to keep the listener self-contained.
func (l *AFAlgAuditListener) EventCount() uint64 { return l.eventCount.Load() }
// NewAFAlgAuditListener opens the audit log, seeks to its current end,
// and registers an inotify watch on the file. The watch fires for
// IN_MODIFY (new bytes appended), IN_MOVE_SELF (logrotate moves the
// file), and IN_DELETE_SELF (file unlinked) so we can re-open the
// rotated/replaced file.
//
// Returns an error if /var/log/audit/audit.log is missing — caller is
// expected to log a warning and either skip live detection (with
// periodic check still active) or retry later.
func NewAFAlgAuditListener(alertCh chan<- alert.Finding, cfg *config.Config) (*AFAlgAuditListener, error) {
l := &AFAlgAuditListener{
alertCh: alertCh,
cfg: cfg,
path: auditLogPath,
}
if err := l.open(); err != nil {
return nil, err
}
return l, nil
}
// open initialises (or re-initialises after rotation) the audit log fd
// and inotify watch. Seeks to end-of-file so we tail forward, never
// re-alerting historical events.
func (l *AFAlgAuditListener) open() error {
// #nosec G304 -- l.path is /var/log/audit/audit.log (or t.TempDir()
// equivalent); not user-controlled.
f, err := os.Open(l.path)
if err != nil {
return fmt.Errorf("open %s: %w", l.path, err)
}
end, err := f.Seek(0, 2) // SEEK_END
if err != nil {
_ = f.Close()
return fmt.Errorf("seek %s: %w", l.path, err)
}
fd, err := unix.InotifyInit1(unix.IN_CLOEXEC | unix.IN_NONBLOCK)
if err != nil {
_ = f.Close()
return fmt.Errorf("inotify_init: %w", err)
}
mask := uint32(unix.IN_MODIFY | unix.IN_MOVE_SELF | unix.IN_DELETE_SELF)
if _, err := unix.InotifyAddWatch(fd, l.path, mask); err != nil {
_ = unix.Close(fd)
_ = f.Close()
return fmt.Errorf("inotify_add_watch %s: %w", l.path, err)
}
// Replace any existing fds (rotation case).
if l.file != nil {
_ = l.file.Close()
}
if l.inotifyFd != 0 {
_ = unix.Close(l.inotifyFd)
}
l.file = f
l.pos = end
l.inotifyFd = fd
l.leftover = nil
l.droppedOversize = false
l.captureCursorAnchor()
l.reopenPending = false
l.reopenBackoff = 0
l.reopenNotBefore = time.Time{}
return nil
}
// reopenIfDue services a pending re-open. It returns true when the listener is
// ready to tail this tick: either no re-open was pending, or a due re-open
// succeeded. On a failed re-open it keeps reopenPending set and schedules the
// next attempt with exponential backoff, returning false so the caller skips
// tailing a file that is not yet in place. A re-open still gated by its backoff
// window also returns false without touching the filesystem.
func (l *AFAlgAuditListener) reopenIfDue(now time.Time) bool {
if !l.reopenPending {
return true
}
if now.Before(l.reopenNotBefore) {
return false
}
if err := l.open(); err != nil {
l.reopenBackoff = nextAFAlgReopenBackoff(l.reopenBackoff)
l.reopenNotBefore = now.Add(l.reopenBackoff)
csmlog.Warn("af_alg audit listener: re-open failed; will retry",
"err", err, "retry_in", l.reopenBackoff)
return false
}
return true
}
// Run drains inotify events and tails the audit log until ctx is done.
// Polls the inotify fd every poll interval (matches forwarder_watcher's
// approach — no epoll, easier to reason about).
//
// On rotation (IN_MOVE_SELF / IN_DELETE_SELF) the listener re-opens the
// new audit.log and continues from its end. If the replacement file is not
// on disk yet the re-open is retried with backoff on subsequent ticks
// instead of giving up (one failed re-open used to blind the listener until
// the next successful rotation). A broken inotify fd is recovered the same
// way. There is still a brief window during rotation where events written
// between the move and the re-open can be missed; the periodic critical-tier
// check covers that gap via its persistent cursor.
func (l *AFAlgAuditListener) Run(ctx context.Context) {
defer func() {
if l.inotifyFd != 0 {
_ = unix.Close(l.inotifyFd)
}
if l.file != nil {
_ = l.file.Close()
}
}()
inotifyBuf := make([]byte, 4096)
readBuf := make([]byte, 16*1024)
// 500 ms tick gives sub-second average detection latency. The cost
// is ~2 cheap syscalls/sec (inotify Read returns EAGAIN immediately
// when nothing's queued, file Seek+Read returns EOF cheaply when
// no new bytes). Cheap enough to not bother with epoll/select for
// a v1 — we can refactor to true event-driven if we ever need
// hundred-microsecond latency.
ticker := time.NewTicker(500 * time.Millisecond)
defer ticker.Stop()
// Start by draining anything the kernel may have queued during
// startup.
l.tail(readBuf)
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
rotated, ok := l.drainInotify(inotifyBuf)
// A rotation needs a re-open; a broken inotify fd (ok=false) needs
// one too, so a persistently unhealthy watch is rebuilt instead of
// leaving us blind. This is the "safety net" the retry loop provides.
if rotated || !ok {
l.reopenPending = true
}
if !l.reopenIfDue(time.Now()) {
continue
}
l.tail(readBuf)
}
}
}
// drainInotify reads any queued inotify events. Returns:
//
// ok=false on read errors that aren't EAGAIN/EINTR (the listener
// should skip this tick entirely; the next tick re-tries).
// rotated=true if any IN_MOVE_SELF or IN_DELETE_SELF event was seen.
//
// IN_MODIFY events implicitly trigger a tail() in the caller because we
// always read at the end of each tick on success.
func (l *AFAlgAuditListener) drainInotify(buf []byte) (rotated, ok bool) {
for {
n, err := unix.Read(l.inotifyFd, buf)
if err != nil {
// EAGAIN: no more events queued — the normal terminator.
// EINTR: interrupted by signal, also benign — retry next tick.
if errors.Is(err, syscall.EAGAIN) || errors.Is(err, syscall.EWOULDBLOCK) || errors.Is(err, syscall.EINTR) {
return rotated, true
}
// Anything else (EBADF after a stray Close, EIO, EFAULT)
// signals the inotify fd is no longer healthy. Surface it so
// the operator notices, and return ok=false: the caller marks a
// re-open pending and the backoff retry rebuilds the watch.
csmlog.Warn("af_alg audit listener: inotify read error", "err", err)
return rotated, false
}
if n <= 0 {
return rotated, true
}
offset := 0
for offset+unix.SizeofInotifyEvent <= n {
// #nosec G103 -- inotify packed binary stream;
// reinterpretation is required and the bound check above
// guarantees the read is in-range.
ev := (*unix.InotifyEvent)(unsafe.Pointer(&buf[offset]))
if ev.Mask&(unix.IN_MOVE_SELF|unix.IN_DELETE_SELF) != 0 {
rotated = true
}
offset += unix.SizeofInotifyEvent + int(ev.Len)
}
}
}
// tail reads any new bytes since l.pos, splits on newlines, and feeds
// each complete line to handleLine. Partial trailing bytes (no newline
// yet) are buffered in l.leftover for the next tick.
func (l *AFAlgAuditListener) tail(buf []byte) {
for {
// Detect in-place truncation (copytruncate logrotate, or a manual
// truncate). When the file shrinks below our read cursor, seeking
// forward would skip the content rewritten from offset 0 and the
// listener would go blind. Size alone misses a fast truncate+rewrite
// that grows past the old cursor before this tick, so verify the
// cursor anchor too.
if fi, err := l.file.Stat(); err == nil {
switch {
case fi.Size() < l.pos:
l.resetTailCursor()
case l.cursorAnchorChanged():
l.resetTailCursor()
}
}
verifiedPos := l.pos
verifiedAnchor := append([]byte(nil), l.cursorAnchor...)
if _, err := l.file.Seek(l.pos, 0); err != nil {
csmlog.Warn("af_alg audit listener: seek failed", "err", err)
return
}
for {
n, err := l.file.Read(buf)
if n > 0 {
if l.cursorAnchorChangedAt(verifiedPos, verifiedAnchor) {
l.resetTailCursor()
break
}
l.feed(buf[:n])
l.pos += int64(n)
}
if err != nil {
if l.cursorAnchorChangedAt(verifiedPos, verifiedAnchor) {
l.resetTailCursor()
break
}
// io.EOF or EAGAIN: out of data for this tick. Re-check the
// old cursor after capturing the new anchor so a copytruncate
// rewrite cannot land in the tiny gap and get treated as the
// new baseline for the next tick.
if !l.refreshCursorAnchorAfterRead(verifiedPos, verifiedAnchor) {
break
}
return
}
}
}
}
func (l *AFAlgAuditListener) resetTailCursor() {
l.pos = 0
l.leftover = nil
l.droppedOversize = false
l.cursorAnchor = nil
}
func (l *AFAlgAuditListener) cursorAnchorChanged() bool {
return l.cursorAnchorChangedAt(l.pos, l.cursorAnchor)
}
func (l *AFAlgAuditListener) cursorAnchorChangedAt(pos int64, anchor []byte) bool {
if l.file == nil || pos <= 0 || len(anchor) == 0 {
return false
}
if int64(len(anchor)) > pos {
return true
}
buf := make([]byte, len(anchor))
n, err := l.file.ReadAt(buf, pos-int64(len(buf)))
if err != nil || n != len(buf) {
return true
}
return !bytes.Equal(buf, anchor)
}
func (l *AFAlgAuditListener) refreshCursorAnchorAfterRead(verifiedPos int64, verifiedAnchor []byte) bool {
if l.cursorAnchorChangedAt(verifiedPos, verifiedAnchor) {
l.resetTailCursor()
return false
}
l.captureCursorAnchor()
if l.cursorAnchorChangedAt(verifiedPos, verifiedAnchor) {
l.resetTailCursor()
return false
}
return true
}
func (l *AFAlgAuditListener) captureCursorAnchor() {
if l.file == nil || l.pos <= 0 {
l.cursorAnchor = nil
return
}
n := afAlgCursorAnchorBytes
if l.pos < int64(n) {
n = int(l.pos)
}
buf := make([]byte, n)
read, err := l.file.ReadAt(buf, l.pos-int64(n))
if err != nil || read != n {
l.cursorAnchor = nil
return
}
l.cursorAnchor = buf
}
// feed appends a chunk of audit-log bytes to the leftover buffer and
// emits a finding for each complete line containing the csm_af_alg_socket
// key. Lines are kept in the leftover until terminated by '\n' so a
// short read at the end of the buffer does not corrupt a multi-byte
// timestamp split across two reads.
func (l *AFAlgAuditListener) feed(chunk []byte) {
l.leftover = append(l.leftover, chunk...)
for {
idx := bytes.IndexByte(l.leftover, '\n')
if idx < 0 {
// No complete line yet. A real audit record fits well under the
// cap; an unterminated buffer past it is garbage (truncated write,
// binary noise, or an attacker-stretched exe= path). Drop it so the
// accumulator cannot grow without bound, and resync at the next
// newline.
if len(l.leftover) > afAlgMaxLeftoverBytes {
l.leftover = nil
l.droppedOversize = true
}
return
}
line := l.leftover[:idx]
l.leftover = l.leftover[idx+1:]
// Skip the remainder of a line whose head was already dropped for
// exceeding the cap: it is a partial record, not a parseable line.
if l.droppedOversize {
l.droppedOversize = false
continue
}
l.handleLine(string(line))
}
}
// handleLine inspects one audit log line. If it carries the
// csm_af_alg_socket key, parse the event and dispatch a finding.
func (l *AFAlgAuditListener) handleLine(line string) {
ev, ok := checks.ParseAFAlgEventLine(line)
if !ok {
return
}
l.eventCount.Add(1)
finding := alert.Finding{
Severity: alert.Critical,
Check: "af_alg_socket_use",
Message: fmt.Sprintf("AF_ALG socket opened by uid=%s exe=%s", ev.UID, ev.Exe),
TenantID: checks.AFAlgOwner(ev),
Timestamp: time.Now(),
Details: fmt.Sprintf(
"Live audit-log detection: timestamp=%s serial=%s\nauid=%s uid=%s comm=%q exe=%q pid=%s\n"+
"AF_ALG is essentially never used by cPanel/PHP workloads. This is\n"+
"the kernel-level exploit signature for CVE-2026-31431 (\"Copy Fail\").\n"+
"This event was caught by the live audit-log listener (sub-second\n"+
"latency); investigate this process immediately.",
ev.Timestamp, ev.Serial, ev.AUID, ev.UID, ev.Comm, ev.Exe, ev.PID,
),
}
// Non-blocking send — the alert dispatcher buffer is sized for bursts;
// dropping a finding under extreme pressure is preferable to blocking
// the listener loop.
if !alert.TryEnqueue(l.alertCh, finding) {
csmlog.Warn("af_alg audit listener: alert channel full; finding dropped", "uid", ev.UID, "exe", ev.Exe)
}
// Optional reactions (kill, quarantine) gated by config; implemented
// in af_alg_react.go so the BPF path can reuse the same logic.
reactToAFAlgEvent(l.cfg, ev)
}
//go:build !(linux && bpf)
package daemon
import (
"context"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
)
// tryStartBPFLSM is the no-tag fallback. The real implementation lives in
// af_alg_bpf.go behind //go:build linux && bpf.
func tryStartBPFLSM(_ context.Context, _ chan<- alert.Finding, _ *config.Config) (AFAlgLiveMonitor, error) {
return nil, bpf.ErrNotBuilt
}
// Code generated by bpf2go; DO NOT EDIT.
//go:build 386 || amd64
package af_alg_bpfprog
import (
"bytes"
_ "embed"
"fmt"
"io"
"structs"
"github.com/cilium/ebpf"
)
type AFAlgAfAlgEvent struct {
_ structs.HostLayout
Uid uint32
Pid uint32
Ppid uint32
Comm [16]uint8
ParentComm [16]uint8
Exe [256]uint8
}
type AFAlgCsmQueueStats struct {
_ structs.HostLayout
Lost uint64
Submitted uint64
}
// Names of all BPF objects in the ELF.
//
// Used for safe lookups in a Collection or CollectionSpec.
const (
AFAlgMapEvents = "events"
AFAlgMapQueueStats = "queue_stats"
AFAlgProgCsmBlockAfAlg = "csm_block_af_alg"
AFAlgVarUnused = "unused"
)
// LoadAFAlg returns the embedded CollectionSpec for AFAlg.
func LoadAFAlg() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_AFAlgBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load AFAlg: %w", err)
}
return spec, err
}
// LoadAFAlgObjects loads AFAlg and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *AFAlgObjects
// *AFAlgPrograms
// *AFAlgMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func LoadAFAlgObjects(obj any, opts *ebpf.CollectionOptions) error {
spec, err := LoadAFAlg()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// AFAlgSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type AFAlgSpecs struct {
AFAlgProgramSpecs
AFAlgMapSpecs
AFAlgVariableSpecs
}
// AFAlgProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type AFAlgProgramSpecs struct {
CsmBlockAfAlg *ebpf.ProgramSpec `ebpf:"csm_block_af_alg"`
}
// AFAlgMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type AFAlgMapSpecs struct {
Events *ebpf.MapSpec `ebpf:"events"`
QueueStats *ebpf.MapSpec `ebpf:"queue_stats"`
}
// AFAlgVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type AFAlgVariableSpecs struct {
Unused *ebpf.VariableSpec `ebpf:"unused"`
}
// AFAlgObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to LoadAFAlgObjects or ebpf.CollectionSpec.LoadAndAssign.
type AFAlgObjects struct {
AFAlgPrograms
AFAlgMaps
AFAlgVariables
}
func (o *AFAlgObjects) Close() error {
return _AFAlgClose(
&o.AFAlgPrograms,
&o.AFAlgMaps,
)
}
// AFAlgMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to LoadAFAlgObjects or ebpf.CollectionSpec.LoadAndAssign.
type AFAlgMaps struct {
Events *ebpf.Map `ebpf:"events"`
QueueStats *ebpf.Map `ebpf:"queue_stats"`
}
func (m *AFAlgMaps) Close() error {
return _AFAlgClose(
m.Events,
m.QueueStats,
)
}
// AFAlgVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to LoadAFAlgObjects or ebpf.CollectionSpec.LoadAndAssign.
type AFAlgVariables struct {
Unused *ebpf.Variable `ebpf:"unused"`
}
// AFAlgPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to LoadAFAlgObjects or ebpf.CollectionSpec.LoadAndAssign.
type AFAlgPrograms struct {
CsmBlockAfAlg *ebpf.Program `ebpf:"csm_block_af_alg"`
}
func (p *AFAlgPrograms) Close() error {
return _AFAlgClose(
p.CsmBlockAfAlg,
)
}
func _AFAlgClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed afalg_x86_bpfel.o
var _AFAlgBytes []byte
package daemon
import (
"context"
"fmt"
"strings"
"sync"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/metrics"
)
// AFAlgLiveMonitor was the local name for the live-monitor interface;
// it is now an alias of the shared bpf.Backend. Existing call sites in
// reactToAFAlgEvent etc. continue to compile.
type AFAlgLiveMonitor = bpf.Backend
// AFAlgBackend* are the operator-facing cfg.Detection.AFAlgBackend values.
// Reuse bpf constants where the public value already matches.
const (
AFAlgBackendAuto = bpf.BackendAuto
AFAlgBackendBPF = bpf.BackendBPF
AFAlgBackendAuditd = "auditd"
AFAlgBackendNone = bpf.BackendNone
)
var (
afAlgBackendMetricOnce sync.Once
afAlgBackendMetric *metrics.GaugeVec
)
func ensureAFAlgBackendMetric() {
afAlgBackendMetricOnce.Do(func() {
afAlgBackendMetric = metrics.NewGaugeVec(
"csm_af_alg_backend",
"Active AF_ALG (Copy Fail) live-monitor backend; 1 for the selected kind, 0 otherwise.",
[]string{"kind"},
)
metrics.MustRegister("csm_af_alg_backend", afAlgBackendMetric)
})
}
func setAFAlgBackendMetric(active string) {
ensureAFAlgBackendMetric()
for _, k := range []string{"bpf-lsm", "auditd-tail", "none"} {
v := 0.0
if k == active {
v = 1.0
}
afAlgBackendMetric.With(k).Set(v)
}
switch active {
case "bpf-lsm":
bpf.SetActive("af_alg", bpf.BackendBPF)
case "auditd-tail":
bpf.SetActive("af_alg", bpf.BackendLegacy)
default:
bpf.SetActive("af_alg", bpf.BackendNone)
}
}
// StartAFAlgLiveMonitor returns the live monitor selected by
// cfg.Detection.AFAlgBackend. "" / "auto" tries BPF LSM first (when
// compiled in and kernel-supported) and falls back to the audit listener.
// "bpf" requires BPF — no audit fallback if BPF is unavailable, useful for
// hosts where the operator deliberately wants the kernel-side block or
// nothing. "auditd" pins the audit listener even on BPF-capable hosts, the
// kill switch when a BPF-tagged release misbehaves. "none" disables the
// live monitor (the periodic critical-tier check still runs). Returns nil
// when no backend ends up active; the metric csm_af_alg_backend{kind=...}
// reflects whichever path was selected.
func StartAFAlgLiveMonitor(alertCh chan<- alert.Finding, cfg *config.Config) AFAlgLiveMonitor {
choice := strings.ToLower(strings.TrimSpace(cfg.Detection.AFAlgBackend))
if choice == "" {
choice = AFAlgBackendAuto
}
switch choice {
case AFAlgBackendAuto, AFAlgBackendBPF, AFAlgBackendAuditd, AFAlgBackendNone:
default:
csmlog.Warn("af_alg live monitor: unknown backend choice, falling back to auto",
"value", choice,
)
choice = AFAlgBackendAuto
}
if choice == AFAlgBackendNone {
csmlog.Info("af_alg live monitor: disabled by config")
setAFAlgBackendMetric("none")
return nil
}
var bpfErr error
if choice == AFAlgBackendAuto || choice == AFAlgBackendBPF {
if mon, err := tryStartBPFLSMFn(context.Background(), alertCh, cfg); err == nil && mon != nil {
csmlog.Info("af_alg live monitor", "backend", "bpf-lsm", "choice", choice)
setAFAlgBackendMetric("bpf-lsm")
return mon
} else if err != nil {
bpfErr = err
csmlog.Info("af_alg live monitor: BPF LSM unavailable",
"state", "bpf-lsm-unsupported",
"reason", err.Error(),
"choice", choice,
)
if choice == AFAlgBackendBPF {
csmlog.Warn("af_alg live monitor: af_alg_backend=bpf but BPF unavailable; no live detection",
"reason", err.Error(),
)
setAFAlgBackendMetric("none")
emitBPFUnavailableFinding(alertCh, "af_alg", choice, "", err)
return nil
}
}
}
listener, err := NewAFAlgAuditListener(alertCh, cfg) //nolint:staticcheck // The Linux constructor is fallible; the non-Linux stub always returns an error.
if err != nil { //nolint:staticcheck // The Linux constructor can also succeed.
csmlog.Warn("af_alg live monitor: auditd fallback unavailable", "err", err)
setAFAlgBackendMetric("none")
if bpfErr != nil {
emitBPFUnavailableFinding(alertCh, "af_alg", choice, "", fmt.Errorf("BPF unavailable: %w; audit fallback unavailable: %w", bpfErr, err))
}
return nil
}
csmlog.Info("af_alg live monitor", "backend", "auditd-tail", "choice", choice)
setAFAlgBackendMetric("auditd-tail")
if bpfErr != nil {
emitBPFUnavailableFinding(alertCh, "af_alg", choice, "auditd-tail", bpfErr)
}
return listener
}
// tryStartBPFLSMFn is a package-level indirection so tests can substitute a
// fake for the BPF probe path without needing the bpf build tag or kernel
// privileges. Production code goes through tryStartBPFLSM, which is the
// stub on default builds and the real probe on -tags bpf builds.
var tryStartBPFLSMFn = tryStartBPFLSM
//go:build linux
package daemon
import (
"context"
"errors"
"fmt"
"math"
"os"
"strconv"
"strings"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/processhandle"
)
var signalAFAlgProcess = processhandle.Signal
// reactToAFAlgEvent applies opt-in live reactions when an AF_ALG socket
// open is caught by either the audit-log listener or the BPF LSM hook.
// Currently supports a single reaction: SIGKILL the offending process
// (gated by config.AutoResponse.CopyFailKillProcess).
//
// Reactions are intentionally narrow: a critical alert is always emitted
// by the listener itself (this function is for *additional* responses
// beyond alerting). Quarantining the offending exe is a future addition;
// keeping the surface minimal until the kill path has been observed in
// production.
//
// Refuses to act on PID 0/1 to avoid catastrophic mistakes if the parser
// ever returns something unexpected.
func reactToAFAlgEvent(cfg *config.Config, ev checks.AFAlgEvent) {
if cfg == nil || !cfg.AutoResponse.CopyFailKillProcess {
return
}
pid, err := strconv.Atoi(ev.PID)
if err != nil || pid <= 1 {
return
}
refused := false
err = signalAFAlgProcess(context.Background(), pid, unix.SIGKILL, func() error {
_, valid, reason := afAlgKillTarget(ev)
if !valid {
refused = true
return fmt.Errorf("refusing to kill: %s", reason)
}
return nil
})
rec := actionlog.Record{Op: "respond.kill_process", Target: "pid " + strconv.Itoa(pid), Actor: actionlog.Daemon, ActorDetail: ev.Exe, Reason: "AF_ALG socket open", Result: actionlog.Applied}
if err != nil {
rec.Result = actionlog.Failed
rec.Error = err.Error()
}
if refused || errors.Is(err, os.ErrProcessDone) {
rec.Result = actionlog.Refused
}
actionlog.Write(rec)
if err != nil {
csmlog.Warn("af_alg react: kill failed",
"pid", pid, "exe", ev.Exe, "uid", ev.UID,
"err", err,
)
return
}
csmlog.Info("af_alg react: killed offending process",
"pid", pid, "exe", ev.Exe, "uid", ev.UID, "comm", ev.Comm,
)
}
// afAlgKillTarget reports whether ev still names the process it described, and
// may therefore be killed.
//
// The audit record can be a tick old by the time it is read. A PID is recycled
// freely on a busy host, so acting on the number alone means SIGKILLing a
// process that merely inherited it -- as root, on a production server. Three
// facts must agree before the kill: the executable behind the PID is the one
// the record named, the process is owned by the recorded user, and it started
// no later than the event. A process that began after the event cannot be the
// one the event describes.
//
// Anything unverifiable fails closed: an event with no executable, an
// unreadable procfs entry, or a process that has already exited.
func afAlgKillTarget(ev checks.AFAlgEvent) (int, bool, string) {
pid, err := strconv.Atoi(ev.PID)
if err != nil || pid <= 1 {
return 0, false, "implausible pid"
}
if ev.Exe == "" || ev.Exe == "(null)" {
return 0, false, "event carries no executable to verify"
}
eventUID, err := strconv.ParseUint(ev.UID, 10, 32)
if err != nil {
return 0, false, "event carries no uid to verify"
}
exe, err := os.Readlink(fmt.Sprintf("%s/%d/exe", procRootDir, pid))
if err != nil {
return 0, false, "process is gone or its executable is unreadable"
}
// A deleted binary is reported as "<path> (deleted)".
if exe != ev.Exe && strings.TrimSuffix(exe, " (deleted)") != ev.Exe {
return 0, false, "pid now runs a different executable"
}
uid, ok := afAlgProcessUID(pid)
if !ok {
return 0, false, "process ownership unreadable"
}
if uid != eventUID {
return 0, false, "pid now belongs to a different user"
}
eventAt, evErr := strconv.ParseFloat(ev.Timestamp, 64)
if evErr != nil || math.IsNaN(eventAt) || math.IsInf(eventAt, 0) || eventAt <= 0 {
return 0, false, "process start time could not be compared with the event"
}
startedBefore, ok := afAlgProcessStartedBefore(pid, eventAt)
if !ok {
return 0, false, "process start time could not be compared with the event"
}
if !startedBefore {
return 0, false, "pid was recycled after the event"
}
return pid, true, ""
}
func afAlgProcessUID(pid int) (uint64, bool) {
data, err := os.ReadFile(fmt.Sprintf("%s/%d/status", procRootDir, pid))
if err != nil {
return 0, false
}
for _, line := range strings.Split(string(data), "\n") {
rest, found := strings.CutPrefix(line, "Uid:")
if !found {
continue
}
fields := strings.Fields(rest)
if len(fields) != 4 {
return 0, false
}
uid, err := strconv.ParseUint(fields[0], 10, 32)
if err != nil {
return 0, false
}
for _, field := range fields[1:] {
credential, parseErr := strconv.ParseUint(field, 10, 32)
if parseErr != nil || credential != uid {
return 0, false
}
}
return uid, true
}
return 0, false
}
func afAlgProcessStartedBefore(pid int, eventAt float64) (bool, bool) {
statData, err := os.ReadFile(fmt.Sprintf("%s/%d/stat", procRootDir, pid))
if err != nil {
return false, false
}
// The comm field is parenthesised and may contain spaces, so fields are
// counted from after the closing parenthesis.
close := strings.LastIndex(string(statData), ")")
if close < 0 {
return false, false
}
fields := strings.Fields(string(statData)[close+1:])
// starttime is field 22 overall, which is index 19 after pid and comm.
if len(fields) < 20 {
return false, false
}
ticks, err := strconv.ParseUint(fields[19], 10, 64)
if err != nil {
return false, false
}
uptimeData, err := os.ReadFile(fmt.Sprintf("%s/uptime", procRootDir))
if err != nil {
return false, false
}
uptimeFields := strings.Fields(string(uptimeData))
if len(uptimeFields) == 0 {
return false, false
}
uptime, err := strconv.ParseFloat(uptimeFields[0], 64)
if err != nil || math.IsNaN(uptime) || math.IsInf(uptime, 0) || uptime < 0 {
return false, false
}
// Read wall time after uptime so the derived event uptime is conservative.
// Ambiguous boundary cases fail closed instead of accepting a recycled PID.
observedAt := float64(time.Now().UnixNano()) / float64(time.Second)
elapsed := observedAt - eventAt
if elapsed < 0 {
return false, true
}
eventUptime := uptime - elapsed
return eventUptime >= 0 && float64(ticks)/afAlgClockTicks <= eventUptime, true
}
// afAlgClockTicks is USER_HZ, 100 on every architecture Linux ships for the
// platforms CSM runs on.
const afAlgClockTicks = 100.0
package daemon
// Provenance must be available before any detector starts, even when the
// optional BPF monitors never initialize the process-context cache.
func init() { wireAncestryProvenance() }
// pkgManagerComms are the process names CSM treats as evidence that an
// observed sensitive-file write originated from a legitimate root-driven
// package transaction. The list intentionally omits shells (sh, bash) and
// generic utilities (cp, mv) -- attackers reuse those. Matching the package
// manager binary itself anywhere in the parent chain is the discriminator.
var pkgManagerComms = map[string]struct{}{
"dnf": {},
"dnf-3": {},
"microdnf": {},
"yum": {},
"rpm": {},
"dpkg": {},
"apt": {},
"apt-get": {},
"unattended-upgr": {}, // unattended-upgrade is comm-truncated to TASK_COMM_LEN-1.
}
func isPackageManagerComm(comm string) bool {
_, ok := pkgManagerComms[comm]
return ok
}
//go:build !linux || !bpf
package daemon
import "github.com/pidginhost/csm/internal/processctx"
// wireAncestryCache is a no-op on hosts built without the bpf build tag.
// cachedAncestryEvidence stays nil and ancestry falls back to the live /proc
// walk, which is racy for short-lived writers but fails closed.
func wireAncestryCache(*processctx.Cache) {}
package daemon
import (
"path/filepath"
"regexp"
"strings"
)
// atomicWriteStageRE matches the `.temp.<digits>.<rest>` filename pattern
// emitted by cPanel's fileTransfer service and similar atomic-write
// helpers. The digits are a nanosecond timestamp; <rest> is the original
// basename that will be rename(2)d into place.
var atomicWriteStageRE = regexp.MustCompile(`^\.temp\.\d+\..+`)
// looksLikeAtomicWriteStage reports whether a base filename matches the
// atomic-write staging convention `.temp.<digits>.<name>`. Pass the
// basename, not the full path.
func looksLikeAtomicWriteStage(name string) bool {
return atomicWriteStageRE.MatchString(name)
}
// atomicWriteRenameCandidate maps cPanel's .temp.<timestamp>.<name> staging
// path to its intended final path. It is only a location hint: the probe must
// still validate content or the destination identity before trusting it.
func atomicWriteRenameCandidate(path string) string {
base := filepath.Base(path)
if !looksLikeAtomicWriteStage(base) {
return ""
}
rest := strings.TrimPrefix(base, ".temp.")
dot := strings.IndexByte(rest, '.')
if dot < 0 || dot == len(rest)-1 {
return ""
}
name := rest[dot+1:]
if name == "." || name == ".." {
return ""
}
return filepath.Join(filepath.Dir(path), name)
}
// atomicWriteContentPath supplies a type and checksum lookup hint without
// changing the descriptor or the reported location of the file being scanned.
func atomicWriteContentPath(path string) string {
if candidate := atomicWriteRenameCandidate(path); candidate != "" {
return candidate
}
return path
}
package daemon
import (
"fmt"
"os"
"github.com/pidginhost/csm/internal/attackdb"
)
func (d *Daemon) prepareAttackDatabase(adb *attackdb.DB) {
d.registerQueueSource("attackdb", adb)
// Seed from permanent blocklist on first run (when attack DB is empty)
if adb.TotalIPs() == 0 {
if n := adb.SeedFromPermanentBlocklist(d.cfg.StatePath); n > 0 {
fmt.Fprintf(os.Stderr, "[%s] Attack DB seeded %d IPs from permanent blocklist\n", ts(), n)
}
}
fmt.Fprintf(os.Stderr, "[%s] Attack DB initialized (%s)\n", ts(), adb.FormatTopLine())
}
package daemon
import (
"context"
"fmt"
"net"
"os/exec"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
)
const (
// cpdoveauthdSocketPath is the cPanel dovecot auth daemon's unix socket.
cpdoveauthdSocketPath = "/usr/local/cpanel/var/cpdoveauthd.sock"
// mailAuthProbeInterval is how often the prober dials the socket.
mailAuthProbeInterval = 20 * time.Second
// mailAuthRestartCooldown is the minimum gap between restart attempts.
mailAuthRestartCooldown = 2 * time.Minute
)
// dialMailAuthBackend probes the cpdoveauthd unix socket. A successful connect
// means the auth backend can answer; a refused/missing socket means it is down.
func dialMailAuthBackend() bool {
c, err := net.DialTimeout("unix", cpdoveauthdSocketPath, 3*time.Second)
if err != nil {
return false
}
_ = c.Close()
return true
}
// restartMailAuthBackend runs the operator-configured restart command.
func restartMailAuthBackend(command string) error {
if strings.TrimSpace(command) == "" {
return fmt.Errorf("restart command is empty")
}
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
// Run through the shell so operators can use normal service commands with
// arguments, such as "systemctl restart dovecot".
// #nosec G204 -- command is operator-configured in root-owned csm.yaml.
return exec.CommandContext(ctx, "/bin/sh", "-c", command).Run()
}
// authBackendHealth tracks reachability of the mail authentication backend
// (cPanel's cpdoveauthd unix socket) via an active probe. While the backend is
// down every mail and SMTP login fails regardless of credentials, so callers
// gate brute-force auto-block on Degraded() to avoid mass-blocking legitimate
// users. chkservd only checks dovecot's listen ports, which stay up during a
// cpdoveauthd outage, so this probe covers a failure mode cPanel's own
// monitoring misses. When restartEnabled, a sustained outage (continuously down
// for at least downGrace) triggers a rate-limited service restart to self-heal.
type authBackendHealth struct {
mu sync.Mutex
now func() time.Time
probe func() bool // true = backend reachable
restart func() error // run the configured restart command
restartEnabled bool
downGrace time.Duration
restartCooldown time.Duration
maxRestartsPerHour int
downSince time.Time // zero == backend currently up
alerted bool // down alert already emitted for the current outage
lastRestart time.Time
restartsThisHour int
hourKey string
}
func newAuthBackendHealth(
now func() time.Time,
probe func() bool,
restart func() error,
restartEnabled bool,
downGrace time.Duration,
restartCooldown time.Duration,
maxRestartsPerHour int,
) *authBackendHealth {
if now == nil {
now = time.Now
}
return &authBackendHealth{
now: now,
probe: probe,
restart: restart,
restartEnabled: restartEnabled,
downGrace: downGrace,
restartCooldown: restartCooldown,
maxRestartsPerHour: maxRestartsPerHour,
}
}
// Degraded reports whether the mail auth backend is currently unreachable.
// Brute-force trackers consult this and suppress auto-block while it is true.
func (h *authBackendHealth) Degraded() bool {
h.mu.Lock()
defer h.mu.Unlock()
return !h.downSince.IsZero()
}
// Observe runs one probe cycle and returns any findings to dispatch: a one-shot
// down alert when an outage starts, and a restart action once the outage is
// sustained past the grace period (subject to cooldown and the hourly cap).
func (h *authBackendHealth) Observe() []alert.Finding {
now := h.now()
reachable := h.probe != nil && h.probe()
h.mu.Lock()
if reachable {
h.downSince = time.Time{}
h.alerted = false
h.mu.Unlock()
return nil
}
var out []alert.Finding
var runRestart bool
var restart func() error
if h.downSince.IsZero() {
h.downSince = now
}
if !h.alerted {
h.alerted = true
out = append(out, alert.Finding{
Severity: alert.Warning,
Check: "mail_auth_backend_degraded",
Message: "Mail auth backend (cpdoveauthd) unreachable; mail and SMTP brute-force auto-block paused",
Details: "CSM could not connect to the cPanel dovecot auth socket. Logins fail regardless of password. Dovecot's listen ports stay up during this fault, so chkservd does not catch it; investigate the auth daemon.",
Timestamp: now,
})
}
if h.restartEnabled && h.restart != nil && now.Sub(h.downSince) >= h.downGrace && h.allowRestart(now) {
h.lastRestart = now
h.restartsThisHour++
restart = h.restart
runRestart = true
}
h.mu.Unlock()
if runRestart {
if err := restart(); err != nil {
out = append(out, alert.Finding{
Severity: alert.High,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-RESTART failed: mail auth backend down >%v and restart errored", h.downGrace),
Details: fmt.Sprintf("Error: %v. Manual intervention required to restore mail authentication.", err),
Timestamp: now,
})
} else {
out = append(out, alert.Finding{
Severity: alert.Warning,
Check: "auto_response",
Message: fmt.Sprintf("AUTO-RESTART: mail auth backend down >%v, restarted the mail service", h.downGrace),
Details: "cpdoveauthd was unreachable beyond the grace period; CSM restarted the mail service to recover authentication.",
Timestamp: now,
})
}
}
return out
}
// allowRestart reports whether a restart may run now, enforcing the per-hour cap
// and the cooldown between attempts. Caller must hold h.mu.
func (h *authBackendHealth) allowRestart(now time.Time) bool {
key := now.Format("2006-01-02T15")
if h.hourKey != key {
h.hourKey = key
h.restartsThisHour = 0
}
if h.maxRestartsPerHour <= 0 {
return false
}
if h.restartsThisHour >= h.maxRestartsPerHour {
return false
}
if !h.lastRestart.IsZero() && now.Sub(h.lastRestart) < h.restartCooldown {
return false
}
return true
}
package daemon
import (
"fmt"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/blockdigest"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
csmlog "github.com/pidginhost/csm/internal/log"
)
var (
blockDigestSendEmail = alert.SendEmail
blockDigestSendWebhookJSON = alert.SendWebhookJSON
)
// buildBlockDigest constructs the collector from config, or returns nil when
// the feature is disabled. Built once at startup; block_digest is
// hotreload:"restart" so its settings do not change under a live daemon.
func (d *Daemon) buildBlockDigest(cfg *config.Config) *blockdigest.Collector {
if !cfg.Alerts.BlockDigest.Enabled {
return nil
}
configuredCountries := blockdigest.ResolveCountries(cfg.Alerts.BlockDigest.Countries, nil)
countries := blockdigest.ResolveCountries(cfg.Alerts.BlockDigest.Countries, cfg.Suppressions.TrustedCountries)
if len(countries) == 0 {
csmlog.Warn("block_digest enabled with no countries; watching ALL countries (set alerts.block_digest.countries or suppressions.trusted_countries)")
}
var countriesOf func() []string
if len(configuredCountries) == 0 {
countriesOf = func() []string {
live := d.currentCfg()
if live == nil {
live = cfg
}
return blockdigest.ResolveCountries(nil, live.Suppressions.TrustedCountries)
}
}
email, webhook := d.blockDigestSinks(cfg)
collector := blockdigest.New(blockdigest.Options{
Countries: countries,
SendOn: cfg.Alerts.BlockDigest.SendOn,
Interval: cfg.BlockDigestInterval(),
Live: cfg.Alerts.BlockDigest.Live,
MinBlock: cfg.Alerts.BlockDigest.MinBlock,
Host: cfg.Hostname,
Version: d.version,
CountriesOf: countriesOf,
CountryOf: d.countryOf,
EnrichModSec: d.modsecEnricher(cfg),
EmailSink: email,
WebhookSink: webhook,
OnError: func(channel string, err error) {
csmlog.Warn("block_digest delivery failed", "channel", channel, "err", err)
},
DeliveryEnabled: func(channel string) bool {
// Explicit destinations still attempt delivery and report disabled
// channels as errors. Default delivery follows current alert policy.
if cfg.Alerts.BlockDigest.Channel != "" {
return true
}
live := d.currentCfg()
if live == nil {
live = cfg
}
if channel == "email" {
return live.Alerts.Email.Enabled
}
return live.Alerts.Webhook.Enabled
},
})
d.registerQueueSource("block_digest", collector)
return collector
}
// blockDigestSinks selects delivery: an empty channel follows whichever alerts
// channels are enabled; an explicit channel forces just that one.
func (d *Daemon) blockDigestSinks(cfg *config.Config) (func(subject, body string) error, func(blockdigest.WebhookPayload) error) {
ch := cfg.Alerts.BlockDigest.Channel
currentCfg := func() *config.Config {
if live := d.currentCfg(); live != nil {
return live
}
return cfg
}
emailSink := func(requireEnabled bool) func(subject, body string) error {
return func(subject, body string) error {
live := currentCfg()
if !live.Alerts.Email.Enabled {
if requireEnabled {
return fmt.Errorf("email alerts disabled")
}
return nil
}
return blockDigestSendEmail(live, subject, body)
}
}
webhookSink := func(requireEnabled bool) func(blockdigest.WebhookPayload) error {
return func(p blockdigest.WebhookPayload) error {
live := currentCfg()
if !live.Alerts.Webhook.Enabled {
if requireEnabled {
return fmt.Errorf("webhook alerts disabled")
}
return nil
}
return blockDigestSendWebhookJSON(live, p)
}
}
switch ch {
case "":
return emailSink(false), webhookSink(false)
case "email":
return emailSink(true), nil
case "webhook":
return nil, webhookSink(true)
default:
return nil, nil
}
}
// countryOf resolves an IP to its ISO country via the loaded GeoIP mmdb. It
// reads through the atomic accessor so it never races the geoip hot-reload swap.
func (d *Daemon) countryOf(ip string) string {
db := getGeoIPDB()
if db == nil {
return ""
}
return db.Lookup(ip).Country
}
// observeBlocks feeds real auto-block findings to the digest collector.
// PERMBLOCK promotions (no reason details) and non-auto_block findings are
// skipped; dedup-by-IP in the collector makes multiple call sites safe.
func (d *Daemon) observeBlocks(actions []alert.Finding) {
if d.blockDigest == nil {
return
}
for _, f := range actions {
if f.Check != "auto_block" || f.Severity != alert.Critical {
continue
}
reason := strings.TrimPrefix(f.Details, "Reason: ")
reason = strings.TrimSuffix(reason, " (warning: "+firewall.CloudflareCoverageWarning+")")
if reason == "" {
continue
}
ip := checks.ExtractIPFromFinding(f)
if ip == "" {
continue
}
d.blockDigest.Observe(ip, reason, f.Timestamp)
}
}
package daemon
import (
"context"
"net/http"
"path/filepath"
"sync"
"time"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/threatintel"
)
var (
botRangesRefreshTotal *metrics.CounterVec
botRangesPrefixes *metrics.GaugeVec
botRangesLastSuccess *metrics.Gauge
botRangesMetricsOnce sync.Once
botRangesPrefixMu sync.Mutex
botRangesPrefixBots = map[string]struct{}{}
)
// botRangesMetrics lazily registers and returns the AI-crawler range-updater
// metrics, mirroring the signature/geoip update metrics: a refresh
// success/failure counter, a per-bot prefix-count gauge, and the timestamp of
// the last successful refresh.
func botRangesMetrics() (*metrics.CounterVec, *metrics.GaugeVec, *metrics.Gauge) {
botRangesMetricsOnce.Do(func() {
botRangesRefreshTotal = metrics.NewCounterVec(
"csm_botranges_refresh_total",
"AI-crawler IP-range refresh attempts, labelled by result (success when at least one vendor feed updated).",
[]string{"result"},
)
metrics.MustRegister("csm_botranges_refresh_total", botRangesRefreshTotal)
botRangesPrefixes = metrics.NewGaugeVec(
"csm_botranges_prefixes",
"Current number of published IP prefixes per AI-crawler identity in the active overlay.",
[]string{"bot"},
)
metrics.MustRegister("csm_botranges_prefixes", botRangesPrefixes)
botRangesLastSuccess = metrics.NewGauge(
"csm_botranges_last_success_timestamp_seconds",
"Unix timestamp of the last successful AI-crawler IP-range refresh.",
)
metrics.MustRegister("csm_botranges_last_success_timestamp_seconds", botRangesLastSuccess)
})
return botRangesRefreshTotal, botRangesPrefixes, botRangesLastSuccess
}
// setBotRangesPrefixGauges syncs the per-bot prefix gauge to the active overlay.
func setBotRangesPrefixGauges() {
_, prefixes, _ := botRangesMetrics()
snap := threatintel.FetchedRangesSnapshot()
botRangesPrefixMu.Lock()
defer botRangesPrefixMu.Unlock()
for bot := range botRangesPrefixBots {
if _, ok := snap[bot]; !ok {
prefixes.With(bot).Set(0)
delete(botRangesPrefixBots, bot)
}
}
for bot, nets := range snap {
prefixes.With(bot).Set(float64(len(nets)))
botRangesPrefixBots[bot] = struct{}{}
}
}
// observeBotRangesRefresh records the outcome of a refresh attempt. On success
// it bumps the prefix gauges and last-success timestamp; on failure it only
// increments the failure counter so the previous overlay's gauges stand.
func observeBotRangesRefresh(success bool) {
total, _, lastSuccess := botRangesMetrics()
if !success {
total.With("failure").Inc()
return
}
total.With("success").Inc()
setBotRangesPrefixGauges()
ts := threatintel.LastFetchedRangesRefresh()
if ts.IsZero() {
ts = time.Now()
}
lastSuccess.Set(float64(ts.Unix()))
}
func (d *Daemon) botRangesCachePath() string {
return filepath.Join(d.cfg.StatePath, "botranges.json")
}
// reloadBotRanges republishes the on-disk overlay into the running daemon. It
// is the botranges.reload control handler's worker: `csm update-bot-ranges`
// fetches fresh ranges, writes the cache, then asks the daemon to reload so the
// new ranges take effect without a restart.
func (d *Daemon) reloadBotRanges() error {
if err := threatintel.LoadFetchedRangesRequired(d.botRangesCachePath()); err != nil {
observeBotRangesRefresh(false)
return err
}
observeBotRangesRefresh(true)
return nil
}
// botRangesUpdater periodically refreshes the published AI-crawler IP ranges
// (OpenAI, Perplexity) so GPTBot/ChatGPT-User/OAI-SearchBot/PerplexityBot stay
// verifiable by address without a new release. Embedded snapshots cover the gap
// before the first refresh and whenever a fetch fails.
func (d *Daemon) botRangesUpdater() {
defer d.wg.Done()
if !d.cfg.BotRangesAutoUpdate() {
return
}
cachePath := d.botRangesCachePath()
if err := threatintel.LoadFetchedRanges(cachePath); err != nil {
csmlog.Warn("bot-ranges cache load failed", "err", err)
} else {
// Reflect the cached overlay and its persisted refresh time in the
// metrics before the first live refresh, so a restart does not show the
// ranges as never-refreshed.
setBotRangesPrefixGauges()
if ts := threatintel.LastFetchedRangesRefresh(); !ts.IsZero() {
_, _, lastSuccess := botRangesMetrics()
lastSuccess.Set(float64(ts.Unix()))
}
}
interval := 24 * time.Hour
if s := d.cfg.Reputation.BotRanges.UpdateInterval; s != "" {
if v, err := time.ParseDuration(s); err == nil && v >= time.Hour {
interval = v
}
}
select {
case <-d.stopCh:
return
case <-time.After(5 * time.Minute):
}
d.doBotRangesUpdate(cachePath)
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
d.doBotRangesUpdate(cachePath)
}
}
}
func (d *Daemon) doBotRangesUpdate(cachePath string) {
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
// Cancel an in-flight fetch promptly on daemon shutdown.
go func() {
select {
case <-d.stopCh:
cancel()
case <-ctx.Done():
}
}()
client := &http.Client{Timeout: 30 * time.Second}
n, err := threatintel.RefreshFetchedRanges(ctx, client, threatintel.DefaultRangeSources(), cachePath)
if err != nil {
csmlog.Warn("bot-ranges refresh error", "err", err)
}
observeBotRangesRefresh(n > 0)
if n > 0 {
csmlog.Info("bot-ranges refreshed", "bots", n)
}
}
package daemon
// PolicyMapPayload is the wire-shape of struct policy_state in the BPF
// program. Mirrors the C layout exactly: three uint32 fields, no
// alignment surprises (all fields are 4 bytes).
type PolicyMapPayload struct {
Enforce uint32
DryRun uint32
ProtectedPorts uint32
}
// policyMapPayload converts the userspace policy into the wire shape
// the BPF program reads. Pure function; tests drive it without loading
// any actual BPF program.
func policyMapPayload(pol BPFEnforcementPolicy) PolicyMapPayload {
return PolicyMapPayload{
Enforce: pol.Enforce,
DryRun: pol.DryRun,
// #nosec G115 -- bounded by protected_ports BPF map max_entries=16
ProtectedPorts: uint32(len(pol.Ports)),
}
}
package daemon
import (
"sync"
"github.com/pidginhost/csm/internal/metrics"
)
// Decision label values match the BPF DECISION_* numeric codes:
// allow (0), dry_run (1), deny (2). Wire-stable strings; SIEM
// dashboards pin on them.
const (
BPFDecisionAllow = "allow"
BPFDecisionDryRun = "dry_run"
BPFDecisionDeny = "deny"
)
var (
bpfEnfMu sync.Mutex
bpfEnfDecisionsVec *metrics.CounterVec
bpfEnfUIDRefreshTotal *metrics.Counter
bpfEnfUIDRefreshFailures *metrics.Counter
)
// RegisterBPFEnforcementMetrics registers the counters on reg. Production
// callers pass metrics.Default(); tests pass metrics.NewRegistry() to
// keep registration isolated.
func RegisterBPFEnforcementMetrics(reg *metrics.Registry) {
bpfEnfMu.Lock()
defer bpfEnfMu.Unlock()
bpfEnfDecisionsVec = metrics.NewCounterVec(
"csm_bpf_enforcement_decisions_total",
"BPF cgroup-deny decisions by label (allow/dry_run/deny).",
[]string{"decision"},
)
reg.MustRegister("csm_bpf_enforcement_decisions_total", bpfEnfDecisionsVec)
bpfEnfUIDRefreshTotal = metrics.NewCounter(
"csm_bpf_enforcement_uid_map_refresh_total",
"BPF safe-UID map refresh successes.",
)
reg.MustRegister("csm_bpf_enforcement_uid_map_refresh_total", bpfEnfUIDRefreshTotal)
bpfEnfUIDRefreshFailures = metrics.NewCounter(
"csm_bpf_enforcement_uid_map_refresh_failures_total",
"BPF safe-UID map refresh failures.",
)
reg.MustRegister("csm_bpf_enforcement_uid_map_refresh_failures_total", bpfEnfUIDRefreshFailures)
}
// BumpBPFEnforcementDecision advances the per-decision counter. Called
// from the connection consumer when a ConnectionEvent with a decision
// field arrives. Unknown labels are silently ignored (caller bug).
func BumpBPFEnforcementDecision(label string) {
bpfEnfMu.Lock()
cv := bpfEnfDecisionsVec
bpfEnfMu.Unlock()
if cv == nil {
return
}
switch label {
case BPFDecisionAllow, BPFDecisionDryRun, BPFDecisionDeny:
cv.With(label).Inc()
}
}
// BumpUIDRefresh advances the periodic-refresh success counter.
func BumpUIDRefresh() {
bpfEnfMu.Lock()
c := bpfEnfUIDRefreshTotal
bpfEnfMu.Unlock()
if c != nil {
c.Inc()
}
}
// BumpUIDRefreshFailure advances the periodic-refresh failure counter.
func BumpUIDRefreshFailure() {
bpfEnfMu.Lock()
c := bpfEnfUIDRefreshFailures
bpfEnfMu.Unlock()
if c != nil {
c.Inc()
}
}
// resetBPFEnforcementMetricsForTest is a test seam.
func resetBPFEnforcementMetricsForTest() {
bpfEnfMu.Lock()
defer bpfEnfMu.Unlock()
bpfEnfDecisionsVec = nil
bpfEnfUIDRefreshTotal = nil
bpfEnfUIDRefreshFailures = nil
}
package daemon
import (
"bufio"
"os"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
)
// BPFEnforcementPolicy is the userspace-derived state to be loaded into
// the BPF policy + protected_ports + safe_uids maps. Pure data; no IO.
type BPFEnforcementPolicy struct {
Enforce uint32
DryRun uint32
Ports []uint16
}
// BuildBPFEnforcementPolicy translates config into the policy struct
// the BPF program consumes. Disabled config returns zero policy
// (Enforce=0); the in-kernel program then short-circuits to allow.
func BuildBPFEnforcementPolicy(cfg *config.Config) BPFEnforcementPolicy {
if cfg == nil || !cfg.BPFEnforcement.Enabled {
return BPFEnforcementPolicy{}
}
if !connectionTrackerAllowsBPF(cfg) {
return BPFEnforcementPolicy{}
}
p := BPFEnforcementPolicy{Enforce: 1}
if cfg.BPFEnforcementDryRunEnabled() {
p.DryRun = 1
}
if cfg.BPFEnforcement.DirectSMTPEgress && checks.DirectSMTPEgressBackendEnabled(cfg, "bpf") {
for _, port := range cfg.Detection.DirectSMTPEgress.Ports {
if port > 0 && port <= 65535 {
p.Ports = append(p.Ports, uint16(port))
}
}
}
if len(p.Ports) == 0 {
return BPFEnforcementPolicy{}
}
return p
}
func connectionTrackerAllowsBPF(cfg *config.Config) bool {
switch strings.ToLower(strings.TrimSpace(cfg.Detection.ConnectionTrackerBackend)) {
case "", "auto", "bpf":
return true
default:
return false
}
}
// safeUIDsFromPasswd returns a map of UIDs that should be exempt from
// in-kernel deny: UID 0 (root), UIDs <1000 (system accounts), and any
// platform-known MTA users that happen to live above 1000 on some
// distros. Hosted account UIDs (>=1000) are NOT in the safe map; their
// connections will be evaluated by the in-kernel deny path when
// enforcement is active.
//
// Caller passes a /etc/passwd path; production wiring uses /etc/passwd.
func safeUIDsFromPasswd(path string) (map[uint32]bool, error) {
f, err := os.Open(path) // #nosec G304 -- caller-controlled; production passes /etc/passwd
if err != nil {
return nil, err
}
defer f.Close()
mta := platform.LocalMTAIdentities(platform.Detect())
out := map[uint32]bool{}
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := scanner.Text()
fields := strings.Split(line, ":")
if len(fields) < 4 {
continue
}
user := fields[0]
uid64, err := strconv.ParseUint(fields[2], 10, 32)
if err != nil {
continue
}
uid := uint32(uid64)
// UID 0 always safe.
if uid == 0 {
out[uid] = true
continue
}
// System UIDs (<1000): always safe. Daemons, services,
// distro-managed users.
if uid < 1000 {
out[uid] = true
continue
}
// Hosted UIDs (>=1000): NOT safe by default. Exception:
// platform-known MTA users on distros that put them above 1000.
if mta.IsMTAUser(user) {
out[uid] = true
}
}
if err := scanner.Err(); err != nil {
return out, err
}
return out, nil
}
package daemon
import (
"sync"
"sync/atomic"
"time"
)
// UIDRefresherConfig drives a UIDRefresher. Refresh is the function
// called every tick; production wiring re-reads /etc/passwd and
// repopulates the BPF safe_uids map. Interval bounds the period;
// production uses 5 minutes.
type UIDRefresherConfig struct {
Interval time.Duration
Refresh func() error
}
// UIDRefresherStats is a counter snapshot.
type UIDRefresherStats struct {
Refreshes uint64
Failures uint64
}
// UIDRefresher runs the configured Refresh on a fixed interval. Stop
// is idempotent.
type UIDRefresher struct {
cfg UIDRefresherConfig
stop chan struct{}
wg sync.WaitGroup
started atomic.Bool
stopped atomic.Bool
refreshes atomic.Uint64
failures atomic.Uint64
}
// NewUIDRefresher returns a stopped refresher. Call Start to launch.
func NewUIDRefresher(cfg UIDRefresherConfig) *UIDRefresher {
return &UIDRefresher{cfg: cfg, stop: make(chan struct{})}
}
// Start launches the refresh goroutine. Idempotent.
func (r *UIDRefresher) Start() {
if r.started.Swap(true) {
return
}
r.wg.Add(1)
go r.loop()
}
// Stop signals the goroutine and waits. Safe to call multiple times.
func (r *UIDRefresher) Stop() {
if r.stopped.Swap(true) {
return
}
close(r.stop)
r.wg.Wait()
}
// Stats returns a counter snapshot.
func (r *UIDRefresher) Stats() UIDRefresherStats {
return UIDRefresherStats{
Refreshes: r.refreshes.Load(),
Failures: r.failures.Load(),
}
}
func (r *UIDRefresher) loop() {
defer r.wg.Done()
t := time.NewTicker(r.cfg.Interval)
defer t.Stop()
for {
select {
case <-r.stop:
return
case <-t.C:
if err := r.cfg.Refresh(); err != nil {
r.failures.Add(1)
continue
}
r.refreshes.Add(1)
}
}
}
package daemon
import (
"fmt"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
)
// emitBPFUnavailableFinding posts an operator-visible warning when a
// BPF-backed live monitor cannot start the kernel-attached path. fallback is
// the selected non-BPF backend; empty means no live fallback is active.
func emitBPFUnavailableFinding(alertCh chan<- alert.Finding, feature, choice, fallback string, err error) bool {
if alertCh == nil {
return false
}
sev := alert.Warning
if choice == bpf.BackendBPF || fallback == "" {
sev = alert.High
}
status := "no live fallback active"
if fallback != "" {
status = fmt.Sprintf("running on %s fallback", fallback)
}
details := ""
if err != nil {
details = err.Error()
}
f := alert.Finding{
Severity: sev,
Check: "bpf_unavailable",
Message: fmt.Sprintf("BPF backend unavailable for %s (operator choice=%q); %s", feature, choice, status),
Details: details,
Timestamp: time.Now(),
}
return alert.TryEnqueue(alertCh, f)
}
package daemon
import (
"context"
"log"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type centralActionConsumer struct {
stop <-chan struct{}
interval time.Duration
refresh func(context.Context) error
perform func(centralQueuedAction) error
queue *queuehealth.Channel[centralQueuedAction]
mu sync.Mutex
closed bool
}
func newCentralActionConsumer(stop <-chan struct{}, capacity int, interval time.Duration, refresh func(context.Context) error, perform func(centralQueuedAction) error) *centralActionConsumer {
return ¢ralActionConsumer{
stop: stop, interval: interval, refresh: refresh, perform: perform,
queue: queuehealth.NewChannel[centralQueuedAction](capacity, time.Minute),
}
}
func (c *centralActionConsumer) QueueStatuses(now time.Time) map[string]queuehealth.Status {
return map[string]queuehealth.Status{"actions": c.queue.Snapshot(now)}
}
func (c *centralActionConsumer) enqueue(a centralQueuedAction) bool {
c.mu.Lock()
defer c.mu.Unlock()
if c.closed {
c.queue.Lose(1)
return false
}
select {
case <-c.stop:
c.queue.Lose(1)
return false
default:
return c.queue.TrySend(a)
}
}
func (c *centralActionConsumer) run() {
ctx, cancel := context.WithCancel(context.Background())
canceled := make(chan struct{})
go func() {
defer close(canceled)
select {
case <-c.stop:
cancel()
case <-ctx.Done():
}
}()
var loggedLoss uint64
defer func() {
cancel()
<-canceled
// A dispatch may retain the old hook after it is uninstalled. Close
// admission with the producer lock before discarding pending work.
c.mu.Lock()
c.closed = true
c.queue.Close()
c.mu.Unlock()
c.queue.DiscardPending()
c.logDropped(&loggedLoss)
}()
if err := c.refresh(ctx); err != nil {
log.Printf("central-intel: initial pull failed: %v", err)
}
ticker := time.NewTicker(c.interval)
defer ticker.Stop()
for {
select {
case <-c.stop:
return
default:
}
select {
case <-c.stop:
return
case work := <-c.queue.Items():
select {
case <-c.stop:
work.Ticket.Reject(time.Now())
return
default:
}
c.process(work)
case <-ticker.C:
c.logDropped(&loggedLoss)
if err := c.refresh(ctx); err != nil {
log.Printf("central-intel: refresh failed: %v", err)
}
}
}
}
func (c *centralActionConsumer) process(work queuehealth.Work[centralQueuedAction]) {
work.Ticket.Start(time.Now())
settled := false
defer func() {
if !settled {
work.Ticket.Reject(time.Now())
}
}()
err := c.perform(work.Value)
if err != nil && !isCentralBlockRefusal(err) {
work.Ticket.Reject(time.Now())
} else {
work.Ticket.Finish(time.Now())
}
settled = true
if err != nil {
logCentralBlockFailure(work.Value.ip, err)
}
}
func (c *centralActionConsumer) logDropped(previous *uint64) {
total := c.queue.Snapshot(time.Now()).DroppedTotal
if n := total - *previous; n > 0 {
log.Printf("central-intel: %d action(s) incomplete", n)
*previous = total
}
}
package daemon
import (
"crypto/ed25519"
"encoding/hex"
"errors"
"log"
"net"
"os"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/reporting"
"github.com/pidginhost/csm/internal/threatintel"
)
const (
centralRefreshDefault = 6 * time.Hour
centralBlockThreshold = 80
centralChallengeTTL = 6 * time.Hour
centralBlockTTL = 24 * time.Hour
centralActionQueue = 1024
)
type centralQueuedAction struct {
findingID string
decision reporting.Decision
ip string
}
// documentationNets are reserved/non-routable ranges (RFC 5737 documentation,
// RFC 3849 IPv6 documentation, RFC 2544 benchmarking) that must never be acted
// on; they are not routable real attackers.
var documentationNets = mustCIDRs(
"192.0.2.0/24", "198.51.100.0/24", "203.0.113.0/24", "198.18.0.0/15", "2001:db8::/32",
)
func mustCIDRs(cidrs ...string) []*net.IPNet {
out := make([]*net.IPNet, 0, len(cidrs))
for _, c := range cidrs {
if _, n, err := net.ParseCIDR(c); err == nil {
out = append(out, n)
}
}
return out
}
// startCentralConsume wires the central scored-set consumer: it pulls and
// verifies the signed set on an interval and installs alert.CentralHook so a
// finding whose IP is in the set is escalated per the configured action. It
// returns the refresh loop, or nil when disabled/misconfigured.
func (d *Daemon) startCentralConsume() func() {
alert.SetCentralHook(nil)
cc := d.cfg.Reputation.Central
if !cc.Enabled {
return nil
}
if cc.SetURL == "" {
log.Printf("central-intel: enabled but set_url is empty; consumer stays off")
return nil
}
pubHex := os.Getenv(cc.PubkeyEnv)
if raw, err := hex.DecodeString(pubHex); err != nil || len(raw) != ed25519.PublicKeySize {
log.Printf("central-intel: %s must hold a 64-hex-char Ed25519 public key; consumer stays off", cc.PubkeyEnv)
return nil
}
policy := reporting.ParseAction(cc.Action)
if cc.Action != "" && !reporting.IsValidAction(cc.Action) {
log.Printf("central-intel: unrecognized action %q, defaulting to challenge", cc.Action)
}
threshold := cc.BlockThreshold
if threshold <= 0 {
threshold = centralBlockThreshold
}
interval := centralRefreshDefault
if cc.RefreshInterval != "" {
if d2, err := time.ParseDuration(cc.RefreshInterval); err == nil && d2 > 0 {
interval = d2
}
}
store := reporting.NewCentralStore(reporting.NewPuller(nil, cc.SetURL, pubHex))
firebreak := d.centralFirebreak()
consumer := newCentralActionConsumer(d.stopCh, centralActionQueue, interval, store.Refresh, d.performCentralAction)
d.registerQueueSource("central", consumer)
alert.SetCentralHook(func(f alert.Finding) {
a, ok := d.planCentralAction(store, policy, threshold, firebreak, f)
if !ok {
return
}
consumer.enqueue(a)
})
log.Printf("central-intel: enabled (action=%s, threshold=%d, refresh=%s)", policy, threshold, interval)
return func() {
defer alert.SetCentralHook(nil)
consumer.run()
}
}
// applyCentral escalates a finding's IP when it appears in the central set. A
// finding firing on the IP is the node's local corroboration. Firebreaks and
// the action policy gate what happens; central data never blocks on its own.
func (d *Daemon) applyCentral(store *reporting.CentralStore, action reporting.Action, threshold int, firebreak func(string) bool, f alert.Finding) {
a, ok := d.planCentralAction(store, action, threshold, firebreak, f)
if !ok {
return
}
if err := d.performCentralAction(a); err != nil {
logCentralBlockFailure(a.ip, err)
}
}
func (d *Daemon) planCentralAction(store *reporting.CentralStore, action reporting.Action, threshold int, firebreak func(string) bool, f alert.Finding) (centralQueuedAction, bool) {
// Response and coverage-health findings are not independent attacker
// signals. Feeding them back into the consumer can schedule a redundant
// block or attribute service degradation to an unrelated source IP.
if f.Check == "auto_block" || f.Check == "reputation_quota_exhausted" || f.Check == "threat_feed_stale" {
return centralQueuedAction{}, false
}
ip := f.SourceIP
if ip == "" {
return centralQueuedAction{}, false
}
entry, found := store.Lookup(ip)
dec := reporting.Decide(reporting.DecisionInput{
Found: found,
Score: entry.Score,
Protected: firebreak(ip),
LocallyCorroborated: true, // a finding fired on this IP
}, action, threshold)
if dec == reporting.DecisionIgnore {
return centralQueuedAction{}, false
}
return centralQueuedAction{decision: dec, ip: ip, findingID: alert.FindingID(f)}, true
}
func (d *Daemon) performCentralAction(a centralQueuedAction) error {
switch a.decision {
case reporting.DecisionChallenge:
if d.ipList != nil {
d.ipList.AddNonEscalating(a.ip, "central-intel", centralChallengeTTL)
}
case reporting.DecisionBlock:
res, err := checks.ApplyBlock(d.currentCfg(), checks.ApplyBlockRequest{
IP: a.ip,
EngineReason: centralIntelBlockReason,
Reason: centralIntelBlockReason,
TTL: centralBlockTTL,
Source: checks.BlockSourceCentral,
FindingID: a.findingID,
})
d.recordAppliedBlocks(res.Findings)
if err != nil {
return err
}
log.Printf("central-intel: block %s outcome: %s", a.ip, res.Outcome)
}
return nil
}
func logCentralBlockFailure(ip string, err error) {
// Protected IPs are never blockable and a host without a firewall
// engine cannot block; both are expected, not failures.
if isCentralBlockRefusal(err) {
return
}
log.Printf("central-intel: block %s failed: %v", ip, err)
}
func isCentralBlockRefusal(err error) bool {
return isProtectedIPRefusal(err) || errors.Is(err, checks.ErrNoIPBlocker)
}
// centralFirebreak returns a predicate that reports whether an IP must never be
// acted on from central data: loopback/unspecified/private, documentation
// ranges, or an operator infra_ips entry.
func (d *Daemon) centralFirebreak() func(string) bool {
infraEntries := d.cfg.InfraIPs
if d.cfg.Firewall != nil {
infraEntries = mergeInfraIPs(d.cfg.InfraIPs, d.cfg.Firewall.InfraIPs)
}
var infra []*net.IPNet
for _, raw := range infraEntries {
if _, n, err := net.ParseCIDR(raw); err == nil {
infra = append(infra, n)
continue
}
if ip := net.ParseIP(raw); ip != nil {
bits := 32
if ip.To4() == nil {
bits = 128
}
infra = append(infra, &net.IPNet{IP: ip, Mask: net.CIDRMask(bits, bits)})
}
}
return func(s string) bool {
ip := net.ParseIP(s)
if ip == nil {
return true // unparseable: never act
}
if ip.IsLoopback() || ip.IsUnspecified() || ip.IsPrivate() || ip.IsLinkLocalUnicast() {
return true
}
for _, n := range documentationNets {
if n.Contains(ip) {
return true
}
}
for _, n := range infra {
if n.Contains(ip) {
return true
}
}
// A Cloudflare edge or a verified crawler in the scored set is not
// an attacker to act on: challenging or blocking it hits every
// visitor behind the edge, or delists the site.
if checks.IsCloudflareIP(ip) || threatintel.IPInAnyVerifiedBotRange(ip) {
return true
}
return false
}
}
package daemon
import (
"fmt"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/eximlog"
"github.com/pidginhost/csm/internal/obs"
)
// --- Cloud-relay compromise detection -----------------------------------
//
// Detects the pattern where a mailbox's SMTP AUTH credentials are being
// abused from a rented botnet of cloud VMs. Characteristic signature:
// multiple authenticated sends in a short window, from several distinct
// cloud-provider IPs (Google Cloud, AWS, Azure, etc.), for the SAME
// mailbox user.
//
// A normal user on a residential/ISP IP never matches (cloud-PTR check).
// A legitimate self-hosted script on a single VPS won't match either
// (requires ≥2 distinct source IPs within the window). Credential abuse
// from a rotating fleet trips the detector within minutes of the first
// distinct IPs showing up.
//
// Action: one Critical finding per user per window. Customer-impacting
// response actions follow the configured auto-response and dry-run settings.
// The finding message embeds a source IP so autoblock can evaluate it.
// cloudProviderPTRSuffixes is an intentionally-conservative list of
// hostname suffixes that strongly indicate a cloud-VM source. Adding to
// this list increases detection coverage; removing reduces it. Keep
// entries specific enough to avoid catching ISP-transit ASNs.
var cloudProviderPTRSuffixes = []string{
// Google Cloud Platform
".googleusercontent.com",
".gce.internal",
// AWS EC2 — public PTRs only. `.compute.internal` is intentionally
// excluded: it is an AWS VPC-internal PTR but also appears on
// corporate VPN and self-hosted lab networks that have nothing to
// do with AWS, so matching it would risk suspending mailboxes
// that just happen to have an internal-looking reverse DNS.
".compute.amazonaws.com",
".compute-1.amazonaws.com",
// Microsoft Azure
".cloudapp.net",
".cloudapp.azure.com",
// Oracle Cloud
".oraclecloud.com",
".oraclevcn.com",
// DigitalOcean
".digitalocean.com",
".digitaloceanspaces.com",
// Linode / Akamai Cloud
".members.linode.com",
".linodeusercontent.com",
// Vultr
".vultr.com",
".vultrusercontent.com",
// Hetzner
".hetzner.com",
".your-server.de",
// OVH / OVHcloud
".ovh.net",
".ovhcloud.com",
".ovh.ca",
".ovh.us",
// Contabo
".contabo.net",
".contabo.host",
".contaboserver.net",
}
// isCloudRelayAllowed reports whether the AUTH user is opted out of the
// email_cloud_relay_abuse detector via the operator-managed allowlists.
// users matches whole mailboxes; domains matches the domain part. Both
// comparisons are case-insensitive. An empty user is never considered
// allowed (defense against malformed log lines).
func isCloudRelayAllowed(user string, users, domains []string) bool {
if user == "" {
return false
}
for _, u := range users {
if strings.EqualFold(user, u) {
return true
}
}
if len(domains) == 0 {
return false
}
at := strings.LastIndexByte(user, '@')
if at < 0 || at >= len(user)-1 {
return false
}
dom := user[at+1:]
for _, d := range domains {
if strings.EqualFold(dom, d) {
return true
}
}
return false
}
// isCloudProviderPTR reports whether the given PTR hostname belongs to a
// recognized public-cloud provider. Case-insensitive suffix match.
func isCloudProviderPTR(ptr string) bool {
if ptr == "" {
return false
}
p := strings.ToLower(ptr)
for _, suffix := range cloudProviderPTRSuffixes {
if strings.HasSuffix(p, suffix) {
return true
}
}
return false
}
// extractEximHostname parses the H=<hostname> field from an exim log line.
// The field often looks like "H=hostname.example (helo.string) [IP]:port"
// — we want the PTR-derived hostname before the HELO-in-parens.
func extractEximHostname(line string) string {
idx, ok := eximlog.HFieldStart(line)
if !ok {
return ""
}
rest := line[idx:]
// Terminate at first space, tab, or opening paren (HELO string).
end := len(rest)
for i, r := range rest {
if r == ' ' || r == '\t' || r == '(' {
end = i
break
}
}
return strings.TrimSpace(rest[:end])
}
// cloudRelayWindow tracks authenticated sends from cloud IPs for one user.
// Bounded so memory can't grow unbounded from a misbehaving log stream.
type cloudRelayWindow struct {
mu sync.Mutex
events []cloudRelayEvent
firedAt time.Time // last Critical emission; dedup guard
lastEvent time.Time // last append — used to garbage-collect idle entries
}
type cloudRelayEvent struct {
at time.Time
ip string
ptr string
}
// cloudRelayWindows tracks per-user cloud-relay activity.
var cloudRelayWindows sync.Map // map[string]*cloudRelayWindow
func lockCloudRelayWindowForUpdate(user string, now time.Time) *cloudRelayWindow {
for {
val, _ := cloudRelayWindows.LoadOrStore(user, &cloudRelayWindow{lastEvent: now})
w, ok := val.(*cloudRelayWindow)
if !ok {
cloudRelayWindows.Delete(user)
continue
}
w.mu.Lock()
// Re-check while holding w.mu so the eviction sweep cannot delete
// a stale entry and orphan a parser update made through that pointer.
current, ok := cloudRelayWindows.Load(user)
currentWindow, currentOK := current.(*cloudRelayWindow)
if ok && currentOK && currentWindow == w {
return w
}
w.mu.Unlock()
}
}
// Detection thresholds. Two OR-combined signals within the same 60-min
// sliding window:
//
// A. Multi-IP burst: ≥ cloudRelayMinEvents sends from ≥ cloudRelayMinDistinctIP
// distinct cloud IPs. Catches rented-fleet abuse rotating IPs per-send
// (the typical credential-stuffing spam pattern).
//
// B. Volume burst: ≥ cloudRelayHighVolumeEvents sends regardless of
// distinct-IP count. Catches paced attacks that deliberately use one
// cloud IP per day to evade signal A. Threshold sits well above any
// legitimate SaaS integration seen on production (SmartBill ~2/hr,
// Nylas ~2/hr, WP transactional ≤3/hr).
//
// Tuning rationale: these values were chosen from the Apr 2026 incident
// analysis. A single-mailbox user with a legit single-VPS cron averaging
// ≤14 mails/hr stays silent; anything above that is either a compromised
// relay or a SaaS integration that should be added to
// `email_protection.high_volume_senders`.
const (
cloudRelayWindow_ = 60 * time.Minute
cloudRelayMinEvents = 3
cloudRelayMinDistinctIP = 2
cloudRelayHighVolumeEvents = 15
cloudRelayDedupCooldown = 60 * time.Minute
cloudRelayMaxEvents = 256 // per-user cap; prevents unbounded growth
)
// parseCloudRelayFinding evaluates an exim acceptance line for the
// cloud-relay compromise pattern. Called from parseEximLogLine. Returns
// zero or one finding. Never auto-suspends on its own — emits a finding
// whose Message embeds the source IP; the existing autoblock + suspend
// pipeline picks it up by check name.
func parseCloudRelayFinding(line string, cfg *config.Config) (findings []alert.Finding) {
// Only care about authenticated outbound acceptance lines.
if !strings.Contains(line, " <= ") || !strings.Contains(line, "A=dovecot_") {
return nil
}
user := extractAuthUser(line)
if user == "" {
return nil
}
if isHighVolumeSender(user, cfg.EmailProtection.HighVolumeSenders) {
return nil
}
if isCloudRelayAllowed(user, cfg.EmailProtection.CloudRelay.AllowUsers, cfg.EmailProtection.CloudRelay.AllowDomains) {
return nil
}
ptr := extractEximHostname(line)
if !isCloudProviderPTR(ptr) {
return nil
}
ip := eximlog.ClientIP(line)
if ip == "" {
// Without an IP we can't dedup distinct sources; bail silently
// so we don't count half-parsed records toward the threshold.
return nil
}
// Registered before the unlock defer so owner I/O runs after it.
defer func() { stampMailAccountOwner(findings, user) }()
now := time.Now()
w := lockCloudRelayWindowForUpdate(user, now)
defer w.mu.Unlock()
// Prune anything older than the window.
cutoff := now.Add(-cloudRelayWindow_)
kept := w.events[:0]
for _, e := range w.events {
if e.at.After(cutoff) {
kept = append(kept, e)
}
}
w.events = kept
// Append this event (with cap).
if len(w.events) < cloudRelayMaxEvents {
w.events = append(w.events, cloudRelayEvent{at: now, ip: ip, ptr: ptr})
}
w.lastEvent = now
// Already fired recently for this user? Dedup.
if !w.firedAt.IsZero() && now.Sub(w.firedAt) < cloudRelayDedupCooldown {
return nil
}
// Evaluate thresholds.
distinctIPs := make(map[string]struct{}, len(w.events))
for _, e := range w.events {
distinctIPs[e.ip] = struct{}{}
}
multiIPBurst := len(w.events) >= cloudRelayMinEvents && len(distinctIPs) >= cloudRelayMinDistinctIP
volumeBurst := len(w.events) >= cloudRelayHighVolumeEvents
if !multiIPBurst && !volumeBurst {
return nil
}
w.firedAt = now
// Build an IP list for the details (newest first, deduped).
seen := make(map[string]struct{}, len(w.events))
ips := make([]string, 0, len(distinctIPs))
for i := len(w.events) - 1; i >= 0; i-- {
e := w.events[i]
if _, dup := seen[e.ip]; dup {
continue
}
seen[e.ip] = struct{}{}
ips = append(ips, e.ip)
}
// The block source is carried in the structured SourceIP field below; the
// IP is included in the message only for operator-facing readability.
message := fmt.Sprintf(
"Email account %s sent %d authenticated messages from %d cloud-provider IPs in %d minutes - credentials compromised - from %s",
user, len(w.events), len(distinctIPs), int(cloudRelayWindow_.Minutes()), ips[0],
)
details := fmt.Sprintf(
"Authenticated SMTP submissions for %s in the last %d minutes:\n"+
" total sends: %d\n"+
" distinct source IPs: %d\n"+
" most recent PTR: %s\n"+
" recent IPs: %s\n\n"+
"Legitimate users do not send mail from rented cloud VMs. "+
"This pattern is characteristic of credential abuse by a bulk "+
"phishing operator. Outgoing mail hold and source-IP blocking "+
"follow the configured auto-response and dry-run settings.",
user,
int(cloudRelayWindow_.Minutes()),
len(w.events),
len(distinctIPs),
ptr,
strings.Join(truncateIPList(ips, 8), ", "),
)
mailbox, domain, _ := splitMailAccount(user)
return []alert.Finding{{
Severity: alert.Critical,
Check: "email_cloud_relay_abuse",
Message: message,
Details: truncateDaemon(details, 800),
SourceIP: ips[0],
Mailbox: mailbox,
Domain: domain,
}}
}
func handleCloudRelayCredentialAbuse(cfg *config.Config, authUser string) {
if domain := extractDomainFromEmail(authUser); domain != "" {
maybeHoldOutgoingMail(cfg, authUser)
// This is correlation state, not an auto-response action; keep it
// active even when the mail hold is disabled or dry-run gated.
RecordCompromisedDomain(domain)
}
}
func truncateIPList(ips []string, n int) []string {
if len(ips) <= n {
return ips
}
return ips[:n]
}
// cloudRelayEvictWindow is how long a per-user cloudRelayWindow can stay
// idle before it is evicted. Picked at 2x cloudRelayWindow_ so a user who
// just barely cleared the threshold does not lose their entry before the
// dedup cooldown can suppress a repeat.
const cloudRelayEvictWindow = 2 * cloudRelayWindow_
// StartCloudRelayEviction periodically prunes per-user cloud-relay
// windows that have not seen activity in cloudRelayEvictWindow. Without
// this sweep the cloudRelayWindows sync.Map grows linearly with every
// distinct authenticated sender ever seen, including users deleted by
// the operator. Pairs with the firedAt dedup guard so repeat alerts on
// long-lived attackers still pass.
func StartCloudRelayEviction(stopCh <-chan struct{}) {
obs.Go("cloud-relay-eviction", func() {
ticker := time.NewTicker(10 * time.Minute)
defer ticker.Stop()
for {
select {
case <-stopCh:
return
case now := <-ticker.C:
evictCloudRelayWindows(now)
}
}
})
}
// evictCloudRelayWindows deletes per-user entries whose lastEvent is
// older than cloudRelayEvictWindow. Safe to call from tests.
func evictCloudRelayWindows(now time.Time) {
cutoff := now.Add(-cloudRelayEvictWindow)
cloudRelayWindows.Range(func(key, val any) bool {
w, ok := val.(*cloudRelayWindow)
if !ok {
cloudRelayWindows.Delete(key)
return true
}
w.mu.Lock()
if w.lastEvent.Before(cutoff) {
cloudRelayWindows.CompareAndDelete(key, val)
}
w.mu.Unlock()
return true
})
}
package daemon
import (
"bufio"
"errors"
"fmt"
"io"
"os"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/eximlog"
"github.com/pidginhost/csm/internal/store"
)
// --- Retrospective scan for cloud-relay credential abuse -----------------
//
// The realtime watcher in cloud_relay.go only catches traffic arriving
// AFTER CSM starts. A credential-abuse spam run can go for weeks before
// an operator notices (a real incident saw 230 outbound sends from one
// compromised account over 20 days before being flagged). This scanner replays
// the last N hours of exim_mainlog through the same rule, so on CSM
// startup any in-progress or recent compromise is surfaced immediately.
//
// Runs once at daemon startup. Not part of the tiered-check registry
// because it is purely an event-log replay; the realtime watcher owns
// live state thereafter.
// cloudRelayScanPathDefault is where cPanel exim writes its mainlog.
const cloudRelayScanPathDefault = "/var/log/exim_mainlog"
// Memory caps for the retro scan. A compromised account or a crafted log
// can otherwise grow byUser without bound. The thresholds are set well
// above the volume detector (cloudRelayHighVolumeEvents=15 in a 60-min
// window) so legitimate detection is unaffected; the caps only kick in
// for pathological volumes that would balloon memory.
const (
cloudRelayScanMaxEventsPerUser = 5000
cloudRelayScanMaxUsers = 10000
)
// CloudRelayScanPath is the log file path scanned at startup. Exported
// via var (not const) for tests.
var CloudRelayScanPath = cloudRelayScanPathDefault
// ScanEximHistoryForCloudRelay replays the tail of exim_mainlog for the
// last `lookback` duration and returns a finding per mailbox that
// exceeds the cloud-relay thresholds. Safe to call from goroutines.
//
// The scanner respects EmailProtection.HighVolumeSenders and the
// detector-scoped EmailProtection.CloudRelay.AllowUsers / .AllowDomains
// allowlists, mirrors the realtime detector's thresholds exactly, and
// uses a per-user persistent marker in the global store to avoid
// re-emitting the same finding on successive restarts.
func ScanEximHistoryForCloudRelay(cfg *config.Config, logPath string, now time.Time, lookback time.Duration) []alert.Finding {
if logPath == "" {
logPath = CloudRelayScanPath
}
// #nosec G304 -- logPath is operator-configured / hardcoded to cPanel default.
f, err := os.Open(logPath)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
since := now.Add(-lookback)
byUser := make(map[string]*cloudRelayScanAccumulator)
// Use bufio.Reader rather than bufio.Scanner: an exim line can
// occasionally exceed whatever fixed Scanner buffer we set (e.g.
// a spam run with a huge Base64 subject). Scanner returns an
// ErrTooLong which aborts the whole loop — missing every later
// compromise event. Reader.ReadString lets us skip oversized
// lines and keep going.
reader := bufio.NewReaderSize(f, 256*1024)
for {
line, rerr := reader.ReadString('\n')
if len(line) > 0 {
if line[len(line)-1] == '\n' {
line = line[:len(line)-1]
}
processCloudRelayScanLine(line, cfg, since, byUser)
}
if rerr == nil {
continue
}
if errors.Is(rerr, io.EOF) {
break
}
if errors.Is(rerr, bufio.ErrBufferFull) {
// Line longer than 256 KB — drain it and move on.
// Real exim acceptance lines are well under 10 KB;
// anything longer is almost certainly a pathological
// subject we can't usefully parse anyway.
if drainErr := drainUntilNewline(reader); drainErr != nil {
break
}
continue
}
// Any other I/O error: stop cleanly, don't panic.
break
}
users := make([]string, 0, len(byUser))
for u := range byUser {
users = append(users, u)
}
sort.Strings(users) // stable finding order
var findings []alert.Finding
for _, user := range users {
acc := byUser[user]
if acc == nil || !acc.reportable {
continue
}
maxSends, maxDistinctIPs, fireAt, peakPTR := acc.bestSends, acc.bestDistinctIPs, acc.bestAt, acc.bestPTR
multiIP := maxSends >= cloudRelayMinEvents && maxDistinctIPs >= cloudRelayMinDistinctIP
volume := maxSends >= cloudRelayHighVolumeEvents
if !multiIP && !volume {
continue
}
// Persistent dedup: skip if we've already fired for this user
// and no new event has landed since then.
latestEvent := acc.latestEvent
if alreadyReportedRetro(user, latestEvent) {
continue
}
ips := acc.recentIPs
if len(ips) == 0 {
continue
}
recentIP := ips[0]
msg := fmt.Sprintf(
"RETRO: account %s sent %d authenticated messages from %d cloud-provider IPs (peak 60-min burst) in the last %d hours - credentials compromised - from %s",
user, maxSends, maxDistinctIPs, int(lookback.Hours()), recentIP,
)
details := fmt.Sprintf(
"Retrospective exim_mainlog scan at %s found a cloud-relay pattern:\n"+
" user: %s\n"+
" total cloud-PTR sends (%dh): %d\n"+
" peak 60-min window: %d sends / %d distinct IPs ending at %s\n"+
" peak PTR: %s\n"+
" distinct source IPs observed: %s\n\n"+
"Outgoing mail hold and source-IP blocking follow the configured "+
"auto-response and dry-run settings. Older IPs are left as context "+
"because rented-fleet addresses tend to be recycled outside a 2-hour window.",
now.Format("2006-01-02 15:04:05"),
user,
int(lookback.Hours()),
acc.total,
maxSends, maxDistinctIPs, fireAt.Format("2006-01-02 15:04:05"),
peakPTR,
strings.Join(ips, ", "),
)
mailbox, domain, _ := splitMailAccount(user)
tenant := mailAccountOwner(user)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_cloud_relay_abuse",
Message: msg,
Details: truncateDaemon(details, 900),
Timestamp: now,
SourceIP: recentIP,
Mailbox: mailbox,
Domain: domain,
TenantID: tenant,
})
markReportedRetro(user, latestEvent)
}
return findings
}
// processCloudRelayScanLine parses a single exim log line and, if it is
// an authenticated cloud-PTR acceptance within the lookback window,
// records it under the AUTH user in `byUser`.
func processCloudRelayScanLine(line string, cfg *config.Config, since time.Time, byUser map[string]*cloudRelayScanAccumulator) {
if !strings.Contains(line, " <= ") || !strings.Contains(line, "A=dovecot_") {
return
}
ts, ok := parseEximTimestamp(line)
if !ok || ts.Before(since) {
return
}
user := extractAuthUser(line)
if user == "" || isHighVolumeSender(user, cfg.EmailProtection.HighVolumeSenders) {
return
}
if isCloudRelayAllowed(user, cfg.EmailProtection.CloudRelay.AllowUsers, cfg.EmailProtection.CloudRelay.AllowDomains) {
return
}
ptr := extractEximHostname(line)
if !isCloudProviderPTR(ptr) {
return
}
ip := eximlog.ClientIP(line)
if ip == "" {
return
}
acc, exists := byUser[user]
if !exists {
if len(byUser) >= cloudRelayScanMaxUsers {
pruneCloudRelayScanUsers(byUser, ts)
}
if len(byUser) >= cloudRelayScanMaxUsers {
evictOldestCloudRelayScanUser(byUser)
}
if len(byUser) >= cloudRelayScanMaxUsers {
return
}
acc = newCloudRelayScanAccumulator()
byUser[user] = acc
}
acc.record(cloudRelayScanEvent{at: ts, ip: ip, ptr: ptr})
}
// drainUntilNewline reads from reader and discards bytes until a newline
// is consumed or EOF is hit. Returns io.EOF if the reader is exhausted.
func drainUntilNewline(reader *bufio.Reader) error {
for {
_, err := reader.ReadSlice('\n')
if err == nil {
return nil
}
if errors.Is(err, bufio.ErrBufferFull) {
// Still inside the oversized line — keep draining.
continue
}
return err
}
}
// cloudRelayScanEvent is a single timestamped cloud-PTR AUTH send
// replayed from the log.
type cloudRelayScanEvent struct {
at time.Time
ip string
ptr string
}
type cloudRelayScanAccumulator struct {
events []cloudRelayScanEvent
ipCounts map[string]int
recentIPs []string
total int
latestEvent time.Time
bestSends int
bestDistinctIPs int
bestAt time.Time
bestPTR string
reportable bool
}
func newCloudRelayScanAccumulator() *cloudRelayScanAccumulator {
return &cloudRelayScanAccumulator{
ipCounts: make(map[string]int),
}
}
func (acc *cloudRelayScanAccumulator) record(event cloudRelayScanEvent) {
acc.total++
if acc.latestEvent.IsZero() || event.at.After(acc.latestEvent) {
acc.latestEvent = event.at
}
acc.rememberRecentIP(event.ip)
cutoff := event.at.Add(-cloudRelayWindow_)
drop := 0
for drop < len(acc.events) && acc.events[drop].at.Before(cutoff) {
acc.removeWindowIP(acc.events[drop].ip)
drop++
}
if drop > 0 {
clear(acc.events[:drop])
acc.events = acc.events[drop:]
}
if len(acc.events) >= cloudRelayScanMaxEventsPerUser {
acc.removeWindowIP(acc.events[0].ip)
var zero cloudRelayScanEvent
acc.events[0] = zero
acc.events = acc.events[1:]
}
acc.events = append(acc.events, event)
acc.ipCounts[event.ip]++
sends := len(acc.events)
distinctIPs := len(acc.ipCounts)
if sends > acc.bestSends || (sends == acc.bestSends && distinctIPs > acc.bestDistinctIPs) {
acc.bestSends = sends
acc.bestDistinctIPs = distinctIPs
acc.bestAt = event.at
acc.bestPTR = event.ptr
}
if acc.bestSends >= cloudRelayHighVolumeEvents ||
(acc.bestSends >= cloudRelayMinEvents && acc.bestDistinctIPs >= cloudRelayMinDistinctIP) {
acc.reportable = true
}
}
func (acc *cloudRelayScanAccumulator) removeWindowIP(ip string) {
count := acc.ipCounts[ip]
if count <= 1 {
delete(acc.ipCounts, ip)
return
}
acc.ipCounts[ip] = count - 1
}
func (acc *cloudRelayScanAccumulator) rememberRecentIP(ip string) {
if ip == "" {
return
}
for i, existing := range acc.recentIPs {
if existing != ip {
continue
}
copy(acc.recentIPs[1:i+1], acc.recentIPs[:i])
acc.recentIPs[0] = ip
return
}
acc.recentIPs = append(acc.recentIPs, "")
copy(acc.recentIPs[1:], acc.recentIPs[:len(acc.recentIPs)-1])
acc.recentIPs[0] = ip
if len(acc.recentIPs) > 10 {
acc.recentIPs = acc.recentIPs[:10]
}
}
func pruneCloudRelayScanUsers(byUser map[string]*cloudRelayScanAccumulator, now time.Time) {
cutoff := now.Add(-cloudRelayWindow_)
for user, acc := range byUser {
if acc == nil || (!acc.reportable && acc.latestEvent.Before(cutoff)) {
delete(byUser, user)
}
}
}
func evictOldestCloudRelayScanUser(byUser map[string]*cloudRelayScanAccumulator) {
var oldestUser string
var oldestSeen time.Time
for user, acc := range byUser {
if acc == nil {
delete(byUser, user)
return
}
if acc.reportable {
continue
}
if oldestUser == "" || acc.latestEvent.Before(oldestSeen) {
oldestUser = user
oldestSeen = acc.latestEvent
}
}
if oldestUser != "" {
delete(byUser, oldestUser)
}
}
// maxCloudRelayBurst finds the strongest 60-min window in a sorted event
// list. Returns (sends, distinctIPs, peakEnd, peakPTR), where peakEnd is
// the timestamp of the LAST event in the best window so operators see
// when the burst peaked, not when it started.
func maxCloudRelayBurst(events []cloudRelayScanEvent) (int, int, time.Time, string) {
if len(events) == 0 {
return 0, 0, time.Time{}, ""
}
bestSends, bestDistinct := 0, 0
var bestAt time.Time
var bestPTR string
left := 0
ipCounts := make(map[string]int)
for right, event := range events {
ipCounts[event.ip]++
for event.at.Sub(events[left].at) > cloudRelayWindow_ {
leftIP := events[left].ip
if ipCounts[leftIP] <= 1 {
delete(ipCounts, leftIP)
} else {
ipCounts[leftIP]--
}
left++
}
sends := right - left + 1
distinct := len(ipCounts)
// "Best" = highest send count; tie-break by distinct IPs.
if sends > bestSends || (sends == bestSends && distinct > bestDistinct) {
bestSends = sends
bestDistinct = distinct
bestAt = event.at
bestPTR = event.ptr
}
}
return bestSends, bestDistinct, bestAt, bestPTR
}
// parseEximTimestamp extracts the "YYYY-MM-DD HH:MM:SS" timestamp prefix
// from an exim log line. Returns false on any parse failure.
func parseEximTimestamp(line string) (time.Time, bool) {
if len(line) < 19 {
return time.Time{}, false
}
t, err := time.ParseInLocation("2006-01-02 15:04:05", line[:19], time.Local)
if err != nil {
return time.Time{}, false
}
return t, true
}
// alreadyReportedRetro returns true when the latest event for this user
// is older than or equal to the persisted marker (meaning: nothing new
// since we last alerted).
func alreadyReportedRetro(user string, latestEvent time.Time) bool {
db := store.Global()
if db == nil {
return false
}
raw := db.GetMetaString("cloudrelay_retro:" + user)
if raw == "" {
return false
}
prev, err := time.Parse(time.RFC3339, raw)
if err != nil {
return false
}
return !latestEvent.After(prev)
}
// extractSenderFromCloudRelayMessage pulls the sender mailbox out of a
// finding message emitted by ScanEximHistoryForCloudRelay. Returns ""
// when the message is not from this check (defensive — never panics on
// unexpected input).
func extractSenderFromCloudRelayMessage(msg string) string {
const marker = "account "
idx := strings.Index(msg, marker)
if idx < 0 {
return ""
}
rest := msg[idx+len(marker):]
sp := strings.IndexByte(rest, ' ')
if sp <= 0 {
return ""
}
candidate := rest[:sp]
if !strings.Contains(candidate, "@") {
return ""
}
return candidate
}
func markReportedRetro(user string, latestEvent time.Time) {
db := store.Global()
if db == nil {
return
}
_ = db.SetMetaString("cloudrelay_retro:"+user, latestEvent.Format(time.RFC3339))
}
//go:build linux
package daemon
import (
"crypto/sha256"
"encoding/hex"
"golang.org/x/sys/unix"
)
// hashEventFD returns the SHA-256 of exactly expectedSize bytes behind an
// event descriptor. prefix must be the bytes the scanner already read at
// offset zero. Reusing them prevents an in-place overwrite from making the CMS
// decision cover different leading bytes than signature and YARA scanning.
// The fixed expected size bounds the work and prevents a concurrently growing
// file from keeping an analyzer busy indefinitely. An empty result means the
// file could not be read consistently.
func hashEventFD(fd int, prefix []byte, expectedSize int64) string {
if expectedSize < int64(len(prefix)) || expectedSize < 0 {
return ""
}
var before unix.Stat_t
if err := unix.Fstat(fd, &before); err != nil || before.Size != expectedSize {
return ""
}
h := sha256.New()
_, _ = h.Write(prefix)
buf := make([]byte, 64*1024)
offset := int64(len(prefix))
for offset < expectedSize {
want := int64(len(buf))
if remaining := expectedSize - offset; remaining < want {
want = remaining
}
n, err := unix.Pread(fd, buf[:int(want)], offset)
if n > 0 {
_, _ = h.Write(buf[:n])
offset += int64(n)
}
if err != nil {
if err == unix.EINTR && n == 0 {
continue
}
return ""
}
if n == 0 {
return ""
}
}
var after unix.Stat_t
if err := unix.Fstat(fd, &after); err != nil || after.Size != expectedSize {
return ""
}
return hex.EncodeToString(h.Sum(nil))
}
//go:build !(linux && bpf)
package daemon
import (
"context"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
)
// connectionBPF is the no-tag placeholder for the BPF cgroup/connect backend.
// The real type with map handles, links, and the ringbuf reader lives in
// connection_bpf.go behind //go:build linux && bpf. On any other build, the
// coordinator never reaches this stub: startConnectionBPF returns
// bpf.ErrNotBuilt before a value is constructed.
type connectionBPF struct{}
func (c *connectionBPF) Mode() string { return "bpf" }
func (c *connectionBPF) EventCount() uint64 { return 0 }
func (c *connectionBPF) Run(_ context.Context) {}
func startConnectionBPF(_ context.Context, _ chan<- alert.Finding, _ *config.Config) (*connectionBPF, error) {
return nil, bpf.ErrNotBuilt
}
package daemon
import (
"context"
"errors"
"fmt"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/verdict"
)
var errBPFVerdictDisabled = errors.New("verdict callback not configured")
// bpfVerdictEnabled reports whether this event is one the operator's callback
// should be asked about. Cheap enough for the ring-buffer consumer loop; the
// request itself is not, and runs on the enricher's workers.
func bpfVerdictEnabled(cfg *config.Config, ev ConnectionEvent) bool {
if cfg == nil || !cfg.BPFEnforcement.VerdictCallback || !cfg.AutoResponse.VerdictCallback.Enabled {
return false
}
return ev.Decision == 1 || ev.Decision == 2
}
// askBPFVerdict performs one callback. The client is built per call from the
// live configuration so a hot reload of the URL or secret is honoured.
func askBPFVerdict(ctx context.Context, cfg *config.Config, req verdict.Request) (verdict.Response, error) {
if cfg == nil {
return verdict.Response{}, errBPFVerdictDisabled
}
vcCfg := cfg.AutoResponse.VerdictCallback
vc := verdict.New(verdict.Config{
URL: vcCfg.URL,
HMACSecret: vcCfg.HMACSecret,
HMACSecretEnv: vcCfg.HMACSecretEnv,
RequireResponseSignature: vcCfg.RequireResponseSignature,
AllowUnsigned: vcCfg.AllowUnsigned,
Timeout: time.Duration(vcCfg.TimeoutSec) * time.Second,
})
return vc.Ask(ctx, req)
}
// bpfVerdictReason is the callback's reason string for one event.
func bpfVerdictReason(check string, port uint16) string {
return fmt.Sprintf("bpf_enforcement:%s:%d", check, port)
}
func appendFindingDetail(f *alert.Finding, detail string) {
if detail == "" {
return
}
if f.Details == "" {
f.Details = detail
return
}
f.Details += ", " + detail
}
// Code generated by bpf2go; DO NOT EDIT.
//go:build 386 || amd64
package connection_bpfprog
import (
"bytes"
_ "embed"
"fmt"
"io"
"structs"
"github.com/cilium/ebpf"
)
type ConnectionConnEvent struct {
_ structs.HostLayout
Uid uint32
Pid uint32
Family uint32
DstPort uint32
DstIp4 uint32
DstIp6 [16]uint8
Comm [16]uint8
Decision uint32
}
type ConnectionCsmQueueStats struct {
_ structs.HostLayout
Lost uint64
Submitted uint64
}
type ConnectionPolicyState struct {
_ structs.HostLayout
Enforce uint32
DryRun uint32
ProtectedPorts uint32
}
// Names of all BPF objects in the ELF.
//
// Used for safe lookups in a Collection or CollectionSpec.
const (
ConnectionMapEvents = "events"
ConnectionMapPolicy = "policy"
ConnectionMapProtectedPorts = "protected_ports"
ConnectionMapQueueStats = "queue_stats"
ConnectionMapSafeUids = "safe_uids"
ConnectionProgCsmConnect4 = "csm_connect4"
ConnectionProgCsmConnect6 = "csm_connect6"
ConnectionVarUnused = "unused"
ConnectionVarUnusedPolicy = "unused_policy"
)
// LoadConnection returns the embedded CollectionSpec for Connection.
func LoadConnection() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_ConnectionBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load Connection: %w", err)
}
return spec, err
}
// LoadConnectionObjects loads Connection and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *ConnectionObjects
// *ConnectionPrograms
// *ConnectionMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func LoadConnectionObjects(obj any, opts *ebpf.CollectionOptions) error {
spec, err := LoadConnection()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// ConnectionSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ConnectionSpecs struct {
ConnectionProgramSpecs
ConnectionMapSpecs
ConnectionVariableSpecs
}
// ConnectionProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ConnectionProgramSpecs struct {
CsmConnect4 *ebpf.ProgramSpec `ebpf:"csm_connect4"`
CsmConnect6 *ebpf.ProgramSpec `ebpf:"csm_connect6"`
}
// ConnectionMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ConnectionMapSpecs struct {
Events *ebpf.MapSpec `ebpf:"events"`
Policy *ebpf.MapSpec `ebpf:"policy"`
ProtectedPorts *ebpf.MapSpec `ebpf:"protected_ports"`
QueueStats *ebpf.MapSpec `ebpf:"queue_stats"`
SafeUids *ebpf.MapSpec `ebpf:"safe_uids"`
}
// ConnectionVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ConnectionVariableSpecs struct {
Unused *ebpf.VariableSpec `ebpf:"unused"`
UnusedPolicy *ebpf.VariableSpec `ebpf:"unused_policy"`
}
// ConnectionObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to LoadConnectionObjects or ebpf.CollectionSpec.LoadAndAssign.
type ConnectionObjects struct {
ConnectionPrograms
ConnectionMaps
ConnectionVariables
}
func (o *ConnectionObjects) Close() error {
return _ConnectionClose(
&o.ConnectionPrograms,
&o.ConnectionMaps,
)
}
// ConnectionMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to LoadConnectionObjects or ebpf.CollectionSpec.LoadAndAssign.
type ConnectionMaps struct {
Events *ebpf.Map `ebpf:"events"`
Policy *ebpf.Map `ebpf:"policy"`
ProtectedPorts *ebpf.Map `ebpf:"protected_ports"`
QueueStats *ebpf.Map `ebpf:"queue_stats"`
SafeUids *ebpf.Map `ebpf:"safe_uids"`
}
func (m *ConnectionMaps) Close() error {
return _ConnectionClose(
m.Events,
m.Policy,
m.ProtectedPorts,
m.QueueStats,
m.SafeUids,
)
}
// ConnectionVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to LoadConnectionObjects or ebpf.CollectionSpec.LoadAndAssign.
type ConnectionVariables struct {
Unused *ebpf.Variable `ebpf:"unused"`
UnusedPolicy *ebpf.Variable `ebpf:"unused_policy"`
}
// ConnectionPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to LoadConnectionObjects or ebpf.CollectionSpec.LoadAndAssign.
type ConnectionPrograms struct {
CsmConnect4 *ebpf.Program `ebpf:"csm_connect4"`
CsmConnect6 *ebpf.Program `ebpf:"csm_connect6"`
}
func (p *ConnectionPrograms) Close() error {
return _ConnectionClose(
p.CsmConnect4,
p.CsmConnect6,
)
}
func _ConnectionClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed connection_x86_bpfel.o
var _ConnectionBytes []byte
package daemon
import (
"net"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
)
var (
directSMTPRDNSOnce sync.Once
directSMTPRDNSCache *checks.RDNSCache
)
// rdnsCache is the daemon-wide rDNS cache used by direct SMTP egress
// detection. TTL 30 min, per-lookup deadline 1 second. Resolver wraps
// net.LookupAddr; negative results cached so a slow upstream does not
// stall the connection consumer.
func rdnsCache() *checks.RDNSCache {
directSMTPRDNSOnce.Do(func() {
directSMTPRDNSCache = checks.NewRDNSCache(checks.RDNSCacheConfig{
TTL: 30 * time.Minute,
ResolveDeadline: time.Second,
Resolve: func(ip net.IP) (string, error) {
names, err := net.LookupAddr(ip.String())
if err != nil || len(names) == 0 {
return "", err
}
return strings.TrimSuffix(names[0], "."), nil
},
})
})
return directSMTPRDNSCache
}
// evaluateConnectionEvent runs every per-event detector and returns the
// findings that should be emitted. Pure-ish: no IO and no alertCh
// access. Caller is responsible for attaching process context (which
// MAY do IO via the enricher) and shipping to alertCh.
//
// The function exists in a non-build-tagged file so unit tests can
// drive a synthetic ConnectionEvent through the same policy logic the
// live BPF Run loop uses, without requiring the linux+bpf build tag.
func evaluateConnectionEvent(cfg *config.Config, mta platform.MTAIdents, ev ConnectionEvent, user string) []alert.Finding {
switch ev.Decision {
case 0:
BumpBPFEnforcementDecision(BPFDecisionAllow)
case 1:
BumpBPFEnforcementDecision(BPFDecisionDryRun)
case 2:
BumpBPFEnforcementDecision(BPFDecisionDeny)
}
// Phase 4 note: bpf_enforcement.verdict_callback is applied by the
// BPF Run loop after this evaluator returns. The in-kernel hook
// NEVER waits on HTTP; cgroup/connect is synchronous and a remote
// callback would add latency to every connect.
now := time.Now()
// Phase 3 note: DryRun knobs are not consulted here. Detection runs
// regardless. The knobs gate the Phase 4 auto-response action that
// has not landed yet.
var out []alert.Finding
if checks.DirectSMTPEgressBackendEnabled(cfg, "bpf") {
// Direct SMTP egress (Phase 3). Distinct Check value; the inbound
// smtp_probe meters never see this traffic.
if f, ok := checks.EvaluateDirectSMTPEgress(cfg, checks.DirectSMTPEgressInput{
UID: ev.UID,
User: user,
PID: ev.PID,
Comm: ev.Comm,
DstIP: ev.DstIP,
DstPort: ev.DstPort,
MTA: mta,
}); ok {
if domain := rdnsCache().Lookup(ev.DstIP); domain != "" {
f.Details += ", Domain: " + domain
}
f.Timestamp = now
out = append(out, f)
checks.BumpDirectSMTPEgressFindings()
}
}
// Pre-existing user_outbound_connection detector. SMTP destinations
// are filtered out by checks.safeRemotePorts inside this evaluator,
// so it does not double-fire for a 25/465/587 connect.
if f, ok := checks.EvaluateConnection(cfg, ev.UID, ev.DstIP, ev.DstPort, 0, protoFromFamily(ev.Family), user); ok {
f.Timestamp = now
out = append(out, f)
}
// Bad-ASN egress (host-takeover chain leg). The BPF program emits
// non-root connects only; root egress is covered by the periodic
// /proc/net scan so root-heavy hosts do not flood the ringbuf.
if ev.UID != 0 && cfg.Detection.BadASNOutbound.Enabled {
if lookup := checks.CurrentASNLookup(); lookup != nil {
asn, org := lookup(ev.DstIP.String())
if f, ok := checks.EvaluateBadASNOutbound(cfg, ev.DstIP, asn, org); ok {
checks.AttributeSocketOwner(&f, ev.UID)
f.Timestamp = now
out = append(out, f)
}
}
}
return out
}
// protoFromFamily maps a sockaddr family int to a string label used in
// finding details. Lives here (not in connection_bpf.go) so the
// evaluator helper compiles on darwin without the bpf tag.
func protoFromFamily(f uint32) string {
if f == 10 {
return "tcp6"
}
return "tcp"
}
package daemon
import (
"encoding/binary"
"errors"
"net"
)
// ConnectionEvent is the userspace shape of a struct conn_event emitted by
// the cgroup/connect BPF program. Field layout matches connection.bpf.c
// byte for byte: scalars are little-endian (host order on amd64/arm64),
// dst_ip4 is network-order, dst_ip6 is the raw 16-byte address.
type ConnectionEvent struct {
UID uint32
PID uint32
Family uint32 // AF_INET=2, AF_INET6=10
DstPort uint16 // host order; BPF program calls bpf_ntohs
DstIP net.IP // resolved from dst_ip4 (v4) or dst_ip6 (v6) per Family
Comm string // null-terminated, up to 16 bytes
Decision uint32 // Phase 4: DECISION_* code (0=allow, 1=dry_run_deny, 2=deny)
}
const connectionEventSize = 4 + 4 + 4 + 4 + 4 + 16 + 16 + 4
func decodeConnectionEvent(b []byte) (ConnectionEvent, error) {
if len(b) < connectionEventSize {
return ConnectionEvent{}, errors.New("connection event short buffer")
}
// The BPF program stores dst_port as __u32 for alignment but calls
// bpf_ntohs() before writing, which guarantees the value fits in 16 bits.
// The narrowing is safe by construction.
dstPort := binary.LittleEndian.Uint32(b[12:16]) & 0xffff
ev := ConnectionEvent{
UID: binary.LittleEndian.Uint32(b[0:4]),
PID: binary.LittleEndian.Uint32(b[4:8]),
Family: binary.LittleEndian.Uint32(b[8:12]),
DstPort: uint16(dstPort), // #nosec G115 -- masked to low 16 bits above
}
switch ev.Family {
case 2: // AF_INET
ipv4 := make(net.IP, 4)
binary.BigEndian.PutUint32(ipv4, binary.BigEndian.Uint32(b[16:20]))
ev.DstIP = ipv4
case 10: // AF_INET6
v6 := make(net.IP, 16)
copy(v6, b[20:36])
ev.DstIP = v6
default:
return ConnectionEvent{}, errors.New("unknown family")
}
ev.Comm = nullTerm(b[36 : 36+16])
ev.Decision = binary.LittleEndian.Uint32(b[52:56])
return ev, nil
}
func indexNull(b []byte) int {
for i, c := range b {
if c == 0 {
return i
}
}
return -1
}
// nullTerm returns the prefix of b up to (but not including) the first NUL,
// or the whole slice if there is none. Helper for fixed-size character
// fields the BPF programs emit (comm[16], filename[256], exe[256]).
func nullTerm(b []byte) string {
if i := indexNull(b); i >= 0 {
return string(b[:i])
}
return string(b)
}
package daemon
import (
"context"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
)
// connectionPoller is the userspace fallback. It runs CheckOutboundUserConnections
// on a fixed interval and forwards any findings to the alert channel. Used
// when the BPF backend is unavailable (no bpf tag, kernel rejects program,
// or operator pinned legacy via Detection.ConnectionTrackerBackend).
type connectionPoller struct {
cfg *config.Config
alertCh chan<- alert.Finding
count atomic.Uint64
}
func newConnectionPoller(cfg *config.Config, alertCh chan<- alert.Finding) *connectionPoller {
return &connectionPoller{cfg: cfg, alertCh: alertCh}
}
func (p *connectionPoller) Mode() string { return "legacy" }
func (p *connectionPoller) EventCount() uint64 { return p.count.Load() }
func (p *connectionPoller) Run(ctx context.Context) {
interval := pollerInterval(p.cfg)
t := time.NewTicker(interval)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
findings := checks.CheckOutboundUserConnections(ctx, activeConnectionCfg(p.cfg), nil)
for _, f := range findings {
p.count.Add(1)
if !alert.TryEnqueue(p.alertCh, f) {
csmlog.Warn("connection legacy: alert channel full, dropping finding")
}
}
}
}
}
// pollerInterval returns the configured polling interval, falling back to a
// 30-second default when Detection.ConnectionPollInterval is unset.
func pollerInterval(cfg *config.Config) time.Duration {
if d := cfg.Detection.ConnectionPollInterval; d > 0 {
return d
}
return 30 * time.Second
}
package daemon
import (
"context"
"errors"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/processctx"
)
// StartConnectionTracker selects the active connection-tracker backend based
// on cfg.Detection.ConnectionTrackerBackend and host capability:
//
// "auto" (default) -- try BPF, fall back to legacy polling.
// "bpf" -- require BPF; return nil if unavailable (no fallback).
// "legacy" -- pin legacy polling.
// "none" -- disable the live tracker (the periodic check still runs).
//
// Unknown values fall back to "auto" with a warning. The metric
// csm_bpf_backend{feature="connection_tracker", kind="..."} reflects the
// chosen path.
func StartConnectionTracker(alertCh chan<- alert.Finding, cfg *config.Config) bpf.Backend {
choice := strings.ToLower(strings.TrimSpace(cfg.Detection.ConnectionTrackerBackend))
if choice == "" {
choice = bpf.BackendAuto
}
switch choice {
case bpf.BackendAuto, bpf.BackendBPF, bpf.BackendLegacy, bpf.BackendNone:
default:
csmlog.Warn("connection_tracker: unknown backend choice, using auto", "value", choice)
choice = bpf.BackendAuto
}
if choice == bpf.BackendNone {
csmlog.Info("connection_tracker: disabled by config")
bpf.SetActive("connection_tracker", bpf.BackendNone)
return nil
}
var bpfErr error
if choice == bpf.BackendAuto || choice == bpf.BackendBPF {
if b, err := tryStartConnectionBPFFn(context.Background(), alertCh, cfg); err == nil && b != nil {
csmlog.Info("connection_tracker", "backend", "bpf", "choice", choice)
bpf.SetActive("connection_tracker", bpf.BackendBPF)
return b
} else if err != nil {
bpfErr = err
level := "bpf-unsupported"
if errors.Is(err, bpf.ErrNotBuilt) {
level = "bpf-not-built"
}
csmlog.Info("connection_tracker: BPF unavailable", "state", level, "reason", err.Error(), "choice", choice)
if choice == bpf.BackendBPF {
csmlog.Warn("connection_tracker: backend=bpf but BPF unavailable; no live tracker", "reason", err.Error())
bpf.SetActive("connection_tracker", bpf.BackendNone)
emitBPFUnavailableFinding(alertCh, "connection_tracker", choice, "", err)
return nil
}
}
}
poller := newConnectionPoller(cfg, alertCh)
csmlog.Info("connection_tracker", "backend", "legacy", "choice", choice)
bpf.SetActive("connection_tracker", bpf.BackendLegacy)
if bpfErr != nil {
emitBPFUnavailableFinding(alertCh, "connection_tracker", choice, bpf.BackendLegacy, bpfErr)
}
return poller
}
func activeConnectionCfg(startup *config.Config) *config.Config {
if cfg := config.Active(); cfg != nil {
return cfg
}
return startup
}
// tryStartConnectionBPFFn is the package-level indirection so tests can
// substitute a fake without the bpf build tag.
var tryStartConnectionBPFFn = tryStartConnectionBPF
func tryStartConnectionBPF(ctx context.Context, ch chan<- alert.Finding, cfg *config.Config) (bpf.Backend, error) {
b, err := startConnectionBPF(ctx, ch, cfg)
if err != nil {
return nil, err
}
return b, nil
}
// attachProcessCtxToFinding sets f.Process from the cache when present, or
// enqueues a /proc enrichment so the next finding for the same PID benefits.
// Cache miss is the common case for short-lived processes; the finding is
// emitted with whatever context already exists (often none).
func attachProcessCtxToFinding(cache *processctx.Cache, enr *processctx.Enricher, f *alert.Finding, ev ConnectionEvent) {
if ev.PID == 0 {
return
}
req := processctxRequestFromConnection(ev)
if pc, needsEnrichment := cache.MaterializeVerifiedSnapshot(req); pc != nil {
f.Process = pc
if pc.Account != "" && (f.Check == "direct_smtp_egress" || f.TenantID == "") {
f.TenantID = pc.Account
}
if needsEnrichment {
enr.Enqueue(req)
}
return
}
enr.Enqueue(req)
}
func processctxRequestFromConnection(ev ConnectionEvent) processctx.EnrichRequest {
pid := int(ev.PID)
return processctx.EnrichRequest{
PID: pid,
UID: int(ev.UID),
UIDKnown: true,
Comm: ev.Comm,
StartedAt: processCtxStartedAt(pid),
}
}
package daemon
import (
"encoding/json"
"fmt"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/store"
)
// handleBaseline clears existing state and captures the current host as
// the new known-good reference. Mirrors the old `csm baseline` flow but
// runs inside the daemon so no external lock coordination is needed.
//
// Concurrency: a sync.Mutex on the daemon serialises baselines against
// each other. The baseline sweep still uses checks.ForceAll to bypass
// throttles; dry-run state is scoped to RunAllDryRun.
func (c *ControlListener) handleBaseline(argsRaw json.RawMessage) (any, error) {
var args control.BaselineArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
c.d.baselineMu.Lock()
defer c.d.baselineMu.Unlock()
histCount := 0
if sdb := store.Global(); sdb != nil {
histCount = sdb.HistoryCount()
}
if histCount > 0 && !args.Confirm {
return control.BaselineResult{
HistoryCleared: histCount,
NeedsConfirm: true,
}, nil
}
// Force-all bypasses throttles for the baseline sweep. Dry-run threads
// through RunAllDryRun so a concurrent periodic scanner running in live
// mode is never silenced by this caller.
prevForceAll := checks.ForceAll
checks.ForceAll = true
defer func() { checks.ForceAll = prevForceAll }()
cfg := c.d.currentCfg()
findings, _ := checks.RunAllDryRun(cfg, c.d.store)
c.d.store.SetBaseline(findings)
binaryHash, err := integrity.HashFile(c.d.binaryPath)
if err != nil {
return nil, fmt.Errorf("hashing binary: %w", err)
}
signed, err := integrity.SignConfigFilePreservingSnapshot(cfg.ConfigFile, cfg.ConfigDir, binaryHash)
if err != nil {
return nil, fmt.Errorf("saving integrity: %w", err)
}
publishSignedBaselineConfig(cfg, signed)
return control.BaselineResult{
Findings: len(findings),
HistoryCleared: histCount,
BinaryHash: binaryHash,
ConfigHash: signed.Integrity.ConfigHash,
}, nil
}
func publishSignedBaselineConfig(live, signedMain *config.Config) {
resynced := *live
resynced.Integrity = signedMain.Integrity
// confd_hash was selected by the on-disk main config's exemption list.
// Publish both together without mutating the config other goroutines read.
resynced.ConfD = signedMain.ConfD
config.SetActive(&resynced)
}
package daemon
import (
"bytes"
"encoding/json"
"fmt"
"net"
"os"
"os/exec"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/firewall"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/obs"
)
var (
firewallDeadmanMu sync.Mutex
firewallDeadmanSeq atomic.Uint64
)
// Subnet / batch / meta firewall handlers: operations that either span
// multiple IPs (deny-file/allow-file, subnet ops) or reshape the whole
// ruleset (flush/restart/apply-confirmed/confirm). `restart` and
// `apply-confirmed` require a live fwEngine — a dead engine means
// "systemctl restart csm" rather than rebuild-from-handler.
func (c *ControlListener) handleFirewallDenySubnet(argsRaw json.RawMessage) (any, error) {
var args control.FirewallSubnetArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if _, _, err := net.ParseCIDR(args.CIDR); err != nil {
return nil, fmt.Errorf("invalid cidr: %q", args.CIDR)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Blocked via CLI"
}
if err := c.d.fwEngine.BlockSubnet(args.CIDR, reason, 0); err != nil {
return nil, fmt.Errorf("block subnet %s: %w", args.CIDR, err)
}
return control.FirewallAckResult{
Message: fmt.Sprintf("Blocked subnet %s - %s", args.CIDR, reason),
}, nil
}
func (c *ControlListener) handleFirewallRemoveSubnet(argsRaw json.RawMessage) (any, error) {
var args control.FirewallSubnetArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if _, _, err := net.ParseCIDR(args.CIDR); err != nil {
return nil, fmt.Errorf("invalid cidr: %q", args.CIDR)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
if err := c.d.fwEngine.UnblockSubnet(args.CIDR); err != nil {
return nil, fmt.Errorf("remove subnet %s: %w", args.CIDR, err)
}
return control.FirewallAckResult{
Message: fmt.Sprintf("Removed subnet block %s", args.CIDR),
}, nil
}
func (c *ControlListener) handleFirewallDenyFile(argsRaw json.RawMessage) (any, error) {
var args control.FirewallFileArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if len(args.IPs) == 0 {
return nil, fmt.Errorf("no ips in batch")
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Bulk block via CLI"
}
blocked, failed, skipped := 0, 0, 0
for _, ip := range args.IPs {
if net.ParseIP(ip) == nil {
skipped++
continue
}
// Operator-initiated batch: bypass auto_response.dry_run gate.
if err := operatorForceBlock(c.d.fwEngine, ip, reason, 0); err != nil {
failed++
continue
}
blocked++
}
msg := fmt.Sprintf("Blocked %d, skipped %d invalid", blocked, skipped)
if failed > 0 {
msg = fmt.Sprintf("Blocked %d, failed %d, skipped %d invalid", blocked, failed, skipped)
}
return control.FirewallAckResult{Message: msg}, nil
}
func (c *ControlListener) handleFirewallAllowFile(argsRaw json.RawMessage) (any, error) {
var args control.FirewallFileArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if len(args.IPs) == 0 {
return nil, fmt.Errorf("no ips in batch")
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Bulk allow via CLI"
}
allowed, failed, skipped := 0, 0, 0
for _, ip := range args.IPs {
if net.ParseIP(ip) == nil {
skipped++
continue
}
if err := c.d.fwEngine.AllowIP(ip, reason); err != nil {
failed++
continue
}
allowed++
}
msg := fmt.Sprintf("Allowed %d, skipped %d invalid", allowed, skipped)
if failed > 0 {
msg = fmt.Sprintf("Allowed %d, failed %d, skipped %d invalid", allowed, failed, skipped)
}
return control.FirewallAckResult{Message: msg}, nil
}
func (c *ControlListener) handleFirewallFlush(_ json.RawMessage) (any, error) {
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
cfg := c.d.currentCfg()
// Clear the auto-block bookkeeping for the flushed IPs; a surviving
// ThreatDB temp row re-blocks the IP on the next scan and silently
// undoes the flush.
result, err := checks.FlushAutoBlockState(cfg.StatePath, c.d.fwEngine.FlushBlocked)
if result.SnapshotErr != nil {
csmlog.Warn("firewall flush could not snapshot persisted blocks", "err", result.SnapshotErr)
}
if err != nil {
if result.Flushed {
return nil, fmt.Errorf("firewall flushed but auto-block cleanup failed: %w", err)
}
return nil, err
}
message := "Flushed blocked IPs (subnet blocks kept; use remove-subnet)"
if result.SnapshotErr == nil {
message = fmt.Sprintf("Flushed %d blocked IPs (subnet blocks kept; use remove-subnet)", result.BlockedCount)
}
return control.FirewallAckResult{
Message: message,
}, nil
}
func (c *ControlListener) handleFirewallRestart(_ json.RawMessage) (any, error) {
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall engine not running; restart the csm daemon")
}
firewallDeadmanMu.Lock()
defer firewallDeadmanMu.Unlock()
cfg := c.d.currentCfg()
confirmFile, _, _ := firewallRollbackFiles(cfg.StatePath)
if err := rejectPendingFirewallConfirmation(confirmFile); err != nil {
return nil, err
}
// Re-read the firewall block from disk: the engine holds the copy taken
// at daemon start, and a restart that re-applied it reported success
// while the operator's edit stayed unapplied.
previous, err := c.d.refreshFirewallFromDisk()
if err != nil {
return nil, err
}
if err := c.d.fwEngine.Apply(); err != nil {
c.d.fwEngine.SetConfig(previous)
return nil, fmt.Errorf("applying ruleset: %w", err)
}
state, _ := firewall.LoadState(cfg.StatePath)
return control.FirewallAckResult{
Message: fmt.Sprintf("Firewall restarted. %d blocked, %d allowed IPs restored.", len(state.Blocked), len(state.Allowed)),
}, nil
}
func (c *ControlListener) handleFirewallApplyConfirmed(argsRaw json.RawMessage) (any, error) {
var args control.FirewallApplyConfirmedArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
minutes := args.Minutes
if minutes <= 0 || minutes > 60 {
minutes = 3
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall engine not running; restart the csm daemon")
}
cfg := c.d.currentCfg()
confirmFile, rollbackFile, legacyRollbackFile := firewallRollbackFiles(cfg.StatePath)
// Read the edited firewall block only after the deadman has established
// that no other confirmation window owns the rollback files. The
// preparation and snapshot run under one lock so a rejected second
// attempt cannot replace the active window's engine rollback hook.
if err := applyFirewallDeadmanPrepared(confirmFile, rollbackFile, legacyRollbackFile,
time.Duration(minutes)*time.Minute, c.d.refreshFirewallFromDisk,
func(previous *firewall.FirewallConfig) { c.d.fwEngine.SetConfig(previous) },
c.d.fwEngine.Apply); err != nil {
return nil, err
}
state, _ := firewall.LoadState(cfg.StatePath)
return control.FirewallAckResult{
Message: fmt.Sprintf("Firewall applied with %d-minute rollback timer. %d blocked, %d allowed. Run `csm firewall confirm` to keep.", minutes, len(state.Blocked), len(state.Allowed)),
}, nil
}
func (c *ControlListener) handleFirewallConfirm(_ json.RawMessage) (any, error) {
firewallDeadmanMu.Lock()
defer firewallDeadmanMu.Unlock()
cfg := c.d.currentCfg()
confirmFile, rollbackFile, legacyRollbackFile := firewallRollbackFiles(cfg.StatePath)
if _, err := os.Stat(confirmFile); err != nil {
if !os.IsNotExist(err) {
return nil, fmt.Errorf("checking confirm marker: %w", err)
}
if cleanupErr := removeFirewallRollbackFiles(rollbackFile, legacyRollbackFile); cleanupErr != nil {
return nil, cleanupErr
}
return control.FirewallAckResult{
Message: "No pending confirmation. Firewall is already confirmed.",
}, nil
}
if err := removeFirewallRollbackFiles(confirmFile, rollbackFile, legacyRollbackFile); err != nil {
return nil, err
}
// Confirmed: the candidate ruleset input stays in the engine.
firewallDeadmanOnRollback = nil
return control.FirewallAckResult{
Message: "Firewall confirmed. Rollback timer cancelled.",
}, nil
}
// applyFirewallDeadman runs the tentative-apply protocol: snapshot the live
// ruleset, persist the confirm deadline, then apply the candidate. The
// marker is written BEFORE the kernel apply on purpose: a crash between the
// two leaves a valid deadline on disk for startup recovery to settle,
// whereas the reverse order would leave the candidate applied with no
// record that it was never confirmed (a permanent lockout).
func applyFirewallDeadman(confirmFile, rollbackFile, legacyRollbackFile string, window time.Duration, apply func() error) error {
return applyFirewallDeadmanPrepared(confirmFile, rollbackFile, legacyRollbackFile, window, nil, nil, apply)
}
// applyFirewallDeadmanPrepared checks ownership before prepare installs the
// candidate engine configuration. restore reverses that installation on any
// pre-apply failure and becomes the live rollback hook once the window is
// armed. Both callbacks run while firewallDeadmanMu is held.
func applyFirewallDeadmanPrepared(confirmFile, rollbackFile, legacyRollbackFile string, window time.Duration,
prepare func() (*firewall.FirewallConfig, error), restore func(*firewall.FirewallConfig), apply func() error,
) error {
firewallDeadmanMu.Lock()
defer firewallDeadmanMu.Unlock()
if err := os.MkdirAll(filepath.Dir(rollbackFile), 0700); err != nil {
return fmt.Errorf("creating firewall rollback dir: %w", err)
}
if err := rejectPendingFirewallConfirmation(confirmFile); err != nil {
return err
}
var rollbackConfig *firewall.FirewallConfig
if prepare != nil {
var err error
rollbackConfig, err = prepare()
if err != nil {
return err
}
firewallDeadmanOnRollback = func() { restore(rollbackConfig) }
}
restorePrepared := func() {
if prepare != nil {
runFirewallDeadmanOnRollbackLocked()
}
}
if err := removeFirewallRollbackFiles(rollbackFile, legacyRollbackFile); err != nil {
restorePrepared()
return err
}
if err := writeFirewallRollbackFile(rollbackFile); err != nil {
restorePrepared()
return err
}
if err := snapshotFirewallConfig(rollbackFile, rollbackConfig); err != nil {
_ = removeFirewallRollbackFiles(rollbackFile)
restorePrepared()
return err
}
deadline := time.Now().Add(window)
marker := newFirewallConfirmMarker(deadline)
if err := os.WriteFile(confirmFile, marker, 0600); err != nil {
_ = removeFirewallRollbackFiles(confirmFile, rollbackFile)
restorePrepared()
return fmt.Errorf("writing confirm marker: %w", err)
}
if err := apply(); err != nil {
// Apply may have partially mutated the kernel; re-applying the
// snapshot is a no-op when it did not, and the undo when it did.
if restoreErr := applyFirewallRollbackFile(rollbackFile); restoreErr != nil {
// Keep marker + snapshot: the deadline stands, so the armed
// deadman retries the restore at expiry (and startup recovery
// does the same after a crash).
armFirewallDeadman(confirmFile, rollbackFile, marker, time.Until(deadline))
return fmt.Errorf("applying ruleset: %w; rollback restore failed: %v", err, restoreErr)
}
if cleanupErr := removeFirewallRollbackFiles(confirmFile, rollbackFile); cleanupErr != nil {
restorePrepared()
return fmt.Errorf("applying ruleset: %w; %v", err, cleanupErr)
}
restorePrepared()
return fmt.Errorf("applying ruleset: %w; previous ruleset restored", err)
}
armFirewallDeadman(confirmFile, rollbackFile, marker, time.Until(deadline))
return nil
}
func rejectPendingFirewallConfirmation(confirmFile string) error {
if _, err := os.Stat(confirmFile); err == nil {
return fmt.Errorf("firewall confirmation already pending; run `csm firewall confirm` or wait for rollback before applying another ruleset")
} else if !os.IsNotExist(err) {
return fmt.Errorf("checking confirm marker: %w", err)
}
return nil
}
// armFirewallDeadman restores the pre-apply ruleset once wait elapses unless
// the operator confirms first. The goroutine lives in the daemon (long-lived,
// so it survives CLI exit); a daemon restart kills it, which is why
// recoverFirewallApplyConfirmed re-arms from the persisted deadline.
// firewallDeadmanOnRollback runs after a rollback restores the kernel
// snapshot (deadline expiry or explicit revert), so the engine's ruleset
// input reverts with it; otherwise the next Apply would rebuild the rules
// the operator never confirmed. Cleared on confirm. Guarded by
// firewallDeadmanMu except in setFirewallDeadmanOnRollback, which takes it.
var firewallDeadmanOnRollback func()
func setFirewallDeadmanOnRollback(fn func()) {
firewallDeadmanMu.Lock()
defer firewallDeadmanMu.Unlock()
firewallDeadmanOnRollback = fn
}
// runFirewallDeadmanOnRollbackLocked invokes and clears the rollback hook.
// Caller holds firewallDeadmanMu.
func runFirewallDeadmanOnRollbackLocked() {
if firewallDeadmanOnRollback != nil {
firewallDeadmanOnRollback()
firewallDeadmanOnRollback = nil
}
}
func armFirewallDeadman(confirmFile, rollbackFile string, marker []byte, wait time.Duration) {
obs.SafeGo("fw-apply-confirmed-rollback", func() {
time.Sleep(wait)
firewallDeadmanMu.Lock()
defer firewallDeadmanMu.Unlock()
if err := restoreFirewallRollback(confirmFile, rollbackFile, marker); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Firewall rollback failed: %v\n", ts(), err)
}
})
}
// recoverFirewallApplyConfirmed settles an apply-confirmed window that a
// daemon restart interrupted. Without it the restart kills the deadman
// goroutine and startFirewall re-applies the unconfirmed candidate, making
// permanent exactly the lockout the command exists to prevent. Must run
// after startFirewall so an expired-window restore lands on top of the
// candidate ruleset the startup Apply just re-applied, and before the
// control listener starts so confirm/cancel cannot race the recovery.
func (d *Daemon) recoverFirewallApplyConfirmed() {
firewallDeadmanMu.Lock()
defer firewallDeadmanMu.Unlock()
confirmFile, rollbackFile, legacyRollbackFile := firewallRollbackFiles(d.cfg.StatePath)
marker, err := os.ReadFile(confirmFile) // #nosec G304 -- CSM-owned marker under the state dir.
if err != nil {
if !os.IsNotExist(err) {
csmlog.Warn("firewall confirm marker unreadable; leaving apply-confirmed state untouched", "err", err)
return
}
// No marker means no candidate reached the kernel with a pending
// window (the marker is written before the apply), so a leftover
// snapshot is debris from an aborted handler run.
if cleanupErr := removeFirewallRollbackFiles(rollbackFile, legacyRollbackFile); cleanupErr != nil {
csmlog.Warn("firewall rollback snapshot cleanup failed", "err", cleanupErr)
}
return
}
if err := installFirewallConfigRollbackHook(d, rollbackFile); err != nil {
csmlog.Warn("firewall configuration rollback snapshot unreadable", "err", err)
}
deadline, parseErr := parseFirewallConfirmDeadline(marker)
if parseErr != nil {
// Fail safe: an unconfirmed ruleset must never outlive its
// deadline, and a corrupt marker gives no deadline to honour.
csmlog.Warn("firewall confirm marker corrupt; restoring previous ruleset", "err", parseErr)
if restoreErr := restoreFirewallRollback(confirmFile, rollbackFile, marker); restoreErr != nil {
csmlog.Warn("firewall rollback failed", "err", restoreErr)
}
return
}
if remaining := time.Until(deadline); remaining > 0 {
// startFirewall already re-applied the candidate, matching the
// kernel state from before the restart, so the operator keeps the
// verification window they asked for; the re-armed deadman still
// bounds it with the original deadline.
armFirewallDeadman(confirmFile, rollbackFile, marker, remaining)
csmlog.Info("firewall apply-confirmed window resumed after restart",
"deadline", deadline.Format(time.RFC3339))
return
}
if restoreErr := restoreFirewallRollback(confirmFile, rollbackFile, marker); restoreErr != nil {
csmlog.Warn("firewall rollback failed", "err", restoreErr)
return
}
csmlog.Warn("firewall apply-confirmed window expired during restart; previous ruleset restored",
"deadline", deadline.Format(time.RFC3339))
}
func newFirewallConfirmMarker(deadline time.Time) []byte {
seq := firewallDeadmanSeq.Add(1)
return []byte(fmt.Sprintf("%s\n%d", deadline.UTC().Format(time.RFC3339Nano), seq))
}
func parseFirewallConfirmDeadline(marker []byte) (time.Time, error) {
text := strings.TrimSpace(string(marker))
if i := strings.IndexByte(text, '\n'); i >= 0 {
text = strings.TrimSpace(text[:i])
}
return time.Parse(time.RFC3339Nano, text)
}
func firewallRollbackFiles(statePath string) (confirmFile, rollbackFile, legacyRollbackFile string) {
firewallDir := filepath.Join(statePath, "firewall")
return filepath.Join(firewallDir, "confirm_pending"),
filepath.Join(firewallDir, "rollback.nft"),
filepath.Join(firewallDir, "rollback.sh")
}
func writeFirewallRollbackFile(rollbackFile string) error {
// #nosec G204 -- "nft list ruleset" is literal.
nftDump, err := exec.Command("nft", "list", "ruleset").Output()
if err != nil {
return fmt.Errorf("capturing rollback ruleset: %w", err)
}
// The dump is nft syntax, so store it as data consumed by nft -f.
// An empty live ruleset still needs a rollback file; flush ruleset
// restores that state.
payload := make([]byte, 0, len("flush ruleset\n")+len(nftDump)+1)
payload = append(payload, "flush ruleset\n"...)
payload = append(payload, nftDump...)
if len(nftDump) > 0 && nftDump[len(nftDump)-1] != '\n' {
payload = append(payload, '\n')
}
// #nosec G306 -- root-only state dir; this is data, not an executable.
if err := os.WriteFile(rollbackFile, payload, 0600); err != nil {
_ = removeFileIfExists(rollbackFile)
return fmt.Errorf("writing rollback ruleset: %w", err)
}
if err := snapshotFirewallState(rollbackFile); err != nil {
_ = removeFileIfExists(rollbackFile)
return err
}
return nil
}
func restoreFirewallRollback(confirmFile, rollbackFile string, expectedMarker []byte) error {
raw, err := os.ReadFile(confirmFile) // #nosec G304 -- CSM-owned marker under the state dir.
if err != nil {
if os.IsNotExist(err) {
return nil
}
return fmt.Errorf("checking confirm marker: %w", err)
}
// A different marker means a newer apply-confirmed window owns the
// files now; this caller's window was superseded and firing would
// roll the newer window back before its own deadline.
if !bytes.Equal(raw, expectedMarker) {
return nil
}
if err := applyFirewallRollbackFile(rollbackFile); err != nil {
return err
}
// The kernel is back on the snapshot; put the engine's ruleset input
// back too, or the next Apply rebuilds the unconfirmed candidate.
runFirewallDeadmanOnRollbackLocked()
return removeFirewallRollbackFiles(confirmFile, rollbackFile, legacyRollbackFileFor(rollbackFile))
}
func applyFirewallRollbackFile(rollbackFile string) (resultErr error) {
rec := actionlog.Record{Op: "operate.manual_firewall", Action: "rollback", Target: rollbackFile, Command: []string{"nft", "-f", rollbackFile}, Result: actionlog.Failed}
defer func() {
if resultErr != nil {
rec.Error = resultErr.Error()
}
actionlog.Write(rec)
}()
if _, err := os.Stat(rollbackFile); err != nil {
if os.IsNotExist(err) {
return fmt.Errorf("rollback ruleset missing")
}
return fmt.Errorf("checking rollback ruleset: %w", err)
}
// #nosec G204 -- nft is hardcoded; rollbackFile is a CSM-written path.
out, err := exec.Command("nft", "-f", rollbackFile).CombinedOutput()
if err != nil {
out = bytes.TrimSpace(out)
if len(out) > 0 {
return fmt.Errorf("restoring rollback ruleset: %w: %s", err, out)
}
return fmt.Errorf("restoring rollback ruleset: %w", err)
}
rec.Result = actionlog.Applied
// Kernel is back on the snapshot; state.json must follow, or the UI
// keeps describing the window's mutations as live.
return restoreFirewallStateSnapshot(rollbackFile)
}
func removeFirewallRollbackFiles(paths ...string) error {
for _, path := range paths {
if err := removeFileIfExists(path); err != nil {
return err
}
// A confirmed or superseded window drops its state snapshot with
// its ruleset snapshot.
if filepath.Base(path) == "rollback.nft" {
if err := removeFileIfExists(firewallStateSnapshotPath(path)); err != nil {
return err
}
if err := removeFileIfExists(firewallConfigSnapshotPath(path)); err != nil {
return err
}
}
}
return nil
}
func firewallConfigSnapshotPath(rollbackFile string) string {
return rollbackFile + ".config.json"
}
func snapshotFirewallConfig(rollbackFile string, cfg *firewall.FirewallConfig) error {
if cfg == nil {
return nil
}
if err := atomicio.AtomicWriteJSON(firewallConfigSnapshotPath(rollbackFile), 0o600, cfg); err != nil {
return fmt.Errorf("snapshotting firewall configuration: %w", err)
}
return nil
}
var restoreFirewallEngineConfig = func(d *Daemon, cfg *firewall.FirewallConfig) {
if d.fwEngine != nil {
d.fwEngine.SetConfig(cfg)
}
}
// installFirewallConfigRollbackHook restores the engine input after the
// kernel and state snapshots. The persisted copy matters only across daemon
// restart; the original process keeps an equivalent in-memory hook.
func installFirewallConfigRollbackHook(d *Daemon, rollbackFile string) error {
raw, err := os.ReadFile(firewallConfigSnapshotPath(rollbackFile)) // #nosec G304 -- CSM-owned snapshot under the state dir.
if err != nil {
if os.IsNotExist(err) {
return nil
}
return err
}
var cfg firewall.FirewallConfig
if err := json.Unmarshal(raw, &cfg); err != nil {
return fmt.Errorf("decoding snapshot: %w", err)
}
firewallDeadmanOnRollback = func() { restoreFirewallEngineConfig(d, &cfg) }
return nil
}
func legacyRollbackFileFor(rollbackFile string) string {
return filepath.Join(filepath.Dir(rollbackFile), "rollback.sh")
}
func removeFileIfExists(path string) error {
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("removing %s: %w", filepath.Base(path), err)
}
return nil
}
package daemon
import (
"encoding/json"
"fmt"
"net"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/store"
)
// dropAutoBlockThreatRow removes the auto-block threat row for ip after an
// operator unblock. Operator permanent blocks are left in place: a
// firewall-only unblock must not silently clear a deliberate block. Under
// older builds a stale auto-block row could outlive the firewall block and
// ip_reputation would re-flag the IP into a new block loop. Clearing the
// persisted row when present and the in-memory temp copy stops that.
func dropAutoBlockThreatRow(ip string) {
if parsed := net.ParseIP(ip); parsed != nil {
ip = parsed.String()
}
if sdb := store.Global(); sdb != nil {
_, _ = sdb.RemoveTemporaryBlock(ip)
}
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.RemoveTemporary(ip)
}
}
// Single-IP firewall mutation handlers. Each validates args, guards on
// c.d.fwEngine != nil, calls the matching engine method, and returns a
// FirewallAckResult with a human-readable message the CLI prints verbatim.
// operatorForceBlock runs an operator-initiated force block and reports it
// to the shared firewall outcome metric alongside auto-response blocks.
func operatorForceBlock(e interface {
BlockIPForce(ip string, reason string, timeout time.Duration) error
}, ip, reason string, timeout time.Duration) error {
err := e.BlockIPForce(ip, reason, timeout)
checks.ObserveOperatorBlock(err, checks.BlockSourceCLI)
return err
}
func (c *ControlListener) handleFirewallBlock(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Blocked via CLI"
}
// Operator-initiated: bypass auto_response.dry_run gate.
if err := operatorForceBlock(c.d.fwEngine, args.IP, reason, 0); err != nil {
return nil, fmt.Errorf("block %s: %w", args.IP, err)
}
msg := fmt.Sprintf("Blocked %s - %s", args.IP, reason)
msg += cloudflareCoverageSuffix(c.d.fwEngine, args.IP)
return control.FirewallAckResult{Message: msg}, nil
}
// cloudflareCoverageSuffix warns the operator when a just-blocked IP sits
// inside a Cloudflare allow range: the input chain accepts CF edges on TCP
// 80/443 before the blocked drop, so the block does not stop web traffic.
func cloudflareCoverageSuffix(e interface{ CloudflareCovers(string) bool }, ip string) string {
if e != nil && e.CloudflareCovers(ip) {
return " (warning: " + firewall.CloudflareCoverageWarning + ")"
}
return ""
}
func (c *ControlListener) handleFirewallUnblock(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
if err := c.d.fwEngine.UnblockIP(args.IP); err != nil {
return nil, fmt.Errorf("unblock %s: %w", args.IP, err)
}
dropAutoBlockThreatRow(args.IP)
msg := fmt.Sprintf("Unblocked %s", args.IP)
// A blocked subnet covering the address keeps dropping it whatever
// happens to the per-IP element; say so instead of reporting success.
if cidr, covered := c.d.fwEngine.BlockedSubnetCovering(args.IP); covered {
msg += fmt.Sprintf(" (still dropped by blocked subnet %s; unblock that subnet to restore access)", cidr)
}
return control.FirewallAckResult{Message: msg}, nil
}
func (c *ControlListener) handleFirewallAllow(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Allowed via CLI"
}
if err := c.d.fwEngine.AllowIP(args.IP, reason); err != nil {
return nil, fmt.Errorf("allow %s: %w", args.IP, err)
}
msg := fmt.Sprintf("Allowed %s - %s", args.IP, reason)
if cidr, covered := c.d.fwEngine.BlockedSubnetCovering(args.IP); covered {
msg += fmt.Sprintf(" (WARNING: still dropped by blocked subnet %s; unblock the subnet for this allow to take effect)", cidr)
}
return control.FirewallAckResult{Message: msg}, nil
}
func (c *ControlListener) handleFirewallRemoveAllow(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
if err := c.d.fwEngine.RemoveAllowIP(args.IP); err != nil {
return nil, fmt.Errorf("remove-allow %s: %w", args.IP, err)
}
return control.FirewallAckResult{Message: fmt.Sprintf("Removed %s from allow list", args.IP)}, nil
}
func (c *ControlListener) handleFirewallAllowPort(argsRaw json.RawMessage) (any, error) {
var args control.FirewallPortArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if args.Port <= 0 || args.Port > 65535 {
return nil, fmt.Errorf("invalid port: %d", args.Port)
}
proto := args.Proto
if proto == "" {
proto = "tcp"
}
if proto != "tcp" && proto != "udp" {
return nil, fmt.Errorf("invalid proto: %q (want tcp or udp)", args.Proto)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Port-allowed via CLI"
}
if err := c.d.fwEngine.AllowIPPort(args.IP, args.Port, proto, reason); err != nil {
return nil, fmt.Errorf("allow-port %s %s:%d: %w", args.IP, proto, args.Port, err)
}
return control.FirewallAckResult{
Message: fmt.Sprintf("Saved allow for %s on %s:%d - %s; takes effect after firewall reload", args.IP, proto, args.Port, reason),
}, nil
}
func (c *ControlListener) handleFirewallRemovePort(argsRaw json.RawMessage) (any, error) {
var args control.FirewallPortArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if args.Port <= 0 || args.Port > 65535 {
return nil, fmt.Errorf("invalid port: %d", args.Port)
}
proto := args.Proto
if proto == "" {
proto = "tcp"
}
if proto != "tcp" && proto != "udp" {
return nil, fmt.Errorf("invalid proto: %q (want tcp or udp)", args.Proto)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
if err := c.d.fwEngine.RemoveAllowIPPort(args.IP, args.Port, proto); err != nil {
return nil, fmt.Errorf("remove-port %s %s:%d: %w", args.IP, proto, args.Port, err)
}
return control.FirewallAckResult{
Message: fmt.Sprintf("Saved port-allow removal for %s on %s:%d; takes effect after firewall reload", args.IP, proto, args.Port),
}, nil
}
func (c *ControlListener) handleFirewallTempBan(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
// Parse timeout FIRST so callers get a duration-parse error before
// the engine-nil check (the unit test depends on this ordering).
if args.Timeout == "" {
return nil, fmt.Errorf("tempban requires timeout")
}
timeout, err := time.ParseDuration(args.Timeout)
if err != nil {
return nil, fmt.Errorf("parsing duration %q: %w", args.Timeout, err)
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Temp-banned via CLI"
}
// Operator-initiated: bypass auto_response.dry_run gate.
if err := operatorForceBlock(c.d.fwEngine, args.IP, reason, timeout); err != nil {
return nil, fmt.Errorf("tempban %s: %w", args.IP, err)
}
msg := fmt.Sprintf("Temp-banned %s for %s - %s", args.IP, timeout, reason)
msg += cloudflareCoverageSuffix(c.d.fwEngine, args.IP)
return control.FirewallAckResult{Message: msg}, nil
}
func (c *ControlListener) handleFirewallTempAllow(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if args.Timeout == "" {
return nil, fmt.Errorf("tempallow requires timeout")
}
timeout, err := time.ParseDuration(args.Timeout)
if err != nil {
return nil, fmt.Errorf("parsing duration %q: %w", args.Timeout, err)
}
if net.ParseIP(args.IP) == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
if c.d.fwEngine == nil {
return nil, fmt.Errorf("firewall disabled in csm.yaml")
}
reason := args.Reason
if reason == "" {
reason = "Temp-allowed via CLI"
}
if err := c.d.fwEngine.TempAllowIP(args.IP, reason, timeout); err != nil {
return nil, fmt.Errorf("tempallow %s: %w", args.IP, err)
}
return control.FirewallAckResult{
Message: fmt.Sprintf("Temp-allowed %s for %s - %s", args.IP, timeout, reason),
}, nil
}
package daemon
import (
"encoding/json"
"fmt"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/firewall"
)
// Read-only firewall handlers: no state mutation, just surface what
// firewall.LoadState already has.
// fmtPortsSlice converts a slice of int ports into a slice of strings,
// one entry per port. Mirrors the atomic rendering the wire schema
// expects — the CLI joins them with commas for display. Empty input
// returns nil so the JSON wire form is `null` / omitted rather than an
// empty array carrying a placeholder.
func fmtPortsSlice(ports []int) []string {
if len(ports) == 0 {
return nil
}
out := make([]string, len(ports))
for i, p := range ports {
out[i] = strconv.Itoa(p)
}
return out
}
func (c *ControlListener) handleFirewallStatus(_ json.RawMessage) (any, error) {
cfg := c.d.currentCfg()
// LoadState tolerates a missing file (returns empty state).
state, err := firewall.LoadState(cfg.StatePath)
if err != nil {
return nil, fmt.Errorf("loading firewall state: %w", err)
}
fwCfg := config.EffectiveFirewallConfig(cfg)
result := control.FirewallStatusResult{
Enabled: fwCfg.Enabled,
TCPIn: fmtPortsSlice(fwCfg.TCPIn),
TCPOut: fmtPortsSlice(fwCfg.TCPOut),
UDPIn: fmtPortsSlice(fwCfg.UDPIn),
UDPOut: fmtPortsSlice(fwCfg.UDPOut),
Restricted: fmtPortsSlice(fwCfg.RestrictedTCP),
PassiveFTPStart: fwCfg.PassiveFTPStart,
PassiveFTPEnd: fwCfg.PassiveFTPEnd,
TCPOutAllow: fmtOutAllow(fwCfg.TCPOutAllow),
InfraIPCount: len(fwCfg.InfraIPs),
BlockedCount: len(state.Blocked),
BlockedNetCount: len(state.BlockedNet),
AllowedCount: len(state.Allowed),
SYNFlood: fwCfg.SYNFloodProtection,
ConnRateLimit: fwCfg.ConnRateLimit,
LogDropped: fwCfg.LogDropped,
LogRate: fwCfg.LogRate,
}
// Recent blocked: last 10, newest-first, matching cmd/csm/firewall.go:fwStatus.
// Time values emitted as RFC3339 so the wire schema is not tied to
// Go's time.Time encoding; the CLI renders "N ago" on receive.
shown := 0
for i := len(state.Blocked) - 1; i >= 0 && shown < 10; i-- {
b := state.Blocked[i]
entry := control.FirewallBlockedEntry{
IP: b.IP,
Reason: b.Reason,
BlockedAt: b.BlockedAt.UTC().Format(time.RFC3339),
}
if !b.ExpiresAt.IsZero() {
entry.ExpiresAt = b.ExpiresAt.UTC().Format(time.RFC3339)
}
result.RecentBlocked = append(result.RecentBlocked, entry)
shown++
}
return result, nil
}
func (c *ControlListener) handleFirewallPorts(_ json.RawMessage) (any, error) {
cfg := c.d.currentCfg()
var lines []string
fwCfg := config.EffectiveFirewallConfig(cfg)
lines = append(lines, "TCP Inbound (public):")
lines = append(lines, " "+joinPorts(fwCfg.TCPIn))
lines = append(lines, "")
if len(fwCfg.RestrictedTCP) > 0 {
lines = append(lines, "TCP Restricted (infra only):")
lines = append(lines, " "+joinPorts(fwCfg.RestrictedTCP))
lines = append(lines, "")
}
lines = append(lines, "TCP Outbound:")
lines = append(lines, " "+joinPorts(fwCfg.TCPOut))
lines = append(lines, "")
lines = append(lines, "UDP Inbound:")
lines = append(lines, " "+joinPorts(fwCfg.UDPIn))
lines = append(lines, "")
lines = append(lines, "UDP Outbound:")
lines = append(lines, " "+joinPorts(fwCfg.UDPOut))
lines = append(lines, "")
if fwCfg.PassiveFTPStart > 0 {
lines = append(lines, "Passive FTP:")
lines = append(lines, fmt.Sprintf(" %d-%d", fwCfg.PassiveFTPStart, fwCfg.PassiveFTPEnd))
}
if out := fmtOutAllow(fwCfg.TCPOutAllow); len(out) > 0 {
lines = append(lines, "")
lines = append(lines, "TCP Outbound (destination-scoped):")
for _, line := range out {
lines = append(lines, " "+line)
}
}
return control.FirewallListResult{Lines: lines}, nil
}
// fmtOutAllow renders tcp_out_allow one rule per line. A single-port range
// prints as the bare port so the common case does not read as "49152-49152".
func fmtOutAllow(rules []firewall.OutAllowRule) []string {
if len(rules) == 0 {
return nil
}
out := make([]string, 0, len(rules))
for _, r := range rules {
ports := fmt.Sprintf("%d-%d", r.PortStart, r.PortEnd)
if r.PortStart == r.PortEnd {
ports = strconv.Itoa(r.PortStart)
}
out = append(out, fmt.Sprintf("%s tcp %s", r.Dst, ports))
}
return out
}
// joinPorts returns a comma-separated rendering of ports matching
// cmd/csm/firewall.go:fmtPortsWrap's behaviour for the ports handler.
// Wrapping stays client-side; the wire payload keeps the full CSV so
// the CLI can format for whatever terminal width it runs in.
func joinPorts(ports []int) string {
if len(ports) == 0 {
return "(none)"
}
strs := make([]string, len(ports))
for i, p := range ports {
strs[i] = strconv.Itoa(p)
}
return strings.Join(strs, ", ")
}
func (c *ControlListener) handleFirewallGrep(argsRaw json.RawMessage) (any, error) {
var args control.FirewallGrepArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
cfg := c.d.currentCfg()
// Empty pattern matches nothing — the CLI used to require a
// positional arg and exit with usage; mirror that by returning an
// empty result rather than dumping everything.
if args.Pattern == "" {
return control.FirewallListResult{}, nil
}
pattern := strings.ToLower(args.Pattern)
state, err := firewall.LoadState(cfg.StatePath)
if err != nil {
return nil, fmt.Errorf("loading firewall state: %w", err)
}
var lines []string
now := time.Now()
for _, b := range state.Blocked {
if strings.Contains(strings.ToLower(b.IP), pattern) ||
strings.Contains(strings.ToLower(b.Reason), pattern) {
ago := now.Sub(b.BlockedAt).Truncate(time.Minute)
expires := "permanent"
if !b.ExpiresAt.IsZero() {
remaining := b.ExpiresAt.Sub(now).Truncate(time.Minute)
expires = fmt.Sprintf("%s left", remaining)
}
lines = append(lines, fmt.Sprintf("BLOCKED %-18s (%s ago, %s) %s",
b.IP, ago, expires, b.Reason))
}
}
for _, a := range state.Allowed {
if strings.Contains(strings.ToLower(a.IP), pattern) ||
strings.Contains(strings.ToLower(a.Reason), pattern) {
port := ""
if a.Port > 0 {
port = fmt.Sprintf(" port:%d", a.Port)
}
lines = append(lines, fmt.Sprintf("ALLOWED %-18s%s %s",
a.IP, port, a.Reason))
}
}
for _, s := range state.BlockedNet {
if strings.Contains(strings.ToLower(s.CIDR), pattern) ||
strings.Contains(strings.ToLower(s.Reason), pattern) {
ago := now.Sub(s.BlockedAt).Truncate(time.Minute)
lines = append(lines, fmt.Sprintf("SUBNET %-18s (%s ago) %s",
s.CIDR, ago, s.Reason))
}
}
for _, ip := range config.EffectiveFirewallConfig(cfg).InfraIPs {
if strings.Contains(strings.ToLower(ip), pattern) {
lines = append(lines, fmt.Sprintf("INFRA %s", ip))
}
}
return control.FirewallListResult{Lines: lines}, nil
}
func (c *ControlListener) handleFirewallAudit(argsRaw json.RawMessage) (any, error) {
var args control.FirewallAuditArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
limit := args.Limit
if limit <= 0 {
limit = 50
}
cfg := c.d.currentCfg()
entries := firewall.ReadAuditLog(cfg.StatePath, limit)
lines := make([]string, 0, len(entries))
for _, e := range entries {
ts := e.Timestamp.Format("2006-01-02 15:04:05")
dur := ""
if e.Duration != "" {
dur = fmt.Sprintf(" (%s)", e.Duration)
}
reason := ""
if e.Reason != "" {
reason = fmt.Sprintf(" %s", e.Reason)
}
lines = append(lines, fmt.Sprintf("%s %-13s %-18s%s%s",
ts, e.Action, e.IP, dur, reason))
}
return control.FirewallListResult{Lines: lines}, nil
}
package daemon
import (
"context"
"encoding/json"
"fmt"
"time"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/firewall/rollback"
)
func (c *ControlListener) handleFirewallRollbackStatus(_ json.RawMessage) (any, error) {
mgr := rollback.Global()
if mgr == nil {
return control.FirewallRollbackStatus{}, nil
}
st := mgr.Status()
out := control.FirewallRollbackStatus{
Pending: st.Pending,
AppliedBy: st.AppliedBy,
PrevHash: st.PrevHash,
NewHash: st.NewHash,
SecondsRemaining: st.SecondsRemaining,
}
if !st.AppliedAt.IsZero() {
out.AppliedAtRFC3339 = st.AppliedAt.Format(time.RFC3339)
}
if !st.ExpiresAt.IsZero() {
out.ExpiresAtRFC3339 = st.ExpiresAt.Format(time.RFC3339)
}
return out, nil
}
func (c *ControlListener) handleFirewallRollbackConfirm(_ json.RawMessage) (any, error) {
mgr := rollback.Global()
if mgr == nil {
return control.FirewallAckResult{Message: "rollback manager not initialised"}, nil
}
st := mgr.Status()
if !st.Pending {
return control.FirewallAckResult{Message: "no pending rollback"}, nil
}
if err := mgr.ConfirmIfCurrent(st); err != nil {
return nil, fmt.Errorf("confirm rollback: %w", err)
}
return control.FirewallAckResult{Message: "rollback confirmed; pending change is now permanent"}, nil
}
func (c *ControlListener) handleFirewallRollbackRevert(_ json.RawMessage) (any, error) {
mgr := rollback.Global()
if mgr == nil {
return control.FirewallAckResult{Message: "rollback manager not initialised"}, nil
}
st := mgr.Status()
if !st.Pending {
return control.FirewallAckResult{Message: "no pending rollback"}, nil
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
if err := mgr.RevertIfCurrent(ctx, st); err != nil {
return nil, fmt.Errorf("revert rollback: %w", err)
}
return control.FirewallAckResult{Message: "rollback reverted; previous config restored, daemon restart issued"}, nil
}
package daemon
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// dispatch parses a raw request line, routes to the right handler, and
// wraps the handler's (result, error) pair in the response envelope.
// Unknown commands fail cleanly rather than crash the listener.
func (c *ControlListener) dispatch(line []byte) control.Response {
var req control.Request
if err := json.Unmarshal(line, &req); err != nil {
return control.Response{OK: false, Error: fmt.Sprintf("bad request: %v", err)}
}
var (
result any
err error
)
switch req.Cmd {
case control.CmdTierRun:
result, err = c.handleTierRun(req.Args)
case control.CmdStatus:
result, err = c.handleStatus(req.Args)
case control.CmdHistoryRead:
result, err = c.handleHistoryRead(req.Args)
case control.CmdRulesReload:
result, err = c.handleRulesReload(req.Args)
case control.CmdGeoIPReload:
result, err = c.handleGeoIPReload(req.Args)
case control.CmdBotRangesReload:
result, err = c.handleBotRangesReload(req.Args)
case control.CmdBaseline:
result, err = c.handleBaseline(req.Args)
case control.CmdThreatForget:
result, err = c.handleThreatForget(req.Args)
case control.CmdFirewallStatus:
result, err = c.handleFirewallStatus(req.Args)
case control.CmdFirewallPorts:
result, err = c.handleFirewallPorts(req.Args)
case control.CmdFirewallGrep:
result, err = c.handleFirewallGrep(req.Args)
case control.CmdFirewallActions:
result, err = c.handleFirewallActions(req.Args)
case control.CmdFirewallActionResolve:
result, err = c.handleFirewallActionResolve(req.Args)
case control.CmdFirewallAudit:
result, err = c.handleFirewallAudit(req.Args)
case control.CmdFirewallBlock:
result, err = c.handleFirewallBlock(req.Args)
case control.CmdFirewallUnblock:
result, err = c.handleFirewallUnblock(req.Args)
case control.CmdFirewallAllow:
result, err = c.handleFirewallAllow(req.Args)
case control.CmdFirewallRemoveAllow:
result, err = c.handleFirewallRemoveAllow(req.Args)
case control.CmdFirewallAllowPort:
result, err = c.handleFirewallAllowPort(req.Args)
case control.CmdFirewallRemovePort:
result, err = c.handleFirewallRemovePort(req.Args)
case control.CmdFirewallTempBan:
result, err = c.handleFirewallTempBan(req.Args)
case control.CmdFirewallTempAllow:
result, err = c.handleFirewallTempAllow(req.Args)
case control.CmdFirewallDenySubnet:
result, err = c.handleFirewallDenySubnet(req.Args)
case control.CmdFirewallRemoveSubnet:
result, err = c.handleFirewallRemoveSubnet(req.Args)
case control.CmdFirewallDenyFile:
result, err = c.handleFirewallDenyFile(req.Args)
case control.CmdFirewallAllowFile:
result, err = c.handleFirewallAllowFile(req.Args)
case control.CmdFirewallFlush:
result, err = c.handleFirewallFlush(req.Args)
case control.CmdFirewallRestart:
result, err = c.handleFirewallRestart(req.Args)
case control.CmdFirewallApplyConfirmed:
result, err = c.handleFirewallApplyConfirmed(req.Args)
case control.CmdFirewallConfirm:
result, err = c.handleFirewallConfirm(req.Args)
case control.CmdFirewallRollbackStatus:
result, err = c.handleFirewallRollbackStatus(req.Args)
case control.CmdFirewallRollbackOK:
result, err = c.handleFirewallRollbackConfirm(req.Args)
case control.CmdFirewallRollbackRevert:
result, err = c.handleFirewallRollbackRevert(req.Args)
case control.CmdStoreExport:
result, err = c.handleStoreExport(req.Args)
case control.CmdHistorySince:
result, err = c.handleHistorySince(req.Args)
case control.CmdPHPRelayStatus:
result, err = c.handlePHPRelayStatus(req.Args)
case control.CmdPHPRelayIgnoreScript:
result, err = c.handlePHPRelayIgnoreScript(req.Args)
case control.CmdPHPRelayUnignore:
result, err = c.handlePHPRelayUnignore(req.Args)
case control.CmdPHPRelayIgnoreList:
result, err = c.handlePHPRelayIgnoreList(req.Args)
case control.CmdPHPRelayDryRun:
result, err = c.handlePHPRelayDryRun(req.Args)
case control.CmdPHPRelayThaw:
result, err = c.handlePHPRelayThaw(req.Args)
case control.CmdIncidentsList:
result, err = c.handleIncidentsList(req.Args)
case control.CmdIncidentsShow:
result, err = c.handleIncidentsShow(req.Args)
case control.CmdIncidentsStatus:
result, err = c.handleIncidentsStatus(req.Args)
case control.CmdIncidentsBulkStatus:
result, err = c.handleIncidentsBulkStatus(req.Args)
case control.CmdScanEnqueue:
result, err = c.handleScanEnqueue(req.Args)
case control.CmdScanStatus:
result, err = c.handleScanStatus(req.Args)
case control.CmdScanReport:
result, err = c.handleScanReport(req.Args)
case control.CmdScanCancel:
result, err = c.handleScanCancel(req.Args)
default:
return control.Response{OK: false, Error: fmt.Sprintf("unknown command: %q", req.Cmd)}
}
if err != nil {
return control.Response{OK: false, Error: err.Error()}
}
payload, mErr := json.Marshal(result)
if mErr != nil {
return control.Response{OK: false, Error: "result marshal: " + mErr.Error()}
}
return control.Response{OK: true, Result: payload}
}
// parseTier maps the wire string onto the checks.Tier constants. Kept
// local to the listener so the protocol package does not depend on the
// checks package.
func parseTier(s string) (checks.Tier, error) {
switch s {
case "critical":
return checks.TierCritical, nil
case "deep":
return checks.TierDeep, nil
case "all", "":
return checks.TierAll, nil
}
return "", fmt.Errorf("unknown tier: %q", s)
}
// handleTierRun runs a tier synchronously and reports the result. The
// flow mirrors Daemon.runPeriodicChecks: integrity verify, RunTier,
// purge-and-merge, then hand findings to the alert pipeline. When
// Alerts=false the handler absorbs the old `csm check*` behaviour:
// skip auto-response for this run, append the raw findings to history,
// and return them in the response body so the CLI can render them
// verbatim.
func (c *ControlListener) handleTierRun(argsRaw json.RawMessage) (any, error) {
var args control.TierRunArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
tier, err := parseTier(args.Tier)
if err != nil {
return nil, err
}
// dryRun ties together the three Alerts=false side effects:
// auto-response suppression, the post-run history append, and the
// FindingList in the response. Hoisting it makes the invariant
// "these three happen together" visually obvious.
dryRun := !args.Alerts
if vErr := c.verifyTierRunIntegrity(); vErr != nil {
// Integrity failures are escalated through the normal alert
// pipeline so the on-call path sees them regardless of who
// kicked the tier run. The client also gets an error so the
// systemd timer unit fails loudly.
if args.Alerts {
if !alert.TryEnqueue(c.d.alertCh, alert.Finding{
Severity: alert.Critical,
Check: "integrity",
Message: fmt.Sprintf("BINARY/CONFIG TAMPER DETECTED: %v", vErr),
Timestamp: time.Now(),
}) {
atomic.AddInt64(&c.d.droppedAlerts, 1)
}
}
return nil, fmt.Errorf("integrity verify failed: %w", vErr)
}
// Dry-run threads through RunTierDryRun, so a concurrent periodic
// scanner running in live mode never sees this caller's dry-run state.
start := time.Now()
var (
findings []alert.Finding
purgeChecks []string
)
cfg := c.d.currentCfg()
scanCtx, gaps := checks.WithCoverageGaps(c.d.scanContext())
if dryRun {
findings, purgeChecks = checks.RunTierDryRunWithContext(scanCtx, cfg, c.d.store, tier)
} else {
findings, purgeChecks = checks.RunTierWithContext(scanCtx, cfg, c.d.store, tier)
}
c.recordTierRunFindings(cfg, findings, purgeChecks, gaps.Snapshot(), !dryRun, args.Alerts)
// Dry-run history + FindingList: the live path writes history via
// Daemon.runPeriodicChecks when the internal scanners fire; the
// pre-phase-2 `csm check*` wrote it via store.AppendHistory in
// cmd/csm/main.go. Preserve that quirk on the socket path.
if dryRun {
c.d.store.AppendHistory(findings)
}
newCount := len(c.d.store.FilterNew(findings))
result := control.TierRunResult{
Findings: len(findings),
NewFindings: newCount,
ElapsedMs: time.Since(start).Milliseconds(),
}
if dryRun {
if findings == nil {
result.FindingList = []alert.Finding{}
} else {
result.FindingList = findings
}
}
return result, nil
}
// recordTierRunFindings persists a control-socket tier run's findings. Auto-fix
// is gated on a live run so a dry run never edits a customer's wp-config.php,
// and the alert push is gated separately on whether the caller asked for alerts.
func (c *ControlListener) recordTierRunFindings(cfg *config.Config, findings []alert.Finding, purgeChecks []string, coverage *state.ScanCoverage, live, alerts bool) {
checks.StoreLatestScanFindingsWithCoverage(c.d.store, purgeChecks, findings, coverage)
if live {
c.d.applyWPCronAutoFix(cfg, findings)
}
if alerts {
c.d.enqueueScanAlerts(findings, "control")
}
}
func (c *ControlListener) verifyTierRunIntegrity() error {
return integrity.Verify(c.d.binaryPath, c.d.currentCfg())
}
// handleStatus reports what `csm status` historically printed from
// disk, sourced from the live daemon instead of re-opening the store.
func (c *ControlListener) handleStatus(_ json.RawMessage) (any, error) {
latest := c.d.store.LatestFindings()
latestTime := c.d.store.LatestScanTime()
var latestStr string
if !latestTime.IsZero() {
latestStr = latestTime.UTC().Format(time.RFC3339)
}
var historyCount int
if sdb := store.Global(); sdb != nil {
historyCount = sdb.HistoryCount()
}
uptime := int64(0)
if !c.d.startTime.IsZero() {
uptime = int64(time.Since(c.d.startTime).Seconds())
}
result := control.StatusResult{
Version: c.d.version,
UptimeSec: uptime,
LatestScanTime: latestStr,
LatestFindings: len(latest),
HistoryCount: historyCount,
DroppedAlerts: c.d.DroppedAlerts(),
}
snap := health.Build(c.d, c.d.version, health.Capabilities())
result.Snapshot = &snap
return result, nil
}
// handleHistoryRead paginates bbolt history. Clamps Limit so a buggy
// client cannot ask for everything at once; 1000 is well above the
// dashboard page size.
func (c *ControlListener) handleHistoryRead(argsRaw json.RawMessage) (any, error) {
var args control.HistoryReadArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if args.Limit <= 0 || args.Limit > 1000 {
args.Limit = 100
}
if args.Offset < 0 {
args.Offset = 0
}
findings, total := c.d.store.ReadHistory(args.Limit, args.Offset)
return control.HistoryReadResult{Findings: findings, Total: total}, nil
}
// handleRulesReload replaces `kill -HUP $(pidof csm)`. Returns after
// the reload completes so the client can confirm it happened.
func (c *ControlListener) handleRulesReload(_ json.RawMessage) (any, error) {
c.d.reloadSignatures()
return map[string]string{"status": "reloaded"}, nil
}
// handleGeoIPReload is the GeoIP equivalent of rules.reload.
func (c *ControlListener) handleGeoIPReload(_ json.RawMessage) (any, error) {
c.d.publishGeoIP()
return map[string]string{"status": "reloaded"}, nil
}
// handleBotRangesReload republishes the on-disk AI-crawler range overlay that
// `csm update-bot-ranges` just refreshed, so the new ranges apply without a
// daemon restart.
func (c *ControlListener) handleBotRangesReload(_ json.RawMessage) (any, error) {
if err := c.d.reloadBotRanges(); err != nil {
return nil, err
}
return map[string]string{"status": "reloaded"}, nil
}
// handleHistorySince streams every history-bucket finding newer than
// the supplied cutoff. Used by `csm export --since` for SIEM backfill;
// the daemon-side bbolt cursor seek is materially faster than
// pagination + client-side filtering for hosts with large histories.
func (c *ControlListener) handleHistorySince(argsRaw json.RawMessage) (any, error) {
var args control.HistorySinceArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if args.Since == "" {
return nil, fmt.Errorf("since is required (RFC 3339)")
}
since, err := time.Parse(time.RFC3339, args.Since)
if err != nil {
return nil, fmt.Errorf("parsing since: %w", err)
}
sdb := store.Global()
if sdb == nil {
return nil, fmt.Errorf("bbolt store not available")
}
findings := sdb.ReadHistorySince(since)
// ReadHistorySince returns newest-first; reverse for chronological
// output so JSONL consumers see the same order they would from a
// live tail.
for i, j := 0, len(findings)-1; i < j; i, j = i+1, j-1 {
findings[i], findings[j] = findings[j], findings[i]
}
return control.HistorySinceResult{Findings: findings}, nil
}
// handleStoreExport writes a tar+zstd backup containing the live bbolt
// snapshot, the state directory, and the signature-rules cache. The
// daemon is the single source of truth for paths; the CLI only supplies
// where to write the archive. Import deliberately does NOT route through
// the socket -- it requires a stopped daemon.
const exportStagingMaxAge = 24 * time.Hour
// prepareExportStagingPath creates a private per-request directory so
// concurrent exports of the same basename cannot overwrite each other. Old
// request directories are removed here so a client that dies after the daemon
// replies cannot leave state storage growing forever.
func prepareExportStagingPath(statePath, dstPath string, now time.Time) (string, error) {
base := filepath.Base(dstPath)
if base == "." || base == ".." || base == string(filepath.Separator) || base == "" {
return "", fmt.Errorf("destination must name an archive file")
}
exportDir := filepath.Join(statePath, "exports")
if err := os.MkdirAll(exportDir, 0o700); err != nil {
return "", err
}
// #nosec G302 -- exportDir is a directory; 0700 is already the tightest
// mode that still lets the daemon traverse into it.
if err := os.Chmod(exportDir, 0o700); err != nil {
return "", err
}
entries, err := os.ReadDir(exportDir)
if err != nil {
return "", err
}
cutoff := now.Add(-exportStagingMaxAge)
for _, entry := range entries {
info, infoErr := entry.Info()
if infoErr == nil && info.ModTime().Before(cutoff) {
_ = os.RemoveAll(filepath.Join(exportDir, entry.Name()))
}
}
requestDir, err := os.MkdirTemp(exportDir, "export-")
if err != nil {
return "", err
}
return filepath.Join(requestDir, base), nil
}
func (c *ControlListener) handleStoreExport(argsRaw json.RawMessage) (any, error) {
var args control.StoreExportArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if args.DstPath == "" {
return nil, fmt.Errorf("dst_path is required")
}
sdb := store.Global()
if sdb == nil {
return nil, fmt.Errorf("bbolt store not available")
}
cfg := c.d.currentCfg()
hostname, _ := os.Hostname()
pi := platform.Detect()
dstPath := args.DstPath
staged := false
var err error
if args.Stage {
dstPath, err = prepareExportStagingPath(cfg.StatePath, args.DstPath, time.Now())
if err != nil {
return nil, fmt.Errorf("creating export staging dir: %w", err)
}
staged = true
}
res, err := sdb.Export(store.ExportOptions{
StatePath: cfg.StatePath,
RulesPath: cfg.Signatures.RulesDir,
DstPath: dstPath,
Manifest: store.Manifest{
CSMVersion: c.d.version,
SourceHostname: hostname,
SourcePlatform: map[string]string{
"os": string(pi.OS),
"os_version": pi.OSVersion,
"panel": string(pi.Panel),
"webserver": string(pi.WebServer),
},
},
})
if err != nil {
if staged {
_ = os.RemoveAll(filepath.Dir(dstPath))
}
return nil, err
}
return control.StoreExportResult{
Path: res.Path,
Bytes: res.Bytes,
ArchiveSHA256: res.ArchiveSHA256,
BboltSHA256: res.BboltSHA256,
}, nil
}
package daemon
import (
"encoding/json"
"fmt"
"strings"
"time"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/incident"
)
const (
defaultIncidentListLimit = 100
maxIncidentListLimit = 1000
defaultIncidentBulkLimit = 100
maxIncidentBulkLimit = 1000
)
// handleIncidentsList returns a bounded, newest-first incident page.
func (c *ControlListener) handleIncidentsList(argsRaw json.RawMessage) (any, error) {
var args control.IncidentListArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, err
}
}
statuses, statusLabel, err := incidentListStatusFilter(args.Status)
if err != nil {
return nil, err
}
offset := args.Offset
if offset < 0 {
offset = 0
}
limit := args.Limit
if args.All {
limit = 0
} else {
if limit <= 0 {
limit = defaultIncidentListLimit
}
if limit > maxIncidentListLimit {
limit = maxIncidentListLimit
}
}
co := IncidentCorrelator()
items, total := co.SnapshotPageStatuses(statuses, offset, limit)
return control.IncidentListResult{
Items: items,
Total: total,
Offset: offset,
Limit: limit,
Status: statusLabel,
}, nil
}
func incidentListStatusFilter(status string) ([]incident.Status, string, error) {
switch strings.ToLower(strings.TrimSpace(status)) {
case "", "all":
return nil, "all", nil
case "active":
return []incident.Status{incident.StatusOpen, incident.StatusContained}, "active", nil
case string(incident.StatusOpen):
return []incident.Status{incident.StatusOpen}, string(incident.StatusOpen), nil
case string(incident.StatusContained):
return []incident.Status{incident.StatusContained}, string(incident.StatusContained), nil
case string(incident.StatusResolved):
return []incident.Status{incident.StatusResolved}, string(incident.StatusResolved), nil
case string(incident.StatusDismissed):
return []incident.Status{incident.StatusDismissed}, string(incident.StatusDismissed), nil
default:
return nil, "", fmt.Errorf("unknown status: %q", status)
}
}
func incidentBulkStatusFilter(status string) ([]incident.Status, string, error) {
switch strings.ToLower(strings.TrimSpace(status)) {
case "", "active":
return []incident.Status{incident.StatusOpen, incident.StatusContained}, "active", nil
case string(incident.StatusOpen):
return []incident.Status{incident.StatusOpen}, string(incident.StatusOpen), nil
case string(incident.StatusContained):
return []incident.Status{incident.StatusContained}, string(incident.StatusContained), nil
default:
return nil, "", fmt.Errorf("bulk status source must be active, open, or contained")
}
}
// handleIncidentsShow returns one incident by id; ErrIncidentNotFound on miss.
func (c *ControlListener) handleIncidentsShow(argsRaw json.RawMessage) (any, error) {
var args control.IncidentShowArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, err
}
}
co := IncidentCorrelator()
inc, ok := co.Get(args.ID)
if !ok {
return nil, incident.ErrIncidentNotFound
}
return inc, nil
}
// handleIncidentsStatus transitions an incident's status. Returns
// {"ok": true} on success; ErrIncidentNotFound or validation error on
// failure.
func (c *ControlListener) handleIncidentsStatus(argsRaw json.RawMessage) (any, error) {
var args control.IncidentStatusArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, err
}
}
co := IncidentCorrelator()
if err := co.SetStatus(args.ID, incident.Status(args.Status), args.Details); err != nil {
return nil, err
}
return map[string]bool{"ok": true}, nil
}
func (c *ControlListener) handleIncidentsBulkStatus(argsRaw json.RawMessage) (any, error) {
var args control.IncidentBulkStatusArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, err
}
}
statuses, statusLabel, err := incidentBulkStatusFilter(args.Status)
if err != nil {
return nil, err
}
to := incident.Status(strings.ToLower(strings.TrimSpace(args.To)))
if to == "" {
to = incident.StatusResolved
}
if to != incident.StatusResolved && to != incident.StatusDismissed {
return nil, fmt.Errorf("bulk status target must be resolved or dismissed")
}
if args.OlderThanSeconds < 0 {
return nil, fmt.Errorf("older-than must be positive")
}
olderThan := time.Duration(args.OlderThanSeconds) * time.Second
if olderThan <= 0 && args.LastSeenBefore.IsZero() {
return nil, fmt.Errorf("bulk status requires --older-than or --last-seen-before")
}
limit := args.Limit
if limit <= 0 {
limit = defaultIncidentBulkLimit
}
if limit > maxIncidentBulkLimit {
limit = maxIncidentBulkLimit
}
if args.Apply && !args.Confirm {
return nil, fmt.Errorf("bulk status apply requires confirmation")
}
dryRun := !args.Apply
co := IncidentCorrelator()
res, err := co.BulkSetStatus(incident.BulkStatusFilter{
FromStatuses: statuses,
To: to,
OlderThan: olderThan,
LastSeenBefore: args.LastSeenBefore,
Kind: incident.Kind(strings.TrimSpace(args.Kind)),
Domain: strings.TrimSpace(args.Domain),
Account: strings.TrimSpace(args.Account),
Mailbox: strings.TrimSpace(args.Mailbox),
Limit: limit,
DryRun: dryRun,
Details: strings.TrimSpace(args.Details),
})
if err != nil {
return nil, err
}
return control.IncidentBulkStatusResult{
DryRun: dryRun,
Matched: res.Matched,
Updated: res.Updated,
Limit: limit,
Status: statusLabel,
To: string(to),
OlderThanSeconds: args.OlderThanSeconds,
LastSeenBefore: args.LastSeenBefore,
Items: res.Items,
}, nil
}
package daemon
import (
"encoding/json"
"fmt"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/store"
)
// handleScanEnqueue validates the request and submits a new full-scan job to
// the ScanJobManager. Scope="account" scans a single account; Scope="all"
// enqueues a server-wide scan. The control payload is translated into
// checks.AccountScanOptions with the full-scan option set (MaxFiles=0,
// ForceContent=true, ForceFileIndex=true).
func (c *ControlListener) handleScanEnqueue(argsRaw json.RawMessage) (any, error) {
if c.scanJobs == nil {
return nil, fmt.Errorf("scan job manager not available")
}
var req control.ScanEnqueueRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
cfg := c.d.currentCfg()
opts := checks.FullScanOptions(cfg, req.RespectIgnores)
switch req.Scope {
case "account":
if req.Target == "" {
return nil, fmt.Errorf("target is required")
}
if !control.ValidScanAccountTarget(req.Target) {
return nil, fmt.Errorf("invalid account target %q", req.Target)
}
id, err := c.scanJobs.Enqueue("account", req.Target, opts, req.Quarantine)
if err != nil {
return nil, fmt.Errorf("enqueue: %w", err)
}
return control.ScanEnqueueResponse{JobID: id, State: "queued"}, nil
case "all":
// Target must be empty or the literal "all"; it must not look like an
// account name or path component — the daemon normalises it to "all".
if req.Target != "" && req.Target != "all" {
return nil, fmt.Errorf("invalid target %q for scope \"all\": must be empty or \"all\"", req.Target)
}
// Defense-in-depth (the CLI also rejects this): server-wide quarantine
// would remediate across every account from one audit pass. Refuse it;
// quarantine is a per-account post-review action.
if req.Quarantine {
return nil, fmt.Errorf("quarantine is not supported with scope \"all\"")
}
id, err := c.scanJobs.Enqueue("all", "all", opts, req.Quarantine)
if err != nil {
return nil, fmt.Errorf("enqueue: %w", err)
}
return control.ScanEnqueueResponse{JobID: id, State: "queued"}, nil
default:
return nil, fmt.Errorf("unsupported scope %q: must be \"account\" or \"all\"", req.Scope)
}
}
// handleScanStatus returns the status of one job (when JobID is set) or the
// full list of jobs ordered newest-first. Unknown IDs produce an error.
func (c *ControlListener) handleScanStatus(argsRaw json.RawMessage) (any, error) {
if c.scanJobs == nil {
return nil, fmt.Errorf("scan job manager not available")
}
var req control.ScanStatusRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if req.JobID != "" {
rec, ok := c.scanJobs.Progress(req.JobID)
if !ok {
return nil, fmt.Errorf("job not found: %q", req.JobID)
}
return control.ScanStatusResponse{Job: &rec}, nil
}
jobs, err := c.scanJobs.ListJobs()
if err != nil {
return nil, fmt.Errorf("listing jobs: %w", err)
}
if jobs == nil {
jobs = []store.ScanJobRecord{}
}
return control.ScanStatusResponse{Jobs: jobs}, nil
}
// handleScanReport returns the job record plus a paginated slice of its
// findings. JobID is required; Offset and Limit follow the usual page semantics
// (Limit=0 returns all findings).
func (c *ControlListener) handleScanReport(argsRaw json.RawMessage) (any, error) {
if c.scanJobs == nil {
return nil, fmt.Errorf("scan job manager not available")
}
var req control.ScanReportRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if req.JobID == "" {
return nil, fmt.Errorf("job_id is required")
}
rec, ok := c.scanJobs.Progress(req.JobID)
if !ok {
return nil, fmt.Errorf("job not found: %q", req.JobID)
}
if req.Offset < 0 {
req.Offset = 0
}
findings, total, err := c.scanJobs.ListFindings(req.JobID, req.Offset, req.Limit)
if err != nil {
return nil, fmt.Errorf("listing findings: %w", err)
}
if findings == nil {
findings = []alert.Finding{}
}
return control.ScanReportResponse{
Job: rec,
Findings: findings,
Total: total,
}, nil
}
// handleScanCancel cancels the job with the given ID. If the job is queued it
// will be marked canceled before the worker processes it; if it is running its
// context is canceled and partial findings are retained. Unknown or already
// terminal IDs produce an error.
func (c *ControlListener) handleScanCancel(argsRaw json.RawMessage) (any, error) {
if c.scanJobs == nil {
return nil, fmt.Errorf("scan job manager not available")
}
var req control.ScanCancelRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if req.JobID == "" {
return nil, fmt.Errorf("job_id is required")
}
if err := c.scanJobs.Cancel(req.JobID); err != nil {
return nil, fmt.Errorf("cancel: %w", err)
}
// Return the job's current state. The worker may not have transitioned it
// yet; the caller polls via scan.status for the terminal state.
rec, ok := c.scanJobs.Progress(req.JobID)
state := "canceling"
if ok {
state = rec.State
}
return control.ScanCancelResponse{JobID: req.JobID, State: state}, nil
}
package daemon
import (
"bufio"
"encoding/json"
"fmt"
"net"
"os"
"path/filepath"
"time"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/obs"
)
// controlSocketPath is the Unix socket the daemon binds for
// CLI-to-daemon IPC. A var (not const) so tests can redirect under
// t.TempDir(); production default matches internal/control.DefaultSocketPath.
var controlSocketPath = "/var/run/csm/control.sock"
// controlRequestTimeout caps how long a single client request can block
// the listener. Handlers that legitimately take longer (tier.run on a
// large server) run on the accepting goroutine, so the timeout applies
// to reading the request line and writing the response, not to the
// handler body itself.
const controlRequestTimeout = 2 * time.Second
// ControlListener serves the local command-line client over a Unix
// socket. One request and one response per connection, line-framed JSON.
// The daemon keeps exclusive ownership of the bbolt store; this listener
// is the only reason any CLI command needs to reach into daemon state.
type ControlListener struct {
d *Daemon
listener net.Listener
phprelay *PHPRelayController // wired by Phase O2; may be nil in tests / pre-wiring
scanJobs *ScanJobManager // wired after startScanJobManager(); nil until then
}
// NewControlListener creates the socket, enforces 0600 perms, and
// returns a listener that the daemon wires into its goroutine pool.
func NewControlListener(d *Daemon) (*ControlListener, error) {
socketDir := filepath.Dir(controlSocketPath)
if err := os.MkdirAll(socketDir, 0750); err != nil {
return nil, fmt.Errorf("creating socket dir: %w", err)
}
// Stale socket from a previous crash would make Listen fail with
// EADDRINUSE. Remove before binding; the file is process-owned.
_ = os.Remove(controlSocketPath)
ln, err := net.Listen("unix", controlSocketPath)
if err != nil {
return nil, fmt.Errorf("listening on %s: %w", controlSocketPath, err)
}
// 0600 root-only: the CLI client also runs as root (fanotify,
// nftables, cpanel APIs all require it), so no group is needed.
if err := os.Chmod(controlSocketPath, 0600); err != nil {
_ = ln.Close()
return nil, fmt.Errorf("chmod socket: %w", err)
}
return &ControlListener{d: d, listener: ln}, nil
}
// Run accepts connections until stopCh closes. Each connection is
// handled on its own goroutine so a slow request never stalls the
// accept loop.
func (c *ControlListener) Run(stopCh <-chan struct{}) {
for {
conn, err := c.listener.Accept()
if err != nil {
select {
case <-stopCh:
return
default:
csmlog.Warn("control listener accept error", "err", err)
time.Sleep(100 * time.Millisecond)
continue
}
}
// SO_PEERCRED defence-in-depth: socket perms are already 0600
// (root-only), but if a future install pattern relaxes that --
// or a confused-deputy mount makes the socket reachable from a
// non-root namespace -- the kernel-supplied credentials force
// us to refuse any non-root caller before reading a byte of
// payload. The socket perms remain the primary defence; this
// is an extra rejection layer that the kernel cannot lie about.
if err := verifyControlPeer(conn); err != nil {
csmlog.Warn("control listener rejecting peer", "err", err)
_ = conn.Close()
continue
}
obs.SafeGo("control-conn", func() { c.handleConnection(conn) })
}
}
// verifyControlPeer is defined per-OS:
// - linux: reads SO_PEERCRED and refuses non-root callers
// - other: no-op (the listener is Linux-only in production; macOS dev
// builds still need the package to compile)
// Stop closes the listener and removes the socket file. Safe to call
// after Run has already returned.
func (c *ControlListener) Stop() {
_ = c.listener.Close()
_ = os.Remove(controlSocketPath)
}
// handleConnection reads one request, dispatches it, writes one
// response, and closes. The short timeout applies to I/O only; the
// handler itself can take as long as the underlying work requires.
func (c *ControlListener) handleConnection(conn net.Conn) {
defer func() { _ = conn.Close() }()
// Read deadline. Writes use a separate deadline set after the
// handler returns so slow scans don't count against the reader.
_ = conn.SetReadDeadline(time.Now().Add(controlRequestTimeout))
scanner := bufio.NewScanner(conn)
// Requests and responses are single-line JSON; on a compromised
// server a dry-run tier scan can legitimately return a findings
// list that exceeds the old 1 MiB cap (think: many WordPress
// installs, each with multiple infected files). The socket is
// root-only 0600 so the original DoS guard no longer applies —
// cap the buffer at 16 MiB so large finding lists round-trip.
scanner.Buffer(make([]byte, 0, 64*1024), 16*1024*1024)
if !scanner.Scan() {
return
}
line := scanner.Bytes()
resp := c.dispatch(line)
payload, err := json.Marshal(resp)
if err != nil {
// Marshalling a Response should not fail; if it does, fall back
// to a minimal error response the client can still parse.
payload = []byte(`{"ok":false,"error":"internal: response marshal failed"}`)
}
payload = append(payload, '\n')
_ = conn.SetWriteDeadline(time.Now().Add(controlRequestTimeout))
_, _ = conn.Write(payload)
}
//go:build linux
package daemon
import (
"fmt"
"net"
"golang.org/x/sys/unix"
)
var controlPeerRequiredUID uint32
// verifyControlPeer reads SO_PEERCRED and refuses any caller whose
// effective uid is not root. Returns nil when the peer is acceptable.
func verifyControlPeer(conn net.Conn) error {
uc, ok := conn.(*net.UnixConn)
if !ok || uc == nil {
return fmt.Errorf("peer credentials: unsupported connection %T", conn)
}
raw, err := uc.SyscallConn()
if err != nil {
return fmt.Errorf("peer raw conn: %w", err)
}
var ucred *unix.Ucred
var optErr error
if err := raw.Control(func(fd uintptr) {
// #nosec G115 -- POSIX fd fits in int on Linux.
ucred, optErr = unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED)
}); err != nil {
return fmt.Errorf("peer credentials: %w", err)
}
if optErr != nil {
return fmt.Errorf("peer credentials: %w", optErr)
}
if ucred == nil {
return fmt.Errorf("peer credentials: empty result")
}
if ucred.Uid != controlPeerRequiredUID {
return fmt.Errorf("peer uid=%d, want %d", ucred.Uid, controlPeerRequiredUID)
}
return nil
}
package daemon
import (
"encoding/json"
"fmt"
"net"
"github.com/pidginhost/csm/internal/attackdb"
"github.com/pidginhost/csm/internal/control"
)
// handleThreatForget drops one address's record from the attack database.
//
// Most score contributions last until the record expires after 90 days.
// Fixing a detection that attributed events to the wrong address stops new
// events but leaves those contributions behind.
//
// The Web UI's clear and whitelist actions also change enforcement.
// This command leaves enforcement and event history intact, so a cleared
// address is scored again from scratch on its next finding.
func (c *ControlListener) handleThreatForget(argsRaw json.RawMessage) (any, error) {
var args control.FirewallIPArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
ip := net.ParseIP(args.IP)
if ip == nil {
return nil, fmt.Errorf("invalid ip: %q", args.IP)
}
args.IP = ip.String()
adb := attackdb.Global()
if adb == nil {
return nil, fmt.Errorf("attack database unavailable")
}
// Read and remove under one lock so concurrent requests cannot claim
// the same record, or report counts from before a concurrent finding.
res := control.ThreatForgetResult{IP: args.IP}
for _, rec := range adb.ForgetIP(ip) {
res.Found = true
res.Score = max(res.Score, attackdb.ComputeScore(rec))
res.Events += rec.EventCount
}
if !res.Found {
res.Message = fmt.Sprintf("No local threat record for %s; nothing cleared", args.IP)
return res, nil
}
// Persist immediately rather than waiting for the 30s background saver:
// Flush does not report persistence errors, so this remains best effort.
// A lookup afterwards cannot verify persistence, and new findings may
// legitimately have created a fresh record by then.
_ = adb.Flush()
res.Message = fmt.Sprintf(
"Cleared local threat record for %s (was score %d/100, %d attack events); block, allow and whitelist entries are unchanged",
args.IP, res.Score, res.Events)
return res, nil
}
package daemon
import (
"net"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/platform"
)
// infraHostnames partitions the cfg.InfraIPs operator list into the
// subset that is hostnames (not literal IPs or CIDRs). Hostnames get
// DNS-refreshed into the engine's infra-block guard so operators can
// list panel hosts by name and have them stay protected as the
// underlying address rotates.
func infraHostnames(entries []string) []string {
out := make([]string, 0, len(entries))
for _, e := range entries {
s := strings.TrimSpace(e)
if s == "" {
continue
}
if _, _, err := net.ParseCIDR(s); err == nil {
continue
}
if ip := net.ParseIP(s); ip != nil {
continue
}
out = append(out, s)
}
return out
}
// containsString returns true when s appears in haystack. Linear scan
// because the call site loops short DynDNS host lists where a map
// would not pay for itself.
func containsString(haystack []string, s string) bool {
for _, h := range haystack {
if h == s {
return true
}
}
return false
}
// expandWithCorrelation runs cross-account correlation over a dispatch
// batch and appends any synthesized findings that are not already present.
// The scan runner may have already produced the same synthetic findings, so
// this helper must be idempotent to avoid double-alerting the first batch.
func expandWithCorrelation(findings []alert.Finding, now time.Time) []alert.Finding {
if len(findings) == 0 {
return findings
}
platform.Detect()
seen := make(map[string]struct{})
for i := range findings {
if !checks.IsDerivedCorrelationCheck(findings[i].Check) {
continue
}
if findings[i].Timestamp.IsZero() {
findings[i].Timestamp = now
}
seen[findings[i].Key()] = struct{}{}
}
res := checks.CorrelateBatchFindings(findings)
for i := range res.Derived {
if res.Derived[i].Timestamp.IsZero() {
res.Derived[i].Timestamp = now
}
key := res.Derived[i].Key()
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
findings = append(findings, res.Derived[i])
}
checks.ReportUnattributedCorrelation(res.Unattributed)
return findings
}
package daemon
import (
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/emailspool"
"gopkg.in/yaml.v3"
)
// userDomainsResolver resolves a cPanel user to the lowercased, IDN-normalised
// set of domains the account owns. TTL-based cache; safe for concurrent use.
type userDomainsResolver struct {
root string
ttl time.Duration
mu sync.Mutex
cache map[string]userDomainsCacheEntry
}
type userDomainsCacheEntry struct {
domains map[string]struct{}
fetched time.Time
err error
}
// newUserDomainsResolver returns a resolver reading from /var/cpanel/userdata/.
// Wired by daemon startup in O2; kept now so that future call sites compile.
//
//nolint:unused // consumed by daemon wiring (Task O2)
func newUserDomainsResolver() *userDomainsResolver {
return newUserDomainsResolverWithRoot("/var/cpanel/userdata", 5*time.Minute)
}
func newUserDomainsResolverWithRoot(root string, ttl time.Duration) *userDomainsResolver {
return &userDomainsResolver{
root: root,
ttl: ttl,
cache: make(map[string]userDomainsCacheEntry),
}
}
// Domains returns the cPanel user's authorised domain set. Returns the
// (possibly cached) error if the user's userdata is unreadable; callers
// must treat an error result as "skip the From-mismatch signal" rather than
// falsely amplifying.
func (r *userDomainsResolver) Domains(user string) (map[string]struct{}, error) {
if user == "" {
return nil, errors.New("empty user")
}
r.mu.Lock()
if e, ok := r.cache[user]; ok && time.Since(e.fetched) < r.ttl {
r.mu.Unlock()
return e.domains, e.err
}
r.mu.Unlock()
set, err := r.read(user)
r.mu.Lock()
r.cache[user] = userDomainsCacheEntry{domains: set, fetched: time.Now(), err: err}
r.mu.Unlock()
return set, err
}
// Invalidate removes the cached entry for user. Callers wire this to
// inotify on /var/cpanel/userdata/<user>/.
func (r *userDomainsResolver) Invalidate(user string) {
r.mu.Lock()
delete(r.cache, user)
r.mu.Unlock()
}
func (r *userDomainsResolver) read(user string) (map[string]struct{}, error) {
path := filepath.Join(r.root, user, "main")
// #nosec G304 -- r.root is fixed at /var/cpanel/userdata/, user is a
// cPanel-managed account name validated by the resolver's caller; the
// resulting path is constrained to the cpanel userdata tree.
data, err := os.ReadFile(path)
if err != nil {
return map[string]struct{}{}, fmt.Errorf("read %s: %w", path, err)
}
var raw struct {
MainDomain string `yaml:"main_domain"`
AddonDomains map[string]string `yaml:"addon_domains"`
ParkedDomains []string `yaml:"parked_domains"`
SubDomains []string `yaml:"sub_domains"`
}
if err := yaml.Unmarshal(data, &raw); err != nil {
return map[string]struct{}{}, fmt.Errorf("parse %s: %w", path, err)
}
set := make(map[string]struct{}, 8)
add := func(d string) {
d = strings.TrimSpace(d)
if d == "" {
return
}
// ExtractDomain handles IDN normalisation and lowercasing for free.
// It expects an addr-style input but happily round-trips bare hosts.
norm := emailspool.ExtractDomain("anyone@" + d)
if norm == "" {
norm = strings.ToLower(d)
}
set[norm] = struct{}{}
}
add(raw.MainDomain)
for k := range raw.AddonDomains {
add(k)
}
for _, d := range raw.ParkedDomains {
add(d)
}
for _, d := range raw.SubDomains {
add(d)
}
return set, nil
}
// IsAuthorisedFromDomain reports whether fromDomain is one of the user's
// domains, accounting for subdomain inclusion (a sub.example.com From is
// authorised if example.com is in the set, but the reverse is NOT true).
func IsAuthorisedFromDomain(fromDomain string, authSet map[string]struct{}) bool {
if fromDomain == "" || len(authSet) == 0 {
return false
}
for base := range authSet {
if emailspool.IsSubdomainOrEqual(fromDomain, base) {
return true
}
}
return false
}
package daemon
import (
"sort"
"time"
)
// credentialStuffingDetector tracks, per source IP, the set of distinct
// accounts hit by auth failures inside a sliding window and flags an IP that
// targets many distinct accounts. This is the breadth signal of credential
// stuffing / password spraying -- one source trying many accounts, often with
// only one or two attempts each -- which the count-based pam_bruteforce
// detector (depth: many failures, any account) does not capture. CSM's auth
// sources never expose the attempted password, so the detector keys on the
// distinct-account behavioral signature rather than a password fingerprint.
//
// Concurrency: callers serialize access (the PAM listener holds its mutex
// across Record), so the detector takes no lock of its own.
type credentialStuffingDetector struct {
distinctAccounts int
window time.Duration
now func() time.Time
perIP map[string]*credStuffState
// maxTrackedIPs bounds the live map under sustained source-IP churn.
// Once at the cap the oldest-by-lastSeen entry is evicted before insert.
maxTrackedIPs int
}
type credStuffState struct {
accounts map[string]time.Time
lastSeen time.Time
// fired is set once the IP crosses the distinct-account threshold so a
// single active window yields one finding, not one per additional account.
fired bool
}
// newCredentialStuffingDetector builds a detector. A distinctAccounts
// threshold below 2 has no breadth meaning and is clamped to 2.
func newCredentialStuffingDetector(distinctAccounts int, window time.Duration, now func() time.Time) *credentialStuffingDetector {
if distinctAccounts < 2 {
distinctAccounts = 2
}
if now == nil {
now = time.Now
}
return &credentialStuffingDetector{
distinctAccounts: distinctAccounts,
window: window,
now: now,
perIP: make(map[string]*credStuffState),
maxTrackedIPs: 10000,
}
}
// Record ingests one auth failure for (ip, account). It returns the sorted
// distinct-account list and true exactly once -- when the IP first crosses
// the distinct-account threshold inside the window. Empty ip or account is
// ignored. A failure whose IP has been silent longer than the window resets
// that IP's distinct set so a fresh campaign does not inherit cold counts.
func (d *credentialStuffingDetector) Record(ip, account string) ([]string, bool) {
if d == nil || ip == "" || account == "" {
return nil, false
}
now := d.now()
state, ok := d.perIP[ip]
if ok {
d.pruneAccountWindow(state, now)
}
if ok && len(state.accounts) == 0 {
state = nil
delete(d.perIP, ip)
}
if state == nil {
if d.maxTrackedIPs > 0 && len(d.perIP) >= d.maxTrackedIPs {
d.evictOldest()
}
state = &credStuffState{accounts: make(map[string]time.Time)}
d.perIP[ip] = state
}
state.accounts[account] = now
state.lastSeen = now
if len(state.accounts) < d.distinctAccounts {
state.fired = false
return nil, false
}
if state.fired {
return nil, false
}
state.fired = true
accounts := make([]string, 0, len(state.accounts))
for a := range state.accounts {
accounts = append(accounts, a)
}
sort.Strings(accounts)
return accounts, true
}
// ClearAccount forgets one account's failures from ip: that account logged
// in, so its earlier failures were the user's own. Failures against other
// accounts stay counted; the entry goes only when none remain, and a set
// that drops back under the threshold may fire again once it regrows.
func (d *credentialStuffingDetector) ClearAccount(ip, account string) {
if d == nil {
return
}
state, ok := d.perIP[ip]
if !ok {
return
}
delete(state.accounts, account)
if len(state.accounts) == 0 {
delete(d.perIP, ip)
return
}
if len(state.accounts) < d.distinctAccounts {
state.fired = false
}
}
// PruneStale drops per-IP entries whose lastSeen is older than the window.
// Called from the PAM listener cleanup loop so the detector does not grow
// without bound between window resets. Returns the number pruned.
func (d *credentialStuffingDetector) PruneStale(now time.Time) int {
if d == nil {
return 0
}
pruned := 0
for ip, state := range d.perIP {
d.pruneAccountWindow(state, now)
if len(state.accounts) == 0 {
delete(d.perIP, ip)
pruned++
continue
}
if len(state.accounts) < d.distinctAccounts {
state.fired = false
}
}
return pruned
}
// Clear removes all tracked breadth state for ip after a successful login.
func (d *credentialStuffingDetector) Clear(ip string) {
if d == nil || ip == "" {
return
}
delete(d.perIP, ip)
}
// Configure applies live threshold/window changes. Callers hold their
// external mutex, matching Record's synchronization contract.
func (d *credentialStuffingDetector) Configure(distinctAccounts int, window time.Duration, now time.Time) {
if d == nil {
return
}
if distinctAccounts < 2 {
distinctAccounts = 2
}
d.distinctAccounts = distinctAccounts
d.window = window
for ip, state := range d.perIP {
d.pruneAccountWindow(state, now)
if len(state.accounts) == 0 {
delete(d.perIP, ip)
continue
}
if len(state.accounts) < d.distinctAccounts {
state.fired = false
}
}
}
func (d *credentialStuffingDetector) pruneAccountWindow(state *credStuffState, now time.Time) {
if state == nil {
return
}
for account, seen := range state.accounts {
if now.Sub(seen) > d.window {
delete(state.accounts, account)
}
}
}
// evictOldest removes the entry with the smallest lastSeen so a fresh insert
// stays within maxTrackedIPs. Linear scan, bounded by the cap.
func (d *credentialStuffingDetector) evictOldest() {
var oldestIP string
var oldestAt time.Time
first := true
for ip, state := range d.perIP {
if first || state.lastSeen.Before(oldestAt) {
oldestIP, oldestAt, first = ip, state.lastSeen, false
}
}
if oldestIP != "" {
delete(d.perIP, oldestIP)
}
}
//go:build linux
package daemon
import "github.com/pidginhost/csm/internal/platform"
// cronSpoolDir returns the per-user crontab directory the watcher marks and
// matches event paths against: cronie's /var/spool/cron, or Debian cron's
// /var/spool/cron/crontabs. platform.Detect caches its answer, so this is
// cheap enough for the per-event path checks.
func cronSpoolDir() string {
if cronSpoolWatchDir != "" {
return cronSpoolWatchDir
}
return platform.Detect().CronSpoolDir()
}
package daemon
import (
"context"
"errors"
"fmt"
"net"
"net/http"
"os"
"os/exec"
"os/signal"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/attackdb"
"github.com/pidginhost/csm/internal/auditd"
"github.com/pidginhost/csm/internal/blockdigest"
"github.com/pidginhost/csm/internal/broadcast"
"github.com/pidginhost/csm/internal/challenge"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailav"
"github.com/pidginhost/csm/internal/emailspool"
"github.com/pidginhost/csm/internal/eximlog"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/firewall/rollback"
"github.com/pidginhost/csm/internal/geoip"
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/integrity"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/maillog"
"github.com/pidginhost/csm/internal/mailranges"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/modsec"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/phptaintworker"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/sdnotify"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/threatintel"
"github.com/pidginhost/csm/internal/updatecheck"
"github.com/pidginhost/csm/internal/verdict"
"github.com/pidginhost/csm/internal/webui"
"github.com/pidginhost/csm/internal/yara"
"github.com/pidginhost/csm/internal/yaraworker"
)
const eximMainlogPath = "/var/log/exim_mainlog"
// Daemon is the main persistent monitoring process.
type Daemon struct {
cfg *config.Config
store *state.Store
lock *state.LockFile
binaryPath string
binaryHash binaryHashCache
logWatchers []*LogWatcher
logWatchersMu sync.Mutex
fileMonitor *FileMonitor
fileMonitorMu sync.RWMutex
hijackDetector *PasswordHijackDetector
pamListener *PAMListener
controlListener *ControlListener
spoolWatcher *SpoolWatcher
spoolWatcherMu sync.Mutex
forwarderWatcher *ForwarderWatcher
emailQuarantine *emailav.Quarantine
webServer *webui.Server
challengeServer *challenge.Server
ipList *challenge.IPList
challengeGate challenge.PortGate
fwEngine *firewall.Engine
// fwActions retains the durable-action boundary even after failed startup,
// so pending actions remain recoverable while the firewall is unmanaged.
fwActions firewallActionBoundary
fwStartupError string // finalized before status servers start
baselineMu sync.Mutex // serialises CmdBaseline handler runs
geoipDB *geoip.DB
geoipMu sync.Mutex // protects geoipDB for publishGeoIP
version string
blockDigest *blockdigest.Collector
alertCh chan alert.Finding
alertQueue *queuehealth.Tracker
queueSourcesMu sync.RWMutex
queueSources map[string]queueSource
// alertHold, while open, keeps the dispatcher draining alertCh into its
// batch without dispatching, so realtime producers (which never block)
// lose nothing during the synchronous startup baseline. The ingest queue
// health is held with it, so the baseline is not reported as a stall. Closed by
// releaseAlertDispatch once the baseline has published; nil means the
// dispatcher never holds.
alertHold chan struct{}
alertReleaseOnce sync.Once
droppedAlerts int64 // atomic counter for alert channel backpressure drops
stopCh chan struct{}
scanCtx context.Context
scanCancel context.CancelFunc // cancels in-flight periodic scans on shutdown
modsecReload checks.ModSecReloadReconciler
// modsecRegistry carries what the last rule-action refresh learned, so a
// refresh that changes nothing costs content hashing instead of a
// platform probe plus a full reparse of the vendor rule tree.
modsecRegistry modsecRegistryState
abuseReportStop chan struct{}
abuseReportDone chan struct{}
wg sync.WaitGroup
smtpAuthTracker *smtpAuthTracker
smtpProbeTracker *smtpProbeTracker
mailAuthTracker *mailAuthTracker
authBackend *authBackendHealth
startTime time.Time
// yaraSup is the supervised YARA-X worker, wired up when
// yaraWorkerOn(cfg) is true (the default; see ROADMAP item 2).
// Nil when the in-process scanner is in use
// (cfg.Signatures.YaraWorkerEnabled explicitly set to false).
yaraSup *yaraworker.Supervisor
yaraCrashMu sync.Mutex
yaraLastCrashAlert time.Time
realtimeRulesMu sync.Mutex
realtimeRulesState string
// phpTaintSup is the mandatory process boundary around the PHP taint
// parser. It starts its child lazily on the first admitted source.
phpTaintSup *phptaintworker.Supervisor
// forceFullRescan is armed by the signature watcher
// (sig_watch.go) when any tracked rule file's content changes.
// The deep-tier scheduler reads + clears the flag at the start
// of each tick; when set, the tick bypasses the fanotify
// short-list and runs the full account tree against the new
// ruleset.
forceFullRescan atomic.Bool
// policies holds the email PHP-relay pattern policies
// (suspicious/safe x-mailer classes, HTTP proxy ranges) loaded
// from EmailProtection.PHPRelay.PoliciesDir. Initialised in O2
// (daemon wiring); stays nil until then. The SIGHUP path
// nil-guards the Reload call so this commit is a no-op at
// runtime until O2 lands.
policies *emailspool.Policies
// botVerifier is the async rDNS bot verifier. Retained so the SIGHUP
// path can push reloaded reputation.verified_bots into it. nil when
// bot verification is disabled or no store is available.
botVerifier *threatintel.AsyncBotVerifier
// PHP-relay components wired by startPHPRelay (Linux only). The fields are
// declared cross-platform but stay nil on non-cPanel or non-Linux hosts.
autoFreezer *autoFreezer
phpRelayShutdown []func() // ordered shutdown hooks
// watcherStatus tracks which top-level watchers have successfully attached.
// Keys are short stable names ("fanotify", "audit", "spool", "modsec",
// "afalg"). Values flip from false-to-true when the watcher's setup
// function completes without error. Used by /api/v1/status and the
// sd_notify gate.
//
// watcherChangedAt records the wall-clock time of the most recent state
// transition for the same key. Driven from MarkWatcher; consumed by the
// /api/v1/components endpoint so operators can see how long a watcher
// has been in its current state.
watcherMu sync.RWMutex
watcherStatus map[string]bool
watcherChangedAt map[string]time.Time
// watcherUpstream maps a watcher name to a probe function that reports
// whether the upstream feeding the watcher is still active. The probe
// runs at most once per /api/v1/components scrape; results compose with
// WatcherStatuses to surface "deaf" (attached but no upstream traffic)
// distinct from "idle" (attached and quiet) in the dashboard.
watcherUpstream map[string]UpstreamProbe
// findingBus fans out dispatched findings to passive observers like
// the SSE event stream. Initialized in Run(); closed on shutdown.
findingBus *broadcast.Bus
// updateChecker polls upstream for new CSM releases. Wired in Run()
// when updates.check_enabled is true (default). Nil when disabled
// or before Run starts; UpdateInfo() handles that.
updateChecker *updatecheck.Checker
// scanJobs is the full-scan job manager. Wired in Run() after the global
// bbolt store is open. Nil when the store is unavailable at startup.
scanJobs *ScanJobManager
// lastAutomationActionCache memoises the newest automation-emitted
// finding so /api/v1/status does not run a 100-row history cursor on
// every poll. Invalidated after lastAutomationActionTTL elapses.
automationActionMu sync.Mutex
automationActionCache *health.AutomationAction
automationActionCached time.Time
}
const lastAutomationActionTTL = 5 * time.Second
var logWatcherRetryInterval = 60 * time.Second
// New creates a new daemon instance.
func New(cfg *config.Config, store *state.Store, lock *state.LockFile, binaryPath string) *Daemon {
d := &Daemon{
cfg: cfg,
store: store,
lock: lock,
binaryPath: binaryPath,
alertCh: make(chan alert.Finding, 500),
stopCh: make(chan struct{}),
}
// Remediation records what it wrote here, so the sensitive-file detectors
// can tell CSM's own change from a third party's after a restart.
checks.SetSelfWriteStore(store)
if store != nil {
d.registerQueueSource("state", store)
}
d.smtpAuthTracker = newSMTPAuthTracker(
cfg.Thresholds.SMTPBruteForceThreshold,
cfg.Thresholds.SMTPBruteForceSubnetThresh,
cfg.Thresholds.SMTPAccountSprayThreshold,
time.Duration(cfg.Thresholds.SMTPBruteForceWindowMin)*time.Minute,
time.Duration(cfg.Thresholds.SMTPBruteForceSuppressMin)*time.Minute,
cfg.Thresholds.SMTPBruteForceSlowThreshold,
time.Duration(cfg.Thresholds.SMTPBruteForceSlowWindowMin)*time.Minute,
cfg.Thresholds.SMTPBruteForceMaxTracked,
time.Now,
)
d.smtpProbeTracker = newSMTPProbeTracker(
cfg.Thresholds.SMTPProbeThreshold,
time.Duration(cfg.Thresholds.SMTPProbeWindowMin)*time.Minute,
time.Duration(cfg.Thresholds.SMTPProbeSuppressMin)*time.Minute,
cfg.Thresholds.SMTPProbeMaxTracked,
time.Now,
smtpProbeBlockExpiryString,
)
d.mailAuthTracker = newMailAuthTracker(
cfg.Thresholds.MailBruteForceThreshold,
cfg.Thresholds.MailBruteForceSubnetThresh,
cfg.Thresholds.MailAccountSprayThreshold,
time.Duration(cfg.Thresholds.MailBruteForceWindowMin)*time.Minute,
time.Duration(cfg.Thresholds.MailBruteForceSuppressMin)*time.Minute,
cfg.Thresholds.MailBruteForceSlowThreshold,
time.Duration(cfg.Thresholds.MailBruteForceSlowWindowMin)*time.Minute,
cfg.Thresholds.MailBruteForceMaxTracked,
time.Now,
)
// Mail auth backend (cpdoveauthd) health: an active socket probe drives
// brute-force suppression for both trackers and optional self-heal restart.
// The probe goroutine only runs on cPanel (see Start); elsewhere Degraded()
// stays false, so the trackers behave exactly as before.
graceDur := 10 * time.Minute
if g, perr := time.ParseDuration(cfg.AutoResponse.MailAuthRecovery.DownGrace); perr == nil && g > 0 {
graceDur = g
}
restartCmd := cfg.AutoResponse.MailAuthRecovery.RestartCommand
d.authBackend = newAuthBackendHealth(
time.Now,
dialMailAuthBackend,
func() error { return restartMailAuthBackend(restartCmd) },
cfg.AutoResponse.MailAuthRecovery.RestartEnabled,
graceDur,
mailAuthRestartCooldown,
cfg.AutoResponse.MailAuthRecovery.MaxRestartsPerHour,
)
d.mailAuthTracker.SetBackendDownCheck(d.authBackend.Degraded)
d.smtpAuthTracker.SetBackendDownCheck(d.authBackend.Degraded)
return d
}
// reconcileBruteThresholds pushes the live config's SMTP/mail brute-force
// thresholds into the running trackers. The `thresholds` block is tagged
// hotreload:"safe", so a SIGHUP that only changes those fields reports success;
// without this push the trackers would keep their startup values until a full
// restart. Called from the reload success path, mirroring reconcileVerifiedBots.
func (d *Daemon) reconcileBruteThresholds() {
cfg := d.activeOrStartupCfg()
if cfg == nil {
return
}
th := cfg.Thresholds
if d.smtpAuthTracker != nil {
d.smtpAuthTracker.SetThresholds(
th.SMTPBruteForceThreshold,
th.SMTPBruteForceSubnetThresh,
th.SMTPAccountSprayThreshold,
time.Duration(th.SMTPBruteForceWindowMin)*time.Minute,
time.Duration(th.SMTPBruteForceSuppressMin)*time.Minute,
th.SMTPBruteForceSlowThreshold,
time.Duration(th.SMTPBruteForceSlowWindowMin)*time.Minute,
th.SMTPBruteForceMaxTracked,
)
}
if d.smtpProbeTracker != nil {
d.smtpProbeTracker.SetThresholds(
th.SMTPProbeThreshold,
time.Duration(th.SMTPProbeWindowMin)*time.Minute,
time.Duration(th.SMTPProbeSuppressMin)*time.Minute,
th.SMTPProbeMaxTracked,
)
}
if d.mailAuthTracker != nil {
d.mailAuthTracker.SetThresholds(
th.MailBruteForceThreshold,
th.MailBruteForceSubnetThresh,
th.MailAccountSprayThreshold,
time.Duration(th.MailBruteForceWindowMin)*time.Minute,
time.Duration(th.MailBruteForceSuppressMin)*time.Minute,
th.MailBruteForceSlowThreshold,
time.Duration(th.MailBruteForceSlowWindowMin)*time.Minute,
th.MailBruteForceMaxTracked,
)
}
}
// SetThresholds swaps the SMTP auth-brute detector thresholds under the
// tracker mutex so a SIGHUP reload of the safe `thresholds` block reaches this
// live tracker. Co-located with reconcileBruteThresholds, its only caller.
func (t *smtpAuthTracker) SetThresholds(perIP, subnet, accountSpray int, window, suppression time.Duration, slowThreshold int, slowWindow time.Duration, maxTracked int) {
t.mu.Lock()
defer t.mu.Unlock()
t.perIPThreshold = perIP
t.subnetThreshold = subnet
t.accountSprayThreshold = accountSpray
t.window = window
t.suppression = suppression
t.slowThreshold = slowThreshold
t.slowWindow = slowWindow
t.maxTracked = maxTracked
}
// SetThresholds swaps the mail auth-brute detector thresholds under the tracker
// mutex for a live SIGHUP reload.
func (t *mailAuthTracker) SetThresholds(perIP, subnet, accountSpray int, window, suppression time.Duration, slowThreshold int, slowWindow time.Duration, maxTracked int) {
t.mu.Lock()
defer t.mu.Unlock()
t.perIPThreshold = perIP
t.subnetThreshold = subnet
t.accountSprayThreshold = accountSpray
t.window = window
t.suppression = suppression
t.slowThreshold = slowThreshold
t.slowWindow = slowWindow
t.maxTracked = maxTracked
}
// SetThresholds swaps the SMTP connect-probe detector thresholds under the
// tracker mutex for a live SIGHUP reload.
func (t *smtpProbeTracker) SetThresholds(threshold int, window, suppression time.Duration, maxTracked int) {
t.mu.Lock()
defer t.mu.Unlock()
t.threshold = threshold
t.window = window
t.suppression = suppression
t.maxTracked = maxTracked
}
// SetVersion sets the application version for display in the web UI.
func (d *Daemon) SetVersion(v string) {
d.version = v
}
// MarkWatcher records the attachment state of a named watcher.
// Call from each watcher's startup path: true on success, false on failure.
// The first record and any subsequent state transition stamps
// watcherChangedAt so the components view can show "since".
func (d *Daemon) MarkWatcher(name string, attached bool) {
d.watcherMu.Lock()
defer d.watcherMu.Unlock()
if d.watcherStatus == nil {
d.watcherStatus = make(map[string]bool)
}
if d.watcherChangedAt == nil {
d.watcherChangedAt = make(map[string]time.Time)
}
prev, existed := d.watcherStatus[name]
if !existed || prev != attached {
d.watcherChangedAt[name] = time.Now()
}
d.watcherStatus[name] = attached
}
// WatcherStatuses returns a snapshot of every recorded watcher.
func (d *Daemon) WatcherStatuses() map[string]bool {
d.watcherMu.RLock()
defer d.watcherMu.RUnlock()
out := make(map[string]bool, len(d.watcherStatus))
for k, v := range d.watcherStatus {
out[k] = v
}
return out
}
// WatcherChangedAt returns the wall-clock time at which each watcher last
// transitioned state. Watchers without a recorded change return the zero
// value.
func (d *Daemon) WatcherChangedAt() map[string]time.Time {
d.watcherMu.RLock()
defer d.watcherMu.RUnlock()
out := make(map[string]time.Time, len(d.watcherChangedAt))
for k, v := range d.watcherChangedAt {
out[k] = v
}
return out
}
// UpstreamProbe reports whether a watcher's upstream input source is
// still feeding it. Return Fresh=false when the watcher is attached but
// can no longer hear from its source (PAM module not installed, log file
// rotated and never reappeared, fanotify marks lost, etc.) so the
// dashboard can surface "deaf" instead of conflating it with "idle".
type UpstreamProbe func() health.UpstreamResult
// RegisterUpstreamProbe wires a probe for a named watcher. Safe to call
// from any watcher's startup path; repeated calls overwrite the previous
// probe. Probes run on the request thread of /api/v1/components, so they
// must be cheap (single stat / atomic load, not a syscall storm).
func (d *Daemon) RegisterUpstreamProbe(name string, probe UpstreamProbe) {
d.watcherMu.Lock()
defer d.watcherMu.Unlock()
if d.watcherUpstream == nil {
d.watcherUpstream = make(map[string]UpstreamProbe)
}
d.watcherUpstream[name] = probe
}
// WatcherUpstream returns a snapshot of every probed watcher's upstream
// state. Watchers without a registered probe are absent from the map; the
// components API treats absence as "no probe wired, do not surface a
// deaf verdict for this watcher".
func (d *Daemon) WatcherUpstream() map[string]health.UpstreamResult {
d.watcherMu.RLock()
probes := make(map[string]UpstreamProbe, len(d.watcherUpstream))
for k, v := range d.watcherUpstream {
probes[k] = v
}
d.watcherMu.RUnlock()
out := make(map[string]health.UpstreamResult, len(probes))
for name, probe := range probes {
if probe == nil {
continue
}
out[name] = probe()
}
return out
}
// buildInfoOnce guards process-wide registration of the build_info
// gauge so repeated daemon construction in tests does not panic.
var buildInfoOnce sync.Once
// storeSizeOnce guards the csm_store_size_bytes gauge hook so tests
// that create multiple daemons share a single registration.
var storeSizeOnce sync.Once
// registerBuildInfo exposes build metadata on /metrics in the
// conventional Prometheus shape: a gauge fixed at 1, with the
// interesting fields as labels so scrapers can join on them.
func (d *Daemon) registerBuildInfo() {
buildInfoOnce.Do(func() {
g := metrics.NewGaugeVec(
"csm_build_info",
"CSM build metadata. Value is always 1; read version from the label.",
[]string{"version"},
)
version := d.version
if version == "" {
version = "unknown"
}
g.With(version).Set(1)
metrics.MustRegister("csm_build_info", g)
})
}
// registerStoreSizeMetric exposes the bbolt on-disk size as a gauge
// that stats the file at scrape time. No caching: the expected scrape
// interval is 15+ seconds and stat is cheap.
func (d *Daemon) registerStoreSizeMetric() {
storeSizeOnce.Do(func() {
metrics.RegisterGaugeFunc(
"csm_store_size_bytes",
"On-disk size of the bbolt state database in bytes.",
func() float64 {
db := store.Global()
if db == nil {
return 0
}
info, err := os.Stat(db.Path())
if err != nil {
return 0
}
return float64(info.Size())
},
)
})
}
// firewallMetricsOnce guards /metrics registration of the firewall +
// blocked-IP gauges so repeated daemon starts in a test binary are
// idempotent.
var firewallMetricsOnce sync.Once
var firewallMetricsMu sync.RWMutex
var firewallMetricsEngine *firewall.Engine
func setFirewallMetricsEngine(engine *firewall.Engine) {
firewallMetricsMu.Lock()
defer firewallMetricsMu.Unlock()
firewallMetricsEngine = engine
}
func firewallMetricsRuleCounts() firewall.RuleCounts {
firewallMetricsMu.RLock()
engine := firewallMetricsEngine
firewallMetricsMu.RUnlock()
if engine == nil {
return firewall.RuleCounts{}
}
return engine.RuleCounts()
}
func (d *Daemon) setFirewallEngine(engine *firewall.Engine) {
d.fwEngine = engine
d.fwActions = nil
if boundary, ok := any(engine).(firewallActionBoundary); ok {
d.fwActions = boundary
}
setFirewallMetricsEngine(engine)
}
// registerFirewallMetrics exposes the count of blocked IPs and the
// total number of firewall rules (IPs + allowed + subnets + port
// allow entries). Both gauges read the firewall engine at scrape time:
// the engine state file is authoritative, while the parallel bbolt
// fw:* buckets are written only at migration, so reading the store
// would freeze the gauge at the migration-time snapshot.
func (d *Daemon) registerFirewallMetrics() {
setFirewallMetricsEngine(d.fwEngine)
firewallMetricsOnce.Do(func() {
metrics.RegisterGaugeFunc(
"csm_blocked_ips_total",
"Number of IPs currently on the firewall block list (excluding expired temp bans).",
func() float64 {
return float64(firewallMetricsRuleCounts().Blocked)
},
)
metrics.RegisterGaugeFunc(
"csm_firewall_rules_total",
"Total firewall rules across all categories (blocked IPs, allowed IPs, blocked subnets, port-specific allows).",
func() float64 {
return float64(firewallMetricsRuleCounts().Total())
},
)
})
}
// Run starts the daemon and blocks until stopped.
func (d *Daemon) Run() error {
d.alertQueue = queuehealth.New(cap(d.alertCh), time.Minute)
defer alert.RegisterQueue(d.alertCh, d.alertQueue)()
if err := d.checkObserveStartupRecovery(); err != nil {
return err
}
d.startTime = time.Now()
defer alert.ClosePhpanelQueues()
if d.store != nil {
d.store.EnsureBaseline(d.startTime)
}
// Initialize structured logging from environment (CSM_LOG_FORMAT,
// CSM_LOG_LEVEL). The default text handler preserves the legacy
// "[YYYY-MM-DD HH:MM:SS] msg" format so operators mixing csmlog
// with legacy fmt.Fprintf call sites see a uniform log stream.
// CSM_LOG_FORMAT=json switches to structured JSON for log shipping.
csmlog.Init()
csmlog.Info("CSM daemon starting")
if err := alert.ConfigurePhpanelQueue(d.cfg); err != nil {
return fmt.Errorf("initializing phpanel webhook queue: %w", err)
}
// Expose Go runtime memory/scheduler stats on /metrics (heap growth is
// observable for leak triage), and start the loopback pprof listener when
// the operator opts in via debug.pprof_listen.
metrics.RegisterRuntimeMetrics()
if d.cfg != nil && d.cfg.Debug.PprofListen != "" {
d.startPprofListener(d.cfg.Debug.PprofListen)
}
// Periodic scans use this context so shutdown can abort an in-flight tier
// instead of blocking the worker drain for the full check budget.
scanCtx, scanCancel := context.WithCancel(context.Background())
d.scanCtx = scanCtx
d.scanCancel = scanCancel
defer scanCancel()
// Wire the active config to the incident-singleton's auto-close
// loop BEFORE the singleton is constructed, so the loop reads the
// operator-supplied thresholds on its first sweep. The closure
// captures the Daemon to pick up reloaded configs without restart.
SetIncidentConfigSource(func() *config.Config { return d.currentCfg() })
// Initialize the findings broadcast bus so passive observers (SSE, etc.)
// can subscribe before any findings are dispatched.
d.installFindingBus()
// Install config-supplied platform overrides BEFORE the first Detect()
// call so every check sees the merged view. The daemon command installs
// them even earlier (before the challenge-snippet refresh detects the
// platform); this call re-asserts them and reports if detection ran
// first without them.
InstallPlatformOverrides(d.cfg)
// Log detected platform as a structured record. In text mode this
// comes out as "[ts] platform detected os=X panel=Y ..."; in
// JSON mode as {"msg":"platform detected","os":"X","panel":"Y",...}.
pi := platform.Detect()
csmlog.Info("platform detected",
"os", orUnknown(string(pi.OS)),
"os_version", orUnknown(pi.OSVersion),
"panel", orNone(string(pi.Panel)),
"webserver", orNone(string(pi.WebServer)),
)
// Build the ModSec rule-action registry before any log watcher starts
// so the LiteSpeed classifier can tell pass-action vendor rules apart
// from real denies on the very first parsed line.
d.initModSecRegistry()
// Startup rollback can restore config and restart before reaching the
// watchers or integrity verification, so its sink must already be ready.
d.installActionLog()
// Wire the firewall tentative-apply manager. Recovery has to run
// before integrity.Verify because a pending rollback whose deadline
// passed while the daemon was down restores the previous csm.yaml
// (and its integrity hash) to disk; verifying first would fail
// against the still-on-disk new config the operator never confirmed.
if sdb := store.Global(); sdb != nil {
mgr := rollback.NewManager(sdb, d.cfg.ConfigFile, rollback.SystemctlRestart, time.Now)
rollback.SetGlobal(mgr)
recoveryCtx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
reverted, rerr := mgr.RecoverOnStartup(recoveryCtx)
cancel()
if rerr != nil {
return fmt.Errorf("firewall rollback recovery failed: %w", rerr)
}
if reverted {
csmlog.Warn("firewall rollback expired during downtime; previous config restored, restart issued")
// systemctl restart csm.service from the manager will tear
// us down momentarily; bail out cleanly so we do not race
// the watchers we are about to start against the restart.
return nil
}
}
// Verify integrity on startup
if err := integrity.Verify(d.binaryPath, d.cfg); err != nil {
tamper := alert.Finding{
Severity: alert.Critical,
Check: "integrity",
Message: fmt.Sprintf("BINARY/CONFIG TAMPER DETECTED: %v", err),
Timestamp: time.Now(),
}
_ = alert.Dispatch(d.cfg, []alert.Finding{tamper})
return fmt.Errorf("integrity check failed: %w", err)
}
// Publish the verified config as the process-wide live pointer.
// Hot paths (check ticks, alert dispatch, etc.) call
// config.Active() to pick up the current snapshot so a SIGHUP
// reload is visible on the next call without restart.
publishActiveConfig(d.cfg, "startup")
// Install the mail-brute account-key extractor selected by config.
// Validation in config.Load() already rejected invalid specs, so the
// error path here is defense-in-depth only.
if err := installAccountExtractorFromConfig(d.cfg); err != nil {
return err
}
d.applyStartupIntegrations()
// Initialize signature scanners and threat DB (fast, no I/O scan)
d.registerBuildInfo()
d.registerStoreSizeMetric()
d.registerFirewallMetrics()
checks.RegisterDirectSMTPEgressMetrics(metrics.Default())
RegisterBPFEnforcementMetrics(metrics.Default())
mailranges.RegisterMailrangesMetrics(metrics.Default())
if err := d.initYaraBackend(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] YARA backend init: %v\n", ts(), err)
}
// A realtime engine with zero rules is a silent outage; say so now that
// both engines have had their startup load.
d.reportRealtimeRuleCoverageNow()
if err := d.initPHPTaintAnalyzer(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] PHP taint worker init: %v\n", ts(), err)
} else {
defer d.stopPHPTaintAnalyzer()
}
checks.InitThreatDB(d.cfg.StatePath, d.cfg.Reputation.Whitelist)
if db := checks.GetThreatDB(); db != nil {
fmt.Fprintf(os.Stderr, "[%s] Threat DB initialized (%d entries)\n", ts(), db.Count())
}
if adb := attackdb.Init(d.cfg.StatePath); adb != nil {
d.prepareAttackDatabase(adb)
}
// Install operator verified-bot ranges before the firewall engine is
// exposed and before the initial auto-response pass. The firewall's
// soft-allow gate consults this registry from its first block decision.
startupBotEntries := verifiedBotEntries(d.cfg)
threatintel.SetOperatorBots(startupBotEntries)
// Load the mail-provider IP range cache synchronously before startFirewall()
// so engine.SetDOSExemptProviderNets receives a full provider set and the
// first Apply() builds dos_exempt_nets with current data. The background
// refresh goroutine is registered here as well.
d.initMailRanges()
// Start firewall engine if enabled
d.startFirewall()
// Settle any apply-confirmed window the restart interrupted. Must run
// after startFirewall (an expired-window restore has to land on top of
// the candidate ruleset the startup Apply just re-applied) and before
// startControlListener (so confirm/cancel cannot race the recovery).
d.recoverFirewallApplyConfirmed()
// Construct the incident correlator after the firewall starts so the
// credential_spray block hand-off is installed before the singleton
// captures its config. This still happens before any finding is
// dispatched.
_ = IncidentCorrelator()
// Start challenge server if enabled (gray listing)
d.startChallengeServer()
// Start challenge escalation ticker
if d.ipList != nil {
d.wg.Add(1)
obs.Go("challenge-escalator", d.challengeEscalator)
}
// Create password hijack detector
d.hijackDetector = NewPasswordHijackDetector(d.cfg, d.alertCh, d.stopCh)
// Start the alert dispatcher before any producer, held until the baseline
// scan below has published. Producers send non-blocking, so a dispatcher
// that started only after the baseline silently lost every realtime
// finding beyond the channel buffer during those minutes.
d.holdAlertDispatch()
d.wg.Add(1)
obs.Go("alert-dispatcher", d.alertDispatcher)
// Start inotify log watchers
d.startLogWatchers()
// Start PAM listener for real-time brute-force detection
d.startPAMListener()
// Wire the full-scan job manager. Must run after the global bbolt store is
// open so NewScanJobManager can resolve store.Global(). Failure here is
// non-fatal for the rest of the daemon: periodic scans and the webui
// continue; only explicit full-scan commands are unavailable.
if db := store.Global(); db != nil {
if sjm, err := d.startScanJobManager(); err != nil {
csmlog.Warn("scan job manager unavailable", "err", err)
} else {
d.scanJobs = sjm
}
}
// Start control socket listener for the thin-client CLI. The
// daemon is the sole bbolt owner; CLI commands that previously
// raced for the lock now route through this socket.
d.startControlListener()
// Start fanotify file monitor (real-time detection starts immediately)
d.startFileMonitor()
// Start email AV spool watcher (separate fanotify for Exim spool).
// Spool and forwarder watchers are cPanel-only; they watch paths
// (/var/spool/exim, /etc/valiases) that only exist on cPanel hosts.
if platform.Detect().IsCPanel() {
d.startSpoolWatcher()
d.startForwarderWatcher()
}
// Wire the update checker before the Web UI starts so the
// /api/v1/status handler always sees a non-nil checker (the
// goroutine that polls upstream still warms up for 5 minutes
// before the first poll, but UpdateInfo() returns zero values
// without racing on the field assignment).
d.startUpdateChecker()
// Start Web UI server - available immediately, before initial scan
d.startWebUI()
// Wire email quarantine to web server (after both start)
d.syncEmailAVWebState()
// Initialize GeoIP databases (after webServer so SetGeoIPDB can attach)
d.initGeoIP()
// Signal systemd we're up as soon as the real-time watchers and the
// control surfaces are attached. The initial baseline scan, the
// kernel-state probes, and the BPF tracker wiring below all run
// inline but no longer block `systemctl is-active` / `systemctl
// restart`. The watchdog notifier has to start in the same step so
// systemd's WatchdogSec doesn't trip while the baseline scan is
// still running on a large host.
d.wg.Add(1)
obs.Go("watchdog-notifier", d.watchdogNotifier)
d.wg.Add(1)
obs.Go("queue-health", d.monitorQueueHealth)
if sent, err := sdnotify.Ready(); err != nil {
fmt.Fprintf(os.Stderr, "sd_notify READY failed: %v\n", err)
} else if sent {
fmt.Fprintf(os.Stderr, "sd_notify: daemon ready\n")
}
_, _ = sdnotify.Status(fmt.Sprintf("watchers attached: %d", countAttachedWatchers(d.WatcherStatuses())))
// Reconcile the opt-in email forward-guard to the current config (installs
// the exim rule when enabled+enforcing, removes it otherwise), then keep its
// bad-IP lookup fresh. Both no-op off cPanel and fail open on error.
d.reconcileForwardGuard()
d.wg.Add(1)
obs.Go("forward-guard-refresh", d.forwardGuardRefresher)
// Snapshot the live config ONCE for the entire initial-scan tick
// (detection + auto-response). Earlier code called d.currentCfg()
// for RunTier and then again later for the auto-response batch;
// a SIGHUP landing between the two reads split the same tick
// across old detection policy and new response policy.
initialCfg := d.currentCfg()
// Build the block-digest collector before the baseline auto-response so
// startup blocks are captured too. GeoIP is already initialised above.
d.blockDigest = d.buildBlockDigest(initialCfg)
if d.blockDigest != nil {
bdInterval := initialCfg.BlockDigestInterval()
d.wg.Add(1)
obs.Go("block-digest", func() {
defer d.wg.Done()
t := time.NewTicker(bdInterval)
defer t.Stop()
d.blockDigest.Run(d.stopCh, t.C)
})
}
// Run the initial scan synchronously while alert dispatch is held.
fmt.Fprintf(os.Stderr, "[%s] Running initial baseline scan...\n", ts())
initialFindings, initialPurge := checks.RunTierWithContext(d.scanContext(), initialCfg, d.store, checks.TierCritical)
// Seed the attack database with initial scan findings
if adb := attackdb.Global(); adb != nil {
for _, f := range initialFindings {
adb.RecordFinding(f)
}
}
newFindings, permFixedKeys := d.respondToInitialScan(initialCfg, initialFindings)
// Remove auto-fixed findings before storing to UI
if len(permFixedKeys) > 0 {
fixedSet := make(map[string]bool, len(permFixedKeys))
for _, k := range permFixedKeys {
fixedSet[k] = true
}
var filtered []alert.Finding
for _, f := range initialFindings {
key := f.Check + ":" + f.Message
if !fixedSet[key] {
filtered = append(filtered, f)
}
}
initialFindings = filtered
}
d.store.Update(initialFindings)
d.store.MarkAlerted(newFindings)
// Merge initial scan findings into the existing set. Previous deep scan
// results (outdated_plugins, wp_core, etc.) persist across restarts until
// the next deep scan replaces them. ClearLatestFindings is NOT called
// here - it would wipe deep scan findings that haven't re-run yet.
checks.StoreLatestScanFindings(d.store, initialPurge, initialFindings)
csmlog.Info("initial scan complete", "findings", len(initialFindings), "new", len(newFindings))
// NOW let the dispatcher dispatch - no more race with the initial scan,
// and nothing queued while it ran was lost.
d.releaseAlertDispatch()
// Findings the previous shutdown drained but could not act on get their
// auto-response and alerts now, through the same pipeline.
d.replayPendingFindings()
// Retrospective cloud-relay scan: replay the last 24h of exim_mainlog
// through the compromise-detection rule so any in-progress credential
// abuse is surfaced within seconds of daemon start, not after the
// realtime watcher sees a new line. Gated on cPanel because
// exim_mainlog is cPanel-specific; safe to run in a goroutine so
// startup isn't delayed by parsing a large log.
if platform.Detect().IsCPanel() {
d.wg.Add(1)
obs.Go("cloud-relay-retro-scan", func() {
defer d.wg.Done()
cfg := d.currentCfg()
retro := ScanEximHistoryForCloudRelay(cfg, "", time.Now(), 24*time.Hour)
for i, f := range retro {
// Enqueue the finding FIRST; only after it is
// accepted by the dispatcher do we trigger the
// account-suspend side-effect. This prevents a
// silent mailbox suspension if the daemon begins
// shutting down between these two operations.
if !alert.Enqueue(d.alertCh, f, d.stopCh) {
alert.RecordQueueLoss(d.alertCh, uint64(len(retro[i+1:])))
return
}
sender := extractSenderFromCloudRelayMessage(f.Message)
if sender == "" {
continue
}
handleCloudRelayCredentialAbuse(cfg, sender)
}
})
}
// Start periodic scanners
d.wg.Add(1)
obs.Go("critical-scanner", d.criticalScanner)
d.wg.Add(1)
obs.Go("deep-scanner", d.deepScanner)
// Async bot-rDNS verifier: runs PTR+forward-A verification for
// claimed search-engine bot IPs that are not in a static range.
// Gated on reputation.bot_verify_enabled (default true) and on a
// non-nil store so the result can be persisted.
// Operator-configured verified bots extend the built-in allowlist.
// They were installed before firewall startup so auto-block soft-allow
// checks see them during the initial baseline pass; refresh from the
// current config here in case a SIGHUP landed during startup.
botEntries := verifiedBotEntries(d.currentCfg())
threatintel.SetOperatorBots(botEntries)
if d.cfg.BotVerifyEnabled() {
if db := store.Global(); db != nil {
ver := threatintel.OperatorBotsCacheVersion(threatintel.LogicVersion, botEntries)
if dropped, err := db.EnsureBotVerifyLogicVersion(ver); err != nil {
csmlog.Warn("bot-verify cache version check failed", "err", err)
} else if dropped {
csmlog.Info("bot-verify cache dropped after logic or verified_bots change")
}
d.startBotVerifier(db, botEntries)
}
}
// Auto-clear stale content and web-exposure findings when their re-check
// logic has changed since the last start.
// Runs in a goroutine so it does not block startup; each allowlisted family
// keeps its own fail-closed dismissal invariant.
if db := store.Global(); db != nil && d.store != nil {
token := checks.FindingReverifyVersion()
d.startContentReverifySweepIfChanged(db, token, func() ([]checks.ContentReverifyDismissal, checks.ReverifySweepStats, bool) {
return checks.ReverifyStaleFindingsStats(d.scanContext(), d.store)
})
}
// Live AF_ALG listener (Copy Fail / CVE-2026-31431) — only started
// when the kernel is actually exploitable. Hosts with a KernelCare
// livepatch covering CVE-2026-31431, OR built without the AF_ALG
// aead interface entirely, skip the listener: there's nothing to
// detect, and the inotify watch + 500ms tick would just burn cycles.
// The hardening audit + periodic critical-tier check stay active
// either way, so re-introduction of the vulnerability (e.g., a
// kernel rollback) is still surfaced via the slower path.
kstate := checks.ObserveAFAlgKernelState()
switch {
case !kstate.IsCopyFailExploitable():
csmlog.Info("af_alg live listener: skipped",
"reason", "kernel not exploitable",
"state", kstate.String(),
)
default:
if mon := StartAFAlgLiveMonitor(d.alertCh, d.cfg); mon == nil {
csmlog.Warn("af_alg live listener: not started",
"reason", "no backend available",
"state", kstate.String(),
)
d.MarkWatcher("afalg", false)
} else {
csmlog.Info("af_alg live listener: started",
"backend", mon.Mode(),
"state", kstate.String(),
)
d.registerBackendQueues("bpf.af_alg", mon)
d.wg.Add(1)
obs.Go("af-alg-listener", func() {
defer d.wg.Done()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { <-d.stopCh; cancel() }()
mon.Run(ctx)
})
d.MarkWatcher("afalg", true)
}
}
d.startPHPRelay()
if mon := StartConnectionTracker(d.alertCh, d.cfg); mon != nil {
d.registerBackendQueues("bpf.connection", mon)
csmlog.Info("connection_tracker: started", "backend", mon.Mode())
d.wg.Add(1)
obs.Go("connection-tracker", func() {
defer d.wg.Done()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { <-d.stopCh; cancel() }()
mon.Run(ctx)
})
}
if mon := StartExecMonitor(d.alertCh, d.cfg); mon != nil {
d.registerBackendQueues("bpf.execution", mon)
csmlog.Info("exec_monitor: started", "backend", mon.Mode())
d.wg.Add(1)
obs.Go("exec-monitor", func() {
defer d.wg.Done()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { <-d.stopCh; cancel() }()
mon.Run(ctx)
})
}
if mon := StartSensitiveFileMonitor(d.alertCh, d.cfg, d.store); mon != nil {
d.registerBackendQueues("bpf.sensitive_files", mon)
csmlog.Info("sensitive_files: started", "backend", mon.Mode())
d.wg.Add(1)
obs.Go("sensitive-files", func() {
defer d.wg.Done()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { <-d.stopCh; cancel() }()
mon.Run(ctx)
})
}
// Start automatic signature updates
d.wg.Add(1)
obs.Go("signature-updater", d.signatureUpdater)
// Start signature watcher: arms forceFullRescan when any rule
// file's content changes. Disabled wholesale via
// detection.rescan_on_signature_update: false.
d.wg.Add(1)
obs.Go("signature-watcher", d.signatureWatcher)
d.wg.Add(1)
obs.Go("geoip-updater", d.geoipUpdater)
d.wg.Add(1)
obs.Go("bot-ranges-updater", d.botRangesUpdater)
// Start heartbeat
d.wg.Add(1)
obs.Go("heartbeat", d.heartbeat)
// One-shot: stagger CSM-managed WP-Cron crontab lines left behind by
// older releases. The perf_wp_cron finding never re-fires once a site
// is fixed, so this is the only path that reaches them.
d.wg.Add(1)
obs.Go("wpcron-migrate", func() {
defer d.wg.Done()
if n := checks.MigrateWPCronCrontabs(d.cfg); n > 0 {
csmlog.Info("wp-cron: staggered legacy system cron entries", "upgraded", n)
}
})
// Start abuse reporting (opt-in). startAbuseReporting installs the alert
// hook and returns the spool drain loop, or nil when disabled. The reporter
// stops after the final shutdown alert flush so late findings can still be
// queued before the bbolt spool closes.
if reportLoop := d.startAbuseReporting(); reportLoop != nil {
obs.Go("abuse-reporter", reportLoop)
}
// Start the central scored-set consumer (opt-in). Maintains the verified
// set and escalates findings whose IP is listed; nil when disabled.
if centralLoop := d.startCentralConsume(); centralLoop != nil {
d.wg.Add(1)
obs.Go("central-intel", func() {
defer d.wg.Done()
centralLoop()
})
}
// Start the retention sweep only when opted in. Compaction is not
// run from this goroutine; see internal/daemon/retention.go for why.
if d.cfg != nil && d.cfg.Retention.Enabled {
d.wg.Add(1)
obs.Go("retention-scanner", d.retentionScanner)
}
// Refresh the sd_notify status line now that AF_ALG, BPF trackers,
// and the periodic scanners have all reported in. The initial
// READY=1 fired earlier so systemctl restart didn't have to wait
// on the baseline scan; this is the operator-visible summary.
_, _ = sdnotify.Status(fmt.Sprintf("watchers attached: %d", countAttachedWatchers(d.WatcherStatuses())))
csmlog.Info("CSM daemon running")
// Wait for signals
sigCh := make(chan os.Signal, 1)
signal.Notify(sigCh, syscall.SIGTERM, syscall.SIGINT, syscall.SIGHUP)
for sig := range sigCh {
if sig == syscall.SIGHUP {
fmt.Fprintf(os.Stderr, "[%s] SIGHUP received - reloading config and rules\n", ts())
d.reloadConfig()
d.reloadSignatures()
if d.policies != nil {
if err := d.policies.Reload(d.currentCfg().EmailProtection.PHPRelay.PoliciesDir); err != nil {
// Previous valid version stays in effect; surface the failure
// via the existing alert pipeline so operators see partial reload.
d.emitReloadFinding(alert.Warning, "email_php_relay_policies_reload",
fmt.Sprintf("policies/email reload encountered errors: %v", err))
}
}
d.publishGeoIP()
if d.fwEngine != nil {
if err := d.fwEngine.Apply(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Firewall reload error: %v\n", ts(), err)
} else {
fmt.Fprintf(os.Stderr, "[%s] Firewall rules reloaded\n", ts())
}
}
continue
}
break // SIGTERM or SIGINT
}
csmlog.Info("shutting down")
shutdownStart := time.Now()
close(d.stopCh)
// Abort any in-flight periodic scan so d.wg.Wait below is not held for a
// whole tier. Scanners observe d.stopCh between cycles; this cancels the
// scan currently executing inside RunTier.
if d.scanCancel != nil {
d.scanCancel()
}
// Log watchers own their files and close them when their Run loop exits on
// d.stopCh; d.wg.Wait below blocks until that happens. Calling Stop here
// would race the still-running Run on w.file.
if d.webServer != nil {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
_ = d.webServer.Shutdown(ctx)
cancel()
}
if d.challengeServer != nil {
d.challengeServer.Shutdown()
}
if d.challengeGate != nil {
if err := d.challengeGate.Close(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] challenge port-gate close: %v\n", ts(), err)
}
d.challengeGate = nil
}
if fm := d.getFileMonitor(); fm != nil {
fm.Stop()
}
if sw := d.getSpoolWatcher(); sw != nil {
sw.Stop()
}
if d.pamListener != nil {
d.pamListener.Stop()
}
if d.controlListener != nil {
d.controlListener.Stop()
}
if d.scanJobs != nil {
d.scanJobs.Stop()
}
d.stopYaraBackend()
d.stopPHPTaintAnalyzer()
csmlog.Info("watchers signalled", "elapsed_ms", time.Since(shutdownStart).Milliseconds())
d.wg.Wait()
stopProcessCtx()
csmlog.Info("workers drained", "elapsed_ms", time.Since(shutdownStart).Milliseconds())
// Some producers can finish a tick after alertDispatcher observes stopCh.
// Drain again once tracked workers are gone and before state is closed.
d.flushPendingAlertsOnShutdown()
alert.ClosePhpanelQueues()
d.stopAbuseReporting()
d.closeFindingBus()
for i := len(d.phpRelayShutdown) - 1; i >= 0; i-- {
d.phpRelayShutdown[i]()
}
d.phpRelayShutdown = nil
if adb := attackdb.Global(); adb != nil {
adb.Stop()
}
// Stop the incident auto-close and retention goroutines before closing the
// store so neither writes to an already-closed bbolt database.
StopIncidentBackgroundLoops()
if err := d.store.Close(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] error closing state store: %v\n", ts(), err)
}
d.lock.Release()
csmlog.Info("daemon stopped", "elapsed_ms", time.Since(shutdownStart).Milliseconds())
fmt.Fprintf(os.Stderr, "[%s] CSM daemon stopped\n", ts())
return nil
}
type contentLogicVersionStore interface {
ContentLogicVersionChanged(token string) (bool, error)
SetContentLogicVersion(token string) error
}
func (d *Daemon) startContentReverifySweepIfChanged(db contentLogicVersionStore, token string, run func() ([]checks.ContentReverifyDismissal, checks.ReverifySweepStats, bool)) {
changed, err := db.ContentLogicVersionChanged(token)
if err != nil {
csmlog.Warn("finding re-verification version check failed", "err", err)
return
}
if changed {
d.startContentReverifySweep(func() ([]checks.ContentReverifyDismissal, checks.ReverifySweepStats, bool) {
dismissed, stats, complete := run()
if !complete {
return dismissed, stats, false
}
if err := db.SetContentLogicVersion(token); err != nil {
csmlog.Warn("finding re-verification version update failed", "err", err)
return dismissed, stats, false
}
return dismissed, stats, true
})
}
}
func (d *Daemon) startContentReverifySweep(run func() ([]checks.ContentReverifyDismissal, checks.ReverifySweepStats, bool)) {
d.wg.Add(1)
obs.Go("content-reverify-sweep", func() {
defer d.wg.Done()
outcomes, stats, complete := run()
for _, dm := range outcomes {
csmlog.Info(contentReverifyOutcomeMessage(dm),
"check", dm.Check, "path", dm.Path, "detail", dm.Detail)
}
if !complete {
csmlog.Info("finding re-verification sweep will retry on next start")
return
}
// Always log the summary. A sweep that produced nothing used to log
// nothing at all, which made "ran and found nothing" indistinguishable
// from "never ran" and from "failed on every finding" -- and that
// ambiguity cost more than one wrong conclusion about why findings
// were not draining.
csmlog.Info("finding re-verification sweep complete",
"considered", stats.Considered, "cleared", stats.Cleared, "demoted", stats.Demoted,
"promoted", stats.Promoted, "unchecked", stats.Unchecked,
"unchecked_reason", stats.TopUncheckedReason)
})
}
func contentReverifyOutcomeMessage(outcome checks.ContentReverifyDismissal) string {
switch {
case outcome.Promoted:
return "finding severity restored after re-verification"
case outcome.Demoted:
return "remediated finding demoted"
default:
return "stale finding auto-cleared"
}
}
// DroppedAlerts returns the total number of alerts dropped due to
// channel backpressure since the daemon started.
func (d *Daemon) DroppedAlerts() int64 {
if d.alertQueue != nil {
return int64(d.alertQueue.Snapshot(time.Now()).DroppedTotal) // #nosec G115 -- a daemon cannot produce MaxInt64 findings within its lifetime.
}
return atomic.LoadInt64(&d.droppedAlerts)
}
func (d *Daemon) scanContext() context.Context {
if d.scanCtx != nil {
return checks.WithModSecReload(d.scanCtx, &d.modsecReload)
}
return checks.WithModSecReload(context.Background(), &d.modsecReload)
}
// FindingBus returns the per-daemon broadcast.Bus used by passive
// observers like the SSE event stream. Returns nil before Run starts.
func (d *Daemon) FindingBus() *broadcast.Bus {
return d.findingBus
}
// alertDispatcher batches and dispatches alerts.
// alertBatchInterval is how long the dispatcher collects findings before
// dispatching them as one batch. A variable so tests can shorten it.
var alertBatchInterval = 5 * time.Second
// alertHoldMaxBatch bounds the batch the dispatcher accumulates while held.
// Beyond it, further findings are dropped and counted rather than growing
// memory without limit on a host whose baseline runs for minutes.
const alertHoldMaxBatch = 5000
// holdAlertDispatch arms the startup hold. Must run before the dispatcher
// goroutine starts.
func (d *Daemon) holdAlertDispatch() {
d.alertHold = make(chan struct{})
if d.alertQueue != nil {
d.alertQueue.Hold(time.Now())
}
}
// releaseAlertDispatch lets the dispatcher start dispatching its batches.
// Idempotent; a no-op when no hold was armed.
func (d *Daemon) releaseAlertDispatch() {
if d.alertHold == nil {
return
}
d.alertReleaseOnce.Do(func() {
if d.alertQueue != nil {
d.alertQueue.Release(time.Now())
}
close(d.alertHold)
})
}
func (d *Daemon) alertDispatcher() {
defer d.wg.Done()
ticker := time.NewTicker(alertBatchInterval)
defer ticker.Stop()
var batch []alert.Finding
held := d.alertHold // nil: never selected, so the dispatcher runs freely
heldDropped := 0
for {
select {
case <-d.stopCh:
batch = d.drainAlertChannel(batch)
d.persistPendingFindingsOnShutdown(batch)
return
case f := <-d.alertCh:
alert.StartQueued(f)
if held != nil && len(batch) >= alertHoldMaxBatch {
alert.RejectQueued(f)
heldDropped++
atomic.AddInt64(&d.droppedAlerts, 1)
continue
}
batch = append(batch, f)
case <-held:
held = nil
if heldDropped > 0 {
fmt.Fprintf(os.Stderr, "[%s] Alert dispatcher: %d realtime findings dropped while dispatch was held for the baseline scan\n", ts(), heldDropped)
}
case <-ticker.C:
if held != nil {
continue
}
if len(batch) > 0 {
d.dispatchBatch(batch)
batch = nil
}
}
}
}
func (d *Daemon) drainAlertChannel(batch []alert.Finding) []alert.Finding {
for {
select {
case f, ok := <-d.alertCh:
if !ok {
return batch
}
alert.StartQueued(f)
batch = append(batch, f)
default:
return batch
}
}
}
func (d *Daemon) flushPendingAlertsOnShutdown() {
d.persistPendingFindingsOnShutdown(d.drainAlertChannel(nil))
}
// persistPendingFindingsOnShutdown records findings still queued at shutdown to
// the history log for forensics and parks them for the next start. It
// deliberately does NOT run the auto-response pipeline (nftables blocks,
// permission fixes, kill/quarantine, DB cleanup) or network alert dispatch
// here: that work blocked the service stop for tens of seconds -- up to twice,
// once here and once in the dispatcher's stop branch -- while systemd waited.
// The baseline scan re-detects scheduled checks after a restart, but a
// realtime-only finding (a webshell written in the last seconds) is seen by
// nothing else, so the parked batch is replayed through dispatchBatch once
// the next start releases the dispatcher. Nothing here is marked sent via
// store.Update, so the replay's dispatch is not suppressed.
func (d *Daemon) persistPendingFindingsOnShutdown(batch []alert.Finding) {
defer alert.FinishQueued(batch)
if len(batch) == 0 {
return
}
alert.FillTimestamps(batch, time.Now())
d.store.AppendHistory(batch)
if err := d.store.AppendPendingFindings(batch); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Could not park %d pending finding(s) for the next start: %v\n", ts(), len(batch), err)
}
}
func isOperatorAlertableCheck(check string) bool {
switch config.CanonicalCheckName(check) {
case "modsec_block_realtime", "modsec_warning_realtime", "modsec_block_escalation", "modsec_csm_block_escalation":
return false // Fully automated and visible on the ModSecurity page.
case "outdated_plugins":
return false // Routine update posture remains on the findings page.
case "email_dkim_failure", "email_spf_rejection":
return false // Operational email authentication issues are informational.
case "email_auth_failure_realtime", "pam_bruteforce", "exim_frozen_realtime":
return false // Failed logins and frozen bounces need no operator action.
case "ftp_login", "cpanel_file_upload_realtime":
// On shared hosting every customer connects from a non-infra address,
// so a successful FTP login or a File Manager write is ordinary use of
// a core feature. Their value is correlation with other evidence on
// the same account, which the findings page, history, incidents and
// the attack database all still get. The brute-force escalations
// (ftp_auth_failure_realtime, ftp_bruteforce,
// ftp_login_after_bruteforce) stay alertable.
return false
default:
return true
}
}
func operatorAlertableFindings(findings []alert.Finding) []alert.Finding {
alertable := make([]alert.Finding, 0, len(findings))
for _, f := range findings {
if isOperatorAlertableCheck(f.Check) {
alertable = append(alertable, f)
}
}
return alertable
}
func (d *Daemon) dispatchBatch(findings []alert.Finding) {
defer alert.FinishQueued(findings)
// Realtime producers may hand over findings without a Timestamp; stamp
// them once here so history, incidents, the latest set and every alert
// sink see the same time.
alert.FillTimestamps(findings, time.Now())
auditSources := append([]alert.Finding(nil), findings...)
// Snapshot the live config once at the top of the batch. Every
// cfg.X read below picks up the last-reloaded value (ROADMAP
// item 7); taking one snapshot avoids the weirder case of a
// SIGHUP landing mid-batch and splitting some auto-response
// actions between old and new policy.
cfg := d.currentCfg()
findings = alert.Deduplicate(findings)
suppressions := d.store.LoadSuppressions()
remediableFindings := filterUnsuppressedFindings(d.store, findings, suppressions)
// Record ALL findings in attack database (before filtering -
// repeated attacks from the same IP must still be counted even if
// the alert is suppressed by FilterNew or a suppression rule).
if adb := attackdb.Global(); adb != nil {
for _, f := range findings {
adb.RecordFinding(f)
}
}
// Auto-block and permission fix run on ALL findings (not just new ones).
// These must execute BEFORE FilterNew because repeat offender IPs and
// recurring permission issues need to be fixed even if the alert was
// already sent in a previous cycle.
// Challenge routing runs FIRST - claims eligible IPs before hard-blocking.
// One ordered helper guarantees that ordering on every auto-response path.
// Suppression rules do not gate IP responses: they mute a check, and a
// check-wide rule would otherwise leave every attacker it reports
// unblocked. An IP false positive belongs on the allowlist.
challengeActions, blockActions := checks.ChallengeThenBlock(cfg, findings)
permActions, permFixedKeys := checks.AutoFixPermissions(cfg, remediableFindings)
// Mark auto-blocked IPs in attack database
if adb := attackdb.Global(); adb != nil {
for _, f := range blockActions {
if ip := checks.ExtractIPFromFinding(f); ip != "" {
adb.MarkBlocked(ip)
}
}
}
// Dismiss auto-fixed findings from the Findings page
for _, key := range permFixedKeys {
d.store.DismissLatestFinding(key)
}
// Filter through state - only new findings get alerted and logged
unfilteredNew := d.store.FilterNew(findings)
// responseFindings is every new observation and action, suppressed or not:
// incidents and central enforcement act on it. newFindings drops what
// suppression rules match, which mutes notifications and remediation.
// Suppressions are stored in state/suppressions.json, not in rule files.
responseFindings := append([]alert.Finding(nil), unfilteredNew...)
responseFindings = append(responseFindings, blockActions...)
responseFindings = append(responseFindings, challengeActions...)
responseFindings = append(responseFindings, permActions...)
// PHP-relay AutoFreeze: emit any new findings produced by post-emit
// freeze decisions back into the dispatched batch so operators see
// the action outcome alongside the original finding. Nil-guard for
// non-cPanel / non-linux hosts where wiring is skipped.
if d.autoFreezer != nil {
if freezeFindings := d.autoFreezer.Apply(remediableFindings); len(freezeFindings) > 0 {
responseFindings = append(responseFindings, freezeFindings...)
}
}
if len(responseFindings) == 0 {
_ = alert.DispatchWithSources(cfg, nil, auditSources)
d.store.Update(findings)
return
}
// Copy: with no rules the filter returns its input, and both slices grow.
newFindings := append([]alert.Finding(nil), filterUnsuppressedFindings(d.store, responseFindings, suppressions)...)
d.store.AppendHistory(newFindings)
d.observeBlocks(blockActions)
// Kill and quarantine only run on new, unsuppressed findings.
killActions := checks.AutoKillProcesses(d.scanContext(), cfg, newFindings)
quarantineActions := checks.AutoQuarantineFiles(cfg, newFindings)
// Database response also discovers attacker session IPs. Suppressions
// stop SQL writes and session revocation, but not those IP blocks.
dbActions := autoRespondDBMalware(cfg, unfilteredNew, func(f alert.Finding) bool {
return !d.store.IsSuppressed(f, suppressions)
})
for _, actions := range [][]alert.Finding{killActions, quarantineActions, dbActions} {
responseFindings = append(responseFindings, actions...)
newFindings = append(newFindings, filterUnsuppressedFindings(d.store, actions, suppressions)...)
}
// Correlation derives notifications, so it reads only unsuppressed
// findings; its derived findings still reach incidents and enforcement.
uncorrelated := len(newFindings)
newFindings = expandWithCorrelation(newFindings, time.Now())
responseFindings = append(responseFindings, newFindings[uncorrelated:]...)
co := IncidentCorrelator()
for _, f := range alert.Deduplicate(responseFindings) {
_, _, _ = co.OnFinding(f)
}
// Derived findings may themselves be suppressed. Keep them in the
// response set while applying their rules before notification fanout.
newFindings = filterUnsuppressedFindings(d.store, newFindings, suppressions)
// Apply notification policy after the phpanel stream and passive
// observers receive the findings, so muting email does not lose evidence.
auditSources = append(auditSources, responseFindings...)
if err := alert.DispatchWithNotificationFilter(cfg, newFindings, auditSources, responseFindings, operatorAlertableFindings); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Alert dispatch error: %v\n", ts(), err)
}
d.store.Update(findings)
d.store.MarkAlerted(newFindings)
}
// respondToInitialScan records the baseline scan, runs its auto-response and
// dispatches the resulting alerts. It returns the alerted findings and the keys
// of findings the permission auto-fix repaired.
func (d *Daemon) respondToInitialScan(cfg *config.Config, initialFindings []alert.Finding) ([]alert.Finding, []string) {
d.store.AppendHistory(initialFindings)
unfilteredNew := d.store.FilterNew(initialFindings)
suppressions := d.store.LoadSuppressions()
// Copy: with no rules the filter returns its input, and newFindings grows.
newFindings := append([]alert.Finding(nil), filterUnsuppressedFindings(d.store, unfilteredNew, suppressions)...)
// Permission auto-fix runs on ALL findings (not just new) because
// it's safe/idempotent and should fix baseline findings too.
permActions, permFixedKeys := checks.AutoFixPermissions(cfg, filterUnsuppressedFindings(d.store, initialFindings, suppressions))
// Challenge routing runs on ALL findings unconditionally when enabled, so an
// eligible IP is on the challenge list before AutoBlockIPs (below, guarded by
// new findings) checks membership. Not folded into ChallengeThenBlock here:
// challenge must route even with no new findings (re-establishing challenges
// on restart) while the block stage stays gated on new findings. As in
// dispatchBatch, suppression rules do not gate IP responses.
challengeActions := checks.ChallengeRouteIPs(cfg, initialFindings)
// Other auto-response only on new findings. As in dispatchBatch,
// responseFindings keeps suppressed observations for incidents and
// central enforcement while newFindings carries what may notify.
var responseFindings []alert.Finding
if len(unfilteredNew) > 0 {
killActions := checks.AutoKillProcesses(d.scanContext(), cfg, newFindings)
quarantineActions := checks.AutoQuarantineFiles(cfg, newFindings)
blockActions := checks.AutoBlockIPs(cfg, initialFindings)
d.observeBlocks(blockActions)
responseFindings = append(responseFindings, unfilteredNew...)
for _, actions := range [][]alert.Finding{killActions, quarantineActions, permActions, challengeActions, blockActions} {
responseFindings = append(responseFindings, actions...)
newFindings = append(newFindings, filterUnsuppressedFindings(d.store, actions, suppressions)...)
}
// Cross-account correlation runs on the initial batch too, not
// just on subsequent ticks. Otherwise three account compromises
// landing in the first scan slip past with no synthetic alert.
uncorrelated := len(newFindings)
newFindings = expandWithCorrelation(newFindings, time.Now())
responseFindings = append(responseFindings, newFindings[uncorrelated:]...)
co := IncidentCorrelator()
for _, f := range alert.Deduplicate(responseFindings) {
_, _, _ = co.OnFinding(f)
}
}
newFindings = filterUnsuppressedFindings(d.store, newFindings, suppressions)
initialAuditSources := append(append([]alert.Finding(nil), initialFindings...), responseFindings...)
_ = alert.DispatchWithNotificationFilter(cfg, newFindings, initialAuditSources, responseFindings, operatorAlertableFindings)
return newFindings, permFixedKeys
}
// autoFixWPCron lets daemon wiring tests avoid real wp-config.php and crontab
// edits; the checks package covers those side effects directly.
var autoFixWPCron = checks.AutoFixWPCron
// Database wiring tests exercise suppression policy without a live MySQL host.
var autoRespondDBMalware = checks.AutoRespondDBMalwareWithPolicy
// processScanFindings handles the output of a deep or periodic scan: it persists
// the findings to the latest-findings surface, runs the auto-responses that act
// on warning-severity perf findings, then forwards the remaining findings to the
// alert dispatcher. Warning-severity perf findings stay off the alert channel so
// they never page an operator; that is exactly why the WP-Cron auto-fix runs
// here and not in dispatchBatch, which only ever sees what the channel carries.
func (d *Daemon) processScanFindings(cfg *config.Config, findings []alert.Finding, purgeChecks []string, label string) {
d.processScanFindingsWithCoverage(cfg, findings, purgeChecks, nil, label)
}
func (d *Daemon) processScanFindingsWithCoverage(cfg *config.Config, findings []alert.Finding, purgeChecks []string, coverage *state.ScanCoverage, label string) {
checks.StoreLatestScanFindingsWithCoverage(d.store, purgeChecks, findings, coverage)
d.applyWPCronAutoFix(cfg, findings)
d.enqueueScanAlerts(findings, label)
}
// scanAlertEnqueueTimeout bounds how long a scan enqueue can make no progress
// while waiting for room on the shared alert channel. The undelivered tail is
// dropped as an absolute last resort when that timeout expires. The
// dispatcher stops draining only while it runs a batch, and auto-response can
// hold that for tens of seconds; a scan burst that fills the buffer in that
// window must apply backpressure rather than silently drop security findings
// the way an earlier non-blocking send did -- a false-positive flood once
// dropped a real rogue-admin finding. Set well above the dispatcher's
// worst-case batch time so a drop here means the dispatcher is genuinely wedged.
const scanAlertEnqueueTimeout = 90 * time.Second
// enqueueScanAlerts forwards a scan's findings to the alert dispatcher.
// Warning-severity perf findings stay off the channel so they never page an
// operator; that is also why the WP-Cron auto-fix runs against the scan
// findings directly rather than in dispatchBatch, which only sees the channel.
//
// Unlike the real-time kernel monitors, which must never block a BPF or
// fanotify handler, the scan path is not latency-sensitive, so it applies
// backpressure when the buffer is full instead of dropping. A finding burst can
// otherwise fill the buffer while the dispatcher is mid-batch and evict
// findings from this or any other producer sharing the channel.
func (d *Daemon) enqueueScanAlerts(findings []alert.Finding, label string) {
d.enqueueScanAlertsWithin(findings, label, scanAlertEnqueueTimeout)
}
func (d *Daemon) enqueueScanAlertsWithin(findings []alert.Finding, label string, timeout time.Duration) {
for i, f := range findings {
if !scanFindingIsAlertable(f) {
continue
}
err := alert.EnqueueWithin(d.alertCh, f, d.stopCh, timeout)
if err == nil {
continue
}
remaining := countAlertableScanFindings(findings[i+1:])
// EnqueueWithin already counted the rejected send, including shutdown.
alert.RecordQueueLoss(d.alertCh, uint64(remaining)) // #nosec G115 -- countAlertableScanFindings returns a nonnegative count bounded by the slice length.
dropped := remaining + 1
atomic.AddInt64(&d.droppedAlerts, int64(dropped))
if errors.Is(err, alert.ErrQueueTimeout) {
fmt.Fprintf(os.Stderr, "[%s] alert channel jammed for %s, dropping %d remaining %s findings (first: %s)\n", ts(), timeout, dropped, label, f.Check)
}
return
}
}
func countAlertableScanFindings(findings []alert.Finding) int {
count := 0
for _, f := range findings {
if scanFindingIsAlertable(f) {
count++
}
}
return count
}
func scanFindingIsAlertable(f alert.Finding) bool {
return !strings.HasPrefix(f.Check, "perf_") || f.Severity != alert.Warning
}
// applyWPCronAutoFix disables WP-Cron and installs a per-user system cron for
// every perf_wp_cron finding, then clears the fixed findings from the
// latest-findings surface and records the actions in history. Findings the
// operator has suppressed are left untouched, so a suppression also stops the
// automated edit of that account's wp-config.php.
func (d *Daemon) applyWPCronAutoFix(cfg *config.Config, findings []alert.Finding) {
if d.store != nil {
if suppressions := d.store.LoadSuppressions(); len(suppressions) > 0 {
findings = filterUnsuppressedFindings(d.store, findings, suppressions)
}
}
actions, fixedKeys := autoFixWPCron(cfg, findings)
for _, key := range fixedKeys {
d.store.DismissLatestFinding(key)
}
if len(actions) > 0 {
d.store.AppendHistory(actions)
}
}
func filterUnsuppressedFindings(store *state.Store, findings []alert.Finding, suppressions []state.SuppressionRule) []alert.Finding {
if len(suppressions) == 0 {
return findings
}
var filtered []alert.Finding
for _, f := range findings {
if !store.IsSuppressed(f, suppressions) {
filtered = append(filtered, f)
}
}
return filtered
}
// criticalScanner runs critical checks every 10 minutes.
func (d *Daemon) criticalScanner() {
defer d.wg.Done()
ticker := time.NewTicker(10 * time.Minute)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
d.runPeriodicChecks(checks.TierCritical)
}
}
}
// deepScanner runs deep checks at the configured interval (default 60 min).
// If fanotify is active, runs only the checks it can't replace (reduced set).
// If fanotify is NOT active (fallback mode), runs the full deep tier for timer-mode parity.
func (d *Daemon) deepScanner() {
defer d.wg.Done()
// Re-read the interval on each iteration so a SIGHUP that changes
// thresholds.deep_scan_interval_min takes effect on the next scan.
// A ticker captured at startup can't be re-sized cleanly without
// a reset path; time.After recomputes.
for {
interval := time.Duration(d.currentCfg().Thresholds.DeepScanIntervalMin) * time.Minute
if interval <= 0 {
// Defensive: an operator who zeroes the threshold would
// otherwise get a tight spin loop. 60 minutes matches the
// default from config.Load.
interval = 60 * time.Minute
}
select {
case <-d.stopCh:
return
case <-time.After(interval):
// Re-verify findings whose condition someone else resolved. The
// startup sweep is gated on the re-check logic version, which only
// moves on deploy: an operator cleaning a file, or a virtual patch
// closing an exposure, changes the world without changing CSM, and
// a finding gated only on that would keep its severity until the
// next upgrade happened to land.
if d.store != nil {
d.startContentReverifySweep(func() ([]checks.ContentReverifyDismissal, checks.ReverifySweepStats, bool) {
return checks.ReverifyStaleFindingsStats(d.scanContext(), d.store)
})
}
// Update threat intelligence feeds (once per day)
if db := checks.GetThreatDB(); db != nil {
_ = db.UpdateFeeds()
}
// Prune expired attack DB records (90-day retention)
if adb := attackdb.Global(); adb != nil {
adb.PruneExpired()
}
// If fanotify is active, only run checks it can't replace.
// If fanotify is NOT active, run the full deep tier.
//
// One exception: forceFullRescan is armed by the
// signature watcher when any rule file's content changes.
// In that case we bypass the fanotify short-list so the
// new ruleset gets a full sweep against existing files;
// without this, only files that change AFTER the rule
// update would catch the new patterns.
cfg := d.currentCfg()
rescan := d.forceFullRescan.CompareAndSwap(true, false)
scanCtx, gaps := checks.WithCoverageGaps(d.scanContext())
var findings []alert.Finding
var purgeChecks []string
switch {
case rescan:
findings, purgeChecks = checks.RunTierWithContext(scanCtx, cfg, d.store, checks.TierDeep)
observeSignatureRescan()
case d.getFileMonitor() != nil:
findings, purgeChecks = checks.RunReducedDeepWithContext(scanCtx, cfg, d.store)
default:
findings, purgeChecks = checks.RunTierWithContext(scanCtx, cfg, d.store, checks.TierDeep)
}
d.processScanFindingsWithCoverage(cfg, findings, purgeChecks, gaps.Snapshot(), "deep")
}
}
}
func (d *Daemon) runPeriodicChecks(tier checks.Tier) {
// Snapshot the live config ONCE for the whole tick. Calling
// d.currentCfg() twice (once for integrity, once for RunTier)
// lets a SIGHUP land between the two reads and split the tick
// between old-policy integrity verification and new-policy
// detection. Matches the snapshot pattern in dispatchBatch.
cfg := d.currentCfg()
// Verify integrity against the snapshot. A SIGHUP reload re-signs
// integrity.config_hash on disk and updates config.Active; using
// d.cfg (the startup snapshot) here would fire a Critical tamper
// alert on every tick after a successful reload because the stored
// hash in d.cfg is stale. If a reload completes while Verify is
// hashing, retry once against the latest live config to avoid a
// false tamper alert from a stale snapshot.
var err error
cfg, err = d.verifyPeriodicIntegritySnapshot(cfg)
if err != nil {
if !alert.TryEnqueue(d.alertCh, alert.Finding{
Severity: alert.Critical,
Check: "integrity",
Message: fmt.Sprintf("BINARY/CONFIG TAMPER DETECTED: %v", err),
Timestamp: time.Now(),
}) {
atomic.AddInt64(&d.droppedAlerts, 1)
fmt.Fprintf(os.Stderr, "[%s] alert channel full, dropping integrity finding\n", ts())
}
return
}
// Age out stale dry-run-block records so the status surface
// reflects recent activity instead of months-old entries left
// over from a previous dry-run window. Keeping a 7-day rolling
// window matches the operator workflow of reviewing a week of
// would-have-been-blocks before flipping to live.
if sdb := store.Global(); sdb != nil {
sdb.PurgeDryRunBlocksOlderThan(time.Now().Add(-7 * 24 * time.Hour))
}
scanCtx, gaps := checks.WithCoverageGaps(d.scanContext())
findings, purgeChecks := checks.RunTierWithContext(scanCtx, cfg, d.store, tier)
d.processScanFindingsWithCoverage(cfg, findings, purgeChecks, gaps.Snapshot(), "periodic")
}
func (d *Daemon) verifyPeriodicIntegritySnapshot(cfg *config.Config) (*config.Config, error) {
if err := integrity.Verify(d.binaryPath, cfg); err != nil {
latest := d.currentCfg()
if latest != nil && latest != cfg {
retryErr := integrity.Verify(d.binaryPath, latest)
if retryErr == nil {
return latest, nil
}
return latest, retryErr
}
return cfg, err
}
return cfg, nil
}
// heartbeat sends periodic pings to dead man's switch.
func (d *Daemon) heartbeat() {
defer d.wg.Done()
ticker := time.NewTicker(10 * time.Minute)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
alert.SendHeartbeat(d.currentCfg())
d.hijackDetector.Cleanup()
// Failed startup can leave only the recovery boundary available.
// Settle pending outcomes before attempting cleanup mutations.
recoverFirewallActions(d.fwActions)
// Clean expired temporary allows
if d.fwEngine != nil {
d.fwEngine.CleanExpiredAllows()
d.fwEngine.CleanExpiredSubnets()
}
// Clean expired temporary whitelist and threat entries
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.PruneExpiredWhitelist()
tdb.PruneExpiredThreats()
}
}
}
}
// startPHPRelay implements the platform gate from Stage 1 spec section
// 9 (O1). Emits a Warning if the host is not cPanel; otherwise locates
// the exim binary for AutoFreeze and (in O2) wires the spool watcher,
// pipeline, and Flow E ticker.
func (d *Daemon) startPHPRelay() {
info := platform.Detect()
if !info.IsCPanel() {
alert.TryEnqueue(d.alertCh, alert.Finding{
Severity: alert.Warning,
Check: "email_php_relay_disabled",
Message: "php_relay disabled: not a cPanel host",
Timestamp: time.Now(),
})
return
}
if !d.cfg.EmailProtection.PHPRelay.Enabled {
return
}
if path, err := exec.LookPath("exim"); err == nil {
eximBinary = path
} else {
alert.TryEnqueue(d.alertCh, alert.Finding{
Severity: alert.Warning,
Check: "email_php_relay_no_exim",
Message: "php_relay auto-action disabled: exim binary not in PATH",
Timestamp: time.Now(),
})
}
// Bridge to the linux-only wiring (Phase O2). On non-linux GOOS
// the stub in php_relay_wiring_other.go is a no-op.
startPHPRelayLinux(d)
}
func (d *Daemon) startLogWatchers() {
d.loadMailGoodSource()
// Session log handler wrapper - feeds events to both the alert handler and hijack detector
sessionHandler := func(line string, cfg *config.Config) []alert.Finding {
// Feed to hijack detector (tracks password changes + correlates with logins)
ParseSessionLineForHijack(line, d.hijackDetector)
// Regular session log handling
return parseSessionLogLine(line, cfg)
}
hostInfo := platform.Detect()
type logFile struct {
name string
path string
handler func(string, *config.Config) []alert.Finding
}
var logFiles []logFile
// Generic Linux auth log. RHEL-family uses /var/log/secure, Debian
// family uses /var/log/auth.log. Only register the log appropriate
// for the detected OS so we don't spam "not found, retrying" forever.
if hostInfo.IsDebianFamily() {
logFiles = append(logFiles, logFile{"", "/var/log/auth.log", parseSecureLogLine})
} else {
logFiles = append(logFiles, logFile{"", "/var/log/secure", parseSecureLogLine})
}
// eximHandler wraps parseEximLogLine (unchanged) and augments the result
// with smtpAuthTracker findings for dovecot authenticator failures and
// smtpProbeTracker findings for raw connect-rate abuse (scanners that
// probe-and-disconnect without ever reaching AUTH).
eximHandler := func(line string, cfg *config.Config) []alert.Finding {
findings := parseEximLogLine(line, cfg)
// Connect-rate signal fires before any AUTH attempt.
if probeIP := parseEximSMTPConnectIP(line); probeIP != "" {
if parsed := net.ParseIP(probeIP); parsed != nil {
if v4 := parsed.To4(); v4 != nil {
probeIP = v4.String()
}
}
if !isInfraIPDaemon(probeIP, cfg.InfraIPs) && !isPrivateOrLoopback(probeIP) {
if d.smtpProbeTracker != nil {
findings = append(findings, d.smtpProbeTracker.Record(probeIP)...)
}
}
}
if strings.Contains(line, "authenticator failed") && strings.Contains(line, "dovecot") {
ip := eximlog.ClientIP(line)
account := extractSetID(line)
// Canonicalize IPv4-mapped IPv6 (::ffff:a.b.c.d) to plain IPv4 so the
// tracker doesn't double-count the same attacker as two IPs.
if ip != "" {
if parsed := net.ParseIP(ip); parsed != nil {
if v4 := parsed.To4(); v4 != nil {
ip = v4.String()
}
}
}
if ip != "" && !isInfraIPDaemon(ip, cfg.InfraIPs) && !isPrivateOrLoopback(ip) {
if d.smtpAuthTracker != nil {
findings = append(findings, d.smtpAuthTracker.Record(ip, account)...)
}
}
}
// Authenticated deliveries prove the source can still log in, which
// disqualifies it from the slow-brute block (an office NAT with one
// stale device keeps working devices too; a walker never succeeds).
recordEximSMTPAuthSuccess(line, cfg, d.smtpAuthTracker)
return findings
}
// mailHandler composes parseDovecotLogLine (preserving email_suspicious_geo)
// with mailAuthTracker augmentation for IMAP/POP3/ManageSieve brute-force,
// subnet spray, account spray, and compromise detection.
mailHandler := func(line string, cfg *config.Config) []alert.Finding {
findings := parseDovecotLogLine(line, cfg)
authLine := isMailAuthLine(line)
// Auth-backend failures (dovecot cannot reach the credential backend)
// arrive on their own lines, not login lines. Feed them to the degraded
// gate so a backend outage pauses brute-force auto-block instead of
// mass-blocking every legitimate user whose login now fails.
if d.mailAuthTracker != nil && !authLine && isMailAuthBackendError(line) {
return append(findings, d.mailAuthTracker.RecordBackendFailure()...)
}
if !authLine {
return findings
}
ip, account, success := extractMailLoginEvent(line)
if ip == "" {
return findings
}
if parsed := net.ParseIP(ip); parsed != nil {
if v4 := parsed.To4(); v4 != nil {
ip = v4.String()
}
}
if isInfraIPDaemon(ip, cfg.InfraIPs) || isPrivateOrLoopback(ip) {
return findings
}
if d.mailAuthTracker == nil {
return findings
}
if success {
findings = append(findings, d.mailAuthTracker.RecordSuccess(ip, account)...)
} else {
findings = append(findings, recordDovecotFailure(d.mailAuthTracker, ip, account, line)...)
}
return findings
}
// cPanel-specific logs only watch these on cPanel hosts. On plain
// Ubuntu/AlmaLinux they do not exist and the old code spammed
// "not found, will retry every 60s" forever.
if hostInfo.IsCPanel() {
logFiles = append(logFiles,
logFile{"", "/usr/local/cpanel/logs/session_log", sessionHandler},
logFile{"", "/usr/local/cpanel/logs/access_log", parseAccessLogLineEnhanced},
logFile{"", "/var/log/messages", parseFTPLogLine},
)
}
if shouldWatchEximMainlog(hostInfo, os.Stat) {
logFiles = append(logFiles, logFile{"", eximMainlogPath, eximHandler})
}
d.startMailLogReader(hostInfo.MailLogPath(), mailHandler)
// Only receive PHP Shield events if enabled AND actually installed. A stale
// php_shield.enabled flag (e.g. after an upgrade wiped /opt/csm) would
// otherwise spin the missing-socket retry forever; warn once with
// a remediation hint instead.
if watch, warnNotInstalled := phpShieldWatchDecision(d.cfg.PHPShield.Enabled, phpShieldInstalled()); watch {
d.wg.Add(1)
obs.Go("php-shield-events", d.watchPHPShieldEvents)
} else if warnNotInstalled {
csmlog.Warn(phpShieldMissingScriptWarning, "shield", phpShieldScriptPath)
d.MarkWatcher("php_shield", false)
}
// ModSecurity error log - auto-discover path based on detected web server.
if modsecPath := discoverModSecLogPath(d.cfg); modsecPath != "" {
logFiles = append(logFiles, logFile{"modsec", modsecPath, parseModSecLogLineDeduped})
} else if hostInfo.WebServer != platform.WSNone {
// Only bother with the retry loop if a web server is actually
// present. Headless hosts don't need this.
fmt.Fprintf(os.Stderr, "[%s] ModSecurity error log not found (checked %v), will retry every 60s\n", ts(), hostInfo.ErrorLogPaths)
d.MarkWatcher("modsec", false)
d.wg.Add(1)
obs.Go("logwatch-modsec-retry", func() {
defer d.wg.Done()
ticker := time.NewTicker(logWatcherRetryInterval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
path := discoverModSecLogPath(d.cfg)
if path == "" {
continue
}
w, err := NewLogWatcher(path, d.cfg, parseModSecLogLineDeduped, d.alertCh)
if err != nil {
continue
}
d.logWatchersMu.Lock()
d.logWatchers = append(d.logWatchers, w)
d.logWatchersMu.Unlock()
d.wg.Add(1)
obs.Go("logwatch-modsec", func() {
defer d.wg.Done()
w.Run(d.stopCh)
})
csmlog.Info("watching log (appeared after retry)", "path", path)
d.MarkWatcher("modsec", true)
return
}
}
})
}
// Real-time access log watcher for wp-login/xmlrpc brute force detection.
// Auto-discover path from platform info (Apache/Nginx/cPanel aware).
if accessLogPath := discoverAccessLogPath(); accessLogPath != "" {
logFiles = append(logFiles, logFile{"", accessLogPath, parseAccessLogBruteForce})
} else if hostInfo.WebServer != platform.WSNone && len(hostInfo.AccessLogPaths) > 0 {
csmlog.Warn("access log not found, will retry every 60s", "candidates", fmt.Sprintf("%v", hostInfo.AccessLogPaths))
d.wg.Add(1)
accessPath := hostInfo.AccessLogPaths[0]
obs.Go("logwatch-access-retry", func() { d.retryLogWatcher(accessPath, parseAccessLogBruteForce) })
}
// Start background eviction for modsec dedup/escalation state
StartModSecEviction(d.stopCh, func() *config.Config { return d.currentCfg() })
// Start background eviction for access log brute force state
StartAccessLogEviction(d.stopCh)
// Start background eviction for email rate limiting state
StartEmailRateEviction(d.stopCh)
// Start background eviction for cloud-relay per-user windows so the
// sync.Map does not grow linearly with every distinct authenticated
// sender ever seen.
StartCloudRelayEviction(d.stopCh)
// Start background purge for SMTP brute-force tracker
d.wg.Add(1)
obs.Go("smtp-tracker-purge", func() {
defer d.wg.Done()
ticker := time.NewTicker(1 * time.Minute)
defer ticker.Stop()
tick := 0
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
if d.smtpAuthTracker != nil {
d.smtpAuthTracker.Purge()
}
if d.smtpProbeTracker != nil {
d.smtpProbeTracker.Purge()
}
// Diagnostic: surface whether the SMTP/mail brute-force
// trackers are actually seeing auth failures and emitting
// blockable findings. A nonzero record_calls with zero
// findings_emitted over a sustained attack means the
// threshold path, not the wiring, is the gap to chase.
tick++
if tick%10 == 0 {
if d.smtpAuthTracker != nil {
sc, se := d.smtpAuthTracker.Stats()
csmlog.Info("smtp brute tracker stats",
"record_calls", sc, "findings_emitted", se, "tracked", d.smtpAuthTracker.Size())
}
if d.mailAuthTracker != nil {
mc, me := d.mailAuthTracker.Stats()
csmlog.Info("mail brute tracker stats",
"record_calls", mc, "findings_emitted", me, "tracked", d.mailAuthTracker.Size())
}
}
}
}
})
// Start background purge for mail (IMAP/POP3) brute-force tracker. It also
// persists established good-source standing every minute and on shutdown, so
// a daemon restart does not re-open the brute-force false-positive
// cold-start window.
d.wg.Add(1)
obs.Go("mail-tracker-purge", func() {
defer d.wg.Done()
ticker := time.NewTicker(1 * time.Minute)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
d.persistMailGoodSource()
return
case <-ticker.C:
if d.mailAuthTracker != nil {
d.mailAuthTracker.Purge()
d.persistMailGoodSource()
}
}
}
})
// Mail auth backend probe: actively detect a cpdoveauthd outage that
// chkservd's port checks miss (dovecot stays up while auth is broken),
// suppress brute-force auto-block while it lasts, and optionally self-heal.
// cPanel only -- elsewhere there is no such socket to probe.
if hostInfo.IsCPanel() && d.authBackend != nil {
if !d.emitAuthBackendFindings() {
return
}
d.wg.Add(1)
obs.Go("mail-auth-backend-probe", func() {
defer d.wg.Done()
ticker := time.NewTicker(mailAuthProbeInterval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
if !d.emitAuthBackendFindings() {
return
}
}
}
})
}
for _, lf := range logFiles {
w, err := NewLogWatcher(lf.path, d.cfg, lf.handler, d.alertCh)
if err != nil {
if os.IsNotExist(err) {
// File doesn't exist yet - retry periodically until it appears
if lf.name != "" {
d.MarkWatcher(lf.name, false)
}
d.wg.Add(1)
path, handler, name := lf.path, lf.handler, lf.name
obs.Go("logwatch-retry", func() { d.retryLogWatcherNamed(path, handler, name) })
} else {
fmt.Fprintf(os.Stderr, "[%s] Warning: could not watch %s: %v\n", ts(), lf.path, err)
if lf.name != "" {
d.MarkWatcher(lf.name, false)
}
}
continue
}
d.logWatchers = append(d.logWatchers, w)
d.wg.Add(1)
watcher := w
obs.Go("logwatch", func() {
defer d.wg.Done()
watcher.Run(d.stopCh)
})
csmlog.Info("watching log", "path", lf.path)
if lf.name != "" {
d.MarkWatcher(lf.name, true)
}
}
}
// recordEximSMTPAuthSuccess feeds authenticated Exim acceptance lines to the
// slow SMTP-brute guard. extractAuthUser verifies a real top-level A= field;
// the IP is deliberately read only from H=, whose bracketed address is the
// connecting client rather than a HELO/subject/forwarded-header address.
func recordEximSMTPAuthSuccess(line string, cfg *config.Config, tracker *smtpAuthTracker) {
if tracker == nil || extractAuthUser(line) == "" {
return
}
hStart, hasHField := eximlog.HFieldStart(line)
if !hasHField {
return
}
ip := eximlog.HFieldClientIP(line[hStart:])
if ip == "" {
return
}
if parsed := net.ParseIP(ip); parsed != nil {
if v4 := parsed.To4(); v4 != nil {
ip = v4.String()
}
}
if !isInfraIPDaemon(ip, cfg.InfraIPs) && !isPrivateOrLoopback(ip) {
tracker.RecordSuccess(ip)
}
}
func (d *Daemon) emitAuthBackendFindings() bool {
if d.authBackend == nil {
return true
}
findings := d.authBackend.Observe()
for i, f := range findings {
if !alert.Enqueue(d.alertCh, f, d.stopCh) {
alert.RecordQueueLoss(d.alertCh, uint64(len(findings[i+1:])))
return false
}
}
return true
}
func (d *Daemon) handleMailLogSourceGone(err error) {
d.MarkWatcher("maillog", false)
finding := alert.Finding{
Severity: alert.Warning,
Check: "mail_log_source_unavailable",
Message: fmt.Sprintf("Mail log source unavailable: %v; brute-force and rate detection degraded while attachment is retried", err),
Timestamp: time.Now(),
}
select {
case <-d.stopCh:
return
default:
}
if !alert.TryEnqueue(d.alertCh, finding) {
atomic.AddInt64(&d.droppedAlerts, 1)
fmt.Fprintf(os.Stderr, "[%s] alert channel full, dropping maillog source finding\n", ts())
}
}
func (d *Daemon) handleMailLogSourceRestored() {
d.MarkWatcher("maillog", true)
}
func (d *Daemon) dispatchMailLogLine(line maillog.Line, handler LogLineHandler) bool {
findings := handler(line.Message, d.currentCfg())
for i, f := range findings {
if !alert.Enqueue(d.alertCh, f, d.stopCh) {
alert.RecordQueueLoss(d.alertCh, uint64(len(findings[i+1:])))
return false
}
}
return true
}
// retryLogWatcher polls for a missing log file every 60 seconds.
// When the file appears, it starts a watcher and returns.
func (d *Daemon) retryLogWatcher(path string, handler LogLineHandler) {
d.retryLogWatcherNamed(path, handler, "")
}
func shouldWatchEximMainlog(hostInfo platform.Info, stat func(string) (os.FileInfo, error)) bool {
if hostInfo.IsCPanel() {
return true
}
if stat == nil {
stat = os.Stat
}
if _, err := stat(eximMainlogPath); err != nil {
return !os.IsNotExist(err)
}
return true
}
func (d *Daemon) retryLogWatcherNamed(path string, handler LogLineHandler, name string) {
defer d.wg.Done()
csmlog.Warn("log not found, will retry every 60s", "path", path)
ticker := time.NewTicker(logWatcherRetryInterval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
w, err := NewLogWatcher(path, d.cfg, handler, d.alertCh)
if err != nil {
continue // still missing, keep retrying
}
d.logWatchersMu.Lock()
d.logWatchers = append(d.logWatchers, w)
d.logWatchersMu.Unlock()
d.wg.Add(1)
watcher := w
obs.Go("logwatch-late", func() {
defer d.wg.Done()
watcher.Run(d.stopCh)
})
csmlog.Info("watching log (appeared after retry)", "path", path)
if name != "" {
d.MarkWatcher(name, true)
}
return
}
}
}
func (d *Daemon) startWebUI() {
if !d.cfg.WebUI.Enabled {
return
}
srv, err := webui.New(d.cfg, d.store)
if err != nil {
csmlog.Error("webui init error", "err", err)
return
}
// Set version and signature count for status API
srv.SetVersion(d.version)
if scanner := signatures.Global(); scanner != nil {
srv.SetSigCount(scanner.RuleCount())
}
d.webServer = srv
srv.SetHealthProvider(d)
srv.SetFindingBus(d.findingBus)
if d.scanJobs != nil {
srv.SetScanJobs(d.scanJobs)
}
srv.SetIncidentCorrelator(IncidentCorrelator())
// Push web-UI verified_bots edits into the live registry + verifier so they
// take effect without a restart, the same path SIGHUP uses.
srv.SetVerifiedBotsReloader(func() error { d.reconcileVerifiedBots(); return nil })
srv.SetHealthInfo(d.FanotifyActive, d.LogWatcherCount)
if d.fwEngine != nil {
srv.SetIPBlocker(d.fwEngine)
}
d.wg.Add(1)
obs.Go("webui", func() {
defer d.wg.Done()
if err := srv.Start(); err != nil {
csmlog.Error("webui server error", "err", err)
}
})
}
func (d *Daemon) startPAMListener() {
pl, err := NewPAMListener(d.cfg, d.alertCh)
if err != nil {
csmlog.Warn("PAM listener not available", "err", err)
d.MarkWatcher("pamlistener", false)
return
}
d.pamListener = pl
d.MarkWatcher("pamlistener", true)
d.RegisterUpstreamProbe("pamlistener", pl.UpstreamResult)
d.wg.Add(1)
obs.Go("pam-listener", func() {
defer d.wg.Done()
pl.Run(d.stopCh)
})
csmlog.Info("PAM listener active", "socket", pamSocketPath)
}
func (d *Daemon) startControlListener() {
cl, err := NewControlListener(d)
if err != nil {
// The daemon can still function without the socket — periodic
// scans and webui keep running — but the CLI will hard-error
// because the socket is the expected path. Log loudly.
csmlog.Error("control listener not available", "err", err)
return
}
cl.scanJobs = d.scanJobs
d.controlListener = cl
d.wg.Add(1)
obs.Go("control-listener", func() {
defer d.wg.Done()
cl.Run(d.stopCh)
})
csmlog.Info("control listener active", "socket", controlSocketPath)
}
func (d *Daemon) startFileMonitor() {
fm, err := NewFileMonitor(d.cfg, d.alertCh) //nolint:staticcheck // The Linux constructor is fallible; the non-Linux stub always returns an error.
if err != nil { //nolint:staticcheck // The Linux constructor can also succeed.
csmlog.Warn("fanotify not available, falling back to periodic deep scan", "err", err)
d.MarkWatcher("fanotify", false)
return
}
fm.registerMetrics()
d.setFileMonitor(fm)
d.wg.Add(1)
obs.Go("fanotify", func() {
defer d.wg.Done()
fm.Run(d.stopCh)
})
csmlog.Info("fanotify file monitor active", "roots", fm.WatchScopeSummary())
d.MarkWatcher("fanotify", true)
}
func (d *Daemon) startSpoolWatcher() {
if !d.cfg.EmailAV.Enabled {
return
}
// Create ClamAV scanner. The socket path belongs to whoever packaged
// clamd, so a setting naming the wrong one falls back to a location that
// is actually answering: mail that is silently never scanned looks exactly
// like mail that came back clean.
clamdSocket, discovered := config.ResolveClamdSocket(d.cfg.EmailAV.ClamdSocket)
if discovered {
csmlog.Warn("email av: configured clamd socket is not answering; using a discovered one",
"configured", d.cfg.EmailAV.ClamdSocket, "using", clamdSocket)
}
clamScanner := emailav.NewClamdScanner(clamdSocket)
// YARA-X scanner over whichever backend initYaraBackend installs.
// The worker backend can come online after startup through the boot-retry
// path, so resolve yara.Active() lazily instead of capturing a nil backend
// forever while email AV is being wired.
yaraScanner := emailav.NewActiveYaraXScanner()
// Create orchestrator with both engines
scanners := []emailav.Scanner{clamScanner, yaraScanner}
orch := emailav.NewOrchestrator(scanners, d.cfg.EmailAV.ScanTimeoutDuration())
d.registerQueueSource("email_av", orch)
// Create quarantine
quar := emailav.NewQuarantine("/opt/csm/quarantine/email")
d.emailQuarantine = quar
// Create and start spool watcher
sw, err := NewSpoolWatcher(d.cfg, d.alertCh, orch, quar) //nolint:staticcheck // The Linux constructor is fallible; the non-Linux stub always returns an error.
if err != nil { //nolint:staticcheck // The Linux constructor can also succeed.
fmt.Fprintf(os.Stderr, "[%s] Email AV spool watcher not available: %v\n", ts(), err)
d.MarkWatcher("email_av_spool", false)
return
}
d.setSpoolWatcher(sw)
d.MarkWatcher("email_av_spool", true)
d.wg.Add(1)
obs.Go("spool-watcher", func() {
defer d.wg.Done()
d.runSpoolWatcherLoop(sw, orch, quar)
})
// Start quarantine cleanup goroutine
d.wg.Add(1)
obs.Go("email-quarantine-cleanup", d.emailQuarantineCleanup)
fmt.Fprintf(os.Stderr, "[%s] Email AV spool watcher active\n", ts())
}
// superviseWatcherRun runs run() to completion. If daemonStop closes while
// run() is still blocked, stop() is invoked so run() can return. The helper
// goroutine is reaped when run() returns on its own. This guarantees the live
// watcher instance is stopped on shutdown even after a crash-restart swapped a
// fresh instance in, which the external shutdown path (it only stops the
// instance registered via setSpoolWatcher) can miss, hanging wg.Wait forever.
func superviseWatcherRun(daemonStop <-chan struct{}, run, stop func()) {
done := make(chan struct{})
helperDone := make(chan struct{})
go func() {
defer close(helperDone)
select {
case <-daemonStop:
stop()
case <-done:
}
}()
run()
close(done)
<-helperDone
}
type spoolWatcherRuntime interface {
Run()
Stop()
}
func (d *Daemon) runSpoolWatcherLoop(sw *SpoolWatcher, orch *emailav.Orchestrator, quar *emailav.Quarantine) {
d.runSpoolWatcherLoopWithFactory(sw, 2*time.Second, func() (spoolWatcherRuntime, error) {
next, err := NewSpoolWatcher(d.cfg, d.alertCh, orch, quar) //nolint:staticcheck // The Linux constructor is fallible; the non-Linux stub always returns an error.
if err != nil { //nolint:staticcheck // The Linux constructor can also succeed.
return nil, err
}
d.setSpoolWatcher(next)
return next, nil
})
}
func (d *Daemon) runSpoolWatcherLoopWithFactory(current spoolWatcherRuntime, restartDelay time.Duration, newWatcher func() (spoolWatcherRuntime, error)) {
for {
superviseWatcherRun(d.stopCh, current.Run, current.Stop)
select {
case <-d.stopCh:
return
default:
}
fmt.Fprintf(os.Stderr, "[%s] Email AV spool watcher stopped unexpectedly; restarting in %s\n", ts(), restartDelay)
for {
select {
case <-d.stopCh:
return
case <-time.After(restartDelay):
}
next, err := newWatcher()
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] Email AV spool watcher restart failed: %v\n", ts(), err)
continue
}
current = next
break
}
}
}
func (d *Daemon) setSpoolWatcher(sw *SpoolWatcher) {
d.spoolWatcherMu.Lock()
if d.spoolWatcher != nil {
sw.inheritQueueHealth(d.spoolWatcher)
}
d.spoolWatcher = sw
d.spoolWatcherMu.Unlock()
d.syncEmailAVWebState()
}
func (d *Daemon) getSpoolWatcher() *SpoolWatcher {
d.spoolWatcherMu.Lock()
defer d.spoolWatcherMu.Unlock()
return d.spoolWatcher
}
// startForwarderWatcher starts the inotify watcher for /etc/valiases/.
func (d *Daemon) startForwarderWatcher() {
fw, err := NewForwarderWatcher(d.alertCh, d.cfg.EmailProtection.KnownForwarders) //nolint:staticcheck // The Linux constructor is fallible; the non-Linux stub always returns an error.
if err != nil { //nolint:staticcheck // The Linux constructor can also succeed.
fmt.Fprintf(os.Stderr, "[%s] Warning: forwarder watcher not started: %v\n", ts(), err)
d.MarkWatcher("forwarder", false)
return
}
d.forwarderWatcher = fw
d.registerQueueSource("forwarder", fw)
d.wg.Add(1)
obs.Go("forwarder-watcher", func() {
defer d.wg.Done()
fw.Run(d.stopCh)
})
csmlog.Info("watching log (inotify forwarder watcher)", "path", "/etc/valiases/")
d.MarkWatcher("forwarder", true)
}
func (d *Daemon) syncEmailAVWebState() {
if d.webServer == nil || d.emailQuarantine == nil {
return
}
d.webServer.SetEmailQuarantine(d.emailQuarantine)
if sw := d.getSpoolWatcher(); sw != nil {
if sw.PermissionMode() {
d.webServer.SetEmailAVWatcherMode("permission")
} else {
d.webServer.SetEmailAVWatcherMode("notification")
}
}
}
// emailQuarantineCleanup periodically removes expired quarantined email messages.
func (d *Daemon) emailQuarantineCleanup() {
defer d.wg.Done()
ticker := time.NewTicker(1 * time.Hour)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
if d.emailQuarantine != nil {
cleaned, err := d.emailQuarantine.CleanExpired(30 * 24 * time.Hour)
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] Email quarantine cleanup error: %v\n", ts(), err)
} else if cleaned > 0 {
fmt.Fprintf(os.Stderr, "[%s] Email quarantine cleanup: removed %d expired messages\n", ts(), cleaned)
}
}
}
}
}
func (d *Daemon) startChallengeServer() {
if d.challengeServer != nil {
return
}
// Challenge suppression must fail open until the listener is known to be
// available. This also clears stale package wiring in repeated daemon
// construction paths used by tests and embedding callers.
checks.SetChallengeIPList(nil)
alert.ChallengedIPFunc = nil
if !d.cfg.Challenge.Enabled {
return
}
if d.fwEngine == nil {
fmt.Fprintf(os.Stderr, "[%s] Challenge server requires firewall to be enabled (for escalation). Skipping.\n", ts())
return
}
d.ipList = challenge.NewIPList(challenge.DefaultMapPath)
if platform.Detect().WebServer == platform.WSNginx {
d.ipList.SetNginxMap(challenge.DefaultNginxMapPath, d.reloadChallengeNginxMap)
}
d.attachChallengePortGate()
srv := challenge.New(d.cfg, d.ipList)
listener, err := srv.Listen()
if err != nil {
if d.challengeGate != nil {
if closeErr := d.challengeGate.Close(); closeErr != nil {
fmt.Fprintf(os.Stderr, "[%s] challenge port-gate cleanup: %v\n", ts(), closeErr)
}
d.challengeGate = nil
}
d.ipList = nil
csmlog.Error("challenge server unavailable; challenge routing disabled", "err", err)
return
}
checks.SetChallengeIPList(d.ipList)
alert.ChallengedIPFunc = d.ipList.Contains
d.challengeServer = srv
d.wg.Add(1)
obs.Go("challenge-server", func() {
defer d.wg.Done()
csmlog.Info("challenge server active", "port", d.cfg.Challenge.ListenPort)
if err := srv.Serve(listener); err != nil && !errors.Is(err, http.ErrServerClosed) {
csmlog.Error("challenge server error", "err", err)
}
})
}
// attachChallengePortGate installs the nftables port-gate for the
// challenge listener when the operator opts in. The gate is silently
// absent when the listener is loopback-only (no off-host traffic can
// reach it anyway) or on non-Linux builds.
func (d *Daemon) attachChallengePortGate() {
if !d.cfg.Challenge.PortGate.Enabled {
return
}
gate, err := challenge.NewPortGate(challenge.PortGateConfig{
ListenAddr: d.cfg.Challenge.ListenAddr,
ListenPort: d.cfg.Challenge.ListenPort,
InfraCIDRs: d.cfg.InfraIPs,
})
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] challenge port-gate install failed: %v (listener stays publicly reachable)\n", ts(), err)
return
}
if gate == nil {
csmlog.Info("challenge port-gate skipped (loopback listener or non-Linux build)",
"listen_addr", d.cfg.Challenge.ListenAddr)
return
}
d.challengeGate = gate
d.ipList.SetPortGate(gate)
csmlog.Info("challenge port-gate active", "port", d.cfg.Challenge.ListenPort)
}
func (d *Daemon) reloadChallengeNginxMap() error {
// #nosec G204 -- static binary and arguments; no operator input is passed.
out, err := exec.Command("nginx", "-s", "reload").CombinedOutput()
if err != nil {
return fmt.Errorf("nginx -s reload: %w: %s", err, strings.TrimSpace(string(out)))
}
return nil
}
func (d *Daemon) challengeEscalator() {
defer d.wg.Done()
ticker := time.NewTicker(60 * time.Second)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
// Re-read the block expiry each tick so a SIGHUP that
// changes auto_response.block_expiry takes effect on the
// next escalation without requiring a restart.
expiry := parseBlockExpiry(d.currentCfg().AutoResponse.BlockExpiry)
if d.challengeServer != nil {
d.challengeServer.CleanExpired()
}
d.escalateExpiredChallenges(expiry)
}
}
}
// escalateExpiredChallenges hard-blocks every challenge-list entry whose
// window lapsed. Blocks route through the checks.ApplyBlock chokepoint so
// escalations leave the same evidence trail as scan auto-blocks (threat-DB
// row, tracker entry, digest visibility, permblock counting) instead of
// only a stderr line.
func (d *Daemon) escalateExpiredChallenges(expiry time.Duration) {
if d.ipList == nil {
return
}
expired := d.ipList.ExpiredEntries()
if len(expired) == 0 {
return
}
cfg := d.currentCfg()
var recorded []alert.Finding
for _, e := range expired {
res, err := checks.ApplyBlock(cfg, checks.ApplyBlockRequest{
IP: e.IP,
EngineReason: fmt.Sprintf("CSM challenge-timeout: %s", truncateStr(e.Reason, 100)),
Reason: challengeTimeoutReasonPrefix + truncateStr(e.Reason, 100),
TTL: expiry,
Source: checks.BlockSourceChallenge,
FindingID: e.FindingID,
})
recorded = append(recorded, res.Findings...)
if err != nil {
// Own-interface / infra IPs are never blockable, and a host
// without a firewall engine cannot escalate; both are expected
// no-ops, not failures worth logging.
if !isProtectedIPRefusal(err) && !errors.Is(err, checks.ErrNoIPBlocker) {
fmt.Fprintf(os.Stderr, "[%s] challenge-escalate: error blocking %s: %v\n", ts(), e.IP, err)
}
if res.Outcome != firewall.BlockOutcomeLive || !errors.Is(err, firewall.ErrActionAuditPending) {
continue
}
}
observeChallengeEscalated(res.Outcome)
fmt.Fprintf(os.Stderr, "[%s] %s\n", ts(), challengeEscalateLogLine(e.IP, res.Outcome))
}
d.recordAppliedBlocks(recorded)
}
// recordAppliedBlocks routes chokepoint findings from async block sources
// (challenge escalation, central intel, incident spray) through the same
// bookkeeping dispatchBatch gives scan blocks: block digest, attack-DB
// blocked marker, finding history, and alert dispatch, which applies the
// standard suppression rules. Deduplication happens before the fanout so every
// sink sees the same single copy, and a failure in one sink does not skip the
// sinks that follow it.
func (d *Daemon) recordAppliedBlocks(findings []alert.Finding) {
alert.FillTimestamps(findings, time.Now())
findings = alert.Deduplicate(findings)
if len(findings) == 0 {
return
}
d.observeBlocks(findings)
if adb := attackdb.Global(); adb != nil {
for _, f := range findings {
if f.Check != "auto_block" || f.Severity != alert.Critical {
continue
}
if ip := checks.ExtractIPFromFinding(f); ip != "" {
adb.MarkBlocked(ip)
}
}
}
if d.store != nil {
d.store.AppendHistory(findings)
}
alertable := findings
if d.store != nil {
alertable = filterUnsuppressedFindings(d.store, findings, d.store.LoadSuppressions())
}
if err := alert.DispatchWithEnforcement(d.currentCfg(), alertable, findings, findings); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Applied-block alert dispatch error: %v\n", ts(), err)
}
}
// applyIncidentSprayBlock is the incident correlator's firewall hand-off,
// routed through the chokepoint so spray blocks leave evidence and reach
// the digest.
func (d *Daemon) applyIncidentSprayBlock(ip, reason string, timeout time.Duration, findingID string) (bool, error) {
res, err := checks.ApplyBlock(d.currentCfg(), checks.ApplyBlockRequest{
IP: ip,
EngineReason: reason,
Reason: reason,
TTL: timeout,
Source: checks.BlockSourceIncident,
FindingID: findingID,
})
d.recordAppliedBlocks(res.Findings)
live := res.Outcome == firewall.BlockOutcomeLive && (err == nil || errors.Is(err, firewall.ErrActionAuditPending))
return live, err
}
var (
challengeEscalatedMetric *metrics.CounterVec
challengeEscalatedMetricOnce sync.Once
)
// challengeEscalatedCounter returns the lazily registered metric. Count paths
// use the same Once gate as increments so first status reads cannot race the
// first escalation.
func challengeEscalatedCounter() *metrics.CounterVec {
challengeEscalatedMetricOnce.Do(func() {
challengeEscalatedMetric = metrics.NewCounterVec(
"csm_challenge_escalated_total",
"Challenge-timeout escalations, by firewall outcome (live, noop, dry_run, allowed, allowlisted).",
[]string{"outcome"},
)
metrics.MustRegister("csm_challenge_escalated_total", challengeEscalatedMetric)
})
return challengeEscalatedMetric
}
// observeChallengeEscalated counts one challenge-timeout escalation, labelled by
// the firewall outcome (live=a new hard block landed, noop=the IP was already
// blocked, plus dry_run/allowed/allowlisted). It lets operators see how many
// challenges became real blocks versus no-ops. Registered lazily on first use.
func observeChallengeEscalated(outcome firewall.BlockOutcome) {
challengeEscalatedCounter().With(string(outcome)).Inc()
}
// challengeEscalatedCount returns how many challenge timeouts escalated to a new
// hard block (outcome=live) since daemon start, for the web UI challenge panel.
func challengeEscalatedCount() int {
return int(challengeEscalatedCounter().With(string(firewall.BlockOutcomeLive)).Value())
}
// challengeEscalateLogLine renders the stderr line for one challenge-timeout
// escalation. Only a live block claims "hard-blocked"; a no-op (the IP was
// already hard-blocked, e.g. a confirmed-threat finding blocked it while it sat
// on the challenge list) or a verdict downgrade must not, so incident review is
// not misled by a block that never landed.
func challengeEscalateLogLine(ip string, outcome firewall.BlockOutcome) string {
switch outcome {
case firewall.BlockOutcomeLive:
return fmt.Sprintf("CHALLENGE-ESCALATE: %s timed out, hard-blocked", ip)
case firewall.BlockOutcomeDryRun:
return fmt.Sprintf("CHALLENGE-ESCALATE [dry-run]: %s timed out, would be hard-blocked", ip)
default:
return fmt.Sprintf("challenge-escalate: %s timed out, no new block (outcome: %s)", ip, outcome)
}
}
func parseBlockExpiry(s string) time.Duration {
if s == "" {
return 24 * time.Hour
}
d, err := time.ParseDuration(s)
if err != nil {
return 24 * time.Hour
}
return d
}
func truncateStr(s string, max int) string {
if len(s) <= max {
return s
}
return s[:max]
}
func (d *Daemon) initGeoIP() {
dbDir := filepath.Join(d.cfg.StatePath, "geoip")
db := geoip.Open(dbDir)
if db != nil {
d.geoipDB = db
setGeoIPDB(db) // make available to log watcher handlers for country filtering
if d.webServer != nil {
d.webServer.SetGeoIPDB(db)
}
}
}
// publishGeoIP reloads existing GeoIP databases or creates a new DB
// if databases were downloaded for the first time.
// Mutex-protected: safe to call from geoipUpdater goroutine and SIGHUP handler concurrently.
func (d *Daemon) publishGeoIP() {
d.geoipMu.Lock()
defer d.geoipMu.Unlock()
if d.geoipDB != nil {
if err := d.geoipDB.Reload(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] GeoIP reload error: %v\n", ts(), err)
} else {
fmt.Fprintf(os.Stderr, "[%s] GeoIP databases reloaded\n", ts())
}
return
}
// First-time: no DB existed at startup, try to open freshly downloaded files
dbDir := filepath.Join(d.cfg.StatePath, "geoip")
db := geoip.OpenFresh(dbDir)
if db != nil {
d.geoipDB = db
setGeoIPDB(db)
if d.webServer != nil {
d.webServer.SetGeoIPDB(db)
}
fmt.Fprintf(os.Stderr, "[%s] GeoIP databases loaded for the first time\n", ts())
}
}
// geoipUpdater periodically downloads updated GeoLite2 databases.
func (d *Daemon) geoipUpdater() {
defer d.wg.Done()
// Skip if no credentials configured
if d.cfg.GeoIP.AccountID == "" || d.cfg.GeoIP.LicenseKey == "" {
return
}
// Skip if auto_update is explicitly false
if d.cfg.GeoIP.AutoUpdate != nil && !*d.cfg.GeoIP.AutoUpdate {
return
}
interval := 24 * time.Hour
if d.cfg.GeoIP.UpdateInterval != "" {
if parsed, err := time.ParseDuration(d.cfg.GeoIP.UpdateInterval); err == nil && parsed >= time.Hour {
interval = parsed
}
}
// Wait 5 minutes before first update attempt (let the daemon stabilize)
select {
case <-d.stopCh:
return
case <-time.After(5 * time.Minute):
}
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
d.doGeoIPUpdate()
select {
case <-d.stopCh:
return
case <-ticker.C:
}
}
}
func (d *Daemon) doGeoIPUpdate() {
results := geoip.Update(
filepath.Join(d.cfg.StatePath, "geoip"),
d.cfg.GeoIP.AccountID,
d.cfg.GeoIP.LicenseKey,
d.cfg.GeoIP.Editions,
)
if results == nil {
return
}
anyUpdated := false
for _, r := range results {
switch r.Status {
case "updated":
fmt.Fprintf(os.Stderr, "[%s] GeoIP auto-update: %s updated\n", ts(), r.Edition)
anyUpdated = true
case "up_to_date":
// silent
case "error":
fmt.Fprintf(os.Stderr, "[%s] GeoIP auto-update: %s error: %v\n", ts(), r.Edition, r.Err)
}
}
if anyUpdated {
d.publishGeoIP()
}
}
func (d *Daemon) autoResponseDryRunEnabled() bool {
return d.activeOrStartupCfg().AutoResponseDryRunEnabled()
}
func (d *Daemon) askVerdictCallback(ctx context.Context, ip, reason string) (string, string, string, error) {
cfg := d.activeOrStartupCfg()
if cfg == nil || !cfg.AutoResponse.VerdictCallback.Enabled {
return "", "", "", nil
}
vcCfg := cfg.AutoResponse.VerdictCallback
vc := verdict.New(verdict.Config{
URL: vcCfg.URL,
HMACSecret: vcCfg.HMACSecret,
HMACSecretEnv: vcCfg.HMACSecretEnv,
RequireResponseSignature: vcCfg.RequireResponseSignature,
AllowUnsigned: vcCfg.AllowUnsigned,
Timeout: time.Duration(vcCfg.TimeoutSec) * time.Second,
})
resp, err := vc.Ask(ctx, verdict.Request{
IP: ip,
Reason: reason,
Severity: "auto",
Source: "auto_response",
})
if err != nil {
return "", "", "", err
}
return resp.Verdict, resp.TenantID, resp.Note, nil
}
func dynDNSUnresolvableFinding(host string) alert.Finding {
return alert.Finding{
Check: "infra_ips_unresolvable",
Severity: alert.Warning,
Message: fmt.Sprintf("dynamic firewall host %s has not resolved within grace period", host),
Details: "Verify DNS for the host or remove it from infra_ips or firewall.dyndns_hosts. While unresolvable, the previous resolved IP remains protected and a rotated IP will not be protected.",
Timestamp: time.Now(),
}
}
// cloudflareRefreshLoop fetches Cloudflare IPs and updates the firewall sets periodically.
func (d *Daemon) cloudflareRefreshLoop() {
defer d.wg.Done()
interval := time.Duration(d.cfg.Cloudflare.RefreshHours) * time.Hour
// Every fetch is cancelled by shutdown. Without this the startup fetch runs
// before the select below is ever reached, so a host that cannot reach
// cloudflare.com holds shutdown for the HTTP timeout.
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() {
select {
case <-d.stopCh:
cancel()
case <-ctx.Done():
}
}()
// Fetch immediately on startup
d.refreshCloudflareIPs(ctx)
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
d.refreshCloudflareIPs(ctx)
}
}
}
// fetchCloudflareIPs downloads the current Cloudflare ranges. Var so tests
// can feed the refresh an unusable result.
var fetchCloudflareIPs = firewall.FetchCloudflareIPs
func (d *Daemon) refreshCloudflareIPs(ctx context.Context) {
ipv4, ipv6, err := fetchCloudflareIPs(ctx)
if err != nil {
csmlog.Error("cloudflare IP fetch error", "err", err)
}
fresh := len(ipv4) > 0 || len(ipv6) > 0
cached4, cached6 := firewall.LoadCFState(d.cfg.StatePath)
if len(ipv4) == 0 {
ipv4 = cached4
}
if len(ipv6) == 0 {
ipv6 = cached6
}
if len(ipv4) == 0 && len(ipv6) == 0 {
csmlog.Error("cloudflare IP refresh has no fetched or cached ranges")
return
}
// Restore the local auto-block guard even when nftables is unavailable.
// Otherwise the daemon's first refresh failure leaves cached edges
// blockable until the network recovers.
allCF := make([]string, 0, len(ipv4)+len(ipv6))
allCF = append(allCF, ipv4...)
allCF = append(allCF, ipv6...)
checks.SetCloudflareNets(allCF)
if d.fwEngine != nil {
if err := d.fwEngine.UpdateCloudflareSet(ipv4, ipv6); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Cloudflare set update error: %v\n", ts(), err)
return
}
}
if fresh {
if err := firewall.SaveCFState(d.cfg.StatePath, ipv4, ipv6, time.Now()); err != nil {
csmlog.Error("cloudflare state save error", "err", err)
}
}
}
// signatureUpdater periodically downloads new rules and reloads scanners.
func (d *Daemon) signatureUpdater() {
defer d.wg.Done()
yamlEnabled := d.cfg.Signatures.UpdateURL != ""
forgeEnabled := d.cfg.Signatures.YaraForge.Enabled && yara.Available()
if !yamlEnabled && !forgeEnabled {
return
}
select {
case <-d.stopCh:
return
case <-time.After(5 * time.Minute):
}
yamlInterval := 24 * time.Hour
if d.cfg.Signatures.UpdateInterval != "" {
if parsed, err := time.ParseDuration(d.cfg.Signatures.UpdateInterval); err == nil && parsed >= time.Hour {
yamlInterval = parsed
}
}
forgeInterval := 168 * time.Hour
if d.cfg.Signatures.YaraForge.UpdateInterval != "" {
if parsed, err := time.ParseDuration(d.cfg.Signatures.YaraForge.UpdateInterval); err == nil && parsed >= time.Hour {
forgeInterval = parsed
}
}
tickInterval := yamlInterval
if forgeEnabled && forgeInterval < tickInterval {
tickInterval = forgeInterval
}
if !yamlEnabled {
tickInterval = forgeInterval
}
var lastYAML, lastForge time.Time
ticker := time.NewTicker(tickInterval)
defer ticker.Stop()
for {
now := time.Now()
if yamlEnabled && now.Sub(lastYAML) >= yamlInterval {
d.doSignatureUpdate()
lastYAML = now
}
if forgeEnabled && now.Sub(lastForge) >= forgeInterval {
d.doForgeUpdate()
lastForge = now
}
select {
case <-d.stopCh:
return
case <-ticker.C:
}
}
}
func (d *Daemon) doSignatureUpdate() {
count, err := signatures.Update(d.cfg.Signatures.RulesDir, d.cfg.Signatures.UpdateURL, d.cfg.Signatures.SigningKey, signatures.UpdateOptions{
AllowRuleCountDecrease: d.cfg.Signatures.AllowRuleCountDecrease,
})
if err != nil {
d.reportSignatureUpdateError(err)
return
}
fmt.Fprintf(os.Stderr, "[%s] Signature auto-update: %d rules downloaded\n", ts(), count)
d.reloadSignatures()
}
func (d *Daemon) reportSignatureUpdateError(err error) {
fmt.Fprintf(os.Stderr, "[%s] Signature auto-update failed: %v\n", ts(), err)
if errors.Is(err, signatures.ErrUpdateRollback) {
d.emitReloadFinding(alert.Critical, "signature_update_rollback",
fmt.Sprintf("Signed rule update refused by rollback protection: %v", err))
}
}
// forgeRollbackNeeded reports whether a freshly installed Forge ruleset has to
// be undone, given the scanner's rule total after reload and the number of
// rules the new archive carries. A total that cannot even account for the new
// archive means it did not compile in.
//
// The pre-reload total is deliberately not consulted. It includes the Forge
// rules being replaced, so an upstream release carrying fewer rules than the
// installed one reads as a conflict, and the recovery action then deletes the
// entire Forge file -- discarding thousands of working rules over a decrease of
// a few dozen.
func forgeRollbackNeeded(newCount, forgeRuleCount int) bool {
return newCount < forgeRuleCount
}
func (d *Daemon) doForgeUpdate() {
// yara.Active() resolves to the in-process scanner or the worker
// supervisor depending on signatures.yara_worker_enabled; both
// satisfy the Reload/RuleCount calls this routine makes, so the
// forge update path is backend-agnostic.
yaraScanner := yara.Active()
if yaraScanner == nil {
// No YARA backend active (build without yara tag or no rules
// dir). Skip Forge update - rules can't be loaded anyway.
return
}
db := store.Global()
currentVersion := ""
if db != nil {
currentVersion = db.GetMetaString("forge_version_" + d.cfg.Signatures.YaraForge.Tier)
}
newVersion, count, err := signatures.ForgeUpdateFromURL(
d.cfg.Signatures.RulesDir,
d.cfg.Signatures.YaraForge.Tier,
currentVersion,
d.cfg.Signatures.SigningKey,
d.cfg.Signatures.YaraForge.DownloadURL,
d.cfg.Signatures.DisabledRules,
)
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] YARA Forge update failed: %v\n", ts(), err)
return
}
if count == 0 {
return
}
fmt.Fprintf(os.Stderr, "[%s] YARA Forge update: %d rules (version %s)\n", ts(), count, newVersion)
forgeFile := filepath.Join(d.cfg.Signatures.RulesDir, fmt.Sprintf("yara-forge-%s.yar", d.cfg.Signatures.YaraForge.Tier))
if !d.settleForgeInstall(yaraScanner, forgeFile, count) {
return // don't store version
}
if db != nil {
_ = db.SetMetaString("forge_version_"+d.cfg.Signatures.YaraForge.Tier, newVersion)
}
}
// forgeReloader is the slice of the YARA backend the Forge settle step uses.
type forgeReloader interface {
Reload() error
RuleCount() int
}
// settleForgeInstall reloads the merged rules directory after a Forge tier
// was written and decides whether the tier stays. It goes when the merged
// reload fails (the tier compiled alone but conflicts with the shipped
// rules) or when the loaded count collapses; either way the file is
// removed, the backend reloads without it, and a Critical finding says so.
// Leaving a non-compiling tier on disk kept the live rules for now but made
// the next worker restart compile the same directory, fail, and run with
// zero rules until an operator deleted the file by hand. Returns true when
// the tier stays and its version may be recorded.
func (d *Daemon) settleForgeInstall(backend forgeReloader, forgeFile string, count int) bool {
reason := ""
if err := backend.Reload(); err != nil {
reason = fmt.Sprintf("the merged rules failed to compile (%v)", err)
} else if newCount := backend.RuleCount(); forgeRollbackNeeded(newCount, count) {
reason = fmt.Sprintf("%d rules downloaded but only %d loaded", count, newCount)
} else {
fmt.Fprintf(os.Stderr, "[%s] Reloaded %d YARA rules after Forge update\n", ts(), newCount)
return true
}
_ = os.Remove(forgeFile)
if err := backend.Reload(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] YARA reload after Forge rollback error: %v\n", ts(), err)
}
// Losing a ruleset this size is a coverage collapse, so it has to be
// alertable rather than a line on stderr that the journal rotates away.
d.emitReloadFinding(alert.Critical, "yara_forge_rollback", fmt.Sprintf(
"YARA Forge update rolled back: %s, so %s was removed. Scanning continues on the remaining rules.",
reason, forgeFile))
return false
}
func (d *Daemon) reloadSignatures() {
if scanner := signatures.Global(); scanner != nil {
if err := scanner.Reload(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] YAML rule reload error: %v\n", ts(), err)
} else {
fmt.Fprintf(os.Stderr, "[%s] Reloaded %d YAML rules (version %d)\n", ts(), scanner.RuleCount(), scanner.Version())
if d.webServer != nil {
d.webServer.SetSigCount(scanner.RuleCount())
}
}
}
yaraRules, yaraActive := 0, false
if yaraScanner := yara.Active(); yaraScanner != nil {
if err := yaraScanner.Reload(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] YARA rule reload error: %v\n", ts(), err)
} else {
fmt.Fprintf(os.Stderr, "[%s] Reloaded %d YARA rule file(s)\n", ts(), yaraScanner.RuleCount())
yaraActive = true
yaraRules = yaraScanner.RuleCount()
}
}
d.reportRealtimeRuleCoverage(yamlRuleCount(), yaraRules, yaraActive)
}
// The startup integrations are indirected so the observe-mode gate around them
// can be tested without a host to write to or a web server to reload.
var (
ensureAuditdRules = auditd.EnsureDeployed
deployHostConfigs = deployConfigs
reconcileModSecReload = (*checks.ModSecReloadReconciler).Reconcile
)
// applyStartupIntegrations refreshes the host-side files CSM owns: the auditd
// rules, the WHM plugin, the ModSecurity section and the deploy script. None
// of them has a switch of its own, so observe mode is what turns them off.
//
// Self-healing the auditd rules matters because package upgrades sometimes
// ship a new csm binary without re-running auditd.Deploy() (postinstall hooks
// differ across apt/dnf and across operator deploy automation), which leaves
// new rules -- including detection layers like csm_af_alg_socket -- silently
// inactive on the upgraded host. Errors are non-fatal: if auditd is absent or
// augenrules fails, the rest of CSM still runs.
func (d *Daemon) applyStartupIntegrations() {
if d.cfg.ObserveMode() {
csmlog.Info("observe mode: skipping host integration deploy (auditd rules, WHM plugin, ModSecurity section, deploy script)")
return
}
if redeployed, err := ensureAuditdRules(); err != nil {
csmlog.Warn("auditd rules ensure failed", "err", err)
} else if redeployed {
csmlog.Info("auditd rules redeployed (drift from embedded constant)")
}
deployHostConfigs()
if err := reconcileModSecReload(&d.modsecReload, d.cfg.ModSec.ReloadCommand); err != nil {
csmlog.Warn("CSM ModSecurity rule activation could not be confirmed", "err", err)
}
}
// deployConfigs writes embedded config files to their system locations on startup.
// Ensures WHM plugin CGI and ModSec rules stay current after binary upgrades.
//
// Every file written here is a system integration point consumed by a
// different process (WHM, Apache, nginx); the permissions intentionally
// allow the right external reader. Gosec G301/G306 warnings on this
// function are suppressed inline with the specific integration target.
func deployConfigs() {
// WHM plugin CGI - embedded in binary
if _, err := os.Stat("/usr/local/cpanel"); err == nil {
dst := "/usr/local/cpanel/whostmgr/docroot/cgi/addon_csm.cgi"
// #nosec G306 -- WHM CGI endpoint; 0755 is required so cPanel's
// webserver can execute it.
if err := os.WriteFile(dst, embeddedWHMCGI, 0755); err == nil {
csmlog.Info("WHM plugin CGI deployed", "path", dst)
}
// Write the AppConfig file, then register it with WHM.
// Writing the file alone does NOT make the plugin appear in the
// sidebar — WHM's AppConfig system maintains a registration
// database that is updated via `register_appconfig`. Skipping
// that step was a long-standing bug; the plugin file existed on
// disk but never showed up in the menu.
// #nosec G301 -- cPanel standard /var/cpanel/apps directory.
_ = os.MkdirAll("/var/cpanel/apps", 0755)
confPath := "/var/cpanel/apps/csm.conf"
// #nosec G306 -- WHM AppConfig; read by cPanel tooling, 0644 is convention.
if err := os.WriteFile(confPath, embeddedWHMConf, 0644); err != nil {
csmlog.Error("WHM AppConfig write failed", "path", confPath, "err", err)
} else if err := registerWHMPlugin(confPath); err != nil {
// Non-fatal: the conf is on disk, register_appconfig failure is
// logged so operators can fix it manually. Most common failure
// is register_appconfig not being in PATH (old cPanel versions).
csmlog.Warn("WHM plugin registration failed", "err", err)
} else {
csmlog.Info("WHM plugin registered with AppConfig")
}
}
// Deploy script (self-updating)
// #nosec G306 -- Shell script executed by operators and by the CSM
// upgrade path; needs to be executable, not private.
_ = os.WriteFile("/opt/csm/deploy.sh", embeddedDeployScript, 0755)
// ModSecurity virtual patches. modsec2.user.conf is shared with
// operator-maintained rules, so the embedded rules go through
// checks.MergeModSecUserConfSection: this startup deploy only ever
// creates or rewrites CSM's marker-delimited section and every byte
// outside it is preserved verbatim.
for _, dst := range []string{
"/etc/apache2/conf.d/modsec/modsec2.user.conf",
"/usr/local/apache/conf/modsec2.user.conf",
} {
if _, err := os.Stat(filepath.Dir(dst)); err == nil {
// #nosec G304 -- dst iterates the literal ModSecurity config paths above.
existing, readErr := os.ReadFile(dst)
if readErr != nil && !os.IsNotExist(readErr) {
// Present but unreadable: rewriting blind could destroy
// operator rules, so leave the file alone.
} else if merged, changed := checks.MergeModSecUserConfSection(existing, embeddedModSec); changed {
// #nosec G306 G703 -- dst iterates the literal ModSecurity config
// paths above; Apache reads this config from a different user.
_ = os.WriteFile(dst, merged, 0644)
}
overridesFile := filepath.Join(filepath.Dir(dst), "modsec2.csm-overrides.conf")
modsec.EnsureOverridesInclude(dst, overridesFile)
break
}
}
}
// registerWHMPlugin runs cPanel's register_appconfig helper to add the CSM
// plugin to the WHM sidebar. WHM maintains a cached registration database
// separate from the /var/cpanel/apps/ conf files; without running this
// helper, the plugin file exists on disk but the menu never shows it.
//
// Idempotent: re-running against an already-registered plugin just updates
// the entry. Non-fatal on failure — deployment continues and the operator
// can rerun manually.
func registerWHMPlugin(confPath string) error {
bin := "/usr/local/cpanel/bin/register_appconfig"
if _, err := os.Stat(bin); err != nil {
return fmt.Errorf("register_appconfig not found at %s: %w", bin, err)
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// #nosec G204 -- bin is the fixed cPanel path validated by os.Stat above;
// confPath was just written by deployConfigs from an embedded constant.
cmd := exec.CommandContext(ctx, bin, confPath)
out, err := cmd.CombinedOutput()
if err != nil {
return fmt.Errorf("%s %s: %w: %s", bin, confPath, err, strings.TrimSpace(string(out)))
}
return nil
}
// watchdogInterval keeps a safety margin under the interval systemd expects,
// and never pings faster than every ten seconds: a unit configured with a very
// short WatchdogSec would otherwise spend the daemon's time on keepalives.
func watchdogInterval(timeout time.Duration) time.Duration {
interval := timeout / 2
if interval < 10*time.Second {
interval = 10 * time.Second
}
return interval
}
// watchdogNotifier sends systemd watchdog keepalives on its own ticker.
// Runs at half the WatchdogSec interval so there's always margin.
// Completely independent of scan goroutines — never blocks.
func (d *Daemon) watchdogNotifier() {
defer d.wg.Done()
timeout, configured := sdnotify.WatchdogTimeout()
if !configured || !sdnotify.Enabled() {
return // watchdog not configured, or nothing to notify
}
interval := watchdogInterval(timeout)
csmlog.Info("systemd watchdog active", "interval", interval.String())
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
if _, err := sdnotify.Watchdog(); err != nil {
csmlog.Warn("systemd watchdog notify failed", "err", err)
}
}
}
}
func ts() string {
return time.Now().Format("2006-01-02 15:04:05")
}
func orUnknown(v string) string {
if v == "" {
return "unknown"
}
return v
}
func orNone(v string) string {
if v == "" {
return "none"
}
return v
}
// countAttachedWatchers returns how many watchers are currently attached
// (value == true). Used for the systemd one-line status string.
func countAttachedWatchers(statuses map[string]bool) int {
n := 0
for _, attached := range statuses {
if attached {
n++
}
}
return n
}
package daemon
import (
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
)
// maxDropperProbeAttempts bounds how many times an inconclusive probe (a
// permission or transient I/O failure, not a confirmed absence) is retried
// before the candidate is dropped. Dropping is detection-coverage loss, so
// the engine counts it, but retrying forever would pin a permanently
// unreadable path into every probe tick.
const maxDropperProbeAttempts = 5
// dropperFSProber resolves a tracked candidate against the live filesystem.
// The real implementation is platform-specific (statx birth time, quarantine
// ledger); tests inject a fake so the engine's orchestration is verifiable
// without a kernel.
type dropperFSProber interface {
probe(c dropperCandidate) dropperProbe
}
// dropperEmitFn delivers a finding. In production it is bound to
// FileMonitor.sendAlertWithPath; tests capture the arguments.
type dropperEmitFn func(sev alert.Severity, check, msg, details, path string)
const dropperCheckName = "self_deleting_dropper_realtime"
type dropperEngineConfig struct {
ttl time.Duration
selfPID int32
// ignorePath reports whether a path is covered by
// suppressions.ignore_paths. Nil means nothing is suppressed.
ignorePath func(string) bool
}
// dropperEngine owns the tracker and drives the observe -> probe -> hold ->
// flush lifecycle. admit is called from the analyzer worker pool (via the
// tracker's own locking); probeStep is called only from the single probe
// goroutine, so its attempt bookkeeping needs no lock.
type dropperEngine struct {
tr *dropperTracker
ttl time.Duration
selfPID int32
emit dropperEmitFn
attempts map[queuehealth.Ticket]int
// ignorePath mirrors the suppression every other content check already
// honours. Applied at admit so a suppressed path never consumes tracker
// capacity a real candidate could have used.
ignorePath func(string) bool
}
func newDropperEngine(cfg dropperEngineConfig) *dropperEngine {
return &dropperEngine{
tr: newDropperTracker(cfg.ttl),
ttl: cfg.ttl,
selfPID: cfg.selfPID,
attempts: make(map[queuehealth.Ticket]int),
ignorePath: cfg.ignorePath,
}
}
// admit records a candidate if it passes the freshness/type gate. Returns
// false when the candidate was rejected by the gate or dropped by the
// tracker capacity bound (the caller surfaces the latter as coverage loss).
func (e *dropperEngine) admit(c dropperCandidate) bool {
if e.ignorePath != nil && e.ignorePath(c.Path) {
return false
}
if !shouldTrackDropper(c, e.selfPID, e.ttl) {
return false
}
// Retain known executable probes until assessment: a late CREATE must
// merge with their completed CLOSE_WRITE instead of becoming a new pending
// write. Inert files still avoid consuming tracker capacity.
// ContentSuspicious wins: a realtime content or signature hit already
// found structure, and no later heuristic may demote that.
if !c.WritePending && !c.ContentSuspicious && !c.ContentMayExecute && !c.ContentUnsettled && dropperCandidateIsInert(c) {
return false
}
return e.tr.Observe(c)
}
// probeStep resolves every candidate whose TTL elapsed at probeNow, then
// flushes any held findings whose grace window closed at flushNow. Callers
// pass the same clock for both; the two parameters exist so tests can drive
// the grace window independently of the TTL.
func (e *dropperEngine) probeStep(probeNow time.Time, prober dropperFSProber, flushNow time.Time) {
for _, c := range e.tr.Due(probeNow) {
// A close-write can strengthen the filesystem identity while queued.
// The work ticket survives that change and owns its attempt budget.
key := c.ticket
if e.ignorePath != nil && e.ignorePath(c.Path) {
delete(e.attempts, key)
c.ticket.Finish(e.tr.now())
continue
}
verdict := assessDropper(c, prober.probe(c))
if verdict == dropperInconclusive {
if e.attempts[key]+1 >= maxDropperProbeAttempts {
delete(e.attempts, key)
c.ticket.Reject(e.tr.now())
continue
}
attempts := e.attempts[key] + 1
delete(e.attempts, key)
if retained, ok := e.tr.Retry(c); ok {
e.attempts[retained] = max(e.attempts[retained], attempts)
}
continue
}
delete(e.attempts, key)
c.ticket.Finish(e.tr.now())
e.tr.HoldGone(c, verdict, flushNow)
}
for _, f := range e.tr.FlushDue(flushNow) {
e.flushFinding(f)
}
}
func (e *dropperEngine) flushFinding(f dropperFinding) {
defer func() {
for _, item := range f.Items {
item.ticket.Finish(e.tr.now())
}
}()
items := make([]dropperGone, 0, len(f.Items))
for _, item := range f.Items {
if e.ignorePath == nil || !e.ignorePath(item.Cand.Path) {
items = append(items, item)
}
}
// Suppress before deciding burst severity. Otherwise excluded files
// can turn a remaining solitary dropper into a lower-severity burst.
for _, grouped := range groupDropperFindings(f.Docroot, items) {
e.emitFinding(grouped)
}
}
func (e *dropperEngine) emitFinding(f dropperFinding) {
if e.emit != nil {
sev, msg, details, path := dropperAlertParams(f)
e.emit(sev, dropperCheckName, msg, details, path)
}
}
package daemon
import (
"strings"
"github.com/pidginhost/csm/internal/checks"
)
func dropperCandidateIsInert(c dropperCandidate) bool {
// Executable files may be interpreted by a shell rather than PHP.
// PHP comments do not prove that such a script has no commands.
if c.Mode&0o111 != 0 {
return c.Size >= 0 && c.Size <= int64(len(c.Head)) && strings.Trim(string(c.Head), " \t\n") == ""
}
if dropperContentIsInert(c.Head, c.Size) {
return true
}
// A file whose first statement halts the interpreter carries data, not
// code, however large the tail is. Plugins keep state and WAF data in
// .php files of that shape and rewrite them constantly.
//
// Shift-based source encodings cannot reach this shape the way they reach
// a comment: the accepted bytes are the opening tag, PHP whitespace, one
// terminator keyword and a plain-ASCII literal with every encoding-shift
// byte rejected, so an ASCII-transparent encoding leaves them unchanged
// and a non-transparent one never matches the raw opening tag. A whole-file
// transport encoding is the exception: under a BASE64 source encoding this
// header is discarded as padding and an encoded tail becomes the program.
// Refuse the exemption when the tail could be that, which costs at most a
// plugin state file staying a candidate.
end, ok := checks.PHPTerminatesImmediatelyAt(c.Head)
return ok && !dropperTailCouldDecodeToSource(c.Head[end:])
}
// dropperTailCouldDecodeToSource reports whether the unreachable tail is made
// only of transport-encoding alphabet, which a source-encoding conversion
// could turn back into PHP. Real data files carry punctuation or text that no
// such alphabet contains.
func dropperTailCouldDecodeToSource(tail []byte) bool {
digits := 0
for _, b := range tail {
switch {
case b == ' ' || b == '\t' || b == '\n' || b == '\r':
case (b >= 'A' && b <= 'Z') || (b >= 'a' && b <= 'z') || (b >= '0' && b <= '9') ||
b == '+' || b == '/' || b == '=' || b == '-' || b == '_':
digits++
default:
return false
}
}
// Eight alphabet characters carry the six bytes of a "<?php " opener.
return digits >= 8
}
// dropperContentIsInert only exempts complete blank content. PHP can decode
// source before tokenization, including transfer encodings such as Base64
// and quoted-printable. Even ASCII text that resembles a comment can contain
// statements after that conversion. Without the interpreter's effective
// encoding settings, a PHP comment parser cannot prove those files inert.
func dropperContentIsInert(head []byte, size int64) bool {
return size >= 0 && size <= int64(len(head)) && strings.Trim(string(head), " \t\r\n") == ""
}
package daemon
// dropperUploadExecutionProbes are the exact scripts Really Simple Security
// copies into the uploads directory, requests over HTTP to learn whether PHP
// runs there, and deletes again. The plugin renamed itself once, so both
// shipped versions of the comment are listed.
//
// Only exact bytes qualify. The script is code, and no parser has to judge
// it: any added statement, encoding change or trailing byte stops the match,
// and a fixed byte string cannot carry a payload an attacker chose under any
// source encoding.
var dropperUploadExecutionProbes = []string{
"<?php\n/**\n * Test file for Really Simple SSL to check if uploads directory has code execution permissions\n *\n */\n\necho \"RSSSL CODE EXECUTION MARKER\";\n",
"<?php\n/**\n * Test file for Really Simple Security to check if uploads directory has code execution permissions\n *\n */\n\necho \"RSSSL CODE EXECUTION MARKER\";\n",
}
// dropperCandidateIsKnownProbe reports whether the snapshot is, in full, a
// plugin's server capability test script.
func dropperCandidateIsKnownProbe(c dropperCandidate) bool {
if c.Mode&0o111 != 0 || c.Size != int64(len(c.Head)) {
return false
}
for _, probe := range dropperUploadExecutionProbes {
if string(c.Head) == probe {
return true
}
}
return false
}
// dropperCandidateIsHarmless reports whether the snapshot cannot be a dropper
// payload: it has no executable statement, or it is a known test script.
func dropperCandidateIsHarmless(c dropperCandidate) bool {
return dropperCandidateIsInert(c) || dropperCandidateIsKnownProbe(c)
}
//go:build linux
package daemon
import (
"bytes"
"crypto/md5" // #nosec G501 -- wordpress.org publishes MD5 digests for core files
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"hash"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/wpcheck"
)
// dropperDigestMax bounds how many bytes the admission hash covers. A
// self-deleting dropper is small; this cap keeps a large legitimate PHP file
// from stalling the analyzer worker while still giving a full digest for
// realistic candidates. Files past the cap keep DigestKnown=false and rely on
// device/inode identity for rename matching.
const dropperDigestMax = 8 << 20
const dropperDigestChunk = 64 << 10
const dropperPHPHandlerCacheMax = 4096
const dropperPHPHandlerCacheTTL = 15 * time.Second
type dropperPHPHandlerCacheEntry struct {
generation uint64
overlay checks.PHPExecutionOverlay
loaded time.Time
}
// dropperProbeInterval derives the probe cadence from the tracking TTL so the
// loop reacts within a fraction of the window without busy-spinning.
func dropperProbeInterval(ttl time.Duration) time.Duration {
iv := ttl / 4
if iv < 5*time.Second {
return 5 * time.Second
}
if iv > time.Minute {
return time.Minute
}
return iv
}
func (fm *FileMonitor) initDropperDetector(cfg *config.Config) {
if cfg == nil || !cfg.Thresholds.DropperDetection {
return
}
ttl := time.Duration(cfg.Thresholds.DropperUnlinkTTLSec) * time.Second
if ttl <= 0 {
ttl = time.Duration(config.DefaultDropperUnlinkTTLSec) * time.Second
}
// #nosec G115 -- os.Getpid returns this process's PID, bounded by
// /proc/sys/kernel/pid_max (<= 2^22 on Linux), so it always fits in int32.
selfPID := int32(os.Getpid())
e := newDropperEngine(dropperEngineConfig{
ttl: ttl,
selfPID: selfPID,
ignorePath: func(p string) bool {
return checks.PathMatchesIgnore(p, fm.currentCfg().Suppressions.IgnorePaths)
},
})
e.emit = func(sev alert.Severity, check, msg, details, path string) {
fm.sendAlertWithPath(sev, check, msg, details, path, "")
}
fm.dropper = e
fm.dropperQuarantines = newDropperQuarantineLedger(
ttl + time.Duration(maxDropperProbeAttempts)*dropperProbeInterval(ttl) + dropperGraceWindow,
)
fm.dropperHandlerCache = make(map[string]dropperPHPHandlerCacheEntry)
fm.dropperDocroots.Store(checks.ResolveWebRoots(cfg))
}
func (fm *FileMonitor) currentDropperDocroots() []string {
if v, ok := fm.dropperDocroots.Load().([]string); ok {
return v
}
return nil
}
// isDropperInteresting admits paths that the normal content filter deliberately
// excludes but the dropper detector still needs: inherited .htaccess PHP
// handlers and regular files carrying an executable mode under a document
// root. It does not take ownership of fd.
func (fm *FileMonitor) isDropperInteresting(path string, fd int) (interesting, phpExecutable bool) {
if fm.dropper == nil {
return false, false
}
docroot := dropperDocrootFor(path, fm.currentDropperDocroots())
if docroot == "" {
return false, false
}
name := strings.ToLower(filepath.Base(path))
if contenttype.IsExecutablePHPName(name) {
return true, false
}
var st unix.Stat_t
if unix.Fstat(fd, &st) != nil || st.Mode&unix.S_IFMT != unix.S_IFREG {
return false, false
}
if st.Mode&0o111 != 0 {
return true, false
}
phpExecutable = fm.dropperPHPHandlerFor(docroot, filepath.Dir(path)).Executes(name)
return phpExecutable, phpExecutable
}
func (fm *FileMonitor) invalidateDropperPHPHandlerCache(path string) {
if fm.dropper != nil && strings.EqualFold(filepath.Base(path), ".htaccess") {
atomic.AddUint64(&fm.dropperHandlerGeneration, 1)
}
}
func (fm *FileMonitor) dropperPHPHandlerFor(docroot, dir string) checks.PHPExecutionOverlay {
key := docroot + "\x00" + dir
generation := atomic.LoadUint64(&fm.dropperHandlerGeneration)
fm.dropperHandlerMu.Lock()
entry, ok := fm.dropperHandlerCache[key]
fm.dropperHandlerMu.Unlock()
if ok && entry.generation == generation && time.Since(entry.loaded) < dropperPHPHandlerCacheTTL {
return entry.overlay
}
overlay := checks.ResolvePHPExecutionOverlay(docroot, dir)
// A .htaccess event that arrived while the filesystem snapshot was being
// rebuilt invalidates this result. Do not cache it; the next event retries.
if atomic.LoadUint64(&fm.dropperHandlerGeneration) != generation {
return checks.ResolvePHPExecutionOverlay(docroot, dir)
}
fm.dropperHandlerMu.Lock()
if fm.dropperHandlerCache == nil || len(fm.dropperHandlerCache) >= dropperPHPHandlerCacheMax {
fm.dropperHandlerCache = make(map[string]dropperPHPHandlerCacheEntry)
}
fm.dropperHandlerCache[key] = dropperPHPHandlerCacheEntry{
generation: generation,
overlay: overlay,
loaded: time.Now(),
}
fm.dropperHandlerMu.Unlock()
return overlay
}
// observeDropperCandidate snapshots an admitted write into the tracker. It
// runs on the analyzer worker while the event fd is still open. Fstat, statx,
// and Pread do not change the shared file offset, and this function never
// closes the event fd. The returned copy lets downstream content checks add a
// suspicious-content verdict without reading or hashing the file again.
func (fm *FileMonitor) observeDropperCandidate(event fileEvent, procInfo string) *dropperCandidate {
if fm.dropper == nil {
return nil
}
docroot := dropperDocrootFor(event.path, fm.currentDropperDocroots())
if docroot == "" {
return nil
}
var st unix.Stat_t
if err := unix.Fstat(event.fd, &st); err != nil {
return nil
}
c := dropperCandidate{
Path: event.path,
Docroot: docroot,
Observed: time.Now(),
Device: uint64(st.Dev),
Inode: st.Ino,
Size: st.Size,
UID: st.Uid,
Mode: uint32(st.Mode),
PID: event.pid,
ProcInfo: procInfo,
Created: event.mask&FAN_CREATE != 0,
WritePending: event.mask&FAN_CREATE != 0 && event.mask&FAN_CLOSE_WRITE == 0,
PHPExecutable: event.phpExecutable,
}
if birth, ok := statxBirthFromFD(event.fd); ok {
c.Birth = birth
c.BirthKnown = true
}
// Reject non-regular and non-executable names before allocating a head or
// hashing content. HTML, archives, and credential logs also reach the
// analyzer, but they are not dropper candidates.
name := strings.ToLower(filepath.Base(event.path))
if c.Mode&unix.S_IFMT != unix.S_IFREG ||
(!contenttype.IsExecutablePHPName(name) && !c.PHPExecutable && c.Mode&0o111 == 0) ||
c.PID == fm.dropper.selfPID {
return nil
}
trackFresh := shouldTrackDropper(c, fm.dropper.selfPID, fm.dropper.ttl)
// A known old birth time proves this CLOSE_WRITE is not the completion of
// a fresh create entry. Filesystems without birth time still need the full
// snapshot so Refresh can join a separate FAN_CREATE event to its close.
if !trackFresh && c.BirthKnown {
return nil
}
if parent, err := statDropperCandidateParent(c.Path, c.Device, c.Inode); err == nil {
c.Parent = parent
}
wpCopy := len(wpUpgradeCopyDestinations(c.Path, c.Docroot)) > 0
staged := wpUpgradeStagedPackageFile(c.Path, c.Docroot)
_, _, core := wpUpgradeCorePackageFile(c.Path, c.Docroot)
read, limit := readFromFd, dropperTrackedHeadMax
// Oversized staged files still need a head for the executable content
// check, even though they cannot supply a complete digest.
if wpCopy || (staged && st.Size <= dropperDigestMax) {
read, limit = readCompleteFromFd, dropperDigestMax
}
// The data proof and retained head must share the opening stat. Separate
// head/body snapshots leave a gap in which a writer can replace a payload.
body, size, stable := readDropperSnapshot(event.fd, st, limit, read)
c.Head, c.Size = body, size
if len(c.Head) > dropperTrackedHeadMax {
c.Head = bytes.Clone(c.Head[:dropperTrackedHeadMax])
}
c.ContentUnsettled = !stable
if staged && body != nil && stable && int64(len(body)) == c.Size {
// Hash the retained snapshot, including CREATE observations, rather
// than rereading an fd whose bytes may already have been replaced.
c.Digest, c.DigestKnown = sha256.Sum256(body), true
if core {
// #nosec G401 -- compared with the MD5 digests wordpress.org publishes
c.CoreMD5, c.CoreMD5Known = md5.Sum(body), true
}
}
// Copy exceptions must check even CREATE snapshots: a benign CLOSE_WRITE
// cannot erase an earlier payload. Blank snapshots carry no such evidence.
if wpCopy && (!stable || !dropperContentIsInert(c.Head, c.Size)) {
// Keep only the proof and hash, not a large translation body in each
// tracker entry. Both must describe the same complete snapshot.
if body != nil && stable && int64(len(body)) == c.Size {
c.Digest, c.DigestKnown = sha256.Sum256(body), true
if filepath.Base(c.Path) == "version-current.php" {
c.WPInstallData = checks.IsWPVersionDataBytesComplete(body, true)
if c.WPInstallData {
c.WPCoreRelease = fm.wpCoreReleaseOf(body)
}
} else {
c.WPInstallData = checks.IsWPTranslationCacheBytesComplete(body, true)
}
}
c.WPInstallUnsafe = !c.WPInstallData || c.Mode&0o111 != 0
} else if !wpCopy && !staged && !c.WritePending && atomicWriteRenameCandidate(c.Path) != "" {
c.Digest, c.DigestKnown = digestFromFD(event.fd, st.Size)
}
// FAN_CREATE and FAN_CLOSE_WRITE normally arrive as separate records. The
// create proves freshness on filesystems without STATX_BTIME; the close
// supplies the final bytes. Worker scheduling may deliver either one first,
// so Refresh is attempted for every non-create snapshot before admission.
if !c.Created && fm.dropper.tr.Refresh(c) {
return &c
}
if !trackFresh {
return nil
}
fm.dropper.admit(c)
// Even an inert snapshot must reach the content pass: a signature hit
// can override the admission gate without another filesystem read.
return &c
}
// readDropperSnapshot reads bytes together with proof that nothing
// changed the file while they were read. Size from before the read cannot
// prove completeness if another writer changed the file in the meantime, even
// when the retained head is empty.
//
// Even a ctime-only change is uncertain: a writer can restore bytes and
// mtime, leaving the same evidence as a harmless unlink or rename. A later
// quiet read cannot establish what executed in that interval.
func readDropperSnapshot(fd int, before unix.Stat_t, limit int, read func(int, int) []byte) ([]byte, int64, bool) {
body := read(fd, limit)
var after unix.Stat_t
stable := unix.Fstat(fd, &after) == nil && sameReadSnapshot(before, after) && before.Mode == after.Mode
return body, before.Size, stable
}
func (fm *FileMonitor) newDropperFSProbe() *dropperFSProbe {
return &dropperFSProbe{quarantines: fm.dropperQuarantines, coreChecksums: fm.wpCache}
}
// wpCoreReleaseOf names the release a version-probe data snapshot declares.
// Asking the cache now starts a missing checksum download in the background,
// so the deletion probe normally finds it cached; the analyzer never waits.
func (fm *FileMonitor) wpCoreReleaseOf(body []byte) *wpcheck.Verification {
if fm.wpCache == nil {
return nil
}
version, locale, err := wpcheck.ParseVersionContent(body)
if err != nil {
return nil
}
// #nosec G401 -- compared with the MD5 digests wordpress.org publishes
sum := md5.Sum(body)
v := &wpcheck.Verification{
Kind: wpcheck.KindCore, Version: version, Locale: locale,
Rel: "wp-includes/version.php", Digest: hex.EncodeToString(sum[:]), Staged: true,
}
fm.wpCache.Verify(*v)
return v
}
// dropperProbeLoop probes overdue candidates for deletion and flushes findings.
// It also refreshes the cached docroot set so account changes are picked up.
func (fm *FileMonitor) dropperProbeLoop() {
defer fm.wg.Done()
prober := fm.newDropperFSProbe()
ticker := time.NewTicker(dropperProbeInterval(fm.dropper.ttl))
defer ticker.Stop()
refresh := time.NewTicker(5 * time.Minute)
defer refresh.Stop()
for {
select {
case <-fm.stopCh:
return
case <-refresh.C:
fm.dropperDocroots.Store(checks.ResolveWebRoots(fm.currentCfg()))
case <-ticker.C:
now := time.Now()
fm.dropper.probeStep(now, prober, now)
fm.reportDropperOverflow()
}
}
}
// reportDropperOverflow surfaces only capacity losses that occurred since the
// previous report and limits a sustained storm to one warning per minute. The
// tracker counter is cumulative for diagnostics.
func (fm *FileMonitor) reportDropperOverflow() {
total := fm.dropper.tr.overflowDropped()
if total <= fm.dropperOverflowReported {
return
}
now := time.Now()
if !fm.lastDropperOverflowReport.IsZero() && now.Sub(fm.lastDropperOverflowReport) < time.Minute {
return
}
dropped := total - fm.dropperOverflowReported
fm.dropperOverflowReported = total
fm.lastDropperOverflowReport = now
fm.sendAlert(alert.Warning, "self_deleting_dropper_overflow",
"self-deleting-dropper tracker is full; some short-lived files were not tracked",
fmt.Sprintf("A create/delete storm exceeded the tracker capacity. %d candidate(s) were dropped since the prior report; those files are only covered by the next scheduled deep scan.", dropped))
}
// statxBirthFromFD returns the file birth time for an open fd when the
// filesystem records it (ext4, xfs, btrfs). Returns ok=false on filesystems
// without STATX_BTIME so the caller falls back to the create-event signal.
func statxBirthFromFD(fd int) (time.Time, bool) {
var stx unix.Statx_t
if err := unix.Statx(fd, "", unix.AT_EMPTY_PATH|unix.AT_SYMLINK_NOFOLLOW, unix.STATX_BTIME, &stx); err != nil {
return time.Time{}, false
}
if stx.Mask&unix.STATX_BTIME == 0 {
return time.Time{}, false
}
return time.Unix(stx.Btime.Sec, int64(stx.Btime.Nsec)), true
}
// digestFromFD hashes up to dropperDigestMax bytes of the open fd. Files larger
// than the cap return ok=false; rename matching then relies on device/inode
// identity, which still covers rename(2) within a filesystem.
func digestFromFD(fd int, size int64) ([32]byte, bool) {
if size < 0 || size > dropperDigestMax {
return [32]byte{}, false
}
h := sha256.New()
if !hashFDRange(h, fd, size) {
return [32]byte{}, false
}
var after unix.Stat_t
if err := unix.Fstat(fd, &after); err != nil || after.Size != size {
return [32]byte{}, false
}
var sum [sha256.Size]byte
copy(sum[:], h.Sum(nil))
return sum, true
}
func hashFDRange(h hash.Hash, fd int, size int64) bool {
if size == 0 {
return true
}
bufSize := int64(dropperDigestChunk)
if size < bufSize {
bufSize = size
}
buf := make([]byte, int(bufSize))
for offset := int64(0); offset < size; {
want := size - offset
if want > int64(len(buf)) {
want = int64(len(buf))
}
n, err := unix.Pread(fd, buf[:int(want)], offset)
if n > 0 {
_, _ = h.Write(buf[:n])
offset += int64(n)
}
if err != nil && !errors.Is(err, unix.EINTR) {
return false
}
if n == 0 && err == nil {
return false
}
}
return true
}
type dropperPathState struct {
file dropperFileState
mode uint32
}
// statPathToFileState opens path first, then derives identity, birth time,
// size, and an optional digest from that one fd. O_PATH is a fallback for
// symlinks or unreadable objects; it still provides stable identity without
// pretending a digest was available.
func statPathToFileState(path string, includeDigest bool) (dropperPathState, error) {
fd, err := unix.Open(path, unix.O_RDONLY|unix.O_NONBLOCK|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0)
readable := err == nil
if err != nil {
fd, err = unix.Open(path, unix.O_PATH|unix.O_CLOEXEC|unix.O_NOFOLLOW, 0)
if err != nil {
return dropperPathState{}, err
}
}
defer func() { _ = unix.Close(fd) }()
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err != nil {
return dropperPathState{}, err
}
state := dropperPathState{
file: dropperFileState{
Path: path, Device: uint64(st.Dev), Inode: st.Ino, Size: st.Size,
IsRegular: st.Mode&unix.S_IFMT == unix.S_IFREG,
},
mode: uint32(st.Mode),
}
if birth, ok := statxBirthFromFD(fd); ok {
state.file.Birth = birth
state.file.BirthKnown = true
}
if includeDigest && readable && st.Mode&unix.S_IFMT == unix.S_IFREG {
if digest, ok := digestFromFD(fd, st.Size); ok {
state.file.Digest = digest
state.file.DigestKnown = true
}
}
return state, nil
}
// openDropperDirNoSymlinks opens an absolute directory path one component at
// a time. Refusing symlinks anywhere in the chain is important: reopening a
// path through a retargeted ancestor could otherwise make an unchanged parent
// look replaced and turn attacker-controlled path churn into demotion evidence.
func openDropperDirNoSymlinks(path string) (int, error) {
if !filepath.IsAbs(path) || filepath.Clean(path) != path {
return -1, unix.EINVAL
}
flags := unix.O_PATH | unix.O_DIRECTORY | unix.O_CLOEXEC | unix.O_NOFOLLOW
fd, err := unix.Open(string(filepath.Separator), flags, 0)
if err != nil {
return -1, err
}
for _, component := range strings.Split(strings.TrimPrefix(path, string(filepath.Separator)), string(filepath.Separator)) {
if component == "" {
continue
}
next, openErr := unix.Openat(fd, component, flags, 0)
_ = unix.Close(fd)
if openErr != nil {
return -1, openErr
}
fd = next
}
return fd, nil
}
// statDropperParent snapshots a real parent directory without following any
// symlink in its path. Symlinked parents deliberately provide no removal
// evidence: a dangling or retargeted link does not prove that the directory
// which contained the event fd was removed.
func statDropperParent(path string) (dropperParentIdentity, error) {
fd, err := openDropperDirNoSymlinks(path)
if err != nil {
return dropperParentIdentity{}, err
}
defer func() { _ = unix.Close(fd) }()
return statDropperParentFD(fd)
}
// statDropperCandidateParent also proves that the candidate fd identity is
// still the entry in this parent. If the path was replaced while the event was
// being admitted, the unrelated successor directory cannot later supply
// removal evidence for the original candidate.
func statDropperCandidateParent(path string, device, inode uint64) (dropperParentIdentity, error) {
fd, err := openDropperDirNoSymlinks(filepath.Dir(path))
if err != nil {
return dropperParentIdentity{}, err
}
defer func() { _ = unix.Close(fd) }()
var child unix.Stat_t
if err := unix.Fstatat(fd, filepath.Base(path), &child, unix.AT_SYMLINK_NOFOLLOW); err != nil {
return dropperParentIdentity{}, err
}
if uint64(child.Dev) != device || child.Ino != inode {
return dropperParentIdentity{}, unix.ESTALE
}
return statDropperParentFD(fd)
}
func statDropperParentFD(fd int) (dropperParentIdentity, error) {
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err != nil {
return dropperParentIdentity{}, err
}
if st.Mode&unix.S_IFMT != unix.S_IFDIR {
return dropperParentIdentity{}, unix.ENOTDIR
}
identity := dropperParentIdentity{Device: uint64(st.Dev), Inode: st.Ino}
if birth, ok := statxBirthFromFD(fd); ok {
identity.BirthKnown = true
identity.BirthNanos = birth.UnixNano()
}
return identity, nil
}
const dropperQuarantineLedgerMax = 4096
type dropperQuarantineRecord struct {
state dropperFileState
expires time.Time
}
type dropperQuarantineLedger struct {
mu sync.Mutex
keepFor time.Duration
count int
byPath map[string][]dropperQuarantineRecord
}
func newDropperQuarantineLedger(keepFor time.Duration) *dropperQuarantineLedger {
return &dropperQuarantineLedger{
keepFor: keepFor,
byPath: make(map[string][]dropperQuarantineRecord),
}
}
func (l *dropperQuarantineLedger) record(originalPath string, state dropperFileState, now time.Time) {
if l == nil {
return
}
l.mu.Lock()
defer l.mu.Unlock()
l.pruneLocked(now)
if l.count >= dropperQuarantineLedgerMax {
return
}
l.byPath[originalPath] = append(l.byPath[originalPath], dropperQuarantineRecord{
state: state, expires: now.Add(l.keepFor),
})
l.count++
}
func (l *dropperQuarantineLedger) matched(c dropperCandidate, now time.Time) bool {
if l == nil {
return false
}
l.mu.Lock()
defer l.mu.Unlock()
l.pruneLocked(now)
records := l.byPath[c.Path]
for i, record := range records {
// Require the same inode generation. Digest-only matching would let a
// later identical drop at the same path consume an older quarantine
// record and evade detection. Cross-filesystem quarantine may therefore
// produce a duplicate alert, which is safer than suppressing a replay.
if !dropperSameIdentity(c, record.state) {
continue
}
records = append(records[:i], records[i+1:]...)
if len(records) == 0 {
delete(l.byPath, c.Path)
} else {
l.byPath[c.Path] = records
}
l.count--
return true
}
return false
}
func (l *dropperQuarantineLedger) pruneLocked(now time.Time) {
for path, records := range l.byPath {
keep := records[:0]
for _, record := range records {
if now.Before(record.expires) {
keep = append(keep, record)
} else {
l.count--
}
}
if len(keep) == 0 {
delete(l.byPath, path)
} else {
l.byPath[path] = keep
}
}
}
func (fm *FileMonitor) recordDropperQuarantine(originalPath, quarantinePath string) {
if fm.dropperQuarantines == nil {
return
}
state, err := statPathToFileState(quarantinePath, false)
if err != nil || state.mode&unix.S_IFMT != unix.S_IFREG {
return
}
state.file.Path = originalPath
fm.dropperQuarantines.record(originalPath, state.file, time.Now())
}
//go:build linux
package daemon
import (
"encoding/hex"
"errors"
"path/filepath"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/wpcheck"
)
// dropperFSProbe resolves a tracked candidate against the live filesystem for
// the probe loop. It distinguishes a confirmed deletion (ENOENT) from a
// permission or transient I/O failure so the engine can requeue the latter
// rather than reporting a phantom self-delete.
type dropperFSProbe struct {
quarantines *dropperQuarantineLedger
// coreChecksums compares a version probe with the official file of the
// release it declares, and a core package file with the release now
// installed. Nil leaves every such candidate unverified.
coreChecksums interface {
Describe(path string) wpcheck.Verification
Verify(wpcheck.Verification) wpcheck.Verdict
}
}
func (p dropperFSProbe) probe(c dropperCandidate) dropperProbe {
state, err := statPathToFileState(c.Path, false)
switch {
case err == nil:
// Every identity field came from one open fd. A matching identity means
// the tracked file survived; a different one means it was replaced.
return dropperProbe{Conclusive: true, AtPath: &state.file}
case !errors.Is(err, unix.ENOENT):
// Permission or I/O error: cannot prove deletion. Inconclusive.
return dropperProbe{Conclusive: false}
}
// Confirmed absent. Attribute it to an install rename or a removed docroot
// before treating it as a self-delete.
result := dropperProbe{Conclusive: true}
if p.quarantines.matched(c, time.Now()) {
result.QuarantineMatched = true
return result
}
if target, ts, ok, renameErr := dropperFindRenameTarget(c); renameErr != nil {
return dropperProbe{Conclusive: false}
} else if ok {
result.RenamedTo = target
result.RenameTarget = &ts
}
if c.WPCoreRelease != nil && p.coreChecksums != nil {
result.OfficialWPCoreFile = p.coreChecksums.Verify(*c.WPCoreRelease) == wpcheck.VerdictVerified
}
result.OfficialWPCorePackageFile = p.officialCorePackageFile(c)
var dst unix.Stat_t
if derr := unix.Stat(c.Docroot, &dst); derr != nil && errors.Is(derr, unix.ENOENT) {
result.DocrootRemoved = true
}
if c.Parent.known() {
if current, perr := statDropperParent(filepath.Dir(c.Path)); perr == nil {
result.ParentRemoved = dropperParentChanged(c.Parent, current)
} else if errors.Is(perr, unix.ENOENT) {
result.ParentRemoved = true
}
}
return result
}
// officialCorePackageFile compares a vanished file of an unpacked core
// release with the release installed at its WordPress root. The staged tree,
// and with it the staged version header, is gone by the time of the probe.
// A completed update has installed that release, so its header names the
// manifest to check. An aborted one leaves the old release, whose manifest
// does not match the new bytes, and the candidate stays reported.
func (p dropperFSProbe) officialCorePackageFile(c dropperCandidate) bool {
if p.coreChecksums == nil || !c.CoreMD5Known {
return false
}
wpRoot, rel, ok := wpUpgradeCorePackageFile(c.Path, c.Docroot)
if !ok {
return false
}
v := p.coreChecksums.Describe(filepath.Join(wpRoot, "wp-includes", "version.php"))
if v.Kind != wpcheck.KindCore || v.Root != wpRoot || v.Version == "" {
return false
}
v.Rel, v.Digest = rel, hex.EncodeToString(c.CoreMD5[:])
return p.coreChecksums.Verify(v) == wpcheck.VerdictVerified
}
// dropperFindRenameTarget snapshots the install destinations WordPress and the
// atomic-write helper may move or copy a staged file to. A matching destination
// wins; otherwise the first regular destination is returned as replacement
// evidence.
func dropperFindRenameTarget(c dropperCandidate) (string, dropperFileState, bool, error) {
return dropperFindRenameTargetWithStat(c, statPathToFileState)
}
func dropperFindRenameTargetWithStat(c dropperCandidate, stat func(string, bool) (dropperPathState, error)) (string, dropperFileState, bool, error) {
targets := wpUpgradeInstallDestinations(c.Path, c.Docroot)
if atomic := atomicWriteRenameCandidate(c.Path); atomic != "" {
targets = append(targets, atomic)
}
var firstTarget string
var firstState dropperFileState
var transientErr error
for _, target := range targets {
if !dropperRenameTargetAllowed(c, target) {
continue
}
state, err := stat(target, false)
// An unreachable destination cannot prove a benign move. Treating it
// as transient would let a broken install path exhaust probe retries.
if dropperInstallPathUnreachable(err) {
continue
}
if err != nil {
transientErr = err
continue
}
if state.mode&unix.S_IFMT != unix.S_IFREG {
continue
}
if dropperSameIdentity(c, state.file) {
return target, state.file, true, nil
}
if c.DigestKnown && c.Size == state.file.Size {
state, err = stat(target, true)
if dropperInstallPathUnreachable(err) {
continue
}
if err != nil {
transientErr = err
continue
}
if dropperRenameMatch(c, state.file) {
return target, state.file, true, nil
}
}
if firstTarget == "" {
firstTarget, firstState = target, state.file
}
}
if transientErr != nil {
return "", dropperFileState{}, false, transientErr
}
if firstTarget != "" {
return firstTarget, firstState, true, nil
}
return "", dropperFileState{}, false, nil
}
func dropperInstallPathUnreachable(err error) bool {
return errors.Is(err, unix.ENOENT) || errors.Is(err, unix.ENOTDIR) || errors.Is(err, unix.ELOOP)
}
package daemon
import (
"bytes"
"fmt"
"path/filepath"
"regexp"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/wpcheck"
)
// dropperCandidate captures the fstat/read state of a file at close-write
// time. The fanotify fd is not retained: everything needed for the later
// TTL probe is copied here while the fd is still open, so probe decisions
// never race the attacker deleting or swapping the file.
type dropperCandidate struct {
Path string
Docroot string
Observed time.Time
Birth time.Time
BirthKnown bool
Created bool
Device uint64
Inode uint64
Size int64
UID uint32
Mode uint32
PID int32
ProcInfo string
// PHPExecutable is set by the Linux analyzer when an inherited
// .htaccess handler makes a non-standard extension executable as PHP.
PHPExecutable bool
// ContentSuspicious prevents FP heuristics from demoting a file whose
// realtime content/signature pass already found malicious structure.
ContentSuspicious bool
// A create event can precede the writer's first bytes. Retain its
// freshness evidence until a close-write supplies the payload.
WritePending bool
// Sticky across refreshes: truncating a previously executable snapshot
// must not turn its later deletion into a harmless empty guard.
ContentMayExecute bool
// ContentUnsettled is sticky: a later harmless snapshot cannot rule out
// code that ran while an earlier read raced a writer.
ContentUnsettled bool
Digest [32]byte
DigestKnown bool
// CoreMD5 is the same snapshot hashed the way wordpress.org publishes
// core checksums. Only files inside an unpacked core release carry it.
CoreMD5 [16]byte
CoreMD5Known bool
// ContentRewritten is sticky: a file whose written content changed
// cannot prove from its last bytes what an earlier version ran.
ContentRewritten bool
// The first nonempty snapshot of a staged package file, kept across empty
// writes and delayed analyzer verdicts: its digest, or that its bytes could
// not be read whole. The latest snapshot alone cannot prove its history.
stagedHistorySet bool
stagedHistoryKnown bool
stagedHistoryDigest [32]byte
stagedHistoryObserved time.Time
// Observed is the TTL origin; snapshotObserved orders the metadata even
// after a merge has moved Observed back to the earliest event.
snapshotObserved time.Time
// WPInstallData proves the complete, stable snapshot used for Digest was
// a translation return literal or version assignments, with no payload.
WPInstallData bool
// Sticky across rewrites: data copied later cannot erase an earlier
// nonempty snapshot that carried code or whose full content was unknown.
WPInstallUnsafe bool
// WPCoreRelease names the WordPress release a version-probe data snapshot
// declares, with that snapshot's digest in the form wordpress.org
// publishes. It belongs to the same snapshot as WPInstallData.
WPCoreRelease *wpcheck.Verification
Head []byte
// Parent identifies the real, non-symlink directory that contained the
// candidate while its event fd was open. The later probe uses this stable
// identity instead of inferring directory removal from two path stats.
Parent dropperParentIdentity
ticket queuehealth.Ticket
}
// dropperParentIdentity is deliberately compact because one is retained for
// every tracked candidate. BirthNanos disambiguates inode reuse when statx
// exposes it; without birth time, a reused identity is treated as unchanged
// and cannot earn a false-positive demotion.
type dropperParentIdentity struct {
Device uint64
Inode uint64
BirthNanos int64
BirthKnown bool
Conflicted bool
}
func (i dropperParentIdentity) known() bool {
return !i.Conflicted && i.Device != 0 && i.Inode != 0
}
func mergeDropperParentIdentity(a, b dropperParentIdentity) dropperParentIdentity {
switch {
case a.Conflicted || b.Conflicted:
return dropperParentIdentity{Conflicted: true}
case !a.known():
return b
case !b.known():
return a
case a.Device != b.Device || a.Inode != b.Inode:
return dropperParentIdentity{Conflicted: true}
case a.BirthKnown && b.BirthKnown && a.BirthNanos != b.BirthNanos:
return dropperParentIdentity{Conflicted: true}
case b.BirthKnown:
return b
default:
return a
}
}
func dropperParentChanged(observed, current dropperParentIdentity) bool {
if !observed.known() || !current.known() {
return false
}
if observed.Device != current.Device || observed.Inode != current.Inode {
return true
}
return observed.BirthKnown && current.BirthKnown && observed.BirthNanos != current.BirthNanos
}
// shouldTrackDropper reports whether a close-write event is a freshly
// created PHP or executable file inside a web document root, i.e. a
// candidate for self-deleting-dropper tracking. Modifications of
// pre-existing files are excluded via either a FAN_CREATE observation or a
// recent statx birth time. The explicit create bit keeps the detector useful
// on filesystems that do not expose STATX_BTIME; the Linux event path must
// preserve FAN_CREATE rather than throwing the event mask away.
func shouldTrackDropper(c dropperCandidate, selfPID int32, freshFor time.Duration) bool {
if c.Docroot == "" {
return false
}
if c.PID == selfPID {
return false
}
if c.Mode&unixSIFMT != unixSIFREG {
return false
}
name := strings.ToLower(filepath.Base(c.Path))
if !contenttype.IsExecutablePHPName(name) && !c.PHPExecutable && c.Mode&0o111 == 0 {
return false
}
if !c.Created {
age := c.Observed.Sub(c.Birth)
if !c.BirthKnown || age < 0 || age > freshFor {
return false
}
}
return true
}
// S_IFMT constants mirrored from the unix package so this file stays free
// of //go:build linux and the decision logic remains testable on any OS.
const (
unixSIFMT = 0o170000
unixSIFREG = 0o100000
)
// dropperMaxTracked bounds the tracker map. A cPanel package restore or a
// WP Toolkit site clone can close-write tens of thousands of PHP files in
// seconds, and every one stays tracked for the whole unlink TTL; entries
// beyond the cap are dropped (and counted) rather than evicting older
// candidates, because the oldest entries are the ones closest to their probe
// and losing them would blind the detector exactly when a bulk write storm
// provides cover.
const dropperMaxTracked = 16384
// Keep each waiting or detached batch's head-byte budget at 16 MiB.
// Candidate/map/path metadata is additional bounded memory and grows with the
// entry cap; this constant only accounts for copied content. Representative
// Twig and Smarty headers place all required markers inside this window.
const (
dropperTrackedHeadBudget = 16 << 20
dropperTrackedHeadMax = dropperTrackedHeadBudget / dropperMaxTracked
)
type dropperCandidateKey struct {
path string
device uint64
inode uint64
birthNanos int64
birthKnown bool
}
func candidateKey(c dropperCandidate) dropperCandidateKey {
key := dropperCandidateKey{
path: c.Path,
device: c.Device,
inode: c.Inode,
birthKnown: c.BirthKnown,
}
if c.BirthKnown {
key.birthNanos = c.Birth.UnixNano()
}
return key
}
func ownDropperCandidate(c dropperCandidate) dropperCandidate {
if c.snapshotObserved.Before(c.Observed) {
c.snapshotObserved = c.Observed
}
if c.Size != 0 && !c.stagedHistorySet && wpUpgradeStagedPackageFile(c.Path, c.Docroot) {
c.stagedHistorySet = true
c.stagedHistoryKnown, c.stagedHistoryDigest = c.DigestKnown, c.Digest
c.stagedHistoryObserved = c.snapshotObserved
}
// Torn bytes that already look like code are evidence, not noise.
c.ContentMayExecute = c.ContentMayExecute || !dropperCandidateIsHarmless(c)
if len(c.Head) > dropperTrackedHeadMax {
c.Head = c.Head[:dropperTrackedHeadMax]
}
c.Head = bytes.Clone(c.Head)
return c
}
// dropperStagedHistoryDiffers reports whether two snapshots of a staged
// package file may hold different content. Bytes that could not be read whole
// only count once a second snapshot exists to compare them with; a single
// unreadable snapshot moved into place is still one write.
func dropperStagedHistoryDiffers(a, b dropperCandidate) bool {
if !a.stagedHistorySet || !b.stagedHistorySet {
return false
}
if !a.stagedHistoryKnown || !b.stagedHistoryKnown {
// Analyzer verdicts and detached probes can replay the same unreadable
// snapshot after newer empty writes. Compare the history's own time,
// not the latest metadata time or the merged TTL origin.
return !a.stagedHistoryObserved.Equal(b.stagedHistoryObserved)
}
return a.stagedHistoryDigest != b.stagedHistoryDigest
}
func mergeDropperCandidate(prev, next dropperCandidate) dropperCandidate {
merged := next
if next.snapshotObserved.Before(prev.snapshotObserved) {
merged = prev
}
merged.Observed = prev.Observed
if next.Observed.Before(prev.Observed) {
merged.Observed = next.Observed
}
merged.Created = prev.Created || next.Created
merged.PHPExecutable = prev.PHPExecutable || next.PHPExecutable
merged.ContentSuspicious = prev.ContentSuspicious || next.ContentSuspicious
merged.ContentMayExecute = prev.ContentMayExecute || next.ContentMayExecute
merged.ContentUnsettled = prev.ContentUnsettled || next.ContentUnsettled
merged.WPInstallUnsafe = prev.WPInstallUnsafe || next.WPInstallUnsafe
merged.ContentRewritten = prev.ContentRewritten || next.ContentRewritten || dropperStagedHistoryDiffers(prev, next)
history := next
if prev.stagedHistorySet {
history = prev
}
merged.stagedHistorySet = history.stagedHistorySet
merged.stagedHistoryKnown, merged.stagedHistoryDigest = history.stagedHistoryKnown, history.stagedHistoryDigest
merged.stagedHistoryObserved = history.stagedHistoryObserved
// CREATE may reach an analyzer after CLOSE_WRITE for the same inode.
merged.WritePending = prev.WritePending && next.WritePending
merged.Parent = mergeDropperParentIdentity(prev.Parent, next.Parent)
if !merged.BirthKnown {
switch {
case prev.BirthKnown:
merged.Birth = prev.Birth
merged.BirthKnown = true
case next.BirthKnown:
merged.Birth = next.Birth
merged.BirthKnown = true
}
}
merged.ticket = prev.ticket
return merged
}
// dropperTracker holds candidates between their close-write observation and
// the TTL probe. All methods are safe for concurrent use by the analyzer
// workers and the probe loop.
type dropperTracker struct {
mu sync.Mutex
ttl time.Duration
maxTracked int
entries map[dropperCandidateKey]dropperCandidate
pending []dropperGone
overflow uint64
now func() time.Time
healthOnce sync.Once
health *queuehealth.Tracker
heldHealth *queuehealth.Tracker
}
func newDropperTracker(ttl time.Duration) *dropperTracker {
return &dropperTracker{
ttl: ttl,
maxTracked: dropperMaxTracked,
entries: make(map[dropperCandidateKey]dropperCandidate),
now: time.Now,
}
}
func (t *dropperTracker) initQueueHealth() {
t.healthOnce.Do(func() {
t.health = queuehealth.New(t.maxTracked, time.Minute)
t.heldHealth = queuehealth.New(dropperMaxTracked, time.Minute)
})
}
func (t *dropperTracker) queueStatuses(now time.Time) (queuehealth.Status, queuehealth.Status) {
t.initQueueHealth()
return t.health.Snapshot(now), t.heldHealth.Snapshot(now)
}
// Observe records a candidate. Re-observing the same file identity keeps the
// earliest event time (so rewrites cannot postpone the probe) and the newest
// metadata snapshot. A replacement inode at the same path is a separate
// candidate; otherwise an attacker could overwrite a vanished drop with a
// benign survivor before the probe. The return value reports whether the
// candidate was retained rather than rejected by the capacity bound. A false
// result is detection coverage loss and the Linux wiring must surface it as
// a metric and operator-facing warning.
func (t *dropperTracker) Observe(c dropperCandidate) bool {
t.initQueueHealth()
c = ownDropperCandidate(c)
key := candidateKey(c)
t.mu.Lock()
defer t.mu.Unlock()
if prev, ok := t.entries[key]; ok {
t.entries[key] = mergeDropperCandidate(prev, c)
prev.ticket.RetainQueuedAt(c.Observed.Add(t.ttl))
return true
}
if len(t.entries) >= t.maxTracked {
t.overflow++
t.health.Lose(t.now(), 1)
return false
}
c.ticket = t.health.BeginAt(c.Observed.Add(t.ttl), t.now())
t.entries[key] = c
return true
}
// Retry returns a detached probe to the waiting set. A new observation may
// already occupy its identity or the available slot, so this transfer must
// share the admission lock with Observe.
func (t *dropperTracker) Retry(c dropperCandidate) (queuehealth.Ticket, bool) {
t.mu.Lock()
defer t.mu.Unlock()
key := candidateKey(c)
now := t.now()
if waiting, ok := t.entries[key]; ok {
waiting.ticket.MergeRunning(c.ticket, now)
t.entries[key] = mergeDropperCandidate(waiting, c)
return waiting.ticket, true
}
if len(t.entries) >= t.maxTracked {
t.overflow++
c.ticket.Reject(now)
return queuehealth.Ticket{}, false
}
c.ticket.Requeue(now)
t.entries[key] = c
return c.ticket, true
}
// Refresh updates a previously admitted candidate without creating a new
// entry. The Linux wiring uses this for a CLOSE_WRITE that follows a separate
// FAN_CREATE event: the create proves freshness, while the close supplies the
// final size, digest, head, and content verdict. CLOSE_WRITE handlers should
// call Refresh first, then use shouldTrackDropper plus Observe only when no
// prior create entry matched. A birth-time availability change between the
// two events is allowed only when path, device, and inode still match.
func (t *dropperTracker) Refresh(c dropperCandidate) bool {
c = ownDropperCandidate(c)
key := candidateKey(c)
t.mu.Lock()
defer t.mu.Unlock()
if prev, ok := t.entries[key]; ok {
t.entries[key] = mergeDropperCandidate(prev, c)
prev.ticket.RetainQueuedAt(c.Observed.Add(t.ttl))
return true
}
if c.Inode == 0 {
return false
}
for prevKey, prev := range t.entries {
if prev.Path != c.Path || prev.Device != c.Device || prev.Inode != c.Inode {
continue
}
// Exact known/known and unknown/unknown identities were handled by
// the direct key lookup. Only strengthen unknown -> known here;
// weakening a known identity could merge an inode-reuse generation.
if prev.BirthKnown || !c.BirthKnown {
continue
}
merged := mergeDropperCandidate(prev, c)
prev.ticket.RetainQueuedAt(c.Observed.Add(t.ttl))
delete(t.entries, prevKey)
mergedKey := candidateKey(merged)
t.entries[mergedKey] = merged
return true
}
return false
}
// Due removes and returns every candidate whose TTL has elapsed at now.
func (t *dropperTracker) Due(now time.Time) []dropperCandidate {
t.mu.Lock()
defer t.mu.Unlock()
started := t.now()
var due []dropperCandidate
for key, c := range t.entries {
if now.Sub(c.Observed) >= t.ttl {
c.ticket.Start(started)
due = append(due, c)
delete(t.entries, key)
}
}
return due
}
func (t *dropperTracker) trackedCount() int {
t.mu.Lock()
defer t.mu.Unlock()
return len(t.entries)
}
func (t *dropperTracker) overflowDropped() uint64 {
t.mu.Lock()
defer t.mu.Unlock()
return t.overflow
}
// discardPending runs after both the probe loop and analyzer workers join.
// Analyzer work finishing during shutdown can still admit fresh candidates.
func (t *dropperTracker) discardPending(now time.Time) {
t.mu.Lock()
defer t.mu.Unlock()
for _, c := range t.entries {
c.ticket.Reject(now)
}
clear(t.entries)
for _, g := range t.pending {
g.ticket.Reject(now)
}
t.pending = nil
}
// dropperFileState is the identity and content evidence captured when the
// probe opens a path. All fields must come from the same open fd. Device plus
// inode handles rename(2); birth time guards against inode reuse; a full
// digest handles copy-delete moves across filesystems. Head bytes are
// deliberately not identity evidence.
type dropperFileState struct {
Path string
Device uint64
Inode uint64
Size int64
Birth time.Time
BirthKnown bool
Digest [32]byte
DigestKnown bool
// IsRegular distinguishes a file that took over the path from a directory
// or symlink left behind there. Only a regular file can be the result of
// an atomic write.
IsRegular bool
}
// dropperProbe is what the TTL probe learned about a candidate. AtPath and
// RenameTarget carry enough evidence for this platform-free core to validate
// identity. A bare path-exists or destination-exists boolean would let a
// replacement file hide the vanished inode.
type dropperProbe struct {
// Conclusive is set only after the probe distinguished absence from a
// permission, I/O, or other transient failure.
Conclusive bool
AtPath *dropperFileState
// DocrootRemoved is true only for a confirmed ENOENT on the document
// root, not for permission or transient I/O failures.
DocrootRemoved bool
// ParentRemoved is true only when a snapshotted, non-symlink parent is now
// absent or a different directory identity. A file whose whole directory
// went away was not singled out for deletion.
ParentRemoved bool
RenamedTo string
RenameTarget *dropperFileState
// QuarantineMatched requires an exact ledger identity/fingerprint match,
// not merely a prior quarantine entry for the same path.
QuarantineMatched bool
// OfficialWPCoreFile reports that the candidate's WPCoreRelease digest
// equals the wordpress.org checksum for that release. Checksums that are
// not cached yet leave it false.
OfficialWPCoreFile bool
// OfficialWPCorePackageFile reports that the candidate, a file of an
// unpacked core release, is byte for byte the file of that path in the
// release installed at the WordPress root.
OfficialWPCorePackageFile bool
}
type dropperVerdict int
const (
dropperBenign dropperVerdict = iota
dropperInconclusive
dropperDemotedTemplate
dropperDemotedAtomicWrite
dropperDemotedWPUpgrade
dropperDemotedDocroot
dropperDemotedDirRemoved
dropperDemotedReplaced
dropperDemotedBackupState
dropperSuspect
)
func dropperVerdictDemoted(v dropperVerdict) bool {
return v >= dropperDemotedTemplate && v <= dropperDemotedBackupState
}
func dropperSameIdentity(c dropperCandidate, current dropperFileState) bool {
if c.Inode == 0 || current.Inode == 0 || c.Device != current.Device || c.Inode != current.Inode {
return false
}
if c.BirthKnown != current.BirthKnown {
return false
}
return !c.BirthKnown || c.Birth.Equal(current.Birth)
}
// dropperReplacedInPlace reports whether the file now at the candidate's path
// is a regular file that came into existence after the candidate was observed.
// That is an atomic write completing (write temp, rename over the live path),
// which repeats every few minutes for WAF and cache state files. The successor
// stays on disk and is scanned in its own right, and the candidate's own bytes
// were already read by the content pass, so the evidence loss that makes a
// self-delete Critical does not apply.
//
// This does not weaken the detector against an attacker who leaves a benign
// file behind: overwriting the same inode already returns dropperBenign above,
// which is both cheaper and quieter than unlink plus rename. A birth time is
// required, so a filesystem without STATX_BTIME keeps the suspect verdict.
func dropperReplacedInPlace(c dropperCandidate, current dropperFileState) bool {
if !current.IsRegular || !current.BirthKnown {
return false
}
return !current.Birth.Before(c.Observed)
}
// assessDropper turns a probe result into a verdict for one candidate.
func assessDropper(c dropperCandidate, p dropperProbe) dropperVerdict {
if !p.Conclusive {
// Due removed this candidate from the tracker. The probe loop should
// reinsert it with Retry and handle a false capacity result.
return dropperInconclusive
}
if p.QuarantineMatched {
return dropperBenign
}
if p.AtPath != nil {
if p.AtPath.Path != c.Path {
return dropperSuspect
}
if dropperSameIdentity(c, *p.AtPath) {
return dropperBenign
}
}
if p.RenamedTo != "" || p.RenameTarget != nil {
if p.RenamedTo == "" || p.RenameTarget == nil || p.RenameTarget.Path != p.RenamedTo {
return dropperSuspect
}
if dropperRenameTargetAllowed(c, p.RenamedTo) && dropperRenameMatch(c, *p.RenameTarget) {
return dropperBenign
}
}
if c.ContentSuspicious {
return dropperSuspect
}
if dropperOfficialVersionProbe(c, p) {
return dropperBenign
}
if dropperOfficialCorePackageFile(c, p) {
return dropperBenign
}
if !c.WritePending && !c.ContentMayExecute && !c.ContentUnsettled && dropperCandidateIsHarmless(c) {
return dropperBenign
}
if p.AtPath != nil && dropperReplacedInPlace(c, *p.AtPath) {
return dropperDemotedReplaced
}
if p.DocrootRemoved {
return dropperDemotedDocroot
}
if p.ParentRemoved {
return dropperDemotedDirRemoved
}
if atomicWriteRenameCandidate(c.Path) != "" {
return dropperDemotedAtomicWrite
}
if len(wpUpgradeRenameCandidates(c.Path, c.Docroot)) > 0 {
return dropperDemotedWPUpgrade
}
if looksLikeCompiledTemplate(c.Head) {
return dropperDemotedTemplate
}
if looksLikeBackWPupJobState(c.Path, c.Head) {
return dropperDemotedBackupState
}
return dropperSuspect
}
var backwpupFolderListName = regexp.MustCompile(`^backwpup-[A-Za-z0-9]+-folder\.php$`)
// looksLikeBackWPupJobState recognises the state files the BackWPup plugin
// writes while a backup job runs and deletes when it ends. The file name and
// the head must both match the writer's format: backwpup-working.php is
// "<?php //" plus the job as JSON on one line, and backwpup-<hash>-folder.php
// is "<?php" and then one "//<absolute folder>" comment per line. Every
// visible byte after the opening tag must stay inside those comments; the
// tracked head cannot see the rest of the file, so this only demotes.
func looksLikeBackWPupJobState(path string, head []byte) bool {
name := filepath.Base(path)
switch {
case name == "backwpup-working.php":
rest, ok := bytes.CutPrefix(head, []byte(`<?php //{"`))
return ok && !bytes.ContainsAny(rest, "\r\n") && !bytes.Contains(rest, []byte("?>"))
case backwpupFolderListName.MatchString(name):
rest, ok := bytes.CutPrefix(head, []byte("<?php\n"))
if !ok {
if rest, ok = bytes.CutPrefix(head, []byte("<?php\r\n")); !ok {
return false
}
}
lines := bytes.Split(rest, []byte("\n"))
folders := 0
for i, line := range lines {
last := i == len(lines)-1
// A full retained head may stop between CR and LF. Removing
// only the terminal CR still rejects any bare CR before code.
if !last || len(head) == dropperTrackedHeadMax {
line = bytes.TrimSuffix(line, []byte("\r"))
}
if bytes.ContainsRune(line, '\r') || bytes.Contains(line, []byte("?>")) {
return false
}
switch {
case bytes.HasPrefix(line, []byte("///")):
folders++
case last && len(line) < 3 && bytes.Equal(line, []byte("///")[:len(line)]):
// The head can end inside the next comment marker.
default:
return false
}
}
return folders > 0
}
return false
}
// looksLikeCompiledTemplate recognises template-engine compile artifacts
// (Twig class caches as written by phpMyAdmin/Symfony/Drupal, Smarty
// compile dirs). These are legitimately created and unlinked in short
// windows during cache rebuilds.
func looksLikeCompiledTemplate(head []byte) bool {
head = bytes.TrimSpace(bytes.TrimPrefix(head, []byte{0xef, 0xbb, 0xbf}))
if !bytes.HasPrefix(head, []byte("<?php")) {
return false
}
if bytes.Contains(head, []byte("class __TwigTemplate_")) &&
bytes.Contains(head, []byte(" extends Template")) {
return true
}
if bytes.Contains(head, []byte("/* Smarty version ")) &&
bytes.Contains(head, []byte(", created on ")) &&
bytes.Contains(head, []byte("from '")) {
return true
}
return false
}
// wpUpgradeStagedPath splits a clean absolute path under
// <wpRoot>/wp-content/upgrade/ into the WordPress root and the part below
// upgrade/. The root must be the configured docroot or inside it.
func wpUpgradeStagedPath(path, configuredDocroot string) (wpRoot, rest string, ok bool) {
const marker = "/wp-content/upgrade/"
if !filepath.IsAbs(path) || !filepath.IsAbs(configuredDocroot) ||
filepath.Clean(path) != path || filepath.Clean(configuredDocroot) != configuredDocroot {
return "", "", false
}
idx := strings.Index(path, marker)
if idx < 0 {
return "", "", false
}
wpRoot = path[:idx]
if wpRoot != configuredDocroot && !strings.HasPrefix(wpRoot, configuredDocroot+string(filepath.Separator)) {
return "", "", false
}
return wpRoot, path[idx+len(marker):], true
}
// wpUpgradeRenameCandidates maps a path inside a WordPress upgrade staging
// dir (wp-content/upgrade/<staging>/<package>/<rest>) to the destinations
// WordPress moves it to on success: the plugin and theme dirs, or the
// docroot itself for core packages. The fanotify mask has no FAN_MOVED_TO,
// so a successful rename-based install makes the staged path vanish; the
// probe checks these destinations before calling it a self-deleting drop.
func wpUpgradeRenameCandidates(path, configuredDocroot string) []string {
wpRoot, rest, ok := wpUpgradeStagedPath(path, configuredDocroot)
if !ok {
return nil
}
parts := strings.SplitN(rest, "/", 3)
if len(parts) < 3 || parts[1] == "" || parts[1] == "." || parts[1] == ".." ||
parts[2] == "" || filepath.Clean(parts[2]) != parts[2] {
// A file directly under upgrade/<staging>/ has no package dir to
// move. The flat language-pack copy is handled by
// wpUpgradeInstallDestinations.
return nil
}
pkg, tail := parts[1], parts[2]
if pkg == "wordpress" {
return []string{filepath.Join(wpRoot, tail)}
}
return []string{
filepath.Join(wpRoot, "wp-content", "plugins", pkg, tail),
filepath.Join(wpRoot, "wp-content", "themes", pkg, tail),
}
}
// wpUpgradeStagedPackageFile reports whether path lies inside an unpacked
// core, plugin or theme package under wp-content/upgrade/.
func wpUpgradeStagedPackageFile(path, configuredDocroot string) bool {
return len(wpUpgradeRenameCandidates(path, configuredDocroot)) > 0
}
// wpUpgradeCorePackageFile splits a path inside an unpacked core release,
// <wpRoot>/wp-content/upgrade/<working>/wordpress/<rel>, into the WordPress
// root and the file's key in the release checksum manifest.
func wpUpgradeCorePackageFile(path, configuredDocroot string) (wpRoot, rel string, ok bool) {
wpRoot, rest, ok := wpUpgradeStagedPath(path, configuredDocroot)
if !ok {
return "", "", false
}
parts := strings.SplitN(rest, "/", 3)
if len(parts) < 3 || parts[0] == "" || parts[1] != "wordpress" ||
parts[2] == "" || filepath.Clean(parts[2]) != parts[2] {
return "", "", false
}
return wpRoot, parts[2], true
}
// wpUpgradeInstallDestinations lists every place the WordPress updater puts
// the bytes of a file it wrote under wp-content/upgrade/ before removing it:
// the package-tree moves above, plus two copy-then-delete steps that leave no
// package directory behind. Language packs unzip flat into
// upgrade/<working>/ and are copied into wp-content/languages/ (plugins/ and
// themes/ for those pack types). A core update copies the staged
// wp-includes/version.php to upgrade/version-current.php, reads it, deletes
// it, and later installs the same file as wp-includes/version.php.
//
// Copy destinations also require a complete data-only snapshot. An identical
// payload planted in languages/ or version.php does not prove updater activity.
func wpUpgradeInstallDestinations(path, configuredDocroot string) []string {
if dests := wpUpgradeRenameCandidates(path, configuredDocroot); dests != nil {
return dests
}
return wpUpgradeCopyDestinations(path, configuredDocroot)
}
func wpUpgradeCopyDestinations(path, configuredDocroot string) []string {
wpRoot, rest, ok := wpUpgradeStagedPath(path, configuredDocroot)
if !ok {
return nil
}
parts := strings.Split(rest, "/")
switch {
case len(parts) == 1 && parts[0] == "version-current.php":
return []string{filepath.Join(wpRoot, "wp-includes", "version.php")}
case len(parts) == 2 && strings.HasSuffix(parts[1], ".l10n.php"):
// Clean already rejected empty, "." and ".." components.
languages := filepath.Join(wpRoot, "wp-content", "languages")
return []string{
filepath.Join(languages, parts[1]),
filepath.Join(languages, "plugins", parts[1]),
filepath.Join(languages, "themes", parts[1]),
}
}
return nil
}
func dropperRenameTargetAllowed(c dropperCandidate, target string) bool {
if atomicTarget := atomicWriteRenameCandidate(c.Path); atomicTarget != "" && target == atomicTarget {
return true
}
for _, candidate := range wpUpgradeRenameCandidates(c.Path, c.Docroot) {
if target == candidate {
return true
}
}
if !dropperWPCopyDataEligible(c) {
return false
}
for _, candidate := range wpUpgradeCopyDestinations(c.Path, c.Docroot) {
if target == candidate {
return true
}
}
return false
}
// dropperWPCopyDataEligible reports whether every snapshot of c was complete
// translation or version data with nothing that could have run.
func dropperWPCopyDataEligible(c dropperCandidate) bool {
return c.WPInstallData && !c.WPInstallUnsafe && c.DigestKnown && !c.WritePending &&
!c.ContentSuspicious && c.Mode&0o111 == 0
}
// dropperOfficialVersionProbe covers a core update that stops after reading
// the new version file, for example on a failed PHP or database requirement.
// The installed version file then belongs to the old release and cannot match,
// so the removed probe must itself be the official file of the release it
// declares.
func dropperOfficialVersionProbe(c dropperCandidate, p dropperProbe) bool {
return p.OfficialWPCoreFile && c.WPCoreRelease != nil && dropperWPCopyDataEligible(c) &&
filepath.Base(c.Path) == "version-current.php" &&
len(wpUpgradeCopyDestinations(c.Path, c.Docroot)) > 0
}
// dropperOfficialCorePackageFile covers the files a core update unpacks but
// never installs: bundled themes and plugins under wp-content/, and files
// unchanged from the running release. WordPress deletes them with the working
// tree. Official release bytes carry nothing an attacker chose, so removing
// them loses no evidence.
func dropperOfficialCorePackageFile(c dropperCandidate, p dropperProbe) bool {
return p.OfficialWPCorePackageFile && c.CoreMD5Known && !c.ContentRewritten &&
!c.WritePending && !c.ContentUnsettled && c.Mode&0o111 == 0
}
// dropperRenameMatch reports whether a probe of a rename-destination path
// identifies the same file as the tracked candidate: identical device,
// inode, and birth time for rename(2), or identical size plus a full SHA-256
// digest for a copy-delete fallback across filesystems.
func dropperRenameMatch(c dropperCandidate, dest dropperFileState) bool {
// Neither the moved file nor a surviving copy of the last staged snapshot
// can account for an earlier payload overwritten before the move.
staged := wpUpgradeStagedPackageFile(c.Path, c.Docroot)
if staged && c.ContentRewritten {
return false
}
if dropperSameIdentity(c, dest) {
return true
}
if staged && (c.ContentUnsettled || c.WritePending) {
return false
}
return c.DigestKnown && dest.DigestKnown && c.Size == dest.Size && c.Digest == dest.Digest
}
// dropperGraceWindow is how long a vanished candidate is held before its
// finding flushes. The hold lets a bulk operation (plugin upgrade fallback,
// cache purge, deploy rollback) accumulate its siblings so the whole batch
// collapses into one Warning for classified churn or one High signal for an
// unclassified burst, instead of paging Critical per file.
const dropperGraceWindow = 45 * time.Second
// dropperBurstThreshold is the group size at which held candidates from one
// docroot are reported as a single bulk-churn aggregate instead of
// individual findings.
const dropperBurstThreshold = 8
type dropperGone struct {
Cand dropperCandidate
Verdict dropperVerdict
held time.Time
ticket queuehealth.Ticket
}
// dropperFinding is one flush decision: a single vanished file, a removed
// directory group, or a per-docroot aggregate of a create/delete burst.
type dropperFinding struct {
Aggregate bool
Docroot string
Items []dropperGone
// RemovedDir groups files that were each demoted because this directory
// was removed with them: one event, reported once.
RemovedDir string
}
// HoldGone parks a vanished candidate until FlushDue decides whether it is
// reported alone or as part of a bulk-churn aggregate.
func (t *dropperTracker) HoldGone(c dropperCandidate, v dropperVerdict, now time.Time) {
if v == dropperBenign || v == dropperInconclusive {
return
}
t.initQueueHealth()
t.mu.Lock()
defer t.mu.Unlock()
if len(t.pending) >= dropperMaxTracked {
t.heldHealth.Lose(t.now(), 1)
return
}
c.ticket = queuehealth.Ticket{}
t.pending = append(t.pending, dropperGone{
Cand: ownDropperCandidate(c), Verdict: v, held: now,
ticket: t.heldHealth.BeginAt(now.Add(dropperGraceWindow), t.now()),
})
}
// FlushDue emits findings for docroot groups whose oldest held entry has
// aged past the grace window. The whole group flushes together so entries
// arriving late in a burst still fold into the aggregate.
func (t *dropperTracker) FlushDue(now time.Time) []dropperFinding {
t.mu.Lock()
defer t.mu.Unlock()
started := t.now()
type groupKey struct {
docroot string
demoted bool
}
keyFor := func(g dropperGone) groupKey {
return groupKey{docroot: g.Cand.Docroot, demoted: dropperVerdictDemoted(g.Verdict)}
}
oldest := make(map[groupKey]time.Time)
for _, g := range t.pending {
key := keyFor(g)
if first, ok := oldest[key]; !ok || g.held.Before(first) {
oldest[key] = g.held
}
}
groups := make(map[groupKey][]dropperGone)
var keep []dropperGone
for _, g := range t.pending {
key := keyFor(g)
if now.Sub(oldest[key]) >= dropperGraceWindow {
g.ticket.Start(started)
groups[key] = append(groups[key], g)
} else {
keep = append(keep, g)
}
}
t.pending = keep
var out []dropperFinding
for key, items := range groups {
out = append(out, groupDropperFindings(key.docroot, items)...)
}
return out
}
// groupDropperFindings also runs after ignore-path filtering: a burst can
// shrink into removed-directory groups, and a directory group into one file.
// Callers keep docroots and demoted/unclassified batches separate.
func groupDropperFindings(docroot string, items []dropperGone) []dropperFinding {
if len(items) >= dropperBurstThreshold {
return []dropperFinding{{Aggregate: true, Docroot: docroot, Items: items}}
}
var out []dropperFinding
byDir := make(map[string][]dropperGone)
for _, item := range items {
if item.Verdict == dropperDemotedDirRemoved {
dir := filepath.Dir(item.Cand.Path)
byDir[dir] = append(byDir[dir], item)
continue
}
out = append(out, dropperFinding{Docroot: docroot, Items: []dropperGone{item}})
}
for dir, members := range byDir {
f := dropperFinding{Docroot: docroot, Items: members}
if len(members) > 1 {
f.RemovedDir = dir
}
out = append(out, f)
}
return out
}
// dropperHeadExcerptMax caps how many leading file bytes a finding's
// details reproduce as evidence. The bounded tracked head stays in memory
// only until the flush; findings carry just enough to triage without a file
// (the file is gone by definition).
const dropperHeadExcerptMax = 160
// dropperAggregateSampleMax caps how many member paths an aggregate
// finding lists before truncating with a count.
const dropperAggregateSampleMax = 10
const (
dropperPathExcerptMax = 512
dropperProcExcerptMax = 256
)
// dropperAlertParams renders one flush decision into alert parameters:
// severity, message, details and the finding path.
func dropperAlertParams(f dropperFinding) (alert.Severity, string, string, string) {
if f.RemovedDir != "" {
var b strings.Builder
b.WriteString("Files created and removed within the tracking TTL together with the directory holding them; each was demoted for that reason:\n")
for i, g := range f.Items {
if i == dropperAggregateSampleMax {
fmt.Fprintf(&b, "... and %d more", len(f.Items)-dropperAggregateSampleMax)
break
}
fmt.Fprintf(&b, "%s (uid=%d size=%d)\n",
dropperPrintable([]byte(g.Cand.Path), dropperPathExcerptMax), g.Cand.UID, g.Cand.Size)
}
msg := fmt.Sprintf("%d short-lived PHP/executable files removed with their directory %s",
len(f.Items), dropperPrintable([]byte(f.RemovedDir), dropperPathExcerptMax))
return alert.Warning, msg, strings.TrimRight(b.String(), "\n"), f.RemovedDir
}
if f.Aggregate {
severity := alert.Warning
unclassified := 0
for _, g := range f.Items {
if !dropperVerdictDemoted(g.Verdict) {
unclassified++
}
}
if unclassified > 0 {
severity = alert.High
}
var b strings.Builder
fmt.Fprintf(&b, "Files created and removed within the tracking TTL; %d remained unclassified after false-positive checks:\n", unclassified)
for i, g := range f.Items {
if i == dropperAggregateSampleMax {
fmt.Fprintf(&b, "... and %d more", len(f.Items)-dropperAggregateSampleMax)
break
}
fmt.Fprintf(&b, "%s (uid=%d size=%d)\n",
dropperPrintable([]byte(g.Cand.Path), dropperPathExcerptMax), g.Cand.UID, g.Cand.Size)
}
displayDocroot := dropperPrintable([]byte(f.Docroot), dropperPathExcerptMax)
msg := fmt.Sprintf("%d short-lived PHP/executable files created and removed under %s", len(f.Items), displayDocroot)
return severity, msg, strings.TrimRight(b.String(), "\n"), f.Docroot
}
g := f.Items[0]
c := g.Cand
lifetime := "unknown"
if c.BirthKnown {
lifetime = c.Observed.Sub(c.Birth).Truncate(time.Second).String()
}
details := fmt.Sprintf(
"File appeared and was removed before the TTL probe. uid=%d size=%d mode=%04o write-age=%s",
c.UID, c.Size, c.Mode&0o7777, lifetime)
if c.ProcInfo != "" {
details += " writer=[" + dropperPrintable([]byte(c.ProcInfo), dropperProcExcerptMax) + "]"
}
sev := alert.Critical
if dropperVerdictDemoted(g.Verdict) {
sev = alert.Warning
switch g.Verdict {
case dropperDemotedTemplate:
details += "\nDemoted: content matches a compiled-template artifact (Twig/Smarty cache churn)."
case dropperDemotedAtomicWrite:
details += "\nDemoted: filename matches an atomic-write staging artifact."
case dropperDemotedWPUpgrade:
details += "\nDemoted: path is structurally inside a WordPress upgrade staging tree."
case dropperDemotedDocroot:
details += "\nDemoted: the containing document root was removed before the probe."
case dropperDemotedDirRemoved:
details += "\nDemoted: the original containing directory was removed before the probe."
case dropperDemotedReplaced:
details += "\nDemoted: the path was replaced in place by a newer file (atomic write), not emptied."
case dropperDemotedBackupState:
details += "\nDemoted: content matches BackWPup job state written behind a PHP comment; only the leading bytes were seen."
}
}
if len(c.Head) > 0 {
details += "\nLeading bytes: " + dropperPrintable(c.Head, dropperHeadExcerptMax)
}
msg := fmt.Sprintf("Self-deleting file under web root: %s",
dropperPrintable([]byte(c.Path), dropperPathExcerptMax))
return sev, msg, details, c.Path
}
// dropperPrintable renders up to max bytes of b with control and non-ASCII
// bytes replaced by '.' so binary heads (ELF droppers) cannot corrupt
// alert transports or terminal output.
func dropperPrintable(b []byte, max int) string {
if len(b) > max {
b = b[:max]
}
out := make([]byte, len(b))
for i, c := range b {
if c >= 0x20 && c < 0x7f {
out[i] = c
} else {
out[i] = '.'
}
}
return string(out)
}
// dropperDocrootFor returns the longest configured web document root that
// contains path, or "" when the path is not inside any docroot. Matching is
// component-safe: /home/a/public_html_old is not inside /home/a/public_html.
func dropperDocrootFor(path string, docroots []string) string {
path = filepath.Clean(path)
if !filepath.IsAbs(path) {
return ""
}
best := ""
for _, configuredRoot := range docroots {
root := filepath.Clean(configuredRoot)
if !filepath.IsAbs(root) {
continue
}
rel, err := filepath.Rel(root, path)
if err != nil || rel == "." || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) {
continue
}
if len(root) > len(best) {
best = root
}
}
return best
}
//go:build !(linux && bpf)
package daemon
import (
"context"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
)
// execBPF is the no-tag placeholder for the BPF exec-monitor backend. The
// real type with a tracepoint link and ringbuf reader lives in exec_bpf.go
// behind //go:build linux && bpf.
type execBPF struct{}
func (e *execBPF) Mode() string { return "bpf" }
func (e *execBPF) EventCount() uint64 { return 0 }
func (e *execBPF) Run(_ context.Context) {}
func startExecBPF(_ context.Context, _ chan<- alert.Finding, _ *config.Config) (*execBPF, error) {
return nil, bpf.ErrNotBuilt
}
// Code generated by bpf2go; DO NOT EDIT.
//go:build 386 || amd64
package exec_bpfprog
import (
"bytes"
_ "embed"
"fmt"
"io"
"structs"
"github.com/cilium/ebpf"
)
type ExecCsmQueueStats struct {
_ structs.HostLayout
Lost uint64
Submitted uint64
}
type ExecExecEvent struct {
_ structs.HostLayout
Uid uint32
Pid uint32
Ppid uint32
Comm [16]uint8
ParentComm [16]uint8
Filename [256]uint8
}
// Names of all BPF objects in the ELF.
//
// Used for safe lookups in a Collection or CollectionSpec.
const (
ExecMapEvents = "events"
ExecMapQueueStats = "queue_stats"
ExecProgCsmOnExec = "csm_on_exec"
ExecVarUnused = "unused"
)
// LoadExec returns the embedded CollectionSpec for Exec.
func LoadExec() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_ExecBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load Exec: %w", err)
}
return spec, err
}
// LoadExecObjects loads Exec and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *ExecObjects
// *ExecPrograms
// *ExecMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func LoadExecObjects(obj any, opts *ebpf.CollectionOptions) error {
spec, err := LoadExec()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// ExecSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ExecSpecs struct {
ExecProgramSpecs
ExecMapSpecs
ExecVariableSpecs
}
// ExecProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ExecProgramSpecs struct {
CsmOnExec *ebpf.ProgramSpec `ebpf:"csm_on_exec"`
}
// ExecMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ExecMapSpecs struct {
Events *ebpf.MapSpec `ebpf:"events"`
QueueStats *ebpf.MapSpec `ebpf:"queue_stats"`
}
// ExecVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type ExecVariableSpecs struct {
Unused *ebpf.VariableSpec `ebpf:"unused"`
}
// ExecObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to LoadExecObjects or ebpf.CollectionSpec.LoadAndAssign.
type ExecObjects struct {
ExecPrograms
ExecMaps
ExecVariables
}
func (o *ExecObjects) Close() error {
return _ExecClose(
&o.ExecPrograms,
&o.ExecMaps,
)
}
// ExecMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to LoadExecObjects or ebpf.CollectionSpec.LoadAndAssign.
type ExecMaps struct {
Events *ebpf.Map `ebpf:"events"`
QueueStats *ebpf.Map `ebpf:"queue_stats"`
}
func (m *ExecMaps) Close() error {
return _ExecClose(
m.Events,
m.QueueStats,
)
}
// ExecVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to LoadExecObjects or ebpf.CollectionSpec.LoadAndAssign.
type ExecVariables struct {
Unused *ebpf.Variable `ebpf:"unused"`
}
// ExecPrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to LoadExecObjects or ebpf.CollectionSpec.LoadAndAssign.
type ExecPrograms struct {
CsmOnExec *ebpf.Program `ebpf:"csm_on_exec"`
}
func (p *ExecPrograms) Close() error {
return _ExecClose(
p.CsmOnExec,
)
}
func _ExecClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed exec_x86_bpfel.o
var _ExecBytes []byte
package daemon
import (
"encoding/binary"
"errors"
)
// ExecEvent matches struct exec_event in exec.bpf.c byte for byte:
// scalars are little-endian (host order on amd64/arm64), comm/parent_comm
// are 16-byte null-padded, filename is 256-byte null-padded.
type ExecEvent struct {
UID uint32
PID uint32
PPID uint32
Comm string
ParentComm string
Filename string
}
const execEventSize = 4 + 4 + 4 + 16 + 16 + 256
func decodeExecEvent(b []byte) (ExecEvent, error) {
if len(b) < execEventSize {
return ExecEvent{}, errors.New("exec event short buffer")
}
ev := ExecEvent{
UID: binary.LittleEndian.Uint32(b[0:4]),
PID: binary.LittleEndian.Uint32(b[4:8]),
PPID: binary.LittleEndian.Uint32(b[8:12]),
}
ev.Comm = nullTerm(b[12:28])
ev.ParentComm = nullTerm(b[28:44])
ev.Filename = nullTerm(b[44 : 44+256])
return ev, nil
}
package daemon
import (
"context"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/state"
)
// execPoller wraps the periodic CheckSuspiciousProcesses + CheckFakeKernelThreads
// in a goroutine. Used when the BPF backend is unavailable or operator-disabled.
// Detection latency for already-running processes equals the poll interval (default
// 30 minutes, matching the existing deep-tier cadence).
type execPoller struct {
cfg *config.Config
alertCh chan<- alert.Finding
count atomic.Uint64
}
func newExecPoller(cfg *config.Config, alertCh chan<- alert.Finding) *execPoller {
return &execPoller{cfg: cfg, alertCh: alertCh}
}
func (p *execPoller) Mode() string { return "legacy" }
func (p *execPoller) EventCount() uint64 { return p.count.Load() }
func (p *execPoller) Run(ctx context.Context) {
interval := execPollerInterval(p.cfg)
t := time.NewTicker(interval)
defer t.Stop()
emit := func(fs []alert.Finding) {
for _, f := range fs {
p.count.Add(1)
if !alert.TryEnqueue(p.alertCh, f) {
csmlog.Warn("exec legacy: alert channel full, dropping finding")
}
}
}
for {
select {
case <-ctx.Done():
return
case <-t.C:
emit(checks.CheckSuspiciousProcesses(ctx, p.cfg, (*state.Store)(nil)))
emit(checks.CheckFakeKernelThreads(ctx, p.cfg, (*state.Store)(nil)))
}
}
}
func execPollerInterval(cfg *config.Config) time.Duration {
if d := cfg.Detection.ExecMonitorPollInterval; d > 0 {
return d
}
return 30 * time.Minute
}
package daemon
import (
"context"
"errors"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/processctx"
)
// StartExecMonitor selects the active exec-monitor backend based on
// cfg.Detection.ExecMonitorBackend and host capability:
//
// "auto" (default) -- try BPF, fall back to legacy polling.
// "bpf" -- require BPF; return nil if unavailable (no fallback).
// "legacy" -- pin legacy polling.
// "none" -- disable the live monitor (periodic checks still run).
//
// Unknown values fall back to "auto" with a warning. The metric
// csm_bpf_backend{feature="exec_monitor", kind="..."} reflects the
// chosen path.
func StartExecMonitor(alertCh chan<- alert.Finding, cfg *config.Config) bpf.Backend {
choice := strings.ToLower(strings.TrimSpace(cfg.Detection.ExecMonitorBackend))
if choice == "" {
choice = bpf.BackendAuto
}
switch choice {
case bpf.BackendAuto, bpf.BackendBPF, bpf.BackendLegacy, bpf.BackendNone:
default:
csmlog.Warn("exec_monitor: unknown backend choice, using auto", "value", choice)
choice = bpf.BackendAuto
}
if choice == bpf.BackendNone {
csmlog.Info("exec_monitor: disabled by config")
bpf.SetActive("exec_monitor", bpf.BackendNone)
return nil
}
var bpfErr error
if choice == bpf.BackendAuto || choice == bpf.BackendBPF {
if b, err := tryStartExecBPFFn(context.Background(), alertCh, cfg); err == nil && b != nil {
csmlog.Info("exec_monitor", "backend", "bpf", "choice", choice)
bpf.SetActive("exec_monitor", bpf.BackendBPF)
return b
} else if err != nil {
bpfErr = err
level := "bpf-unsupported"
if errors.Is(err, bpf.ErrNotBuilt) {
level = "bpf-not-built"
}
csmlog.Info("exec_monitor: BPF unavailable", "state", level, "reason", err.Error(), "choice", choice)
if choice == bpf.BackendBPF {
csmlog.Warn("exec_monitor: backend=bpf but BPF unavailable; no live monitor", "reason", err.Error())
bpf.SetActive("exec_monitor", bpf.BackendNone)
emitBPFUnavailableFinding(alertCh, "exec_monitor", choice, "", err)
return nil
}
}
}
poller := newExecPoller(cfg, alertCh)
csmlog.Info("exec_monitor", "backend", "legacy", "choice", choice)
bpf.SetActive("exec_monitor", bpf.BackendLegacy)
if bpfErr != nil {
emitBPFUnavailableFinding(alertCh, "exec_monitor", choice, bpf.BackendLegacy, bpfErr)
}
return poller
}
// tryStartExecBPFFn is the package-level indirection so tests can substitute
// a fake without the bpf build tag.
var tryStartExecBPFFn = tryStartExecBPF
func tryStartExecBPF(ctx context.Context, ch chan<- alert.Finding, cfg *config.Config) (bpf.Backend, error) {
b, err := startExecBPF(ctx, ch, cfg)
if err != nil {
return nil, err
}
return b, nil
}
// populateProcessCtxFromExec writes the BPF exec event into the process
// context cache without doing identity work in the event loop. Zero PID is
// ignored (synthetic boot-time noise).
func populateProcessCtxFromExec(cache *processctx.Cache, ev ExecEvent, startedAt time.Time) {
if ev.PID == 0 {
return
}
cache.PutFromExecStartedAt(int(ev.PID), int(ev.PPID), int(ev.UID), ev.Comm, ev.Filename, startedAt)
}
func attachProcessCtxToExecFinding(cache *processctx.Cache, f *alert.Finding, ev ExecEvent) {
req := processctx.EnrichRequest{
PID: int(ev.PID),
UID: int(ev.UID),
UIDKnown: true,
Comm: ev.Comm,
StartedAt: processCtxStartedAt(int(ev.PID)),
}
if pc, _ := cache.MaterializeVerifiedSnapshot(req); pc != nil {
f.Process = pc
}
}
// processctxRequestFromExec builds the EnrichRequest snapshot used by the BPF
// exec consumer. Lives here (no build tag) so unit tests on darwin can verify
// the field mapping without a Linux+bpf build.
func processctxRequestFromExec(ev ExecEvent) processctx.EnrichRequest {
pid := int(ev.PID)
return processctx.EnrichRequest{
PID: pid,
UID: int(ev.UID),
UIDKnown: true,
Comm: ev.Comm,
StartedAt: processCtxStartedAt(pid),
}
}
//go:build linux
package daemon
import (
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"runtime"
"runtime/debug"
"strings"
"sync"
"sync/atomic"
"time"
"unsafe"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/wpcheck"
"github.com/pidginhost/csm/internal/yara"
)
// fanotify constants (not all in Go stdlib)
const (
FAN_MARK_ADD = 0x00000001
// FAN_MARK_MOUNT covers a single vfsmount. FAN_MARK_FILESYSTEM marks the
// whole superblock, so a write that reaches the same inode through a bind
// mount is reported too. EL8 backported the flag into 4.18, which is what
// CloudLinux 8 runs, so cages are reachable on production kernels.
FAN_MARK_MOUNT = 0x00000010
FAN_MARK_FILESYSTEM = 0x00000100
FAN_CLOSE_WRITE = 0x00000008
FAN_CREATE = 0x00000100
FAN_CLASS_NOTIF = 0x00000000
FAN_CLOEXEC = 0x00000001
FAN_NONBLOCK = 0x00000002
)
// markFunc is the fanotify_mark syscall, passed in so the ladder can be
// exercised without a kernel and without a mutable package-level seam that
// concurrent tests would race on.
type markFunc func(fd int, flags uint, mask uint64, dirFd int, path string) error
// markScope records how widely a watch root ended up being marked.
type markScope int
const (
markScopeNone markScope = iota
markScopeFilesystem
markScopeMount
)
func (s markScope) String() string {
switch s {
case markScopeFilesystem:
return "filesystem"
case markScopeMount:
return "mount"
default:
return "none"
}
}
// markWatchRoot watches path, preferring a filesystem-scoped mark.
//
// A mount-scoped mark sees only the vfsmount it was added to. Every CloudLinux
// CageFS account reaches its files through a bind mount of the same superblock,
// so writes inside a cage produced no event at all and the realtime scanner was
// blind for precisely the accounts most likely to be compromised. Marking the
// superblock covers every mount of it.
//
// The ladder degrades in two independent directions: kernels without
// FAN_MARK_FILESYSTEM fall back to the mount mark, and kernels without
// FAN_CREATE (EL8 among them) keep their scope and drop that event bit.
func markWatchRoot(fd int, path string, mark markFunc) (markScope, error) {
var lastErr error
for _, attempt := range []struct {
scope markScope
flags uint
}{
{markScopeFilesystem, FAN_MARK_ADD | FAN_MARK_FILESYSTEM},
{markScopeMount, FAN_MARK_ADD | FAN_MARK_MOUNT},
} {
for _, mask := range []uint64{FAN_CLOSE_WRITE | FAN_CREATE, FAN_CLOSE_WRITE} {
if err := mark(fd, attempt.flags, mask, -1, path); err != nil {
lastErr = err
continue
}
return attempt.scope, nil
}
}
return markScopeNone, lastErr
}
// watchRootMark records what a watch root ended up being watched by. coveredBy
// names the root whose filesystem-scoped mark covers this one. A later
// filesystem mark can also cover an earlier mount mark.
type watchRootMark struct {
path string
scope markScope
device uint64
coveredBy string
// rootFS records that this root sits on the same filesystem as /, so a
// filesystem-scoped mark on it raises an event for every write on the
// machine. That is device identity, not an inference about containment.
rootFS bool
}
// devStatFunc reports the device a path lives on. Injected so the mark ladder
// can be exercised without a filesystem laid out like a production host.
type devStatFunc func(path string) (uint64, error)
func statDevice(path string) (uint64, error) {
var st unix.Stat_t
if err := unix.Stat(path, &st); err != nil {
return 0, err
}
return uint64(st.Dev), nil // #nosec G115 -- device ids are non-negative on Linux
}
// markWatchRoots skips roots covered by a successful filesystem mark, never
// by a mount mark. The returned error includes partial failures so startup can
// report unwatched roots while continuing with the marks that succeeded.
func markWatchRoots(fd int, roots []string, mark markFunc, statDev devStatFunc) ([]watchRootMark, error) {
var marks []watchRootMark
covering := make(map[uint64]string)
var failures []error
rootDevice, rootErr := statDev("/")
for _, path := range roots {
device, statErr := statDev(path)
if statErr != nil {
if !errors.Is(statErr, os.ErrNotExist) {
failures = append(failures, fmt.Errorf("cannot inspect watch root %s: %w", path, statErr))
}
continue
}
if owner, covered := covering[device]; covered {
marks = append(marks, watchRootMark{
path: path, scope: markScopeFilesystem, device: device, coveredBy: owner,
rootFS: rootErr == nil && device == rootDevice,
})
continue
}
scope, err := markWatchRoot(fd, path, mark)
if err != nil {
failures = append(failures, fmt.Errorf("cannot watch %s: %w", path, err))
continue
}
if scope == markScopeFilesystem {
covering[device] = path
}
marks = append(marks, watchRootMark{
path: path, scope: scope, device: device,
rootFS: rootErr == nil && device == rootDevice,
})
}
// A later filesystem mark makes an earlier mount-only warning obsolete.
for i := range marks {
if owner, covered := covering[marks[i].device]; covered && owner != marks[i].path {
marks[i].scope = markScopeFilesystem
marks[i].coveredBy = owner
}
}
if len(marks) == 0 && len(failures) == 0 {
failures = append(failures, fmt.Errorf("no watch root exists"))
}
return marks, errors.Join(failures...)
}
// EventStats reports completed admission decisions for delivered file events.
type EventStats struct {
Received int64
Admitted int64
Dropped int64
}
// Filtered counts events rejected before queue admission, including events
// whose file descriptor could not be resolved to a path.
func (s EventStats) Filtered() int64 { return s.Received - s.Admitted - s.Dropped }
// EventStats returns a coherent snapshot. In-progress admission decisions are
// published with their outcome so they cannot appear as filtered events.
func (fm *FileMonitor) EventStats() EventStats {
fm.eventStatsMu.Lock()
defer fm.eventStatsMu.Unlock()
return fm.eventStats
}
func (fm *FileMonitor) recordEvent(admitted, dropped bool) {
fm.eventStatsMu.Lock()
defer fm.eventStatsMu.Unlock()
fm.eventStats.Received++
if fanotifyEventsTotal != nil {
fanotifyEventsTotal.Inc()
}
if admitted {
fm.eventStats.Admitted++
if fanotifyEventsAdmittedTotal != nil {
fanotifyEventsAdmittedTotal.Inc()
}
}
if dropped {
fm.eventStats.Dropped++
if fanotifyDroppedTotal != nil {
fanotifyDroppedTotal.Inc()
}
}
}
// WatchScopeSummary describes what the monitor actually watches.
func (fm *FileMonitor) WatchScopeSummary() string { return watchScopeSummary(fm.watchRoots) }
// A mount point does not establish filesystem containment: symlinks and bind
// mounts can name a subtree of the same superblock. Describe the scope itself
// instead of inferring a boundary by comparing a root's device to its parent's.
func watchScopeSummary(marks []watchRootMark) string {
parts := make([]string, 0, len(marks))
for _, m := range marks {
if m.coveredBy != "" {
parts = append(parts, fmt.Sprintf("%s (covered by %s)", m.path, m.coveredBy))
continue
}
coverage := "whole filesystem, including all bind mounts"
if m.scope == markScopeMount {
coverage = "whole containing mount, excluding other bind mounts"
}
if m.rootFS {
coverage += ", same filesystem as /"
}
parts = append(parts, fmt.Sprintf("%s (%s scope, device %d: %s)", m.path, m.scope, m.device, coverage))
}
return strings.Join(parts, ", ")
}
// fanotifyEventMetadata is the header for each fanotify event.
type fanotifyEventMetadata struct {
EventLen uint32
Vers uint8
Reserved uint8
MetadataLen uint16
Mask uint64
Fd int32
Pid int32
}
const metadataSize = int(unsafe.Sizeof(fanotifyEventMetadata{}))
const htaccessRealtimeMaxBytes = 1 << 20
// M1 - webshells map at package level (avoid per-call allocation)
var knownWebshells = map[string]bool{
"h4x0r.php": true, "c99.php": true, "r57.php": true,
"wso.php": true, "alfa.php": true, "b374k.php": true,
"shell.php": true, "cmd.php": true, "backdoor.php": true,
"webshell.php": true,
}
// M3 - WordPress path stat cache with TTL
type wpPathCacheEntry struct {
exists bool
ts time.Time
}
var wpPathStatCache sync.Map // key: path string → value: wpPathCacheEntry
const wpPathCacheTTL = 5 * time.Minute
// alertDedupTTL is the cooldown period for duplicate alerts on the same
// check+filepath combination. Prevents alert storms from rapid writes.
const alertDedupTTL = 30 * time.Second
// FileMonitor watches mount points for file creation/modification using fanotify.
type FileMonitor struct {
fd int
cfg *config.Config
alertCh chan<- alert.Finding
// panicMu / lastPanicAt rate-limit the realtime_scanner_panic finding
// raised when an analyzer panics on one event (see analyzeFileSafe).
panicMu sync.Mutex
lastPanicAt time.Time
analyzerCh chan fileEvent
queueHealthOnce sync.Once
analyzerHealth *queuehealth.Tracker
reconcileHealth *queuehealth.Tracker
kernelQueueHealth *queuehealth.Tracker
kernelQueue *notificationQueue
// Publish each completed decision atomically with its lifetime totals.
// The minute overflow alert drains droppedEvents independently.
eventStatsMu sync.Mutex
eventStats EventStats
droppedEvents int64
droppedAlerts int64
// queueOverflows counts FAN_Q_OVERFLOW events: the kernel notification
// queue filled and events were dropped before userspace ever saw them.
// Distinct from droppedEvents (analyzer-queue backpressure in userspace)
// because a kernel overflow carries no fd, so the affected files are
// unknown and cannot be reconciled by path.
queueOverflows int64
// overflowReportMu rate-limits the operator-facing kernel-overflow finding
// so a sustained storm does not flood the alert channel.
overflowReportMu sync.Mutex
lastOverflowReport time.Time
yaraErrorReportMu sync.Mutex
lastYARAError time.Time
// C4 - pipe for epoll stop signaling
pipeFds [2]int // [0]=read, [1]=write
pipeClosed int32 // atomic flag: 1 = pipe fds closed by drainAndClose
// drainDeadline is the unix-nano instant after which analyzer workers
// stop scanning queued events and only release them. Zero until the
// shutdown drain arms it.
drainDeadline atomic.Int64
// C2 - sync.Once for safe Stop
stopOnce sync.Once
drainOnce sync.Once
stopCh chan struct{} // internal stop channel
wg sync.WaitGroup
// Per-path alert deduplication: "check:filepath" → last alert time
alertDedup sync.Map
// accountRootPatterns and docRootPatterns describe where accounts and their
// document roots live on this platform. The realtime detectors used to
// hardcode /home and /public_html, which made every one of them dead on
// Plesk and DirectAdmin and on cPanel accounts outside /home.
accountRootPatterns []string
docRootPatterns []string
// webRootPatterns is the immutable PHP configuration root set captured at
// startup from account_roots and platform discovery.
webRootPatterns []string
// watchRoots records what each watch root ended up watching, including
// roots an earlier filesystem-scoped mark already covers.
watchRoots []watchRootMark
// WordPress checksum verifier: skips detection on unmodified core and
// plugin files and judges staged update packages file by file.
wpCache wpVerifier
// wpPending holds staged package files whose checksums are still being
// fetched; stagedPackageLoop resolves them once a second.
wpPending *stagedPackageQueue
wpPendingInit sync.Once
// Drop-recovery reconcile: directories that had fanotify events dropped
// because the analyzer queue was full. The overflow reporter schedules
// bounded passes over interesting files within the original recovery
// window so bulk filesystem operations (unzip, backup restore)
// do not blind detection to actual threats landing in the storm.
reconcileMu sync.Mutex
reconcileDirs map[string]reconcileDirectory
// reconcileSig is a buffered cap-1 channel that lets sendEvent's drop
// branch nudge overflowReporter to run reconcileDrops out of cycle
// when sustained drops cross eagerReconcileDropThreshold. The cap-1
// shape collapses multiple triggers in the same window into one and
// keeps sendEvent non-blocking on the event-loop hot path.
reconcileSig chan struct{}
// reconcileRunning is set while a recovery pass is in flight, and
// reconcileLastPass is the unix-nano instant the last one ended. Passes
// are single-flight and rate-limited: recovery that keeps up with a
// storm's trigger rate feeds the storm.
reconcileRunning atomic.Bool
reconcileLastPass atomic.Int64
// metricsOnce guards one-time registration of the fanotify-scoped
// Prometheus metrics. Each FileMonitor registers its own hooks when
// it first starts; subsequent calls are a no-op.
metricsOnce sync.Once
// dropper drives the self-deleting-dropper detector: candidates are
// admitted from the analyzer path and probed for deletion after a TTL by
// dropperProbeLoop. nil when thresholds.dropper_detection is off.
dropper *dropperEngine
// dropperDocroots is the cached web-document-root list the admission
// gate matches paths against, refreshed on the probe loop's cadence so
// account add/remove is picked up without a restart.
dropperDocroots atomic.Value // []string
// dropperQuarantines records exact snapshots of files CSM moved out of
// their original paths, so the later absence probe does not misclassify
// CSM's own remediation as attacker self-deletion.
dropperQuarantines *dropperQuarantineLedger
// dropperOverflowReported is the last cumulative tracker-overflow count
// surfaced to the operator. Only the probe goroutine accesses it.
dropperOverflowReported uint64
lastDropperOverflowReport time.Time
// dropperHandlerCache shares immutable inherited .htaccess PHP mappings
// across analyzer events. The generation changes as soon as a .htaccess
// event reaches the reader, invalidating every inherited snapshot.
dropperHandlerMu sync.Mutex
dropperHandlerCache map[string]dropperPHPHandlerCacheEntry
dropperHandlerGeneration uint64
}
const (
// reconcileDirCap bounds the dirty-region tracker fed by sendEvent's
// drop branch. A 2026-04-28 cpanel package restore overflowed the
// previous 64-entry cap inside seconds (every wp-content subdir was
// a distinct parent), evicting older dirs before reconcileDrops ran.
// 1024 entries fits typical cpanel restore bursts while bounding the
// retained paths, scan cursors and queue tickets.
reconcileDirCap = 1024
// reconcileWindow scopes which files reconcileDrops will rescan: only
// files whose mtime is within this window of the first pass. The cutoff
// stays fixed across retries. Sized just over
// the minute tick so a drop near the start of a tick is still picked
// up by the reconcile that runs at tick end.
reconcileWindow = 70 * time.Second
// reconcileBudget bounds one recovery pass. The pass reads whole
// directories and scans their recent files inline, outside the analyzer
// pool's backpressure, so an unbounded pass competes with the very
// workers whose backlog caused the drops.
reconcileBudget = 10 * time.Second
// reconcileMinInterval is the shortest gap between recovery passes. The
// eager trigger fires once per few hundred drops, which during a
// sustained storm meant a pass every few seconds: recovery rescanned
// files that were still being written, spent CPU the analyzers needed,
// and produced the next round of drops.
reconcileMinInterval = 30 * time.Second
// analyzerChBufferSize sizes the channel feeding the analyzer pool.
// A cpanel package restore in production observed ~4189 events in a
// few seconds; a 16 KiB buffer absorbs that burst plus headroom
// without ever overflowing. Memory cost is bounded (fileEvent is a
// path string + fd + pid, ~40 bytes each, so <1 MiB at full buffer).
analyzerChBufferSize = 16384
// analyzerDrainBudget bounds how long the analyzer pool keeps scanning
// queued events after shutdown starts. The queue is deep on purpose, so
// a host that is saturated at SIGTERM has thousands of events left; the
// unit gives the daemon 90 seconds in total, and draining that backlog
// to completion spent all of it and ended in SIGKILL.
analyzerDrainBudget = 15 * time.Second
// eagerReconcileDropThreshold triggers an out-of-cycle reconcile
// when sustained drops cross this count within a single minute tick.
// Without this, drops happening just after a tick wait the full
// interval before reconcileDrops walks them - long enough for the
// reconcileWindow to expire on the earliest dropped files.
eagerReconcileDropThreshold = 500
)
// Package-level Prometheus metrics for fanotify. Instantiated once per
// process; one FileMonitor per daemon instance reuses them.
var (
fanotifyEventsTotal *metrics.Counter
fanotifyEventsAdmittedTotal *metrics.Counter
fanotifyDroppedTotal *metrics.Counter
fanotifyKernelOverflowTotal *metrics.Counter
fanotifyReconcileDur *metrics.Histogram
contentScanTruncated *metrics.CounterVec
)
// registerFanotifyMetrics is called once per FileMonitor via
// fm.metricsOnce. Safe to call multiple times at the FileMonitor
// layer; the package-level sync.Once guards the actual registrations.
var fanotifyMetricsInit sync.Once
func (fm *FileMonitor) registerMetrics() {
fm.metricsOnce.Do(func() {
fanotifyMetricsInit.Do(func() {
fanotifyEventsTotal = metrics.NewCounter(
"csm_fanotify_events_total",
"Delivered fanotify file events with completed admission decisions, including events rejected before analysis. Excludes kernel queue overflow notifications, which have their own counter.",
)
metrics.MustRegister("csm_fanotify_events_total", fanotifyEventsTotal)
fanotifyEventsAdmittedTotal = metrics.NewCounter(
"csm_fanotify_events_admitted_total",
"Fanotify file events queued for analysis, including dropper tracking. Subtract this and csm_fanotify_events_dropped_total from csm_fanotify_events_total to count events rejected before queue admission.",
)
metrics.MustRegister("csm_fanotify_events_admitted_total", fanotifyEventsAdmittedTotal)
fanotifyDroppedTotal = metrics.NewCounter(
"csm_fanotify_events_dropped_total",
"Fanotify events dropped because the analyzer queue was full. Sustained growth indicates an event storm (bulk unzip, backup restore) or an attack producing more file activity than the scanner can analyse; the reconcile pass still rescans affected directories, so dropped events do not vanish from detection, they arrive delayed.",
)
metrics.MustRegister("csm_fanotify_events_dropped_total", fanotifyDroppedTotal)
fanotifyKernelOverflowTotal = metrics.NewCounter(
"csm_fanotify_kernel_queue_overflow_total",
"FAN_Q_OVERFLOW events: the kernel fanotify notification queue filled and dropped events before userspace read them. Unlike analyzer-queue drops these carry no fd, so the affected files are unknown; the next scheduled deep scan is the backstop. Sustained growth means a storm (bulk unzip, backup restore) or an attack producing more file activity than the reader can drain.",
)
metrics.MustRegister("csm_fanotify_kernel_queue_overflow_total", fanotifyKernelOverflowTotal)
fanotifyReconcileDur = metrics.NewHistogram(
"csm_fanotify_reconcile_latency_seconds",
"How long the post-overflow reconcile pass takes to walk drop-affected directories and rescan recent files. Buckets sized for the observed range; alert if p99 crosses tens of seconds (reconcile is stealing CPU from real-time analysis).",
[]float64{0.01, 0.05, 0.1, 0.5, 1, 5, 10, 30, 60},
)
metrics.MustRegister("csm_fanotify_reconcile_latency_seconds", fanotifyReconcileDur)
metrics.RegisterGaugeFunc(
"csm_fanotify_queue_depth",
"Current number of queued fanotify events waiting for the analyzer pool. Capacity is 4000; queue approaching that cap means drops are imminent.",
func() float64 {
if fm == nil || fm.analyzerCh == nil {
return 0
}
return float64(len(fm.analyzerCh))
},
)
contentScanTruncated = metrics.NewCounterVec(
"csm_realtime_content_scan_truncated_total",
"Real-time fanotify content checks whose file was larger than the main read window, so the full-rule pass saw only the leading window. Labels: check (phpcontent_inline, phpcontent_uploads, php_check, crontab, htaccess, user_ini, html_phishing, cgi_backdoor).",
[]string{"check"},
)
metrics.MustRegister("csm_realtime_content_scan_truncated_total", contentScanTruncated)
})
})
}
// recordReadTruncation increments csm_realtime_content_scan_truncated_total
// when the file behind fd is larger than maxBytes. Cheap fstat per scan.
// No-op if the counter has not been registered (test setups that skip
// registerMetrics).
func recordReadTruncation(fd int, maxBytes int, check string) {
if contentScanTruncated == nil {
return
}
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err != nil {
return
}
if st.Size > int64(maxBytes) {
contentScanTruncated.With(check).Inc()
}
}
type fileEvent struct {
queueTicket queuehealth.Ticket
path string
fd int
pid int32
mask uint64
dropperOnly bool
phpExecutable bool
}
func (fm *FileMonitor) currentCfg() *config.Config {
if cfg := config.Active(); cfg != nil {
return cfg
}
if fm == nil {
return nil
}
return fm.cfg
}
// NewFileMonitor creates a fanotify-based file monitor.
// Returns error if the kernel doesn't support the required features.
func NewFileMonitor(cfg *config.Config, alertCh chan<- alert.Finding) (*FileMonitor, error) {
// H1 - use golang.org/x/sys/unix for fanotify_init
fd, err := unix.FanotifyInit(FAN_CLASS_NOTIF|FAN_CLOEXEC|FAN_NONBLOCK, unix.O_RDONLY)
if err != nil {
return nil, fmt.Errorf("fanotify_init: %w (kernel may not support fanotify)", err)
}
// Mark mount points; M2 - track successful mounts
webRootPatterns := checks.PHPConfigRealtimeRootPatterns(cfg)
mountPaths := fanotifyMountPaths(webRootPatterns)
marks, markErr := markWatchRoots(fd, mountPaths, unix.FanotifyMark, statDevice)
if len(marks) == 0 {
_ = unix.Close(fd)
return nil, fmt.Errorf("no mount points could be watched (tried %v): %w", mountPaths, markErr)
}
if markErr != nil {
fmt.Fprintf(os.Stderr, "[%s] Warning: %v\n", ts(), markErr)
}
var mountScoped []string
for _, m := range marks {
if m.coveredBy == "" && m.scope == markScopeMount {
mountScoped = append(mountScoped, m.path)
}
}
if len(mountScoped) > 0 {
// Worth saying out loud: on these roots a write that arrives through a
// bind mount (a CageFS cage) raises no event, and only the rolling
// content scan will meet it.
fmt.Fprintf(os.Stderr, "[%s] Warning: watching %v per-mount only; writes through bind mounts on them are not seen in real time\n",
ts(), mountScoped)
}
// Directory-scoped watch on /var/spool/cron so any user crontab write
// reaches analyzeFile in real time. Best-effort: cron may live under a
// different path on non-cPanel hosts (the platform layer normalises),
// and we'd rather lose the realtime crontab signal than fail daemon
// startup. The polled CheckCrontabs run still covers this case via
// the next scheduled scan. Mask matches spoolwatch.go (the proven
// production pattern for directory-scoped marks): FAN_CLOSE_WRITE
// alone catches both new and modified crontabs, since the close
// after O_CREAT|O_WRONLY|... fires the close-write event. FAN_CREATE
// is omitted because it has stricter kernel requirements with
// directory-scoped (non-MOUNT) marks and adds no coverage here.
if _, statErr := os.Stat(cronSpoolDir()); statErr == nil {
if err := unix.FanotifyMark(fd, FAN_MARK_ADD,
FAN_CLOSE_WRITE|FAN_EVENT_ON_CHILD, -1, cronSpoolDir()); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Warning: cannot watch %s: %v\n", ts(), cronSpoolDir(), err)
}
}
// C4 - create pipe for epoll stop signaling
var pipeFds [2]int
if err := unix.Pipe2(pipeFds[:], unix.O_NONBLOCK|unix.O_CLOEXEC); err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("pipe2: %w", err)
}
fm := &FileMonitor{
fd: fd,
watchRoots: marks,
cfg: cfg,
alertCh: alertCh,
analyzerCh: make(chan fileEvent, analyzerChBufferSize),
pipeFds: pipeFds,
stopCh: make(chan struct{}),
reconcileDirs: make(map[string]reconcileDirectory),
reconcileSig: make(chan struct{}, 1),
webRootPatterns: webRootPatterns,
accountRootPatterns: checks.AccountHomePatterns(),
docRootPatterns: checks.RealtimeDocumentRootPatterns(cfg),
}
wpCache := wpcheck.NewCache(cfg.StatePath)
wpCache.SetStopCh(fm.stopCh)
fm.wpCache = wpCache
fm.wpPending = newStagedPackageQueue(stagedPackageQueueMax)
fm.initDropperDetector(cfg)
return fm, nil
}
func fanotifyMountPaths(webRootPatterns []string) []string {
paths := []string{"/home", "/tmp", "/dev/shm", "/var/tmp"}
seen := make(map[string]struct{}, len(paths)+len(webRootPatterns))
for _, path := range paths {
seen[path] = struct{}{}
}
for _, pattern := range webRootPatterns {
anchor := fanotifyMountAnchor(pattern)
if anchor == "" {
continue
}
if _, exists := seen[anchor]; exists {
continue
}
seen[anchor] = struct{}{}
paths = append(paths, anchor)
}
return paths
}
func fanotifyMountAnchor(pattern string) string {
if strings.TrimSpace(pattern) == "" {
return ""
}
clean := filepath.Clean(pattern)
if !filepath.IsAbs(clean) {
return ""
}
meta := strings.IndexAny(clean, "*?[")
if meta < 0 {
return clean
}
prefix := clean[:meta]
if strings.HasSuffix(prefix, string(filepath.Separator)) {
return filepath.Clean(prefix)
}
return filepath.Dir(prefix)
}
const (
// minAnalyzerWorkers keeps one write from stalling every other event
// while it is scanned, without putting four content scans on a host that
// has two cores to serve requests with.
minAnalyzerWorkers = 2
// maxAnalyzerWorkers caps the pool on large hosts: past this the workers
// contend on disk rather than finishing sooner.
maxAnalyzerWorkers = 16
)
// analyzerWorkerCount sizes the analyzer pool from the host's core count.
func analyzerWorkerCount(cpus int) int {
if cpus < minAnalyzerWorkers {
return minAnalyzerWorkers
}
if cpus > maxAnalyzerWorkers {
return maxAnalyzerWorkers
}
return cpus
}
// Run starts the file monitor event loop and analyzer workers.
func (fm *FileMonitor) Run(stopCh <-chan struct{}) {
numWorkers := analyzerWorkerCount(runtime.NumCPU())
for i := 0; i < numWorkers; i++ {
fm.wg.Add(1)
obs.Go("fanotify-analyzer", fm.analyzerWorker)
}
// Start overflow reporter
fm.wg.Add(1)
obs.Go("fanotify-overflow", fm.overflowReporter)
// Resolve staged WordPress package files once their checksums land.
fm.wg.Add(1)
obs.Go("fanotify-wp-package", fm.stagedPackageLoop)
// Start the self-deleting-dropper probe loop when the detector is enabled.
if fm.dropper != nil {
fm.wg.Add(1)
obs.Go("fanotify-dropper", fm.dropperProbeLoop)
}
// C4 - create epoll instance, watch fanotify fd + pipe read end
epfd, err := unix.EpollCreate1(unix.EPOLL_CLOEXEC)
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] epoll_create1 failed: %v, falling back to poll loop\n", ts(), err)
fm.runPollFallback(stopCh)
return
}
defer func() { _ = unix.Close(epfd) }()
// Add fanotify fd to epoll
if err := unix.EpollCtl(epfd, unix.EPOLL_CTL_ADD, fm.fd, &unix.EpollEvent{
Events: unix.EPOLLIN,
// #nosec G115 -- POSIX fd fits in int32 (rlimit caps fds at ~1024).
Fd: int32(fm.fd),
}); err != nil {
fmt.Fprintf(os.Stderr, "[%s] epoll_ctl(fanotify): %v\n", ts(), err)
fm.runPollFallback(stopCh)
return
}
// Add pipe read end to epoll (for stop signaling)
if err := unix.EpollCtl(epfd, unix.EPOLL_CTL_ADD, fm.pipeFds[0], &unix.EpollEvent{
Events: unix.EPOLLIN,
// #nosec G115 -- POSIX fd fits in int32.
Fd: int32(fm.pipeFds[0]),
}); err != nil {
fmt.Fprintf(os.Stderr, "[%s] epoll_ctl(pipe): %v\n", ts(), err)
fm.runPollFallback(stopCh)
return
}
// Forward external stopCh to our internal mechanism
obs.SafeGo("fanotify-stop-forward", func() {
select {
case <-stopCh:
fm.Stop()
case <-fm.stopCh:
}
})
buf := make([]byte, 4096*24) // Large buffer for event batches
events := make([]unix.EpollEvent, 4)
for {
n, err := unix.EpollWait(epfd, events, 500) // 500ms timeout
if err != nil {
if err == unix.EINTR {
continue
}
// Check if we've been stopped
select {
case <-fm.stopCh:
fm.drainAndClose()
return
default:
}
fmt.Fprintf(os.Stderr, "[%s] epoll_wait error: %v\n", ts(), err)
time.Sleep(1 * time.Second)
continue
}
// Check for stop first
select {
case <-fm.stopCh:
fm.drainAndClose()
return
default:
}
for i := 0; i < n; i++ {
// #nosec G115 -- POSIX fd fits in int32; comparing against epoll event fd.
if events[i].Fd == int32(fm.pipeFds[0]) {
// Stop signal received via pipe
fm.drainAndClose()
return
}
// #nosec G115 -- POSIX fd fits in int32.
if events[i].Fd == int32(fm.fd) {
// fanotify events ready — single read per epoll wake
fm.initQueueHealth()
_, readErr := fm.kernelQueue.read(buf, fm.processEvents)
if readErr != nil {
if readErr != unix.EAGAIN && readErr != unix.EINTR {
fmt.Fprintf(os.Stderr, "[%s] fanotify read error: %v\n", ts(), readErr)
}
}
}
}
}
}
// runPollFallback is used when epoll setup fails; falls back to sleep-based polling.
func (fm *FileMonitor) runPollFallback(stopCh <-chan struct{}) {
// Forward external stopCh to our internal mechanism
obs.SafeGo("fanotify-stop-forward", func() {
select {
case <-stopCh:
fm.Stop()
case <-fm.stopCh:
}
})
buf := make([]byte, 4096*24)
for {
select {
case <-fm.stopCh:
fm.drainAndClose()
return
default:
}
fm.initQueueHealth()
_, err := fm.kernelQueue.read(buf, fm.processEvents)
if err != nil {
if err == unix.EAGAIN || err == unix.EINTR {
time.Sleep(100 * time.Millisecond)
continue
}
fmt.Fprintf(os.Stderr, "[%s] fanotify read error: %v\n", ts(), err)
time.Sleep(1 * time.Second)
continue
}
}
}
// processEvents parses a buffer of fanotify event metadata and dispatches each event.
func (fm *FileMonitor) processEvents(buf []byte) {
offset := 0
for offset+metadataSize <= len(buf) {
// #nosec G103 -- fanotify delivers a packed binary stream on the
// fd; we must reinterpret the byte buffer as the kernel struct.
// The metadataSize bounds check above guarantees we have enough
// bytes for the struct.
event := (*fanotifyEventMetadata)(unsafe.Pointer(&buf[offset]))
eventLen := int(event.EventLen)
if eventLen < metadataSize || offset+eventLen > len(buf) {
break
}
if event.Mask&unix.FAN_Q_OVERFLOW != 0 {
fm.handleQueueOverflow()
} else if event.Fd >= 0 {
fm.handleEvent(int(event.Fd), event.Pid, event.Mask)
}
offset += eventLen
}
}
// handleQueueOverflow reacts to a FAN_Q_OVERFLOW record. The kernel dropped
// events we will never see and gave us no fd, so we cannot reconcile the exact
// files. Count it, emit a rate-limited Warning so operators learn coverage was
// lost, and nudge the reconcile pass to rescan directories that also saw
// analyzer-queue drops during the same storm.
func (fm *FileMonitor) handleQueueOverflow() {
fm.initQueueHealth()
fm.kernelQueueHealth.Lose(time.Now(), 1)
atomic.AddInt64(&fm.queueOverflows, 1)
if fanotifyKernelOverflowTotal != nil {
fanotifyKernelOverflowTotal.Inc()
}
fm.reportQueueOverflow()
if fm.reconcileSig != nil {
select {
case fm.reconcileSig <- struct{}{}:
default:
}
}
}
// reportQueueOverflow emits the kernel-overflow finding at most once per minute.
func (fm *FileMonitor) reportQueueOverflow() {
fm.overflowReportMu.Lock()
if !fm.lastOverflowReport.IsZero() && time.Since(fm.lastOverflowReport) < time.Minute {
fm.overflowReportMu.Unlock()
return
}
fm.lastOverflowReport = time.Now()
fm.overflowReportMu.Unlock()
fm.sendAlert(alert.Warning, "fanotify_kernel_overflow",
"fanotify kernel event queue overflowed; file events were dropped by the kernel and cannot be recovered by path",
"The kernel notification queue filled during a storm (bulk unzip, backup restore, or high-volume attack) and dropped events before the reader could drain them. Files touched during the overflow that are not written again are only covered by the next scheduled deep scan. A reconcile of directories that also saw analyzer-queue drops has been triggered.")
}
// drainAndClose drains the analyzerCh and waits for workers to finish.
// C1 - ensures no fd leak on shutdown.
func (fm *FileMonitor) drainAndClose() {
fm.drainOnce.Do(func() {
fm.drainDeadline.Store(time.Now().Add(analyzerDrainBudget).UnixNano())
close(fm.analyzerCh)
fm.wg.Wait()
fm.discardReconcilePending()
fm.stagedPackages().discardPending(time.Now())
if fm.dropper != nil {
fm.dropper.tr.discardPending(time.Now())
clear(fm.dropper.attempts)
}
// Mark pipe as closed before actually closing, so Stop() won't
// write to an already-closed fd (H2 fix).
atomic.StoreInt32(&fm.pipeClosed, 1)
_ = unix.Close(fm.pipeFds[0])
_ = unix.Close(fm.pipeFds[1])
})
}
// Stop signals the monitor to shut down.
// C2 - sync.Once ensures safe concurrent calls; does not close analyzerCh directly.
func (fm *FileMonitor) Stop() {
fm.stopOnce.Do(func() {
close(fm.stopCh)
// Wake epoll so Run() exits and calls drainAndClose.
// Only write if pipe hasn't been closed by drainAndClose yet.
if atomic.LoadInt32(&fm.pipeClosed) == 0 {
_, _ = unix.Write(fm.pipeFds[1], []byte{0})
}
// Close fanotify fd - causes any pending Read/EpollWait to return
fm.initQueueHealth()
_ = fm.kernelQueue.close()
})
}
func (fm *FileMonitor) handleEvent(fd int, pid int32, mask uint64) {
var admitted, dropped bool
defer func() { fm.recordEvent(admitted, dropped) }()
// Get the file path from the fd via /proc/self/fd/N
path, err := os.Readlink(fmt.Sprintf("/proc/self/fd/%d", fd))
if err != nil {
_ = unix.Close(fd)
return
}
path = normalizeFanotifyEventPath(path)
// M5 - skip directory events
if strings.HasSuffix(path, "/") {
_ = unix.Close(fd)
return
}
// The dropper tracker also needs arbitrary executable names that have
// no path-only content signal.
fm.invalidateDropperPHPHandlerCache(path)
contentInteresting := fm.isInteresting(path)
// Filtering must not reroute writes retained for dropper tracking: the
// dropper-only path bypasses normal suppressions and checksum verification.
contentNeedsAnalysis := contentInteresting
if contentInteresting && underTempRoot(path) {
contentNeedsAnalysis = fm.tempRootEventNeedsAnalysis(path, fd)
}
dropperInteresting, phpExecutable := fm.isDropperInteresting(path, fd)
if !contentNeedsAnalysis && !dropperInteresting {
_ = unix.Close(fd)
return
}
// Send to analyzer pool (with backpressure)
fm.initQueueHealth()
ticket := fm.analyzerHealth.Begin(time.Now())
select {
case fm.analyzerCh <- fileEvent{
queueTicket: ticket,
path: path, fd: fd, pid: pid, mask: mask,
dropperOnly: !contentInteresting, phpExecutable: phpExecutable,
}:
admitted = true
default:
// Queue full - drop event, count, and record the parent dir so the
// reconcile pass in overflowReporter can rescan it. Without this
// every file in a bulk burst past buffer capacity is invisible to
// detection forever.
ticket.Reject(time.Now())
dropped = true
n := atomic.AddInt64(&fm.droppedEvents, 1)
if n%100 == 0 {
fmt.Fprintf(os.Stderr, "[%s] fanotify: %d events dropped (analyzer queue full)\n", ts(), n)
}
fm.recordDroppedDir(path)
fm.maybeTriggerEagerReconcile(n)
_ = unix.Close(fd)
}
}
// normalizeFanotifyEventPath removes procfs's synthetic " (deleted)" suffix
// when the event fd's directory entry was already unlinked. Linux does not
// disambiguate that marker from a literal filename suffix. Treat it as the
// kernel marker: otherwise an attacker can create a same-inode hardlink with
// the literal suffix and make an immediate self-delete miss the PHP path gate.
func normalizeFanotifyEventPath(path string) string {
const deletedSuffix = " (deleted)"
return strings.TrimSuffix(path, deletedSuffix)
}
// maybeTriggerEagerReconcile nudges overflowReporter to run reconcileDrops
// immediately when sustained drops cross eagerReconcileDropThreshold within
// a single minute window. Delegates to the free function so the trigger
// logic stays testable from a cross-platform test file.
func (fm *FileMonitor) maybeTriggerEagerReconcile(droppedSoFar int64) {
signalEagerReconcile(fm.reconcileSig, droppedSoFar, eagerReconcileDropThreshold)
}
// isInteresting is the fast filter - zero I/O, pure string matching.
func (fm *FileMonitor) isInteresting(path string) bool {
path = atomicWriteContentPath(path)
lower := strings.ToLower(path)
// PHP source files. This is intentionally broader than the executable-PHP
// predicate used by the location and dropper checks: .phps is inert under a
// stock handler, but still needs signature/YARA analysis while staged.
if isPHPSourceExtension(filepath.Base(lower)) {
return true
}
// Webshell extensions
if strings.HasSuffix(lower, ".haxor") || strings.HasSuffix(lower, ".cgix") {
return true
}
// CGI scripts in hosted trees - detect Perl/Python/Bash backdoors.
// An explicit document root may live outside the platform's account homes.
if fm.underAccountOrConfiguredDocRoot(path) {
if strings.HasSuffix(lower, ".pl") || strings.HasSuffix(lower, ".cgi") ||
strings.HasSuffix(lower, ".py") || strings.HasSuffix(lower, ".sh") ||
strings.HasSuffix(lower, ".rb") {
return true
}
}
// .htaccess and .user.ini files (any location), and php.ini under a
// configured or detected web root. An attacker plants php.ini files there
// to weaken disable_functions.
if strings.HasSuffix(lower, ".htaccess") || strings.HasSuffix(lower, ".user.ini") {
return true
}
if filepath.Base(lower) == "php.ini" && pathMatchesWebRootPatterns(path, fm.webRootPatterns) {
return true
}
// HTML files in an account or explicitly configured document tree.
if fm.underAccountOrConfiguredDocRoot(path) &&
(strings.HasSuffix(lower, ".html") || strings.HasSuffix(lower, ".htm")) {
return true
}
// Images in an account or explicitly configured document tree. A real
// image container is a working payload store: PHP appended to a valid
// PNG still opens as a picture, and a one-line include elsewhere in the
// site executes it. Admitting the write is what lets checkImagePayload
// look at the bytes; the extension only routes the event, and the
// verdict comes from the container magic, so a renamed payload is still
// caught by the other branches.
if fm.underAccountOrConfiguredDocRoot(path) && contenttype.IsImageExt(filepath.Ext(lower)) {
return true
}
// Credential log files - known phishing harvest filenames
base := filepath.Base(lower)
if credentialLogNames[base] {
return true
}
// ZIP archives in an account or explicitly configured document tree.
if fm.underAccountOrConfiguredDocRoot(path) && strings.HasSuffix(lower, ".zip") {
return true
}
// Anything in .config directories
if strings.Contains(path, "/.config/") {
return true
}
// User crontabs surfaced via the directory-scoped fanotify mark in
// NewFileMonitor. Each write to /var/spool/cron/<user> dispatches to
// checkCrontab in real time.
if strings.HasPrefix(path, cronSpoolDir()+"/") {
return true
}
// Writes in the shared temporary trees. The path alone cannot say whether
// one carries signal; the reader decides that from the event descriptor in
// tempRootEventNeedsAnalysis.
if underTempRoot(path) {
return true
}
// PHP in sensitive directories that should never contain PHP
if inSensitivePHPDir(path) && isPHPExtension(strings.ToLower(filepath.Base(path))) {
return true
}
return false
}
// underAccountRoot reports whether path sits inside a hosting account's tree.
// Falls back to the historical /home spelling when the platform offers no
// patterns, so an unconfigured plain-Linux host keeps the behaviour it had.
func (fm *FileMonitor) underAccountRoot(path string) bool {
if len(fm.accountRootPatterns) == 0 {
return strings.HasPrefix(path, "/home/")
}
return pathMatchesWebRootPatterns(path, fm.accountRootPatterns)
}
// underDocRoot reports whether path sits inside a served document root.
func (fm *FileMonitor) underDocRoot(path string) bool {
if len(fm.docRootPatterns) == 0 {
return strings.Contains(path, "/public_html/")
}
return pathMatchesWebRootPatterns(path, fm.docRootPatterns)
}
func (fm *FileMonitor) underAccountOrConfiguredDocRoot(path string) bool {
return fm.underAccountRoot(path) ||
(len(fm.docRootPatterns) > 0 && fm.underDocRoot(path))
}
// tempRootPrefixes are the shared temporary trees. Anything may be written
// there by anyone, so they are watched, but most of what lands there carries no
// signal any detector acts on.
var tempRootPrefixes = []string{"/tmp/", "/dev/shm/", "/var/tmp/"}
// underTempRoot reports whether path is inside one of those trees.
func underTempRoot(path string) bool {
for _, prefix := range tempRootPrefixes {
if strings.HasPrefix(path, prefix) {
return true
}
}
return false
}
// sensitivePHPDirs are directories that should never contain PHP, so a PHP
// file in one is reported wherever the directory itself lives.
var sensitivePHPDirs = []string{"/.ssh/", "/.cpanel/", "/mail/", "/.gnupg/", "/.cagefs/"}
func inSensitivePHPDir(path string) bool {
for _, dir := range sensitivePHPDirs {
if strings.Contains(path, dir) {
return true
}
}
return false
}
// tempRootEventNeedsAnalysis decides from the descriptor the reader already
// holds whether a write under a temp root can produce a finding. The analyzer
// reaches the same verdict, but only after the event has taken a queue slot, a
// descriptor and a worker wake-up. Every branch here mirrors a check the
// analyzer runs before or inside its own temp-root branch; an event this
// rejects is one the analyzer would have returned on without reporting.
//
// A descriptor it cannot stat is admitted: losing the mode is not a reason to
// stop looking at a file.
func (fm *FileMonitor) tempRootEventNeedsAnalysis(path string, fd int) bool {
// Judge the file the write is producing, not the staging name it is being
// written under: an editor saving .htaccess writes .temp.1..htaccess first,
// and both isInteresting and the analyzer resolve that before deciding.
content := atomicWriteContentPath(path)
lower := strings.ToLower(content)
name := filepath.Base(lower)
switch {
case strings.HasPrefix(path, cronSpoolDir()+"/"),
isPHPSourceExtension(name),
name == ".htaccess", name == ".user.ini", name == "php.ini",
knownWebshells[name],
strings.HasSuffix(lower, ".haxor"), strings.HasSuffix(lower, ".cgix"),
strings.Contains(content, "/.config/"),
isPHPExtension(name) && inSensitivePHPDir(content),
contenttype.IsImageExt(filepath.Ext(lower)) && fm.underAccountOrConfiguredDocRoot(content):
return true
}
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err != nil {
return true
}
return st.Mode&unix.S_IFMT != unix.S_IFDIR && st.Mode&0o111 != 0
}
// credentialLogNames are filenames commonly used by phishing kits to store
// harvested credentials. Checked in isInteresting() for real-time detection.
var credentialLogNames = map[string]bool{
"results.txt": true, "result.txt": true, "log.txt": true,
"logs.txt": true, "emails.txt": true, "data.txt": true,
"passwords.txt": true, "creds.txt": true, "credentials.txt": true,
"victims.txt": true, "output.txt": true, "harvested.txt": true,
"results.log": true, "emails.log": true, "data.log": true,
"results.csv": true, "emails.csv": true, "data.csv": true,
}
// analyzerWorker processes file events from the bounded channel.
// C1 - on channel close, drains remaining events and closes their fds.
// Past the shutdown drain budget the remaining events are released without
// being scanned: their files are still covered by the next scheduled deep
// scan, while a backlog scanned to completion costs the daemon its whole
// stop timeout and ends in SIGKILL.
func (fm *FileMonitor) analyzerWorker() {
defer fm.wg.Done()
for event := range fm.analyzerCh {
if fm.drainBudgetSpent(time.Now()) {
event.queueTicket.Reject(time.Now())
_ = unix.Close(event.fd)
continue
}
fm.analyzeFileSafe(event)
_ = unix.Close(event.fd)
}
}
// drainBudgetSpent reports whether the shutdown drain has run out of time.
func (fm *FileMonitor) drainBudgetSpent(now time.Time) bool {
deadline := fm.drainDeadline.Load()
return deadline != 0 && now.UnixNano() >= deadline
}
// fileAnalyzer analyzes one queued event. Var so tests can substitute a
// panicking analyzer.
var fileAnalyzer = (*FileMonitor).analyzeFile
// analyzeFileSafe runs one event and contains a panic: one crafted file
// must not restart the daemon and reopen the detection gap for every other
// write in flight. The caller still closes the event fd.
func (fm *FileMonitor) analyzeFileSafe(event fileEvent) {
defer func() {
if r := recover(); r != nil {
fm.reportScannerPanic(event.path, r)
}
}()
work := queuehealth.Work[fileEvent]{Value: event, Ticket: event.queueTicket}
work.Process(func(queued fileEvent) { fileAnalyzer(fm, queued) })
}
// reportScannerPanic logs the panic with its stack, forwards it to
// observability and raises a critical finding at most once per ten minutes.
func (fm *FileMonitor) reportScannerPanic(path string, r interface{}) {
obs.CaptureMsg("fanotify-analyzer", fmt.Sprintf("panic analyzing %s: %v", path, r))
fmt.Fprintf(os.Stderr, "[%s] file monitor: recovered panic analyzing %s: %v\n%s", ts(), path, r, debug.Stack())
fm.panicMu.Lock()
defer fm.panicMu.Unlock()
if !fm.lastPanicAt.IsZero() && time.Since(fm.lastPanicAt) < 10*time.Minute {
return
}
fm.lastPanicAt = time.Now()
fm.sendAlertWithPath(alert.Critical, "realtime_scanner_panic",
fmt.Sprintf("Realtime scanner panicked on %s and skipped it: %v", path, r), "", path, "")
}
// readFromFd reads up to maxBytes from a file descriptor at position 0.
// C3 - avoids TOCTOU by reading from the original fanotify event fd.
// readFromFd reads up to maxBytes from the fanotify event fd using pread
// at offset 0. Uses unix.Pread directly to avoid os.NewFile's GC finalizer
// which would close the fd out-of-band, racing with the worker's explicit close.
func readFromFd(fd int, maxBytes int) []byte {
buf := make([]byte, maxBytes)
n, err := unix.Pread(fd, buf, 0)
if n <= 0 || (err != nil && n == 0) {
return nil
}
return buf[:n]
}
const readCompleteMaxInterrupts = 8
// readExactSize reads exactly the snapshotted size. Its buffer is fixed before
// the first read, so a concurrently growing source cannot extend the loop; a
// bounded EINTR retry count also prevents a pathological signal storm from
// pinning an analyzer worker.
func readExactSize(size int64, maxBytes int, pread func([]byte, int64) (int, error)) []byte {
if size <= 0 || maxBytes <= 0 || size > int64(maxBytes) {
return nil
}
buf := make([]byte, int(size))
interrupts := 0
for off := 0; off < len(buf); {
n, err := pread(buf[off:], int64(off))
if n < 0 || n > len(buf)-off {
return nil
}
if n > 0 {
off += n
interrupts = 0
}
if err != nil && !errors.Is(err, unix.EINTR) {
return nil
}
if n > 0 {
continue
}
if !errors.Is(err, unix.EINTR) {
return nil
}
interrupts++
if interrupts > readCompleteMaxInterrupts {
return nil
}
}
return buf
}
func sameReadSnapshot(before, after unix.Stat_t) bool {
return before.Dev == after.Dev && before.Ino == after.Ino && before.Size == after.Size &&
before.Mtim == after.Mtim && before.Ctim == after.Ctim
}
// readCompleteFromFd returns a stable snapshot of the entire file behind fd
// when it fits within maxBytes. A short read, concurrent size/content change,
// or excessive interruption fails closed so whole-file recognizers never
// accept a stale prefix.
func readCompleteFromFd(fd, maxBytes int) []byte {
var before unix.Stat_t
if err := unix.Fstat(fd, &before); err != nil {
return nil
}
buf := readExactSize(before.Size, maxBytes, func(p []byte, off int64) (int, error) {
return unix.Pread(fd, p, off)
})
if buf == nil {
return nil
}
var after unix.Stat_t
if err := unix.Fstat(fd, &after); err != nil || !sameReadSnapshot(before, after) {
return nil
}
return buf
}
func isBenignPHPStubData(fd int, data []byte) bool {
if len(data) == 0 {
return false
}
complete := false
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err == nil {
complete = st.Size <= int64(len(data))
}
return checks.IsBenignPHPStubBytesComplete(data, complete)
}
func isWPTranslationCacheData(fd int, data []byte) bool {
if len(data) == 0 {
return false
}
complete := false
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err == nil {
complete = st.Size <= int64(len(data))
}
return checks.IsWPTranslationCacheBytesComplete(data, complete)
}
// readTailFromFd reads the last maxBytes of a file via its fd using pread.
// Returns nil if the file is smaller than maxBytes (head scan already covers it).
func readTailFromFd(fd int, maxBytes int) []byte {
var stat unix.Stat_t
if err := unix.Fstat(fd, &stat); err != nil {
return nil
}
size := stat.Size
if size <= int64(maxBytes) {
return nil // head read already covers the entire file
}
offset := size - int64(maxBytes)
buf := make([]byte, maxBytes)
n, err := unix.Pread(fd, buf, offset)
if n <= 0 || (err != nil && n == 0) {
return nil
}
return buf[:n]
}
// resolveProcessInfo reads /proc/<pid>/comm and /proc/<pid>/status
// to build a "pid=N cmd=name uid=N" string for alert enrichment.
// Returns empty string on any error (process may have exited).
func resolveProcessInfo(pid int32) string {
if pid <= 0 {
return ""
}
procDir := fmt.Sprintf("/proc/%d", pid)
// Read process name
// #nosec G304 -- /proc/<pid>/comm; kernel pseudo-FS, pid is int32 from fanotify event.
comm, err := os.ReadFile(procDir + "/comm")
if err != nil {
return ""
}
name := strings.TrimSpace(string(comm))
info := fmt.Sprintf("pid=%d cmd=%s", pid, name)
// Read UID from status to map to cPanel username
// #nosec G304 -- /proc/<pid>/status; kernel pseudo-FS.
statusData, err := os.ReadFile(procDir + "/status")
if err != nil {
return info
}
for _, line := range strings.Split(string(statusData), "\n") {
if strings.HasPrefix(line, "Uid:") {
fields := strings.Fields(line)
if len(fields) >= 2 {
info += fmt.Sprintf(" uid=%s", fields[1])
}
break
}
}
return info
}
func (fm *FileMonitor) analyzeFile(event fileEvent) {
path := event.path
contentPath := atomicWriteContentPath(path)
name := filepath.Base(contentPath)
nameLower := strings.ToLower(name)
imageInHostedTree := contenttype.IsImageExt(filepath.Ext(nameLower)) && fm.underAccountOrConfiguredDocRoot(contentPath)
// Resolve process info from PID (best-effort - process may have exited)
procInfo := resolveProcessInfo(event.pid)
// Snapshot this write as a possible self-deleting dropper before the
// per-type checks below can early-return. Reads are positional (Pread/
// Fstat/Statx) so they do not disturb the fd for later content checks.
cand := fm.observeDropperCandidate(event, procInfo)
markDropperContentSuspicious := func() {
if cand == nil {
return
}
cand.ContentSuspicious = true
if !fm.dropper.tr.Refresh(*cand) {
fm.dropper.admit(*cand)
}
}
// Some events are admitted only for dropper tracking. Handler-mapped PHP
// still needs the normal PHP scanner; arbitrary executables retain the
// strongest cheap content signal for the later deletion verdict.
if event.dropperOnly {
if event.phpExecutable {
if fm.checkPHPContent(event.fd, path, procInfo) {
markDropperContentSuspicious()
}
} else if cand != nil && looksLikePHPWebshell(cand.Head) {
markDropperContentSuspicious()
}
return
}
// H2 - suppression path matching using filepath.Match. Read the live config
// (config.Active via currentCfg) so a SIGHUP change to suppressions.ignore_paths
// takes effect without a restart, matching the rest of this analyzer.
for _, ignore := range fm.currentCfg().Suppressions.IgnorePaths {
if matchSuppression(ignore, path) {
return
}
}
// Skip unmodified WordPress core and plugin files: the hash matches the
// official wordpress.org checksums for the version the install or
// package declares. Stops signature/YARA FPs on stock code, installed or
// staged: a realtime Critical feeds inline quarantine, and a byte-for-byte
// copy of the official release is not what that is for. A cache miss
// triggers a background fetch and falls through to rule evaluation; the
// description is kept for the update-staging branch below, which judges
// a staged package by these verdicts. For atomic writes, the intended
// basename only selects the checksum entry. Trust requires hashing the
// complete original event descriptor.
var wpVerdict wpcheck.Verification
if fm.wpCache != nil {
wpVerdict = fm.wpCache.VerifyFile(event.fd, contentPath)
if wpVerdict.Verdict == wpcheck.VerdictVerified {
return
}
}
// User crontab written under /var/spool/cron/<user>. Scan content
// from the event fd via the shared deep matcher and emit Critical on
// any hit. The polled CheckCrontabs run still tracks root crontab
// hash drift, so we skip root here to avoid duplicate signal.
if strings.HasPrefix(path, cronSpoolDir()+"/") {
fm.checkCrontab(event.fd, path, procInfo)
return
}
// Location-based severity escalation: PHP in dirs that should NEVER have PHP
if isPHPExtension(nameLower) {
for _, sensitive := range sensitivePHPDirs {
if strings.Contains(path, sensitive) {
fm.sendAlertWithPath(alert.Critical, "php_in_sensitive_dir_realtime",
fmt.Sprintf("PHP file in critical directory: %s", path),
fmt.Sprintf("PHP should never exist in %s - likely webshell or backdoor", sensitive), path, procInfo)
return
}
}
}
// Known webshell filenames (M1 - package-level var). Filename alone is
// too weak: WordPress core ships wp-includes/Text/Diff/Engine/shell.php
// (the Pear Text_Diff library using shell_exec to call Unix `diff`).
// Confirm with content: the file must also exhibit webshell markers
// (request superglobal flowing into a dangerous function, or an
// eval/assert wrapping a base64/gzinflate decoder).
if knownWebshells[nameLower] {
recordReadTruncation(event.fd, 65536, "phpcontent_inline")
if data := readFromFd(event.fd, 65536); looksLikePHPWebshell(data) {
markDropperContentSuspicious()
fm.sendAlertWithPath(alert.Critical, "webshell_realtime",
fmt.Sprintf("Webshell file created: %s", path), "", path, procInfo)
return
}
}
// Webshell extensions
if strings.HasSuffix(nameLower, ".haxor") || strings.HasSuffix(nameLower, ".cgix") {
fm.sendAlertWithPath(alert.Critical, "webshell_realtime",
fmt.Sprintf("Suspicious CGI file created: %s", path), "", path, procInfo)
return
}
// .htaccess modification - check for injection (C3 - read from fd).
// Checked before the /tmp early-return so a malicious .htaccess anywhere
// (including /tmp) is still analyzed for dangerous directives.
if nameLower == ".htaccess" {
fm.checkHtaccess(event.fd, path, procInfo)
return
}
// .user.ini modification - check for dangerous PHP settings (C3 - read from fd).
// Also checked before /tmp so malicious .user.ini is detected anywhere.
if nameLower == ".user.ini" || nameLower == "php.ini" {
fm.checkUserINI(event.fd, path, procInfo)
return
}
// Executables in .config - checked before the /tmp block so a miner
// dropped at /tmp/.config/* is flagged as executable_in_config_realtime
// (more specific) rather than executable_in_tmp_realtime.
// Uses unix.Fstat on the event fd (not os.Stat by path) for TOCTOU
// safety: an attacker cannot chmod -x or swap the file after the event.
if strings.Contains(path, "/.config/") {
var cfgStat unix.Stat_t
if err := unix.Fstat(event.fd, &cfgStat); err == nil {
isDir := cfgStat.Mode&unix.S_IFMT == unix.S_IFDIR
if !isDir && cfgStat.Mode&0111 != 0 {
fm.sendAlertWithPath(alert.Critical, "executable_in_config_realtime",
fmt.Sprintf("Executable created in .config: %s", path),
fmt.Sprintf("Size: %d", cfgStat.Size), path, procInfo)
}
}
if !imageInHostedTree {
return
}
}
// Executables in /tmp or /dev/shm - detect dropped malware/miners
// Uses unix.Fstat on event fd for TOCTOU safety (attacker can't chmod -x after event)
if underTempRoot(path) {
var tmpStat unix.Stat_t
if err := unix.Fstat(event.fd, &tmpStat); err == nil {
isDir := tmpStat.Mode&unix.S_IFMT == unix.S_IFDIR
isExec := tmpStat.Mode&0111 != 0
if !isDir && isExec {
// Skip known root-owned work directories:
// - cPanel: SpamAssassin compiles .so regex modules, UPCP stages scripts
// - dracut: rebuilds initramfs after kernel updates, copies system binaries
// Non-root files in these paths are still suspicious.
isCpanelWork := strings.Contains(path, "/cpanel.TMP.work.") || strings.Contains(path, "/cPanel-")
isDracutWork := strings.Contains(path, "/dracut.")
if (isCpanelWork || isDracutWork) && tmpStat.Uid == 0 {
// Root-owned executable in system work dir - legitimate, skip
} else {
// Root-owned drops from a live package transaction
// (e.g. weak-modules extracting initramfs via cpio
// after a kernel update) are rescored to Warning,
// never suppressed. See demoteTmpExec for the gates.
severity := alert.Critical
details := fmt.Sprintf("Size: %d, Mode: %04o, UID: %d", tmpStat.Size, tmpStat.Mode&0777, tmpStat.Uid)
if ok, reason := tmpExecDemote(tmpStat.Uid, event.pid, time.Now()); ok {
severity = alert.Warning
details += " [demoted: " + reason + "]"
}
fm.sendAlertWithPath(severity, "executable_in_tmp_realtime",
fmt.Sprintf("Executable created in %s: %s", filepath.Dir(path), path),
details, path, procInfo)
}
}
}
// Hosted images still need payload checks when the configured root
// lives in a temporary directory, as do PHP source files anywhere.
if !isPHPSourceExtension(nameLower) && !imageInHostedTree {
return
}
}
// PHP in uploads directories.
// Any PHP file here is anomalous: /wp-content/uploads/ is meant for
// media, not code. Plugin-update temp dirs are recognised structurally
// via looksLikePluginUpdate only after content checks have had first
// refusal, so a decoy update directory cannot downgrade a webshell.
// Operators suppress legitimate caching daemons through the path-scoped
// suppressions_api, not an implicit substring allowlist in the daemon.
if strings.Contains(path, "/wp-content/uploads/") && isPHPExtension(nameLower) {
// Content-aware severity: PHP in uploads is anomalous but not
// always malicious (TinyMCE smile_fonts/charmap.php is glyph
// data shipped by WP's bundled editor). Emit Critical for
// direct webshell markers, otherwise run the broader PHP
// content/signature/YARA path before downgrading clean PHP to
// a Warning.
recordReadTruncation(event.fd, 65536, "phpcontent_uploads")
data := readFromFd(event.fd, 65536)
if looksLikePHPWebshell(data) {
markDropperContentSuspicious()
fm.sendAlertWithPath(alert.Critical, "php_in_uploads_realtime",
fmt.Sprintf("PHP file created in uploads: %s", path),
"Webshell markers in content (request superglobal -> dangerous function, or eval/assert + decoder chain)",
path, procInfo)
} else {
if fm.checkPHPContent(event.fd, path, procInfo) {
markDropperContentSuspicious()
return
}
// Content-shape gate: file whose reachable code is
// whitespace+comments, or that terminates with
// die/exit/__halt_compiler before any statement,
// cannot execute attacker-controlled code via web
// request. BackWPup writes its working-job and
// folder-cache state files this way. The earlier
// signature/YARA pass and the path-only warning
// below are the layers that fire on real droppers;
// a structurally inert stub adds no signal.
if isBenignPHPStubData(event.fd, data) {
return
}
if looksLikePluginUpdate(path) {
// Verified plugin update - emit one low-severity alert per temp directory.
uploadsIdx := strings.Index(path, "/wp-content/uploads/")
afterUploads := path[uploadsIdx+len("/wp-content/uploads/"):]
tempDir := afterUploads
if slashIdx := strings.Index(afterUploads, "/"); slashIdx > 0 {
tempDir = afterUploads[:slashIdx]
}
updateDir := path[:uploadsIdx] + "/wp-content/uploads/" + tempDir
fm.sendAlertWithPath(alert.Warning, "php_in_uploads_realtime",
fmt.Sprintf("Plugin update in uploads: %s", updateDir),
"Verified: matching plugin exists in plugins/", updateDir, procInfo)
return
}
// Suppress the path-only "anomalous location" warning
// when the file is structurally a duplicate (cPanel
// restore staging) or a known plugin probe shape that
// never carries executable input. Signature/YARA scans
// already ran above, so any real malicious content is
// reported through its own pipeline.
if looksLikeCpanelRestoreStaging(path) {
return
}
if looksLikeWPOptimizeProbe(path, data) {
return
}
fm.sendAlertWithPath(alert.Warning, "php_in_uploads_realtime",
fmt.Sprintf("PHP file in uploads (no webshell markers): %s", path),
"Anomalous location for PHP, but content is clean",
path, procInfo)
}
return
}
// PHP in languages/upgrade directories.
// Path-only Critical buried real alerts under location noise (WPML
// translation queues, WP auto-update staging). Run content analysis
// on every file -- a real rule fires Critical, clean real code gets a
// Warning, and inert stubs stay quiet. No filename allowlist: an attacker
// must not be able to hide a backdoor by naming it like a translation or
// index file.
if (strings.Contains(path, "/wp-content/languages/") || strings.Contains(path, "/wp-content/upgrade/")) &&
isPHPExtension(nameLower) {
if fm.checkPHPContent(event.fd, path, procInfo) {
markDropperContentSuspicious()
} else {
// Every staged file reaches content analysis first. Its original
// digest then decides the path-only warning, even for inert files
// absent from the official manifest.
if fm.handleStagedPackageFile(path, wpVerdict, procInfo) {
return
}
// Translation caches and comment-only stubs require a stable,
// complete body. A no-argument PHP terminator is safe from a
// prefix because all following bytes are unreachable, so retain
// the old bounded-head fallback for oversized files.
data := readCompleteFromFd(event.fd, checks.MaxInertPHPScanBytes)
if data != nil && checks.IsBenignPHPStubBytesComplete(data, true) {
return
}
if data == nil && checks.IsBenignPHPStubBytesComplete(readFromFd(event.fd, 65536), false) {
return
}
// WordPress 6.5+ writes *.l10n.php translation caches here as pure
// data return arrays. Suppress by content structure, not filename.
if isWPTranslationCacheData(event.fd, data) {
return
}
// A core update copies the new version file here to read the
// release requirements. Literal assignments to the version
// variables cannot run code; the whole body must be seen.
if data != nil && checks.IsWPVersionDataBytesComplete(data, true) {
return
}
fm.sendAlertWithPath(alert.Warning, "php_in_sensitive_dir_realtime",
fmt.Sprintf("PHP file created in sensitive WP directory (content clean): %s", path), "", path, procInfo)
}
return
}
// PHP content analysis (C3 - read from fd; M4 - 32KB scan size).
// .htaccess, .user.ini, and .config executable checks are handled
// earlier in this function (before the /tmp early-return) so specific
// file types take precedence over the /tmp generic block.
if isPHPSourceExtension(nameLower) {
if fm.checkPHPContent(event.fd, path, procInfo) {
markDropperContentSuspicious()
}
return
}
// HTML phishing page detection (uses event fd for content, unix.Fstat for size)
if strings.HasSuffix(nameLower, ".html") || strings.HasSuffix(nameLower, ".htm") {
fm.checkHTMLPhishing(event.fd, path, procInfo)
return
}
// PHP carried inside an image file (uses event fd for content).
if contenttype.IsImageExt(filepath.Ext(nameLower)) {
// Handler-mapped images previously entered through dropperOnly and
// received the full PHP scan. Content admission must retain it.
if event.phpExecutable && fm.checkPHPContent(event.fd, path, procInfo) {
markDropperContentSuspicious()
return
}
if fm.checkImagePayload(event.fd, path, procInfo) {
markDropperContentSuspicious()
}
return
}
// Credential log files (content read from the event fd)
if credentialLogNames[nameLower] {
fm.checkCredentialLog(event.fd, path, procInfo)
return
}
// Phishing kit ZIP archives (path-based)
if strings.HasSuffix(nameLower, ".zip") {
fm.checkPhishingZip(path, nameLower, procInfo)
return
}
// CGI scripts in web-accessible directories (Perl, Python, Bash, Ruby)
// Detect backdoor toolkits like LEVIATHAN that use non-PHP scripts.
if fm.underAccountOrConfiguredDocRoot(path) && isCGIExtension(nameLower) {
fm.checkCGIBackdoor(event.fd, path, procInfo)
return
}
}
// Structural exclusions for checkHtaccess. Both anchor to the actual
// directive or regex context, not to loose substrings that an attacker
// can paste anywhere on the line.
var (
// Legit auto_(prepend|append)_file directive targets: known product
// files shipped by security plugins. Match is anchored to the
// directive argument, so a trailing "# litespeed" comment cannot
// forge safety.
htaccessAutoPrependSafeTarget = regexp.MustCompile(
`(?i)auto_(?:prepend|append)_file\s*=?\s*['"]?(?:[^\s'"]*/)?` +
`(?:wordfence-waf|sucuri|advanced-headers)\.php(?:['"]|\s|$)`,
)
// Apache mod_rewrite directives. base64_decode / eval( appearing
// inside a RewriteCond or RewriteRule is a pattern in an attack-query
// blocklist (e.g. Really Simple SSL hardening), not PHP code.
htaccessRewriteDirective = regexp.MustCompile(
`(?i)^\s*Rewrite(?:Cond|Rule)\s`,
)
)
// checkCrontab scans a freshly-written /var/spool/cron/<user> file for the
// known persistence-marker patterns (literal + base64-decoded). Reads from
// the event fd, not the path, so an attacker swapping the file post-event
// cannot redirect us. Root crontab drift is tracked separately via
// hash-baseline by the polled CheckCrontabs.
func (fm *FileMonitor) checkCrontab(fd int, path, procInfo string) {
user := filepath.Base(path)
if user == "" || user == "root" || user == filepath.Base(cronSpoolDir()) {
return
}
recordReadTruncation(fd, 65536, "crontab")
data := readFromFd(fd, 65536)
if data == nil {
return
}
matched := checks.MatchCrontabPatternsDeep(string(data), fm.currentCfg())
if len(matched) == 0 {
return
}
fm.sendAlertWithPath(alert.Critical, "suspicious_crontab",
fmt.Sprintf("Suspicious crontab written for user %s: %v", user, matched),
fmt.Sprintf("File: %s\nPatterns matched: %v", path, matched),
path, procInfo)
}
// checkHtaccess reads .htaccess content from the event fd and checks for injection.
// C3 - reads from fd, not path.
func (fm *FileMonitor) checkHtaccess(fd int, path, procInfo string) {
recordReadTruncation(fd, htaccessRealtimeMaxBytes, "htaccess")
data := readFromFd(fd, htaccessRealtimeMaxBytes+1)
if data == nil {
return
}
if len(data) > htaccessRealtimeMaxBytes {
fm.sendAlertWithPath(alert.High, "htaccess_injection_realtime",
fmt.Sprintf(".htaccess too large to inspect in real time: %s", path),
"The file exceeds the real-time .htaccess scan limit and may hide malicious directives.",
path, procInfo)
return
}
for _, rawLine := range strings.Split(string(data), "\n") {
line := strings.TrimSpace(rawLine)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
lower := strings.ToLower(line)
// auto_prepend_file / auto_append_file: suspicious unless the
// directive target matches a known-legit security plugin file.
if strings.Contains(lower, "auto_prepend_file") || strings.Contains(lower, "auto_append_file") {
if htaccessAutoPrependSafeTarget.MatchString(line) {
continue
}
fm.sendAlertWithPath(alert.High, "htaccess_injection_realtime",
fmt.Sprintf("Suspicious .htaccess modification: %s", path),
"auto_prepend_file/auto_append_file target not recognised", path, procInfo)
continue
}
// eval( / base64_decode outside a RewriteCond/RewriteRule is a
// tamper signal: .htaccess is not a PHP execution context, so
// the only legit appearance of these tokens is as regex patterns
// inside mod_rewrite attack-blocklists.
if strings.Contains(lower, "eval(") || strings.Contains(lower, "base64_decode") {
if htaccessRewriteDirective.MatchString(line) {
continue
}
fm.sendAlertWithPath(alert.High, "htaccess_injection_realtime",
fmt.Sprintf("Suspicious .htaccess modification: %s", path),
"PHP function reference outside RewriteCond/RewriteRule", path, procInfo)
}
}
// Run the full .htaccess detector registry so realtime detection matches
// the depth of the scheduled scan: CGI-handler webshell arming, ModSecurity
// disable, PHP-in-uploads, handler remaps, cloaks, and redirect hijacks.
hardened, _ := checks.AuditHtaccessContent(path, data)
for _, f := range hardened {
fm.sendAlertWithPath(f.Severity, f.Check, f.Message, f.Details, path, procInfo)
}
// Run signature/YARA scanning on .htaccess content
fm.runEventSignatureScan(fd, data, path, ".htaccess", procInfo)
}
// checkUserINI reads the event fd so a path replacement cannot change the
// content being analyzed.
func (fm *FileMonitor) checkUserINI(fd int, path, procInfo string) {
recordReadTruncation(fd, checks.PHPConfigMaxBytes, "user_ini")
data := readFromFd(fd, checks.PHPConfigMaxBytes+1)
if data == nil {
return
}
if len(data) > checks.PHPConfigMaxBytes {
fm.sendAlertWithPath(alert.High, "php_config_realtime",
fmt.Sprintf("PHP configuration too large to inspect: %s", path),
"The file exceeds the PHP configuration scan limit and may hide dangerous directives.",
path, procInfo)
return
}
if dangerous := checks.PHPConfigSecurityBypasses(string(data)); len(dangerous) > 0 {
fm.sendAlertWithPath(alert.Critical, "php_config_realtime",
fmt.Sprintf("PHP security configuration weakened: %s", path),
fmt.Sprintf("Dangerous settings:\n- %s", strings.Join(dangerous, "\n- ")), path, procInfo)
return
}
// Run signature/YARA scanning on PHP configuration content.
fm.runEventSignatureScan(fd, data, path, ".ini", procInfo)
}
// checkPHPContent reads PHP content from the event fd and checks for malicious patterns.
// C3 - reads from fd, not path. M4 - 32KB scan size.
func (fm *FileMonitor) checkPHPContent(fd int, path, procInfo string) bool {
recordReadTruncation(fd, 32768, "php_check")
data := readFromFd(fd, 32768)
if data == nil {
return false
}
content := strings.ToLower(string(data))
if looksLikePHPWebshell(data) {
fm.sendAlertWithPath(alert.Critical, "webshell_content_realtime",
fmt.Sprintf("Webshell pattern detected: %s", path),
"Request input reaches a dangerous PHP execution primitive", path, procInfo)
return true
}
// Remote payload fetching — paste sites are always suspicious.
// GitHub raw URLs only flag when combined with a dangerous call on
// the same line (legitimate plugins use GitHub for update checks).
pasteURLs := []string{"pastebin.com/raw", "paste.ee/r/", "ghostbin.co/paste/", "hastebin.com/raw/"}
for _, p := range pasteURLs {
if strings.Contains(content, p) {
fm.sendAlertWithPath(alert.Critical, "php_dropper_realtime",
fmt.Sprintf("PHP dropper with paste site URL: %s", path),
fmt.Sprintf("Fetches from: %s", p), path, procInfo)
return true
}
}
githubURLs := []string{"gist.githubusercontent.com", "raw.githubusercontent.com"}
dangerousFns := []string{"file_put_contents(", "fwrite(", "shell_", "passthru(", "popen("}
for _, gh := range githubURLs {
if !strings.Contains(content, gh) {
continue
}
for _, line := range strings.Split(content, "\n") {
if !strings.Contains(line, gh) {
continue
}
for _, fn := range dangerousFns {
if strings.Contains(line, fn) {
fm.sendAlertWithPath(alert.Critical, "php_dropper_realtime",
fmt.Sprintf("PHP dropper fetching from GitHub with dangerous call: %s", path),
fmt.Sprintf("URL: %s, Function: %s", gh, fn), path, procInfo)
return true
}
}
}
}
// eval + decoder combo — require same-line nesting to avoid FPs on
// legitimate plugins that use these functions in unrelated contexts.
evalStr := "eval(" // search target for PHP eval function calls
assertStr := "assert(" // search target for PHP assert function calls
decoders := []string{"base64_decode", "gzinflate", "gzuncompress", "str_rot13", "gzdecode"}
for _, line := range strings.Split(content, "\n") {
lineHasEval := strings.Contains(line, evalStr) || strings.Contains(line, assertStr)
if !lineHasEval {
continue
}
for _, dec := range decoders {
if strings.Contains(line, dec) {
fm.sendAlertWithPath(alert.Critical, "obfuscated_php_realtime",
fmt.Sprintf("Obfuscated PHP detected: %s", path),
fmt.Sprintf("PHP code execution with %s on same line", dec), path, procInfo)
return true
}
}
}
// Fragmented base64 evasion: $a="base"; $b="64_decode"; $c=$a.$b;
if strings.Contains(content, "\"base\"") || strings.Contains(content, "'base'") {
if strings.Contains(content, "64_dec") && strings.Contains(content, evalStr) {
fm.sendAlertWithPath(alert.Critical, "obfuscated_php_realtime",
fmt.Sprintf("Fragmented base64_decode evasion detected: %s", path),
"base64_decode function name split across string variables", path, procInfo)
return true
}
}
// Massive variable concatenation payload ($z .= "xxxx"; repeated thousands of times)
concatCount := strings.Count(content, ".= \"")
if concatCount > 50 && strings.Contains(content, evalStr) {
fm.sendAlertWithPath(alert.Critical, "obfuscated_php_realtime",
fmt.Sprintf("Concatenation payload detected: %s (%d concat ops)", path, concatCount),
"Variable built from hundreds of string concatenations then executed", path, procInfo)
return true
}
// Shell execution with request input
// Uses containsFunc to avoid substring false positives
// (e.g. "WP_Filesystem(" matching "exec(", "preg_match(" matching "exec(")
shellFuncs := []string{"system(", "passthru(", "exec(", "shell_exec(", "popen("}
requestVars := []string{"$_request", "$_post", "$_get", "$_cookie", "$_server"}
hasShell := false
hasInput := false
for _, sf := range shellFuncs {
if containsFunc(content, sf) {
hasShell = true
}
}
for _, rv := range requestVars {
if strings.Contains(content, rv) {
hasInput = true
}
}
if hasShell && hasInput {
// Require shell function + request variable on the SAME line.
// Same-line narrowing is the actual detection: admin panels with
// both tokens in unrelated contexts stay quiet because they never
// co-occur on one line. A file-wide allowlist (e.g. "skip when
// 'wp_filesystem' appears anywhere") would be forgeable by any
// webshell that pastes the token into a comment.
for _, line := range strings.Split(content, "\n") {
lineHasShell := false
lineHasInput := false
for _, sf := range shellFuncs {
if containsFunc(line, sf) {
lineHasShell = true
break
}
}
for _, rv := range requestVars {
if strings.Contains(line, rv) {
lineHasInput = true
break
}
}
if lineHasShell && lineHasInput {
fm.sendAlertWithPath(alert.Critical, "webshell_content_realtime",
fmt.Sprintf("Webshell pattern detected: %s", path),
fmt.Sprintf("Shell execution with request input on same line: %s", strings.TrimSpace(line)), path, procInfo)
return true
}
}
}
// Tail scan: for large files, also check the last 32KB.
// Attackers append payloads (eval+base64) at the end of legitimate PHP files,
// beyond the head scan window. Only do the cheap heuristic checks, not full
// signature scanning (which would be too slow on every large PHP file).
if tailData := readTailFromFd(fd, 32768); tailData != nil {
tail := strings.ToLower(string(tailData))
// Check for eval+decoder on same line in tail
for _, line := range strings.Split(tail, "\n") {
lineHasEval := strings.Contains(line, evalStr) || strings.Contains(line, assertStr)
if !lineHasEval {
continue
}
for _, dec := range decoders {
if strings.Contains(line, dec) {
fm.sendAlertWithPath(alert.Critical, "obfuscated_php_realtime",
fmt.Sprintf("Obfuscated PHP appended to file tail: %s", path),
fmt.Sprintf("PHP code execution with %s found at end of file", dec), path, procInfo)
return true
}
}
}
// Fragmented base64 in tail
if strings.Contains(tail, "\"base\"") || strings.Contains(tail, "'base'") {
if strings.Contains(tail, "64_dec") && strings.Contains(tail, evalStr) {
fm.sendAlertWithPath(alert.Critical, "obfuscated_php_realtime",
fmt.Sprintf("Fragmented base64_decode evasion in file tail: %s", path),
"Payload appended at end of legitimate PHP file", path, procInfo)
return true
}
}
// Concat payload with eval in tail
tailConcatCount := strings.Count(tail, ".= \"")
if tailConcatCount > 50 && strings.Contains(tail, evalStr) {
fm.sendAlertWithPath(alert.Critical, "obfuscated_php_realtime",
fmt.Sprintf("Concatenation payload in file tail: %s (%d concat ops)", path, tailConcatCount),
"Payload appended at end of legitimate PHP file", path, procInfo)
return true
}
}
// Skip signature/YARA scanning for verified CMS core files.
// The wp_core periodic check validates files against official checksums;
// if a file's hash matches a known-clean core file, signature matches
// on it are false positives (e.g. $_POST in wp-includes, mail() in
// PHPMailer, fsockopen() in POP3.php).
// Hashed from the event descriptor, not by re-opening the path: the path
// can resolve to clean core content while the bytes just scanned were
// malicious, which would skip signature and YARA scanning for the file
// that was actually examined.
contentSize := int64(len(data))
var stat unix.Stat_t
if err := unix.Fstat(fd, &stat); err == nil && stat.Size > contentSize {
contentSize = stat.Size
}
if !checks.CMSCacheEmpty() && checks.CMSCacheMayContainSize(contentSize) &&
checks.IsVerifiedCMSHash(hashEventFD(fd, data, contentSize)) {
return false
}
// External signature + YARA scanning. The YAML engine sees the complete
// event-file size even though realtime analysis scans a bounded prefix, so
// per-rule file-size limits cannot be defeated by prefix truncation.
return fm.runSignatureScanWithSize(data, contentSize, path, filepath.Ext(path), procInfo, scannedIdentity(fd))
}
// Images fitting the combined read budget are scanned in one piece so a
// payload cannot straddle two independently evaluated windows. Larger files
// get a head and tail window; the middle is left to the deep scan, subject
// to thresholds.full_scan_max_file_mb.
const (
imagePayloadHeadBytes = 65536
imagePayloadTailBytes = 65536
)
// checkImagePayload looks for executable PHP inside a file served as an image.
// Returns true when a finding was raised.
//
// Two shapes reach the same verdict. A polyglot is a genuine image container
// with PHP appended or stored in a metadata chunk: it renders in a browser,
// passes an upload filter that trusts getimagesize, and executes the moment
// any PHP file includes its path. A file that only wears an image name and
// holds PHP source is the same backdoor without the disguise. Neither is
// legitimate under a served tree, so the container is reported as context
// rather than used as a gate.
//
// A PHP opening tag by itself is not evidence. Plugin screenshots quote one
// in their description chunks, so an execution, inclusion or remote-fetch
// construct is required alongside it.
func (fm *FileMonitor) checkImagePayload(fd int, path, procInfo string) bool {
var st unix.Stat_t
if err := unix.Fstat(fd, &st); err != nil || st.Mode&unix.S_IFMT != unix.S_IFREG || st.Size <= 0 {
return false
}
const readBudget = imagePayloadHeadBytes + imagePayloadTailBytes
headBytes := imagePayloadHeadBytes
if st.Size <= readBudget {
headBytes = int(st.Size)
}
recordReadTruncation(fd, readBudget, "image_payload")
head := readFromFd(fd, headBytes)
if len(head) == 0 {
return false
}
container, _ := contenttype.ImageContainer(head)
evidence, found := phpExecutableContent(head)
if !found && st.Size > readBudget {
// The container verdict came from the head, so the tail is examined
// for the payload alone.
if tail := readTailFromFd(fd, imagePayloadTailBytes); tail != nil {
evidence, found = phpExecutableContent(tail)
}
}
if !found {
return false
}
describedContainer := container
if describedContainer == "" {
describedContainer = "none (file is not a valid image)"
}
fm.sendAlertWithPath(alert.Critical, "php_in_image_realtime",
fmt.Sprintf("Executable PHP inside image file: %s", path),
fmt.Sprintf("Container: %s\nEvidence: %s\nRemediation: this path is the payload; find and remove the PHP file that includes it", describedContainer, evidence),
path, procInfo)
return true
}
// checkHTMLPhishing reads an HTML file and checks for phishing indicators:
// brand impersonation + credential input + redirect/exfiltration.
// Uses event fd for content read and unix.Fstat for size (TOCTOU-safe).
func (fm *FileMonitor) checkHTMLPhishing(fd int, path, procInfo string) {
// Only check files in web-accessible directories.
//
// No path-allowlist below this point: the content gates (credential
// inputs + brand impersonation + exfil/trust-badge) reject legitimate
// framework HTML on their own. A previous allowlist for /wp-admin/,
// /wp-content/themes/, /wp-content/plugins/, /node_modules/, /vendor/,
// /.well-known/ let an attacker who compromised any of those dirs drop
// a credential-harvesting page with full suppression.
if !fm.underDocRoot(path) {
return
}
var stat unix.Stat_t
if err := unix.Fstat(fd, &stat); err != nil {
return
}
size := stat.Size
if size < 500 || size > 500000 {
return
}
recordReadTruncation(fd, 16384, "html_phishing")
data := readFromFd(fd, 16384)
if data == nil {
return
}
content := strings.ToLower(string(data))
// Must have a form with credential inputs
if !strings.Contains(content, "<form") && !strings.Contains(content, "<input") {
return
}
hasCredInput := strings.Contains(content, "type=\"email\"") ||
strings.Contains(content, "type=\"password\"") ||
strings.Contains(content, "type='email'") ||
strings.Contains(content, "type='password'") ||
strings.Contains(content, "name=\"email\"") ||
strings.Contains(content, "name=\"password\"") ||
strings.Contains(content, "placeholder=\"you@")
if !hasCredInput {
return
}
// Check for brand impersonation
brands := []struct {
name string
patterns []string
}{
{"Microsoft/SharePoint", []string{"sharepoint", "onedrive", "microsoft 365", "outlook web", "office 365"}},
{"Google", []string{"google drive", "google docs", "accounts.google", "gmail"}},
{"Dropbox", []string{"dropbox"}},
{"DocuSign", []string{"docusign"}},
{"Adobe", []string{"adobe sign", "adobe document"}},
{"Apple/iCloud", []string{"icloud", "apple id"}},
{"Webmail", []string{"roundcube", "horde", "webmail login", "zimbra"}},
{"Generic", []string{"secure access", "verify your", "confirm your identity", "account verification"}},
}
brandMatch := ""
for _, b := range brands {
for _, p := range b.patterns {
if strings.Contains(content, p) {
brandMatch = b.name
break
}
}
if brandMatch != "" {
break
}
}
if brandMatch == "" {
return
}
// Check for redirect/exfiltration patterns
exfilPatterns := []string{
"window.location.href", "window.location.replace", "window.location =",
".workers.dev", "fetch(", "xmlhttprequest",
}
hasExfil := false
for _, p := range exfilPatterns {
if strings.Contains(content, p) {
hasExfil = true
break
}
}
// Also check for trust badges (strong phishing signal)
hasTrustBadge := strings.Contains(content, "secured by microsoft") ||
strings.Contains(content, "secured by google") ||
strings.Contains(content, "256-bit encrypted") ||
strings.Contains(content, "256‑bit encrypted")
if hasExfil || hasTrustBadge {
fm.sendAlertWithPath(alert.Critical, "phishing_realtime",
fmt.Sprintf("Phishing page created (%s impersonation): %s", brandMatch, path),
fmt.Sprintf("Size: %d bytes", size), path, procInfo)
return
}
// Run signature/YARA scanning on HTML content not caught by phishing heuristics
fm.runEventSignatureScan(fd, data, path, ".html", procInfo)
}
// checkCredentialLog reads a text file and checks if it contains harvested
// email:password pairs - output from an active phishing kit. The path is used
// only for the suppression/location checks; the content is read from the
// fanotify event fd (not re-opened by path) so an attacker cannot swap the
// file between the event and the read.
func (fm *FileMonitor) checkCredentialLog(fd int, path, procInfo string) {
if !fm.underDocRoot(path) {
return
}
// Exclude known config file paths - these legitimately contain email-like patterns.
if strings.HasPrefix(path, "/etc/") {
return
}
for _, suffix := range []string{".conf", ".cfg", ".ini", ".yaml", ".yml"} {
if strings.HasSuffix(path, suffix) {
return
}
}
data := readFromFd(fd, 4096)
if data == nil {
return
}
content := string(data)
lines := strings.Split(content, "\n")
credLines := 0
emailCount := 0
for _, line := range lines {
line = strings.TrimSpace(line)
if line == "" {
continue
}
if strings.Contains(line, "@") {
emailCount++
for _, delim := range []string{":", "|", "\t", ","} {
parts := strings.SplitN(line, delim, 3)
if len(parts) >= 2 {
p0 := strings.TrimSpace(parts[0])
p1 := strings.TrimSpace(parts[1])
if strings.Contains(p0, "@") && len(p1) > 0 && !strings.Contains(p1, " ") {
credLines++
break
}
}
}
}
}
if credLines >= 5 {
fm.sendAlertWithPath(alert.Critical, "credential_log_realtime",
fmt.Sprintf("Harvested credential log detected: %s", path),
fmt.Sprintf("%d credential lines (email:password format) found", credLines), path, procInfo)
} else if emailCount >= 10 {
fm.sendAlertWithPath(alert.High, "credential_log_realtime",
fmt.Sprintf("Possible harvested email list: %s", path),
fmt.Sprintf("%d email addresses found in %s", emailCount, filepath.Base(path)), path, procInfo)
}
}
// checkPhishingZip checks if a newly created ZIP file matches known
// phishing-kit archive name patterns. The signal is the COMBINATION of a
// brand name and a phishing-suggestive token in the same filename --
// "office365-login.zip", "paypal-verify.zip", "microsoft-secure.zip".
// Plain plugin distribution backups (google-site-kit.zip, mailchimp.zip)
// have a brand without an action verb and don't fire.
func (fm *FileMonitor) checkPhishingZip(path, nameLower, procInfo string) {
if !fm.underDocRoot(path) {
return
}
// Brand impersonation targets: filenames mimicking a service users log in to.
brands := []string{
"office365", "office 365", "sharepoint", "onedrive",
"microsoft", "outlook", "google", "gmail",
"dropbox", "docusign", "adobe", "wetransfer",
"paypal", "apple", "icloud", "netflix",
"facebook", "instagram", "linkedin",
"webmail", "roundcube", "cpanel",
}
// Phishing-suggestive verbs/nouns. These must co-occur with a brand
// for the rule to fire. "kit" is intentionally NOT here -- many
// official WordPress plugin slugs end in -kit (google-site-kit,
// mailchimp-for-wp-kit) and were the dominant FP source.
phishingIndicators := []string{
"login", "signin", "sign-in", "sign_in",
"verify", "verification",
"secure", "security",
"phish", "scam",
"bank", "account",
"capture", "harvest", "steal",
}
var matchedBrand string
for _, b := range brands {
if strings.Contains(nameLower, b) {
matchedBrand = b
break
}
}
if matchedBrand == "" {
return
}
var matchedIndicator string
for _, p := range phishingIndicators {
if strings.Contains(nameLower, p) {
matchedIndicator = p
break
}
}
if matchedIndicator == "" {
return
}
fm.sendAlertWithPath(alert.High, "phishing_kit_realtime",
fmt.Sprintf("Suspected phishing kit archive uploaded: %s", path),
fmt.Sprintf("Filename combines brand '%s' with phishing indicator '%s'", matchedBrand, matchedIndicator),
path, procInfo)
}
// runSignatureScan runs YAML and YARA signature scanning on file content.
// Returns true if a match was found and an alert was sent.
// Non-critical YAML matches use directory-level dedup to avoid alert floods
// when a plugin directory has many files matching the same rule.
// Critical matches (backdoors, webshells) always alert per-file.
// scannedIdentity describes the object behind an event descriptor. Stat of the
// /proc magic link resolves the open file itself rather than walking the path
// again, so it still names the scanned inode after the path has been replaced.
// os.NewFile is avoided deliberately: its finalizer can close a descriptor the
// daemon still owns.
func scannedIdentity(fd int) os.FileInfo {
info, err := os.Stat(fmt.Sprintf("/proc/self/fd/%d", fd))
if err != nil {
return nil
}
return info
}
func (fm *FileMonitor) runSignatureScan(data []byte, path, ext, procInfo string) bool {
return fm.runSignatureScanWithSize(data, int64(len(data)), path, ext, procInfo, nil)
}
func (fm *FileMonitor) runEventSignatureScan(fd int, data []byte, path, ext, procInfo string) bool {
return fm.runSignatureScanWithSize(data, int64(len(data)), path, ext, procInfo, scannedIdentity(fd))
}
func (fm *FileMonitor) runSignatureScanWithSize(data []byte, contentSize int64, path, ext, procInfo string, scanned os.FileInfo) bool {
// Both engines see every file. A .yml hit used to end the scan here, so
// a file matching a High .yml rule never met the Critical YARA rule and
// the inline quarantine that only a Critical match triggers. Only a file
// the .yml path already moved to quarantine is not handed to YARA.
matched := false
if scanner := signatures.Global(); scanner != nil {
matches := scanner.ScanContentWithSize(data, ext, contentSize)
if len(matches) > 0 {
matched = true
m := matches[0]
sev := alert.High
if m.Severity == "critical" {
sev = alert.Critical
}
// Non-critical: dedup by rule+directory so 30 files in the same
// plugin matching the same rule produce one alert, not 30.
// Critical matches always alert per-file (real path for quarantine).
suppressed := false
if sev != alert.Critical {
dirKey := m.RuleName + ":" + filepath.Dir(path)
suppressed = !fm.shouldAlert("signature_match_realtime", dirKey)
}
if !suppressed {
details := fmt.Sprintf("Category: %s\nDescription: %s\nMatched: %s",
m.Category, m.Description, strings.Join(m.Matched, ", "))
details += signatures.ReferencedPayloadDetail(data)
finding := alert.Finding{
Severity: sev,
Check: "signature_match_realtime",
Message: fmt.Sprintf("Signature match [%s]: %s", m.RuleName, path),
Details: details,
FilePath: path,
ProcessInfo: procInfo,
}
var qPath string
var quarantined bool
var paused *alert.Finding
if sev == alert.Critical {
// Capture provenance before remediation can remove the source.
checks.StampContentFingerprint(&finding)
qPath, quarantined, paused = checks.InlineQuarantineGatedIdentified(fm.currentCfg(), &finding, path, data, scanned)
}
// Publish after the inline decision so delivery sees its budget
// provenance. A rejected window can still get full-file validation.
fm.sendFileFinding(finding)
if paused != nil && !alert.TryEnqueue(fm.alertCh, *paused) {
atomic.AddInt64(&fm.droppedAlerts, 1)
}
if quarantined {
fm.recordDropperQuarantine(path, qPath)
fm.sendAlert(alert.Critical, "auto_response",
fmt.Sprintf("AUTO-QUARANTINE (inline): %s moved to quarantine", path),
fmt.Sprintf("Quarantined to: %s\nRule: %s", qPath, m.RuleName))
return true
}
}
}
}
if yaraScanner := yara.Active(); yaraScanner != nil {
matches, scannedSHA, err := scanRealtimeYARA(yaraScanner, path, data)
if err != nil {
fm.reportYARAScanError(path, err)
return matched
}
if len(matches) > 0 {
fm.sendFileFinding(alert.Finding{
Severity: alert.Critical,
Check: "yara_match_realtime",
Message: fmt.Sprintf("YARA rule match [%s]: %s", matches[0].RuleName, path),
Details: fmt.Sprintf("Matched %d YARA rule(s)", len(matches)) + signatures.ReferencedPayloadDetail(data),
FilePath: path,
ProcessInfo: procInfo,
ContentSHA256: scannedSHA,
})
return true
}
}
return matched
}
// stopping reports whether this monitor has been signalled to stop. A nil
// stopCh (a monitor built directly in a test) is never stopping, because a
// receive on a nil channel cannot proceed.
func (fm *FileMonitor) stopping() bool {
select {
case <-fm.stopCh:
return true
default:
return false
}
}
// reportYARAScanError names a changed file the scanner could not inspect.
// It is its own check rather than the deep scan's "yara_scan_incomplete",
// which reports scheduled coverage: that report fires for every archive past
// the scan size limit, roughly thirteen times a day forever on a live host,
// and sharing the name left a real scanning outage indistinguishable from
// routine backlog.
//
// Shutdown is not an outage. The daemon stops the YARA backend while this
// monitor's goroutine is still draining events, because the wait for workers
// comes after the teardown, so a clean restart otherwise reported a
// High-severity scanning failure every time. The teardown cannot move after
// that wait, which is unbounded and would hang on a wedged worker. The
// return happens before the rate-limit window is taken, so a suppressed
// shutdown report cannot swallow the first genuine failure afterwards.
func (fm *FileMonitor) reportYARAScanError(path string, err error) {
if fm.stopping() {
return
}
fm.yaraErrorReportMu.Lock()
if !fm.lastYARAError.IsZero() && time.Since(fm.lastYARAError) < time.Minute {
fm.yaraErrorReportMu.Unlock()
return
}
fm.lastYARAError = time.Now()
fm.yaraErrorReportMu.Unlock()
fm.sendAlert(alert.High, "yara_realtime_scan_error",
"YARA real-time scan could not inspect a changed file",
fmt.Sprintf("File: %s\nError: %v", path, err))
}
// M7 - sendAlert uses droppedAlerts counter, separate from droppedEvents.
// No dedup - only used for system-level alerts (overflow reporting) that are
// already ticker-gated. File-related alerts should use sendAlertWithPath.
func (fm *FileMonitor) sendAlert(severity alert.Severity, check, message, details string) {
finding := alert.Finding{
Severity: severity,
Check: check,
Message: message,
Details: details,
Timestamp: time.Now(),
}
if !alert.TryEnqueue(fm.alertCh, finding) {
atomic.AddInt64(&fm.droppedAlerts, 1)
}
}
// sendAlertWithPath is like sendAlert but also sets the FilePath and
// ProcessInfo fields for structured propagation to auto-response.
// Applies per-path deduplication to prevent alert storms from rapid writes.
func (fm *FileMonitor) sendAlertWithPath(severity alert.Severity, check, message, details, filePath, processInfo string) {
fm.sendFileFinding(alert.Finding{
Severity: severity,
Check: check,
Message: message,
Details: details,
FilePath: filePath,
ProcessInfo: processInfo,
})
}
func (fm *FileMonitor) sendFileFinding(finding alert.Finding) {
dedupPath := finding.FilePath
if finding.DedupKey != "" {
// A later header or digest can change a staged finding's identity
// within the path cooldown. Let persistent dedup see that evidence.
dedupPath = finding.Key()
}
if !fm.shouldAlert(finding.Check, dedupPath) {
return
}
finding.Timestamp = time.Now()
checks.StampContentFingerprint(&finding)
if !alert.TryEnqueue(fm.alertCh, finding) {
atomic.AddInt64(&fm.droppedAlerts, 1)
}
}
// shouldAlert returns true if this check+path combination hasn't been alerted
// recently. Prevents duplicate alerts from rapid writes to the same file.
// Uses LoadOrStore for atomic initial insertion to avoid TOCTOU races
// between concurrent analyzer workers.
func (fm *FileMonitor) shouldAlert(check, filePath string) bool {
if filePath == "" {
return true // no path = no dedup possible
}
key := check + ":" + filePath
now := time.Now()
if v, loaded := fm.alertDedup.LoadOrStore(key, now); loaded {
if now.Sub(v.(time.Time)) < alertDedupTTL {
return false
}
fm.alertDedup.Store(key, now) // refresh TTL on expiry
}
return true
}
// M7 - overflowReporter reports dropped events and alerts separately.
//
// Periodic ticks and eager signals feed this loop:
// - 1-minute ticker: emits the periodic fanotify_overflow alert,
// resets drop counters, runs reconcileDrops, and evicts stale alert
// dedup entries.
// - retry ticker: retries pending work even after new drops stop.
// - reconcileSig: out-of-cycle reconcile triggered by sendEvent when
// sustained drops cross eagerReconcileDropThreshold within the
// current tick. Closes the latency gap between a drop and its
// reconcile read so the file's mtime is still inside reconcileWindow.
func (fm *FileMonitor) overflowReporter() {
defer fm.wg.Done()
ticker := time.NewTicker(1 * time.Minute)
defer ticker.Stop()
// Retried work must not depend on another drop or a nonzero alert counter.
retry := time.NewTicker(reconcileMinInterval)
defer retry.Stop()
for {
select {
case <-fm.stopCh:
return
case <-retry.C:
fm.startReconcile()
case <-fm.reconcileSig:
// Eager reconcile: do not reset counters, do not emit the
// minute-tick alert. Just walk the tracked dirs and surface
// any interesting file inside reconcileWindow. The minute
// tick will still fire its alert + drain the counters when
// it arrives.
fm.startReconcile()
case <-ticker.C:
droppedEv := atomic.SwapInt64(&fm.droppedEvents, 0)
droppedAl := atomic.SwapInt64(&fm.droppedAlerts, 0)
if droppedEv > 0 {
fm.sendAlert(alert.Warning, "fanotify_overflow",
fmt.Sprintf("fanotify event queue overflowed: %d events dropped in last minute", droppedEv),
"Possible event storm (backup, bulk update) or high-volume attack")
// Recover coverage: scan files in directories that saw drops
// so a threat landing during the storm is still detected.
fm.startReconcile()
}
if droppedAl > 0 {
fmt.Fprintf(os.Stderr, "[%s] alert channel full: %d alerts dropped in last minute\n", ts(), droppedAl)
}
// Evict stale dedup entries every minute
now := time.Now()
fm.alertDedup.Range(func(key, value any) bool {
if now.Sub(value.(time.Time)) > alertDedupTTL {
fm.alertDedup.Delete(key)
}
return true
})
evictStaleWPPathStatCache(now)
}
}
}
// evictStaleWPPathStatCache bounds the package-level WordPress path stat cache.
// The compare-delete keeps the minute sweep from removing a fresh stat result
// stored by an analyzer worker after Range observed an older entry.
func evictStaleWPPathStatCache(now time.Time) {
cutoff := 2 * wpPathCacheTTL
wpPathStatCache.Range(func(key, value any) bool {
entry, ok := value.(wpPathCacheEntry)
if !ok {
wpPathStatCache.Delete(key)
return true
}
if now.Sub(entry.ts) > cutoff {
wpPathStatCache.CompareAndDelete(key, entry)
}
return true
})
}
// isPHPExtension returns true for all PHP file extensions that can execute code.
// containsFunc checks if content contains a function call that isn't part of
// a longer identifier. Prevents "WP_Filesystem(" matching "exec(" or
// "preg_match(" matching "exec(". Checks the character before the match
// is not a letter, digit, or underscore.
func containsFunc(content, funcCall string) bool {
idx := 0
for {
pos := strings.Index(content[idx:], funcCall)
if pos < 0 {
return false
}
absPos := idx + pos
if absPos == 0 {
return true
}
prev := content[absPos-1]
if (prev < 'a' || prev > 'z') && (prev < 'A' || prev > 'Z') &&
(prev < '0' || prev > '9') && prev != '_' {
return true
}
idx = absPos + len(funcCall)
if idx >= len(content) {
return false
}
}
}
func isPHPExtension(nameLower string) bool {
// Single source of truth shared with the periodic content scanners and the
// rule engines so no path drifts on which extensions execute PHP.
return contenttype.IsExecutablePHPName(nameLower)
}
func isPHPSourceExtension(nameLower string) bool {
return contenttype.IsPHPSourceName(nameLower)
}
func isCGIExtension(nameLower string) bool {
return strings.HasSuffix(nameLower, ".pl") ||
strings.HasSuffix(nameLower, ".cgi") ||
strings.HasSuffix(nameLower, ".py") ||
strings.HasSuffix(nameLower, ".sh") ||
strings.HasSuffix(nameLower, ".rb")
}
// checkCGIBackdoor reads a CGI script and checks for backdoor patterns.
// Detects Perl/Python/Bash backdoors like the LEVIATHAN toolkit.
func (fm *FileMonitor) checkCGIBackdoor(fd int, path, procInfo string) {
recordReadTruncation(fd, 32768, "cgi_backdoor")
data := readFromFd(fd, 32768)
if data == nil {
return
}
content := strings.ToLower(string(data))
// Backdoor indicators in CGI scripts
indicators := 0
var matched []string
// Indicators weighted by suspicion level. Generic patterns like
// "request_method" and "cmd" removed — they match every CGI script.
shellPatterns := []struct {
pattern string
desc string
}{
{"system(", "system() call"},
{"os.popen", "os.popen() call"},
{"`$", "backtick execution with variable"},
{"content_length", "reads POST body length"},
{"base64_decode", "base64 decoding"},
{"$_post", "PHP POST input"},
{"$_get", "PHP GET input"},
{"param(", "CGI parameter read"},
{"qs.parse", "query string parsing"},
}
for _, sp := range shellPatterns {
if strings.Contains(content, sp.pattern) {
indicators++
matched = append(matched, sp.desc)
}
}
// 4+ indicators = likely backdoor
if indicators >= 4 {
fm.sendAlertWithPath(alert.Critical, "cgi_backdoor_realtime",
fmt.Sprintf("CGI backdoor detected: %s", path),
fmt.Sprintf("Indicators (%d): %s", indicators, strings.Join(matched, ", ")), path, procInfo)
return
}
// CGI scripts in unusual locations (images, css, js directories)
if strings.Contains(path, "/img/") || strings.Contains(path, "/images/") ||
strings.Contains(path, "/css/") || strings.Contains(path, "/js/") ||
strings.Contains(path, "/fonts/") || strings.Contains(path, "/icons/") {
fm.sendAlertWithPath(alert.High, "cgi_suspicious_location_realtime",
fmt.Sprintf("CGI script in non-CGI directory: %s", path),
"Scripts should not exist in image/css/js directories", path, procInfo)
return
}
// Run signature scan on the content
fm.runEventSignatureScan(fd, data, path, filepath.Ext(path), procInfo)
}
// matchSuppression checks if a file path matches a suppression glob pattern.
// Supports patterns like "*/cache/*", "*/vendor/*", "*.log".
func matchSuppression(pattern, path string) bool {
if pattern == "" {
return false
}
// Direct match against full path
if m, _ := filepath.Match(pattern, path); m {
return true
}
// Match against basename (e.g. "*.log")
if m, _ := filepath.Match(pattern, filepath.Base(path)); m {
return true
}
if !strings.ContainsAny(pattern, "*?[") {
return strings.Contains(path, pattern)
}
if strings.ContainsAny(pattern, "?[") || !hasLeadingAnyDepthSuppressionGlob(pattern) {
return false
}
residue := strings.ReplaceAll(pattern, "*", "")
if strings.Contains(residue, "/") && strings.Trim(residue, "/") != "" {
return strings.Contains(path, residue)
}
return false
}
func hasLeadingAnyDepthSuppressionGlob(pattern string) bool {
firstSlash := strings.Index(pattern, "/")
if firstSlash <= 0 {
return false
}
return strings.Trim(pattern[:firstSlash], "*") == ""
}
// readHead opens a file by path and reads the first maxBytes.
// Kept for path-based checks (HTML phishing, credential logs, ZIP checks)
// that need os.Stat for file size anyway.
func readHead(path string, maxBytes int) []byte {
if maxBytes <= 0 {
return nil
}
// #nosec G304 -- readHead scans files surfaced by fanotify/scanner;
// reading user files for signature analysis is the daemon's purpose.
f, err := os.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
// ReadAll over a LimitReader, not a single f.Read into a pre-sized
// buffer: a short read would hand only a prefix to the detectors and
// silently miss content deeper in the file.
buf, err := io.ReadAll(io.LimitReader(f, int64(maxBytes)))
if err != nil || len(buf) == 0 {
return nil
}
return buf
}
// looksLikePluginUpdate checks if a PHP file in uploads looks like a plugin
// update temp directory (e.g., elementor_t0q9y). Returns true if it matches
// the pattern of a known plugin extracting an update.
// M3 - uses sync.Map cache with 5-minute TTL for plugin directory stat results.
func looksLikePluginUpdate(path string) bool {
// WordPress plugin updates extract to /uploads/{pluginname}_{random}/
// Detect by extracting the directory name under uploads/ and checking
// if a matching plugin exists in wp-content/plugins/.
// No hardcoded whitelist - works for all 60,000+ WP plugins.
uploadsIdx := strings.Index(path, "/wp-content/uploads/")
if uploadsIdx < 0 {
return false
}
wpRoot := path[:uploadsIdx]
afterUploads := path[uploadsIdx+len("/wp-content/uploads/"):]
// Extract the first directory component: "header-footer_7ocsd"
slashIdx := strings.Index(afterUploads, "/")
if slashIdx < 0 {
return false
}
dirName := afterUploads[:slashIdx]
// Strip the random suffix (e.g. "_7ocsd") - WordPress appends _XXXXX
// The plugin name is everything before the last underscore-followed-by-random
pluginName := dirName
if lastUnderscore := strings.LastIndex(dirName, "_"); lastUnderscore > 0 {
suffix := dirName[lastUnderscore+1:]
// Random suffixes are short alphanumeric strings (5-8 chars)
if len(suffix) >= 4 && len(suffix) <= 10 {
pluginName = dirName[:lastUnderscore]
}
}
// Check if a matching plugin directory exists in plugins/
return cachedPathExists(wpRoot + "/wp-content/plugins/" + pluginName)
}
// cachedPathExists answers whether path exists, memoised for wpPathCacheTTL
// when it does and for wpPathNegativeTTL when it does not. The realtime path
// asks this once per file event during an update, so an uncached stat per
// staged file would be paid thousands of times per package; a missing path
// during an update is transient, so its answer must expire quickly.
func cachedPathExists(path string) bool {
if cached, ok := wpPathStatCache.Load(path); ok {
if entry, ok := cached.(wpPathCacheEntry); ok {
ttl := wpPathCacheTTL
if !entry.exists {
ttl = wpPathNegativeTTL
}
if time.Since(entry.ts) < ttl {
return entry.exists
}
}
}
_, err := os.Stat(path)
exists := err == nil
wpPathStatCache.Store(path, wpPathCacheEntry{
exists: exists,
ts: time.Now(),
})
return exists
}
//go:build linux
package daemon
import (
"errors"
"io/fs"
"os"
"path/filepath"
"time"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/queuehealth"
)
type reconcileDirectory struct {
firstDrop time.Time
lastDrop time.Time
ticket queuehealth.Ticket
// Keep the original scan window and progress when a pass yields.
cutoff time.Time
after string
failed bool
}
// Repeated drops refresh eviction priority without postponing recovery age.
func (fm *FileMonitor) recordDroppedDir(path string) {
fm.initQueueHealth()
dir := filepath.Dir(path)
fm.reconcileMu.Lock()
defer fm.reconcileMu.Unlock()
if fm.reconcileDirs == nil {
fm.reconcileDirs = make(map[string]reconcileDirectory)
}
now := time.Now()
entry, exists := fm.reconcileDirs[dir]
if !exists {
entry.firstDrop = now
entry.ticket = fm.reconcileHealth.Begin(now)
}
// A new write may precede the saved cursor and must be examined again.
entry.after = ""
entry.lastDrop = now
fm.reconcileDirs[dir] = entry
fm.evictOverflowDirLocked(now)
}
// evictOverflowDirLocked drops the least recently refreshed directory once the
// tracker is over its cap. The caller holds reconcileMu.
func (fm *FileMonitor) evictOverflowDirLocked(now time.Time) {
if len(fm.reconcileDirs) <= reconcileDirCap {
return
}
var oldestKey string
var oldestTime time.Time
first := true
for path, candidate := range fm.reconcileDirs {
if first || candidate.lastDrop.Before(oldestTime) {
oldestKey, oldestTime, first = path, candidate.lastDrop, false
}
}
fm.reconcileDirs[oldestKey].ticket.Reject(now)
delete(fm.reconcileDirs, oldestKey)
}
// startReconcile runs one recovery pass off the caller's goroutine. Passes
// never overlap and never run closer together than reconcileMinInterval: the
// eager trigger fires on drop volume, and a storm produces that volume far
// faster than a pass can absorb it.
func (fm *FileMonitor) startReconcile() {
select {
case <-fm.stopCh:
return
default:
}
if !fm.reconcileRunning.CompareAndSwap(false, true) {
return
}
// Acquire single-flight ownership before reading the completion time:
// the previous pass may finish while this caller is being scheduled.
now := time.Now()
if last := fm.reconcileLastPass.Load(); last != 0 && now.Sub(time.Unix(0, last)) < reconcileMinInterval {
fm.reconcileRunning.Store(false)
return
}
fm.reconcileMu.Lock()
pending := len(fm.reconcileDirs) > 0
fm.reconcileMu.Unlock()
if !pending {
fm.reconcileRunning.Store(false)
return
}
// The overflow reporter owns a wg count until its last call returns,
// so shutdown cannot observe zero while this Add is possible.
fm.wg.Add(1)
obs.Go("fanotify-reconcile", func() {
defer fm.wg.Done()
defer func() {
fm.reconcileLastPass.Store(time.Now().UnixNano())
fm.reconcileRunning.Store(false)
}()
fm.reconcileDrops()
})
}
func (fm *FileMonitor) reconcileDrops() {
fm.initQueueHealth()
fm.reconcileMu.Lock()
dirs := fm.reconcileDirs
fm.reconcileDirs = make(map[string]reconcileDirectory)
started := time.Now()
for dir, entry := range dirs {
entry.ticket.Start(started)
if entry.cutoff.IsZero() {
entry.cutoff = started.Add(-reconcileWindow)
}
dirs[dir] = entry
}
fm.reconcileMu.Unlock()
if len(dirs) == 0 {
return
}
// A panic abandons every unfinished directory in this detached batch.
defer func() {
now := time.Now()
for _, entry := range dirs {
entry.ticket.Reject(now)
}
}()
if fanotifyReconcileDur != nil {
defer func() { fanotifyReconcileDur.Observe(time.Since(started).Seconds()) }()
}
deadline := started.Add(reconcileBudget)
for dir, entry := range dirs {
if fm.reconcilePassDone(deadline) {
fm.deferReconcileDirs(dirs)
return
}
if !fm.reconcileDirectory(dir, &entry, deadline) {
dirs[dir] = entry
fm.deferReconcileDirs(dirs)
return
}
// A refreshed drop cannot recover an older obligation outside this scan's window.
if !entry.failed && !entry.firstDrop.Before(entry.cutoff) {
entry.ticket.Finish(time.Now())
} else {
entry.ticket.Reject(time.Now())
}
delete(dirs, dir)
}
}
// reconcilePassDone reports whether this pass must stop: its budget is spent,
// or the monitor is shutting down.
func (fm *FileMonitor) reconcilePassDone(deadline time.Time) bool {
select {
case <-fm.stopCh:
return true
default:
}
return !time.Now().Before(deadline)
}
// deferReconcileDirs returns directories this pass did not reach to the
// tracker so the next pass takes them, and empties the detached batch so the
// panic guard does not count them as lost. A directory that took a new drop
// while the pass ran is already represented by a waiting ticket; merging keeps
// the older admission age rather than restarting the clock on it.
func (fm *FileMonitor) deferReconcileDirs(dirs map[string]reconcileDirectory) {
now := time.Now()
fm.reconcileMu.Lock()
defer fm.reconcileMu.Unlock()
if fm.reconcileDirs == nil {
fm.reconcileDirs = make(map[string]reconcileDirectory)
}
for dir, entry := range dirs {
delete(dirs, dir)
if waiting, exists := fm.reconcileDirs[dir]; exists {
waiting.ticket.MergeRunning(entry.ticket, now)
if entry.firstDrop.Before(waiting.firstDrop) {
waiting.firstDrop = entry.firstDrop
}
if waiting.cutoff.IsZero() || entry.cutoff.Before(waiting.cutoff) {
waiting.cutoff = entry.cutoff
}
// The new admission may concern a file behind the old cursor.
waiting.after = ""
waiting.failed = waiting.failed || entry.failed
fm.reconcileDirs[dir] = waiting
continue
}
entry.ticket.Requeue(now)
fm.reconcileDirs[dir] = entry
fm.evictOverflowDirLocked(now)
}
}
// A tree removed before recovery ran holds nothing left to scan, so the
// kernel loss that started the recovery is the only loss. Reads that failed
// for any other reason left work the operator can still act on.
// Returns false only when the pass yields; read failures stay on the
// obligation until completion so a later pass cannot conceal partial loss.
func (fm *FileMonitor) reconcileDirectory(dir string, work *reconcileDirectory, deadline time.Time) bool {
entries, err := os.ReadDir(dir)
if err != nil {
work.failed = work.failed || !errors.Is(err, fs.ErrNotExist)
return true
}
for _, entry := range entries {
if fm.reconcilePassDone(deadline) {
return false
}
// os.ReadDir sorts names, so continuation does not rescan the prefix.
if entry.Name() <= work.after {
continue
}
work.after = entry.Name()
if entry.IsDir() {
continue
}
path := filepath.Join(dir, entry.Name())
if !fm.isInteresting(path) {
continue
}
info, err := entry.Info()
if err != nil {
if !errors.Is(err, fs.ErrNotExist) {
work.failed = true
}
continue
}
if info.ModTime().Before(work.cutoff) {
continue
}
if !fm.reconcileFile(path) {
work.failed = true
}
}
return true
}
func (fm *FileMonitor) reconcileFile(path string) bool {
// #nosec G304 -- path is a candidate in a directory recorded after a dropped event; its original event fd is no longer available.
file, err := os.Open(path)
if err != nil {
return errors.Is(err, fs.ErrNotExist)
}
// Close per file, including panic, rather than retaining the whole batch's fds.
defer func() { _ = file.Close() }()
// #nosec G115 -- Linux file descriptors are nonnegative int32 values and fit Go int.
fileAnalyzer(fm, fileEvent{path: path, fd: int(file.Fd())})
return true
}
// Reader and worker producers have joined before shutdown discards this map.
func (fm *FileMonitor) discardReconcilePending() {
fm.reconcileMu.Lock()
defer fm.reconcileMu.Unlock()
now := time.Now()
for dir, entry := range fm.reconcileDirs {
entry.ticket.Reject(now)
delete(fm.reconcileDirs, dir)
}
}
//go:build linux
package daemon
import (
"fmt"
"os"
"strings"
"sync"
"syscall"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/wpcheck"
)
// WordPress unpacks every core, plugin and theme package under
// wp-content/upgrade/<package>/<unpacked>/ and moves the unpacked tree into
// place a few hundred milliseconds later. The analyzer sees each staged PHP
// file in that window, usually before the checksum fetch for the package
// version has finished and sometimes before the header file that names the
// version has even been written.
//
// A staged file is therefore judged by hash, not by where it sits. Files the
// verifier can compare are compared; files whose checksums are still on the
// way wait here and are compared once they land, from the installed location
// if WordPress has moved them by then. Only a file that fails the comparison
// gets its own finding. Packages with no checksum source, and packages the
// fetch cannot resolve in time, collapse to one finding per staging directory.
// wpVerifier is the wordpress.org checksum verifier the analyzer consults.
type wpVerifier interface {
Describe(path string) wpcheck.Verification
Verify(v wpcheck.Verification) wpcheck.Verdict
VerifyFile(fd int, path string) wpcheck.Verification
}
const (
// stagedPackageTimeout bounds how long a staged file waits for its
// package checksums. The fetch is one ZIP download; a package still
// unresolved after this is reported at the directory level instead.
stagedPackageTimeout = 60 * time.Second
// stagedPackageQueueMax caps queued files across all packages. A full
// core or plugin update is a few thousand files; the cap only matters
// when many updates run at once while the fetches stall.
stagedPackageQueueMax = 20000
// stagedPackageIdleTTL retires per-package state that saw no new file.
stagedPackageIdleTTL = 10 * time.Minute
// wpPathNegativeTTL is how long a missing path stays cached. WordPress
// moves the old plugin to upgrade-temp-backup/ and renames the new tree
// in a few hundred milliseconds later, so a negative answer is transient
// by definition and must not hold for the five minutes a positive one does.
wpPathNegativeTTL = 2 * time.Second
)
// wpStagedPackage is a parsed wp-content/upgrade/<package>/<unpacked>/ path.
type wpStagedPackage struct {
dir string
unpacked string
wpRoot string
}
// parseWPStagedPackage identifies the staging directory and the unpacked tree
// a path sits in. Pure string work: the staging name is generated by
// WordPress in several shapes and carries no authority anyway, since anything
// able to write a staged tree chooses both names.
func parseWPStagedPackage(path string) wpStagedPackage {
const marker = "/wp-content/upgrade/"
idx := strings.Index(path, marker)
if idx < 0 {
return wpStagedPackage{}
}
// An installed plugin may ship upgrade-layout fixtures. Its outer root
// owns those files, just as it does when wpcheck describes their manifest.
if root, _ := wpcheck.DetectPluginRoot(path); root != "" && len(root) <= idx {
return wpStagedPackage{}
}
rest := path[idx+len(marker):]
pkg, inner, ok := strings.Cut(rest, "/")
if !ok || pkg == "" {
return wpStagedPackage{} // a file sitting directly in upgrade/ belongs to no package
}
unpacked, _, ok := strings.Cut(inner, "/")
if !ok || unpacked == "" {
return wpStagedPackage{} // the package directory holds no unpacked tree
}
return wpStagedPackage{
dir: path[:idx+len(marker)] + pkg,
unpacked: unpacked,
wpRoot: path[:idx],
}
}
// wpPackageInstalled reports whether the unpacked tree names something the
// site already has: a plugin or theme directory, its rollback copy under
// upgrade-temp-backup/ while an update is mid-flight, or a WordPress root for
// a core package.
func wpPackageInstalled(wpRoot, unpacked string) bool {
if unpacked == "wordpress" {
return cachedPathExists(wpRoot + "/wp-includes/version.php")
}
for _, dir := range []string{
"/wp-content/plugins/", "/wp-content/themes/",
"/wp-content/upgrade-temp-backup/plugins/", "/wp-content/upgrade-temp-backup/themes/",
} {
if cachedPathExists(wpRoot + dir + unpacked) {
return true
}
}
return false
}
type stagedPackageFile struct {
path string
procInfo string
v wpcheck.Verification
pkg wpStagedPackage
key stagedPackageKey
queuedAt time.Time
ticket queuehealth.Ticket
// coreHeader is the installed version.php as it was when a core file
// was queued; a zero stamp means there was none.
coreHeader fileStamp
}
// A pathname can be reused by a later unpack. Retained headers only belong
// to the tree whose directory identity was observed with the queued file.
type stagedPackageKey struct {
root string
dev uint64
ino uint64
}
// stagedPackageInfo is what the directory-level finding says about a package.
// installed is decided on the package's first file, while the old tree or its
// rollback copy is still on disk.
type stagedPackageInfo struct {
installed bool
identity wpcheck.Verification
lastSeen time.Time
}
type stagedPackageQueue struct {
mu sync.Mutex
limit int
files []stagedPackageFile
draining []stagedPackageFile
packages map[stagedPackageKey]stagedPackageInfo
health *queuehealth.Tracker
now func() time.Time
}
func newStagedPackageQueue(limit int) *stagedPackageQueue {
return &stagedPackageQueue{
limit: limit, packages: make(map[stagedPackageKey]stagedPackageInfo),
health: queuehealth.NewSharedCapacity(limit, time.Minute), now: time.Now,
}
}
func (q *stagedPackageQueue) snapshot(now time.Time) queuehealth.Status {
return q.health.Snapshot(now)
}
func (q *stagedPackageQueue) pendingCount() int {
q.mu.Lock()
defer q.mu.Unlock()
return len(q.files) + len(q.draining)
}
// note records a package the analyzer saw a file of and returns its info,
// resolving installed-ness on first sight and the identity as soon as a
// description carries one.
func (q *stagedPackageQueue) note(pkg wpStagedPackage, v wpcheck.Verification, now time.Time) (stagedPackageKey, stagedPackageInfo) {
key := stagedPackageKey{root: pkg.dir + "/" + pkg.unpacked}
// Use the identity captured with the header, before content scanning:
// by now WordPress may already have renamed or replaced this pathname.
if v.RootInfo != nil && v.Root == key.root {
if st, ok := v.RootInfo.Sys().(*syscall.Stat_t); ok {
key.dev, key.ino = st.Dev, st.Ino
}
}
return key, q.record(pkg, key, v, now)
}
func (q *stagedPackageQueue) record(pkg wpStagedPackage, key stagedPackageKey, v wpcheck.Verification, now time.Time) stagedPackageInfo {
q.mu.Lock()
defer q.mu.Unlock()
info, ok := q.packages[key]
if !ok {
info.installed = wpPackageInstalled(pkg.wpRoot, pkg.unpacked)
}
if v.Kind != wpcheck.KindNone && (info.identity.Kind == wpcheck.KindNone || v.Version != "" ||
(v.Kind == wpcheck.KindTheme && info.identity.Version == "")) {
info.identity = v
info.identity.Rel, info.identity.Digest = "", ""
// Drains can discover the header at its installed location. Keep
// its identity tied to the original tree for other queued files.
info.identity.Root = key.root
}
info.lastSeen = now
q.packages[key] = info
return info
}
func (q *stagedPackageQueue) info(key stagedPackageKey) stagedPackageInfo {
q.mu.Lock()
defer q.mu.Unlock()
return q.packages[key]
}
func (q *stagedPackageQueue) push(f stagedPackageFile) bool {
q.mu.Lock()
defer q.mu.Unlock()
if len(q.files)+len(q.draining) >= q.limit {
q.health.Lose(q.now(), 1)
return false
}
f.ticket = q.health.BeginAt(f.queuedAt, q.now())
q.files = append(q.files, f)
return true
}
// take hands every queued file to the single drain loop, which requeues the
// ones still waiting. Reserve their capacity until requeue completes; new
// arrivals must not overfill the queue while the slice is detached.
func (q *stagedPackageQueue) take(now time.Time) []stagedPackageFile {
q.mu.Lock()
defer q.mu.Unlock()
files := q.files
q.files = nil
q.draining = files
for _, f := range files {
f.ticket.Start(now)
}
return files
}
func (q *stagedPackageQueue) requeue(files []stagedPackageFile, now time.Time) {
q.mu.Lock()
defer q.mu.Unlock()
keep := make(map[queuehealth.Ticket]bool, len(files))
for _, f := range files {
keep[f.ticket] = true
}
for _, f := range q.draining {
if keep[f.ticket] {
f.ticket.Requeue(now)
} else {
f.ticket.Finish(now)
}
}
q.files = append(files, q.files...)
q.draining = nil
}
// The monitor calls this after joining the analyzer and verifier workers,
// when no producer can append behind the final shutdown accounting.
func (q *stagedPackageQueue) discardPending(now time.Time) {
q.mu.Lock()
defer q.mu.Unlock()
for _, f := range q.files {
f.ticket.Reject(now)
}
q.files = nil
}
func (q *stagedPackageQueue) evictIdle(now time.Time) {
q.mu.Lock()
defer q.mu.Unlock()
for dir, info := range q.packages {
if now.Sub(info.lastSeen) > stagedPackageIdleTTL {
delete(q.packages, dir)
}
}
}
func (fm *FileMonitor) stagedPackages() *stagedPackageQueue {
fm.wpPendingInit.Do(func() {
if fm.wpPending == nil {
fm.wpPending = newStagedPackageQueue(stagedPackageQueueMax)
}
})
return fm.wpPending
}
// handleStagedPackageFile takes over a content-clean PHP file written inside
// a staged package. It returns false when path is not inside one, so the
// caller keeps its per-file warning for loose files under upgrade/.
func (fm *FileMonitor) handleStagedPackageFile(path string, v wpcheck.Verification, procInfo string) bool {
pkg := parseWPStagedPackage(path)
if pkg.dir == "" {
return false
}
q := fm.stagedPackages()
now := q.now()
key, info := q.note(pkg, v, now)
switch v.Verdict {
case wpcheck.VerdictVerified:
case wpcheck.VerdictMismatch, wpcheck.VerdictUnverifiable:
fm.alertStagedFileMismatch(path, pkg, v, procInfo)
case wpcheck.VerdictPending, wpcheck.VerdictNoVersion, wpcheck.VerdictReady:
file := stagedPackageFile{path: path, procInfo: procInfo, v: v, pkg: pkg, key: key, queuedAt: now}
if v.Kind == wpcheck.KindCore {
file.coreHeader, _ = statFileStamp(pkg.wpRoot + "/wp-includes/version.php")
}
if !q.push(file) {
fm.alertStagedPackage(pkg, info, "verification queue full", procInfo)
}
default:
fm.alertStagedPackage(pkg, info, stagedPackageReason(v), procInfo)
}
return true
}
func stagedPackageReason(v wpcheck.Verification) string {
switch {
case v.Verdict == wpcheck.VerdictUnavailable && v.Kind == wpcheck.KindTheme:
return "themes have no checksum source"
case v.Verdict == wpcheck.VerdictUnavailable:
return "not published on wordpress.org"
}
return "no plugin, theme or core header found"
}
func stagedPackageWhat(info stagedPackageInfo, pkg wpStagedPackage) string {
name := stagedPackageName(info.identity.Slug, pkg)
var what string
switch info.identity.Kind {
case wpcheck.KindCore:
what = "WordPress core"
case wpcheck.KindPlugin:
what = "plugin " + name
case wpcheck.KindTheme:
what = "theme " + name
default:
what = "package " + name
}
if info.identity.Version != "" {
what += " " + info.identity.Version
}
return what
}
func stagedPackageKind(kind wpcheck.PackageKind) string {
switch kind {
case wpcheck.KindCore:
return "core"
case wpcheck.KindPlugin:
return "plugin"
case wpcheck.KindTheme:
return "theme"
}
return "package"
}
// stagedPackageName is the package's own name: its slug once a header named
// it, otherwise the unpacked directory the archive created.
func stagedPackageName(slug string, pkg wpStagedPackage) string {
if slug == "" {
return pkg.unpacked
}
return slug
}
// alertStagedPackage raises the one directory-level finding for a package
// that cannot be verified file by file. Every file of the package lands on
// the same finding. The staging directory name is random per upload, so the
// alert identity is the site, package and reason: uploading the same
// unverifiable release again is a repeat, not a new condition. Incomplete
// headers cannot identify a release, so those warnings stay upload-specific.
func (fm *FileMonitor) alertStagedPackage(pkg wpStagedPackage, info stagedPackageInfo, reason, procInfo string) {
what := stagedPackageWhat(info, pkg)
state := "new install"
if info.installed {
state = "updates installed " + strings.TrimSuffix(what, " "+info.identity.Version)
}
var dedupKey string
if info.identity.Kind != wpcheck.KindNone && info.identity.Version != "" {
dedupKey = fmt.Sprintf("staged-package site=%q type=%s name=%q version=%q reason=%q",
pkg.wpRoot, stagedPackageKind(info.identity.Kind), stagedPackageName(info.identity.Slug, pkg),
info.identity.Version, reason)
}
fm.sendFileFinding(alert.Finding{
Severity: alert.Warning,
Check: "php_in_sensitive_dir_realtime",
Message: fmt.Sprintf("WordPress package staged, not verified against wordpress.org: %s", pkg.dir),
Details: fmt.Sprintf("%s; %s; %s. Content findings are reported separately for individual files.", what, state, reason),
DedupKey: dedupKey,
FilePath: pkg.dir,
ProcessInfo: procInfo,
})
}
// stagedFileInstalledPath is where WordPress puts a staged file once the
// package is moved into place.
func stagedFileInstalledPath(pkg wpStagedPackage, v wpcheck.Verification) string {
if v.Rel == "" {
return ""
}
switch v.Kind {
case wpcheck.KindPlugin:
if v.Slug == "" {
return ""
}
return pkg.wpRoot + "/wp-content/plugins/" + v.Slug + "/" + v.Rel
case wpcheck.KindCore:
return pkg.wpRoot + "/" + v.Rel
}
return ""
}
func pathPresent(path string) bool {
_, err := os.Lstat(path)
return err == nil
}
// alertStagedFileMismatch raises the per-file finding for a staged file that
// is not what the official package ships. If WordPress has already moved the
// tree, the finding names the installed path so it can be acted on there.
func (fm *FileMonitor) alertStagedFileMismatch(stagedPath string, pkg wpStagedPackage, v wpcheck.Verification, procInfo string) {
reportPath := stagedPath
moved := false
if !pathPresent(stagedPath) {
if installed := stagedFileInstalledPath(pkg, v); installed != "" && pathPresent(installed) {
reportPath, moved = installed, true
}
}
what := v.Slug + " " + v.Version
if v.Kind == wpcheck.KindCore {
what = "WordPress " + v.Version
}
message := fmt.Sprintf("Staged file does not match wordpress.org %s: %s", what, reportPath)
details := "Content rules found nothing. The file differs from the official package or is not part of it."
if v.Verdict == wpcheck.VerdictUnverifiable {
message = fmt.Sprintf("Staged file could not be verified against wordpress.org %s: %s", what, reportPath)
details = "Content rules found nothing. The file could not be hashed completely, so the official checksum could not be compared."
}
if moved {
details += fmt.Sprintf(" Staged at %s, since installed.", stagedPath)
}
// Identify the file by its place in the package, not the random staging
// directory, and by its digest so different bytes stay a new finding.
// An absent digest is not evidence of equal content across uploads.
file := strings.TrimPrefix(stagedPath, pkg.dir+"/")
var dedupKey string
if v.Digest != "" {
dedupKey = fmt.Sprintf("staged-file site=%q type=%s name=%q version=%q file=%q verdict=%s digest=%q",
pkg.wpRoot, stagedPackageKind(v.Kind), stagedPackageName(v.Slug, pkg), v.Version, file, v.Verdict, v.Digest)
}
fm.sendFileFinding(alert.Finding{
Severity: alert.Warning,
Check: "php_in_sensitive_dir_realtime",
Message: message,
Details: details,
DedupKey: dedupKey,
FilePath: reportPath,
ProcessInfo: procInfo,
})
}
// redescribeStaged retries identifying a file queued before its package
// header existed, from the staged path and then from the installed one.
// The digest taken at event time is kept: a rename does not change content,
// and a later write is a new event.
func (fm *FileMonitor) redescribeStaged(f stagedPackageFile) wpcheck.Verification {
// A later event may have identified this same tree before WordPress
// removed it. Prefer that header to an installed copy that may still be
// the old release; never borrow an identity from a sibling unpacked tree.
identity := fm.stagedPackages().info(f.key).identity
if f.key.ino != 0 && identity.Root == f.v.Root && identity.Version != "" {
identity.Rel, identity.Digest = f.v.Rel, f.v.Digest
// Another file's comparison result says nothing about this digest.
if identity.Kind == wpcheck.KindCore || identity.Kind == wpcheck.KindPlugin {
identity.Verdict = wpcheck.VerdictPending
}
return identity
}
unresolved := func(v wpcheck.Verification) bool {
return v.Verdict == wpcheck.VerdictNoVersion || v.Verdict == wpcheck.VerdictUnknown
}
var st unix.Stat_t
err := unix.Lstat(f.v.Root, &st)
if err == nil && (f.key.ino == 0 || st.Dev != f.key.dev || st.Ino != f.key.ino) {
return f.v // a replacement tree cannot identify an older event
}
v := fm.wpCache.Describe(f.path)
if !unresolved(v) {
if f.v.RootInfo == nil || v.RootInfo == nil || !os.SameFile(f.v.RootInfo, v.RootInfo) {
return f.v // the tree changed during header detection
}
v.Digest = f.v.Digest
return v
}
if !os.IsNotExist(err) {
return f.v
}
if f.v.Kind == wpcheck.KindCore {
return fm.redescribeCoreFromInstalled(f)
}
// Only a renamed plugin directory can link an installed header to this
// staging tree. A plugin copy install loses that link: a recent ctime
// (including chmod or child creation) is not evidence of origin.
if f.v.Kind != wpcheck.KindPlugin || f.v.RootInfo == nil {
return f.v
}
if installed := stagedFileInstalledPath(f.pkg, f.v); installed != "" {
root := f.pkg.wpRoot + "/wp-content/plugins/" + f.v.Slug
before, statErr := os.Lstat(root)
if statErr != nil || !before.IsDir() || !os.SameFile(f.v.RootInfo, before) {
return f.v
}
if iv := fm.wpCache.Describe(installed); !unresolved(iv) {
after, statErr := os.Lstat(root)
if statErr != nil || !after.IsDir() || !os.SameFile(before, after) ||
iv.Root != root || iv.Kind != wpcheck.KindPlugin || iv.Slug != f.v.Slug {
return f.v // the installed tree changed during header detection
}
v = iv
v.Rel = f.v.Rel
}
}
if unresolved(v) {
return f.v
}
v.Digest, v.Staged = f.v.Digest, true
return v
}
// redescribeCoreFromInstalled identifies a core file whose staged tree is gone
// by the installed version.php. A core update copies files into place, so no
// inode links the two trees; the installed header names this release only if
// it was replaced or rewritten since the file was queued. A refused or failed
// update leaves the old release, whose manifest would call every new file
// modified. Borrowing cannot make a modified file pass: its digest must still
// equal the official bytes of that path in the borrowed release.
func (fm *FileMonitor) redescribeCoreFromInstalled(f stagedPackageFile) wpcheck.Verification {
header := f.pkg.wpRoot + "/wp-includes/version.php"
now, ok := statFileStamp(header)
if !ok || now == f.coreHeader {
return f.v // keep waiting: a late event may still identify the tree
}
iv := fm.wpCache.Describe(header)
if iv.Kind != wpcheck.KindCore || iv.Version == "" ||
iv.Verdict == wpcheck.VerdictNoVersion || iv.Verdict == wpcheck.VerdictUnknown {
return f.v
}
if after, ok := statFileStamp(header); !ok || after != now {
return f.v // the header changed during detection
}
iv.Rel, iv.Digest, iv.Staged = f.v.Rel, f.v.Digest, true
return iv
}
// fileStamp identifies one version of a regular file: a replacement changes
// the inode, a rewrite or metadata change moves the change time.
type fileStamp struct {
dev, ino uint64
ctime int64
}
func statFileStamp(path string) (fileStamp, bool) {
var st unix.Stat_t
if err := unix.Lstat(path, &st); err != nil || st.Mode&unix.S_IFMT != unix.S_IFREG {
return fileStamp{}, false
}
return fileStamp{dev: uint64(st.Dev), ino: st.Ino, ctime: st.Ctim.Nano()}, true // #nosec G115 -- device numbers are non-negative
}
// drainStagedPackages resolves every queued file whose checksums have
// arrived and reports packages that ran out of time.
func (fm *FileMonitor) drainStagedPackages(now time.Time) {
q := fm.stagedPackages()
q.evictIdle(now)
files := q.take(time.Now())
if len(files) == 0 {
return
}
var keep []stagedPackageFile
for _, f := range files {
v := f.v
if v.Verdict == wpcheck.VerdictNoVersion || v.Version == "" {
v = fm.redescribeStaged(f)
q.record(f.pkg, f.key, v, now)
}
verdict := v.Verdict
if verdict == wpcheck.VerdictPending || verdict == wpcheck.VerdictReady {
verdict = fm.wpCache.Verify(v)
}
switch verdict {
case wpcheck.VerdictVerified:
case wpcheck.VerdictMismatch, wpcheck.VerdictUnverifiable:
v.Verdict = verdict
fm.alertStagedFileMismatch(f.path, f.pkg, v, f.procInfo)
case wpcheck.VerdictPending, wpcheck.VerdictNoVersion, wpcheck.VerdictReady:
if now.Sub(f.queuedAt) > stagedPackageTimeout {
if _, err := os.Lstat(f.key.root); v.Version == "" && os.IsNotExist(err) {
// Use the normal alert cooldown, not a permanent claim on a
// path that a later unidentified upload may reuse.
fm.alertStagedPackage(f.pkg, q.info(f.key),
"staged tree removed before its release could be identified; installation could not be confirmed", f.procInfo)
continue
}
fm.alertStagedPackage(f.pkg, q.info(f.key),
fmt.Sprintf("checksums not fetched within %s", stagedPackageTimeout), f.procInfo)
continue
}
f.v = v
keep = append(keep, f)
default:
fm.alertStagedPackage(f.pkg, q.info(f.key), stagedPackageReason(v), f.procInfo)
}
}
q.requeue(keep, time.Now())
}
// stagedPackageLoop drains the queue once a second until the monitor stops.
func (fm *FileMonitor) stagedPackageLoop() {
defer fm.wg.Done()
ticker := time.NewTicker(time.Second)
defer ticker.Stop()
for {
select {
case <-fm.stopCh:
return
case now := <-ticker.C:
fm.drainStagedPackages(now)
}
}
}
//go:build linux
package daemon
import (
"crypto/sha256"
"errors"
"fmt"
"os"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/yara"
"github.com/pidginhost/csm/internal/yaraipc"
)
// scanRealtimeYARA preserves the event's bounded read across an IPC retry.
// Reopening the event path can read a replacement, and a full-file scan rejects
// a larger file when data holds only its prefix. A sealed memfd gives the worker
// exactly the original bytes without creating a file on a monitored mount.
func scanRealtimeYARA(backend yara.Backend, path string, data []byte) ([]yara.Match, string, error) {
matches, err := yara.ScanBytesChecked(backend, path, data)
if !errors.Is(err, yaraipc.ErrPayloadTooLarge) {
return matches, "", err
}
fd, err := unix.MemfdCreate("csm-yara-snapshot", unix.MFD_CLOEXEC|unix.MFD_ALLOW_SEALING)
if err != nil {
return nil, "", fmt.Errorf("create YARA snapshot: %w", err)
}
f := os.NewFile(uintptr(fd), "csm-yara-snapshot")
defer f.Close()
if _, err = f.Write(data); err != nil {
return nil, "", fmt.Errorf("write YARA snapshot: %w", err)
}
if _, err = unix.FcntlInt(f.Fd(), unix.F_ADD_SEALS, unix.F_SEAL_WRITE|unix.F_SEAL_GROW|unix.F_SEAL_SHRINK|unix.F_SEAL_SEAL); err != nil {
return nil, "", fmt.Errorf("seal YARA snapshot: %w", err)
}
// The worker has a different descriptor table. Keep our descriptor alive
// until its synchronous scan returns, and name our process explicitly.
snapshotPath := fmt.Sprintf("/proc/%d/fd/%d", os.Getpid(), fd)
result, err := yara.ScanFileChecked(backend, snapshotPath, len(data))
if err != nil {
return nil, "", err
}
digest := fmt.Sprintf("%x", sha256.Sum256(data))
if result.ContentSHA256 != digest {
return nil, "", errors.New("YARA retry scanned different content than the event snapshot")
}
return result.Matches, digest, nil
}
package daemon
import (
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/broadcast"
)
func (d *Daemon) installFindingBus() {
bus := broadcast.NewBus(64)
d.registerQueueSource("events", bus)
d.findingBus = bus
alert.FindingBus = bus
}
package daemon
// closeFindingBus closes the finding broadcast bus at shutdown and leaves it
// installed in alert.FindingBus. Close turns Publish into a no-op, which is
// all shutdown needs; clearing the package-level interface as well raced
// the untracked control-socket and web UI goroutines that read it inside
// alert.Dispatch (a torn interface read is a nil-receiver panic).
func (d *Daemon) closeFindingBus() {
if d.findingBus != nil {
d.findingBus.Close()
}
}
package daemon
import (
"encoding/json"
"fmt"
"strings"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/firewall"
csmlog "github.com/pidginhost/csm/internal/log"
)
// firewallActionBoundary is the durable-action surface an engine exposes once
// a lifecycle is attached. Keeping it an interface means the daemon compiles
// on platforms whose engine has no lifecycle at all.
type firewallActionBoundary interface {
DurableActionsEnabled() bool
RecoverActions() error
PendingActions() ([]firewall.FirewallAction, error)
ResolveAction(id, outcome, detail string) (firewall.FirewallAction, error)
}
// recoverFirewallActions settles outcomes left uncertain by a crash or a
// kernel that could not answer. Until every action is settled the engine
// refuses new mutations, so this runs at startup and on the maintenance tick.
func recoverFirewallActions(engine any) {
boundary, ok := engine.(firewallActionBoundary)
if !ok || boundary == nil || !boundary.DurableActionsEnabled() {
return
}
if err := boundary.RecoverActions(); err != nil {
csmlog.Error("firewall action recovery incomplete", "err", err)
}
}
func (c *ControlListener) firewallActions() (firewallActionBoundary, error) {
if c.d.fwActions == nil || !c.d.fwActions.DurableActionsEnabled() {
return nil, fmt.Errorf("durable firewall actions are not active")
}
return c.d.fwActions, nil
}
// handleFirewallActions reports the actions that still block firewall
// mutations, with the evidence an operator needs to decide what happened.
func (c *ControlListener) handleFirewallActions(json.RawMessage) (any, error) {
boundary, err := c.firewallActions()
if err != nil {
return nil, err
}
pending, err := boundary.PendingActions()
if err != nil {
return nil, err
}
if len(pending) == 0 {
return control.FirewallListResult{Lines: []string{"No firewall actions are waiting for recovery."}}, nil
}
lines := make([]string, 0, len(pending)*2)
for _, a := range pending {
lines = append(lines,
fmt.Sprintf("%s %s %s %s", a.UpdatedAt.Format("2006-01-02 15:04:05"), a.Phase, a.Request.Operation, a.Request.Target),
fmt.Sprintf(" id %s actor %s source %s", a.Request.ID, a.Request.Actor, a.Request.Source),
)
if a.Request.Reason != "" {
lines = append(lines, " reason "+a.Request.Reason)
}
if a.Detail != "" {
lines = append(lines, " detail "+a.Detail)
}
if a.Phase == "unknown" {
lines = append(lines, fmt.Sprintf(" resolve with: csm firewall actions resolve %s applied|rejected", a.Request.ID))
}
}
return control.FirewallListResult{Lines: lines}, nil
}
// handleFirewallActionResolve records what an operator established by hand.
// The engine still prefers kernel evidence when it can prove the outcome.
func (c *ControlListener) handleFirewallActionResolve(argsRaw json.RawMessage) (any, error) {
var args control.FirewallActionResolveArgs
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &args); err != nil {
return nil, fmt.Errorf("parsing args: %w", err)
}
}
if strings.TrimSpace(args.ID) == "" {
return nil, fmt.Errorf("resolve needs the action id")
}
var outcome string
switch args.Outcome {
case "applied":
outcome = "verified"
case "rejected":
outcome = "failed"
default:
return nil, fmt.Errorf("outcome must be applied or rejected, not %q", args.Outcome)
}
boundary, err := c.firewallActions()
if err != nil {
return nil, err
}
detail := "operator resolved via cli"
if note := strings.TrimSpace(args.Note); note != "" {
detail += ": " + note
}
resolved, err := boundary.ResolveAction(args.ID, outcome, detail)
if err != nil {
return nil, err
}
message := fmt.Sprintf("action %s recorded as %s", args.ID, resolved.Phase)
if resolved.Phase != outcome {
message = fmt.Sprintf("action %s was proven %s by the firewall itself; the operator decision was not used", args.ID, resolved.Phase)
}
return control.FirewallAckResult{Message: message}, nil
}
package daemon
import (
"fmt"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
)
// loadEffectiveFirewallFromDisk reads csm.yaml (and conf.d) again and returns
// the firewall block the daemon would build rules from at its next start.
// A SIGHUP reload deliberately leaves restart-required blocks such as
// firewall untouched in the live config, so the re-apply commands, whose
// whole purpose is to apply an edited firewall block under a deadman,
// must read the file rather than the live snapshot.
func loadEffectiveFirewallFromDisk(configFile, confDir string) (*firewall.FirewallConfig, error) {
if configFile == "" {
return nil, fmt.Errorf("config file path unknown")
}
cfg, err := config.LoadWithDir(configFile, confDir)
if err != nil {
return nil, fmt.Errorf("reading %s: %w", configFile, err)
}
for _, res := range config.Validate(cfg) {
if res.Level == "error" && len(res.Field) >= 8 && res.Field[:8] == "firewall" {
return nil, fmt.Errorf("%s: %s", res.Field, res.Message)
}
}
effective := config.EffectiveFirewallConfig(cfg)
if effective == nil || !effective.Enabled {
return nil, fmt.Errorf("firewall is disabled in %s; re-apply would drop the ruleset", configFile)
}
return effective, nil
}
// refreshFirewallFromDisk installs the on-disk firewall block into the
// running engine and returns the configuration it replaced, so a rollback
// can put it back. The previous configuration is returned even when the
// engine holds none (nil) so callers can always restore.
func (d *Daemon) refreshFirewallFromDisk() (previous *firewall.FirewallConfig, err error) {
if d.fwEngine == nil {
return nil, fmt.Errorf("firewall engine not running")
}
cfg := d.currentCfg()
effective, err := loadEffectiveFirewallFromDisk(cfg.ConfigFile, cfg.ConfigDir)
if err != nil {
return nil, err
}
previous = d.fwEngine.Config()
d.fwEngine.SetConfig(effective)
return previous, nil
}
package daemon
import (
"errors"
"fmt"
"os"
"path/filepath"
)
// firewallStateSnapshotPath is where the apply-confirmed window keeps the
// pre-apply copy of state.json, beside the nft ruleset snapshot.
func firewallStateSnapshotPath(rollbackFile string) string {
return rollbackFile + ".state.json"
}
func firewallStateFileFor(rollbackFile string) string {
return filepath.Join(filepath.Dir(rollbackFile), "state.json")
}
// snapshotFirewallState copies state.json next to the rollback ruleset. The
// deadman used to restore the kernel snapshot alone, so an address
// unblocked inside the window came back blocked in the kernel while
// state.json, the UI and `csm firewall status` still said it was free.
// A missing state.json is recorded as an empty snapshot so the restore
// removes whatever the window wrote.
func snapshotFirewallState(rollbackFile string) error {
data, err := os.ReadFile(firewallStateFileFor(rollbackFile)) // #nosec G304 -- CSM-owned state dir.
if err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("reading firewall state for rollback: %w", err)
}
// #nosec G306 G703 -- root-only state dir.
if err := os.WriteFile(firewallStateSnapshotPath(rollbackFile), data, 0o600); err != nil {
return fmt.Errorf("writing firewall state snapshot: %w", err)
}
return nil
}
// restoreFirewallStateSnapshot puts the snapshotted state.json back and
// removes the snapshot. No snapshot (an older window) is not an error.
func restoreFirewallStateSnapshot(rollbackFile string) error {
snap := firewallStateSnapshotPath(rollbackFile)
data, err := os.ReadFile(snap) // #nosec G304 -- CSM-owned state dir.
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil
}
return fmt.Errorf("reading firewall state snapshot: %w", err)
}
stateFile := firewallStateFileFor(rollbackFile)
if len(data) == 0 {
if err := os.Remove(stateFile); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("removing firewall state written inside the window: %w", err)
}
} else if err := os.WriteFile(stateFile, data, 0o600); err != nil { // #nosec G306 G703 -- root-only state dir.
return fmt.Errorf("restoring firewall state: %w", err)
}
return removeFileIfExists(snap)
}
package daemon
import (
"context"
"fmt"
"net"
"os"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/mailranges"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/threatintel"
)
type firewallStartupOps struct {
newEngine func(*firewall.FirewallConfig, string) (*firewall.Engine, error)
apply func(*firewall.Engine) error
delays []time.Duration
}
func (d *Daemon) startFirewall() {
d.startFirewallUsing(firewallStartupOps{
newEngine: firewall.NewEngine,
apply: (*firewall.Engine).Apply,
delays: []time.Duration{time.Second, 2 * time.Second},
})
}
// Retry before publishing the engine or starting its consumers. Each attempt
// gets a fresh netlink transaction; a failed Apply preserves the kernel rules.
func retryFirewallStartup(stop <-chan struct{}, delays []time.Duration, attempt func() (*firewall.Engine, error)) (*firewall.Engine, error) {
for i := 0; ; i++ {
select {
case <-stop:
return nil, context.Canceled
default:
}
engine, err := attempt()
if err == nil {
return engine, nil
}
if i == len(delays) {
return nil, err
}
timer := time.NewTimer(delays[i])
select {
case <-stop:
timer.Stop()
return nil, context.Canceled
case <-timer.C:
}
}
}
func (d *Daemon) startFirewallUsing(ops firewallStartupOps) {
if err := checks.InitAutoBlockQueueHealth(d.cfg.StatePath); err != nil {
csmlog.Error("auto-block retry state unreadable", "err", err)
}
effectiveFirewall := config.EffectiveFirewallConfig(d.cfg)
if effectiveFirewall == nil || !effectiveFirewall.Enabled {
return
}
engine, err := retryFirewallStartup(d.stopCh, ops.delays, func() (*firewall.Engine, error) {
return d.prepareFirewall(effectiveFirewall, ops)
})
if err != nil {
d.fwStartupError = err.Error()
csmlog.Error("firewall remains unmanaged after startup attempts", "err", err)
return
}
d.fwStartupError = ""
// Apply does not consult the verdict callback. Install the shutdown
// context only after a successful firewall setup so a failed init
// does not leave behind a stopCh waiter.
verdictCtx, cancelVerdict := context.WithCancel(context.Background())
go func() {
<-d.stopCh
cancelVerdict()
}()
engine.SetShutdownContext(verdictCtx)
d.setFirewallEngine(engine)
// Set firewall engine for auto-blocking
checks.SetIPBlocker(engine)
// Prune auto-response subnet blocks that now intersect the DoS-exempt set.
// The mail-provider cache is loaded (initMailRanges ran before startFirewall)
// and Apply has completed, so the exempt set is current.
checks.PruneExemptAutoSubnets(d.cfg, engine)
// Wire the incident firewall hand-off through the ApplyBlock chokepoint
// so the correlator distinguishes live mutation from dry-run and no-op
// outcomes AND spray blocks leave the standard evidence trail.
SetIncidentSprayBlocker(d.applyIncidentSprayBlock)
fwState, _ := firewall.LoadState(d.cfg.StatePath)
csmlog.Info("firewall active",
"blocked_ips", len(fwState.Blocked),
"allowed_ips", len(fwState.Allowed),
)
// Start Dynamic DNS resolver if configured. The same resolver
// loop also services hostnames listed under infra_ips so they get
// DNS-refreshed into the engine's infra-block guard; otherwise the
// hostname entries would only protect operators whose IPs never
// move, which defeats the point of listing them by name.
infraHosts := infraHostnames(effectiveFirewall.InfraIPs)
dynHosts := append([]string{}, effectiveFirewall.DynDNSHosts...)
for _, h := range infraHosts {
if !containsString(dynHosts, h) {
dynHosts = append(dynHosts, h)
}
}
if len(dynHosts) > 0 {
resolver := firewall.NewDynDNSResolver(dynHosts, engine)
resolver.SetInfraEngine(engine)
for _, h := range infraHosts {
resolver.RegisterInfraHost(h)
}
resolver.SetFindingSink(func(host string) {
if !alert.TryEnqueue(d.alertCh, dynDNSUnresolvableFinding(host)) {
atomic.AddInt64(&d.droppedAlerts, 1)
fmt.Fprintf(os.Stderr, "[%s] alert channel full, dropping dyndns guard finding: %s\n", ts(), host)
}
})
d.wg.Add(1)
obs.Go("dyndns-resolver", func() {
defer d.wg.Done()
resolver.Run(d.stopCh)
})
csmlog.Info("DynDNS resolver active", "hosts", len(dynHosts), "infra_hosts", len(infraHosts))
}
// Start Cloudflare IP whitelist refresh if configured
if d.cfg.Cloudflare.Enabled {
d.wg.Add(1)
obs.Go("cloudflare-refresh", d.cloudflareRefreshLoop)
csmlog.Info("cloudflare IP whitelist enabled", "refresh_hours", d.cfg.Cloudflare.RefreshHours)
}
}
func (d *Daemon) prepareFirewall(effectiveFirewall *firewall.FirewallConfig, ops firewallStartupOps) (*firewall.Engine, error) {
engine, err := ops.newEngine(effectiveFirewall, d.cfg.StatePath)
if err != nil {
return nil, fmt.Errorf("initializing firewall: %w", err)
}
// Wire dry-run + verdict callbacks BEFORE Apply() and before the
// engine is exposed via d.fwEngine / checks.SetIPBlocker. The
// auto_response.dry_run safety default is "on": if any code path
// reaches engine.BlockIP while these callbacks are still nil, the
// engine treats dry-run as off and the block lands live, defeating
// the operator's stated intent. Wiring before exposure removes the
// boot-time race window entirely.
engine.SetDryRunRecorder(func(ip, reason string, timeout time.Duration) {
if db := store.Global(); db != nil {
db.RecordDryRunBlock(ip, reason, timeout)
}
})
engine.SetDryRunEnabledFunc(d.autoResponseDryRunEnabled)
engine.SetVerdictAsker(d.askVerdictCallback)
// The auto-block path skips published-crawler IPs so a high-volume bot is
// never re-added to blocked_ips behind the operator allowlist. Built-in and
// operator verified_bots ranges both flow through this lookup.
engine.SetSoftAllowChecker(func(ip string) bool {
parsed := net.ParseIP(ip)
return parsed != nil && threatintel.IPInAnyVerifiedBotRange(parsed)
})
// Push the mail-provider ranges loaded by initMailRanges() into the engine
// before Apply() so the dos_exempt_nets interval sets are populated in the
// first nftables transaction. initMailRanges() runs before startFirewall()
// so ProviderNets() always returns the cached or embedded snapshot here.
engine.SetDOSExemptProviderNets(mailranges.ProviderNets())
// Apply is itself a mutation and refuses pending actions. Recover first,
// retaining the boundary for operator resolution even if Apply fails.
d.fwActions, _ = any(engine).(firewallActionBoundary)
recoverFirewallActions(d.fwActions)
if err := ops.apply(engine); err != nil {
return nil, fmt.Errorf("applying firewall: %w", err)
}
return engine, nil
}
package daemon
import (
"time"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/mailfwd/adapter"
"github.com/pidginhost/csm/internal/mailfwd/guard"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/store"
)
// forwardGuardBadIPScore is the reputation score at/above which a sender IP is
// treated as bad for the forward-guard's bad_sender_ip signal.
const forwardGuardBadIPScore = 50
// forwardGuardRefreshInterval is how often the bad-IP lookup file is refreshed
// from the reputation DB. exim reads the lsearch file per lookup, so this only
// rewrites a file -- no exim rebuild or reload.
const forwardGuardRefreshInterval = 15 * time.Minute
// forwardGuardReconciler builds the reconciler for the current host. The guard
// is only active on cPanel/exim; elsewhere Reconcile/RefreshBadIPs are no-ops.
func (d *Daemon) forwardGuardReconciler() guard.Reconciler {
return guard.Reconciler{
Guard: adapter.NewEximServiceAdapter(),
Active: platform.Detect().IsCPanel() && !d.currentCfg().ObserveMode(),
BadIPs: d.forwardGuardBadIPs,
}
}
// forwardGuardBadIPs returns sender IPs the reputation DB scores as bad. Empty
// when the store is unavailable -- the guard then simply holds nothing on the
// bad-IP signal (the null-sender signal is unaffected).
func (d *Daemon) forwardGuardBadIPs() []string {
db := store.Global()
if db == nil {
return nil
}
var ips []string
for ip, e := range db.AllReputation() {
if e.Score >= forwardGuardBadIPScore {
ips = append(ips, ip)
}
}
return ips
}
// reconcileForwardGuard installs or removes the exim forward-guard to match the
// current config. Errors are logged, never fatal: a guard failure must not take
// the daemon down or block mail (fail-open).
func (d *Daemon) reconcileForwardGuard() {
fg := d.currentCfg().EmailProtection.ForwardGuard
if err := d.forwardGuardReconciler().Reconcile(fg); err != nil {
csmlog.Error("forward-guard reconcile failed", "err", err)
}
}
// forwardGuardRefresher periodically refreshes the bad-IP lookup file while the
// guard is enforcing.
func (d *Daemon) forwardGuardRefresher() {
defer d.wg.Done()
ticker := time.NewTicker(forwardGuardRefreshInterval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
fg := d.currentCfg().EmailProtection.ForwardGuard
if err := d.forwardGuardReconciler().RefreshBadIPs(fg); err != nil {
csmlog.Error("forward-guard bad-IP refresh failed", "err", err)
}
}
}
}
package daemon
import (
"fmt"
"os"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
)
// parseValiasFileForFindings parses a valiases file and returns findings.
// Used by both the realtime watcher and tests.
func parseValiasFileForFindings(path, domain string, localDomains map[string]bool, knownForwarders []string) []alert.Finding {
return parseValiasFileForFindingsFiltered(path, domain, localDomains, knownForwarders, true)
}
func parseValiasFileForFindingsFiltered(path, domain string, localDomains map[string]bool, knownForwarders []string, includeExternal bool) []alert.Finding {
// #nosec G304 -- path from cPanel valiases directory walk; operator-scoped.
f, err := os.Open(path)
if err != nil {
return nil
}
defer f.Close()
// A read error mid-file still reports the entries parsed before it.
entries, _ := checks.ParseValiasEntries(f, domain)
var findings []alert.Finding
for _, e := range entries {
localPart, mailDomain, d := e.LocalPart, e.Domain, e.Dest
if checks.IsKnownForwarder(localPart, mailDomain, d, knownForwarders) {
continue
}
if checks.IsPipeForwarder(d) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_pipe_forwarder",
Message: fmt.Sprintf("Pipe forwarder detected: %s@%s -> %s", localPart, mailDomain, d),
Details: fmt.Sprintf("Domain: %s\nLocal part: %s\nDestination: %s\nFile: %s", mailDomain, localPart, d, path),
})
continue
}
if d == "/dev/null" {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_suspicious_forwarder",
Message: fmt.Sprintf("Mail blackhole: %s@%s -> /dev/null", localPart, mailDomain),
Details: fmt.Sprintf("Domain: %s\nLocal part: %s\nDestination: /dev/null\nFile: %s", mailDomain, localPart, path),
})
continue
}
if includeExternal && checks.IsExternalDest(d, localDomains) {
msg := fmt.Sprintf("External forwarder: %s@%s -> %s", localPart, mailDomain, d)
if localPart == "*" {
msg = fmt.Sprintf("Wildcard catch-all to external: *@%s -> %s", mailDomain, d)
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_suspicious_forwarder",
Message: msg,
Details: fmt.Sprintf("Domain: %s\nLocal part: %s\nDestination: %s\nFile: %s", mailDomain, localPart, d, path),
})
}
}
return findings
}
//go:build linux
package daemon
import (
"crypto/sha256"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"time"
"unsafe"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/store"
)
// valiasesDir is the inotify watch root. It is a var (not const) so
// tests can redirect it under t.TempDir() without touching the real
// /etc/valiases directory; mirrors cronSpoolWatchDir in fanotify.go.
var valiasesDir = "/etc/valiases"
// ForwarderWatcher watches /etc/valiases/ for changes using inotify.
type ForwarderWatcher struct {
alertCh chan<- alert.Finding
knownForwarders []string
inotifyFd int
queueHealthOnce sync.Once
kernelQueue *notificationQueue
}
// NewForwarderWatcher creates a watcher for the valiases directory.
func NewForwarderWatcher(alertCh chan<- alert.Finding, knownForwarders []string) (*ForwarderWatcher, error) {
fd, err := unix.InotifyInit1(unix.IN_CLOEXEC | unix.IN_NONBLOCK)
if err != nil {
return nil, fmt.Errorf("inotify_init1: %w", err)
}
_, err = unix.InotifyAddWatch(fd, valiasesDir, unix.IN_CLOSE_WRITE)
if err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("inotify_add_watch(%s): %w", valiasesDir, err)
}
return &ForwarderWatcher{
alertCh: alertCh,
knownForwarders: knownForwarders,
inotifyFd: fd,
}, nil
}
// Run starts the watch loop. Blocks until stopCh is closed.
func (fw *ForwarderWatcher) Run(stopCh <-chan struct{}) {
fw.initQueueHealth()
defer func() { _ = fw.kernelQueue.close() }()
buf := make([]byte, 4096)
// Use a polling approach since inotify fd + stopCh coordination
// requires either epoll or periodic polling. Keep it simple.
ticker := time.NewTicker(2 * time.Second)
defer ticker.Stop()
for {
select {
case <-stopCh:
return
case <-ticker.C:
fw.readEvents(buf)
}
}
}
func (fw *ForwarderWatcher) readEvents(buf []byte) {
fw.initQueueHealth()
for {
n, err := fw.kernelQueue.read(buf, fw.processEvents)
if err != nil || n <= 0 {
return // EAGAIN or error - no more events
}
}
}
func (fw *ForwarderWatcher) processEvents(buf []byte) {
n := len(buf)
offset := 0
for offset < n {
if offset+unix.SizeofInotifyEvent > n {
break
}
// #nosec G103 -- inotify returns a packed binary stream;
// reinterpretation is required and bounded by the SizeofInotifyEvent check above.
event := (*unix.InotifyEvent)(unsafe.Pointer(&buf[offset]))
if event.Mask&unix.IN_Q_OVERFLOW != 0 {
fw.kernelQueue.losses.Lose(time.Now(), 1)
}
nameLen := int(event.Len)
if nameLen > 0 && offset+unix.SizeofInotifyEvent+nameLen <= n {
nameBytes := buf[offset+unix.SizeofInotifyEvent : offset+unix.SizeofInotifyEvent+nameLen]
// Trim null bytes
name := strings.TrimRight(string(nameBytes), "\x00")
if name != "" && !strings.HasPrefix(name, ".") {
fw.handleFileChange(name)
}
}
offset += unix.SizeofInotifyEvent + nameLen
}
}
func (fw *ForwarderWatcher) handleFileChange(domain string) {
path := filepath.Join(valiasesDir, domain)
// 2026-04-27: suppress alerts on the first observation of a valiases file.
// Account transfers via WHM rsync write the entire file fresh; alerting
// on every pre-existing forwarder buries operators in noise. Hash the
// file: baseline external destinations on first sight, alert on external
// destinations only when the hash changes, and keep pipe/dev-null checks
// active because they are dangerous even on a first observation. Mirrors
// auditValiasFile's behaviour in internal/checks/forwarder.go.
db := store.Global()
includeExternal := true
if db != nil {
// #nosec G304 -- path is filepath.Join(valiasesDir, domain) where the
// inotify event already restricted domain to a single path component.
data, err := os.ReadFile(path)
if err == nil {
currentHash := fmt.Sprintf("%x", sha256.Sum256(data))
oldHash, found := db.GetForwarderHash("valiases:" + domain)
_ = db.SetForwarderHash("valiases:"+domain, currentHash)
switch {
case !found:
includeExternal = false
case oldHash == currentHash:
return
default:
includeExternal = true
}
}
}
// Load local domains for external detection
localDomains := loadLocalDomainsForWatcher()
findings := parseValiasFileForFindingsFiltered(path, domain, localDomains, fw.knownForwarders, includeExternal)
for _, f := range findings {
f.Timestamp = time.Now()
f.Details += "\n(detected in realtime via inotify)"
if !alert.TryEnqueue(fw.alertCh, f) {
fmt.Fprintf(os.Stderr, "[%s] Warning: alert channel full, dropping forwarder finding for %s\n",
time.Now().Format("2006-01-02 15:04:05"), domain)
}
}
}
// loadLocalDomainsForWatcher reads local domain files. Separate from the checks
// package version to avoid import cycles.
func loadLocalDomainsForWatcher() map[string]bool {
domains := make(map[string]bool)
for _, path := range []string{"/etc/localdomains", "/etc/virtualdomains"} {
// #nosec G304 -- path iterates a literal slice of cPanel system files.
data, err := os.ReadFile(path)
if err != nil {
continue
}
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if idx := strings.IndexByte(line, ':'); idx > 0 {
line = strings.TrimSpace(line[:idx])
}
domains[strings.ToLower(line)] = true
}
}
return domains
}
package daemon
import (
"regexp"
"strings"
"sync"
"time"
)
// eximFrozenDedupTTL is how long an inactive queue ID remains tracked. Queue
// runs refresh the timestamp for messages that are still frozen, so one stuck
// message stays suppressed for its entire frozen lifetime. The expiry only
// bounds stale state when the daemon misses the corresponding unfreeze event.
const eximFrozenDedupTTL = 24 * time.Hour
// Sweep stale entries periodically instead of walking the entire map for every
// new frozen message. A burst of distinct frozen messages must remain O(n), not
// degrade to O(n^2) map scans.
const eximFrozenDedupPruneInterval = time.Hour
// Bound attacker-influenced queue state. Reaching the cap evicts an arbitrary
// older identity and therefore fails open to an occasional duplicate alert
// instead of allowing frozen-message churn to exhaust daemon memory.
const eximFrozenDedupMaxEntries = 10_000
// eximMessageIDPattern matches the exim queue ID as its own log field
// (e.g. "1wuZUi-0000000BrCR-0u0H"; older exims use shorter middle segments).
var eximMessageIDPattern = regexp.MustCompile(`^[0-9A-Za-z]{6}-[0-9A-Za-z]{6,11}-[0-9A-Za-z]{2,4}$`)
// eximFrozenDedup keeps the last-observed time per frozen message ID. Exim
// re-logs "Message is frozen" on every queue run for as long as the message
// stays queued, so without this one stuck bounce raises a finding every few
// minutes for days. In-memory on purpose: stale entries are bounded by the TTL,
// and losing the state on restart only costs one duplicate finding per
// still-frozen message.
var eximFrozenDedup = struct {
mu sync.Mutex
seen map[string]time.Time
nextPrune time.Time
}{seen: make(map[string]time.Time)}
func resetEximFrozenDedup() {
eximFrozenDedup.mu.Lock()
defer eximFrozenDedup.mu.Unlock()
eximFrozenDedup.seen = make(map[string]time.Time)
eximFrozenDedup.nextPrune = time.Time{}
}
type eximFrozenEvent uint8
const (
eximFrozenEventNone eximFrozenEvent = iota
eximFrozenEventFreeze
eximFrozenEventUnfreeze
)
// eximFrozenShouldAlert reports whether a mainlog line is a freeze event that
// deserves a finding. Repeated queue-run notices for the same ID are
// suppressed, while an unfreeze event clears the ID so a later re-freeze is a
// new finding. Freeze-shaped lines with no parseable queue ID fail open.
func eximFrozenShouldAlert(line string, now time.Time) bool {
id, event := parseEximFrozenEvent(line)
if event == eximFrozenEventNone {
return false
}
if event == eximFrozenEventUnfreeze {
if id != "" {
eximFrozenDedup.mu.Lock()
delete(eximFrozenDedup.seen, id)
eximFrozenDedup.mu.Unlock()
}
return false
}
if id == "" {
return true
}
eximFrozenDedup.mu.Lock()
defer eximFrozenDedup.mu.Unlock()
if eximFrozenDedup.nextPrune.IsZero() || !now.Before(eximFrozenDedup.nextPrune) {
cutoff := now.Add(-eximFrozenDedupTTL)
for queuedID, lastSeen := range eximFrozenDedup.seen {
if !lastSeen.After(cutoff) {
delete(eximFrozenDedup.seen, queuedID)
}
}
eximFrozenDedup.nextPrune = now.Add(eximFrozenDedupPruneInterval)
}
if lastSeen, ok := eximFrozenDedup.seen[id]; ok && now.Before(lastSeen.Add(eximFrozenDedupTTL)) {
eximFrozenDedup.seen[id] = now
return false
}
if _, tracked := eximFrozenDedup.seen[id]; !tracked && len(eximFrozenDedup.seen) >= eximFrozenDedupMaxEntries {
for queuedID := range eximFrozenDedup.seen {
delete(eximFrozenDedup.seen, queuedID)
break
}
}
eximFrozenDedup.seen[id] = now
return true
}
// releaseEximFrozenDedup re-arms a freeze finding that the log watcher could
// not enqueue. Without this rollback, one full alert channel would discard the
// first finding and suppress every later queue-run reminder for that message.
func releaseEximFrozenDedup(line string) {
id, event := parseEximFrozenEvent(line)
if id == "" || event != eximFrozenEventFreeze {
return
}
eximFrozenDedup.mu.Lock()
delete(eximFrozenDedup.seen, id)
eximFrozenDedup.mu.Unlock()
}
// parseEximFrozenEvent recognizes only Exim's action field, not arbitrary
// occurrences such as an attacker-controlled Subject containing "Frozen".
func parseEximFrozenEvent(line string) (string, eximFrozenEvent) {
if !strings.Contains(line, "frozen") && !strings.Contains(line, "Frozen") {
return "", eximFrozenEventNone
}
fields := strings.Fields(line)
idIndex := eximMessageIDFieldIndex(fields)
if idIndex >= 0 {
return fields[idIndex], parseEximFrozenAction(fields[idIndex+1:])
}
// Preserve fail-open behavior for a future message-ID format, but only
// when the field immediately after the candidate ID is an actual Exim
// freeze action. Treating every ID-less occurrence of "frozen" as an event
// lets subjects and router errors manufacture findings.
payloadIndex := eximLogPayloadFieldIndex(fields)
if payloadIndex < 0 {
return "", eximFrozenEventNone
}
if event := parseEximFrozenAction(fields[payloadIndex:]); event != eximFrozenEventNone {
return "", event
}
if event := parseEximFrozenAction(fields[payloadIndex+1:]); event != eximFrozenEventNone {
return "", event
}
return "", eximFrozenEventNone
}
func parseEximFrozenAction(action []string) eximFrozenEvent {
if len(action) == 0 {
return eximFrozenEventNone
}
if strings.EqualFold(action[0], "unfrozen") {
return eximFrozenEventUnfreeze
}
if strings.EqualFold(action[0], "frozen") {
return eximFrozenEventFreeze
}
if len(action) >= 3 &&
strings.EqualFold(action[0], "message") &&
strings.EqualFold(action[1], "is") &&
strings.EqualFold(action[2], "frozen") {
return eximFrozenEventFreeze
}
return eximFrozenEventNone
}
// eximMessageIDFieldIndex locates the queue ID after the timestamp. Exim can
// insert a timezone and/or PID before it when log_timezone or the pid log
// selector is enabled.
func eximMessageIDFieldIndex(fields []string) int {
index := eximLogPayloadFieldIndex(fields)
if index < 0 || !eximMessageIDPattern.MatchString(fields[index]) {
return -1
}
return index
}
func eximLogPayloadFieldIndex(fields []string) int {
if len(fields) < 3 {
return -1
}
index := 2
for metadataFields := 0; metadataFields < 2 && index < len(fields); metadataFields++ {
if !isEximLogTimezone(fields[index]) && !isEximLogPID(fields[index]) {
break
}
index++
}
if index >= len(fields) {
return -1
}
return index
}
func isEximLogTimezone(field string) bool {
if len(field) != 5 || (field[0] != '+' && field[0] != '-') {
return false
}
for _, c := range field[1:] {
if c < '0' || c > '9' {
return false
}
}
return true
}
func isEximLogPID(field string) bool {
if len(field) < 3 || field[0] != '[' || field[len(field)-1] != ']' {
return false
}
for _, c := range field[1 : len(field)-1] {
if c < '0' || c > '9' {
return false
}
}
return true
}
package daemon
import (
"strings"
"sync/atomic"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/geoip"
)
// daemonGeoIPDB is a package-level pointer to the GeoIP database, set once
// during daemon init and read by log watcher handlers for country filtering.
var daemonGeoIPDB atomic.Pointer[geoip.DB]
// setGeoIPDB stores the GeoIP database for daemon-wide use and wires the
// ASN resolver the bad_asn_outbound detector uses. Clearing the DB (nil)
// disables ASN classification.
func setGeoIPDB(db *geoip.DB) {
daemonGeoIPDB.Store(db)
if db == nil {
checks.SetASNLookup(nil)
return
}
checks.SetASNLookup(func(ip string) (uint, string) {
info := db.Lookup(ip)
return info.ASN, info.ASOrg
})
}
// getGeoIPDB returns the daemon's GeoIP database, or nil.
func getGeoIPDB() *geoip.DB {
return daemonGeoIPDB.Load()
}
// isTrustedCountry checks if an IP's country is in the trusted list.
// Returns false if GeoIP is unavailable or country can't be resolved.
func isTrustedCountry(ip string, trustedCountries []string) bool {
if len(trustedCountries) == 0 {
return false
}
db := getGeoIPDB()
if db == nil {
return false
}
info := db.Lookup(ip)
if info.Country == "" {
return false
}
for _, tc := range trustedCountries {
if strings.EqualFold(info.Country, tc) {
return true
}
}
return false
}
package daemon
import (
"fmt"
"net"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
)
// Enhanced access_log handler - catches File Manager, API failures,
// webmail logins, wp-login brute force, and xmlrpc abuse.
func parseAccessLogLineEnhanced(line string, cfg *config.Config) []alert.Finding {
var findings []alert.Finding
ip, _, _, ok := accessLogIPMethodPath(line)
if !ok {
return nil
}
if isInfraIPDaemon(ip, cfg.InfraIPs) || ip == "127.0.0.1" {
return nil
}
lineLower := strings.ToLower(line)
// File Manager write operations (port 2083)
// Only match actual write actions - not read-only calls like get_homedir.
// Skip 401/403 responses - the server rejected the request, no write occurred.
// Match against the request URI only (between first pair of quotes), not the
// full line which includes the referer URL that can contain "upload" in paths.
if strings.Contains(line, "2083") && !strings.Contains(line, "\" 401 ") && !strings.Contains(line, "\" 403 ") {
requestURI := extractRequestURI(lineLower)
filemanWriteActions := []string{
"fileman/save_file", "fileman/upload_files",
"fileman/paste", "fileman/rename", "fileman/delete",
}
for _, action := range filemanWriteActions {
if strings.Contains(requestURI, action) {
findings = append(findings, alert.Finding{
// Warning, not Critical: this fires on every File Manager
// write by any customer, and 401/403 are already skipped
// above, so it only reports authenticated operations. Its
// value is correlation with other findings on the account,
// not the single event.
Severity: alert.Warning,
Check: "cpanel_file_upload_realtime",
Message: fmt.Sprintf("cPanel File Manager write from non-infra IP: %s", ip),
Details: truncateDaemon(line, 300),
SourceIP: ip,
})
break
}
}
}
// API authentication failures (401/403)
// Suppress 401s that are stale-session artifacts from a recent password change.
// When a user changes their password, in-flight browser AJAX requests (notification
// polls, etc.) will 401 against the now-invalidated session - that's expected, not
// an attack. Real API abuse won't correlate with a recent purge for the same account.
if strings.Contains(line, "\" 401 ") || strings.Contains(line, "\" 403 ") {
if strings.Contains(lineLower, "json-api") || strings.Contains(lineLower, "/execute/") {
if !purgeTracker.isPostPurge401(ip) {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "api_auth_failure_realtime",
Message: fmt.Sprintf("cPanel API auth failure from %s", ip),
Details: truncateDaemon(line, 300),
SourceIP: ip,
})
}
}
}
// Webmail login attempts (port 2095/2096)
if !cfg.Suppressions.SuppressWebmail {
if strings.Contains(line, "2095") || strings.Contains(line, "2096") {
if strings.Contains(lineLower, "post") {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "webmail_login_realtime",
Message: fmt.Sprintf("Webmail login attempt from non-infra IP: %s", ip),
Details: truncateDaemon(line, 200),
SourceIP: ip,
})
}
}
}
// WHM login attempts (port 2086/2087). CVE-2026-41940 step 1 creates the
// preauth session via a POST to the WHM login endpoint; the CRLF
// injection lands when cpsrvd writes that session file. Surfacing every
// non-infra POST gives ops a brute-force/recon signal even on patched
// hosts. Suppressed under the cPanel-login suppression flag because WHM
// is the admin face of cPanel and shares the same noise profile.
isWHMPort := isWHMLogVhost(line)
if isWHMPort && !cfg.Suppressions.SuppressCpanelLogin {
if strings.Contains(lineLower, "post /login/?login_only=1") {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "whm_login_realtime",
Message: fmt.Sprintf("WHM login attempt from non-infra IP: %s", ip),
Details: truncateDaemon(line, 200),
SourceIP: ip,
})
}
}
// CVE-2026-41940 step 4 fingerprint: a tokenless request to a
// token-required WHM path triggers do_token_denied(), which the watchTowr
// PoC abuses to promote a CRLF-injected session record into the JSON
// cache. Legitimate WHM clients always prefix /scripts*/* with the
// /cpsessXXXXXX/ security token, so the bare prefix on a WHM port is a
// hard signature, not a heuristic. Matches both /scripts/ and /scripts2/
// because the do_token_denied() trigger is path-agnostic - the watchTowr
// PoC happens to use listaccts but any token-required endpoint works.
// Fires regardless of suppression - this is an attack IOC, not a login.
if isWHMPort && isUnauthWHMScriptsRequest(extractRequestURI(line)) {
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "whm_unauth_scripts_realtime",
Message: fmt.Sprintf("Tokenless WHM scripts request from %s (CVE-2026-41940 IOC)", ip),
Details: truncateDaemon(line, 300),
SourceIP: ip,
})
}
return findings
}
// isWHMLogVhost reports whether the access-log line was served by WHM (port
// 2086 plain or 2087 SSL). cPanel's combined log format ends every line with
// the served vhost as the final double-quoted field (e.g. "host:2087"); we
// anchor on the suffix of that field to avoid matching a port-like substring
// inside a referer URL or user-agent.
func isWHMLogVhost(line string) bool {
vhost := lastQuotedField(line)
return strings.HasSuffix(vhost, ":2087") || strings.HasSuffix(vhost, ":2086")
}
// lastQuotedField returns the content of the final double-quoted field on the
// line, or "" if there isn't a closed pair. cPanel's log writer always emits
// the served vhost as that final field.
func lastQuotedField(line string) string {
end := strings.LastIndex(line, "\"")
if end <= 0 {
return ""
}
start := strings.LastIndex(line[:end], "\"")
if start < 0 {
return ""
}
return line[start+1 : end]
}
// isUnauthWHMScriptsRequest returns true when the request URI targets a path
// under /scripts/ or /scripts2/ without a /cpsessXXXXXX/ security-token
// prefix - the literal step-4 fingerprint of CVE-2026-41940. Query strings
// are stripped before comparison.
func isUnauthWHMScriptsRequest(requestURI string) bool {
parts := strings.SplitN(requestURI, " ", 3)
if len(parts) < 2 {
return false
}
path := parts[1]
if q := strings.Index(path, "?"); q >= 0 {
path = path[:q]
}
if strings.Contains(path, "/cpsess") {
return false
}
return strings.HasPrefix(path, "/scripts/") || strings.HasPrefix(path, "/scripts2/")
}
// parseFTPLogLine handles FTP log entries from /var/log/messages.
func parseFTPLogLine(line string, cfg *config.Config) []alert.Finding {
var findings []alert.Finding
if !strings.Contains(line, "pure-ftpd") {
return nil
}
// Extract the client address. Pure-ftpd's standard syslog format
// prefixes the client as (user@addr), where addr is either an IP
// (DontResolve=yes) or a reverse-resolved hostname (cPanel's
// default with DontResolve=no). We try the pure-ftpd prefix first
// and fall back to the generic "whitespace field starting with a
// digit" scanner. If the pure-ftpd prefix contains a hostname
// rather than an IP, no finding is emitted — we can't hold an
// attacker accountable by hostname, and reverse-DNS lookups in the
// log hot path are not acceptable.
ip := extractPureFTPDClientIP(line)
if ip == "" {
ip = extractIPFromLogDaemon(line)
}
if ip == "" || isInfraIPDaemon(ip, cfg.InfraIPs) {
return nil
}
// Failed authentication
if strings.Contains(line, "Authentication failed") || strings.Contains(line, "auth failed") {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "ftp_auth_failure_realtime",
Message: fmt.Sprintf("FTP authentication failed from %s", ip),
Details: truncateDaemon(line, 200),
SourceIP: ip,
})
}
// Successful login from non-infra. The scheduled ftp_logins check reads
// the same file and will meet this line again, so the finding comes from
// the shared builder: identical findings collapse in the state store,
// while a login this watcher never saw is still reported by the scan.
// The builder also drops loopback, which panel transfers use; a local
// relay does not make the auth failure above trustworthy.
if f, ok := checks.FTPLoginFinding(line, cfg); ok {
findings = append(findings, f)
}
return findings
}
// extractPureFTPDClientIP parses the "(user@addr)" prefix that pure-ftpd
// prepends to every log message and returns `addr` only if it parses as
// an IP. Returns empty if the log line contains no prefix at all, the
// prefix is malformed, or addr is a reverse-resolved hostname (in which
// case the caller should not emit a finding since we can't block a
// hostname at the firewall).
func extractPureFTPDClientIP(line string) string {
open := strings.Index(line, "(")
if open < 0 {
return ""
}
rest := line[open+1:]
close := strings.Index(rest, ")")
if close < 0 {
return ""
}
inner := rest[:close]
at := strings.IndexByte(inner, '@')
if at < 0 {
return ""
}
addr := inner[at+1:]
if net.ParseIP(addr) == nil {
return "" // hostname, not an IP — nothing we can block
}
return addr
}
// extractRequestURI extracts the request URI from an access log line.
// Format: ... "METHOD /path HTTP/1.1" ... → returns "/path"
// Returns the content between the first pair of quotes (the request line).
func extractRequestURI(line string) string {
start := strings.Index(line, "\"")
if start < 0 {
return ""
}
end := strings.Index(line[start+1:], "\"")
if end < 0 {
return ""
}
return line[start+1 : start+1+end]
}
func extractIPFromLogDaemon(line string) string {
fields := strings.Fields(line)
for _, f := range fields {
if len(f) >= 7 && f[0] >= '0' && f[0] <= '9' && strings.Count(f, ".") == 3 {
return strings.TrimRight(f, ",:;)([]")
}
}
return ""
}
package daemon
import (
"fmt"
"os"
"sort"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
)
// Real-time access log handler for detecting wp-login.php brute force and
// xmlrpc.php abuse. Watches the LiteSpeed/Apache Combined Log Format access
// log and tracks per-IP POST counts using a sliding time window.
//
// Emits the same check names as the periodic CheckWPBruteForce
// (wp_login_bruteforce, xmlrpc_abuse) so the existing auto-block pipeline
// handles them automatically.
const (
// Sliding window for counting requests per IP.
accessLogWindow = 5 * time.Minute
// Threshold within the window. Lower than periodic checks because
// we're watching in real-time and want fast response.
accessLogWPLoginThreshold = 10
// Eviction: how often to prune expired trackers.
accessLogEvictInterval = 5 * time.Minute
// Cooldown after an IP is flagged: don't re-alert for this long.
// Prevents alert spam while auto-block processes the finding.
accessLogBlockCooldown = 30 * time.Minute
)
// accessLogTracker tracks POST timestamps per endpoint for a single IP.
type accessLogTracker struct {
mu sync.Mutex
wpLoginTimes []time.Time
xmlrpcTimes []time.Time
adminPanelTimes []time.Time
wpLoginAlerted bool
xmlrpcAlerted bool
adminPanelAlerted bool
lastSeen time.Time
generation uint64
evicting bool
}
// accessLogTrackers holds per-IP state. sync.Map for concurrent handler access.
var accessLogTrackers sync.Map // key: IP string → value: *accessLogTracker
// accessLogTrackerCount approximates the live entry count in
// accessLogTrackers. sync.Map exposes no Len(); maintaining a side
// counter is the canonical workaround. Used to trigger eager
// eviction during a DDoS burst, where the 5-min timer alone would
// let the map grow into the hundreds of thousands of unique IPs
// before the next prune.
var (
accessLogTrackerCount atomic.Int64
accessLogEagerEvictTrip = make(chan struct{}, 1)
)
// accessLogEvictSoftCap is the live-entry threshold above which the
// hot path nudges the eviction goroutine to run sooner than the
// 5-min ticker. Picked to be well below typical memory limits even
// at the worst per-tracker size (~256 bytes) so a 100k entry
// burst stays under 32 MB.
const (
accessLogEvictSoftCap int64 = 50000
accessLogEvictTargetPercent int64 = 95
)
type accessLogEvictionCandidate struct {
key string
tracker *accessLogTracker
lastSeen time.Time
generation uint64
}
// discoverAccessLogPath returns the first access log path that exists,
// consulting the platform detector for OS/web-server specific candidates.
func discoverAccessLogPath() string {
info := platform.Detect()
for _, p := range info.AccessLogPaths {
if _, err := os.Stat(p); err == nil {
return p
}
}
return ""
}
// parseAccessLogBruteForce is the LogLineHandler for the Apache/LiteSpeed
// Combined Log Format access log. It parses each line, tracks per-IP POST
// counts to wp-login.php and xmlrpc.php, and emits findings when thresholds
// are crossed.
func parseAccessLogBruteForce(line string, cfg *config.Config) []alert.Finding {
// Fast reject: only care about POST requests to known attack targets.
if !strings.Contains(line, "POST") {
return nil
}
ip, method, path, ok := accessLogIPMethodPath(line)
if !ok {
return nil
}
// Skip infra IPs and loopback.
if ip == "127.0.0.1" || ip == "::1" || isInfraIPDaemon(ip, cfg.InfraIPs) {
return nil
}
if method != "POST" {
return nil
}
isWPLogin := strings.Contains(path, "wp-login.php")
isXMLRPC := strings.Contains(path, "xmlrpc.php")
isAdminPanel := isAdminPanelPath(path)
xmlrpcThreshold := effectiveAccessLogXMLRPCThreshold(cfg)
if isXMLRPC && xmlrpcThreshold <= 0 {
isXMLRPC = false
}
if !isWPLogin && !isXMLRPC && !isAdminPanel {
return nil
}
now := time.Now()
tracker := loadAccessLogTracker(ip, now)
defer tracker.mu.Unlock()
tracker.lastSeen = now
tracker.generation++
cutoff := now.Add(-accessLogWindow)
var results []alert.Finding
// Once a per-tier alert has fired, skip pruneAndAppend until cooldown
// clears the `alerted` flag. The slice would otherwise grow on every
// event during a sustained burst (potentially tens of thousands of
// entries for a 5-min window at 100 rps), wasting CPU on prune passes
// whose result is never consumed: the `alerted` flag already prevents
// re-alerts, and the eviction loop trims the slice on its own schedule.
//
// Safety: evictAccessLogState resets `alerted` once `lastSeen` is older
// than `cooldownCutoff` (30 min of silence by default). By that point
// the same eviction call has also pruned the slice to empty (window is
// 5 min, so any remaining timestamp is far past cutoff), so the next
// matching event correctly starts a fresh count from 1.
if isWPLogin && !tracker.wpLoginAlerted {
tracker.wpLoginTimes = pruneAndAppend(tracker.wpLoginTimes, cutoff, now)
if len(tracker.wpLoginTimes) >= accessLogWPLoginThreshold {
tracker.wpLoginAlerted = true
results = append(results, alert.Finding{
Severity: alert.Critical,
Check: "wp_login_bruteforce",
Message: fmt.Sprintf("WordPress login brute force from %s: %d POSTs in %v (real-time)", ip, len(tracker.wpLoginTimes), accessLogWindow),
Details: "Real-time detection: high rate of POST requests to wp-login.php",
Timestamp: now,
SourceIP: ip,
})
}
}
if isXMLRPC && !tracker.xmlrpcAlerted {
tracker.xmlrpcTimes = pruneAndAppend(tracker.xmlrpcTimes, cutoff, now)
if len(tracker.xmlrpcTimes) >= xmlrpcThreshold {
tracker.xmlrpcAlerted = true
results = append(results, alert.Finding{
Severity: alert.Critical,
Check: "xmlrpc_abuse",
Message: fmt.Sprintf("XML-RPC abuse from %s: %d POSTs in %v (real-time)", ip, len(tracker.xmlrpcTimes), accessLogWindow),
Details: "Real-time detection: high rate of POST requests to xmlrpc.php (brute force or amplification)",
Timestamp: now,
SourceIP: ip,
})
}
}
if isAdminPanel && !tracker.adminPanelAlerted {
tracker.adminPanelTimes = pruneAndAppend(tracker.adminPanelTimes, cutoff, now)
if len(tracker.adminPanelTimes) >= accessLogWPLoginThreshold {
tracker.adminPanelAlerted = true
results = append(results, alert.Finding{
Severity: alert.Critical,
Check: "admin_panel_bruteforce",
Message: fmt.Sprintf("Admin panel brute force from %s: %d POSTs in %v (real-time)", ip, len(tracker.adminPanelTimes), accessLogWindow),
Details: "Real-time detection: high rate of POST requests to common admin panel login paths (phpMyAdmin / Joomla)",
Timestamp: now,
SourceIP: ip,
})
}
}
return results
}
func effectiveAccessLogXMLRPCThreshold(cfg *config.Config) int {
if cfg == nil {
return config.DefaultXMLRPCThreshold
}
return cfg.Thresholds.XMLRPCThreshold
}
func loadAccessLogTracker(ip string, now time.Time) *accessLogTracker {
for {
val, loaded := accessLogTrackers.LoadOrStore(ip, &accessLogTracker{lastSeen: now})
tracker := val.(*accessLogTracker)
if !loaded && accessLogTrackerCount.Add(1) > accessLogEvictSoftCap {
signalAccessLogEagerEviction()
}
tracker.mu.Lock()
if !tracker.evicting {
return tracker
}
tracker.mu.Unlock()
if accessLogTrackers.CompareAndDelete(ip, tracker) {
decrementAccessLogTrackerCount()
}
}
}
func signalAccessLogEagerEviction() {
select {
case accessLogEagerEvictTrip <- struct{}{}:
default:
}
}
// pruneAndAppend removes entries older than cutoff and appends now.
func pruneAndAppend(times []time.Time, cutoff, now time.Time) []time.Time {
recent := times[:0]
for _, t := range times {
if !t.Before(cutoff) {
recent = append(recent, t)
}
}
return append(recent, now)
}
// StartAccessLogEviction starts a background goroutine that prunes expired
// tracker entries to prevent unbounded memory growth. Same pattern as
// StartModSecEviction.
func StartAccessLogEviction(stopCh <-chan struct{}) {
obs.Go("accesslog-eviction", func() {
ticker := time.NewTicker(accessLogEvictInterval)
defer ticker.Stop()
for {
select {
case <-stopCh:
return
case now := <-ticker.C:
evictAccessLogState(now)
case <-accessLogEagerEvictTrip:
// Soft-cap signal from the hot path. Run an
// immediate eviction so a DDoS burst of unique
// IPs cannot grow the tracker map past memory
// budget before the next 5-min tick.
evictAccessLogState(time.Now())
}
}
})
}
func evictAccessLogState(now time.Time) {
evictAccessLogStateWithCap(now, accessLogEvictSoftCap)
}
func evictAccessLogStateWithCap(now time.Time, cap int64) {
cutoff := now.Add(-accessLogWindow)
cooldownCutoff := now.Add(-accessLogBlockCooldown)
candidates := make([]accessLogEvictionCandidate, 0)
accessLogTrackers.Range(func(key, value any) bool {
ip := key.(string)
tracker := value.(*accessLogTracker)
tracker.mu.Lock()
if tracker.evicting {
tracker.mu.Unlock()
return true
}
// Prune old timestamps.
tracker.wpLoginTimes = pruneSlice(tracker.wpLoginTimes, cutoff)
tracker.xmlrpcTimes = pruneSlice(tracker.xmlrpcTimes, cutoff)
tracker.adminPanelTimes = pruneSlice(tracker.adminPanelTimes, cutoff)
// Reset alerted flags after cooldown so the IP can be re-detected
// if it comes back after the block expires.
if tracker.wpLoginAlerted && tracker.lastSeen.Before(cooldownCutoff) {
tracker.wpLoginAlerted = false
}
if tracker.xmlrpcAlerted && tracker.lastSeen.Before(cooldownCutoff) {
tracker.xmlrpcAlerted = false
}
if tracker.adminPanelAlerted && tracker.lastSeen.Before(cooldownCutoff) {
tracker.adminPanelAlerted = false
}
empty := len(tracker.wpLoginTimes) == 0 && len(tracker.xmlrpcTimes) == 0 &&
len(tracker.adminPanelTimes) == 0 && tracker.lastSeen.Before(cooldownCutoff)
if empty {
deleteAccessLogTrackerLocked(ip, tracker)
} else {
candidates = append(candidates, accessLogEvictionCandidate{
key: ip,
tracker: tracker,
lastSeen: tracker.lastSeen,
generation: tracker.generation,
})
}
tracker.mu.Unlock()
return true
})
enforceAccessLogTrackerCap(candidates, cap)
}
func enforceAccessLogTrackerCap(candidates []accessLogEvictionCandidate, cap int64) {
if cap <= 0 || accessLogTrackerCount.Load() <= cap {
return
}
sort.Slice(candidates, func(i, j int) bool {
return candidates[i].lastSeen.Before(candidates[j].lastSeen)
})
target := cap * accessLogEvictTargetPercent / 100
for _, candidate := range candidates {
if accessLogTrackerCount.Load() <= target {
return
}
tracker := candidate.tracker
tracker.mu.Lock()
if !tracker.evicting &&
tracker.generation == candidate.generation &&
tracker.lastSeen.Equal(candidate.lastSeen) {
deleteAccessLogTrackerLocked(candidate.key, tracker)
}
tracker.mu.Unlock()
}
}
func deleteAccessLogTrackerLocked(key string, tracker *accessLogTracker) bool {
tracker.evicting = true
if accessLogTrackers.CompareAndDelete(key, tracker) {
decrementAccessLogTrackerCount()
return true
}
return false
}
func decrementAccessLogTrackerCount() {
for {
current := accessLogTrackerCount.Load()
if current <= 0 {
return
}
if accessLogTrackerCount.CompareAndSwap(current, current-1) {
return
}
}
}
// accessLogIPMethodPath extracts client IP, request method, and request path
// from an Apache/LiteSpeed Combined Log Format line without allocating a
// string slice. Hot path: each domlog line that survives the "POST" prefilter
// hits this. strings.Fields allocates len(fields)+1 strings per call; this
// scanner only returns sub-strings that share the input's backing array.
func accessLogIPMethodPath(line string) (ip, method, path string, ok bool) {
var fields [7]string
n := len(line)
i := 0
for f := 0; f < 7; f++ {
for i < n && isAccessLogSpace(line[i]) {
i++
}
if i >= n {
return "", "", "", false
}
start := i
for i < n && !isAccessLogSpace(line[i]) {
i++
}
fields[f] = line[start:i]
}
method = fields[5]
for len(method) > 0 && method[0] == '"' {
method = method[1:]
}
for len(method) > 0 && method[len(method)-1] == '"' {
method = method[:len(method)-1]
}
return fields[0], method, fields[6], true
}
func isAccessLogSpace(b byte) bool {
return b == ' ' || b == '\t' || b == '\n' || b == '\v' || b == '\f' || b == '\r'
}
// isAdminPanelPath returns true for high-confidence non-WP admin panel login
// paths suitable for hard-block auto-response. Drupal /user/login, Tomcat
// /manager/html, generic /admin/login.php, /mysql/ are intentionally EXCLUDED
// because they're either too generic (FP risk on shared hosting) or use a
// different attack shape (Basic auth vs. POST forms). See spec Component 5
// for the full rationale.
func isAdminPanelPath(path string) bool {
return strings.Contains(path, "/phpmyadmin/index.php") ||
strings.Contains(path, "/pma/index.php") ||
strings.Contains(path, "/phpMyAdmin/index.php") ||
strings.Contains(path, "/administrator/index.php")
}
func pruneSlice(times []time.Time, cutoff time.Time) []time.Time {
recent := times[:0]
for _, t := range times {
if !t.Before(cutoff) {
recent = append(recent, t)
}
}
return recent
}
package daemon
import (
"fmt"
"net"
"os"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/geoip"
"github.com/pidginhost/csm/internal/store"
)
const (
geoHistoryMaxAge = 30 * 24 * 60 * 60 // 30 days in seconds
geoMinLoginCount = 5 // minimum logins before alerting on new country
geoAlertCooldownH = 24 // hours between alerts per account
)
// geoLookup resolves an IP through the daemon's GeoIP database, returning an
// empty Info when no database is loaded. A variable so tests can drive the geo
// detector without an mmdb file.
var geoLookup = func(ip string) geoip.Info {
db := getGeoIPDB()
if db == nil {
return geoip.Info{}
}
return db.Lookup(ip)
}
// parseDovecotLogLine handles Dovecot login lines from /var/log/maillog.
// It tracks per-mailbox login countries and alerts on new-country logins.
func parseDovecotLogLine(line string, cfg *config.Config) []alert.Finding {
// Only process successful login lines. dovecotLoginSucceeded accepts both
// the classic "Login: user=<" and the production "Logged in: user=<"
// formats so this detector is not silently dead on whichever one the host
// emits.
if !dovecotLoginSucceeded(line) {
return nil
}
user, ip := parseDovecotLoginFields(line)
if user == "" || ip == "" {
return nil
}
// Skip private/loopback IPs
if isPrivateOrLoopback(ip) {
return nil
}
// Skip infra IPs
if isInfraIPDaemon(ip, cfg.InfraIPs) {
return nil
}
info := geoLookup(ip)
country := info.Country
if country == "" {
return nil
}
// A trusted country suppresses the alert, not the login: the logins from
// home are what carry a mailbox past the login floor and teach its
// baseline. Returning before the history update left every mailbox whose
// owner lives in a trusted country stuck at zero logins, so its first
// foreign login was recorded as a known country and never alerted.
trusted := false
for _, tc := range cfg.Suppressions.TrustedCountries {
if strings.EqualFold(country, tc) {
trusted = true
break
}
}
// Load history from bbolt
boltDB := store.Global()
if boltDB == nil {
return nil
}
now := time.Now().Unix()
history, _ := boltDB.GetGeoHistory(user)
if history.Countries == nil {
history.Countries = make(map[string]int64)
}
// Prune old country entries
history.Countries = pruneOldCountries(history.Countries, now, geoHistoryMaxAge)
// Increment login count
history.LoginCount++
// Check if this is a new country
_, countryKnown := history.Countries[country]
isNewCountry := !trusted && !countryKnown && history.LoginCount >= geoMinLoginCount
// Update country timestamp
history.Countries[country] = now
// Persist updated history
if err := boltDB.SetGeoHistory(user, history); err != nil {
fmt.Fprintf(os.Stderr, "[%s] Warning: failed to save geo history for %s: %v\n",
time.Now().Format("2006-01-02 15:04:05"), user, err)
}
if !isNewCountry {
return nil
}
// Rate limit: max 1 alert per account per 24h
alertKey := "email:geo_alert:" + user
lastAlertStr := boltDB.GetMetaString(alertKey)
if lastAlertStr != "" {
if lastAlert, err := time.Parse(time.RFC3339, lastAlertStr); err == nil {
if time.Since(lastAlert) < time.Duration(geoAlertCooldownH)*time.Hour {
return nil
}
}
}
// Record alert time
_ = boltDB.SetMetaString(alertKey, time.Now().Format(time.RFC3339))
// Build "previously seen" country list
var knownCountries []string
for c := range history.Countries {
if c != country {
knownCountries = append(knownCountries, c)
}
}
previousList := "none"
if len(knownCountries) > 0 {
previousList = strings.Join(knownCountries, ", ")
}
countryName := info.CountryName
if countryName == "" {
countryName = country
}
mailbox, domain, _ := splitMailAccount(user)
tenant := mailAccountOwner(user)
return []alert.Finding{{
Severity: alert.High,
Check: "email_suspicious_geo",
Message: fmt.Sprintf("Suspicious email login for %s from %s (%s) - previously seen: %s",
user, countryName, ip, previousList),
Details: fmt.Sprintf("Country: %s (%s)\nIP: %s\nLogin count: %d\nPreviously seen countries: %s",
country, countryName, ip, history.LoginCount, previousList),
SourceIP: ip,
Mailbox: mailbox,
Domain: domain,
TenantID: tenant,
Timestamp: time.Now(),
}}
}
// parseDovecotLoginFields extracts user and remote IP from a Dovecot login line.
// Expected format: "... Login: user=<user@domain>, ... rip=1.2.3.4, ..."
func parseDovecotLoginFields(line string) (user, ip string) {
// Extract user from user=<...>
userIdx := strings.Index(line, "user=<")
if userIdx < 0 {
return "", ""
}
rest := line[userIdx+6:]
endIdx := strings.IndexByte(rest, '>')
if endIdx < 0 {
return "", ""
}
user = rest[:endIdx]
// Extract remote IP from rip=...
ripIdx := strings.Index(line, "rip=")
if ripIdx < 0 {
return "", ""
}
rest = line[ripIdx+4:]
endIdx = strings.IndexAny(rest, ", \t\n")
if endIdx < 0 {
ip = rest
} else {
ip = rest[:endIdx]
}
if user == "" || ip == "" {
return "", ""
}
return user, ip
}
// isPrivateOrLoopback returns true if the IP is loopback or RFC1918 private.
func isPrivateOrLoopback(ipStr string) bool {
ip := net.ParseIP(ipStr)
if ip == nil {
return true // invalid = skip
}
if ip.IsLoopback() {
return true
}
// RFC1918 checks
private10 := net.IPNet{IP: net.ParseIP("10.0.0.0"), Mask: net.CIDRMask(8, 32)}
private172 := net.IPNet{IP: net.ParseIP("172.16.0.0"), Mask: net.CIDRMask(12, 32)}
private192 := net.IPNet{IP: net.ParseIP("192.168.0.0"), Mask: net.CIDRMask(16, 32)}
return private10.Contains(ip) || private172.Contains(ip) || private192.Contains(ip)
}
// pruneOldCountries removes country entries older than maxAge seconds.
func pruneOldCountries(countries map[string]int64, now, maxAge int64) map[string]int64 {
pruned := make(map[string]int64, len(countries))
cutoff := now - maxAge
for c, ts := range countries {
if ts >= cutoff {
pruned[c] = ts
}
}
return pruned
}
package daemon
import (
"fmt"
"net"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/modsec"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/store"
)
// modsecDenyEvent is one ModSecurity event recorded in an IP's sliding window.
// Warnings (isBlock=false) are recorded only when high-confidence or unknown,
// to supply fail-secure evidence to a later low-confidence anomaly-threshold
// deny in the same burst (anomaly-scoring WAFs log the specific attack rule as
// a warning, then deny on the points-threshold rule).
type modsecDenyEvent struct {
t time.Time
rule int
conf modsecConfidence
isBlock bool
}
// modsecIPCounter tracks recent ModSecurity events for a single IP.
type modsecIPCounter struct {
mu sync.Mutex
events []modsecDenyEvent
escalated bool // latched once escalation fires; reset when the window drains
lastEscalated time.Time // lets a sustained source refresh an expired firewall block
lowBurstEmitted bool // latched once the low-confidence burst finding fires
}
// modsecEscalationOutcome reports what a single recorded event should produce.
type modsecEscalationOutcome struct {
escalate bool // fire the Critical auto-block escalation finding
lowConfBurst bool // fire the non-actioned low-confidence-burst visibility finding
classifierGap bool // fire the non-actioned classifier-gap finding (new unknown rule)
}
type modsecWindowSummary struct {
totalBlock int
lowConfBlock int
highEvidence bool
unknownEvidence bool
}
var (
modsecDedup sync.Map // key: "IP:ruleID" → value: time.Time
modsecBlockCount sync.Map // key: IP → value: *modsecIPCounter
)
// Classifier gaps describe the confidence table, not a source, so they are
// tracked per rule ID for the whole host.
var (
modsecGapMu sync.Mutex
modsecGapReported = map[int]time.Time{}
)
const (
modsecDedupTTL = 60 * time.Second
modsecEvictInterval = 10 * time.Minute
modsecDefaultEscalationWin = 10 * time.Minute
modsecDefaultEscalationHits = 3
// modsecClassifierGapInterval is how long one unknown rule stays quiet after
// its classifier-gap finding before an unresolved gap is reported again.
modsecClassifierGapInterval = 24 * time.Hour
)
// modsecEscalationParams returns the operator-tuned (hits, window) pair,
// falling back to the shipped defaults when either is unset or
// non-positive. nil cfg returns the defaults so test wiring without a
// config still behaves predictably.
func modsecEscalationParams(cfg *config.Config) (int, time.Duration) {
hits := modsecDefaultEscalationHits
win := modsecDefaultEscalationWin
if cfg == nil {
return hits, win
}
if cfg.Thresholds.ModSecEscalationHits > 0 {
hits = cfg.Thresholds.ModSecEscalationHits
}
if cfg.Thresholds.ModSecEscalationWindowMin > 0 {
win = time.Duration(cfg.Thresholds.ModSecEscalationWindowMin) * time.Minute
}
return hits, win
}
// modsecDefaultLowConfEscalationHits is the shipped low-confidence-only backstop:
// how many low-confidence policy/anomaly denies from one IP within the
// escalation window force a firewall escalation even with no attack signature.
// It closes the deliberate "only trip anomaly/content-type rules" bypass while
// staying far above any single legitimate checkout-retry pattern.
const modsecDefaultLowConfEscalationHits = 30
const modsecDefaultEscalationRearmInterval = 24 * time.Hour
// modsecLowConfEscalationHits returns the operator-tuned backstop count,
// falling back to the shipped default when unset or non-positive. Raise it on
// hosts with high-volume legitimate apps that trip policy/anomaly rules in
// bulk; it is not meant to be disabled (that would reopen the bypass).
func modsecLowConfEscalationHits(cfg *config.Config) int {
if cfg == nil {
return modsecDefaultLowConfEscalationHits
}
if cfg.Thresholds.ModSecLowConfidenceEscalationHits > 0 {
return cfg.Thresholds.ModSecLowConfidenceEscalationHits
}
return modsecDefaultLowConfEscalationHits
}
// modsecEscalationRearmInterval mirrors the firewall auto-block expiry so a
// sustained attacker can refresh the ban after it would expire, without
// re-emitting duplicate escalation findings during the active block window.
func modsecEscalationRearmInterval(cfg *config.Config) time.Duration {
if cfg == nil || cfg.AutoResponse.BlockExpiry == "" {
return modsecDefaultEscalationRearmInterval
}
d, err := time.ParseDuration(cfg.AutoResponse.BlockExpiry)
if err != nil || d <= 0 {
return modsecDefaultEscalationRearmInterval
}
return d
}
// parseModSecLogLine parses a ModSecurity log line from Apache or LiteSpeed
// error logs and returns findings for blocked requests or warnings.
func parseModSecLogLine(line string, cfg *config.Config) []alert.Finding {
// Fast reject: not a ModSecurity line.
if !strings.Contains(line, "ModSecurity:") && !strings.Contains(line, "[MODSEC]") {
return nil
}
isLiteSpeed := strings.Contains(line, "[MODSEC]")
var ip, ruleID, msg, hostname, uri string
if isLiteSpeed {
ip = extractLiteSpeedIP(line)
ruleID = extractModSecField(line, `[id "`, `"]`)
msg = extractModSecField(line, `[msg "`, `"]`)
hostname = extractModSecField(line, `[hostname "`, `"]`)
uri = extractModSecField(line, `[uri "`, `"]`)
} else {
// Apache format: [client IP] or [client IP:port]
raw := extractModSecField(line, "[client ", "]")
// Strip port if present (Apache 2.4 uses "IP:port").
if idx := strings.LastIndex(raw, ":"); idx > 0 {
// Make sure it's not an IPv6 address (contains multiple colons).
if strings.Count(raw, ":") == 1 {
raw = raw[:idx]
}
}
ip = raw
ruleID = extractModSecField(line, `[id "`, `"]`)
msg = extractModSecField(line, `[msg "`, `"]`)
hostname = extractModSecField(line, `[hostname "`, `"]`)
uri = extractModSecField(line, `[uri "`, `"]`)
}
// Skip infra and loopback IPs - consistent with other realtime handlers
// (handlers.go:22, autoblock.go:149). Prevents noisy findings and
// false escalation from proxied or locally forwarded Apache traffic.
if ip != "" && (isInfraIPDaemon(ip, cfg.InfraIPs) || ip == "127.0.0.1" || ip == "::1") {
return nil
}
// Determine check name.
//
// Apache mod_security writes the action verbatim into the message, so
// "Access denied" is a reliable block signal. LiteSpeed's mod_security
// front-end writes every match as "triggered!" with no action context,
// regardless of whether the rule's declared action denied the request
// or merely incremented a counter. Without further context every match
// would be counted as a deny, escalating to a 24-hour auto-block after
// three pass-action triggers from the same IP. Consult the rule-action
// registry built at daemon start: pass/log/allow rules produce a
// warning, deny/drop/block produce a block, and an unknown rule ID in a
// populated registry stays conservative. When the registry is empty, the
// line remains a warning until a refresh loads rule actions.
check := "modsec_warning_realtime"
if strings.Contains(line, "Access denied") {
check = "modsec_block_realtime"
} else if isLiteSpeed && strings.Contains(line, "triggered!") {
check = classifyLiteSpeedTrigger(ruleID)
}
ruleNum := 0
if n, err := strconv.Atoi(ruleID); err == nil {
ruleNum = n
}
conf := classifyModSecConfidence(ruleNum, msg, extractAllModSecFields(line, `[tag "`, `"]`), extractModSecRuleFile(line))
// Determine severity from confidence. Individual low-confidence and unknown
// blocks stay Warning so the confidence-gated escalation path remains the
// only firewall-grade signal for those rules.
// Individual blocks are informational - ModSecurity already denied the
// request. Only the escalation finding (3+ from same IP) is CRITICAL
// because it triggers auto-block at the firewall level.
severity := alert.Warning
if conf == modsecConfHigh {
severity = alert.High // specific attack/probe evidence, informational
}
// Build message.
message := fmt.Sprintf("ModSecurity rule %s", ruleID)
if check == "modsec_block_realtime" {
message = fmt.Sprintf("ModSecurity blocked request: rule %s", ruleID)
}
if ip != "" {
message += fmt.Sprintf(" from %s", ip)
}
if hostname != "" {
message += fmt.Sprintf(" on %s", hostname)
}
if uri != "" {
message += fmt.Sprintf(" uri=%s", uri)
}
if msg != "" {
message += fmt.Sprintf(" - %s", msg)
}
// Store structured details so the web UI can extract fields consistently
// regardless of whether the source was Apache or LiteSpeed format.
details := fmt.Sprintf("Rule: %s\nMessage: %s\nHostname: %s\nURI: %s\nRaw: %s",
ruleID, msg, hostname, uri, truncateDaemon(line, 300))
return []alert.Finding{{
Severity: severity,
Check: check,
Message: message,
Details: details,
SourceIP: ip,
Domain: domainOrEmpty(hostname),
}}
}
// domainOrEmpty returns hostname unless it parses as a bare IP address
// (v4 or v6, with or without surrounding brackets). Vhosts served on a
// raw IP would otherwise key the incident bucket on the IP literal,
// causing two unrelated victim sites that happen to be reachable over
// their public IPs to merge into a single bucket.
func domainOrEmpty(hostname string) string {
if hostname == "" {
return ""
}
probe := strings.TrimPrefix(hostname, "[")
probe = strings.TrimSuffix(probe, "]")
if net.ParseIP(probe) != nil {
return ""
}
return hostname
}
// classifyLiteSpeedTrigger decides whether a LiteSpeed mod_security
// "triggered!" line represents a real deny (block_realtime) or merely an
// informational pass-action match (warning_realtime), based on the rule's
// declared action in the rule-action registry. A populated registry with an
// unknown rule defaults to block; a nil or empty registry cannot distinguish
// pass-action matches from denies, so ambiguous lines stay warnings until a
// refresh loads rule actions.
func classifyLiteSpeedTrigger(ruleID string) string {
num, err := strconv.Atoi(ruleID)
if err != nil {
return "modsec_block_realtime"
}
reg := modsec.Global()
// No rule-action knowledge available: the registry has not been built yet,
// or the vendor rule tree was transiently empty (cPanel modsec_assemble
// mid-rewrite, or a boot-time web-server mis-detection). We cannot tell a
// pass-action scoring rule -- e.g. Comodo CWAF 210710 / 214930, which only
// add anomaly points and never deny -- from a real deny. Defaulting every
// "triggered" line to a block in this state false-escalates benign hits
// into 24h auto-bans of real visitors. Explicit "Access denied" lines are
// classified as blocks before this function is reached. LiteSpeed deny
// rules that log only "triggered!" are degraded to warning during this
// empty-registry window, but ambiguous lines must not auto-escalate while
// the registry is unavailable.
if reg == nil || reg.Len() == 0 {
return "modsec_warning_realtime"
}
action, known := reg.Action(num)
if !known {
// Registry IS populated but this specific rule is unrecognised -- stay
// conservative and treat it as a block so a genuinely unknown deny
// rule still escalates.
return "modsec_block_realtime"
}
if modsec.IsBlockingAction(action) {
return "modsec_block_realtime"
}
return "modsec_warning_realtime"
}
// extractModSecField extracts the value between start and end delimiters.
// Returns empty string if delimiters are not found.
func extractModSecField(line, start, end string) string {
idx := strings.Index(line, start)
if idx < 0 {
return ""
}
rest := line[idx+len(start):]
endIdx := strings.Index(rest, end)
if endIdx < 0 {
return ""
}
return rest[:endIdx]
}
// extractModSecRuleFile returns the base name of the rule file that matched.
// Apache writes it as [file "..."]; LiteSpeed as "at [path:line]".
func extractModSecRuleFile(line string) string {
path := extractModSecField(line, `[file "`, `"]`)
if path == "" {
path = extractModSecField(line, "] at [", "]")
if idx := strings.LastIndex(path, ":"); idx > 0 {
path = path[:idx]
}
}
if path == "" {
return ""
}
return filepath.Base(path)
}
// extractAllModSecFields returns every value delimited by start/end, joined by
// a space. ModSecurity lines carry repeated [tag "..."] fields; the confidence
// classifier needs all of them, not just the first.
func extractAllModSecFields(line, start, end string) string {
var vals []string
for {
idx := strings.Index(line, start)
if idx < 0 {
break
}
rest := line[idx+len(start):]
endIdx := strings.Index(rest, end)
if endIdx < 0 {
break
}
vals = append(vals, rest[:endIdx])
line = rest[endIdx+len(end):]
}
return strings.Join(vals, " ")
}
// extractLiteSpeedIP extracts the client IP from a LiteSpeed log line.
// Format: [IP:PORT-CONN#VHOST] e.g. [122.9.114.57:41920-13#APVH_*_server.example.com]
func extractLiteSpeedIP(line string) string {
// Find the field that looks like [IP:PORT-CONN#VHOST]
// It appears as a bracketed field containing # and a port separator.
start := 0
for {
openBracket := strings.Index(line[start:], "[")
if openBracket < 0 {
return ""
}
openBracket += start
closeBracket := strings.Index(line[openBracket:], "]")
if closeBracket < 0 {
return ""
}
closeBracket += openBracket
field := line[openBracket+1 : closeBracket]
// LiteSpeed connection field has # (for VHOST) and contains IP:PORT-CONN#
if strings.Contains(field, "#") && strings.Contains(field, ":") && strings.Contains(field, "-") {
// Extract IP part: everything before the first ':'
colonIdx := strings.Index(field, ":")
if colonIdx > 0 {
ip := field[:colonIdx]
// Validate it looks like an IP (has dots).
if strings.Count(ip, ".") == 3 {
return ip
}
}
}
start = closeBracket + 1
}
}
// parseModSecLogLineDeduped wraps parseModSecLogLine with dedup and block
// threshold escalation. It is the handler registered with the log watcher.
//
// Order of operations (critical for correctness):
// 1. Parse the raw line.
// 2. ALWAYS increment the block escalation counter (even if dedup will suppress).
// 3. Then check dedup - suppress the base finding if a duplicate, but still
// return any escalation finding from step 2.
func parseModSecLogLineDeduped(line string, cfg *config.Config) []alert.Finding {
raw := parseModSecLogLine(line, cfg)
if len(raw) == 0 {
return nil
}
f := raw[0]
now := time.Now()
var results []alert.Finding
// --- Step 1: block escalation (before dedup) ---
// Extract IP and rule ID directly from the raw log line - NOT from the
// finding message, which could be manipulated via log injection.
ip := extractModSecField(line, "[client ", "]")
if ip == "" {
ip = extractLiteSpeedIP(line)
}
// Strip port from Apache 2.4 format (IP:port)
if strings.Count(ip, ":") == 1 {
if idx := strings.LastIndex(ip, ":"); idx > 0 {
ip = ip[:idx]
}
}
ruleID := extractModSecField(line, `[id "`, `"]`)
ruleNum, _ := strconv.Atoi(ruleID)
isBlock := f.Check == "modsec_block_realtime"
isCSM := isBlock && ruleNum >= 900000 && ruleNum <= 900999
// Record hit for per-rule stats (24h hourly buckets)
if ruleNum >= 900000 && ruleNum <= 900999 {
if sdb := store.Global(); sdb != nil {
sdb.IncrModSecRuleHit(ruleNum, now)
}
}
// Classify the rule's attack confidence from its ID, message, tags, and
// severity. This decides whether a deny may auto-escalate to a firewall ban
// (high/unknown) or only feeds the low-confidence visibility/backstop path
// (low). See docs/superpowers/specs/2026-06-27-modsec-escalation-fp-options.md.
msg := extractModSecField(line, `[msg "`, `"]`)
tags := extractAllModSecFields(line, `[tag "`, `"]`)
conf := classifyModSecConfidence(ruleNum, msg, tags, extractModSecRuleFile(line))
// Operator override (Rules page): exclude a rule ID from escalation. Coarse
// and dual-use-unsafe on its own; the classifier is the primary control.
noEscalate := false
if db := store.Global(); db != nil {
noEscalate = db.GetModSecNoEscalateRules()[ruleNum]
}
// Feed the per-IP window with every blocking deny, plus high/unknown
// warnings, which supply attack evidence to a later anomaly-threshold deny
// in the same burst. LiteSpeed often omits msg/tag on pass-action attack
// rules; unknown must fail secure instead of disappearing before the low
// anomaly rule denies. Operator-excluded rules are not recorded at all.
if ip != "" && ruleID != "" && !noEscalate && (isBlock || conf == modsecConfHigh || conf == modsecConfUnknown) {
hits, win := modsecEscalationParams(cfg)
lowConfHits := modsecLowConfEscalationHits(cfg)
rearmAfter := modsecEscalationRearmInterval(cfg)
outcome := recordModSecEventWithRearm(ip, now, ruleNum, conf, isBlock, hits, lowConfHits, win, rearmAfter)
switch {
case outcome.escalate:
check := "modsec_block_escalation"
label := "ModSecurity"
if isCSM {
check = "modsec_csm_block_escalation"
label = "CSM rule"
}
results = append(results, alert.Finding{
Severity: alert.Critical,
Check: check,
Message: fmt.Sprintf("%s escalation: %d+ denies from %s within %v", label, hits, ip, win),
Details: truncateDaemon(line, 400),
SourceIP: ip,
Domain: f.Domain,
})
case outcome.lowConfBurst:
results = append(results, alert.Finding{
Severity: alert.Warning,
Check: "modsec_low_confidence_burst",
Message: fmt.Sprintf("ModSecurity low-confidence burst: %d+ policy/anomaly denies from %s within %v (no attack signature; not auto-blocked)",
hits, ip, win),
Details: truncateDaemon(line, 400),
SourceIP: ip,
Domain: f.Domain,
})
}
if outcome.classifierGap {
results = append(results, alert.Finding{
Severity: alert.Warning,
Check: "modsec_classifier_gap",
Message: fmt.Sprintf("ModSecurity rule %s from %s is unclassified (escalation-eligible; add it to the confidence table if it is a known policy/anomaly rule)",
ruleID, ip),
Details: truncateDaemon(line, 400),
SourceIP: ip,
Domain: f.Domain,
})
}
}
// --- Step 2: Dedup ---
dedupKey := ip + ":" + ruleID
if prev, loaded := modsecDedup.Load(dedupKey); loaded {
if now.Sub(prev.(time.Time)) < modsecDedupTTL {
// Suppress the base finding but still return any escalation.
if len(results) > 0 {
return results
}
return nil
}
}
modsecDedup.Store(dedupKey, now)
results = append(results, f)
return results
}
// recordModSecEvent records one ModSecurity event for an IP's sliding window
// and decides what it should produce. Confidence-gated:
//
// - escalate (Critical auto-block) fires when total blocking denies reach
// hits AND the window has high-confidence or unknown evidence; OR when
// low-confidence-only denies reach the lowConfHits backstop.
// - lowConfBurst (non-actioned visibility) fires when the hit count is reached
// with only low-confidence evidence and the backstop is not yet met.
// - classifierGap (non-actioned visibility) fires once per unknown rule ID
// host-wide per modsecClassifierGapInterval so new vendor packs are noticed
// instead of silently taking the low-confidence path.
//
// Repeating one high-confidence rule still escalates: diversity is never
// required when high-confidence evidence is present. hits/lowConfHits/window are
// operator knobs; callers pull defaults via modsecEscalationParams.
func recordModSecEvent(ip string, now time.Time, rule int, conf modsecConfidence, isBlock bool, hits, lowConfHits int, window time.Duration) modsecEscalationOutcome {
return recordModSecEventWithRearm(ip, now, rule, conf, isBlock, hits, lowConfHits, window, modsecDefaultEscalationRearmInterval)
}
func recordModSecEventWithRearm(ip string, now time.Time, rule int, conf modsecConfidence, isBlock bool, hits, lowConfHits int, window, rearmAfter time.Duration) modsecEscalationOutcome {
val, _ := modsecBlockCount.LoadOrStore(ip, &modsecIPCounter{})
ctr := val.(*modsecIPCounter)
ctr.mu.Lock()
defer ctr.mu.Unlock()
// Prune entries older than the escalation window. Re-arm latches from the
// pre-append window; otherwise a fresh event that takes the IP from below the
// trigger back to the trigger would keep the stale latch and never re-emit.
cutoff := now.Add(-window)
kept := ctr.events[:0]
for _, e := range ctr.events {
if !e.t.Before(cutoff) {
kept = append(kept, e)
}
}
ctr.events = kept
if len(ctr.events) == 0 {
ctr.escalated = false
ctr.lastEscalated = time.Time{}
ctr.lowBurstEmitted = false
} else {
before := summarizeModSecEvents(ctr.events)
normalBefore, backstopBefore := modsecTriggerStates(before, hits, lowConfHits)
if !normalBefore && !backstopBefore {
ctr.escalated = false
ctr.lastEscalated = time.Time{}
}
}
ctr.events = append(ctr.events, modsecDenyEvent{t: now, rule: rule, conf: conf, isBlock: isBlock})
// Aggregate the current window.
summary := summarizeModSecEvents(ctr.events)
var out modsecEscalationOutcome
if conf == modsecConfUnknown {
out.classifierGap = claimModSecClassifierGap(rule, now)
}
normalFire, backstopFire := modsecTriggerStates(summary, hits, lowConfHits)
if ctr.escalated && !modsecEscalationRearmDue(ctr.lastEscalated, now, rearmAfter) {
return out
}
if normalFire || backstopFire {
ctr.escalated = true
ctr.lastEscalated = now
out.escalate = true
return out
}
// Low-confidence-only burst at the normal bar: visibility, never a ban.
if summary.totalBlock >= hits && !summary.highEvidence && !summary.unknownEvidence && !ctr.lowBurstEmitted {
ctr.lowBurstEmitted = true
out.lowConfBurst = true
}
return out
}
func summarizeModSecEvents(events []modsecDenyEvent) modsecWindowSummary {
var summary modsecWindowSummary
for _, e := range events {
if e.conf == modsecConfHigh {
summary.highEvidence = true
}
if e.conf == modsecConfUnknown {
summary.unknownEvidence = true
}
if !e.isBlock {
continue
}
summary.totalBlock++
if e.conf == modsecConfLow {
summary.lowConfBlock++
}
}
return summary
}
func modsecTriggerStates(summary modsecWindowSummary, hits, lowConfHits int) (normalFire, backstopFire bool) {
normalFire = summary.totalBlock >= hits && (summary.highEvidence || summary.unknownEvidence)
backstopFire = !summary.highEvidence && !summary.unknownEvidence && lowConfHits > 0 && summary.lowConfBlock >= lowConfHits
return normalFire, backstopFire
}
func modsecEscalationRearmDue(lastEscalated, now time.Time, rearmAfter time.Duration) bool {
if lastEscalated.IsZero() || rearmAfter <= 0 {
return false
}
return !now.Before(lastEscalated.Add(rearmAfter))
}
// StartModSecEviction starts a background goroutine that prunes expired dedup
// and counter entries every modsecEvictInterval to prevent unbounded memory
// growth. It returns when stopCh is closed. cfgFn supplies the live
// thresholds at each tick so SIGHUP edits to the escalation window take
// effect without restarting the evictor.
func StartModSecEviction(stopCh <-chan struct{}, cfgFn func() *config.Config) {
if cfgFn == nil {
cfgFn = func() *config.Config { return nil }
}
obs.Go("modsec-eviction", func() {
ticker := time.NewTicker(modsecEvictInterval)
defer ticker.Stop()
for {
select {
case <-stopCh:
return
case now := <-ticker.C:
cfg := cfgFn()
hits, win := modsecEscalationParams(cfg)
lowConfHits := modsecLowConfEscalationHits(cfg)
evictModSecStateWithLowConf(now, hits, lowConfHits, win)
}
}
})
}
// discoverModSecLogPath returns the path to the web server error log that
// CSM should tail for ModSecurity denies. Config override wins, then the
// first candidate from platform detection that actually exists.
func discoverModSecLogPath(cfg *config.Config) string {
if cfg.ModSecErrorLog != "" {
return cfg.ModSecErrorLog
}
return firstExistingPath(platform.Detect().ErrorLogPaths)
}
// firstExistingPath returns the first path in the list that exists on disk,
// or "" if none do. Pure function so tests can exercise it directly.
func firstExistingPath(candidates []string) string {
for _, p := range candidates {
if _, err := os.Stat(p); err == nil {
return p
}
}
return ""
}
func evictModSecState(now time.Time, hits int, window time.Duration) {
evictModSecStateWithLowConf(now, hits, modsecDefaultLowConfEscalationHits, window)
}
// evictModSecStateWithLowConf prunes expired entries from modsecDedup and
// modsecBlockCount. hits, lowConfHits, and window mirror the live thresholds so
// latch reset matches what recordModSecEvent would compute on the next event.
func evictModSecStateWithLowConf(now time.Time, hits, lowConfHits int, window time.Duration) {
// Prune dedup entries older than modsecDedupTTL.
modsecDedup.Range(func(key, value any) bool {
if now.Sub(value.(time.Time)) >= modsecDedupTTL {
modsecDedup.Delete(key)
}
return true
})
// Prune counter entries.
cutoff := now.Add(-window)
modsecBlockCount.Range(func(key, value any) bool {
ctr := value.(*modsecIPCounter)
ctr.mu.Lock()
kept := ctr.events[:0]
for _, e := range ctr.events {
if !e.t.Before(cutoff) {
kept = append(kept, e)
}
}
ctr.events = kept
empty := len(kept) == 0
summary := summarizeModSecEvents(kept)
normalFire, backstopFire := modsecTriggerStates(summary, hits, lowConfHits)
if !normalFire && !backstopFire {
ctr.escalated = false
ctr.lastEscalated = time.Time{}
}
if empty {
ctr.lastEscalated = time.Time{}
ctr.lowBurstEmitted = false
}
ctr.mu.Unlock()
if empty {
modsecBlockCount.Delete(key)
}
return true
})
modsecGapMu.Lock()
for rule, reported := range modsecGapReported {
if now.Sub(reported) >= modsecClassifierGapInterval {
delete(modsecGapReported, rule)
}
}
modsecGapMu.Unlock()
}
// claimModSecClassifierGap reports whether this hit of an unknown rule should
// raise the classifier-gap finding: the first hit host-wide, then again once
// per modsecClassifierGapInterval while the rule stays unclassified.
func claimModSecClassifierGap(rule int, now time.Time) bool {
modsecGapMu.Lock()
defer modsecGapMu.Unlock()
if reported, ok := modsecGapReported[rule]; ok && now.Sub(reported) < modsecClassifierGapInterval {
return false
}
modsecGapReported[rule] = now
return true
}
package daemon
import (
"context"
"maps"
"strings"
"time"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall/rollback"
"github.com/pidginhost/csm/internal/health"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/processhandle"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/updatecheck"
)
// QueueStatuses reports protection work independently of alert delivery.
func (d *Daemon) QueueStatuses() map[string]queuehealth.Status {
return d.queueStatuses(time.Now())
}
func (d *Daemon) queueStatuses(now time.Time) map[string]queuehealth.Status {
out := d.registeredQueueStatuses(now)
for name, status := range checks.AutoBlockQueueStatuses(now) {
out["auto_block."+name] = status
}
if incidentCorrelator != nil {
for name, status := range incidentCorrelator.QueueStatuses(now) {
out["incident."+name] = status
}
}
out["actionlog.writes"] = actionlog.QueueStatus(now)
out["phpanel.spool"] = alert.PhpanelQueueStatus(now)
out["checks.executions"] = checks.CheckExecutionQueueStatus(now)
out["checks.plugin_inventory"] = checks.PluginInventoryQueueStatus(now)
out["checks.wordpress_core"] = checks.WPCoreQueueStatus(now)
out["checks.reputation_queries"] = checks.ReputationQueueStatus(now)
for name, status := range checks.FileIndexQueueStatuses(now) {
out["checks.file_index."+name] = status
}
out["checks.dispatch"] = checks.CheckDispatchQueueStatus(now)
for name, status := range rdnsCache().QueueStatuses(now) {
out["smtp_rdns."+name] = status
}
for name, status := range checks.EmailPasswordQueueStatuses(now) {
out["email_password."+name] = status
}
if enr := processCtxPublished.Load(); enr != nil {
for name, state := range enr.QueueStatuses(now) {
out["processctx."+name] = state
}
}
if d.alertQueue != nil {
out["findings.ingest"] = d.alertQueue.Snapshot(now)
}
if fm := d.getFileMonitor(); fm != nil {
maps.Copy(out, fm.queueStatuses(now))
}
if sw := d.getSpoolWatcher(); sw != nil {
maps.Copy(out, sw.queueStatuses(now))
}
return out
}
// Hostname implements health.Provider.
func (d *Daemon) Hostname() string {
cfg := d.currentCfg()
if cfg == nil {
return ""
}
return cfg.Hostname
}
// StartedAt implements health.Provider.
func (d *Daemon) StartedAt() time.Time {
return d.startTime
}
// LatestScan implements health.Provider.
func (d *Daemon) LatestScan() time.Time {
if d.store == nil {
return time.Time{}
}
return d.store.LatestScanTime()
}
// BaselineAt implements health.Provider. Returns the persisted first-start
// timestamp recorded by EnsureBaseline on the daemon's first successful
// boot against this state directory. Reinstalls and upgrades preserve
// the original value.
func (d *Daemon) BaselineAt() time.Time {
if d.store == nil {
return time.Time{}
}
return d.store.BaselineAt()
}
// StoreHealthy implements health.Provider.
func (d *Daemon) StoreHealthy() bool {
s := store.Global()
if s == nil {
return false
}
return s.IsHealthy()
}
// StoreSizeMB implements health.Provider.
func (d *Daemon) StoreSizeMB() float64 {
s := store.Global()
if s == nil {
return 0
}
return float64(s.SizeBytes()) / (1024 * 1024)
}
// SeverityCounts implements health.Provider.
// Buckets the latest findings by severity name.
func (d *Daemon) SeverityCounts() map[string]int {
out := map[string]int{"critical": 0, "high": 0, "warning": 0}
if d.store == nil {
return out
}
for _, f := range d.store.LatestFindings() {
switch f.Severity {
case alert.Critical:
out["critical"]++
case alert.High:
out["high"]++
default:
out["warning"]++
}
}
return out
}
// BlocklistSize implements health.Provider.
//
// The firewall engine is the authoritative source: production code never
// writes the parallel bbolt `fw:blocked` bucket the previous implementation
// read, so /api/v1/status reported a stale count (a production host showed 25
// against 909 in the real engine state). Engine.BlockedCount() reads the
// same state file Status() and `csm firewall status` use, with expired
// entries pruned.
func (d *Daemon) BlocklistSize() int {
if d.fwEngine == nil {
return 0
}
return d.fwEngine.BlockedCount()
}
// IncidentsOpen implements health.Provider. Returns the count of
// open + contained incidents in the correlator. Falls back to 0 if
// the correlator has not been constructed yet (e.g. very early
// startup or shutdown).
func (d *Daemon) IncidentsOpen() int {
if incidentCorrelator == nil {
return 0
}
return incidentCorrelator.OpenCount()
}
// BPFEnforcementActive implements health.Provider. Reports the
// configured enforcement state. Phase 4 of the BPF Incident Response
// Roadmap. Reads the live config via Daemon.currentCfg(); falls back
// to false if config is nil (very early startup).
func (d *Daemon) BPFEnforcementActive() bool {
cfg := d.currentCfg()
if cfg == nil {
return false
}
return cfg.BPFEnforcement.Enabled &&
cfg.BPFEnforcement.DirectSMTPEgress &&
bpf.ActiveKind("connection_tracker") == bpf.BackendBPF
}
// HistoryCount implements health.Provider.
func (d *Daemon) HistoryCount() int {
s := store.Global()
if s == nil {
return 0
}
return s.HistoryCount()
}
// ConfigHash implements health.Provider.
func (d *Daemon) ConfigHash() string {
cfg := d.currentCfg()
if cfg == nil {
return ""
}
return cfg.Integrity.ConfigHash
}
// BinaryHash implements health.Provider. The binary is read again only
// when the file changes; returns empty on error.
func (d *Daemon) BinaryHash() string {
if d.binaryPath == "" {
return ""
}
return d.binaryHash.get(d.binaryPath)
}
// CorrelationAttribution implements health.Provider. Nil until the first
// active-set merge, so a fresh daemon does not claim a clean state it has
// not yet observed.
func (d *Daemon) CorrelationAttribution() *health.CorrelationAttribution {
h := checks.AttributionHealth()
if h.ActiveSetUpdates == 0 {
return nil
}
return &health.CorrelationAttribution{
Current: h.Current,
Cumulative: h.Cumulative,
ActiveSetUpdates: h.ActiveSetUpdates,
Since: h.Since,
}
}
// DryRunBlocksCount implements health.Provider.
// Returns the number of firewall blocks that were intercepted by dry_run.
func (d *Daemon) DryRunBlocksCount() int {
s := store.Global()
if s == nil {
return 0
}
return s.DryRunBlocksCount()
}
// AutomationStatus implements health.Provider.
func (d *Daemon) AutomationStatus() health.AutomationStatus {
cfg := d.currentCfg()
out := health.AutomationStatus{
DryRunBlocks: d.DryRunBlocksCount(),
LastAction: d.lastAutomationAction(),
}
if cfg != nil {
out.AutoResponseEnabled = cfg.AutoResponse.Enabled
out.AutoResponseBlockIPs = cfg.AutoResponse.BlockIPs
out.AutoResponseDryRun = cfg.AutoResponseDryRunEnabled()
// Configured termination that the kernel cannot perform safely is
// inoperative, not merely unused: report the capability either way.
out.ProcessKillEnabled = cfg.AutoResponse.Enabled && cfg.AutoResponse.KillProcesses
if err := processhandle.Available(); err != nil {
out.ProcessSignalError = err.Error()
} else {
out.ProcessSignalSupported = true
}
out.ChallengeEnabled = cfg.Challenge.Enabled
out.ChallengePortGateEnabled = cfg.Challenge.PortGate.Enabled
}
if d.ipList != nil {
out.ChallengePending = d.ipList.Count()
}
out.ChallengeEscalated = challengeEscalatedCount()
out.ChallengePortGateActive = d.challengeGate != nil
if cfg != nil && cfg.Firewall != nil {
out.FirewallEnabled = cfg.Firewall.Enabled
if out.FirewallEnabled {
out.FirewallStartupError = d.fwStartupError
}
}
// FirewallManaged is true only when a live engine is wired. Reporting it
// (alongside FirewallEnabled) lets monitoring detect "enabled but not
// managed" -- the engine-failed-to-apply condition that previously left the
// firewall silently unmanaged.
if engine := d.fwEngine; engine != nil {
out.FirewallManaged = true
rc := engine.RuleCounts()
out.FirewallBlockedIPs = rc.Blocked
out.FirewallBlockedSubnets = rc.Subnets
}
if mgr := rollback.Global(); mgr != nil {
st := mgr.Status()
out.FirewallRollbackPending = st.Pending
out.FirewallRollbackSecondsRemain = st.SecondsRemaining
}
return out
}
func (d *Daemon) lastAutomationAction() *health.AutomationAction {
if d.store == nil {
return nil
}
d.automationActionMu.Lock()
defer d.automationActionMu.Unlock()
if !d.automationActionCached.IsZero() && time.Since(d.automationActionCached) < lastAutomationActionTTL {
return d.automationActionCache
}
d.automationActionCache = d.computeLastAutomationAction()
d.automationActionCached = time.Now()
return d.automationActionCache
}
func (d *Daemon) computeLastAutomationAction() *health.AutomationAction {
var (
best alert.Finding
ok bool
)
consider := func(findings []alert.Finding) {
for _, f := range findings {
if !isAutomationActionCheck(f.Check) {
continue
}
if !ok || f.Timestamp.After(best.Timestamp) {
best = f
ok = true
}
}
}
consider(d.store.LatestFindings())
if history, _ := d.store.ReadHistory(100, 0); len(history) > 0 {
consider(history)
}
if !ok {
return nil
}
return &health.AutomationAction{
Check: best.Check,
Message: best.Message,
Timestamp: best.Timestamp,
}
}
func isAutomationActionCheck(check string) bool {
switch check {
case "auto_block", "auto_response", "challenge_route":
return true
}
return strings.HasPrefix(check, "email_php_relay_action_")
}
// Mode implements health.Provider. It reports the operator's posture so
// status, the API and doctor all show whether this host is allowed to change
// its own state.
func (d *Daemon) Mode() string {
cfg := config.Active()
if cfg == nil {
cfg = d.cfg
}
if cfg == nil {
return config.ModeEnforce
}
if cfg.ObserveMode() {
return config.ModeObserve
}
return config.ModeEnforce
}
// UpdateInfo implements health.Provider. Returns the latest cached
// release-check result, or zero value if the checker was disabled
// or has not completed its first poll yet.
func (d *Daemon) UpdateInfo() health.UpdateInfo {
if d.updateChecker == nil {
return health.UpdateInfo{}
}
info := d.updateChecker.Latest()
return health.UpdateInfo{
LatestVersion: info.LatestVersion,
Available: info.Available,
Source: info.Source,
CheckedAt: info.CheckedAt,
Err: info.Err,
}
}
// startUpdateChecker wires the updatecheck.Checker. No-op when
// updates.check_enabled is false (operator opt-out, e.g. air-gapped
// deployments) or when the running binary is "dev" (still useful but
// the banner will always show).
func (d *Daemon) startUpdateChecker() {
cfg := d.cfg
if cfg == nil || !cfg.UpdatesCheckEnabled() {
return
}
interval := cfg.UpdatesInterval()
pkgProbe := selectPackageProbe(cfg.UpdatesPackageName())
d.updateChecker = updatecheck.New(updatecheck.Options{
CurrentVersion: d.version,
Interval: interval,
GitHubAPIURL: cfg.Updates.GitHubAPIURL,
PackageProbe: pkgProbe,
LogErr: func(source string, err error) {
csmlog.Debug("update check probe failed", "source", source, "err", err)
},
})
d.wg.Add(1)
obs.Go("update-checker", func() {
defer d.wg.Done()
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
go func() { <-d.stopCh; cancel() }()
d.updateChecker.Run(ctx)
})
}
// selectPackageProbe returns an apt or dnf probe based on the
// detected OS family, or nil when the host runs neither (binary
// installs, source builds, etc.).
func selectPackageProbe(packageName string) updatecheck.PackageProbe {
info := platform.Detect()
switch {
case info.IsDebianFamily():
return updatecheck.AptProbe(packageName)
case info.IsRHELFamily():
return updatecheck.DnfProbe(packageName)
default:
return nil
}
}
// compile-time check: Daemon satisfies health.Provider.
var _ health.Provider = (*Daemon)(nil)
package daemon
import (
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/store"
)
// WordPressVerification supplies the same persisted coverage to the CLI and API.
func (d *Daemon) WordPressVerification() map[string]health.WPVerificationCounts {
db := store.Global()
if db == nil {
return nil
}
result := make(map[string]health.WPVerificationCounts)
for _, kind := range []string{"core", "plugins"} {
rows, err := db.WPVerification(kind)
if err != nil {
result[kind] = health.WPVerificationCounts{Error: "verification history unavailable"}
continue
}
if len(rows) == 0 {
continue
}
var counts health.WPVerificationCounts
for _, row := range rows {
switch row.State {
case "verified":
counts.Verified++
case "modified":
counts.Modified++
case "unverified":
counts.Unverified++
case "not_wordpress":
counts.NotWordPress++
default:
counts.Unknown++
}
if row.AttemptAt.After(counts.LastAttempt) {
counts.LastAttempt = row.AttemptAt
}
}
result[kind] = counts
}
return result
}
package daemon
import (
"fmt"
"os"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/store"
)
// Reload outcome labels for csm_config_reloads_total. Keep these in
// sync with docs/src/metrics.md.
const (
reloadResultSuccess = "success"
reloadResultError = "error"
reloadResultRestartRequired = "restart_required"
reloadResultNoop = "noop"
)
var (
reloadMetric *metrics.CounterVec
reloadMetricOnce sync.Once
)
// recordReloadResult bumps csm_config_reloads_total by one for the
// given outcome label, registering the metric on first use.
func recordReloadResult(result string) {
reloadMetricOnce.Do(func() {
reloadMetric = metrics.NewCounterVec(
"csm_config_reloads_total",
"SIGHUP config reload attempts, by outcome. result=success when safe fields were swapped in place; restart_required when the edit touched a field that needs a full restart (live config unchanged); error on YAML parse, validation, or re-sign failure (live config unchanged); noop when the edit was semantically identical to the running config.",
[]string{"result"},
)
metrics.MustRegister("csm_config_reloads_total", reloadMetric)
})
reloadMetric.With(result).Inc()
}
// reloadConfig re-reads the on-disk csm.yaml plus configured drop-ins,
// validates it, diffs against the current live config, and - if every
// change is marked safe for live reload - installs the new config via
// config.SetActive. Only the main config file is re-signed; drop-ins are
// never written back into csm.yaml.
//
// Failure modes all leave the live config untouched:
//
// - YAML parse error: Critical `config_reload_error` finding.
// - Validation error: Critical `config_reload_error` finding.
// - Restart-required fields changed: Warning
// `config_reload_restart_required` finding, listing the offending
// field names.
// - Re-signing failure: Critical `config_reload_error` finding.
//
// ROADMAP item 7.
func (d *Daemon) reloadConfig() {
oldCfg := d.activeOrStartupCfg()
cfgPath := oldCfg.ConfigFile
fmt.Fprintf(os.Stderr, "[%s] SIGHUP: reloading config from %s\n", ts(), cfgPath)
newCfg, err := config.LoadWithDir(cfgPath, oldCfg.ConfigDir)
if err != nil {
recordReloadResult(reloadResultError)
d.emitReloadFinding(alert.Critical, "config_reload_error",
fmt.Sprintf("SIGHUP reload: parse failed (%v); keeping old config", err))
return
}
for _, r := range config.Validate(newCfg) {
if r.Level == "error" {
recordReloadResult(reloadResultError)
d.emitReloadFinding(alert.Critical, "config_reload_error",
fmt.Sprintf("SIGHUP reload: validation error on %q: %s; keeping old config",
r.Field, r.Message))
return
}
}
changes := config.Diff(oldCfg, newCfg)
if len(changes) == 0 {
recordReloadResult(reloadResultNoop)
fmt.Fprintf(os.Stderr, "[%s] SIGHUP: no config changes detected\n", ts())
return
}
if config.RestartRequired(changes) {
var offenders []string
for _, c := range changes {
if c.Tag != config.TagSafe {
offenders = append(offenders, c.Field)
}
}
// The edit passed Validate, so the file on disk is
// loadable. Re-sign integrity.config_hash to match the
// edited content -- otherwise the next daemon restart
// (systemd, manual, crash recovery) trips the startup
// integrity check and crash-loops. Update the live cfg's
// stored ConfigHash in lock-step so the periodic
// integrity.Verify(currentCfg()) does not see a disk /
// memory divergence and fire spurious tamper alerts.
//
// Any error re-signing degrades to "live config unchanged,
// file unchanged, operator must rehash manually"; we still
// emit the warning so they know to act.
if err := d.signAndSaveReloadedConfig(oldCfg, newCfg); err == nil {
resynced := *oldCfg
resynced.Integrity = newCfg.Integrity
// confd_hash was computed under the edited exemption list, so
// the live copy must verify under that same list.
resynced.ConfD = newCfg.ConfD
resynced.ConfigFile = cfgPath
resynced.ConfigDir = oldCfg.ConfigDir
publishActiveConfig(&resynced, "SIGHUP")
} else {
fmt.Fprintf(os.Stderr, "[%s] config_reload_restart_required: re-sign failed (%v); file and live hash remain mismatched until operator runs `csm rehash`\n",
ts(), err)
}
recordReloadResult(reloadResultRestartRequired)
d.emitReloadFinding(alert.Warning, "config_reload_restart_required",
fmt.Sprintf("SIGHUP reload: restart-required fields changed: %v; live config unchanged, main config re-signed if needed for next restart",
offenders))
return
}
if err := d.signAndSaveReloadedConfig(oldCfg, newCfg); err != nil {
recordReloadResult(reloadResultError)
d.emitReloadFinding(alert.Critical, "config_reload_error",
fmt.Sprintf("SIGHUP reload: re-signing config failed: %v; live config unchanged", err))
return
}
newCfg.ConfigFile = cfgPath
newCfg.ConfigDir = oldCfg.ConfigDir
if err := installAccountExtractorFromConfig(newCfg); err != nil {
recordReloadResult(reloadResultError)
d.emitReloadFinding(alert.Critical, "config_reload_error",
fmt.Sprintf("SIGHUP reload: account extractor update failed: %v; live config unchanged", err))
return
}
phpanelEnabled := newCfg.Alerts.Webhook.Enabled && newCfg.Alerts.Webhook.Type == "phpanel"
if phpanelEnabled {
if err := alert.ConfigurePhpanelQueue(newCfg); err != nil {
recordReloadResult(reloadResultError)
d.emitReloadFinding(alert.Critical, "config_reload_error",
fmt.Sprintf("SIGHUP reload: phpanel queue update failed: %v; live config unchanged", err))
return
}
}
publishActiveConfig(newCfg, "SIGHUP")
if !phpanelEnabled {
_ = alert.ConfigurePhpanelQueue(newCfg)
}
recordReloadResult(reloadResultSuccess)
var names []string
for _, c := range changes {
names = append(names, c.Field)
}
fmt.Fprintf(os.Stderr, "[%s] SIGHUP: config reloaded; safe fields updated: %v\n", ts(), names)
// A forward-guard config change is a safe (hot-reload) field; re-reconcile
// so enabling/disabling or toggling enforce takes effect without a restart.
d.reconcileForwardGuard()
// reputation.verified_bots is a safe field too; push the new list into the
// live verifier and re-stamp the PTR cache so changes take effect at once.
d.reconcileVerifiedBots()
// thresholds is a safe block; push the new SMTP/mail brute-force thresholds
// into the live trackers so a SIGHUP takes effect without a restart.
d.reconcileBruteThresholds()
// reputation.whitelist is a safe field; replace the threat database's
// configured whitelist so additions and removals apply at once.
d.reconcileReputationWhitelist()
}
// activeOrStartupCfg returns the current live config, falling back
// to d.cfg (the startup snapshot) if SetActive has not yet been
// called. Reload paths use this so the first reload diffs against
// the startup config, and every subsequent reload diffs against
// whatever the last successful reload installed.
//
// Also used by the tier-run hot paths (via currentCfg below) so a
// SIGHUP-driven threshold change reaches the next tick without a
// restart.
func (d *Daemon) activeOrStartupCfg() *config.Config {
if c := config.Active(); c != nil {
return c
}
return d.cfg
}
// currentCfg is the per-tick config accessor for hot paths. See
// ROADMAP item 7 for the threshold-tuning motivation.
func (d *Daemon) currentCfg() *config.Config {
return d.activeOrStartupCfg()
}
func publishActiveConfig(cfg *config.Config, source string) {
config.SetActive(cfg)
purgeDryRunBlocksIfAutoResponseLive(cfg, source)
}
func purgeDryRunBlocksIfAutoResponseLive(cfg *config.Config, source string) {
if cfg == nil || cfg.AutoResponseDryRunEnabled() {
return
}
if sdb := store.Global(); sdb != nil {
if removed := sdb.PurgeAllDryRunBlocks(); removed > 0 {
fmt.Fprintf(os.Stderr, "[%s] %s: purged %d dry_run_blocks records (auto-response live)\n", ts(), source, removed)
}
}
}
// signAndSaveReloadedConfig re-computes integrity.config_hash for the main
// config file (and integrity.confd_hash for the conf.d fragments) when either
// changed across a SIGHUP. Drop-in fragment CONTENT is intentionally not
// written back into csm.yaml; only their digest is folded into confd_hash so
// a later Verify still passes. The binary hash is preserved from the prior
// live config; a SIGHUP reload cannot upgrade the binary, so it must not drift.
func (d *Daemon) signAndSaveReloadedConfig(oldCfg, newCfg *config.Config) error {
currentHash, err := integrity.HashConfigStable(oldCfg.ConfigFile)
if err != nil {
return err
}
currentConfd, err := integrity.HashConfDir(oldCfg.ConfigDir, oldCfg.ConfD.IntegrityExempt)
if err != nil {
return err
}
if currentHash == oldCfg.Integrity.ConfigHash && currentConfd == oldCfg.Integrity.ConfdHash {
newCfg.Integrity = oldCfg.Integrity
return nil
}
configHash, confdHash, err := integrity.SignConfigFilePreserving(oldCfg.ConfigFile, oldCfg.ConfigDir, oldCfg.Integrity.BinaryHash)
if err != nil {
return err
}
newCfg.Integrity.BinaryHash = oldCfg.Integrity.BinaryHash
newCfg.Integrity.ConfigHash = configHash
newCfg.Integrity.ConfdHash = confdHash
return nil
}
// emitReloadFinding logs to stderr and pushes a Finding into the
// daemon's alert channel. Non-blocking on channel saturation; the
// daemon's existing drop counter tracks those.
func (d *Daemon) emitReloadFinding(sev alert.Severity, check, msg string) {
fmt.Fprintf(os.Stderr, "[%s] %s: %s\n", ts(), check, msg)
finding := alert.Finding{
Severity: sev,
Check: check,
Message: msg,
Timestamp: time.Now(),
}
if !alert.TryEnqueue(d.alertCh, finding) {
atomic.AddInt64(&d.droppedAlerts, 1)
}
}
package daemon
import (
"regexp"
"github.com/pidginhost/csm/internal/contenttype"
)
// imagePayloadConstructs are the constructs that turn PHP bytes riding inside
// an image container into a working backdoor: a code or command sink, or a
// fetch of remote code to run. Each entry is reported verbatim as the
// finding's evidence, so the operator sees why the file was flagged.
//
// Every entry is a multi-character ASCII identifier. Compressed pixel data is
// random bytes, so a short punctuation shape turns up in ordinary plugin
// artwork by chance and cannot carry a rule: a shell-backtick arm measured
// here fired on a stock plugin layout preview whose pixel data also happened
// to hold a three-byte short-echo tag.
var imagePayloadConstructs = []struct {
name string
re *regexp.Regexp
}{
{"code execution sink", regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_>:$])(?:eval|assert|create_function|call_user_func(?:_array)?)\s*\(`)},
{"command execution sink", regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_>:$])(?:system|exec|shell_exec|passthru|proc_open|popen|pcntl_exec)\s*\(`)},
{"file inclusion sink", regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_>:$])(?:include|require)(?:_once)?\s*(?:\(\s*)?(?:@\s*)?[$'"]`)},
{"file write sink", regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_>:$])(?:file_put_contents|fwrite|fputs|move_uploaded_file)\s*\(`)},
{"remote code fetch", regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_>:$])(?:curl_init|curl_exec|curl_setopt(?:_array)?|fsockopen|stream_context_create)\s*\(`)},
{"remote code fetch", regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_>:$])(?:file_get_contents|fopen|readfile)\s*\(\s*['"]?(?:https?|ftp|php):`)},
{"request-controlled input", regexp.MustCompile(`\$_(?:GET|POST|REQUEST|COOKIE|FILES)\s*\[`)},
}
// phpExecutableContent reports a PHP opening tag accompanied by a construct
// that executes or fetches code, and names the construct.
//
// A PHP opening tag alone is not evidence. Plugin screenshots quote one in
// their description chunks, and three bytes of it turn up in compressed pixel
// data by chance, so the construct beside it is what makes the verdict.
//
// It carries no opinion about the file type, so a caller that already knows a
// path has no business holding PHP can apply it to a tail read, where the
// container magic at offset zero is out of view.
func phpExecutableContent(data []byte) (string, bool) {
if !contenttype.HasPHPOpenTag(data) {
return "", false
}
for _, construct := range imagePayloadConstructs {
if construct.re.Match(data) {
return construct.name, true
}
}
return "", false
}
package daemon
import (
"errors"
"net"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/incident"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/threatintel"
)
var (
incidentOnce sync.Once
incidentCorrelator *incident.Correlator
incidentRegistry = metrics.Default
incidentRetentionCancel func()
incidentAutoCloseCancel func()
// incidentSprayBlocker is the firewall hand-off the daemon wires in
// before IncidentCorrelator() is first invoked. The bool reports
// whether nftables was actually mutated; dry-run, already-blocked,
// and verdict-allow outcomes return false so incident timelines do
// not record a live block that never landed.
// nil means "no blocker wired" (early startup or unit tests); the
// singleton then skips wiring OnSprayBlock and the spray detector
// stays detection-only even with BlockAtSeverity set.
incidentSprayBlocker func(ip, reason string, timeout time.Duration, findingID string) (bool, error)
)
// autoResponseBlockExpiry is the operator's configured block duration, the
// first rung of the incident auto-block escalation ladder. An unset or
// unparseable value falls back to a day, matching the firewall default.
func autoResponseBlockExpiry(cfg *config.Config) time.Duration {
if cfg == nil {
return 24 * time.Hour
}
timeout, err := time.ParseDuration(cfg.AutoResponse.BlockExpiry)
if err != nil || timeout <= 0 {
return 24 * time.Hour
}
return timeout
}
// SetIncidentSprayBlocker installs the firewall-side hand-off used by the
// incident auto-block paths. Call once after the firewall engine is built
// and before the first IncidentCorrelator() call.
// Passing nil clears the binding.
func SetIncidentSprayBlocker(fn func(ip, reason string, timeout time.Duration, findingID string) (bool, error)) {
incidentSprayBlocker = fn
}
// incidentAutoCloseInterval is how often the daemon scans Open / Contained
// incidents for staleness. One hour is fast enough that 24h-idle incidents
// close within ~25h worst-case, but slow enough that the per-kind walk
// stays cheap on hosts with thousands of incidents.
const incidentAutoCloseInterval = 1 * time.Hour
// incidentAutoCloseWarmup delays the first sweep just long enough for the
// startup finding burst to settle. Incidents are restored synchronously
// before the loop starts, so a long warm-up only parks a stale backlog as
// "open" after every restart (observed on a frequently-upgraded prod host:
// thousands of >24h incidents sitting open for the full warm-up). Short.
const incidentAutoCloseWarmup = 2 * time.Minute
// incidentAutoCloseDrainDelay is the cadence used when a sweep hit its
// per-sweep cap, so a large backlog drains over a few quick passes instead
// of waiting a full interval between each capped sweep.
const incidentAutoCloseDrainDelay = 30 * time.Second
// incidentAutoCloseMaxPerSweep bounds live closes per sweep so a big
// post-restart backlog does not hold the correlator lock or burst
// thousands of bbolt persists in one tick.
const incidentAutoCloseMaxPerSweep = 1000
// incidentClosedRetention is how long resolved/dismissed incidents are kept
// before compaction prunes them. Named value per project convention; config
// exposure deferred until operators ask.
var incidentClosedRetention = incident.ClosedRetention{
Operator: 30 * 24 * time.Hour,
Auto: 7 * 24 * time.Hour,
}
// incidentOpenThreshold is the number of correlated findings required
// before a finding subject to the threshold opens an incident. Two means an
// isolated probe (a single dictionary-attack guess, one modsec hit
// from a wandering scanner) is treated as a finding only and never
// promoted to an incident on its own; the next correlated event
// inside the merge window does the promotion. Declared as var so
// tests that exercise the correlator wiring (where one finding is
// expected to land in one incident immediately) can pin it to 1
// via resetIncidentForTest. Production code never mutates this.
var incidentOpenThreshold = 2
// IncidentCorrelator returns the daemon-wide incident correlator.
// On first call: builds the correlator, restores prior state from
// the bbolt store (when available), and registers metrics. Safe for
// concurrent callers.
func IncidentCorrelator() *incident.Correlator {
incidentOnce.Do(func() {
db := store.Global()
var persist func(incident.Incident) error
if db != nil {
persist = db.SaveIncident
}
// Resolve spray-suppression knobs from the active config. nil
// config (early test wiring) leaves the detector disabled.
var spray incident.SpraySuppressionConfig
var autoBlock incident.IncidentAutoBlockConfig
var whitelisted func(string) bool
var onSprayBlock func(ip, reason string, ttl time.Duration, findingID string) bool
var onIncidentBlock func(ip, reason string, ttl time.Duration, findingID string) bool
if cfg := globalCfgForIncidents(); cfg != nil {
spray = incident.SpraySuppressionConfig{
Enabled: cfg.Incidents.SpraySuppression.Enabled,
DryRun: cfg.Incidents.SpraySuppression.DryRun,
DistinctMailboxes: cfg.Incidents.SpraySuppression.DistinctMailboxes,
BlockExpiry: autoResponseBlockExpiry(cfg),
SeverityEscalateAt: cfg.Incidents.SpraySuppression.SeverityEscalateAt,
PerCheck: cfg.IncidentsSpraySuppressionPerCheck(),
MaxTrackedIPs: cfg.Incidents.SpraySuppression.MaxTrackedIPs,
BlockAtSeverity: cfg.Incidents.SpraySuppression.BlockAtSeverity,
}
// Only wire the firewall hand-off when block-on-spray is
// configured and the daemon has a blocker installed. The live
// auto_response gate is checked at decision time so SIGHUP
// changes to enabled/block_ips take effect without rebuilding
// the singleton.
if spray.BlockAtSeverity != "" && incidentSprayBlocker != nil {
blocker := incidentSprayBlocker
onSprayBlock = func(ip, reason string, ttl time.Duration, findingID string) bool {
liveCfg := globalCfgForIncidents()
if liveCfg == nil || !liveCfg.AutoResponse.Enabled || !liveCfg.AutoResponse.BlockIPs {
return false
}
// ttl comes from the correlator's escalation ladder; zero
// is a permanent block, which the engine understands.
live, err := blocker(ip, sprayBlockReasonPrefix+reason, ttl, findingID)
if err != nil {
if live && errors.Is(err, firewall.ErrActionAuditPending) {
csmlog.Warn("credential_spray block audit delivery pending", "ip", ip, "err", err)
return true
}
if !isProtectedIPRefusal(err) {
csmlog.Warn("credential_spray block failed", "ip", ip, "err", err)
}
return false
}
return live
}
}
// Generic incident-driven auto-block. Reuses the same firewall
// blocker as the spray path; the reason prefix differs so audit
// log rows distinguish which detector triggered the block.
kindsRaw := cfg.IncidentsAutoBlockKinds()
kinds := make(map[incident.Kind]bool, len(kindsRaw))
for k := range kindsRaw {
kinds[incident.Kind(k)] = true
}
autoBlock = incident.IncidentAutoBlockConfig{
Enabled: cfg.Incidents.AutoBlock.Enabled,
BlockAtSeverity: cfg.Incidents.AutoBlock.BlockAtSeverity,
BlockExpiry: autoResponseBlockExpiry(cfg),
Kinds: kinds,
}
if autoBlock.Enabled && autoBlock.BlockAtSeverity != "" && incidentSprayBlocker != nil {
blocker := incidentSprayBlocker
onIncidentBlock = func(ip, reason string, ttl time.Duration, findingID string) bool {
liveCfg := globalCfgForIncidents()
if liveCfg == nil || !liveCfg.AutoResponse.Enabled || !liveCfg.AutoResponse.BlockIPs {
return false
}
live, err := blocker(ip, incidentReasonPrefix+reason, ttl, findingID)
if err != nil {
if live && errors.Is(err, firewall.ErrActionAuditPending) {
csmlog.Warn("incident auto-block audit delivery pending", "ip", ip, "err", err)
return true
}
// Own-interface / infra IPs are intentionally never
// blockable; the incident still opened, so the operator is
// alerted to activity attributed to a protected address
// (e.g. a compromised site pivoting through the server IP)
// without the noise of a failed-block warning.
if !isProtectedIPRefusal(err) {
csmlog.Warn("incident auto-block failed", "ip", ip, "err", err)
}
return false
}
return live
}
}
// Whitelist accessor: check the static reputation.whitelist
// list, then the bbolt-backed live whitelist operators add at
// runtime. db nil-check inside the closure so store resolution
// stays current across daemon restarts.
staticAllow := make(map[string]bool, len(cfg.Reputation.Whitelist))
for _, ip := range cfg.Reputation.Whitelist {
if ip != "" {
staticAllow[ip] = true
}
}
whitelisted = func(ip string) bool {
if ip == "" {
return false
}
if staticAllow[ip] {
return true
}
if d := store.Global(); d != nil && d.IsWhitelisted(ip) {
return true
}
// Backstop: a verified-crawler IP from a published range
// (Googlebot/Bingbot/Applebot) should never anchor a
// correlated incident. CDN edge ranges are intentionally
// excluded -- see threatintel.IPInAnyBot.
return threatintel.DefaultRanges().IPInAnyBot(net.ParseIP(ip))
}
}
incidentCorrelator = incident.NewCorrelator(incident.CorrelatorConfig{
Persist: persist,
OpenThreshold: incidentOpenThreshold,
SpraySuppression: spray,
AutoBlock: autoBlock,
IsWhitelisted: whitelisted,
CanSprayBlock: func() bool {
cfg := globalCfgForIncidents()
return cfg != nil && cfg.AutoResponse.Enabled && cfg.AutoResponse.BlockIPs
},
CanIncidentBlock: func() bool {
cfg := globalCfgForIncidents()
return cfg != nil && cfg.AutoResponse.Enabled && cfg.AutoResponse.BlockIPs
},
OnSprayBlock: onSprayBlock,
OnIncidentBlock: onIncidentBlock,
})
if db != nil {
list, err := db.ListIncidents()
if err != nil {
csmlog.Warn("incident restore failed", "err", err)
} else {
incidentCorrelator.Restore(list)
}
}
incident.RegisterMetrics(incidentRegistry(), incidentCorrelator)
incidentRetentionCancel = startIncidentRetentionLoop(incidentCorrelator)
// Auto-close runs on its own hourly ticker so the daily retention
// loop is not coupled to the close cadence; a 24h-idle incident
// closes within at most ~25h.
if cfg := globalCfgForIncidents(); cfg != nil {
incidentAutoCloseCancel = startIncidentAutoCloseLoop(incidentCorrelator, cfg)
}
})
return incidentCorrelator
}
// globalCfgForIncidents is overridden in tests to plug a synthetic config
// without touching package-level state. Production wiring sets this to a
// closure over the daemon's loaded config; until that wiring lands the
// auto-close loop simply does not start (no panic). The retention loop
// keeps running unchanged.
var globalCfgForIncidents = func() *config.Config { return nil }
// SetIncidentConfigSource wires the daemon-loaded config so the
// incident singleton can resolve auto-close thresholds at construction.
// Called once from cmd/csm/serve before IncidentCorrelator() is first
// invoked. Subsequent calls overwrite the source so reload paths can
// rebind without restarting the singleton.
func SetIncidentConfigSource(get func() *config.Config) {
if get == nil {
globalCfgForIncidents = func() *config.Config { return nil }
return
}
globalCfgForIncidents = get
}
// startIncidentAutoCloseLoop launches the per-kind idle scan that
// auto-resolves stale incidents. Returns a cancel func. Logs every run
// at info when work was done; silent when nothing closed.
func startIncidentAutoCloseLoop(c *incident.Correlator, cfg *config.Config) func() {
stop := make(chan struct{})
stopped := make(chan struct{})
go func() {
defer close(stopped)
// Single resettable timer: first sweep after a short warm-up, then
// the normal interval -- unless a sweep reports a remaining backlog,
// in which case the next sweep fires on the fast drain cadence so a
// post-restart backlog clears in minutes rather than over hours.
timer := time.NewTimer(incidentAutoCloseWarmup)
defer timer.Stop()
for {
select {
case <-stop:
return
case <-timer.C:
more := runIncidentAutoClose(c, cfg)
next := incidentAutoCloseInterval
if more {
next = incidentAutoCloseDrainDelay
}
timer.Reset(next)
}
}
}()
// The cancel waits for the goroutine to actually exit, not just signal it,
// so the daemon's shutdown sequence can guarantee no sweep is mid-bbolt-
// write when it closes the store immediately after cancelling.
return func() {
close(stop)
<-stopped
}
}
// runIncidentAutoClose is one tick of the auto-close loop. Gated on
// the operator's config and on the per-kind threshold map. dry_run=true
// only increments counters; live mode flips Status -> resolved and
// records "auto:stale" attribution. Returns more=true when the per-sweep
// cap left stale incidents unclosed, so the caller schedules a prompt
// follow-up sweep instead of waiting the full interval.
func runIncidentAutoClose(c *incident.Correlator, cfg *config.Config) (more bool) {
return runIncidentAutoCloseAt(c, cfg, time.Now())
}
// runIncidentAutoCloseAt is runIncidentAutoClose with an explicit clock.
func runIncidentAutoCloseAt(c *incident.Correlator, cfg *config.Config, now time.Time) (more bool) {
// Sub-threshold findings and spray-detector state age out after the
// merge window. Pruning them only in the daily retention compaction
// held a day of stale entries, each carrying a full Finding, on a host
// with sustained one-shot scanner traffic.
_ = c.PruneStalePending(now)
_ = c.PruneStaleSpray(now)
// The safety cap runs on every tick regardless of the operator's
// auto-close toggle or per-kind thresholds. It is a hard backstop against
// unbounded growth of Open/Contained incidents (in memory and bbolt) on a
// host under sustained attack when auto-close is off or a kind is omitted.
capMore := runIncidentSafetyCap(c)
if cfg == nil || !cfg.IncidentsAutoCloseEnabled() {
return capMore
}
rawThresholds := cfg.IncidentsAutoCloseThresholds()
if len(rawThresholds) == 0 {
return capMore
}
thresholds := make(map[incident.Kind]time.Duration, len(rawThresholds))
for k, v := range rawThresholds {
thresholds[incident.Kind(k)] = v
}
dryRun := cfg.Incidents.AutoClose.DryRun
closed, dryRunCount, scanned, more := c.CloseStaleLimited(now, thresholds, dryRun, incidentAutoCloseMaxPerSweep)
if closed > 0 || dryRunCount > 0 {
csmlog.Info("incident auto-close",
"closed", closed,
"dry_run_decisions", dryRunCount,
"scanned", scanned,
"dry_run", dryRun,
"backlog_remaining", more,
)
}
return more || capMore
}
// incidentSafetyMaxAge is the hard age cap: any Open/Contained incident idle
// longer than this is force-closed regardless of auto-close config.
const incidentSafetyMaxAge = 30 * 24 * time.Hour
// incidentSafetyMaxActive bounds how many Open/Contained incidents are held in
// memory at once; the oldest over this are force-closed.
const incidentSafetyMaxActive = 50000
// runIncidentSafetyCap force-closes incidents past the age cap and trims the
// active set back under the size ceiling. Always runs, independent of the
// operator's auto-close settings. Returns more=true if either sweep left a
// backlog so the loop schedules a prompt follow-up.
func runIncidentSafetyCap(c *incident.Correlator) (more bool) {
now := time.Now()
byAge, ageMore := c.CloseStaleByAge(now, incidentSafetyMaxAge, incidentAutoCloseMaxPerSweep)
byCap, capMore := c.EnforceActiveCap(now, incidentSafetyMaxActive, incidentAutoCloseMaxPerSweep)
if byAge > 0 || byCap > 0 {
csmlog.Info("incident safety cap",
"closed_by_age", byAge,
"closed_by_active_cap", byCap,
"backlog_remaining", ageMore || capMore,
)
}
return ageMore || capMore
}
// startIncidentRetentionLoop runs a daily compaction sweep against the
// store. Started from IncidentCorrelator() once after the singleton is
// constructed. Logs errors but never panics. Returns the cancel func.
// The first sweep fires after one hour so the daemon has time to settle
// before touching the store under retention rules.
func startIncidentRetentionLoop(c *incident.Correlator) func() {
stop := make(chan struct{})
stopped := make(chan struct{})
go func() {
defer close(stopped)
t := time.NewTicker(24 * time.Hour)
defer t.Stop()
first := time.NewTimer(time.Hour)
defer first.Stop()
for {
select {
case <-stop:
return
case <-first.C:
runIncidentCompaction(c)
case <-t.C:
runIncidentCompaction(c)
}
}
}()
// Wait for the goroutine to exit, so a compaction in flight finishes
// before the daemon closes the store.
return func() {
close(stop)
<-stopped
}
}
// runIncidentCompaction prunes resolved/dismissed incidents older than
// the retention window and bumps the compacted_total counter so the
// metric reflects actual store work. Errors are logged, never fatal.
func runIncidentCompaction(c *incident.Correlator) {
db := store.Global()
if db == nil {
return
}
runIncidentCompactionWith(c, time.Now(), db.CompactIncidents)
}
func runIncidentCompactionWith(c *incident.Correlator, now time.Time, compact func(time.Time, incident.ClosedRetention) (int, error)) {
pruned, err := compact(now, incidentClosedRetention)
c.IncrementCompactedTotal(pruned)
if err != nil {
csmlog.Warn("incident retention compaction failed", "pruned", pruned, "err", err)
return
}
_ = c.PruneClosedOlderThan(now, incidentClosedRetention)
_ = c.PruneStalePending(now)
_ = c.PruneStaleSpray(now)
if pruned > 0 {
csmlog.Info("incident retention compaction", "pruned", pruned)
}
}
// StopIncidentBackgroundLoops cancels the incident auto-close and retention
// goroutines. The daemon calls it during shutdown, before closing the store,
// so neither loop performs a bbolt write against an already-closed database.
// Safe to call when the singleton was never constructed (both cancels nil).
func StopIncidentBackgroundLoops() {
if incidentRetentionCancel != nil {
incidentRetentionCancel()
incidentRetentionCancel = nil
}
if incidentAutoCloseCancel != nil {
incidentAutoCloseCancel()
incidentAutoCloseCancel = nil
}
if incidentCorrelator != nil {
if flushed := incidentCorrelator.FlushPendingPersists(); flushed > 0 {
csmlog.Info("incident pending persists flushed", "count", flushed)
}
}
}
// resetIncidentForTest is a test seam. Stops any retention worker, zeros
// the singleton, and pins the registry to a private one so subsequent
// IncidentCorrelator() calls do not collide on metrics.Default.
func resetIncidentForTest() {
resetIncidentForTestWithThreshold(1)
}
// resetIncidentForTestWithThreshold is the same seam but lets a test pin
// the open threshold to a specific value. Used by the wiring test that
// proves the production default (2) is honored end-to-end through the
// IncidentCorrelator() singleton constructor.
func resetIncidentForTestWithThreshold(threshold int) {
if incidentRetentionCancel != nil {
incidentRetentionCancel()
incidentRetentionCancel = nil
}
if incidentAutoCloseCancel != nil {
incidentAutoCloseCancel()
incidentAutoCloseCancel = nil
}
incidentCorrelator = nil
incidentOnce = sync.Once{}
incidentRegistry = metrics.NewRegistry
globalCfgForIncidents = func() *config.Config { return nil }
incidentSprayBlocker = nil
// Most tests assert that one finding lands in one incident; the
// production threshold of 2 would defer creation to the second
// correlated event and break those wiring assertions. Pin to the
// caller-supplied value; production callers never invoke this seam.
incidentOpenThreshold = threshold
}
//go:build linux
package daemon
import (
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/queuehealth"
)
type inotifyDescriptor struct{ fanotifyDescriptor }
func (fd inotifyDescriptor) Pending() (int, error) {
return unix.IoctlGetInt(int(fd.fanotifyDescriptor), unix.TIOCINQ)
}
func newInotifyQueue(fd int) *notificationQueue {
q := newNotificationQueue(inotifyDescriptor{fanotifyDescriptor(fd)}, queuehealth.New(0, time.Minute), queuehealth.New(1, time.Minute))
q.sampled = queuehealth.NewSampled(0, "bytes", time.Minute)
q.variableRecords = true
return q
}
func (fw *ForwarderWatcher) initQueueHealth() {
fw.queueHealthOnce.Do(func() { fw.kernelQueue = newInotifyQueue(fw.inotifyFd) })
}
func (fw *ForwarderWatcher) QueueStatuses(_ time.Time) map[string]queuehealth.Status {
fw.initQueueHealth()
kernel, reader := fw.kernelQueue.snapshot(time.Now)
return map[string]queuehealth.Status{"kernel": kernel, "reader": reader}
}
func (w *spoolWatcher) initQueueHealth() {
w.queueHealthOnce.Do(func() { w.kernelQueue = newInotifyQueue(w.fd) })
}
func (w *spoolWatcher) QueueStatuses(_ time.Time) map[string]queuehealth.Status {
w.initQueueHealth()
kernel, reader := w.kernelQueue.snapshot(time.Now)
return map[string]queuehealth.Status{"kernel": kernel, "reader": reader}
}
// Set before publishing or running the replacement. Its descriptor has fresh
// occupancy, but a restart must not clear losses from this daemon's lifetime.
func (w *spoolWatcher) inheritQueueHealth(previous *spoolWatcher) {
previous.initQueueHealth()
w.initQueueHealth()
w.kernelQueue.losses = previous.kernelQueue.losses
w.kernelQueue.batches = previous.kernelQueue.batches
}
package daemon
import "github.com/pidginhost/csm/internal/config"
// liveConfigFn returns an accessor for the live daemon config: the active
// (hot-reloaded) config once one is published, otherwise the startup
// snapshot. Long-lived components that used to capture the startup pointer
// read through it so a reload actually reaches them.
func liveConfigFn(startup *config.Config) func() *config.Config {
return func() *config.Config {
if active := config.Active(); active != nil {
return active
}
return startup
}
}
package daemon
import (
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
)
// mailAccountOwner accepts an authenticated identity, never an envelope
// sender. Bare names need the same hosting-account validation as spool users.
func mailAccountOwner(account string) string {
if strings.Contains(account, "@") {
return checks.MailOwner(account)
}
return checks.HostingAccountForUser(account)
}
// Call after releasing tracker locks: the resolver may refresh host files.
func stampMailAccountOwner(findings []alert.Finding, account string) {
if len(findings) == 0 {
return
}
owner := mailAccountOwner(account)
for i := range findings {
findings[i].TenantID = owner
}
}
// Only cPanel's own permission records establish a held local account.
// Peer names, subjects and remote delivery replies are untrusted log data.
func mailPermissionLogText(line string) string {
if _, ok := parseEximTimestamp(line); !ok {
return ""
}
fields := strings.Fields(line)
if len(fields) < 3 {
return ""
}
fields = fields[2:]
if _, err := time.Parse("-0700", fields[0]); err == nil {
fields = fields[1:]
}
if len(fields) < 2 {
return ""
}
if token := fields[0]; len(token) >= 3 && token[0] == '[' && token[len(token)-1] == ']' &&
strings.Trim(token[1:len(token)-1], "0123456789") == "" {
fields = fields[1:]
}
if len(fields) < 2 {
return ""
}
if fields[0] == "Sender" || fields[0] == "Domain" {
return strings.Join(fields, " ")
}
if len(fields) < 5 || !msgIDPattern.MatchString(fields[0]) ||
(fields[1] != "==" && fields[1] != "**") || fields[3] != "R=enforce_mail_permissions" {
return ""
}
return strings.Join(fields[4:], " ")
}
package daemon
import (
"context"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/maillog"
"github.com/pidginhost/csm/internal/obs"
)
func (d *Daemon) startMailLogReader(platformDefault string, handler LogLineHandler) {
d.MarkWatcher("maillog", false)
queue := maillog.NewQueue()
d.registerQueueSource("mail", queue)
ctx, cancel := context.WithCancel(context.Background())
d.wg.Add(2)
obs.Go("maillog-stop", func() {
defer d.wg.Done()
select {
case <-d.stopCh:
case <-ctx.Done():
}
cancel()
})
obs.Go("maillog-supervisor", func() {
defer d.wg.Done()
defer cancel()
defer d.MarkWatcher("maillog", false)
maillog.Supervise(ctx, func() (maillog.Reader, error) {
return maillog.New(d.currentCfg().MailLogs, platformDefault, queue)
}, func(err error) {
if err != nil {
csmlog.Warn("mail log source unavailable; retrying", "err", err)
d.handleMailLogSourceGone(err)
} else {
d.handleMailLogSourceRestored()
}
}, func(line maillog.Line) bool {
return d.dispatchMailLogLine(line, handler)
})
})
}
package daemon
import (
"fmt"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/store"
)
// mailIPEntry tracks failed-auth timestamps, successful-login timestamps, the
// mailboxes seen on each side, and suppression state for one IP. The mailbox
// context keeps unrelated successful logins from hiding a mailbox attack.
type mailIPEntry struct {
times []time.Time
succ []time.Time
successAccounts map[string][]time.Time
failedAccounts map[string][]time.Time
// Long-horizon state for the slow-brute signal: bounded failure timestamps,
// bounded per-mailbox last-failure times, the most recent success, and the
// most recent failure not represented in the mailbox map. Any in-window
// success marks the source as a live legitimate client and disqualifies the
// slow block. An unattributed or overflowed failure prevents the bounded
// named-mailbox set from exempting activity it cannot vouch for.
slowTimes []time.Time
slowAccounts map[string]time.Time
slowLastSuccess time.Time
slowUnvouched time.Time
// goodFirst/goodLast record, per mailbox, the earliest and most recent
// successful auth from this IP within mailGoodSourceTTL. They outlive the
// short failure window so an established legitimate sender (e.g. a working
// POP3 profile) is not mistaken for an attacker when a second misconfigured
// profile produces a burst of auth failures.
goodFirst map[string]time.Time
goodLast map[string]time.Time
suppressed time.Time
// suspectedSuppressed rate-limits the advisory mail_bruteforce_suspected
// finding independently of the auto-block clock, so a misconfigured client
// is surfaced once per window without firing repeatedly, yet a later
// escalation to a real spray can still reach the auto-block path.
suspectedSuppressed time.Time
lastSeen time.Time
}
// recordGoodAuth notes a successful auth for account at now, maintaining the
// established-sender window. A success after a gap longer than the TTL starts a
// fresh relationship. Caller must hold t.mu.
func (e *mailIPEntry) recordGoodAuth(account string, now time.Time) {
if account == "" {
return
}
if e.goodFirst == nil {
e.goodFirst = make(map[string]time.Time)
e.goodLast = make(map[string]time.Time)
}
if last, ok := e.goodLast[account]; !ok || now.Sub(last) > mailGoodSourceTTL {
e.goodFirst[account] = now
}
e.goodLast[account] = now
if len(e.goodLast) > mailGoodSourceMaxAccountsPerIP {
e.evictOldestGoodSource()
}
}
// evictOldestGoodSource drops the least-recently-successful mailbox so a single
// IP's good-source records stay bounded. Caller must hold t.mu.
func (e *mailIPEntry) evictOldestGoodSource() {
var oldestAcct string
var oldest time.Time
for acct, ts := range e.goodLast {
if oldestAcct == "" || ts.Before(oldest) {
oldestAcct, oldest = acct, ts
}
}
delete(e.goodLast, oldestAcct)
delete(e.goodFirst, oldestAcct)
}
// pruneGoodSource drops good-source records whose most recent success is older
// than cutoff (now - mailGoodSourceTTL). Caller must hold t.mu.
func (e *mailIPEntry) pruneGoodSource(cutoff time.Time) {
for account, last := range e.goodLast {
if last.Before(cutoff) {
delete(e.goodLast, account)
delete(e.goodFirst, account)
}
}
}
// establishedGood reports whether this IP is an established legitimate sender
// for account: it authenticated successfully longer ago than the failure window
// (so the current failure burst is not its first contact) and recently enough
// to still be live. Caller must hold t.mu.
func (e *mailIPEntry) establishedGood(account string, now time.Time, window time.Duration) bool {
last, ok := e.goodLast[account]
if !ok || now.Sub(last) > mailGoodSourceTTL {
return false
}
return now.Sub(e.goodFirst[account]) >= window
}
// establishedOtherGoodAccounts counts mailboxes other than account for which
// this IP holds established good standing. Same-window successes are not
// established, so a stuffing run that already landed other mailboxes earns
// nothing here. Caller must hold t.mu.
func (e *mailIPEntry) establishedOtherGoodAccounts(account string, now time.Time, window time.Duration) int {
n := 0
for acct := range e.goodLast {
if acct == account {
continue
}
if e.establishedGood(acct, now, window) {
n++
}
}
return n
}
// establishedGoodConfined reports whether this IP looks like a misconfigured
// legitimate client rather than a brute-forcer: every current failure is
// attributed to a named mailbox, and every one of those failing mailboxes is a
// mailbox this IP holds ESTABLISHED good standing for (it authenticated
// successfully to that same mailbox longer ago than the failure window and
// recently enough to still be live). When true the per-IP signal is downgraded
// to an advisory rather than a firewall block, so a source whose own working
// mailbox now fails on a stale saved password (e.g. POP3 succeeds while IMAP
// fails on the same box) is not locked out.
//
// Good standing is scoped to the (IP, mailbox) identity that earned it: standing
// on mailbox A does NOT vouch for failures on a different mailbox B the IP has
// never signed into. The earlier count-only bound (failing count <= good count)
// was identity-blind and let a source with good standing on N of its own
// mailboxes brute-force N unrelated victim mailboxes indefinitely without ever
// auto-blocking. A fresh same-window "success" is not established, so padding
// earns no downgrade, and the compromise and spray detectors stay fully armed.
// Caller must hold t.mu and must have pruned
// e.times/e.failedAccounts/e.successAccounts to the window.
func (e *mailIPEntry) establishedGoodConfined(now time.Time, window time.Duration) bool {
failing := len(e.failedAccounts)
if failing == 0 {
return false
}
named := 0
for account, times := range e.failedAccounts {
named += len(times)
if !e.establishedGood(account, now, window) {
return false
}
}
// Accountless failures (no user= in the log line) could target anything; a
// good source's standing must not hide them.
return named == len(e.times)
}
func normalizeMailAuthAccount(account string) string {
local, domain, ok := strings.Cut(account, "@")
if !ok || local == "" || domain == "" {
return account
}
return local + "@" + strings.ToLower(domain)
}
// slowConfinedToEstablishedGood reports whether every mailbox the slow-window
// failures target is one this IP holds established good standing for -- the
// shared-device-with-stale-credentials shape the fast path downgrades to an
// advisory. Caller must hold t.mu and must have pruned e.slowAccounts.
func (e *mailIPEntry) slowConfinedToEstablishedGood(now time.Time, window time.Duration) bool {
if len(e.slowAccounts) == 0 || !e.slowUnvouched.IsZero() {
return false
}
for acct := range e.slowAccounts {
if !e.establishedGood(acct, now, window) {
return false
}
}
return true
}
// recentGoodCount returns how many distinct mailboxes this IP succeeded on
// recently enough to still be on record (last success within
// mailGoodSourceTTL), regardless of whether that standing has aged past the
// failure window. Unlike establishedGood it does not require the
// relationship to predate the burst, so it recognizes a legitimate sender whose
// standing is only seconds old -- a cold-started daemon before persisted
// standing re-ages, or a customer returning after the snapshot expired. Caller
// must hold t.mu.
func (e *mailIPEntry) recentGoodCount(now time.Time) int {
n := 0
for _, last := range e.goodLast {
if now.Sub(last) <= mailGoodSourceTTL {
n++
}
}
return n
}
// looksLikeFreshGoodSourceFP reports the cold-start/misconfiguration
// false-positive shape: the same confined, named, non-cracking failure set that
// establishedGoodConfined downgrades, but bounded by the IP's recent-success
// footprint instead of its established (aged) one. It is true exactly when the
// only reason the burst was not downgraded to an advisory is that the source's
// good standing is too fresh to have aged past the window. The auto-block still
// fires -- a fresh success cannot be trusted to grant a brute-force bypass, or a
// padding login would buy an attacker slack -- but the finding is annotated so
// an operator can spot a likely false positive without reconstructing the mail
// log by hand. Caller must hold t.mu and must have pruned the window state.
func (e *mailIPEntry) looksLikeFreshGoodSourceFP(now time.Time, window time.Duration) bool {
failing := len(e.failedAccounts)
if failing == 0 {
return false
}
named := 0
for account, times := range e.failedAccounts {
named += len(times)
if e.isCrackInProgress(account, now, window) {
return false
}
}
if named != len(e.times) {
return false
}
return failing <= e.recentGoodCount(now)
}
// mailFailTarget is one mailbox this IP's in-window failures hit, with the
// failure count for that mailbox.
type mailFailTarget struct {
Account string
Count int
}
// failTargets summarizes which mailboxes this IP's in-window failures targeted.
// It returns the per-mailbox counts sorted by count (desc) then account name
// (asc) for deterministic output, plus how many failures carried no mailbox
// (dovecot lines with no user=). Caller must hold t.mu and must have pruned
// e.times and e.failedAccounts to the window.
func (e *mailIPEntry) failTargets() ([]mailFailTarget, int) {
named := 0
targets := make([]mailFailTarget, 0, len(e.failedAccounts))
for account, times := range e.failedAccounts {
named += len(times)
targets = append(targets, mailFailTarget{Account: account, Count: len(times)})
}
sort.Slice(targets, func(i, j int) bool {
if targets[i].Count != targets[j].Count {
return targets[i].Count > targets[j].Count
}
return targets[i].Account < targets[j].Account
})
return targets, len(e.times) - named
}
// Unsafe account bytes are hex-escaped so a crafted user= value cannot add
// fake target separators, new alert lines, or path-looking tokens.
func formatMailFailTargetAccount(account string) string {
truncated := false
if len(account) > maxMailTargetAccountDisplayBytes {
account = account[:maxMailTargetAccountDisplayBytes]
truncated = true
}
if isPlainMailFailTargetAccount(account) {
if truncated {
return account + "..."
}
return account
}
var b strings.Builder
b.WriteString(`account="`)
for i := 0; i < len(account); i++ {
c := account[i]
if isPlainMailFailTargetAccountByte(c) {
b.WriteByte(c)
continue
}
fmt.Fprintf(&b, `\x%02x`, c)
}
if truncated {
b.WriteString("...")
}
b.WriteByte('"')
return b.String()
}
func isPlainMailFailTargetAccount(account string) bool {
if account == "" {
return false
}
for i := 0; i < len(account); i++ {
if !isPlainMailFailTargetAccountByte(account[i]) {
return false
}
}
return true
}
func isPlainMailFailTargetAccountByte(c byte) bool {
return c >= 'a' && c <= 'z' ||
c >= 'A' && c <= 'Z' ||
c >= '0' && c <= '9' ||
strings.ContainsRune("@._+-=%*", rune(c))
}
// formatMailFailTargets renders the target summary for a mail brute-force
// finding: the mailboxes hit with their failure counts (at most max named, the
// rest collapsed into "(+N more)") and a count of failures that named no
// mailbox. Returns "" when there is nothing to report.
func formatMailFailTargets(targets []mailFailTarget, accountless, max int) string {
if len(targets) == 0 && accountless <= 0 {
return ""
}
listed := targets
remainder := 0
if max > 0 && len(targets) > max {
listed = targets[:max]
remainder = len(targets) - max
}
var b strings.Builder
b.WriteString("Targets: ")
for i, tg := range listed {
if i > 0 {
b.WriteString(", ")
}
fmt.Fprintf(&b, "%s (%d)", formatMailFailTargetAccount(tg.Account), tg.Count)
}
if remainder > 0 {
fmt.Fprintf(&b, " (+%d more)", remainder)
}
if accountless > 0 {
if len(listed) > 0 {
b.WriteString("; ")
}
fmt.Fprintf(&b, "%d with no mailbox", accountless)
}
return b.String()
}
// successDominant reports whether this IP behaves like a legit busy client
// rather than a brute-forcer: it has successful logins in the window, at least
// as many successes as failures, and the successful mailboxes explain the failed
// mailbox set. Caller must hold t.mu and must have pruned the slices and account
// maps to the current window first.
func (e *mailIPEntry) successDominant() bool {
if len(e.succ) == 0 || len(e.succ) < len(e.times) || len(e.failedAccounts) == 0 {
return false
}
named := 0
for account, failures := range e.failedAccounts {
named += len(failures)
if len(e.successAccounts[account]) < len(failures) {
return false
}
}
return named == len(e.times)
}
func (e *mailIPEntry) accountSuccessDominant(account string) bool {
failures := len(e.failedAccounts[account])
return failures > 0 && len(e.successAccounts[account]) >= failures
}
// isCrackInProgress reports whether same-account successes mixed with failures
// look like a guessing breakthrough rather than a flaky legitimate client.
//
// An account the IP holds established good standing for (it has owned the
// mailbox longer than the window) is exempt: a mix of successes and failures
// there is an intermittent or stale-on-one-device client, the same mistype
// signal RecordSuccess uses to suppress compromise. Only a fresh, non-established
// in-window success that does not dominate the failures is treated as a crack:
// that is the attacker-just-guessed-it shape. Caller must hold t.mu.
func (e *mailIPEntry) isCrackInProgress(account string, now time.Time, window time.Duration) bool {
if e.establishedGood(account, now, window) {
return false
}
cutoff := now.Add(-window)
for _, ts := range e.successAccounts[account] {
if ts.After(cutoff) {
return !e.accountSuccessDominant(account)
}
}
return false
}
// goodSourceTimes is the persisted established-sender window for one mailbox:
// the earliest and most recent successful auth from an IP.
type goodSourceTimes struct {
First time.Time
Last time.Time
}
// goodSourceSnapshot maps ip -> account -> goodSourceTimes. It is persisted
// across daemon restarts so established good standing survives a restart and the
// post-restart cold-start window does not re-open the brute-force false-positive
// window (a customer's working profile would otherwise have to re-authenticate
// and re-age past the failure window before its stale-password profile stops
// being mistaken for an attacker).
type goodSourceSnapshot map[string]map[string]goodSourceTimes
// ExportGoodSource snapshots the established-sender records for persistence.
// Only good-source state is exported; short-lived failure/success window state
// is intentionally not persisted.
func (t *mailAuthTracker) ExportGoodSource() goodSourceSnapshot {
t.mu.Lock()
defer t.mu.Unlock()
snap := make(goodSourceSnapshot)
for ip, e := range t.ips {
if len(e.goodLast) == 0 {
continue
}
accts := make(map[string]goodSourceTimes, len(e.goodLast))
for acct, last := range e.goodLast {
accts[acct] = goodSourceTimes{First: e.goodFirst[acct], Last: last}
}
snap[ip] = accts
}
return snap
}
// LoadGoodSource seeds established-sender records from a persisted snapshot,
// dropping any whose most recent success is already older than the good-source
// TTL. Intended for one-time startup seeding; a record newer than the snapshot
// (a success already observed this run) keeps its Last time, while colliding
// normalized mailbox records keep the earliest First time.
func (t *mailAuthTracker) LoadGoodSource(snap goodSourceSnapshot, now time.Time) {
t.mu.Lock()
defer t.mu.Unlock()
cutoff := now.Add(-mailGoodSourceTTL)
for ip, accts := range snap {
for acct, ts := range accts {
acct = normalizeMailAuthAccount(acct)
if ip == "" || acct == "" || !validGoodSourceTimes(ts, cutoff) {
continue
}
e, ok := t.ips[ip]
if !ok {
e = &mailIPEntry{}
t.ips[ip] = e
}
if e.goodFirst == nil {
e.goodFirst = make(map[string]time.Time)
e.goodLast = make(map[string]time.Time)
}
if cur, ok := e.goodFirst[acct]; !ok || ts.First.Before(cur) {
e.goodFirst[acct] = ts.First
}
if cur, ok := e.goodLast[acct]; !ok || ts.Last.After(cur) {
e.goodLast[acct] = ts.Last
}
if e.goodLast[acct].After(e.lastSeen) {
e.lastSeen = e.goodLast[acct]
}
if len(e.goodLast) > mailGoodSourceMaxAccountsPerIP {
e.evictOldestGoodSource()
}
}
}
t.enforceMaxTracked("")
}
func validGoodSourceTimes(ts goodSourceTimes, cutoff time.Time) bool {
return !ts.First.IsZero() &&
!ts.Last.IsZero() &&
!ts.First.After(ts.Last) &&
!ts.Last.Before(cutoff)
}
// loadMailGoodSource seeds the tracker before log readers start, so first
// post-startup auth failures see the persisted standing instead of cold state.
func (d *Daemon) loadMailGoodSource() {
if d.mailAuthTracker == nil {
return
}
if sdb := store.Global(); sdb != nil {
snap, err := sdb.LoadMailGoodSource()
if err != nil {
csmlog.Warn("mail good-source load failed", "err", err)
return
}
d.mailAuthTracker.LoadGoodSource(storeToGoodSourceSnapshot(snap), d.mailAuthTracker.now())
}
}
// persistMailGoodSource writes the tracker's established good-source snapshot to
// the store so it survives a restart. No-op when the store or tracker is absent.
func (d *Daemon) persistMailGoodSource() {
if d.mailAuthTracker == nil {
return
}
if sdb := store.Global(); sdb != nil {
snap := goodSourceSnapshotToStore(d.mailAuthTracker.ExportGoodSource())
if err := sdb.SaveMailGoodSource(snap); err != nil {
csmlog.Warn("mail good-source persistence failed", "err", err)
}
}
}
func storeToGoodSourceSnapshot(in map[string]map[string]store.GoodSourcePair) goodSourceSnapshot {
out := make(goodSourceSnapshot, len(in))
for ip, accts := range in {
m := make(map[string]goodSourceTimes, len(accts))
for a, p := range accts {
m[a] = goodSourceTimes{First: p.First, Last: p.Last}
}
out[ip] = m
}
return out
}
func goodSourceSnapshotToStore(in goodSourceSnapshot) map[string]map[string]store.GoodSourcePair {
out := make(map[string]map[string]store.GoodSourcePair, len(in))
for ip, accts := range in {
m := make(map[string]store.GoodSourcePair, len(accts))
for a, ts := range accts {
m[a] = store.GoodSourcePair{First: ts.First, Last: ts.Last}
}
out[ip] = m
}
return out
}
// mailSubnetEntry tracks unique attacker IPs within a /24.
type mailSubnetEntry struct {
ips map[string]time.Time
suppressed time.Time
lastSeen time.Time
}
// mailAccountEntry tracks unique attacker IPs per mailbox, plus a separate
// suppression clock for compromise findings emitted by RecordSuccess.
type mailAccountEntry struct {
ips map[string]time.Time
suppressed time.Time
compromiseSuppressed time.Time
lastSeen time.Time
}
// mailBackendDegradedThreshold is how many auth-backend failure observations
// (dovecot unable to reach the credential backend, e.g. cPanel's cpdoveauthd)
// within the tracker window mark the mail auth subsystem as degraded. A healthy
// host produces zero of these; a backend outage produces thousands. While
// degraded, every login fails regardless of credentials, so the per-IP and
// per-subnet brute signals are suppressed to avoid auto-blocking legitimate
// users en masse.
const mailBackendDegradedThreshold = 10
// mailGoodSourceTTL is how long a successful authentication keeps an (IP,
// mailbox) pair on record as an established legitimate sender. A real owner
// authenticates successfully at least this often; once the last success ages
// past the TTL the standing is forgotten, so an IP cannot be permanently
// whitelisted by a single old success.
const mailGoodSourceTTL = 24 * time.Hour
// mailGoodSourceMaxAccountsPerIP caps how many distinct mailboxes one IP keeps
// good-source records for. A carrier-grade NAT can carry many legitimate
// mailboxes; the cap bounds memory without affecting detection (eviction is
// least-recently-successful first).
const mailGoodSourceMaxAccountsPerIP = 256
// mailEstablishedSourceAccounts is how many OTHER mailboxes an IP must hold
// established good standing on before a compromise finding for one more
// mailbox is downgraded to an advisory (office/agency device pattern).
const mailEstablishedSourceAccounts = 2
// maxMailTargetsListed caps how many mailboxes a brute-force finding names
// before collapsing the rest into a "(+N more)" suffix, so a wide spray does
// not dump dozens of names into one alert.
const maxMailTargetsListed = 5
// maxMailTargetAccountDisplayBytes bounds attacker-controlled account text in
// brute-force Details while keeping normal mailbox names intact.
const maxMailTargetAccountDisplayBytes = 128
// mailAuthTracker aggregates dovecot IMAP/POP3/ManageSieve auth events into
// four detection signals: per-IP brute force, per-/24 password spray,
// per-mailbox account spray, and per-account compromise (success after
// recent failures).
//
// Thread-safe; Record/RecordSuccess may be called concurrently from multiple
// log readers.
type mailAuthTracker struct {
mu sync.Mutex
perIPThreshold int
subnetThreshold int
accountSprayThreshold int
window time.Duration
suppression time.Duration
slowThreshold int
slowWindow time.Duration
maxTracked int
now func() time.Time
ips map[string]*mailIPEntry
subnets map[string]*mailSubnetEntry
accounts map[string]*mailAccountEntry
// Diagnostic counters (guarded by mu): cumulative Record invocations and
// findings emitted, logged periodically by the daemon to pin whether the
// non-cPanel dovecot brute-force path sees traffic and escalates.
recordCalls int64
findingsEmitted int64
// backendErr holds recent auth-backend failure timestamps. While the
// windowed count is at or above mailBackendDegradedThreshold the auth
// subsystem is treated as down and brute/subnet auto-block signals are
// suppressed. backendWarnUntil rate-limits the operator warning to one per
// suppression window.
backendErr []time.Time
backendWarnUntil time.Time
// backendDownFn, when set, reports whether the active socket probe currently
// sees the mail auth backend down. It augments the log-derived backendErr
// heuristic so suppression also triggers on the authoritative probe signal.
backendDownFn func() bool
}
// SetBackendDownCheck installs the active-probe callback the tracker consults to
// learn whether the mail auth backend is down. When it returns true, brute-force
// and subnet auto-block are suppressed. Set once at startup before log readers
// begin.
func (t *mailAuthTracker) SetBackendDownCheck(fn func() bool) {
t.mu.Lock()
defer t.mu.Unlock()
t.backendDownFn = fn
}
// newMailAuthTracker constructs a tracker. `now` is injected so tests can
// use deterministic clocks; pass `time.Now` in production.
func newMailAuthTracker(
perIPThreshold int,
subnetThreshold int,
accountSprayThreshold int,
window time.Duration,
suppression time.Duration,
slowThreshold int,
slowWindow time.Duration,
maxTracked int,
now func() time.Time,
) *mailAuthTracker {
if now == nil {
now = time.Now
}
return &mailAuthTracker{
perIPThreshold: perIPThreshold,
subnetThreshold: subnetThreshold,
accountSprayThreshold: accountSprayThreshold,
window: window,
suppression: suppression,
slowThreshold: slowThreshold,
slowWindow: slowWindow,
maxTracked: maxTracked,
now: now,
ips: make(map[string]*mailIPEntry),
subnets: make(map[string]*mailSubnetEntry),
accounts: make(map[string]*mailAccountEntry),
}
}
// Size returns the total number of tracked entities (IPs + subnets + accounts).
func (t *mailAuthTracker) Size() int {
t.mu.Lock()
defer t.mu.Unlock()
return len(t.ips) + len(t.subnets) + len(t.accounts)
}
// Record processes one dovecot IMAP/POP3/ManageSieve auth-failure observation.
// Returns zero or more findings that callers should append.
//
// ip MUST be non-private, non-loopback, and non-infra — callers enforce this
// before invoking Record.
func (t *mailAuthTracker) Record(ip, account string) []alert.Finding {
if ip == "" {
return nil
}
account = normalizeMailAuthAccount(account)
t.mu.Lock()
defer t.mu.Unlock()
t.recordCalls++
now := t.now()
cutoff := now.Add(-t.window)
var findings []alert.Finding
// During an auth-backend outage every login fails regardless of password, so
// the failure-rate signals are meaningless and would mass-block real users.
// Either the log-derived heuristic or the authoritative socket probe trips it.
degraded := t.backendDegraded(now)
if !degraded && t.backendDownFn != nil {
degraded = t.backendDownFn()
}
// --- Per-IP tracker ---
e, ok := t.ips[ip]
if !ok {
e = &mailIPEntry{}
t.ips[ip] = e
}
e.times = pruneTimes(e.times, cutoff)
pruneMailAccountTimes(e.failedAccounts, cutoff)
e.times = append(e.times, now)
e.failedAccounts = appendMailAccountTime(e.failedAccounts, account, now)
e.lastSeen = now
if t.perIPThreshold > 0 && len(e.times) >= t.perIPThreshold {
e.succ = pruneTimes(e.succ, cutoff)
pruneMailAccountTimes(e.successAccounts, cutoff)
if !e.successDominant() && !degraded {
switch {
case e.establishedGoodConfined(now, t.window):
// Established good source fat-fingering a confined set of its own
// mailboxes: surface for visibility but do not auto-block, so a
// stale saved password does not lock out a real customer. A real
// attack from the same source still blocks via the compromise and
// spray detectors. Rate-limited on its own clock so escalation to a
// wider spray can still reach the block path below.
if now.Before(e.suspectedSuppressed) {
break
}
e.suspectedSuppressed = now.Add(t.suppression)
details := "Failures are confined to a few named mailboxes from a source with established successful mail auth history; likely a stale saved password, not a brute-force. Visibility only - no auto-block. Compromise and spray detectors remain active."
targets, accountless := e.failTargets()
if targetSummary := formatMailFailTargets(targets, accountless, maxMailTargetsListed); targetSummary != "" {
details += " " + targetSummary + "."
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "mail_bruteforce_suspected",
Message: fmt.Sprintf("Suspected mail misconfiguration from %s: %d failed auths in %v from an established good source (not auto-blocked)",
ip, len(e.times), t.window),
Details: details,
Timestamp: now,
SourceIP: ip,
})
case !now.Before(e.suppressed):
e.suppressed = now.Add(t.suppression)
details := "Real-time detection of dovecot imap/pop3/managesieve auth failures."
if e.looksLikeFreshGoodSourceFP(now, t.window) {
details += " Note: this source has recent successful mail auth for one or more other mailboxes, so the failures may be a misconfigured client with a stale saved password rather than an attack. Verify before treating it as a confirmed brute-force."
}
targets, accountless := e.failTargets()
if targetSummary := formatMailFailTargets(targets, accountless, maxMailTargetsListed); targetSummary != "" {
details += " " + targetSummary + "."
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "mail_bruteforce",
Message: fmt.Sprintf("Mail auth brute force from %s: %d failed auths in %v",
ip, len(e.times), t.window),
Details: details,
Timestamp: now,
SourceIP: ip,
})
}
}
}
// --- Long-horizon slow-brute tracker ---
// Catches attackers pacing below the fast window. Distinct-mailbox and
// no-recent-success requirements separate a mailbox walk from a stale
// saved password or a busy NAT, and a source whose paced failures stay
// confined to mailboxes it holds established good standing for keeps the
// fast path's advisory treatment instead of escalating to a block.
if t.slowThreshold > 0 && t.slowWindow > 0 {
if degraded {
// Backend failures say nothing about credentials. Discard the
// long-lived evidence so outage traffic cannot trigger a delayed
// block as soon as the backend recovers.
e.slowTimes = nil
e.slowAccounts = nil
e.slowUnvouched = time.Time{}
} else {
slowCutoff := now.Add(-t.slowWindow)
e.slowTimes = pruneTimes(e.slowTimes, slowCutoff)
e.slowTimes = appendSlowFailure(e.slowTimes, now)
var trackedAccount bool
e.slowAccounts, trackedAccount = recordSlowAccount(e.slowAccounts, account, now, slowCutoff)
if !trackedAccount {
e.slowUnvouched = now
}
if e.slowLastSuccess.Before(slowCutoff) {
e.slowLastSuccess = time.Time{}
}
if e.slowUnvouched.Before(slowCutoff) {
e.slowUnvouched = time.Time{}
}
if (len(e.slowTimes) >= t.slowThreshold || len(e.slowAccounts) >= slowBruteWalkAccounts) &&
len(e.slowAccounts) >= slowBruteMinAccounts {
safeSource := !e.slowLastSuccess.IsZero() || e.slowConfinedToEstablishedGood(now, t.window)
switch {
case safeSource && !now.Before(e.suspectedSuppressed):
e.suspectedSuppressed = now.Add(t.suppression)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "mail_bruteforce_suspected",
Message: fmt.Sprintf("Suspected mail misconfiguration from %s: %d paced auth failures across %d mailboxes in %v (not auto-blocked)",
ip, len(e.slowTimes), len(e.slowAccounts), t.slowWindow),
Details: "Long-horizon failures came from a source with recent successful mail auth or established good standing on every named target. Visibility only - no auto-block; compromise and spray detectors remain active.",
Timestamp: now,
SourceIP: ip,
})
case !safeSource && !now.Before(e.suppressed):
e.suppressed = now.Add(t.suppression)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "mail_bruteforce",
Message: fmt.Sprintf("Mail auth brute force from %s: %d failed auths across %d mailboxes in %v",
ip, len(e.slowTimes), len(e.slowAccounts), t.slowWindow),
Details: "Long-horizon detection of paced imap/pop3/managesieve auth failures that stay below the fast per-IP window",
Timestamp: now,
SourceIP: ip,
})
}
}
}
}
// --- Per-/24 subnet tracker (IPv4 only) ---
if prefix := extractPrefix24Daemon(ip); prefix != "" {
s, ok := t.subnets[prefix]
if !ok {
s = &mailSubnetEntry{ips: make(map[string]time.Time)}
t.subnets[prefix] = s
}
for ipKey, ts := range s.ips {
if ts.Before(cutoff) {
delete(s.ips, ipKey)
}
}
s.ips[ip] = now
s.lastSeen = now
if t.subnetThreshold > 0 && len(s.ips) >= t.subnetThreshold && !now.Before(s.suppressed) && !degraded {
s.suppressed = now.Add(t.suppression)
cidr := prefix + ".0/24"
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "mail_subnet_spray",
Message: fmt.Sprintf("Mail password spray from %s.0/24: %d unique IPs in %v",
prefix, len(s.ips), t.window),
Details: "Real-time detection of mail auth failures from many IPs in one /24",
Timestamp: now,
SourceIP: cidr,
})
}
}
// --- Per-account spray tracker ---
if account != "" {
a, ok := t.accounts[account]
if !ok {
a = &mailAccountEntry{ips: make(map[string]time.Time)}
t.accounts[account] = a
}
for ipKey, ts := range a.ips {
if ts.Before(cutoff) {
delete(a.ips, ipKey)
}
}
a.ips[ip] = now
a.lastSeen = now
if t.accountSprayThreshold > 0 && len(a.ips) >= t.accountSprayThreshold && !now.Before(a.suppressed) {
a.suppressed = now.Add(t.suppression)
_, acctDomain := alert.SplitEmail(account)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "mail_account_spray",
Message: fmt.Sprintf("Mail password spray targeting %s: %d unique IPs in %v",
account, len(a.ips), t.window),
Details: "Distributed login attempts across many IPs against one mailbox (visibility only — no auto-block).",
Timestamp: now,
SourceIP: ip,
Domain: acctDomain,
Mailbox: account,
})
}
}
t.enforceMaxTracked(ip)
t.findingsEmitted += int64(len(findings))
return findings
}
// Stats returns cumulative Record invocations and findings emitted since
// startup. Used by the daemon's periodic diagnostic log.
func (t *mailAuthTracker) Stats() (calls, emits int64) {
t.mu.Lock()
defer t.mu.Unlock()
return t.recordCalls, t.findingsEmitted
}
// RecordBackendFailure records one auth-backend failure observation (dovecot
// could not reach the credential backend). It returns a one-shot operator
// warning the first time the subsystem crosses into "degraded" within a
// suppression window, so an outage is visible rather than silently pausing
// detection. Callers feed this from log lines matched by isMailAuthBackendError.
func (t *mailAuthTracker) RecordBackendFailure() []alert.Finding {
t.mu.Lock()
defer t.mu.Unlock()
now := t.now()
t.recordBackendFailureTime(now)
if len(t.backendErr) < mailBackendDegradedThreshold || now.Before(t.backendWarnUntil) {
return nil
}
t.backendWarnUntil = now.Add(t.suppression)
t.findingsEmitted++
return []alert.Finding{{
Severity: alert.Warning,
Check: "mail_auth_backend_degraded",
Message: fmt.Sprintf("Mail auth backend degraded: %d backend failures in %v; mail brute-force auto-block paused",
len(t.backendErr), t.window),
Details: "Dovecot could not reach the credential backend (e.g. cpdoveauthd). Logins fail regardless of password, so brute-force and subnet auto-blocks are paused to avoid locking out legitimate users. Investigate the auth daemon.",
Timestamp: now,
}}
}
// recordBackendFailureTime keeps only the newest observations needed to prove
// the degraded invariant. Backend outages can flood logs; tracking more than
// the threshold would add memory pressure without changing detection.
func (t *mailAuthTracker) recordBackendFailureTime(now time.Time) {
t.backendErr = pruneTimes(t.backendErr, now.Add(-t.window))
if len(t.backendErr) == 0 {
t.backendErr = nil
}
if len(t.backendErr) >= mailBackendDegradedThreshold {
const keep = mailBackendDegradedThreshold - 1
if cap(t.backendErr) > mailBackendDegradedThreshold {
recent := make([]time.Time, mailBackendDegradedThreshold)
copy(recent, t.backendErr[len(t.backendErr)-keep:])
recent[keep] = now
t.backendErr = recent
return
}
copy(t.backendErr, t.backendErr[len(t.backendErr)-keep:])
t.backendErr = t.backendErr[:mailBackendDegradedThreshold]
t.backendErr[keep] = now
return
}
t.backendErr = append(t.backendErr, now)
}
// backendDegraded reports whether the mail auth backend looks down: at least
// mailBackendDegradedThreshold backend-failure observations within the window.
// Caller must hold t.mu.
func (t *mailAuthTracker) backendDegraded(now time.Time) bool {
t.backendErr = pruneTimes(t.backendErr, now.Add(-t.window))
if len(t.backendErr) == 0 {
t.backendErr = nil
}
return len(t.backendErr) >= mailBackendDegradedThreshold
}
func pruneMailAccountTimes(accounts map[string][]time.Time, cutoff time.Time) {
for account, times := range accounts {
times = pruneTimes(times, cutoff)
if len(times) == 0 {
delete(accounts, account)
continue
}
accounts[account] = times
}
}
func appendMailAccountTime(accounts map[string][]time.Time, account string, ts time.Time) map[string][]time.Time {
if account == "" {
return accounts
}
if accounts == nil {
accounts = make(map[string][]time.Time)
}
accounts[account] = append(accounts[account], ts)
return accounts
}
// RecordSuccess processes a successful mail login. Emits mail_account_compromised
// when the successful IP has repeated recent failed auths for the same account
// and that mailbox is failure-dominant from this IP. Prior successful logins for
// the same mailbox look like a legit client that merely mistyped, so they are
// not flagged.
//
// ip and account MUST both be non-empty. Caller filters infra/private/loopback
// IPs before invoking.
func (t *mailAuthTracker) RecordSuccess(ip, account string) (findings []alert.Finding) {
if ip == "" || account == "" {
return nil
}
account = normalizeMailAuthAccount(account)
// Registered before tracker cleanup so resolution sees no held locks.
defer func() { stampMailAccountOwner(findings, account) }()
t.mu.Lock()
defer t.mu.Unlock()
// Successes create per-IP entries too; keep the tracker bounded on every
// path (runs under the held lock, before Unlock).
defer t.enforceMaxTracked(ip)
now := t.now()
cutoff := now.Add(-t.window)
// Track per-IP successes unconditionally. The current success is recorded at
// function exit so the compromise gate can distinguish prior legit activity
// from a padding login that belongs to the event being classified.
e, ok := t.ips[ip]
if !ok {
e = &mailIPEntry{}
t.ips[ip] = e
}
e.succ = pruneTimes(e.succ, cutoff)
pruneMailAccountTimes(e.successAccounts, cutoff)
pruneMailAccountTimes(e.failedAccounts, cutoff)
// A guessed-password success should not train the long-lived good-source cache.
recordGoodSource := true
defer func() {
e.succ = append(e.succ, now)
e.successAccounts = appendMailAccountTime(e.successAccounts, account, now)
if t.slowThreshold > 0 && t.slowWindow > 0 {
e.slowLastSuccess = now
}
if recordGoodSource {
e.recordGoodAuth(account, now)
}
e.lastSeen = now
}()
a, ok := t.accounts[account]
if !ok {
return nil
}
for ipKey, ts := range a.ips {
if ts.Before(cutoff) {
delete(a.ips, ipKey)
}
}
if _, failedRecently := a.ips[ip]; !failedRecently {
return nil
}
// A mailbox that already succeeds from this IP is a legit owner who mistyped,
// not a takeover. Genuine password guessing is failure-dominant for the same
// mailbox.
e.times = pruneTimes(e.times, cutoff)
targetFailures := len(e.failedAccounts[account])
if targetFailures < 2 || e.accountSuccessDominant(account) || e.establishedGood(account, now, t.window) {
return nil
}
recordGoodSource = false
if now.Before(a.compromiseSuppressed) {
return nil
}
a.compromiseSuppressed = now.Add(t.suppression)
_, compDomain := alert.SplitEmail(account)
t.findingsEmitted++
severity := alert.Critical
message := fmt.Sprintf("Mail account compromise: successful login for %s from %s after recent auth failures",
account, ip)
details := "Attacker succeeded after repeated failed attempts from the same IP for this mailbox. Rotate password and revoke sessions."
// An IP with established standing on several other mailboxes is far more
// likely a shared office or agency device carrying one stale credential
// than a takeover: real credential attacks come from sources with no
// legitimate multi-mailbox history on this host. Keep the finding for
// visibility but downgrade it below the auto-block bar.
if n := e.establishedOtherGoodAccounts(account, now, t.window); n >= mailEstablishedSourceAccounts {
severity = alert.High
message += " (established multi-mailbox source)"
details = fmt.Sprintf("Source IP holds established successful logins to %d other mailboxes on this host, so this is more likely a shared office or agency device with a stale credential than a takeover. Verify with the customer before acting; not auto-blocked.", n)
}
return []alert.Finding{{
Severity: severity,
Check: "mail_account_compromised",
Message: message,
Details: details,
Timestamp: now,
SourceIP: ip,
Domain: compDomain,
Mailbox: account,
}}
}
// Purge removes stale tracker entries. Good-source-only IPs live until
// mailGoodSourceTTL; failure and short success history uses the detector window.
// Called from a background goroutine every minute.
func (t *mailAuthTracker) Purge() {
t.mu.Lock()
defer t.mu.Unlock()
now := t.now()
activityCutoff := now.Add(-(t.window + t.suppression))
windowCutoff := now.Add(-t.window)
goodCutoff := now.Add(-mailGoodSourceTTL)
slowCutoff := now.Add(-t.slowWindow)
for k, e := range t.ips {
e.times = pruneTimes(e.times, windowCutoff)
e.succ = pruneTimes(e.succ, windowCutoff)
e.slowTimes = pruneTimes(e.slowTimes, slowCutoff)
pruneSlowAccounts(e.slowAccounts, slowCutoff)
if e.slowLastSuccess.Before(slowCutoff) {
e.slowLastSuccess = time.Time{}
}
if e.slowUnvouched.Before(slowCutoff) {
e.slowUnvouched = time.Time{}
}
pruneMailAccountTimes(e.successAccounts, windowCutoff)
pruneMailAccountTimes(e.failedAccounts, windowCutoff)
e.pruneGoodSource(goodCutoff)
if len(e.times) == 0 && len(e.succ) == 0 && len(e.successAccounts) == 0 &&
len(e.failedAccounts) == 0 && len(e.goodLast) == 0 &&
len(e.slowTimes) == 0 && e.slowLastSuccess.IsZero() && !e.lastSeen.After(activityCutoff) {
delete(t.ips, k)
}
}
for k, s := range t.subnets {
for ip, ts := range s.ips {
if ts.Before(windowCutoff) {
delete(s.ips, ip)
}
}
if len(s.ips) == 0 && !s.lastSeen.After(activityCutoff) {
delete(t.subnets, k)
}
}
for k, a := range t.accounts {
for ip, ts := range a.ips {
if ts.Before(windowCutoff) {
delete(a.ips, ip)
}
}
if len(a.ips) == 0 && !a.lastSeen.After(activityCutoff) {
delete(t.accounts, k)
}
}
t.backendErr = pruneTimes(t.backendErr, windowCutoff)
if len(t.backendErr) == 0 {
t.backendErr = nil
}
}
// enforceMaxTracked evicts the least-recently-seen entries until total tracked
// state is <= 95% of maxTracked. Batch target avoids re-sorting on every
// subsequent insert. Caller must hold t.mu.
// keepIP is the source entry the caller just wrote. It is never evicted:
// dropping the entry a Record is currently accumulating into means a single
// source can never reach its threshold while the table is under pressure, so
// a flood would switch detection off for the very source causing it.
func (t *mailAuthTracker) enforceMaxTracked(keepIP string) {
total := len(t.ips) + len(t.subnets) + len(t.accounts)
if total <= t.maxTracked {
return
}
// Evict to 95% of cap so subsequent inserts don't re-trigger the sort.
target := t.maxTracked * 95 / 100
now := t.now()
// Victims are ranked by what they cost to lose, then by age. Account
// and subnet keys are attacker-chosen (any string after "user=<...>",
// any /24), so a flood of unique mailbox names used to run one LRU over
// the IP entries too and evict a source's good-source standing or
// slow-brute evidence, exactly the state that keeps a legitimate client
// exempt from auto-block. Accounts go first, then subnets, then sources
// with no live failures, and sources with in-window failure or
// slow-brute evidence last.
type victim struct {
kind string // "ip" | "subnet" | "account"
key string
seen time.Time
rank int
}
victims := make([]victim, 0, total)
for k, v := range t.ips {
if k == keepIP {
continue
}
victims = append(victims, victim{"ip", k, v.lastSeen, v.evictionRank(now, t.window, t.slowWindow)})
}
for k, v := range t.subnets {
victims = append(victims, victim{"subnet", k, v.lastSeen, evictionRankSubnet})
}
for k, v := range t.accounts {
victims = append(victims, victim{"account", k, v.lastSeen, evictionRankAccount})
}
sort.Slice(victims, func(i, j int) bool {
if victims[i].rank != victims[j].rank {
return victims[i].rank < victims[j].rank
}
return victims[i].seen.Before(victims[j].seen)
})
for i := 0; i < len(victims); i++ {
if len(t.ips)+len(t.subnets)+len(t.accounts) <= target {
break
}
v := victims[i]
switch v.kind {
case "ip":
delete(t.ips, v.key)
case "subnet":
delete(t.subnets, v.key)
case "account":
delete(t.accounts, v.key)
}
}
}
// Eviction ranks, lowest evicted first.
const (
evictionRankAccount = iota
evictionRankSubnet
evictionRankIdleIP
evictionRankActiveIP
evictionRankProtectedIP
)
// evictionRank keeps established legitimate sources and meaningful slow
// evidence behind disposable one-shot failure entries. Otherwise an attacker
// can send one failure from each fresh IP, make every attacker entry "active",
// and evict the older good-source standing that prevents false blocks.
func (e *mailIPEntry) evictionRank(now time.Time, window, slowWindow time.Duration) int {
for _, last := range e.goodLast {
if !last.Before(now.Add(-mailGoodSourceTTL)) {
return evictionRankProtectedIP
}
}
if slowWindow > 0 {
cutoff := now.Add(-slowWindow)
if (!e.slowLastSuccess.IsZero() && !e.slowLastSuccess.Before(cutoff)) || countTimesAtOrAfter(e.slowTimes, cutoff) > 1 {
return evictionRankProtectedIP
}
}
if hasTimeAtOrAfter(e.times, now.Add(-window)) {
return evictionRankActiveIP
}
return evictionRankIdleIP
}
func hasTimeAtOrAfter(times []time.Time, cutoff time.Time) bool {
for _, ts := range times {
if !ts.Before(cutoff) {
return true
}
}
return false
}
func countTimesAtOrAfter(times []time.Time, cutoff time.Time) int {
n := 0
for _, ts := range times {
if !ts.Before(cutoff) {
n++
}
}
return n
}
// isMailAuthBackendError reports whether a dovecot log line shows the auth
// backend itself failing (could not verify ANY credential), as opposed to an
// ordinary wrong-password failure. During such an outage every login fails, so
// these events must drive the degraded gate, not the brute-force counters.
func isMailAuthBackendError(line string) bool {
if isMailAuthLine(line) {
return false
}
msg := dovecotServiceMessage(line)
if msg == "" {
return false
}
detail := msg
if strings.HasPrefix(detail, "auth-worker") {
detail = authWorkerDetail(detail)
}
lowerDetail := strings.ToLower(detail)
if strings.Contains(lowerDetail, "cpdoveauthd.sock") {
return strings.Contains(lowerDetail, "failed to connect") ||
strings.Contains(lowerDetail, "socket error") ||
strings.Contains(lowerDetail, "connection refused")
}
if strings.Contains(lowerDetail, "temporary authentication failure") {
return strings.HasPrefix(msg, "auth-worker") || strings.Contains(strings.ToLower(msg), "auth:")
}
if !strings.HasPrefix(msg, "auth-worker") {
return false
}
return strings.Contains(lowerDetail, "connection refused") ||
strings.Contains(lowerDetail, "internal error")
}
func dovecotServiceMessage(line string) string {
tagEnd := strings.Index(line, ": ")
if tagEnd < 0 {
return ""
}
prefix := line[:tagEnd]
tagStart := strings.LastIndexAny(prefix, " \t")
if tagStart >= 0 {
prefix = prefix[tagStart+1:]
}
if !isDovecotProgramTag(prefix) {
return ""
}
return strings.TrimSpace(line[tagEnd+2:])
}
func isDovecotProgramTag(tag string) bool {
if tag == "dovecot" {
return true
}
const prefix = "dovecot["
if !strings.HasPrefix(tag, prefix) || !strings.HasSuffix(tag, "]") {
return false
}
pid := tag[len(prefix) : len(tag)-1]
if pid == "" {
return false
}
for i := 0; i < len(pid); i++ {
if pid[i] < '0' || pid[i] > '9' {
return false
}
}
return true
}
func authWorkerDetail(msg string) string {
parenDepth := 0
angleDepth := 0
for i := 0; i < len(msg); i++ {
switch msg[i] {
case '(':
if angleDepth == 0 {
parenDepth++
}
case ')':
if angleDepth == 0 && parenDepth > 0 {
parenDepth--
}
case '<':
if parenDepth == 0 {
angleDepth++
}
case '>':
if parenDepth == 0 && angleDepth > 0 {
angleDepth--
}
case ':':
if parenDepth == 0 && angleDepth == 0 {
return strings.TrimSpace(msg[i+1:])
}
}
}
return ""
}
// isMailAuthLine returns true for dovecot imap/pop3/managesieve login events.
func isMailAuthLine(line string) bool {
msg := dovecotServiceMessage(line)
return strings.HasPrefix(msg, "imap-login:") ||
strings.HasPrefix(msg, "pop3-login:") ||
strings.HasPrefix(msg, "managesieve-login:")
}
// dovecotLoginSucceeded reports whether a dovecot imap/pop3/managesieve line
// records a successful login. Dovecot emits two success formats depending on
// version and configuration:
//
// "<proto>-login: Logged in: user=<...>" (observed on production cPanel)
// "<proto>-login: Login: user=<...>" (classic dovecot)
//
// Both the mailbrute compromise detector and the geo new-country detector must
// accept BOTH; keying each consumer off a different single marker left one of
// them silently dead on whichever format the deployment happened to use.
func dovecotLoginSucceeded(line string) bool {
msg := dovecotServiceMessage(line)
return strings.HasPrefix(msg, "imap-login: Logged in: user=<") ||
strings.HasPrefix(msg, "imap-login: Login: user=<") ||
strings.HasPrefix(msg, "pop3-login: Logged in: user=<") ||
strings.HasPrefix(msg, "pop3-login: Login: user=<") ||
strings.HasPrefix(msg, "managesieve-login: Logged in: user=<") ||
strings.HasPrefix(msg, "managesieve-login: Login: user=<")
}
// extractMailLoginEvent parses a dovecot login line and returns
// (ip, account, success). Returns empty strings and false on parse failure.
//
// Real dovecot wire format (validated against production logs):
//
// Success: "imap-login: Logged in: user=<alice@x.ro>, method=PLAIN, rip=..."
// Failure: "imap-login: Login aborted: ... (auth failed, N attempts ...): user=<...>, method=..., rip=..."
//
// Success is matched via dovecotLoginSucceeded so both dovecot success formats
// count (an earlier version keyed only off "Logged in" and silently skipped
// every classic-format login, which broke RecordSuccess compromise detection).
// The failure marker is "(auth failed" with the opening paren, which
// distinguishes real auth failures from Login-aborted reasons like
// "no auth attempts" or TLS handshake errors.
func extractMailLoginEvent(line string) (ip, account string, success bool) {
switch {
case dovecotLoginSucceeded(line):
success = true
case strings.Contains(line, "(auth failed"):
success = false
default:
return "", "", false
}
// Extract account key via the configured (or default) extractor.
account = currentAccountExtractor().Extract(line)
// Extract rip=... field. Delimited by comma or whitespace.
if i := strings.Index(line, "rip="); i >= 0 {
rest := line[i+len("rip="):]
end := strings.IndexAny(rest, ", \n")
if end < 0 {
end = len(rest)
}
ip = rest[:end]
}
return ip, account, success
}
// dovecotAttemptsCap bounds how many failures one "Login aborted" line may
// record, so a forged or absurd count cannot flood the tracker.
const dovecotAttemptsCap = 20
// dovecotFailedAttempts returns the number of password attempts a Dovecot
// "Login aborted ... (auth failed, N attempts in S secs)" line reports, or 1
// when the line carries no count. Dovecot writes one such line per
// connection; counting it once let a client that tries many passwords per
// connection stay under every per-IP threshold.
func dovecotFailedAttempts(line string) int {
const marker = "(auth failed, "
i := strings.Index(line, marker)
if i < 0 {
return 1
}
rest := line[i+len(marker):]
end := strings.IndexByte(rest, ' ')
if end <= 0 || !strings.HasPrefix(strings.TrimSpace(rest[end:]), "attempts") {
return 1
}
n, err := strconv.Atoi(rest[:end])
if err != nil || n < 1 {
return 1
}
if n > dovecotAttemptsCap {
return dovecotAttemptsCap
}
return n
}
// recordDovecotFailure records one failure per attempt the Dovecot line
// reports and returns every finding those records produced.
func recordDovecotFailure(t *mailAuthTracker, ip, account, line string) []alert.Finding {
var findings []alert.Finding
for i := dovecotFailedAttempts(line); i > 0; i-- {
findings = append(findings, t.Record(ip, account)...)
}
return findings
}
package daemon
import (
"fmt"
"regexp"
"strings"
"sync/atomic"
"github.com/pidginhost/csm/internal/config"
)
// AccountExtractor pulls the account/mailbox identifier out of a mail
// server log line. Used by mailbrute for per-account scoring. Selected
// by cfg.Thresholds.MailBruteAccountKey at daemon startup and SIGHUP
// reload.
type AccountExtractor struct {
mode string
re *regexp.Regexp
}
// NewAccountExtractor parses the spec string from cfg.Thresholds.MailBruteAccountKey.
// Empty spec defaults to "builtin:dovecot-user" (matches the legacy behavior).
func NewAccountExtractor(spec string) (*AccountExtractor, error) {
switch {
case spec == "" || spec == "builtin:dovecot-user":
return &AccountExtractor{mode: "dovecot-user"}, nil
case spec == "builtin:postfix-sasl":
return &AccountExtractor{mode: "postfix-sasl"}, nil
case strings.HasPrefix(spec, "regex:"):
re, err := regexp.Compile(strings.TrimPrefix(spec, "regex:"))
if err != nil {
return nil, fmt.Errorf("invalid regex: %w", err)
}
if re.NumSubexp() < 1 {
return nil, fmt.Errorf("regex must contain at least one capture group")
}
return &AccountExtractor{mode: "regex", re: re}, nil
default:
return nil, fmt.Errorf("unknown extractor spec: %s", spec)
}
}
// Extract returns the account/mailbox key, or "" when no match.
func (e *AccountExtractor) Extract(line string) string {
switch e.mode {
case "dovecot-user":
return extractAngleBracket(line, "user=")
case "postfix-sasl":
return extractEqualsValue(line, "sasl_username=")
case "regex":
m := e.re.FindStringSubmatch(line)
if len(m) >= 2 {
return m[1]
}
}
return ""
}
// extractAngleBracket matches `key<value>` with balanced angle brackets.
func extractAngleBracket(line, key string) string {
idx := strings.Index(line, key+"<")
if idx < 0 {
return ""
}
end := strings.IndexByte(line[idx+len(key)+1:], '>')
if end < 0 {
return ""
}
return line[idx+len(key)+1 : idx+len(key)+1+end]
}
// extractEqualsValue matches `key=<value>` where value is delimited by
// whitespace or comma (postfix log format).
func extractEqualsValue(line, key string) string {
idx := strings.Index(line, key)
if idx < 0 {
return ""
}
rest := line[idx+len(key):]
end := strings.IndexAny(rest, " ,\t\n")
if end < 0 {
return rest
}
return rest[:end]
}
// defaultAccountExtractor is the package-level singleton set at daemon
// startup and safe config reloads.
var defaultAccountExtractor atomic.Pointer[AccountExtractor]
func installAccountExtractorFromConfig(cfg *config.Config) error {
ex, err := NewAccountExtractor(cfg.Thresholds.MailBruteAccountKey)
if err != nil {
return fmt.Errorf("invalid mail_brute_account_key: %w", err)
}
SetAccountExtractor(ex)
return nil
}
// SetAccountExtractor installs the configured extractor; called from
// Daemon.Run() and safe reload after applyDefaults has set the spec.
func SetAccountExtractor(ex *AccountExtractor) {
defaultAccountExtractor.Store(ex)
}
// currentAccountExtractor returns the installed extractor, or lazily
// initializes the default so test code that doesn't call SetAccountExtractor
// gets the legacy dovecot-user behavior.
func currentAccountExtractor() *AccountExtractor {
if ex := defaultAccountExtractor.Load(); ex != nil {
return ex
}
ex, _ := NewAccountExtractor("")
defaultAccountExtractor.Store(ex)
return ex
}
package daemon
import (
"context"
"net"
"path/filepath"
"time"
"github.com/pidginhost/csm/internal/firewall"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/mailranges"
"github.com/pidginhost/csm/internal/obs"
)
const (
mailRangesRefreshInterval = 12 * time.Hour
mailRangesStartupDelay = 5 * time.Minute
)
// mailRangesResolver is the DNS resolver for the periodic SPF-based refresh.
// Replaced in tests to avoid live network calls.
var mailRangesResolver mailranges.Resolver = net.DefaultResolver
// mailRangesReapplyFn applies updated provider nets to the firewall engine's
// DoS-exempt sets. Replaced in tests to inject reapply failures without
// requiring a live nftables connection.
var mailRangesReapplyFn = func(e *firewall.Engine, nets []*net.IPNet) error {
return e.RefreshDOSExemptSets(nets)
}
func (d *Daemon) mailRangesCachePath() string {
return filepath.Join(d.cfg.StatePath, "mailranges.json")
}
// initMailRanges loads the on-disk mail-provider range cache synchronously.
// This must be called before startFirewall() so that engine.SetDOSExemptProviderNets
// receives a full provider set and Apply() builds dos_exempt_nets from day zero.
//
// LoadCache errors are non-fatal: when the cache is absent or corrupt the
// embedded snapshot is published and the daemon continues. The error is
// logged so operators know the fallback is active.
func (d *Daemon) initMailRanges() {
if err := mailranges.LoadCache(d.mailRangesCachePath()); err != nil {
csmlog.Warn("mailranges: cache load failed, using embedded snapshot", "err", err)
}
d.wg.Add(1)
obs.Go("mailranges-refresh", d.mailRangesRefreshLoop)
}
// mailRangesRefreshLoop periodically refreshes the mail-provider IP ranges
// (Google, Microsoft) so the firewall's DoS-exempt sets remain current.
// A short startup delay mirrors the botranges updater so the daemon does not
// hit external DNS before the host is fully settled.
func (d *Daemon) mailRangesRefreshLoop() {
defer d.wg.Done()
select {
case <-d.stopCh:
return
case <-time.After(mailRangesStartupDelay):
}
d.doMailRangesRefresh()
ticker := time.NewTicker(mailRangesRefreshInterval)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
d.doMailRangesRefresh()
}
}
}
// doMailRangesRefresh runs one refresh cycle: resolves SPF records, updates
// the on-disk cache, and reapplies the firewall DoS-exempt sets. On reapply
// failure it restores the previous provider snapshot so the in-memory state
// stays consistent with what is actually live in nftables.
func (d *Daemon) doMailRangesRefresh() {
cachePath := d.mailRangesCachePath()
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
// Cancel in-flight DNS promptly on daemon shutdown.
go func() {
select {
case <-d.stopCh:
cancel()
case <-ctx.Done():
}
}()
// Snapshot the provider set before refresh so it can be restored if the
// subsequent firewall reapply fails. The snapshot and the reapply must be
// consistent: ProviderNets() must reflect what is in nftables.
prev := mailranges.ProviderSnapshot()
n, err := mailranges.Refresh(ctx, mailRangesResolver, cachePath)
if err != nil {
csmlog.Warn("mailranges: refresh error", "err", err)
}
if n == 0 {
// All providers failed; Refresh did not update the active snapshot.
return
}
csmlog.Info("mailranges: providers refreshed", "prefixes", n)
if d.fwEngine == nil {
return
}
if reapplyErr := mailRangesReapplyFn(d.fwEngine, mailranges.ProviderNets()); reapplyErr != nil {
// Restore the previous snapshot so ProviderNets() reflects what is
// actually in nftables, not the candidate update that failed to apply.
// Auto-response subnets for the failed candidate are not pruned.
mailranges.PublishProviderSnapshot(prev)
csmlog.Warn("mailranges: firewall reapply failed; previous snapshot restored", "err", reapplyErr)
}
}
package daemon
import "strings"
// modsecConfidence classifies how strongly a single ModSecurity deny indicates
// a real attack, independent of the rule's blocking action. It drives whether a
// deny may auto-escalate to a 24h firewall ban (high/unknown) or only counts
// toward the low-confidence visibility/backstop path (low). See
// docs/superpowers/specs/2026-06-27-modsec-escalation-fp-options.md.
type modsecConfidence int
const (
// modsecConfUnknown is a blocking deny CSM cannot classify from the parsed
// rule ID, message, or tags. Treated as escalation-eligible at the normal
// bar (fail-secure) so a new vendor rule never gets a silent no-ban path.
modsecConfUnknown modsecConfidence = iota
// modsecConfLow is a policy/anomaly/scoring deny whose own rule is not proof
// of hostile intent (content-type policy, anomaly score threshold). Never
// auto-bans on its own at the normal bar; only via the low-confidence
// backstop.
modsecConfLow
// modsecConfHigh is a specific attack/probe signal (SQLi, RCE, traversal,
// URL-encoding abuse, CSM custom deny, scanner). Escalation-eligible at the
// normal bar even when it is the only distinct rule.
modsecConfHigh
)
func (c modsecConfidence) String() string {
switch c {
case modsecConfLow:
return "low"
case modsecConfHigh:
return "high"
default:
return "unknown"
}
}
// modsecKnownLowConfRules are vendor policy/anomaly/scoring rule IDs verified to
// fire on legitimate-but-unusual traffic as often as on attacks. They are
// low-confidence by exact ID, but a high-confidence attack signal in the same
// message/tags still overrides (see classifyModSecConfidence ordering). Do NOT
// add broad rule-ID ranges here; only exact, fixture-verified IDs.
var modsecKnownLowConfRules = map[int]bool{
// COMODO CWAF policy/anomaly.
210710: true, // request content-type not allowed by policy
214930: true, // inbound points exceeded (anomaly threshold)
211170: true, // outbound points / scoring
211220: true, // outbound points / scoring
// OWASP CRS policy/anomaly.
920100: true, // invalid HTTP request line (protocol enforcement)
920420: true, // request content-type not allowed
920430: true, // HTTP protocol version not allowed (policy)
920440: true, // URL file extension restricted by policy
949110: true, // inbound anomaly score exceeded
959100: true, // outbound anomaly score exceeded
980130: true, // anomaly score reporting
}
// modsecAttackMsgKeywords are lowercase substrings of vendor messages that name
// a specific attack/probe class. Presence means high-confidence.
var modsecAttackMsgKeywords = []string{
"sql injection", "sqli", "remote command", "command injection",
"os command", "remote code execution", "code execution", "code injection",
"file inclusion", "lfi", "rfi", "cross-site scripting",
"xss", "path traversal", "directory traversal", "url encoding abuse",
"web shell", "webshell", "backdoor", "shellshock", "ssrf",
"server-side request forgery", "session fixation", "remote file",
"request smuggling", "response splitting", "crlf injection",
"object injection", "template injection", "xxe", "xml external entity",
"ldap injection", "nosql injection", "scanner", "exploit",
"injection attack", "deserializ",
}
// modsecAttackTagKeywords are lowercase substrings of the ModSecurity rule tag
// taxonomy that name a specific attack class. "attack-protocol" and
// "attack-generic" are deliberately excluded: protocol/anomaly policy rules
// carry them but are low-confidence.
var modsecAttackTagKeywords = []string{
"attack-sqli", "attack-rce", "attack-xss", "attack-lfi", "attack-rfi",
"attack-injection", "attack-disclosure", "attack-fixation", "attack-ssrf",
"attack-shell", "application-attack",
}
// modsecPolicyAnomalyKeywords are lowercase substrings that name a
// policy/anomaly/scoring decision (low-confidence) when no attack signal is
// present.
var modsecPolicyAnomalyKeywords = []string{
"anomaly", "inbound points", "outbound points", "points exceeded",
"total incoming points", "total inbound points", "content-type",
"content type", "not allowed by policy", "not allowed by the policy",
"protocol version", "score exceeded",
}
// classifyModSecConfidence classifies a single deny. Ordering is deliberate:
// a high-confidence attack signal always wins; CSM custom rules are high;
// otherwise an exact known-low ID or policy/anomaly wording (with no attack
// signal) is low; anything else is unknown (fail-secure, escalation-eligible).
//
// The modsec [severity] field is intentionally not an input: anomaly-scoring
// WAFs (COMODO CWAF) emit CRITICAL severity on benign policy/anomaly rules, so
// it is too noisy to separate attacks from policy hits. Rule ID, message, and
// tags carry the reliable signal.
func classifyModSecConfidence(ruleNum int, msg, tags, ruleFile string) modsecConfidence {
lc := strings.ToLower(msg + " " + tags)
// OWASP CRS names each rule file after the tag taxonomy of the rules in it
// (REQUEST-942-APPLICATION-ATTACK-SQLI.conf). LiteSpeed logs only the rule
// ID and file, so the file name is the only place that evidence survives.
attackTags := strings.ToLower(tags + " " + ruleFile)
// CSM custom rules are purpose-built attack/probe detections.
if ruleNum >= 900000 && ruleNum <= 900999 {
return modsecConfHigh
}
// Specific attack/probe evidence wins over any low signal.
if containsAny(lc, modsecAttackMsgKeywords) || containsAny(attackTags, modsecAttackTagKeywords) {
return modsecConfHigh
}
// Positive low evidence only: exact known-low ID or policy/anomaly wording.
if modsecKnownLowConfRules[ruleNum] {
return modsecConfLow
}
if containsAny(lc, modsecPolicyAnomalyKeywords) {
return modsecConfLow
}
return modsecConfUnknown
}
func containsAny(s string, subs []string) bool {
for _, sub := range subs {
if strings.Contains(s, sub) {
return true
}
}
return false
}
package daemon
import (
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/store"
)
const (
// modsecContextScanCap bounds the history walk per enriched IP. Denies that
// drive one escalation number in the dozens, so a few thousand is ample
// headroom while keeping the bbolt read cheap on a hammered host.
modsecContextScanCap = 2000
// modsecContextMaxItems caps the domains and URIs surfaced per block so the
// digest line stays readable.
modsecContextMaxItems = 5
)
// modsecEnricher returns the EnrichModSec lookup wired into the block digest.
// It reuses already-parsed modsec findings from the history store, so the
// digest can name targeted domains and request paths without re-reading the
// (multi-hundred-MB) audit log. The lookback matches the escalation window
// that produced the block. The lookup reads the live threshold so safe reloads
// of the ModSecurity escalation window also change the digest lookback.
func (d *Daemon) modsecEnricher(cfg *config.Config) func(ip string) (domains, uris []string) {
return func(ip string) ([]string, []string) {
live := d.currentCfg()
if live == nil {
live = cfg
}
_, win := modsecEscalationParams(live)
return aggregateModSecContext(ip, time.Now().Add(-win))
}
}
// aggregateModSecContext reads the per-deny modsec findings for one IP since the
// cutoff and returns its most-hit customer domains and request URIs. Escalation
// findings (no per-deny URI) and other IPs are skipped. Returns nil slices when
// no store is loaded or nothing matched.
func aggregateModSecContext(ip string, since time.Time) (domains, uris []string) {
db := store.Global()
if db == nil {
return nil, nil
}
findings := db.SearchHistorySince(since, modsecContextScanCap, func(f alert.Finding) bool {
return f.Check == "modsec_block_realtime" && f.SourceIP == ip
})
domainCounts := make(map[string]int)
uriCounts := make(map[string]int)
for _, f := range findings {
if f.Domain != "" {
domainCounts[f.Domain]++
}
if uri := modsecURIFromDetails(f.Details); uri != "" {
uriCounts[uri]++
}
}
return topByCount(domainCounts, modsecContextMaxItems), topByCount(uriCounts, modsecContextMaxItems)
}
// modsecURIFromDetails pulls the URI line out of a structured modsec finding
// Details blob ("Rule: ...\nURI: ...\n..."). Returns "" when absent.
func modsecURIFromDetails(details string) string {
for _, line := range strings.Split(details, "\n") {
if v, ok := strings.CutPrefix(line, "URI: "); ok {
return v
}
}
return ""
}
// topByCount returns up to n keys ordered by descending count, breaking ties on
// the key itself so output is deterministic. Returns nil when the map is empty.
func topByCount(counts map[string]int, n int) []string {
if len(counts) == 0 {
return nil
}
keys := make([]string, 0, len(counts))
for k := range counts {
keys = append(keys, k)
}
sort.Slice(keys, func(i, j int) bool {
if counts[keys[i]] != counts[keys[j]] {
return counts[keys[i]] > counts[keys[j]]
}
return keys[i] < keys[j]
})
if len(keys) > n {
keys = keys[:n]
}
return keys
}
package daemon
import (
"time"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/modsec"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
)
// modsecRegistryRefresh controls how often the rule-action registry is
// checked against disk. ModSec rule files change rarely (vendor pack updates,
// cPanel modsec_assemble nightly run), so a coarse interval keeps the cost
// negligible while still picking up operator edits within minutes.
const modsecRegistryRefresh = 5 * time.Minute
// Seams: the refresh probes the platform and parses the rule tree, and tests
// need to count both without a web server on the host.
var (
// modsecProbeRuleDirs re-runs detection (not the cached Detect) so a
// web-server mis-detection at boot -- LiteSpeed probed before lsws
// finished starting, which points RuleDirs at non-existent directories --
// self-heals on a later refresh instead of staying wrong for the daemon's
// lifetime. Each call forks one process per candidate unit, which is why
// refreshModSecRegistry stops calling it once a rule set has loaded.
modsecProbeRuleDirs = func() []string { return modsec.RuleDirs(platform.DetectFreshWithOverrides()) }
modsecBuildRegistry = modsec.BuildRegistry
)
// modsecRegistryState carries what the last refresh learned. Only the refresh
// path touches it: once at startup, then from the refresh goroutine.
type modsecRegistryState struct {
dirs []string
fingerprint string
// loaded records whether the last build produced a non-empty rule set
// from these dirs. Otherwise the platform is probed on every refresh,
// because an empty registry is exactly the symptom of detection having
// resolved the wrong directories.
loaded bool
}
// initModSecRegistry builds the rule-action registry once at startup and
// installs it as the package-level singleton. The registry tells the
// LiteSpeed log-line classifier which "triggered!" matches actually denied
// the request and which were pass-action informational rules. Without this,
// pass-action vendor rules (Comodo CWAF id 210710, 214930, ...) would be
// counted as denies, falsely escalating to a 24-hour auto-block of any IP
// that hits them three times in ten minutes.
//
// The build is failure-soft: missing rule directories yield an empty
// registry. With no prior healthy registry, ambiguous LiteSpeed "triggered!"
// lines are warnings until a later refresh loads rule actions.
func (d *Daemon) initModSecRegistry() {
d.refreshModSecRegistry()
d.wg.Add(1)
obs.Go("modsec-registry-refresh", d.modsecRegistryRefreshLoop)
}
func (d *Daemon) refreshModSecRegistry() {
state := &d.modsecRegistry
if !state.loaded || len(state.dirs) == 0 {
state.dirs = modsecProbeRuleDirs()
}
fingerprint, present := modsec.RuleTreeFingerprint(state.dirs)
if state.loaded && !present {
// The directories the last detection resolved are gone: the web
// server was swapped out, or the vendor pack removed. Detection has
// to run again rather than keep reporting the old rule actions.
state.loaded = false
state.dirs = modsecProbeRuleDirs()
fingerprint, _ = modsec.RuleTreeFingerprint(state.dirs)
}
if state.loaded && fingerprint != "" && fingerprint == state.fingerprint {
// Only a complete fingerprint can establish that the rule contents
// match the last successful build.
return
}
reg, err := modsecBuildRegistry(state.dirs)
if err != nil {
csmlog.Warn("modsec rule-action registry build had errors", "err", err, "rules_loaded", reg.Len())
}
state.loaded = reg.Len() > 0
// An empty fingerprint means the build could not read every rule file, so
// the next refresh has to look again. A parse failure over bytes that were
// read in full keeps its fingerprint: reparsing them changes nothing.
state.fingerprint = ""
// ReplaceGlobal keeps a previously-healthy registry rather than blanking
// it to empty: the vendor rule tree is briefly empty during cPanel's
// modsec_assemble rewrite, and a blank registry loses known pass and deny
// actions.
if !modsec.ReplaceGlobal(reg) {
previousRules := 0
if prev := modsec.Global(); prev != nil {
previousRules = prev.Len()
}
csmlog.Warn("modsec rule-action registry refresh returned 0 rules; keeping previous rule actions",
"previous_rules", previousRules, "dirs", len(state.dirs))
return
}
state.fingerprint = reg.Fingerprint()
csmlog.Info("modsec rule-action registry loaded", "rules", reg.Len(), "dirs", len(state.dirs))
}
func (d *Daemon) modsecRegistryRefreshLoop() {
defer d.wg.Done()
ticker := time.NewTicker(modsecRegistryRefresh)
defer ticker.Stop()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
d.refreshModSecRegistry()
}
}
}
//go:build linux
package daemon
import (
"fmt"
"sync"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/queuehealth"
)
type notificationSource interface {
Pending() (int, error)
Read([]byte) (int, error)
Write([]byte) (int, error)
Close() error
}
type fanotifyDescriptor int
func (fd fanotifyDescriptor) Pending() (int, error) {
bytes, err := unix.IoctlGetInt(int(fd), unix.TIOCINQ)
if err != nil {
return 0, err
}
// Fanotify's FIONREAD counts metadata headers, including overflow records.
if bytes < 0 || bytes%metadataSize != 0 {
return 0, fmt.Errorf("invalid fanotify pending byte count: %d", bytes)
}
return bytes / metadataSize, nil
}
func (fd fanotifyDescriptor) Read(buf []byte) (int, error) { return unix.Read(int(fd), buf) }
func (fd fanotifyDescriptor) Write(buf []byte) (int, error) { return unix.Write(int(fd), buf) }
func (fd fanotifyDescriptor) Close() error { return unix.Close(int(fd)) }
// Notification descriptors are nonblocking. Serialize their syscalls with
// close so health reads and permission responses cannot use a recycled fd.
// Parsing runs outside that lock and keeps its own tracked batch.
type notificationQueue struct {
mu sync.Mutex
source notificationSource
sampled *queuehealth.Sampled
unmeasured queuehealth.Dwell
losses *queuehealth.Tracker
batches *queuehealth.Tracker
consumed uint64
closed bool
closeErr error
unavailable bool
variableRecords bool
}
func newNotificationQueue(source notificationSource, losses, batches *queuehealth.Tracker) *notificationQueue {
return ¬ificationQueue{
source: source, losses: losses, batches: batches,
sampled: queuehealth.NewSampled(0, "records", time.Minute),
}
}
func (q *notificationQueue) snapshot(now func() time.Time) (kernel, reader queuehealth.Status) {
q.mu.Lock()
defer q.mu.Unlock()
if !q.closed {
pending, err := q.source.Pending()
q.unavailable = err != nil || pending < 0
if !q.unavailable {
q.sampled.Observe(now(), pending, q.consumed)
}
}
at := now()
kernel = q.sampled.Snapshot(at)
losses := q.losses.Snapshot(at)
kernel.DroppedTotal, kernel.RecentDrops = losses.DroppedTotal, losses.RecentDrops
// The group limit is not exposed. The current sysctl can differ from the
// value copied at creation. An overflow marker also omits its loss count.
kernel.CapacityUnavailable, kernel.DroppedLowerBound = true, true
// A live reading can be retried, so one failure is not yet a degradation.
// The reading taken at close is all there will ever be.
unmeasured := q.unmeasured.Held(at, q.unavailable, queuehealth.MeasurementWindow) || q.unavailable && q.closed
if q.unavailable && !q.closed {
// This reading invalidates the depth the sampled reason came from.
kernel.Depth, kernel.LagSeconds = 0, 0
kernel.DepthUnavailable, kernel.LagBasis = true, "unavailable"
kernel.Status, kernel.Reason = "ok", ""
}
switch {
case losses.Status == "degraded" && kernel.Reason == "":
// Records the kernel already dropped outrank an unreadable depth.
kernel.Status, kernel.Reason = losses.Status, losses.Reason
case unmeasured:
kernel.Status, kernel.Reason = "degraded", "measurement_unavailable"
}
reader = q.batches.Snapshot(at)
reader.DepthUnit = "batches"
return kernel, reader
}
func (q *notificationQueue) read(buf []byte, process func([]byte)) (int, error) {
q.mu.Lock()
if q.closed {
q.mu.Unlock()
return 0, unix.EBADF
}
n, err := q.source.Read(buf)
if err != nil || n <= 0 {
q.mu.Unlock()
return n, err
}
q.consumed++
work := queuehealth.Work[[]byte]{Value: buf[:n], Ticket: q.batches.Begin(time.Now())}
q.mu.Unlock()
work.Process(process)
return n, err
}
func (q *notificationQueue) write(buf []byte) (int, error) {
q.mu.Lock()
defer q.mu.Unlock()
if q.closed {
return 0, unix.EBADF
}
return q.source.Write(buf)
}
// Watch changes and readiness polling share the same descriptor lifetime as reads.
func (q *notificationQueue) useDescriptor(fn func() (int, error)) (int, error) {
q.mu.Lock()
defer q.mu.Unlock()
if q.closed {
return 0, unix.EBADF
}
return fn()
}
func (q *notificationQueue) close() error {
q.mu.Lock()
defer q.mu.Unlock()
if q.closed {
return q.closeErr
}
pending, err := q.source.Pending()
if err != nil || pending < 0 {
q.unavailable = true
} else {
q.unavailable = false
// Events can still arrive before close. Retain only the known minimum.
if q.variableRecords {
// A byte count cannot identify the number of variable-length records.
if pending > 0 {
q.losses.Lose(time.Now(), 1)
}
} else {
q.losses.Lose(time.Now(), uint64(pending))
}
}
q.closeErr = q.source.Close()
q.closed = true
q.sampled.Observe(time.Now(), 0, q.consumed)
return q.closeErr
}
package daemon
import (
"fmt"
"os"
"github.com/pidginhost/csm/internal/store"
)
// A pending manual apply must be settled before switching posture. Recovering
// it could restore an enforcing config or ruleset; ignoring it would abandon
// the operator's rollback deadline.
func (d *Daemon) checkObserveStartupRecovery() error {
if !d.cfg.ObserveMode() {
return nil
}
if db := store.Global(); db != nil {
if _, pending := db.GetFirewallRollback(); pending {
return fmt.Errorf("mode: observe cannot start with pending firewall settings recovery; resolve the pending apply in enforce mode before switching modes")
}
}
marker, _, _ := firewallRollbackFiles(d.cfg.StatePath)
if _, err := os.Stat(marker); err == nil {
return fmt.Errorf("mode: observe cannot start with pending firewall rules recovery; resolve the pending apply in enforce mode before switching modes")
} else if !os.IsNotExist(err) {
return fmt.Errorf("mode: observe cannot check pending firewall recovery: %w", err)
}
return nil
}
package daemon
import (
"bufio"
"fmt"
"net"
"os"
"sort"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/obs"
)
const pamSocketPath = "/var/run/csm/pam.sock"
const (
defaultPAMFailureThreshold = 5
defaultPAMFailureWindowMin = 10
defaultCredStuffingDistinctAccounts = 5
)
// PAMListener listens on a Unix socket for authentication events from the
// pam_csm.so PAM module. Tracks failures per IP and triggers CSM auto-blocking.
type PAMListener struct {
cfg *config.Config
alertCh chan<- alert.Finding
listener net.Listener
mu sync.Mutex
failures map[string]*pamFailureTracker
// stopCh is set by Run so emit can abort a send when the daemon is
// shutting down. Per-connection goroutines are not tracked by a
// WaitGroup, so without this escape a goroutine blocked on the
// undrained alert channel would leak past shutdown. Nil for
// hand-constructed listeners in tests, where emit keeps blocking-send
// semantics.
stopCh <-chan struct{}
// useActiveConfig is true for the daemon-owned listener, so SIGHUP
// threshold changes apply without rebuilding the socket. Unit tests that
// assemble a listener by hand keep using their local cfg.
useActiveConfig bool
// stuffing flags one source IP failing against many distinct accounts
// (credential stuffing / password spraying breadth) -- complementary to
// the per-IP failure-count brute-force trigger below. Guarded by mu.
stuffing *credentialStuffingDetector
// startedAt anchors the upstream-probe verdict so the dashboard does
// not flag a freshly-started daemon as "deaf" before the PAM module
// has had a chance to emit anything. Compared against lastPeerNanos
// to decide whether silence is informative.
startedAt time.Time
// lastPeerNanos is an atomic UnixNano timestamp updated on every
// inbound connection (regardless of payload validity). Used by the
// upstream probe to detect a missing PAM module hook.
lastPeerNanos atomic.Int64
}
type pamFailureTracker struct {
count int
firstSeen time.Time
lastSeen time.Time
users map[string]bool
services map[string]bool
accounts map[string]*pamAccountFailures
blocked bool
}
type pamAccountFailures struct {
count int
firstSeen time.Time
lastSeen time.Time
services map[string]bool
}
// NewPAMListener creates a Unix socket listener for PAM events.
func NewPAMListener(cfg *config.Config, alertCh chan<- alert.Finding) (*PAMListener, error) {
// Ensure socket directory exists
if err := os.MkdirAll("/var/run/csm", 0750); err != nil {
return nil, fmt.Errorf("creating socket dir: %w", err)
}
// Remove stale socket
os.Remove(pamSocketPath)
listener, err := net.Listen("unix", pamSocketPath)
if err != nil {
return nil, fmt.Errorf("listening on %s: %w", pamSocketPath, err)
}
// The PAM module runs inside privileged auth stacks, so the socket can stay
// root-only instead of accepting arbitrary local writers.
_ = os.Chmod(pamSocketPath, 0600)
_, window, distinct := pamThresholds(cfg)
return &PAMListener{
cfg: cfg,
alertCh: alertCh,
listener: listener,
failures: make(map[string]*pamFailureTracker),
startedAt: time.Now(),
useActiveConfig: true,
stuffing: newCredentialStuffingDetector(distinct, window, nil),
}, nil
}
// UpstreamProbe returns an UpstreamResult describing whether the PAM
// module hook is feeding the socket. The probe is cheap (single atomic
// load + clock read) so it is safe to wire into the components API.
//
// Fresh verdict:
// - At least one inbound connection within pamUpstreamFreshWindow.
// - Or the daemon has been up for less than pamUpstreamGracePeriod
// (so a freshly-started daemon is not flagged before the first
// real auth happens).
//
// LastActivity is the most recent connection time, or the daemon start
// time when no connection has arrived. Reason explains a !Fresh verdict
// to operators.
func (p *PAMListener) UpstreamResult() health.UpstreamResult {
last := time.Unix(0, p.lastPeerNanos.Load())
if p.lastPeerNanos.Load() == 0 && time.Since(p.startedAt) < pamUpstreamGracePeriod {
return health.UpstreamResult{Fresh: true, LastActivity: p.startedAt}
}
if p.lastPeerNanos.Load() != 0 && time.Since(last) < pamUpstreamFreshWindow {
return health.UpstreamResult{Fresh: true, LastActivity: last}
}
reason := "no PAM module hook feeding the socket; install pam_csm.so and add `session optional pam_csm.so` to the relevant /etc/pam.d/ files"
activity := last
if p.lastPeerNanos.Load() == 0 {
activity = p.startedAt
}
return health.UpstreamResult{Fresh: false, LastActivity: activity, Reason: reason}
}
const (
// pamUpstreamFreshWindow is how recently a peer connection must have
// arrived before the upstream is considered alive. Sized to comfortably
// span the longest realistic gap between auth events on a host that
// has the PAM module installed.
pamUpstreamFreshWindow = 24 * time.Hour
// pamUpstreamGracePeriod is the post-start window during which a
// silent socket is not yet flagged as deaf.
pamUpstreamGracePeriod = 15 * time.Minute
)
// Run accepts connections and processes PAM events.
func (p *PAMListener) Run(stopCh <-chan struct{}) {
p.stopCh = stopCh
// Start cleanup goroutine to expire old failure records
obs.Go("pam-cleanup", func() { p.cleanupLoop(stopCh) })
// Accept connections
obs.Go("pam-accept", func() {
for {
conn, err := p.listener.Accept()
if err != nil {
select {
case <-stopCh:
return
default:
fmt.Fprintf(os.Stderr, "[%s] PAM listener accept error: %v\n", ts(), err)
time.Sleep(100 * time.Millisecond)
continue
}
}
obs.SafeGo("pam-conn", func() { p.handleConnection(conn) })
}
})
<-stopCh
}
// Stop closes the listener and removes the socket file.
func (p *PAMListener) Stop() {
_ = p.listener.Close()
os.Remove(pamSocketPath)
}
func (p *PAMListener) handleConnection(conn net.Conn) {
defer func() { _ = conn.Close() }()
// Record the connection moment regardless of peer trust so the
// upstream probe can distinguish "PAM hook not installed" (no
// connections at all) from "PAM hook present but rejected as
// untrusted" (connections happen but never deliver payload).
p.lastPeerNanos.Store(time.Now().UnixNano())
if !isTrustedPAMPeer(conn) {
return
}
_ = conn.SetDeadline(time.Now().Add(1 * time.Second))
scanner := bufio.NewScanner(conn)
for scanner.Scan() {
line := scanner.Text()
p.processEvent(line)
}
}
// processEvent handles a single PAM event line.
// Format: FAIL ip=1.2.3.4 user=root service=sshd
//
// OK ip=1.2.3.4 user=root service=sshd
func (p *PAMListener) processEvent(line string) {
parts := strings.SplitN(strings.TrimSpace(line), " ", 2)
if len(parts) < 2 {
return
}
eventType := parts[0]
kvPart := parts[1]
var ip, user, service string
for _, kv := range strings.Fields(kvPart) {
switch {
case strings.HasPrefix(kv, "ip="):
ip = kv[3:]
case strings.HasPrefix(kv, "user="):
user = kv[5:]
case strings.HasPrefix(kv, "service="):
service = kv[8:]
}
}
if ip == "" || ip == "-" || ip == "127.0.0.1" {
return
}
// Skip infra IPs. Read from the live config like the thresholds are: an
// infrastructure address added by reload must stop counting at once.
if isInfraIP(ip, p.currentCfg().InfraIPs) {
return
}
switch eventType {
case "FAIL":
p.emit(p.recordFailure(ip, user, service))
case "OK":
p.clearFailuresForUser(ip, user)
// Successful login from non-infra IP - informational alert
p.emit([]alert.Finding{{
Severity: alert.High,
Check: "pam_login",
Message: fmt.Sprintf("Login success from non-infra IP: %s (user: %s, service: %s)", ip, user, service),
Timestamp: time.Now(),
SourceIP: ip,
}})
}
}
// emit forwards findings to the alert channel. It must be called WITHOUT
// p.mu held: a stalled alert consumer blocks the send, and holding the lock
// across that send would wedge recordFailure, clearFailures, and the cleanup
// loop, letting failure trackers grow without bound.
func (p *PAMListener) emit(findings []alert.Finding) {
for i, f := range findings {
if !alert.Enqueue(p.alertCh, f, p.stopCh) {
alert.RecordQueueLoss(p.alertCh, uint64(len(findings[i+1:])))
// Shutting down and the dispatcher has stopped draining;
// drop the remaining findings rather than leak this
// goroutine. A nil stopCh (hand-constructed listener)
// makes this case unselectable, preserving blocking send.
return
}
}
}
// recordFailure updates per-IP failure state and returns any findings the
// update produced. Findings are returned rather than sent so the caller can
// emit them after releasing p.mu (see emit).
func (p *PAMListener) recordFailure(ip, user, service string) []alert.Finding {
p.mu.Lock()
defer p.mu.Unlock()
var findings []alert.Finding
cfg := p.currentCfg()
threshold, window, distinct := pamThresholds(cfg)
now := time.Now()
tracker, exists := p.failures[ip]
if !exists {
tracker = &pamFailureTracker{
firstSeen: now,
users: make(map[string]bool),
services: make(map[string]bool),
accounts: make(map[string]*pamAccountFailures),
}
p.failures[ip] = tracker
}
if tracker.accounts == nil {
tracker.accounts = make(map[string]*pamAccountFailures)
}
account := tracker.accounts[user]
if account == nil {
account = &pamAccountFailures{firstSeen: now, services: make(map[string]bool)}
tracker.accounts[user] = account
}
tracker.count++
tracker.lastSeen = now
tracker.users[user] = true
tracker.services[service] = true
account.count++
account.lastSeen = now
account.services[service] = true
// Credential-stuffing breadth signal: one source IP failing against many
// distinct accounts. Independent of the per-IP failure-count brute-force
// trigger below, so a low-and-slow campaign that stays under the count
// threshold per account is still caught. Fires once per window.
p.ensureCredentialStuffingDetectorLocked(distinct, window, now)
if accounts, fire := p.stuffing.Record(ip, user); fire {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "credential_stuffing",
Message: fmt.Sprintf("Credential stuffing: %s failed logins against %d distinct accounts", ip, len(accounts)),
Details: fmt.Sprintf("Accounts targeted: %s\nService(s): %s",
strings.Join(accounts, ", "), strings.Join(sortedBoolKeys(tracker.services), ", ")),
Timestamp: now,
SourceIP: ip,
SprayTargets: append([]string(nil), accounts...),
})
}
// Only block if within the time window
if now.Sub(tracker.firstSeen) > window {
// Window expired - reset tracker
tracker.count = 1
tracker.firstSeen = now
tracker.users = map[string]bool{user: true}
tracker.services = map[string]bool{service: true}
tracker.accounts = map[string]*pamAccountFailures{user: {
count: 1, firstSeen: now, lastSeen: now, services: map[string]bool{service: true},
}}
tracker.blocked = false
return findings
}
if tracker.count >= threshold && !tracker.blocked {
tracker.blocked = true
users := sortedBoolKeys(tracker.users)
services := sortedBoolKeys(tracker.services)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "pam_bruteforce",
Message: fmt.Sprintf("PAM brute-force detected: %s (%d failures in %ds)", ip, tracker.count, int(now.Sub(tracker.firstSeen).Seconds())),
Details: fmt.Sprintf("Users targeted: %s\nServices: %s",
strings.Join(users, ", "), strings.Join(services, ", ")),
Timestamp: now,
SourceIP: ip,
SprayTargets: users,
})
}
return findings
}
// clearFailuresForUser forgets the failures attributed to user from ip once
// that user logged in. Failures against other users stay, and so does the
// credential-stuffing breadth: an attacker walking many accounts who finally
// lands one (or who owns one valid account) must not reset the record against
// every other account. The tracker goes only when no failed user remains.
func (p *PAMListener) clearFailuresForUser(ip, user string) {
p.mu.Lock()
defer p.mu.Unlock()
p.stuffing.ClearAccount(ip, user)
tracker, ok := p.failures[ip]
if !ok {
return
}
delete(tracker.users, user)
delete(tracker.accounts, user)
if len(tracker.users) == 0 {
delete(p.failures, ip)
p.stuffing.Clear(ip)
return
}
if tracker.accounts != nil {
tracker.count = 0
tracker.firstSeen = time.Time{}
tracker.lastSeen = time.Time{}
tracker.services = make(map[string]bool)
for _, account := range tracker.accounts {
tracker.count += account.count
if tracker.firstSeen.IsZero() || account.firstSeen.Before(tracker.firstSeen) {
tracker.firstSeen = account.firstSeen
}
if account.lastSeen.After(tracker.lastSeen) {
tracker.lastSeen = account.lastSeen
}
for service := range account.services {
tracker.services[service] = true
}
}
threshold, _, _ := pamThresholds(p.currentCfg())
if tracker.count < threshold {
tracker.blocked = false
}
}
}
// cleanupLoop removes expired failure trackers every minute.
func (p *PAMListener) cleanupLoop(stopCh <-chan struct{}) {
ticker := time.NewTicker(1 * time.Minute)
defer ticker.Stop()
for {
select {
case <-stopCh:
return
case <-ticker.C:
p.cleanupAt(time.Now())
}
}
}
func (p *PAMListener) cleanupAt(now time.Time) {
_, window, distinct := pamThresholds(p.currentCfg())
cutoff := now.Add(-window)
p.mu.Lock()
defer p.mu.Unlock()
for ip, tracker := range p.failures {
if tracker.lastSeen.Before(cutoff) {
delete(p.failures, ip)
}
}
if p.stuffing != nil {
p.stuffing.Configure(distinct, window, now)
p.stuffing.PruneStale(now)
}
}
func (p *PAMListener) currentCfg() *config.Config {
if p.useActiveConfig {
if cfg := config.Active(); cfg != nil {
return cfg
}
}
return p.cfg
}
func pamThresholds(cfg *config.Config) (threshold int, window time.Duration, distinct int) {
threshold = defaultPAMFailureThreshold
windowMin := defaultPAMFailureWindowMin
distinct = defaultCredStuffingDistinctAccounts
if cfg != nil {
if cfg.Thresholds.PAMBruteforceThreshold > 0 {
threshold = cfg.Thresholds.PAMBruteforceThreshold
}
if cfg.Thresholds.PAMBruteforceWindowMin > 0 {
windowMin = cfg.Thresholds.PAMBruteforceWindowMin
}
if cfg.Thresholds.CredStuffingDistinctAccounts > 0 {
distinct = cfg.Thresholds.CredStuffingDistinctAccounts
}
}
return threshold, time.Duration(windowMin) * time.Minute, distinct
}
func (p *PAMListener) ensureCredentialStuffingDetectorLocked(distinct int, window time.Duration, now time.Time) {
if p.stuffing == nil {
p.stuffing = newCredentialStuffingDetector(distinct, window, nil)
return
}
p.stuffing.Configure(distinct, window, now)
}
func sortedBoolKeys(m map[string]bool) []string {
out := make([]string, 0, len(m))
for k := range m {
out = append(out, k)
}
sort.Strings(out)
return out
}
// isInfraIP checks if an IP is in the configured infra IP ranges.
// Duplicated here to avoid import cycle with checks package.
// isInfraIP defers to the shared matcher so an infra_ips entry written as a
// bare address ("203.0.113.5") counts here the way it does everywhere else;
// the local copy accepted CIDRs only and silently ignored bare entries.
func isInfraIP(ip string, infraNets []string) bool {
return checks.IsInfraIP(ip, infraNets)
}
//go:build linux
package daemon
import (
"net"
"golang.org/x/sys/unix"
)
func isTrustedPAMPeer(conn net.Conn) bool {
unixConn, ok := conn.(*net.UnixConn)
if !ok {
return false
}
rawConn, err := unixConn.SyscallConn()
if err != nil {
return false
}
trusted := false
controlErr := rawConn.Control(func(fd uintptr) {
// #nosec G115 -- socket fd from net.Conn.SyscallConn; POSIX fd fits in int.
cred, err := unix.GetsockoptUcred(int(fd), unix.SOL_SOCKET, unix.SO_PEERCRED)
if err == nil && cred != nil && cred.Uid == 0 {
trusted = true
}
})
return controlErr == nil && trusted
}
package daemon
import (
"fmt"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
)
// PasswordHijackDetector tracks WHM password changes from non-infra IPs
// and correlates them with subsequent cPanel logins to detect the attack
// pattern: attacker changes password via WHM → immediately logs in.
//
// Legitimate flow (excluded):
//
// Portal (infra IP) changes password via xml-api → user logs in
//
// Attack flow (detected):
//
// Attacker (non-infra IP) changes password via whostmgr → logs in within 60s
type PasswordHijackDetector struct {
mu sync.Mutex
recentChanges map[string]*passwordChange // account -> change info
cfg *config.Config
alertCh chan<- alert.Finding
stopCh <-chan struct{} // closed on daemon shutdown; lets sends escape
}
type passwordChange struct {
account string
ip string
timestamp time.Time
}
const hijackWindow = 120 * time.Second // time window to correlate password change + login
// NewPasswordHijackDetector creates a new detector.
func NewPasswordHijackDetector(cfg *config.Config, alertCh chan<- alert.Finding, stopCh <-chan struct{}) *PasswordHijackDetector {
return &PasswordHijackDetector{
recentChanges: make(map[string]*passwordChange),
cfg: cfg,
alertCh: alertCh,
stopCh: stopCh,
}
}
// emit delivers a finding without ever blocking the (synchronous) session-log
// watcher goroutine forever. It applies backpressure while the daemon is
// running, but gives up the moment shutdown is signaled -- after the alert
// dispatcher stops draining, a plain send would otherwise wedge d.wg.Wait().
// Callers must NOT hold d.mu while calling this.
func (d *PasswordHijackDetector) emit(f alert.Finding) {
alert.Enqueue(d.alertCh, f, d.stopCh)
}
// HandlePasswordChange records a WHM password change from a non-infra IP.
func (d *PasswordHijackDetector) HandlePasswordChange(account, ip string) {
if isInfraIPDaemon(ip, d.cfg.InfraIPs) || ip == "127.0.0.1" || ip == "internal" {
return // legitimate - portal or admin action
}
d.mu.Lock()
d.recentChanges[account] = &passwordChange{
account: account,
ip: ip,
timestamp: time.Now(),
}
d.mu.Unlock()
// Alert on the password change itself - non-infra WHM password change is
// always suspicious. Emit outside the lock so a saturated alert channel
// can never wedge the mutex.
d.emit(alert.Finding{
Severity: alert.Critical,
Check: "whm_password_change_noninfra",
Message: fmt.Sprintf("WHM password change from non-infra IP: %s (account: %s)", ip, account),
Details: "Password was changed via WHM from an IP outside your infrastructure. This is a strong indicator of account takeover.",
Timestamp: time.Now(),
SourceIP: ip,
TenantID: account,
})
}
// HandleLogin checks if a cPanel login matches a recent non-infra password change.
func (d *PasswordHijackDetector) HandleLogin(account, loginIP string) {
if isInfraIPDaemon(loginIP, d.cfg.InfraIPs) {
return
}
d.mu.Lock()
change, exists := d.recentChanges[account]
if exists {
delete(d.recentChanges, account)
}
d.mu.Unlock()
if !exists {
return
}
// Check if within the hijack window
if time.Since(change.timestamp) > hijackWindow {
return
}
// CONFIRMED ATTACK: password changed from non-infra IP, login within 120s
d.emit(alert.Finding{
Severity: alert.Critical,
Check: "password_hijack_confirmed",
Message: fmt.Sprintf("CONFIRMED ACCOUNT HIJACK: %s - password changed from %s, login from %s within %ds", account, change.ip, loginIP, int(time.Since(change.timestamp).Seconds())),
Details: fmt.Sprintf("Attack pattern: WHM password change from non-infra IP followed by immediate cPanel login.\nPassword change IP: %s\nLogin IP: %s\nTime between: %ds\n\nBoth IPs should be permanently blocked.", change.ip, loginIP, int(time.Since(change.timestamp).Seconds())),
Timestamp: time.Now(),
SourceIP: loginIP,
TenantID: account,
})
}
// Cleanup removes expired entries.
func (d *PasswordHijackDetector) Cleanup() {
d.mu.Lock()
defer d.mu.Unlock()
for account, change := range d.recentChanges {
if time.Since(change.timestamp) > hijackWindow*2 {
delete(d.recentChanges, account)
}
}
}
// ParseSessionLineForHijack extracts password change and login events
// from session log lines and feeds them to the detector.
func ParseSessionLineForHijack(line string, detector *PasswordHijackDetector) {
// WHM password change: [timestamp] info [whostmgr] IP PURGE account:token password_change
if strings.Contains(line, "[whostmgr]") && strings.Contains(line, "PURGE") && strings.Contains(line, "password_change") {
ip, account := parseWHMPurge(line)
if ip != "" && account != "" {
detector.HandlePasswordChange(account, ip)
}
}
// cPanel login: [timestamp] info [cpaneld] IP NEW account:token ...
if strings.Contains(line, "[cpaneld]") && strings.Contains(line, " NEW ") {
// Skip API sessions
if strings.Contains(line, "method=create_user_session") {
return
}
ip, account := parseCpanelSessionLogin(line)
if ip != "" && account != "" {
detector.HandleLogin(account, ip)
}
}
}
func parseWHMPurge(line string) (ip, account string) {
// Format: [timestamp] info [whostmgr] 198.51.100.50 PURGE account:token password_change
idx := strings.Index(line, "[whostmgr]")
if idx < 0 {
return "", ""
}
rest := strings.TrimSpace(line[idx+len("[whostmgr]"):])
fields := strings.Fields(rest)
if len(fields) < 3 {
return "", ""
}
ip = fields[0]
// Find account from PURGE account:token
for i, f := range fields {
if f == "PURGE" && i+1 < len(fields) {
parts := strings.SplitN(fields[i+1], ":", 2)
if len(parts) >= 1 {
account = parts[0]
}
break
}
}
return ip, account
}
package daemon
import (
"fmt"
"os"
"github.com/pidginhost/csm/internal/alert"
)
// replayPendingFindings runs the batch parked by the previous shutdown
// through the normal dispatch pipeline: auto-response, history, correlation
// and alerting. Called once the dispatcher has been released, so the replay
// sees the same ordering as any live batch.
func (d *Daemon) replayPendingFindings() {
if d.store == nil {
return
}
err := d.store.ReplayPendingFindings(func(pending []alert.Finding) {
fmt.Fprintf(os.Stderr, "[%s] Replaying %d finding(s) left queued by the previous shutdown\n", ts(), len(pending))
d.dispatchBatch(pending)
})
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] Cannot replay pending findings: %v\n", ts(), err)
}
}
package daemon
import (
"errors"
"fmt"
"net"
"net/url"
"os"
"path"
"strings"
"syscall"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/phpshield"
"golang.org/x/sys/unix"
)
const (
phpEventMaxBytes = 64 * 1024
phpEventArchiveMaxBytes = 10 * 1024 * 1024
phpEventSocketMode = 0o222
// phpShieldURIMaxBytes is the length the Shield cuts REQUEST_URI to before
// sending it.
phpShieldURIMaxBytes = 200
)
var (
phpEventsLogPath = phpshield.EventLogPath
phpEventsSocketPath = phpshield.EventSocketPath
phpEventRetryInterval = 2 * time.Second
phpShieldEventListener = listenPHPShieldEventSocket
errPHPEventArchiveFull = errors.New("PHP Shield event archive reached its size cap; rotate it before events can be archived again")
)
type phpEventPacketListener interface {
Read([]byte) (int, error)
SetReadDeadline(time.Time) error
Close() error
}
type phpEventUnixgramListener struct {
*net.UnixConn
path string
info os.FileInfo
}
func (l *phpEventUnixgramListener) Close() error {
err := l.UnixConn.Close()
if info, statErr := os.Lstat(l.path); statErr == nil && os.SameFile(l.info, info) {
_ = os.Remove(l.path)
}
return err
}
func listenPHPShieldEventSocket(path string) (phpEventPacketListener, error) {
addr := &net.UnixAddr{Name: path, Net: "unixgram"}
if info, err := os.Lstat(path); err == nil {
if info.Mode()&os.ModeSocket == 0 {
return nil, fmt.Errorf("PHP Shield event socket path is not a socket")
}
// Before this release the directory was tenant-writable. Do not trust an
// active socket a tenant may have planted there before the installer
// hardened it: otherwise Shield events could be delivered to that tenant.
stat, statOK := info.Sys().(*syscall.Stat_t)
foreignSocket := os.Geteuid() == 0 && statOK && stat.Uid != 0
if !foreignSocket {
probe, probeErr := net.DialUnix("unixgram", nil, addr)
if probeErr == nil {
_ = probe.Close()
return nil, fmt.Errorf("PHP Shield event socket is already active")
}
}
if removeErr := os.Remove(path); removeErr != nil {
return nil, fmt.Errorf("removing stale PHP Shield event socket: %w", removeErr)
}
} else if !os.IsNotExist(err) {
return nil, fmt.Errorf("checking PHP Shield event socket: %w", err)
}
conn, err := net.ListenUnixgram("unixgram", addr)
if err != nil {
return nil, fmt.Errorf("listening on PHP Shield event socket: %w", err)
}
if chmodErr := os.Chmod(path, phpEventSocketMode); chmodErr != nil {
_ = conn.Close()
_ = os.Remove(path)
return nil, fmt.Errorf("setting PHP Shield event socket mode: %w", chmodErr)
}
info, err := os.Lstat(path)
if err != nil {
_ = conn.Close()
_ = os.Remove(path)
return nil, fmt.Errorf("checking PHP Shield event socket after bind: %w", err)
}
return &phpEventUnixgramListener{UnixConn: conn, path: path, info: info}, nil
}
func processPHPShieldEventPacket(data []byte, archivePath string, _ *config.Config, alertCh chan<- alert.Finding) (bool, error) {
if len(data) == 0 || len(data) > phpEventMaxBytes {
return false, nil
}
line := strings.TrimSuffix(string(data), "\n")
if line == "" || strings.ContainsAny(line, "\r\n") {
return false, nil
}
finding, quiet := parsePHPShieldEventLine(line)
if finding == nil {
return false, nil
}
_, archiveErr := appendPHPShieldEventArchive(archivePath, line)
if !quiet {
if finding.Timestamp.IsZero() {
finding.Timestamp = time.Now()
}
if !alert.TryEnqueue(alertCh, *finding) {
fmt.Fprintln(os.Stderr, "Warning: alert channel full, dropping PHP Shield finding")
}
}
return true, archiveErr
}
func appendPHPShieldEventArchive(path, line string) (archived bool, retErr error) {
// #nosec G304 G302 -- fixed root-owned archive path; O_NOFOLLOW rejects
// symlink replacement and mode 0600 keeps tenants from reading/truncating it.
fd, err := unix.Open(path, unix.O_WRONLY|unix.O_APPEND|unix.O_CREAT|unix.O_NONBLOCK|unix.O_NOFOLLOW|unix.O_CLOEXEC, 0o600)
if err != nil {
return false, err
}
f := os.NewFile(uintptr(fd), path) // #nosec G115 -- unix.Open returned a non-negative fd
if f == nil {
_ = unix.Close(fd)
return false, fmt.Errorf("opening PHP Shield event archive")
}
defer func() {
if err := f.Close(); err != nil && retErr == nil {
retErr = err
}
}()
var stat unix.Stat_t
if err := unix.Fstat(fd, &stat); err != nil {
return false, err
}
if stat.Mode&unix.S_IFMT != unix.S_IFREG || stat.Nlink != 1 {
return false, fmt.Errorf("PHP Shield event archive is not a single-link regular file")
}
if err := f.Chmod(0o600); err != nil {
return false, err
}
record := line + "\n"
if stat.Size > phpEventArchiveMaxBytes-int64(len(record)) {
return false, errPHPEventArchiveFull
}
if _, err := f.WriteString(record); err != nil {
return false, err
}
return true, nil
}
func waitPHPShieldEventRetry(stopCh <-chan struct{}) bool {
timer := time.NewTimer(phpEventRetryInterval)
defer timer.Stop()
select {
case <-stopCh:
return false
case <-timer.C:
return true
}
}
func (d *Daemon) watchPHPShieldEvents() {
defer d.wg.Done()
buffer := make([]byte, phpEventMaxBytes+1)
lastError := ""
var listener phpEventPacketListener
defer func() {
if listener != nil {
_ = listener.Close()
}
}()
for {
select {
case <-d.stopCh:
return
default:
}
if listener == nil {
var err error
listener, err = phpShieldEventListener(phpEventsSocketPath)
if err != nil {
d.MarkWatcher("php_shield", false)
if err.Error() != lastError {
csmlog.Warn("PHP Shield event socket unavailable", "path", phpEventsSocketPath, "err", err)
lastError = err.Error()
}
if !waitPHPShieldEventRetry(d.stopCh) {
return
}
continue
}
lastError = ""
d.MarkWatcher("php_shield", true)
}
if err := listener.SetReadDeadline(time.Now().Add(phpEventRetryInterval)); err != nil {
_ = listener.Close()
listener = nil
d.MarkWatcher("php_shield", false)
if err.Error() != lastError {
csmlog.Warn("PHP Shield event socket deadline failed", "path", phpEventsSocketPath, "err", err)
lastError = err.Error()
}
if !waitPHPShieldEventRetry(d.stopCh) {
return
}
continue
}
n, err := listener.Read(buffer)
if err == nil {
processed, processErr := processPHPShieldEventPacket(buffer[:n], phpEventsLogPath, d.cfg, d.alertCh)
if processErr != nil {
if processErr.Error() != lastError {
csmlog.Warn("PHP Shield event archive unavailable", "path", phpEventsLogPath, "err", processErr)
lastError = processErr.Error()
}
} else if processed {
lastError = ""
}
continue
}
if netErr, ok := err.(net.Error); ok && netErr.Timeout() {
continue
}
_ = listener.Close()
listener = nil
d.MarkWatcher("php_shield", false)
if err.Error() != lastError {
csmlog.Warn("PHP Shield event socket read failed", "path", phpEventsSocketPath, "err", err)
lastError = err.Error()
}
if !waitPHPShieldEventRetry(d.stopCh) {
return
}
}
}
// parsePHPShieldLogLine wraps parsePHPShieldLine for the log watcher handler signature.
func parsePHPShieldLogLine(line string, _ *config.Config) []alert.Finding { //nolint:unparam
f := parsePHPShieldLine(line)
if f == nil {
return nil
}
return []alert.Finding{*f}
}
// parsePHPShieldLine parses a line from the PHP shield event log and returns
// a finding if it represents a security event.
//
// Format: [2026-03-25 10:00:00] EVENT_TYPE sha256=H ip=X script=Y uri=Z ua=A details=B
// Older Shields omit sha256; those observations cannot be quieted.
func parsePHPShieldLine(line string) *alert.Finding {
finding, quiet := parsePHPShieldEventLine(line)
if quiet {
return nil
}
return finding
}
// Keep the observation even when its alert is quiet: a route mismatch does
// not prove the requested file was absent, or that downstream CMS code is safe.
func parsePHPShieldEventLine(line string) (*alert.Finding, bool) {
line = strings.TrimSpace(line)
if line == "" || !strings.HasPrefix(line, "[") {
return nil, false
}
// Extract event type (first word after the timestamp bracket)
closeBracket := strings.Index(line, "]")
if closeBracket < 0 || closeBracket+2 >= len(line) {
return nil, false
}
rest := strings.TrimSpace(line[closeBracket+1:])
fields := strings.SplitN(rest, " ", 2)
if len(fields) < 1 {
return nil, false
}
eventType := fields[0]
// Extract key=value pairs. The URI and user agent are what identify the
// request: "/alfacgiapi/perl.alfa" from a "Mozlila" agent names the scanner,
// where the bare parameter name does not. Both were parsed and discarded.
var digest, ip, script, uri, ua, details string
if len(fields) > 1 {
kvPart := fields[1]
// Only the producer's first field is content evidence. A URI or user
// agent containing sha256= must not forge proof for a legacy event.
if first, rest, ok := strings.Cut(kvPart, " "); ok && strings.HasPrefix(first, "sha256=") {
digest = strings.TrimPrefix(first, "sha256=")
kvPart = rest
}
for _, kv := range splitKV(kvPart) {
switch kv[0] {
case "ip":
ip = kv[1]
case "script":
script = kv[1]
case "uri":
uri = kv[1]
case "ua":
ua = kv[1]
case "details":
details = kv[1]
}
}
}
context := phpShieldDetails(ip, uri, ua, details)
switch eventType {
case "BLOCK_PATH":
return &alert.Finding{
Severity: alert.Critical,
Check: "php_shield_block",
SourceIP: ip,
FilePath: script,
Message: fmt.Sprintf("PHP Shield blocked execution from dangerous path: %s", script),
Details: context,
}, false
case "WEBSHELL_PARAM":
// Observation, not a denial: for a document-root script the Shield never
// reaches its deny branch, so nothing was blocked. Every public site
// receives these daily, and rating them Critical buries the real blocks.
// A rewrite can reach a real shell, even one called index.php. Use the
// scanner's verified content cache and PHP's event-time fingerprint;
// reopening the path here could inspect a replacement file instead.
quiet := !phpShieldRequestReachedScript(script, uri) && checks.IsVerifiedCMSHash(digest)
return &alert.Finding{
Severity: alert.Warning,
Check: "php_shield_webshell",
SourceIP: ip,
FilePath: script,
Message: fmt.Sprintf("PHP Shield observed a webshell command parameter: %s", script),
Details: context,
}, quiet
case "BLOCK_WEBSHELL":
// Not gated on the request path: the Shield blocks on the executing
// script's own source, so a rewrite into a planted shell is still a
// stopped webshell.
return &alert.Finding{
Severity: alert.Critical,
Check: "php_shield_webshell",
SourceIP: ip,
FilePath: script,
Message: fmt.Sprintf("PHP Shield blocked a webshell signature: %s", script),
Details: context,
}, false
case "EVAL_FATAL":
return &alert.Finding{
Severity: alert.High,
Check: "php_shield_eval",
SourceIP: ip,
FilePath: script,
Message: fmt.Sprintf("PHP Shield detected eval() chain failure: %s", script),
Details: context,
}, false
}
return nil, false
}
// phpShieldRequestReachedScript reports whether the request URI names the
// script that executed, i.e. whether a command parameter was delivered to the
// script the client asked for.
//
// Scanners send cmd= to paths that do not exist. CMS rewrite rules (or a
// 404 handler) can answer with the site's front controller. This is only a
// routing hint: a rewritten request can also execute a real shell. The caller
// must establish content evidence before quieting an alert, and still archive
// the observation. The request names the executing script when a
// leading run of its path segments is a trailing part of the script path (this
// covers PATH_INFO such as /shell.php/extra), or when it names the directory
// holding the script, which is then served as its directory index. A request
// for "/" therefore still fires on the document-root index.php: the client did
// ask for that script, and a shell injected into index.php is reached exactly
// that way.
//
// A leading /~user segment is dropped as well, since that is how a userdir URL
// maps onto the account's document root.
//
// When the path cannot be judged (no URI, a form other than an origin or
// absolute path, a bad escape, or a path cut short by the Shield's truncation)
// the event is kept: silence has to be earned by a path that clearly names a
// different script.
func phpShieldRequestReachedScript(script, uri string) bool {
if !path.IsAbs(script) || uri == "" || uri == "-" {
return true
}
rawPath, _, hasQuery := strings.Cut(uri, "?")
if !strings.HasPrefix(rawPath, "/") {
parsed, err := url.ParseRequestURI(uri)
if err != nil || (parsed.Scheme != "http" && parsed.Scheme != "https") || parsed.Host == "" || parsed.User != nil {
return true
}
if strings.HasPrefix(parsed.Host, "[") && net.ParseIP(parsed.Hostname()) == nil {
return true
}
rawPath = parsed.EscapedPath()
if rawPath == "" {
rawPath = "/"
}
}
decoded, err := url.PathUnescape(rawPath)
if err != nil {
return true
}
for _, c := range decoded {
if c <= ' ' || c == 0x7f || c == '\\' || c == '#' {
return true
}
}
requested := path.Clean(decoded)
script = path.Clean(script)
candidates := []string{requested}
if first, rest, _ := strings.Cut(requested[1:], "/"); strings.HasPrefix(first, "~") {
candidates = append(candidates, "/"+rest)
}
for _, candidate := range candidates {
if phpShieldPathNamesScript(script, candidate) {
return true
}
}
return !hasQuery && len(uri) >= phpShieldURIMaxBytes
}
// phpShieldPathNamesScript applies the matching rule described on
// phpShieldRequestReachedScript to one cleaned request path.
func phpShieldPathNamesScript(script, requested string) bool {
if strings.HasSuffix(path.Dir(script), strings.TrimSuffix(requested, "/")) {
return true
}
for end := len(requested); end > 0; end = strings.LastIndexByte(requested[:end], '/') {
if strings.HasSuffix(script, requested[:end]) {
return true
}
}
return false
}
// phpShieldDetails renders the context an operator needs to judge a Shield
// event: who sent it, what they asked for, and what they claimed to be. Empty
// fields are omitted rather than printed as blanks.
func phpShieldDetails(ip, uri, ua, details string) string {
var b strings.Builder
for _, field := range [][2]string{
{"IP", ip},
{"URI", uri},
{"User-Agent", ua},
} {
if field[1] == "" {
continue
}
if b.Len() > 0 {
b.WriteString("\n")
}
fmt.Fprintf(&b, "%s: %s", field[0], field[1])
}
if details != "" {
if b.Len() > 0 {
b.WriteString("\n")
}
b.WriteString(details)
}
return b.String()
}
// splitKV splits "key1=val1 key2=val2" respecting values with spaces.
func splitKV(s string) [][2]string {
var result [][2]string
keys := []string{"ip=", "script=", "uri=", "ua=", "details="}
for i, key := range keys {
idx := strings.Index(s, key)
if idx < 0 {
continue
}
valStart := idx + len(key)
// Value ends at the next key or end of string
valEnd := len(s)
for _, nextKey := range keys[i+1:] {
nextIdx := strings.Index(s[valStart:], " "+nextKey)
if nextIdx >= 0 {
valEnd = valStart + nextIdx
break
}
}
val := s[valStart:valEnd]
// Preserve the URI byte count and ambiguous whitespace. Trimming a
// truncated path can make it look complete enough to suppress an alert.
if key != "uri=" {
val = strings.TrimSpace(val)
}
keyName := strings.TrimSuffix(key, "=")
result = append(result, [2]string{keyName, val})
}
return result
}
package daemon
import (
"bufio"
"bytes"
"encoding/gob"
"fmt"
"os"
"sort"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailspool"
"github.com/pidginhost/csm/internal/store"
)
var phpRelayEvaluatorRef atomic.Pointer[evaluator]
// SetPHPRelayEvaluator is called by daemon wiring once everything is up.
// nil disables php_relay path dispatch from parseEximLogLine.
func SetPHPRelayEvaluator(e *evaluator) {
phpRelayEvaluatorRef.Store(e)
}
// PHPRelayEvaluator returns the registered evaluator, or nil if none.
func PHPRelayEvaluator() *evaluator {
return phpRelayEvaluatorRef.Load()
}
// scriptKey = host(X-PHP-Script) + ":" + path(X-PHP-Script)
type scriptKey = string
// scriptEvent is one accepted outbound mail from a PHP script. The booleans
// are computed at acceptance time (Flow A); evaluatePaths counts them in
// the sliding window.
type scriptEvent struct {
At time.Time
MsgID string
Subject string
FromMismatch bool
AdditionalSignal bool // Reply-To external mismatch OR X-Mailer suspicious-not-safe
SourceIP string
}
// rejectionEvent records a remote-MTA policy-block rejection for Path 3.
// Stage 1 defines the type for completeness; Stage 2 wires it.
//
//nolint:unused // wired by Path 3 in Stage 2
type rejectionEvent struct {
At time.Time
MsgID string
MTACode string
Snippet string
}
const (
phpRelayMaxEventsPerScript = 256
phpRelayMaxRejectionsPerScript = 64
phpRelayMaxActiveMsgsPerScript = 4096
phpRelayScriptIdleHorizon = 25 * time.Hour //nolint:unused // consumed by Flow E in Task O2
)
// recipientTracker records the distinct envelope recipients seen in some
// window and whether that window is safe to gate on. A window is only safe
// when at least one recipient parse succeeded and no parse gap was seen in
// that same window, so a caller reading it can fail open. Callers hold their
// own lock; the tracker has none.
type recipientTracker struct {
recipients map[string]time.Time
lastAt time.Time
lastUnknownAt time.Time
}
// record folds one message's recipients in. An empty list marks a parse gap.
func (t *recipientTracker) record(recipients []string, at time.Time, max int) {
recorded := false
if len(recipients) > 0 {
if t.recipients == nil {
t.recipients = make(map[string]time.Time, 8)
}
for _, raw := range recipients {
r := normalizeRecipient(raw)
if r == "" {
continue
}
if _, exists := t.recipients[r]; !exists && len(t.recipients) >= max {
evictOldestRecipient(t.recipients)
}
if at.After(t.recipients[r]) {
t.recipients[r] = at
}
recorded = true
}
}
if recorded && at.After(t.lastAt) {
t.lastAt = at
}
if !recorded && at.After(t.lastUnknownAt) {
t.lastUnknownAt = at
}
}
// distinctSince returns how many distinct recipients were seen no earlier than
// since, and whether that answer is trustworthy. known=false means fail open.
func (t *recipientTracker) distinctSince(since time.Time) (count int, known bool) {
known = seenAtOrAfter(t.lastAt, since) && !seenAtOrAfter(t.lastUnknownAt, since)
for _, last := range t.recipients {
if !last.Before(since) {
count++
}
}
return count, known
}
// scriptState tracks one script's recent activity. All fields read or
// written through the embedded mutex.
type scriptState struct {
mu sync.Mutex
events []scriptEvent
rejections []rejectionEvent //nolint:unused // wired by Path 3 in Stage 2
firedAt map[string]time.Time
rcpts recipientTracker
activeMsgs map[string]time.Time // msgID -> acceptedAt
activeMsgsCapped bool
lastEvent time.Time
maxEvents int
maxRejections int //nolint:unused // wired by Path 3 in Stage 2
maxActiveMsgs int
}
func newScriptState() *scriptState {
return &scriptState{
firedAt: make(map[string]time.Time, 4),
activeMsgs: make(map[string]time.Time, 64),
maxEvents: phpRelayMaxEventsPerScript,
maxRejections: phpRelayMaxRejectionsPerScript,
maxActiveMsgs: phpRelayMaxActiveMsgsPerScript,
}
}
func (s *scriptState) append(e scriptEvent) {
s.mu.Lock()
defer s.mu.Unlock()
s.appendLocked(e)
}
func (s *scriptState) appendLocked(e scriptEvent) {
if len(s.events) >= s.maxEvents {
s.events = s.events[1:]
}
s.events = append(s.events, e)
if e.At.After(s.lastEvent) {
s.lastEvent = e.At
}
}
// appendMessage records an event and its recipient parse result as one state
// transition. This prevents a concurrent evaluator from counting the event
// while still trusting recipient data from only the preceding messages.
func (s *scriptState) appendMessage(e scriptEvent, recipients []string) {
s.mu.Lock()
defer s.mu.Unlock()
s.appendLocked(e)
s.rcpts.record(recipients, e.At, maxTrackedRecipients)
}
// distinctRecipientsSince reports the distinct recipients this script reached
// no earlier than since. known=false means callers must fail open.
func (s *scriptState) distinctRecipientsSince(since time.Time) (count int, known bool) {
s.mu.Lock()
defer s.mu.Unlock()
return s.rcpts.distinctSince(since)
}
// qualifyingCount returns the number of events whose At is at or after
// since AND for which match returns true.
func (s *scriptState) qualifyingCount(since time.Time, match func(scriptEvent) bool) int {
s.mu.Lock()
defer s.mu.Unlock()
n := 0
for _, e := range s.events {
if e.At.Before(since) {
continue
}
if match(e) {
n++
}
}
return n
}
// volumeCount returns events on or after since regardless of signal flags.
//
//nolint:unused // consumed by Path 2 evaluator in Task F2
func (s *scriptState) volumeCount(since time.Time) int {
return s.qualifyingCount(since, func(scriptEvent) bool { return true })
}
func (s *scriptState) relayHit(k scriptKey, since time.Time, match func(scriptEvent) bool) (alert.RelayScriptHit, bool) {
s.mu.Lock()
defer s.mu.Unlock()
hit := alert.RelayScriptHit{ScriptKey: string(k)}
var sampleAt time.Time
for _, e := range s.events {
if e.At.Before(since) || !match(e) {
continue
}
hit.Hits++
if e.At.After(hit.LastSeen) {
hit.LastSeen = e.At
}
if e.Subject != "" && (sampleAt.IsZero() || e.At.After(sampleAt)) {
hit.SampleSubject = truncateDaemon(e.Subject, phpRelayBreakdownSubjectMax)
sampleAt = e.At
}
}
if hit.Hits == 0 {
return alert.RelayScriptHit{}, false
}
return hit, true
}
func (s *scriptState) recordActive(msgID string, at time.Time) {
s.mu.Lock()
defer s.mu.Unlock()
if _, exists := s.activeMsgs[msgID]; !exists && len(s.activeMsgs) >= s.maxActiveMsgs {
// Drop oldest.
var oldestID string
var oldest time.Time
first := true
for id, t := range s.activeMsgs {
if first || t.Before(oldest) {
oldestID = id
oldest = t
first = false
}
}
if oldestID != "" {
delete(s.activeMsgs, oldestID)
}
s.activeMsgsCapped = true
}
s.activeMsgs[msgID] = at
}
func (s *scriptState) removeActive(msgID string) {
s.mu.Lock()
delete(s.activeMsgs, msgID)
s.mu.Unlock()
}
// snapshotActiveMsgs returns a copy of activeMsgs keys and the capped flag.
// The returned slice is independent of internal state; callers may mutate
// it without affecting the scriptState.
func (s *scriptState) snapshotActiveMsgs() ([]string, bool) {
s.mu.Lock()
defer s.mu.Unlock()
out := make([]string, 0, len(s.activeMsgs))
for id := range s.activeMsgs {
out = append(out, id)
}
return out, s.activeMsgsCapped
}
// pruneActiveMsgsOlderThan drops activeMsgs entries older than cutoff.
// Called by Flow E's GC to bound the lifetime of unreaped ids.
func (s *scriptState) pruneActiveMsgsOlderThan(cutoff time.Time) int {
s.mu.Lock()
defer s.mu.Unlock()
n := 0
for id, at := range s.activeMsgs {
if at.Before(cutoff) {
delete(s.activeMsgs, id)
n++
}
}
return n
}
// shouldFire returns true and updates firedAt if the cooldown for path has
// elapsed since the last fire.
//
//nolint:unused // consumed by Path 1/2/3 evaluators in Tasks F1/F2/G1
func (s *scriptState) shouldFire(path string, now time.Time, cooldown time.Duration) bool {
s.mu.Lock()
defer s.mu.Unlock()
if last, ok := s.firedAt[path]; ok && now.Sub(last) < cooldown {
return false
}
s.firedAt[path] = now
return true
}
// perScriptWindow keeps a scriptState per scriptKey.
type perScriptWindow struct {
states sync.Map // map[scriptKey]*scriptState
}
func newPerScriptWindow() *perScriptWindow { return &perScriptWindow{} }
func (w *perScriptWindow) getOrCreate(k scriptKey) *scriptState {
if v, ok := w.states.Load(k); ok {
return v.(*scriptState)
}
fresh := newScriptState()
actual, _ := w.states.LoadOrStore(k, fresh)
return actual.(*scriptState)
}
// SweepIdle drops scriptState entries whose lastEvent is before cutoff.
// Returns the number of entries dropped. Called by Flow E.
//
//nolint:unused // consumed by Flow E GC in Task O2
func (w *perScriptWindow) SweepIdle(cutoff time.Time) int {
n := 0
w.states.Range(func(k, v any) bool {
s := v.(*scriptState)
s.mu.Lock()
idle := s.lastEvent.Before(cutoff)
s.mu.Unlock()
if idle {
w.states.Delete(k)
n++
}
return true
})
return n
}
// PruneActiveMsgs iterates retained scriptStates and prunes activeMsgs
// entries older than cutoff. Used by Flow E so still-active scripts don't
// accumulate ghost activeMsgs whose corresponding messages have left the
// queue without a "Completed" log line being parsed. Returns the total
// number of activeMsgs entries removed across all scripts.
func (w *perScriptWindow) PruneActiveMsgs(cutoff time.Time) int {
n := 0
w.states.Range(func(_, v any) bool {
n += v.(*scriptState).pruneActiveMsgsOlderThan(cutoff)
return true
})
return n
}
// Snapshot returns the current per-script states (for csm phprelay status).
//
//nolint:unused // consumed by csm phprelay status command in Task M1
func (w *perScriptWindow) Snapshot() map[scriptKey]*scriptState {
out := make(map[scriptKey]*scriptState)
w.states.Range(func(k, v any) bool {
out[k.(scriptKey)] = v.(*scriptState)
return true
})
return out
}
type ipState struct {
mu sync.Mutex
scripts map[scriptKey]*ipScriptState
rcpts recipientTracker
lastEvent time.Time
}
type ipScriptState struct {
lastSeen time.Time
sampleAt time.Time
sampleSubject string
}
type perIPWindow struct {
states sync.Map // map[string]*ipState
capPerIP int
}
func newPerIPWindow(capPerIP int) *perIPWindow {
if capPerIP <= 0 {
capPerIP = 64
}
return &perIPWindow{capPerIP: capPerIP}
}
func (w *perIPWindow) append(ip string, k scriptKey, at time.Time, subject ...string) {
if ip == "" {
return
}
v, _ := w.states.LoadOrStore(ip, &ipState{scripts: make(map[scriptKey]*ipScriptState, 8)})
s := v.(*ipState)
s.mu.Lock()
defer s.mu.Unlock()
s.appendLocked(k, at, w.capPerIP, subject...)
}
func (s *ipState) appendLocked(k scriptKey, at time.Time, capPerIP int, subject ...string) {
if _, exists := s.scripts[k]; !exists && len(s.scripts) >= capPerIP {
var oldestK scriptKey
var oldest time.Time
first := true
for kk, ss := range s.scripts {
if first || ss.lastSeen.Before(oldest) {
oldestK = kk
oldest = ss.lastSeen
first = false
}
}
delete(s.scripts, oldestK)
}
ss := s.scripts[k]
if ss == nil {
ss = &ipScriptState{}
s.scripts[k] = ss
}
sampleSubject := ""
if len(subject) > 0 {
sampleSubject = truncateDaemon(subject[0], phpRelayBreakdownSubjectMax)
}
if at.After(ss.lastSeen) {
ss.lastSeen = at
}
if sampleSubject != "" && (ss.sampleAt.IsZero() || at.After(ss.sampleAt)) {
ss.sampleAt = at
ss.sampleSubject = sampleSubject
}
if at.After(s.lastEvent) {
s.lastEvent = at
}
}
// appendMessage records a script hit and its recipient parse result while
// holding the same IP-state lock, keeping the Path 4 gate fail-open when
// evaluation runs concurrently with spool processing.
func (w *perIPWindow) appendMessage(ip string, k scriptKey, at time.Time, subject string, recipients []string) {
if ip == "" {
return
}
v, _ := w.states.LoadOrStore(ip, &ipState{scripts: make(map[scriptKey]*ipScriptState, 8)})
s := v.(*ipState)
s.mu.Lock()
defer s.mu.Unlock()
s.appendLocked(k, at, w.capPerIP, subject)
s.rcpts.record(recipients, at, maxTrackedRecipients)
}
func (w *perIPWindow) distinctScriptsSince(ip string, since time.Time) int {
v, ok := w.states.Load(ip)
if !ok {
return 0
}
s := v.(*ipState)
s.mu.Lock()
defer s.mu.Unlock()
n := 0
for _, ss := range s.scripts {
if !ss.lastSeen.Before(since) {
n++
}
}
return n
}
// maxTrackedRecipients bounds each script or source-IP recipient set. A
// genuine high-diversity relay crosses the gate threshold before this cap, so
// the bound protects memory without hiding the current high diversity.
const maxTrackedRecipients = 256
// distinctRecipientsSince returns the number of distinct recipients seen from
// ip no earlier than since, and whether any recipient data was recorded within
// that window. known=false means recipients are unknown; callers must fail open.
func (w *perIPWindow) distinctRecipientsSince(ip string, since time.Time) (count int, known bool) {
v, ok := w.states.Load(ip)
if !ok {
return 0, false
}
s := v.(*ipState)
s.mu.Lock()
defer s.mu.Unlock()
return s.rcpts.distinctSince(since)
}
func seenAtOrAfter(t, since time.Time) bool {
return !t.IsZero() && !t.Before(since)
}
func evictOldestRecipient(m map[string]time.Time) {
var oldestK string
var oldest time.Time
first := true
for k, t := range m {
if first || t.Before(oldest) {
oldestK, oldest, first = k, t, false
}
}
if !first {
delete(m, oldestK)
}
}
func normalizeRecipient(r string) string {
r = strings.TrimSpace(r)
r = strings.TrimPrefix(r, "<")
r = strings.TrimSuffix(r, ">")
return strings.ToLower(strings.TrimSpace(r))
}
func (w *perIPWindow) relaySamplesSince(ip string, since time.Time) []alert.RelayScriptHit {
v, ok := w.states.Load(ip)
if !ok {
return nil
}
s := v.(*ipState)
s.mu.Lock()
defer s.mu.Unlock()
out := make([]alert.RelayScriptHit, 0, len(s.scripts))
for k, ss := range s.scripts {
if ss.lastSeen.Before(since) {
continue
}
sampleSubject := ""
if !ss.sampleAt.Before(since) {
sampleSubject = ss.sampleSubject
}
out = append(out, alert.RelayScriptHit{
ScriptKey: string(k),
Hits: 1,
LastSeen: ss.lastSeen,
SampleSubject: sampleSubject,
})
}
return out
}
func (w *perIPWindow) SweepIdle(cutoff time.Time) int {
n := 0
w.states.Range(func(k, v any) bool {
s := v.(*ipState)
s.mu.Lock()
idle := s.lastEvent.Before(cutoff)
s.mu.Unlock()
if idle {
w.states.Delete(k)
n++
}
return true
})
return n
}
type accountState struct {
mu sync.Mutex
events []time.Time
firedAt time.Time
lastEvent time.Time
maxEvents int
}
type perAccountWindow struct {
states sync.Map
cap int
}
func newPerAccountWindow(capPerAccount int) *perAccountWindow {
if capPerAccount <= 0 {
capPerAccount = 5000
}
return &perAccountWindow{cap: capPerAccount}
}
func (w *perAccountWindow) append(user string, at time.Time) {
if user == "" {
return
}
v, _ := w.states.LoadOrStore(user, &accountState{maxEvents: w.cap})
s := v.(*accountState)
s.mu.Lock()
defer s.mu.Unlock()
if len(s.events) >= s.maxEvents {
s.events = s.events[1:]
}
s.events = append(s.events, at)
if at.After(s.lastEvent) {
s.lastEvent = at
}
}
func (w *perAccountWindow) volumeSince(user string, since time.Time) int {
v, ok := w.states.Load(user)
if !ok {
return 0
}
s := v.(*accountState)
s.mu.Lock()
defer s.mu.Unlock()
n := 0
for _, t := range s.events {
if !t.Before(since) {
n++
}
}
return n
}
func (w *perAccountWindow) shouldFire(user string, now time.Time, cooldown time.Duration) bool {
v, ok := w.states.Load(user)
if !ok {
return false
}
s := v.(*accountState)
s.mu.Lock()
defer s.mu.Unlock()
if !s.firedAt.IsZero() && now.Sub(s.firedAt) < cooldown {
return false
}
s.firedAt = now
return true
}
func (w *perAccountWindow) SweepIdle(cutoff time.Time) int {
n := 0
w.states.Range(func(k, v any) bool {
s := v.(*accountState)
s.mu.Lock()
idle := s.lastEvent.Before(cutoff)
s.mu.Unlock()
if idle {
w.states.Delete(k)
n++
}
return true
})
return n
}
// signals is the per-event boolean fingerprint Flow A appends to
// scriptState. The numeric scriptKey lookup happens against the same
// emailspool helpers used elsewhere so subdomain handling is consistent.
type signals struct {
ScriptKey scriptKey
SourceIP string
FromMismatch bool
AdditionalSignal bool
XMailer string
}
// computeSignals resolves the per-event flags for an accepted message.
// authDomains is the cPanel user's authorised domain set (empty on
// resolver error -- caller treats that as "skip From-mismatch contribution"
// by passing isAuthDomainsKnown=false).
func computeSignals(h emailspool.Headers, authDomains map[string]struct{}, pol *emailspool.Policies) signals {
sk, sourceIP := parseXPHPScript(h.XPHPScript)
s := signals{
ScriptKey: sk,
SourceIP: sourceIP,
XMailer: h.XMailer,
}
if len(authDomains) > 0 {
fromDomain := emailspool.ExtractDomain(h.From)
if fromDomain != "" && !IsAuthorisedFromDomain(fromDomain, authDomains) {
s.FromMismatch = true
}
}
// Reply-To external mismatch contribution.
var replyToDomainMismatch bool
if h.ReplyTo != "" && h.From != "" {
rd := emailspool.ExtractDomain(h.ReplyTo)
fd := emailspool.ExtractDomain(h.From)
if rd != "" && fd != "" && rd != fd {
replyToDomainMismatch = true
}
}
// X-Mailer suspicious contribution.
var mailerSuspicious bool
if pol != nil {
if pol.MailerSuspicious(h.XMailer) && !pol.MailerSafe(h.XMailer) {
mailerSuspicious = true
}
}
s.AdditionalSignal = replyToDomainMismatch || mailerSuspicious
return s
}
// parseXPHPScript splits an X-PHP-Script header value into (scriptKey, sourceIP).
// Format: "<host>/<path> for <ip>". Returns ("", "") on parse failure.
func parseXPHPScript(v string) (scriptKey, string) {
v = strings.TrimSpace(v)
if v == "" {
return "", ""
}
forIdx := strings.LastIndex(v, " for ")
var url, ip string
if forIdx > 0 {
url = strings.TrimSpace(v[:forIdx])
ip = strings.TrimSpace(v[forIdx+5:])
} else {
url = v
}
// Strip any query string.
if q := strings.IndexByte(url, '?'); q > 0 {
url = url[:q]
}
slash := strings.IndexByte(url, '/')
if slash < 0 {
// Bare host with no path.
return scriptKey(url + ":/"), ip
}
host := url[:slash]
path := url[slash:]
return scriptKey(host + ":" + path), ip
}
const (
phpRelayPathCooldown = 30 * time.Minute
phpRelayBreakdownSubjectMax = 160
)
// evaluator combines windows + config + alerter in one object so the
// detector code path is callable from inotify watcher, retro scan, and
// startup spool walker without rebuilding the dependency graph each time.
type evaluator struct {
scripts *perScriptWindow
ips *perIPWindow
accounts *perAccountWindow
cfgFn func() *config.Config // live config; see liveConfigFn
metrics *phpRelayMetrics // optional; nil in unit tests
policies *emailspool.Policies
msgIndex *msgIDIndex // optional; nil in unit tests
effectiveAccountLimit int
accountLimitFn func(*config.Config) int
}
func newEvaluator(s *perScriptWindow, i *perIPWindow, a *perAccountWindow, cfg *config.Config, m *phpRelayMetrics) *evaluator {
return &evaluator{scripts: s, ips: i, accounts: a, cfgFn: liveConfigFn(cfg), metrics: m}
}
// config returns the live config so a reload that retunes or disables the
// relay detector reaches every evaluation; cfg was a startup snapshot before.
func (e *evaluator) config() *config.Config { return e.cfgFn() }
// SetPolicies is called by daemon wiring once the policies file has loaded.
func (e *evaluator) SetPolicies(p *emailspool.Policies) { e.policies = p }
// evaluatePaths inspects the script's window state (and IP window) and
// returns the set of findings that fire at this moment. Cooldowns prevent
// duplicate emissions per (script, path).
func (e *evaluator) evaluatePaths(k scriptKey, sourceIP, cpuser string, now time.Time) []alert.Finding {
cfg := e.config()
relayCfg := cfg.EmailProtection.PHPRelay
if !relayCfg.Enabled {
return nil
}
var findings []alert.Finding
s := e.scripts.getOrCreate(k)
// Path 1: sustained qualifying events.
win := time.Duration(relayCfg.RateWindowMin) * time.Minute
qualifying := s.qualifyingCount(now.Add(-win), func(ev scriptEvent) bool {
return ev.FromMismatch && ev.AdditionalSignal
})
if qualifying >= relayCfg.HeaderScoreVolumeMin {
if s.shouldFire("header", now, phpRelayPathCooldown) {
f := e.makeFinding(k, "header", sourceIP, cpuser, s, fmtHeaderMessage(qualifying, win), now)
f.RelayTotal = qualifying
f.RelayBreakdown = e.scriptRelayBreakdown(k, now.Add(-win), func(ev scriptEvent) bool {
return ev.FromMismatch && ev.AdditionalSignal
})
if e.metrics != nil {
e.metrics.Findings.With("header").Inc()
}
findings = append(findings, f)
}
}
// Path 2: absolute volume per script in the last 60 min.
absVol := s.volumeCount(now.Add(-60 * time.Minute))
if absVol >= relayCfg.AbsoluteVolumePerHour &&
!e.scriptIsLowDiversityNotification(s, now.Add(-60*time.Minute), relayCfg.FanoutDistinctRecipients) {
if s.shouldFire("volume", now, phpRelayPathCooldown) {
f := e.makeFinding(k, "volume", sourceIP, cpuser, s,
fmt.Sprintf("Path 2: %d outbound mails from one script in last 60 min", absVol), now)
f.RelayTotal = absVol
f.RelayBreakdown = e.scriptRelayBreakdown(k, now.Add(-60*time.Minute), func(scriptEvent) bool {
return true
})
if e.metrics != nil {
e.metrics.Findings.With("volume").Inc()
}
findings = append(findings, f)
}
}
// Path 4: HTTP-IP fanout. Skipped silently for proxy IPs.
if sourceIP != "" {
if e.policies == nil || !e.policies.IsProxyIP(sourceIP) {
fwin := time.Duration(relayCfg.FanoutWindowMin) * time.Minute
distinct := e.ips.distinctScriptsSince(sourceIP, now.Add(-fwin))
if distinct >= relayCfg.FanoutDistinctScripts &&
!e.fanoutIsLowDiversityNotification(sourceIP, now.Add(-fwin), relayCfg.FanoutDistinctRecipients) {
if s.shouldFire("fanout", now, phpRelayPathCooldown) {
f := e.makeFinding(k, "fanout", sourceIP, cpuser, s,
fmt.Sprintf("Path 4: HTTP source IP %s triggered %d distinct scripts in last %s", sourceIP, distinct, fwin), now)
f.RelayTotal = distinct
f.RelayBreakdown = e.fanoutRelayBreakdown(sourceIP, now.Add(-fwin))
if e.metrics != nil {
e.metrics.Findings.With("fanout").Inc()
}
findings = append(findings, f)
}
}
}
}
return findings
}
// scriptIsLowDiversityNotification reports whether a script's hourly volume is
// notification mail rather than relay abuse. A security or e-commerce plugin
// can legitimately emit hundreds of mails an hour to the same one or two
// admin addresses; relay abuse reaches many distinct victims. Unknown or
// partially parsed recipients leave the gate failing open.
func (e *evaluator) scriptIsLowDiversityNotification(s *scriptState, since time.Time, minRcpt int) bool {
if minRcpt <= 0 || s == nil {
return false
}
count, known := s.distinctRecipientsSince(since)
return known && count < minRcpt
}
// fanoutIsLowDiversityNotification reports whether a script fanout from this
// source IP looks like WordPress notification mail (comment moderation, contact
// forms): many distinct scripts but a small fixed recipient set. It returns
// true only when recipient data is known AND the distinct recipient count is
// below the configured minimum, so Path 4 still fires whenever recipients are
// diverse (real relay) or unknown (recipient parsing gap -- fail open). A
// non-positive threshold disables the gate and preserves the original behavior.
func (e *evaluator) fanoutIsLowDiversityNotification(sourceIP string, since time.Time, minRcpt int) bool {
if minRcpt <= 0 || e.ips == nil {
return false
}
count, known := e.ips.distinctRecipientsSince(sourceIP, since)
return known && count < minRcpt
}
// makeFinding builds a Critical Finding for the given path.
func (e *evaluator) makeFinding(k scriptKey, path, sourceIP, cpuser string, s *scriptState, message string, now time.Time) alert.Finding {
msgIDs, _ := s.snapshotActiveMsgs()
// Cap the sample shown in the finding so the alert payload stays bounded;
// AutoFreezePHPRelayQueue takes its own complete snapshot.
if len(msgIDs) > 10 {
msgIDs = msgIDs[:10]
}
return alert.Finding{
Severity: alert.Critical,
Check: "email_php_relay_abuse",
Path: path,
Message: message,
ScriptKey: string(k),
SourceIP: sourceIP,
CPUser: cpuser,
TenantID: checks.HostingAccountForUser(cpuser),
MsgIDs: msgIDs,
Timestamp: now,
}
}
func (e *evaluator) scriptRelayBreakdown(k scriptKey, since time.Time, match func(scriptEvent) bool) []alert.RelayScriptHit {
hit, ok := e.scripts.getOrCreate(k).relayHit(k, since, match)
if !ok {
return nil
}
return []alert.RelayScriptHit{hit}
}
func (e *evaluator) fanoutRelayBreakdown(sourceIP string, since time.Time) []alert.RelayScriptHit {
if e.ips == nil || sourceIP == "" {
return nil
}
samples := e.ips.relaySamplesSince(sourceIP, since)
if len(samples) == 0 {
return nil
}
out := make([]alert.RelayScriptHit, 0, len(samples))
for _, sample := range samples {
hit := sample
if e.scripts != nil {
if v, ok := e.scripts.states.Load(scriptKey(sample.ScriptKey)); ok {
if counted, ok := v.(*scriptState).relayHit(scriptKey(sample.ScriptKey), since, func(ev scriptEvent) bool {
return ev.SourceIP == sourceIP
}); ok {
hit = counted
}
}
}
out = append(out, hit)
}
sortRelayScriptHits(out)
return out
}
func sortRelayScriptHits(out []alert.RelayScriptHit) {
sort.Slice(out, func(i, j int) bool {
if out[i].Hits != out[j].Hits {
return out[i].Hits > out[j].Hits
}
if !out[i].LastSeen.Equal(out[j].LastSeen) {
return out[i].LastSeen.After(out[j].LastSeen)
}
return out[i].ScriptKey < out[j].ScriptKey
})
}
func fmtHeaderMessage(qualifying int, win time.Duration) string {
return fmt.Sprintf("Path 1: %d qualifying outbound mails (From-mismatch AND suspicious header) in last %s", qualifying, win)
}
type cpanelLimitStatus int
const (
cpanelLimitOK cpanelLimitStatus = iota
cpanelLimitMissing
cpanelLimitUnparsable
cpanelLimitDisabled
)
// readCpanelHourlyLimit returns (parsed-value, status). The key in
// /var/cpanel/cpanel.config is `maxemailsperhour` (no underscores, matches
// internal/checks/hardening_audit.go usage).
//
// OK -> integer > 0; the cap is in force.
// Missing -> file or key absent; caller assumes default 100 + Warning.
// Unparsable -> key present but not a number; caller assumes default 100 + Warning.
// Disabled -> key present and == 0; cpanel hourly limit explicitly off.
func readCpanelHourlyLimit(path string) (int, cpanelLimitStatus) {
// #nosec G304 -- path is the cPanel config path (default /var/cpanel/cpanel.config); operator-controlled, root-owned.
f, err := os.Open(path)
if err != nil {
return 0, cpanelLimitMissing
}
defer f.Close()
sc := bufio.NewScanner(f)
for sc.Scan() {
line := sc.Text()
eq := strings.IndexByte(line, '=')
if eq < 0 {
continue
}
if strings.TrimSpace(line[:eq]) != "maxemailsperhour" {
continue
}
val := strings.TrimSpace(line[eq+1:])
n, err := strconv.Atoi(val)
if err != nil {
return 0, cpanelLimitUnparsable
}
if n == 0 {
return 0, cpanelLimitDisabled
}
return n, cpanelLimitOK
}
return 0, cpanelLimitMissing
}
// deriveEffectiveAccountLimit implements spec section 6.1's three-step
// derivation. Returns (effective, enabled, cappedFromOperator).
// - cpanelLimit/status come from readCpanelHourlyLimit.
// - missing/unparsable callers should also emit a Warning startup finding.
// - returned enabled=false means Path 2b should not run this session.
func deriveEffectiveAccountLimit(cfg *config.Config, cpanelLimit int, status cpanelLimitStatus) (effective int, enabled bool, capped bool) {
op := cfg.EmailProtection.PHPRelay.AccountVolumePerHour
// Step 1: classify the cPanel limit.
var assumed int
var known bool
switch status {
case cpanelLimitOK:
assumed = cpanelLimit
known = true
case cpanelLimitMissing, cpanelLimitUnparsable:
// Caller emits Warning; we use the cPanel default 100.
assumed = 100
known = true
case cpanelLimitDisabled:
known = false
}
// Step 2: derive effective.
if known {
cap := assumed * 95 / 100
if cap < 1 {
cap = 1
}
if op == 0 {
target := assumed * 60 / 100
if target < 20 {
target = 20
}
if target > 60 {
target = 60
}
effective = target
if effective > cap {
effective = cap
}
} else {
effective = op
if effective > cap {
effective = cap
capped = true
}
}
if effective <= 0 {
return 0, false, false
}
return effective, true, capped
}
// Cpanel limit explicitly disabled.
if op > 0 {
return op, true, false
}
return 0, false, false
}
const (
phpRelayAccountWindowDur = 60 * time.Minute
phpRelayAccountFireCooldown = 30 * time.Minute
)
// SetEffectiveAccountLimit is called by daemon wiring after derivation.
// Tests may also call it directly. A non-positive value disables Path 2b.
func (e *evaluator) SetEffectiveAccountLimit(n int) {
e.effectiveAccountLimit = n
e.accountLimitFn = nil
}
// SetAccountLimitSource keeps the cPanel limit fixed while deriving the
// operator-controlled threshold from the live config on each evaluation.
func (e *evaluator) SetAccountLimitSource(cpanelLimit int, status cpanelLimitStatus) {
e.accountLimitFn = func(cfg *config.Config) int {
effective, enabled, _ := deriveEffectiveAccountLimit(cfg, cpanelLimit, status)
if !enabled {
return 0
}
return effective
}
e.effectiveAccountLimit = e.accountLimit(e.config())
}
func (e *evaluator) accountLimit(cfg *config.Config) int {
if e.accountLimitFn != nil {
return e.accountLimitFn(cfg)
}
return e.effectiveAccountLimit
}
// SetMsgIndex wires the msgID->script index so exim queue-completion log lines
// can reap the corresponding activeMsgs entry.
func (e *evaluator) SetMsgIndex(idx *msgIDIndex) { e.msgIndex = idx }
// reapCompletedMsg drops an activeMsgs entry when exim logs that its message
// completed delivery and left the queue, so a delivered message is not later
// frozen as if still queued nor counted against the freeze rate-limit budget.
func (e *evaluator) reapCompletedMsg(line string) {
if e.msgIndex == nil || e.scripts == nil {
return
}
msgID, ok := eximCompletedMsgID(line)
if !ok {
return
}
entry, found := e.msgIndex.Get(msgID)
if !found {
return
}
if v, loaded := e.scripts.states.Load(scriptKey(entry.ScriptKey)); loaded {
v.(*scriptState).removeActive(msgID)
}
}
// eximCompletedMsgID returns the message id from an exim queue-completion line
// ("YYYY-MM-DD HH:MM:SS <msgid> Completed[ ...]") and true, or ("", false) for
// any other line. Requiring the "Completed" verb to immediately follow the id
// keeps an attacker-controlled Subject that merely contains the word from
// reaping a live message.
func eximCompletedMsgID(line string) (string, bool) {
const tsLen = 19 // "2006-01-02 15:04:05"
if len(line) < tsLen+2 {
return "", false
}
rest := line[tsLen+1:] // skip timestamp and the following space
end := strings.IndexByte(rest, ' ')
if end < 0 {
return "", false
}
id := rest[:end]
if !msgIDPattern.MatchString(id) {
return "", false
}
verb := rest[end+1:]
if verb != "Completed" && !strings.HasPrefix(verb, "Completed ") {
return "", false
}
return id, true
}
// parsePHPRelayAccountVolume processes one outbound `<= ` exim_mainlog line
// seen live. Returns zero or one finding (per cooldown).
func (e *evaluator) parsePHPRelayAccountVolume(line string, now time.Time) []alert.Finding {
// Reap on the live path only: queue-completion lines free the msgID from
// activeMsgs. The retro history replay never sees the live index.
e.reapCompletedMsg(line)
return e.parsePHPRelayAccountVolumeAt(line, now, now)
}
// parsePHPRelayAccountVolumeAt records the account event at eventTime and
// evaluates the volume window relative to now. The live watcher passes
// eventTime == now; the startup history replay passes the line's real exim
// timestamp so days-old log entries are not all stamped "now" and miscounted as
// a single last-hour burst (a false Path 2b Critical on every restart).
func (e *evaluator) parsePHPRelayAccountVolumeAt(line string, eventTime, now time.Time) []alert.Finding {
cfg := e.config()
limit := e.accountLimit(cfg)
if !cfg.EmailProtection.PHPRelay.Enabled || limit <= 0 || e.accounts == nil {
return nil
}
if !strings.Contains(line, " <= ") {
return nil
}
if !strings.Contains(line, " B=redirect_resolver") {
return nil
}
user := extractUField(line)
if user == "" {
return nil
}
e.accounts.append(user, eventTime)
volume := e.accounts.volumeSince(user, now.Add(-phpRelayAccountWindowDur))
if volume < limit {
return nil
}
if !e.accounts.shouldFire(user, now, phpRelayAccountFireCooldown) {
return nil
}
if e.metrics != nil {
e.metrics.Findings.With("volume_account").Inc()
}
return []alert.Finding{{
Severity: alert.Critical,
Check: "email_php_relay_abuse",
Path: "volume_account",
Message: fmt.Sprintf("Path 2b: account %s sent >= %d outbound mails in last hour", user, limit),
CPUser: user,
TenantID: checks.HostingAccountForUser(user),
RelayTotal: volume,
Timestamp: now,
}}
}
// extractUField returns the cpuser from "U=<name>" in an exim log line.
// Returns "" if absent.
func extractUField(line string) string {
idx := strings.Index(line, " U=")
if idx < 0 {
return ""
}
rest := line[idx+3:]
end := len(rest)
for i := 0; i < len(rest); i++ {
c := rest[i]
if c == ' ' || c == '\t' {
end = i
break
}
}
return rest[:end]
}
// ignoreEntry records an operator-issued ignore on a script. A zero
// ExpiresAt means "never expires"; otherwise Has/List/SweepExpired drop
// the entry once now > ExpiresAt.
type ignoreEntry struct {
ScriptKey string
AddedAt time.Time
ExpiresAt time.Time
AddedBy string
Reason string
}
// ignoreList is the in-memory operator allowlist for php_relay scripts.
// L2 (bbolt persistence) wraps this with --persist semantics; Flow E
// (O2) calls SweepExpired periodically.
type ignoreList struct {
mu sync.Mutex
entries map[string]ignoreEntry
db *store.DB
}
func newIgnoreList() *ignoreList {
return &ignoreList{entries: make(map[string]ignoreEntry)}
}
func (il *ignoreList) Add(k scriptKey, expiresAt time.Time, by, reason string) {
il.mu.Lock()
defer il.mu.Unlock()
il.entries[string(k)] = ignoreEntry{
ScriptKey: string(k),
AddedAt: time.Now(),
ExpiresAt: expiresAt,
AddedBy: by,
Reason: reason,
}
}
func (il *ignoreList) Remove(k scriptKey) {
il.mu.Lock()
delete(il.entries, string(k))
il.mu.Unlock()
}
func (il *ignoreList) Has(k scriptKey) bool {
il.mu.Lock()
defer il.mu.Unlock()
e, ok := il.entries[string(k)]
if !ok {
return false
}
if !e.ExpiresAt.IsZero() && time.Now().After(e.ExpiresAt) {
delete(il.entries, string(k))
return false
}
return true
}
func (il *ignoreList) List() []ignoreEntry {
il.mu.Lock()
defer il.mu.Unlock()
out := make([]ignoreEntry, 0, len(il.entries))
now := time.Now()
for k, e := range il.entries {
if !e.ExpiresAt.IsZero() && now.After(e.ExpiresAt) {
delete(il.entries, k)
continue
}
out = append(out, e)
}
return out
}
// SweepExpired drops expired entries. Called by Flow E ticker.
//
//nolint:unused // wired in O2 Flow E ticker
func (il *ignoreList) SweepExpired(now time.Time) int {
il.mu.Lock()
defer il.mu.Unlock()
n := 0
for k, e := range il.entries {
if !e.ExpiresAt.IsZero() && now.After(e.ExpiresAt) {
delete(il.entries, k)
n++
}
}
return n
}
const ignoreBucket = "phprelay:ignore"
func (il *ignoreList) SetStore(db *store.DB) { il.db = db }
func (il *ignoreList) AddPersist(k scriptKey, expiresAt time.Time, by, reason string) error {
il.Add(k, expiresAt, by, reason)
if il.db == nil {
return nil
}
var buf bytes.Buffer
if err := gob.NewEncoder(&buf).Encode(ignoreEntry{
ScriptKey: string(k), AddedAt: time.Now(),
ExpiresAt: expiresAt, AddedBy: by, Reason: reason,
}); err != nil {
return err
}
return il.db.PHPRelayPut(ignoreBucket, string(k), buf.Bytes())
}
//nolint:unused // wired in M2/O2
func (il *ignoreList) RemovePersist(k scriptKey) error {
il.Remove(k)
if il.db == nil {
return nil
}
return il.db.PHPRelayDelete(ignoreBucket, string(k))
}
// Restore re-populates the in-memory list from bbolt at daemon start.
// Corrupt rows are skipped silently; expired rows are skipped (the bbolt
// row stays put until SweepBolt prunes it on the next Flow E tick).
func (il *ignoreList) Restore() error {
if il.db == nil {
return nil
}
rows, err := il.db.PHPRelayList(ignoreBucket)
if err != nil {
return err
}
now := time.Now()
for _, raw := range rows {
var e ignoreEntry
if err := gob.NewDecoder(bytes.NewReader(raw)).Decode(&e); err != nil {
continue
}
if !e.ExpiresAt.IsZero() && now.After(e.ExpiresAt) {
continue
}
il.Add(scriptKey(e.ScriptKey), e.ExpiresAt, e.AddedBy, e.Reason)
}
return nil
}
// SweepBolt drops expired bbolt entries on the Flow E ticker. Corrupt
// rows are also dropped so the bucket stays healthy.
//
//nolint:unused // wired in M2/O2
func (il *ignoreList) SweepBolt(now time.Time) (int, error) {
if il.db == nil {
return 0, nil
}
return il.db.PHPRelaySweep(ignoreBucket, func(_, value []byte) bool {
var e ignoreEntry
if err := gob.NewDecoder(bytes.NewReader(value)).Decode(&e); err != nil {
return true // drop corrupt rows
}
return !e.ExpiresAt.IsZero() && now.After(e.ExpiresAt)
})
}
package daemon
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"regexp"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailspool"
"github.com/pidginhost/csm/internal/systemdrun"
)
// msgIDPattern guards exim -Mf invocations against header-injected garbage
// that slipped past parseHeaders. Exim msgIDs are <= 23 chars in practice
// but we accept up to 32 to allow for future format changes; the lower
// bound of 16 rules out any short string an attacker could try to slip in.
var msgIDPattern = regexp.MustCompile(`^[A-Za-z0-9-]{16,32}$`)
// eximBinary is resolved at module init via exec.LookPath. Empty means
// auto-action is permanently disabled (a Warning finding is emitted at
// startup -- see Phase O).
//
//nolint:unused // populated in K5 by AutoFreezePHPRelayQueue init via exec.LookPath
var eximBinary string
// actionRateLimiter is a sliding-window counter of exim -M* invocations.
// Per spec: at most maxPerMinute actions in any rolling 60s window.
type actionRateLimiter struct {
mu sync.Mutex
maxPerMin int
bucket int
refilledAt time.Time
now func() time.Time
}
func newActionRateLimiter(maxPerMin int) *actionRateLimiter {
return &actionRateLimiter{
maxPerMin: maxPerMin,
bucket: maxPerMin,
now: time.Now,
}
}
// consumeUpTo grants up to n tokens and returns how many were available and
// consumed (0..n). Partial grants let a real outbreak freeze as many messages
// as the per-minute budget allows instead of freezing none the moment the batch
// size exceeds the cap.
func (rl *actionRateLimiter) consumeUpTo(n int) int {
rl.mu.Lock()
defer rl.mu.Unlock()
return rl.consumeUpToLocked(n, rl.maxPerMin)
}
func (rl *actionRateLimiter) consumeUpToLimit(n, maxPerMin int) int {
rl.mu.Lock()
defer rl.mu.Unlock()
return rl.consumeUpToLocked(n, maxPerMin)
}
func (rl *actionRateLimiter) consumeUpToLocked(n, maxPerMin int) int {
if maxPerMin <= 0 {
maxPerMin = 60
}
now := rl.now()
if maxPerMin != rl.maxPerMin {
used := rl.maxPerMin - rl.bucket
rl.maxPerMin = maxPerMin
rl.bucket = maxPerMin - used
if rl.bucket < 0 {
rl.bucket = 0
}
}
if rl.refilledAt.IsZero() || now.Sub(rl.refilledAt) >= time.Minute {
rl.bucket = rl.maxPerMin
rl.refilledAt = now
}
if n <= 0 {
return 0
}
grant := n
if grant > rl.bucket {
grant = rl.bucket
}
rl.bucket -= grant
return grant
}
// freezeErrIsAlreadyGone matches the Exim stderr fragments emitted when
// the message has already left the queue between snapshot and freeze.
// Those are not action failures -- they are normal queue churn.
func freezeErrIsAlreadyGone(stderr string) bool {
s := strings.ToLower(stderr)
return strings.Contains(s, "message not found") ||
strings.Contains(s, "spool file not found") ||
strings.Contains(s, "no such message")
}
// spoolScanMatchingScript walks every -H file under spoolRoot, parses
// headers, and returns msgIDs whose X-PHP-Script host:path matches
// scriptKey. Used by AutoFreezePHPRelayQueue when activeMsgs was capped
// or when a late reputation finding has no in-memory activeMsgs left.
//
// Handles BOTH spool layouts:
//
// 1. Split (cPanel default + Exim's split_spool_directory=true): each
// msgID-H lives under spoolRoot/<hash-char>/. We must descend one
// level into each subdir.
// 2. Unsplit (some self-hosted Exim builds, smaller cPanel installs
// where the operator has disabled split_spool_directory): -H files
// live directly in spoolRoot.
//
// The spec section 5.8 explicitly requires both layouts. We probe each
// entry: if it's a regular -H file at the root, scan it; if it's a
// directory, descend. No probing of /etc/exim or spool config -- the
// filesystem layout is the source of truth.
func spoolScanMatchingScript(spoolRoot string, k scriptKey) []string {
var out []string
// #nosec G304 -- spoolRoot is operator-configured / hardcoded to cPanel default.
entries, err := os.ReadDir(spoolRoot)
if err != nil {
return nil
}
inspect := func(full string, name string) {
if !strings.HasSuffix(name, "-H") {
return
}
h, err := emailspool.ParseHeaders(full)
if err != nil || h.XPHPScript == "" {
return
}
sk, _ := parseXPHPScript(h.XPHPScript)
if sk != k {
return
}
id := strings.TrimSuffix(name, "-H")
if msgIDPattern.MatchString(id) {
out = append(out, id)
}
}
for _, e := range entries {
full := filepath.Join(spoolRoot, e.Name())
if e.IsDir() {
// Split layout: descend one level.
// #nosec G304 -- spoolRoot is operator-configured / hardcoded to cPanel default.
files, err := os.ReadDir(full)
if err != nil {
continue
}
for _, f := range files {
inspect(filepath.Join(full, f.Name()), f.Name())
}
continue
}
// Unsplit layout: -H files at the root of spoolRoot.
inspect(full, e.Name())
}
return out
}
// runner abstracts exec.CommandContext so tests can inject a stub.
type runner interface {
Run(ctx context.Context, bin string, args []string) (stderr string, err error)
}
type defaultRunner struct{}
var (
defaultRunnerLookPath = exec.LookPath
defaultRunnerCommand = func(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- name is either the resolved systemd-run path or the
// resolved Exim binary; args are fixed flags plus validated message IDs.
cmd := exec.CommandContext(ctx, name, args...)
// Only stderr is returned: the freeze audit log records this under
// "stderr" and operators read it as the failure reason. --pipe connects
// the transient unit's stderr to ours, so wrapping does not change what
// is captured.
var stderr bytes.Buffer
cmd.Stderr = &stderr
err := cmd.Run()
return stderr.Bytes(), err
}
)
func (defaultRunner) Run(ctx context.Context, bin string, args []string) (string, error) {
opt := systemdrun.Options{Pipe: true}
if deadline, ok := ctx.Deadline(); ok {
opt.RuntimeMax = time.Until(deadline)
}
out, err := systemdrun.Run(ctx, defaultRunnerLookPath, defaultRunnerCommand, opt, bin, args...)
return string(out), err
}
// auditEntry is the per-action record written by the auditor. K6 will add a
// JSONL serialiser; K5 only constructs the in-memory shape.
type auditEntry struct {
Ts time.Time
MsgID string
ScriptKey string
Path string
DryRun bool
Exit int
Stderr string
Action string // "freeze" | "thaw" | "freeze_dry_run"
}
type auditor interface {
Write(e auditEntry)
}
// autoFreezer holds the wiring needed to translate findings into exim -Mf
// invocations. Constructed once at daemon start; Apply is invoked per
// post-emit AutoResponse pass.
//
// dryRunFn returns the EFFECTIVE dry-run state from the same config snapshot
// Apply uses for enablement and rate limits. The CLI's runtime override,
// bbolt override and csm.yaml fallback are resolved by the controller so
// `csm phprelay dry-run on|off|reset` changes freeze behaviour immediately.
//
//nolint:unused // wired in O2 by daemon controller
type autoFreezer struct {
scripts *perScriptWindow
cfgFn func() *config.Config // live config; see liveConfigFn
spoolRoot string
eximBin string
runner runner
auditor auditor
rateLim *actionRateLimiter
metrics *phpRelayMetrics
dryRunFn func(*config.Config) bool
}
// config returns the live config so a reload that disables auto-response or
// the freeze action stops the next Apply; cfg was a startup snapshot before.
func (a *autoFreezer) config() *config.Config { return a.cfgFn() }
//nolint:unused // wired in O2 by daemon controller
func newAutoFreezer(scripts *perScriptWindow, cfg *config.Config, spoolRoot, eximBin string, r runner, a auditor, m *phpRelayMetrics, dryRunFn func(*config.Config) bool) *autoFreezer {
if r == nil {
r = defaultRunner{}
}
rl := newActionRateLimiter(cfg.AutoResponse.PHPRelay.MaxActionsPerMinute)
if cfg.AutoResponse.PHPRelay.MaxActionsPerMinute <= 0 {
rl = newActionRateLimiter(60)
}
if dryRunFn == nil {
// Defensive default: if no resolver wired, fall back to the safe
// YAML-level dry-run state (PHPRelayDryRunEnabled defaults to TRUE).
dryRunFn = func(liveCfg *config.Config) bool { return liveCfg.PHPRelayDryRunEnabled() }
}
return &autoFreezer{
scripts: scripts, cfgFn: liveConfigFn(cfg), spoolRoot: spoolRoot, eximBin: eximBin,
runner: r, auditor: a, rateLim: rl, metrics: m, dryRunFn: dryRunFn,
}
}
// Apply iterates findings, snapshots each script's activeMsgs, optionally
// extends with a spool-scan fallback, and freezes via exim -Mf. Returns
// any new findings produced (Warning/Critical for action outcomes). Pure
// from the perspective of finding emission -- caller forwards them to the
// alert pipeline.
//
//nolint:unused // wired in O2 by daemon controller
func (a *autoFreezer) Apply(findings []alert.Finding) []alert.Finding {
var emitted []alert.Finding
cfg := a.config()
if !cfg.AutoResponse.Enabled || !cfg.PHPRelayFreezeEnabled() {
return nil
}
if a.eximBin == "" {
return nil
}
dryRun := a.dryRunFn(cfg)
for _, f := range findings {
if f.Check != "email_php_relay_abuse" {
continue
}
if !canActOnPath(f.Path) {
emitted = append(emitted, alert.Finding{
Severity: alert.Warning,
Check: "email_php_relay_action_skipped",
Path: f.Path,
Message: fmt.Sprintf("AutoFreeze skipped: path %q has no scriptKey", f.Path),
Timestamp: time.Now(),
})
continue
}
s := a.scripts.getOrCreate(scriptKey(f.ScriptKey))
ids, capped := s.snapshotActiveMsgs()
if capped || (len(ids) == 0 && f.Path == "reputation") {
if a.metrics != nil {
if capped {
a.metrics.SpoolScanFallbacks.With("capped").Inc()
} else {
a.metrics.SpoolScanFallbacks.With("reputation").Inc()
}
}
extra := spoolScanMatchingScript(a.spoolRoot, scriptKey(f.ScriptKey))
ids = unionStrings(ids, extra)
}
if len(ids) == 0 {
continue
}
if dryRun {
for _, id := range ids {
a.auditor.Write(auditEntry{
Ts: time.Now(), MsgID: id, ScriptKey: f.ScriptKey,
Path: f.Path, DryRun: true, Action: "freeze_dry_run",
})
}
emitted = append(emitted, alert.Finding{
Severity: alert.Warning,
Check: "email_php_relay_action_dry_run",
Path: f.Path,
Message: fmt.Sprintf("AutoFreeze dry-run: would freeze %d msgs from %s", len(ids), f.ScriptKey),
ScriptKey: f.ScriptKey,
Timestamp: time.Now(),
})
continue
}
requested := len(ids)
grant := a.rateLim.consumeUpToLimit(requested, cfg.AutoResponse.PHPRelay.MaxActionsPerMinute)
if grant < requested {
emitted = append(emitted, alert.Finding{
Severity: alert.Warning,
Check: "email_php_relay_rate_limit_hit",
Path: f.Path,
Message: fmt.Sprintf("AutoFreeze rate limit deferred %d of %d freezes for %s", requested-grant, requested, f.ScriptKey),
ScriptKey: f.ScriptKey,
Timestamp: time.Now(),
})
}
if grant == 0 {
continue
}
ids = ids[:grant]
var failed []string
for _, id := range ids {
if !msgIDPattern.MatchString(id) {
continue
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
stderr, err := a.runner.Run(ctx, a.eximBin, []string{"-Mf", id})
cancel()
entry := auditEntry{
Ts: time.Now(), MsgID: id, ScriptKey: f.ScriptKey,
Path: f.Path, DryRun: false, Stderr: stderr, Action: "freeze",
}
if err != nil {
if freezeErrIsAlreadyGone(stderr) {
if a.metrics != nil {
a.metrics.ActionGone.Inc()
}
a.auditor.Write(entry)
continue
}
entry.Exit = 1
failed = append(failed, id)
if a.metrics != nil {
a.metrics.Actions.With("freeze", "fail").Inc()
}
}
a.auditor.Write(entry)
// On successful freeze, drop the id from activeMsgs so we
// don't re-freeze it on the next finding emission.
if err == nil {
if a.metrics != nil {
a.metrics.Actions.With("freeze", "ok").Inc()
}
s.removeActive(id)
}
}
if len(failed) > 0 {
emitted = append(emitted, alert.Finding{
Severity: alert.Critical,
Check: "email_php_relay_action_failed",
Path: f.Path,
Message: fmt.Sprintf("exim -Mf failed for %d msgs from %s", len(failed), f.ScriptKey),
ScriptKey: f.ScriptKey,
MsgIDs: failed,
Timestamp: time.Now(),
})
}
}
return emitted
}
// canActOnPath reports whether AutoFreeze can act on a finding's Path.
// volume_account fires per-cpuser without scriptKey; baseline / reputation /
// header / volume / fanout all carry scriptKey.
//
//nolint:unused // wired in O2 by daemon controller
func canActOnPath(p string) bool {
switch p {
case "header", "volume", "fanout", "baseline", "reputation":
return true
}
return false
}
//nolint:unused // wired in O2 by daemon controller
func unionStrings(a, b []string) []string {
seen := make(map[string]struct{}, len(a)+len(b))
out := make([]string, 0, len(a)+len(b))
for _, s := range append(a, b...) {
if _, ok := seen[s]; ok {
continue
}
seen[s] = struct{}{}
out = append(out, s)
}
return out
}
type structuredAuditor struct {
mu sync.Mutex
w io.Writer
}
func newStructuredAuditor(w io.Writer) *structuredAuditor { return &structuredAuditor{w: w} }
func (a *structuredAuditor) Write(e auditEntry) {
payload := struct {
Ts time.Time `json:"ts"`
MsgID string `json:"msg_id"`
ScriptKey string `json:"script_key"`
Path string `json:"path"`
Action string `json:"action"`
DryRun bool `json:"dry_run"`
Exit int `json:"exit"`
Stderr string `json:"stderr,omitempty"`
}{
Ts: e.Ts.UTC(), MsgID: e.MsgID, ScriptKey: e.ScriptKey,
Path: e.Path, Action: e.Action, DryRun: e.DryRun,
Exit: e.Exit, Stderr: e.Stderr,
}
line, err := json.Marshal(payload)
if err != nil {
return
}
a.mu.Lock()
defer a.mu.Unlock()
_, _ = a.w.Write(line)
_, _ = a.w.Write([]byte("\n"))
}
package daemon
import (
"io"
"os"
"sync"
"github.com/pidginhost/csm/internal/config"
)
// eximAuditWriterAt returns the writer used by the structured JSONL auditor.
// When auto-freeze is disabled, it returns a lazy writer so daemon startup
// does not create a 0-byte orphan log file, while manual thaw commands can
// still create the file on their first real audit entry. When auto-freeze is
// enabled, the file is opened at startup to preserve the existing live-action
// failure visibility.
func eximAuditWriterAt(cfg *config.Config, path string) io.Writer {
if cfg == nil || !cfg.PHPRelayFreezeEnabled() {
return &lazyEximAuditWriter{path: path}
}
return openEximAuditWriterAt(path)
}
type lazyEximAuditWriter struct {
mu sync.Mutex
path string
w io.Writer
}
func (w *lazyEximAuditWriter) Write(p []byte) (int, error) {
w.mu.Lock()
defer w.mu.Unlock()
if w.w == nil {
w.w = openEximAuditWriterAt(w.path)
}
return w.w.Write(p)
}
func openEximAuditWriterAt(path string) io.Writer {
// #nosec G304 G302 -- G304: path is the compile-time constant phpRelayAuditPath in production; tests pass t.TempDir-derived paths. G302: 0640 is intentional; SIEM log shippers (Vector, Filebeat, Fluentbit) commonly run as a non-root user that needs group-read access. 0600 would force the shipper to run as root. Same rationale as internal/alert/audit_jsonl.go.
if f, err := os.OpenFile(path, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0640); err == nil {
return f
}
return os.Stderr
}
package daemon
import (
"context"
"encoding/json"
"errors"
"fmt"
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/store"
)
// PHPRelayController aggregates the phprelay-state references the
// control-socket handlers need. One instance per running daemon;
// constructed in daemon.New (Phase O2) and assigned to ControlListener.phprelay.
//
// The spool-pipeline reference (linux-only spoolPipeline) is intentionally
// not held here so this file stays cross-platform. Phase O2 can attach
// the pipeline through a separate linux-gated wiring file if needed.
type PHPRelayController struct {
eng *evaluator
msgIndex *msgIDIndex
ignores *ignoreList
actionDryRun *runtimeBool
db *store.DB
runner runner
eximBin string
auditor auditor
enabled bool
platform string
}
// runtimeBool is the in-memory dry-run override; effective-value precedence
// resolves CLI > bbolt > csm.yaml at read time.
type runtimeBool struct {
mu sync.Mutex
set bool
value bool
}
func (r *runtimeBool) Set(v bool) { r.mu.Lock(); r.set = true; r.value = v; r.mu.Unlock() }
func (r *runtimeBool) Reset() { r.mu.Lock(); r.set = false; r.mu.Unlock() }
func (r *runtimeBool) Get() (value, set bool) {
r.mu.Lock()
defer r.mu.Unlock()
return r.value, r.set
}
// effectiveDryRun resolves precedence: runtime > bbolt > csm.yaml.
// Returns (effective, source) where source identifies the winning input.
func (c *PHPRelayController) effectiveDryRun() (bool, string) {
var cfg *config.Config
if c.eng != nil {
cfg = c.eng.config()
}
return c.effectiveDryRunForConfig(cfg)
}
func (c *PHPRelayController) effectiveDryRunForConfig(cfg *config.Config) (bool, string) {
if v, set := c.actionDryRun.Get(); set {
return v, "runtime"
}
if c.db != nil {
if v, ok, err := readDryRunOverride(c.db); err == nil && ok {
return v, "bbolt"
}
}
if cfg != nil {
return cfg.PHPRelayDryRunEnabled(), "csm.yaml"
}
return true, "default"
}
// Status returns a snapshot of detector state for `csm phprelay status`.
func (c *PHPRelayController) Status(_ context.Context, _ control.PHPRelayStatusRequest) (control.PHPRelayStatusResponse, error) {
cfg := c.eng.config()
resp := control.PHPRelayStatusResponse{
Enabled: c.enabled,
Platform: c.platform,
EffectiveAccountLimit: c.eng.accountLimit(cfg),
IgnoresActive: len(c.ignores.List()),
RecentFindings: map[string]int{}, // populated by metrics in Phase N
}
eff, _ := c.effectiveDryRunForConfig(cfg)
resp.DryRun = eff
if c.eng.scripts != nil {
resp.ScriptsTracked = len(c.eng.scripts.Snapshot())
}
if c.msgIndex != nil {
resp.MsgIDIndexSize = c.msgIndex.Len()
}
return resp, nil
}
// handlePHPRelayStatus is the dispatcher-side adapter that bridges the
// json.RawMessage args to the typed Status method.
func (c *ControlListener) handlePHPRelayStatus(argsRaw json.RawMessage) (any, error) {
if c.phprelay == nil {
return nil, c.phpRelayUnavailable()
}
var req control.PHPRelayStatusRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("bad args: %w", err)
}
}
return c.phprelay.Status(context.Background(), req)
}
func (c *PHPRelayController) IgnoreScript(_ context.Context, req control.PHPRelayIgnoreScriptRequest) (control.PHPRelayIgnoreScriptResponse, error) {
if req.ScriptKey == "" {
return control.PHPRelayIgnoreScriptResponse{}, errors.New("script_key required")
}
hours := req.ForHours
if hours == 0 {
hours = 24 * 7
}
expires := time.Now().Add(time.Duration(hours) * time.Hour)
by := req.AddedBy
if by == "" {
by = "operator"
}
if req.Persist {
if err := c.ignores.AddPersist(scriptKey(req.ScriptKey), expires, by, req.Reason); err != nil {
return control.PHPRelayIgnoreScriptResponse{}, err
}
} else {
c.ignores.Add(scriptKey(req.ScriptKey), expires, by, req.Reason)
}
return control.PHPRelayIgnoreScriptResponse{ExpiresAt: expires}, nil
}
func (c *PHPRelayController) Unignore(_ context.Context, req control.PHPRelayUnignoreRequest) (struct{}, error) {
if req.ScriptKey == "" {
return struct{}{}, errors.New("script_key required")
}
if req.Persist {
if err := c.ignores.RemovePersist(scriptKey(req.ScriptKey)); err != nil {
return struct{}{}, err
}
} else {
c.ignores.Remove(scriptKey(req.ScriptKey))
}
return struct{}{}, nil
}
func (c *PHPRelayController) IgnoreList(_ context.Context, _ struct{}) (control.PHPRelayIgnoreListResponse, error) {
raw := c.ignores.List()
out := make([]control.PHPRelayIgnoreEntry, 0, len(raw))
for _, e := range raw {
out = append(out, control.PHPRelayIgnoreEntry{
ScriptKey: e.ScriptKey, ExpiresAt: e.ExpiresAt,
AddedBy: e.AddedBy, Reason: e.Reason,
})
}
return control.PHPRelayIgnoreListResponse{Entries: out}, nil
}
func (c *ControlListener) handlePHPRelayIgnoreScript(argsRaw json.RawMessage) (any, error) {
if c.phprelay == nil {
return nil, c.phpRelayUnavailable()
}
var req control.PHPRelayIgnoreScriptRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("bad args: %w", err)
}
}
return c.phprelay.IgnoreScript(context.Background(), req)
}
func (c *ControlListener) handlePHPRelayUnignore(argsRaw json.RawMessage) (any, error) {
if c.phprelay == nil {
return nil, c.phpRelayUnavailable()
}
var req control.PHPRelayUnignoreRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("bad args: %w", err)
}
}
return c.phprelay.Unignore(context.Background(), req)
}
func (c *ControlListener) handlePHPRelayIgnoreList(argsRaw json.RawMessage) (any, error) {
if c.phprelay == nil {
return nil, c.phpRelayUnavailable()
}
_ = argsRaw // no args expected
return c.phprelay.IgnoreList(context.Background(), struct{}{})
}
func (c *PHPRelayController) DryRun(_ context.Context, req control.PHPRelayDryRunRequest) (control.PHPRelayDryRunResponse, error) {
switch req.Mode {
case "on":
c.actionDryRun.Set(true)
if req.Persist {
if err := writeDryRunOverride(c.db, true, "operator"); err != nil {
return control.PHPRelayDryRunResponse{}, err
}
}
case "off":
c.actionDryRun.Set(false)
if req.Persist {
if err := writeDryRunOverride(c.db, false, "operator"); err != nil {
return control.PHPRelayDryRunResponse{}, err
}
}
case "reset":
c.actionDryRun.Reset()
if req.Persist {
if err := deleteDryRunOverride(c.db); err != nil {
return control.PHPRelayDryRunResponse{}, err
}
}
default:
return control.PHPRelayDryRunResponse{}, errors.New("mode must be on|off|reset")
}
eff, src := c.effectiveDryRun()
return control.PHPRelayDryRunResponse{Effective: eff, Source: src}, nil
}
// DryRunFn evaluates the precedence chain against the operation's config
// snapshot. Daemon wiring passes it to newAutoFreezer so `csm phprelay
// dry-run` changes freeze behaviour without rebuilding the freezer.
func (c *PHPRelayController) DryRunFn() func(*config.Config) bool {
return func(cfg *config.Config) bool {
v, _ := c.effectiveDryRunForConfig(cfg)
return v
}
}
func (c *ControlListener) handlePHPRelayDryRun(argsRaw json.RawMessage) (any, error) {
if c.phprelay == nil {
return nil, c.phpRelayUnavailable()
}
var req control.PHPRelayDryRunRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("bad args: %w", err)
}
}
return c.phprelay.DryRun(context.Background(), req)
}
// Thaw runs `exim -Mt <msg_id>` to release a frozen message back to the
// queue. msgIDPattern validation guards against header-injected garbage
// even though only operators can hit this endpoint. The audit entry is
// written for both success and failure so an operator can later prove
// what was thawed.
//
// req.By is accepted on the wire for forward compatibility (future
// auditEntry.By field) but is not used by the M4 handler.
func (c *PHPRelayController) Thaw(ctx context.Context, req control.PHPRelayThawRequest) (control.PHPRelayThawResponse, error) {
if !msgIDPattern.MatchString(req.MsgID) {
return control.PHPRelayThawResponse{}, errors.New("invalid msg_id")
}
sub, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
stderr, err := c.runner.Run(sub, c.eximBin, []string{"-Mt", req.MsgID})
c.auditor.Write(auditEntry{
Ts: time.Now(), MsgID: req.MsgID, Action: "thaw",
Stderr: stderr,
})
if err != nil {
return control.PHPRelayThawResponse{Stderr: stderr}, err
}
return control.PHPRelayThawResponse{Stderr: stderr}, nil
}
func (c *ControlListener) handlePHPRelayThaw(argsRaw json.RawMessage) (any, error) {
if c.phprelay == nil {
return nil, c.phpRelayUnavailable()
}
var req control.PHPRelayThawRequest
if len(argsRaw) > 0 {
if err := json.Unmarshal(argsRaw, &req); err != nil {
return nil, fmt.Errorf("bad args: %w", err)
}
}
return c.phprelay.Thaw(context.Background(), req)
}
package daemon
import (
"sync"
"github.com/pidginhost/csm/internal/metrics"
)
// phpRelayMetrics holds every series the module emits via the local
// internal/metrics OpenMetrics implementation.
//
// Defined as a struct of pointers so callers can pass nil and skip
// observation. All increments at call sites are guarded by
// `if e.metrics != nil` checks.
type phpRelayMetrics struct {
Findings *metrics.CounterVec // labels: path
Actions *metrics.CounterVec // labels: action, result
PathSkipped *metrics.CounterVec // labels: path, reason
WindowsActive *metrics.GaugeVec // labels: kind (script/ip/account)
MsgIDIndexSize *metrics.GaugeVec // labels: layer (memory/bbolt)
MsgindexPersistDropped *metrics.Counter
MsgindexPersistErrors *metrics.Counter
InotifyOverflows *metrics.Counter
SpoolReadErrors *metrics.Counter
UserdataErrors *metrics.Counter
ActiveMsgsCapped *metrics.Counter
SpoolScanFallbacks *metrics.CounterVec // labels: reason
ActionGone *metrics.Counter
}
var (
phpRelayMetricsOnce sync.Once
phpRelayMetricsInstance *phpRelayMetrics
)
// newPHPRelayMetrics constructs (and registers via the package default
// registry) the singleton metric set. Subsequent calls return the same
// instance -- sync.Once protects against the duplicate-name panic from
// metrics.MustRegister.
func newPHPRelayMetrics() *phpRelayMetrics {
phpRelayMetricsOnce.Do(func() {
m := &phpRelayMetrics{
Findings: metrics.NewCounterVec("csm_php_relay_findings_total", "Findings emitted by php_relay paths.", []string{"path"}),
Actions: metrics.NewCounterVec("csm_php_relay_actions_total", "AutoFreeze actions attempted.", []string{"action", "result"}),
PathSkipped: metrics.NewCounterVec("csm_php_relay_path_skipped_total", "Path evaluation skipped.", []string{"path", "reason"}),
WindowsActive: metrics.NewGaugeVec("csm_php_relay_windows_active", "Active windows per kind.", []string{"kind"}),
MsgIDIndexSize: metrics.NewGaugeVec("csm_php_relay_msgid_index_size", "msgIDIndex size by storage layer.", []string{"layer"}),
MsgindexPersistDropped: metrics.NewCounter("csm_php_relay_msgindex_persist_dropped_total", "Persist queue overflow drops."),
MsgindexPersistErrors: metrics.NewCounter("csm_php_relay_msgindex_persist_errors_total", "bbolt commit failures."),
InotifyOverflows: metrics.NewCounter("csm_php_relay_inotify_overflows_total", "IN_Q_OVERFLOW events."),
SpoolReadErrors: metrics.NewCounter("csm_php_relay_spool_read_errors_total", "Spool -H read errors."),
UserdataErrors: metrics.NewCounter("csm_php_relay_userdata_errors_total", "cpanelUserDomains read errors."),
ActiveMsgsCapped: metrics.NewCounter("csm_php_relay_active_msgs_capped_total", "scriptState.activeMsgs cap-hit events."),
SpoolScanFallbacks: metrics.NewCounterVec("csm_php_relay_spool_scan_fallbacks_total", "AutoFreeze spool-scan fallback invocations.", []string{"reason"}),
ActionGone: metrics.NewCounter("csm_php_relay_action_gone_total", "Messages already absent at exim -Mf time."),
}
metrics.MustRegister("csm_php_relay_findings_total", m.Findings)
metrics.MustRegister("csm_php_relay_actions_total", m.Actions)
metrics.MustRegister("csm_php_relay_path_skipped_total", m.PathSkipped)
metrics.MustRegister("csm_php_relay_windows_active", m.WindowsActive)
metrics.MustRegister("csm_php_relay_msgid_index_size", m.MsgIDIndexSize)
metrics.MustRegister("csm_php_relay_msgindex_persist_dropped_total", m.MsgindexPersistDropped)
metrics.MustRegister("csm_php_relay_msgindex_persist_errors_total", m.MsgindexPersistErrors)
metrics.MustRegister("csm_php_relay_inotify_overflows_total", m.InotifyOverflows)
metrics.MustRegister("csm_php_relay_spool_read_errors_total", m.SpoolReadErrors)
metrics.MustRegister("csm_php_relay_userdata_errors_total", m.UserdataErrors)
metrics.MustRegister("csm_php_relay_active_msgs_capped_total", m.ActiveMsgsCapped)
metrics.MustRegister("csm_php_relay_spool_scan_fallbacks_total", m.SpoolScanFallbacks)
metrics.MustRegister("csm_php_relay_action_gone_total", m.ActionGone)
phpRelayMetricsInstance = m
})
return phpRelayMetricsInstance
}
package daemon
import (
"sync"
"time"
)
// indexEntry maps a message ID to the per-message attribution recorded at
// acceptance. Used by Path 3 (Stage 2) to map delivery-failure log lines
// back to the originating script. Public field names because gob-encoded
// for bbolt persistence in Task C3.
type indexEntry struct {
ScriptKey string
HeaderScore int
SourceIP string
CPUser string
At time.Time
}
// msgIDIndex stores indexEntry per msgID, served from memory.
// Bounded by maxEntries; overflow drops the oldest entry by acceptance time.
// Persistence to bbolt is handled by msgIndexPersister (Task C3).
type msgIDIndex struct {
mu sync.Mutex
entries map[string]indexEntry
maxEntries int
persister *msgIndexPersister // nil in unit tests; real in production
}
func newMsgIDIndex(persister *msgIndexPersister, maxEntries int) *msgIDIndex {
if maxEntries <= 0 {
maxEntries = 200_000
}
return &msgIDIndex{
entries: make(map[string]indexEntry, 4096),
maxEntries: maxEntries,
persister: persister,
}
}
// Put records an entry. If the in-memory map exceeds maxEntries, the
// oldest entry by At is evicted from memory. The persister (if non-nil)
// receives the put asynchronously; persistence failure does not affect
// in-memory correctness.
func (i *msgIDIndex) Put(msgID string, e indexEntry) {
i.mu.Lock()
if _, ok := i.entries[msgID]; !ok && len(i.entries) >= i.maxEntries {
i.evictOldestLocked()
}
i.entries[msgID] = e
i.mu.Unlock()
if i.persister != nil {
i.persister.Enqueue(msgID, e)
}
}
// Get returns the entry and whether it was present.
func (i *msgIDIndex) Get(msgID string) (indexEntry, bool) {
i.mu.Lock()
e, ok := i.entries[msgID]
i.mu.Unlock()
return e, ok
}
// Has reports whether msgID is present in memory.
func (i *msgIDIndex) Has(msgID string) bool {
i.mu.Lock()
_, ok := i.entries[msgID]
i.mu.Unlock()
return ok
}
// Len returns the number of entries currently in memory.
func (i *msgIDIndex) Len() int {
i.mu.Lock()
defer i.mu.Unlock()
return len(i.entries)
}
// SweepMemory drops entries whose At is at or before cutoff.
// Called by Flow E's 1-min ticker with cutoff = now - 4h.
func (i *msgIDIndex) SweepMemory(cutoff time.Time) int {
i.mu.Lock()
defer i.mu.Unlock()
n := 0
for id, e := range i.entries {
if !e.At.After(cutoff) {
delete(i.entries, id)
n++
}
}
return n
}
func (i *msgIDIndex) evictOldestLocked() {
var oldestID string
var oldestAt time.Time
first := true
for id, e := range i.entries {
if first || e.At.Before(oldestAt) {
oldestID = id
oldestAt = e.At
first = false
}
}
if oldestID != "" {
delete(i.entries, oldestID)
}
}
package daemon
import (
"bytes"
"encoding/gob"
"fmt"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/store"
)
const msgIndexBucket = "phprelay:msgindex"
// All bbolt access goes through `*store.DB`'s phprelay helpers. The
// underlying `*bolt.DB` is unexported by design (internal/store/db.go);
// daemon code never imports go.etcd.io/bbolt directly.
// msgIndexPersister persists msgIDIndex entries to bbolt off the hot path.
// Enqueue, Flush and Stop are safe for concurrent callers after Start.
// Persistence failure does not affect the in-memory index; it degrades recovery
// after a restart and
// emits a Critical via the alerter callback.
//
// SetErrorCallback must be invoked BEFORE Start to establish a
// happens-before relationship between the writer of `onError` and the
// goroutine that reads it; concurrent callers after Start are not safe.
type msgIndexPersister struct {
db *store.DB
queue chan persistOp
flushEvery time.Duration
batchSize int
stopCh chan struct{}
doneCh chan struct{}
flushReqCh chan chan struct{}
queueStats *queuehealth.Tracker
admission sync.Mutex
stopped bool
droppedTotal uint64
errorsTotal uint64
onError func(alert.Finding)
metrics *phpRelayMetrics
}
type persistOp struct {
msgID string
entry indexEntry
ticket queuehealth.Ticket
}
func newMsgIndexPersister(db *store.DB, queueSize int, flushEvery time.Duration) *msgIndexPersister {
if queueSize <= 0 {
queueSize = 4096
}
if flushEvery <= 0 {
flushEvery = 100 * time.Millisecond
}
return &msgIndexPersister{
db: db,
queue: make(chan persistOp, queueSize),
flushEvery: flushEvery,
batchSize: 256,
stopCh: make(chan struct{}),
doneCh: make(chan struct{}),
flushReqCh: make(chan chan struct{}),
queueStats: queuehealth.New(queueSize, time.Minute),
onError: func(alert.Finding) {},
}
}
// SetErrorCallback wires Critical findings emission for bbolt failures.
// Optional -- nil disables emission (used by tests). Must be called before
// Start; concurrent invocation after Start is not safe.
func (p *msgIndexPersister) SetErrorCallback(fn func(alert.Finding)) {
if fn != nil {
p.onError = fn
}
}
// SetMetrics wires the phpRelayMetrics sink. Optional -- nil disables
// observation (used by tests). Must be called before Start; concurrent
// invocation after Start is not safe.
func (p *msgIndexPersister) SetMetrics(m *phpRelayMetrics) {
p.metrics = m
}
// Start launches the single writer. Call it once, after configuring callbacks.
func (p *msgIndexPersister) Start() {
go p.run()
}
func (p *msgIndexPersister) Stop() {
p.admission.Lock()
if !p.stopped {
p.stopped = true
close(p.queue)
close(p.stopCh)
}
p.admission.Unlock()
<-p.doneCh
}
// Enqueue is non-blocking. Returns immediately; the op is dropped if the
// queue is full or stopped (in which case DroppedCount increments).
func (p *msgIndexPersister) Enqueue(msgID string, e indexEntry) {
p.admission.Lock()
defer p.admission.Unlock()
if p.stopped {
p.queueStats.Lose(time.Now(), 1)
} else {
ticket := p.queueStats.Begin(time.Now())
select {
case p.queue <- persistOp{msgID: msgID, entry: e, ticket: ticket}:
return
default:
ticket.Reject(time.Now())
}
}
atomic.AddUint64(&p.droppedTotal, 1)
if p.metrics != nil {
p.metrics.MsgindexPersistDropped.Inc()
}
}
// Flush blocks until the persister has drained whatever was already
// enqueued at call time. For tests and shutdown.
func (p *msgIndexPersister) Flush() {
done := make(chan struct{})
select {
case p.flushReqCh <- done:
<-done
case <-p.stopCh:
<-p.doneCh
}
}
func (p *msgIndexPersister) QueueStatuses(now time.Time) map[string]queuehealth.Status {
return map[string]queuehealth.Status{"persistence": p.queueStats.Snapshot(now)}
}
func (p *msgIndexPersister) DroppedCount() uint64 {
return atomic.LoadUint64(&p.droppedTotal)
}
func (p *msgIndexPersister) ErrorCount() uint64 {
return atomic.LoadUint64(&p.errorsTotal)
}
// Lookup reads an entry from bbolt by msgID.
func (p *msgIndexPersister) Lookup(msgID string) (indexEntry, bool, error) {
raw, ok, err := p.db.PHPRelayGet(msgIndexBucket, msgID)
if err != nil || !ok {
return indexEntry{}, false, err
}
var e indexEntry
if err := gob.NewDecoder(bytes.NewReader(raw)).Decode(&e); err != nil {
return indexEntry{}, false, fmt.Errorf("decode %s: %w", msgID, err)
}
return e, true, nil
}
// SweepBolt deletes phprelay:msgindex entries whose At <= cutoff.
// Returns the number of entries removed. Called by Flow E's 1-min ticker.
// Corrupt rows (decode failure) are also dropped to keep the bucket
// healthy.
func (p *msgIndexPersister) SweepBolt(cutoff time.Time) (int, error) {
return p.db.PHPRelaySweep(msgIndexBucket, func(_, value []byte) bool {
var e indexEntry
if err := gob.NewDecoder(bytes.NewReader(value)).Decode(&e); err != nil {
return true // drop corrupt rows
}
return !e.At.After(cutoff)
})
}
func (p *msgIndexPersister) run() {
defer close(p.doneCh)
ticker := time.NewTicker(p.flushEvery)
defer ticker.Stop()
var pending []persistOp
for {
select {
case op, ok := <-p.queue:
if !ok {
p.commitBatch(pending)
return
}
op.ticket.Start(time.Now())
pending = append(pending, op)
if len(pending) >= p.batchSize {
p.commitBatch(pending)
pending = pending[:0]
}
case <-ticker.C:
if len(pending) > 0 {
p.commitBatch(pending)
pending = pending[:0]
}
case done := <-p.flushReqCh:
p.flushPending(pending)
pending = pending[:0]
close(done)
}
}
}
func (p *msgIndexPersister) flushPending(pending []persistOp) {
// There is one consumer. Snapshot admission so concurrent producers cannot
// extend this flush, and keep each transaction within the writer's limit.
for queued := len(p.queue); queued > 0; queued-- {
op := <-p.queue
op.ticket.Start(time.Now())
pending = append(pending, op)
if len(pending) >= p.batchSize {
p.commitBatch(pending)
pending = pending[:0]
}
}
p.commitBatch(pending)
}
func (p *msgIndexPersister) commitBatch(ops []persistOp) {
if len(ops) == 0 {
return
}
// A failed error callback must not leave its batch reported as running.
defer rejectPersistenceOps(ops)
kvs := make([]store.PHPRelayKV, 0, len(ops))
var buf bytes.Buffer
for i := range ops {
op := &ops[i]
buf.Reset()
if err := gob.NewEncoder(&buf).Encode(&op.entry); err != nil {
// Encoding failure is a code bug, not a transient I/O issue.
// Skip the offending op and continue with the rest of the batch.
op.ticket.Reject(time.Now())
op.ticket = queuehealth.Ticket{}
atomic.AddUint64(&p.errorsTotal, 1)
if p.metrics != nil {
p.metrics.MsgindexPersistErrors.Inc()
}
p.onError(alert.Finding{
Severity: alert.Critical,
Check: "email_php_relay_msgindex_persist_failed",
Timestamp: time.Now(),
Message: fmt.Sprintf("encode %s: %v", op.msgID, err),
})
continue
}
kvs = append(kvs, store.PHPRelayKV{
Key: []byte(op.msgID),
Value: append([]byte(nil), buf.Bytes()...),
})
}
if err := p.db.PHPRelayPutBatch(msgIndexBucket, kvs); err != nil {
rejectPersistenceOps(ops)
atomic.AddUint64(&p.errorsTotal, 1)
if p.metrics != nil {
p.metrics.MsgindexPersistErrors.Inc()
}
p.onError(alert.Finding{
Severity: alert.Critical,
Check: "email_php_relay_msgindex_persist_failed",
Timestamp: time.Now(),
Message: fmt.Sprintf("phprelay:msgindex commit failed (%d ops): %v", len(kvs), err),
})
return
}
for i := range ops {
ops[i].ticket.Finish(time.Now())
ops[i].ticket = queuehealth.Ticket{}
}
}
func rejectPersistenceOps(ops []persistOp) {
for i := range ops {
ops[i].ticket.Reject(time.Now())
ops[i].ticket = queuehealth.Ticket{}
}
}
package daemon
import (
"bufio"
"bytes"
"context"
"errors"
"io"
"os"
"time"
"github.com/pidginhost/csm/internal/alert"
)
const phpRelayHistoryMaxLineBytes = 1024 * 1024
// ScanEximHistoryForPHPRelayAccountVolume replays exim_mainlog through the
// Path 2b parser, used at daemon startup to populate perAccountWindow with
// recent outbound activity. Each accepted finding is delivered via emit.
//
// Reads the file lazily and skips a single oversized line instead of
// abandoning later entries. Caller passes `now` to keep retro replays
// deterministic for tests; production passes time.Now() once and the parser
// uses it for window math.
//
// ctx scopes the scan to daemon lifetime; nil leaves the scan unbounded for
// direct helper callers.
func ScanEximHistoryForPHPRelayAccountVolume(ctx context.Context, path string, eng *evaluator, now time.Time, emit func(alert.Finding)) {
// #nosec G304 -- path is operator-configured / hardcoded to cPanel default.
f, err := os.Open(path)
if err != nil {
return
}
defer func() { _ = f.Close() }()
reader := bufio.NewReaderSize(f, 64*1024)
var line []byte
oversized := false
for {
if phpRelayScanContextDone(ctx) {
return
}
part, rerr := reader.ReadSlice('\n')
if len(part) > 0 && !oversized {
if len(line)+len(part) > phpRelayHistoryMaxLineBytes {
oversized = true
line = nil
} else {
line = append(line, part...)
}
}
switch {
case rerr == nil:
if !oversized {
emitPHPRelayHistoryLine(line, eng, now, emit)
}
line = nil
oversized = false
case errors.Is(rerr, bufio.ErrBufferFull):
continue
case errors.Is(rerr, io.EOF):
if len(line) > 0 && !oversized {
emitPHPRelayHistoryLine(line, eng, now, emit)
}
return
default:
return
}
}
}
func phpRelayScanContextDone(ctx context.Context) bool {
if ctx == nil {
return false
}
select {
case <-ctx.Done():
return true
default:
return false
}
}
func emitPHPRelayHistoryLine(line []byte, eng *evaluator, now time.Time, emit func(alert.Finding)) {
line = bytes.TrimSuffix(line, []byte("\n"))
line = bytes.TrimSuffix(line, []byte("\r"))
s := string(line)
// Replay only lines inside the account detection window, stamped with their
// real exim timestamp. Without this every historical line was stamped `now`,
// so a whole day of sends collapsed into one hour and fired a false
// "account sent >= N in the last hour" Critical on every daemon start.
ts, ok := parseEximTimestamp(s)
if !ok || ts.Before(now.Add(-phpRelayAccountWindowDur)) || ts.After(now) {
return
}
for _, ev := range eng.parsePHPRelayAccountVolumeAt(s, ts, now) {
emit(ev)
}
}
package daemon
import (
"bytes"
"encoding/gob"
"errors"
"time"
"github.com/pidginhost/csm/internal/store"
)
const (
settingsBucket = "phprelay:settings"
dryRunOverrideKey = "dry_run_override"
)
type dryRunOverrideRow struct {
Value bool
UpdatedAt time.Time
UpdatedBy string
}
func writeDryRunOverride(db *store.DB, val bool, by string) error {
if db == nil {
return errors.New("db nil")
}
var buf bytes.Buffer
if err := gob.NewEncoder(&buf).Encode(dryRunOverrideRow{Value: val, UpdatedAt: time.Now(), UpdatedBy: by}); err != nil {
return err
}
return db.PHPRelayPut(settingsBucket, dryRunOverrideKey, buf.Bytes())
}
func deleteDryRunOverride(db *store.DB) error {
if db == nil {
return nil
}
return db.PHPRelayDelete(settingsBucket, dryRunOverrideKey)
}
func readDryRunOverride(db *store.DB) (bool, bool, error) {
if db == nil {
return false, false, nil
}
raw, ok, err := db.PHPRelayGet(settingsBucket, dryRunOverrideKey)
if err != nil || !ok {
return false, false, err
}
var row dryRunOverrideRow
if err := gob.NewDecoder(bytes.NewReader(raw)).Decode(&row); err != nil {
return false, false, err
}
return row.Value, true, nil
}
//go:build linux
package daemon
import (
"context"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"unsafe"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/emailspool"
)
// spoolWatcher watches /var/spool/exim/input for new -H files. cPanel hashes
// msgIDs into 64+ subdirs; the watcher enumerates them at start and watches
// IN_CREATE on the parent so subdirs that appear later are also picked up.
//
// On every IN_CLOSE_WRITE / IN_MOVED_TO whose name ends in "-H", the
// supplied callback is invoked synchronously with the absolute path. The
// callback must not block long; spawn worker goroutines if needed.
type spoolWatcher struct {
root string
onFile func(path string)
fd int
parentW int
mu sync.Mutex
subDirs map[int]string // watch descriptor -> path
queueHealthOnce sync.Once
kernelQueue *notificationQueue
overflowCount uint64
onOverflow func() // invoked from Run() the moment IN_Q_OVERFLOW arrives
metrics *phpRelayMetrics
}
// SetOverflowHandler wires the recovery scan + Critical finding emission
// into the watcher. Caller passes a closure that calls runRecoveryScan
// against the spool root and emits findings via the daemon alerter.
func (w *spoolWatcher) SetOverflowHandler(fn func()) { w.onOverflow = fn }
// SetMetrics wires the phpRelayMetrics sink. Optional -- nil disables
// observation (used by tests). Must be called before Run; concurrent
// invocation after Run is not safe.
func (w *spoolWatcher) SetMetrics(m *phpRelayMetrics) { w.metrics = m }
func newSpoolWatcher(root string, onFile func(path string)) (*spoolWatcher, error) {
fd, err := unix.InotifyInit1(unix.IN_CLOEXEC | unix.IN_NONBLOCK)
if err != nil {
return nil, fmt.Errorf("inotify_init1: %w", err)
}
parentMask := uint32(unix.IN_CREATE | unix.IN_MOVED_TO)
parentW, err := unix.InotifyAddWatch(fd, root, parentMask)
if err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("inotify_add_watch %s: %w", root, err)
}
w := &spoolWatcher{
root: root,
onFile: onFile,
fd: fd,
parentW: parentW,
subDirs: make(map[int]string),
}
// Enumerate existing subdirs.
entries, err := os.ReadDir(root)
if err != nil {
_ = unix.Close(fd)
return nil, fmt.Errorf("readdir %s: %w", root, err)
}
for _, e := range entries {
if !e.IsDir() {
continue
}
if err := w.addSubdir(filepath.Join(root, e.Name())); err != nil {
// Non-fatal -- continue with what we have.
continue
}
}
return w, nil
}
func (w *spoolWatcher) addSubdir(path string) error {
mask := uint32(unix.IN_CLOSE_WRITE | unix.IN_MOVED_TO)
w.initQueueHealth()
wd, err := w.kernelQueue.useDescriptor(func() (int, error) {
return unix.InotifyAddWatch(w.fd, path, mask)
})
if err != nil {
return err
}
w.mu.Lock()
w.subDirs[wd] = path
w.mu.Unlock()
return nil
}
func (w *spoolWatcher) Close() error {
w.initQueueHealth()
return w.kernelQueue.close()
}
// Run drains inotify events until ctx is cancelled.
func (w *spoolWatcher) Run(ctx context.Context) {
w.initQueueHealth()
defer func() { _ = w.Close() }()
buf := make([]byte, 16*1024)
for {
select {
case <-ctx.Done():
return
default:
}
_, err := w.kernelQueue.read(buf, w.processEvents)
if err != nil {
if errors.Is(err, syscall.EAGAIN) || errors.Is(err, syscall.EINTR) {
// Briefly yield via select so cancellation is responsive.
select {
case <-ctx.Done():
return
default:
_, _ = w.kernelQueue.useDescriptor(func() (int, error) {
// #nosec G115 -- Linux descriptors are nonnegative int32 values.
return unix.Poll([]unix.PollFd{{Fd: int32(w.fd), Events: unix.POLLIN}}, 100)
})
continue
}
}
// Treat other errors as fatal; supervisor will restart us.
return
}
}
}
func (w *spoolWatcher) processEvents(buf []byte) {
n := len(buf)
offset := 0
for offset+unix.SizeofInotifyEvent <= n {
// #nosec G103 -- bounds-checked above; standard inotify decode pattern.
ev := (*unix.InotifyEvent)(unsafe.Pointer(&buf[offset]))
nameBytes := buf[offset+unix.SizeofInotifyEvent : offset+unix.SizeofInotifyEvent+int(ev.Len)]
name := strings.TrimRight(string(nameBytes), "\x00")
offset += unix.SizeofInotifyEvent + int(ev.Len)
if ev.Mask&unix.IN_Q_OVERFLOW != 0 {
w.kernelQueue.losses.Lose(time.Now(), 1)
w.overflowCount++
if w.metrics != nil {
w.metrics.InotifyOverflows.Inc()
}
if w.onOverflow != nil {
w.onOverflow()
}
continue
}
if int(ev.Wd) == w.parentW {
if ev.Mask&(unix.IN_CREATE|unix.IN_MOVED_TO) != 0 && name != "" {
full := filepath.Join(w.root, name)
if fi, err := os.Stat(full); err == nil && fi.IsDir() {
_ = w.addSubdir(full)
}
}
continue
}
w.mu.Lock()
dir, ok := w.subDirs[int(ev.Wd)]
w.mu.Unlock()
if !ok || name == "" {
continue
}
if !strings.HasSuffix(name, "-H") {
continue
}
w.onFile(filepath.Join(dir, name))
}
}
// OverflowCount returns the number of IN_Q_OVERFLOW events observed.
// Used by the daemon to drive recovery scans (Task I3).
//
//nolint:unused // consumed by daemon wiring (Task O2)
func (w *spoolWatcher) OverflowCount() uint64 {
return w.overflowCount
}
// spoolPipeline ties together: parse headers -> compute signals -> update
// windows -> evaluate paths -> emit findings via alerter callback.
type spoolPipeline struct {
eng *evaluator
domains *userDomainsResolver
policies *emailspool.Policies
msgIndex *msgIDIndex
ignores *ignoreList
alerter func(alert.Finding)
rebuilding atomic.Bool
}
func newSpoolPipeline(eng *evaluator, domains *userDomainsResolver, pol *emailspool.Policies, idx *msgIDIndex, ignores *ignoreList, alerter func(alert.Finding)) *spoolPipeline {
eng.SetPolicies(pol)
return &spoolPipeline{
eng: eng, domains: domains, policies: pol, msgIndex: idx, ignores: ignores, alerter: alerter,
}
}
// SetRebuilding gates finding emission during the startup spool-walker
// rebuild pass. When true: state is updated, findings are NOT emitted.
func (p *spoolPipeline) SetRebuilding(v bool) { p.rebuilding.Store(v) }
// OnFile is the inotify callback for live spool events: it stamps the message
// at the current time.
func (p *spoolPipeline) OnFile(path string) {
p.onFileAt(path, time.Now())
}
// onFileAt parses, signals, updates state, and evaluates a single -H file,
// attributing the message to event time `at`. Live callers pass time.Now(); the
// startup walker passes the -H file's ModTime so already-queued mail keeps its
// real age instead of being compressed into one instant.
func (p *spoolPipeline) onFileAt(path string, at time.Time) {
h, err := emailspool.ParseHeaders(path)
if err != nil {
if p.eng != nil && p.eng.metrics != nil {
p.eng.metrics.SpoolReadErrors.Inc()
}
return
}
if h.XPHPScript == "" {
return
}
msgID := msgIDFromPath(path)
if msgID == "" {
return
}
if p.msgIndex != nil && p.msgIndex.Has(msgID) {
return // queue-runner re-write dedup
}
auth, _ := p.domains.Domains(h.EnvelopeUser)
sig := computeSignals(h, auth, p.policies)
if sig.ScriptKey == "" {
return
}
if p.ignores != nil && p.ignores.Has(sig.ScriptKey) {
return
}
now := at
if p.msgIndex != nil {
p.msgIndex.Put(msgID, indexEntry{
ScriptKey: string(sig.ScriptKey),
SourceIP: sig.SourceIP,
CPUser: h.EnvelopeUser,
At: now,
})
}
state := p.eng.scripts.getOrCreate(sig.ScriptKey)
// Path 2 includes cron-driven mail without an HTTP source IP, so every
// script event carries its recipient parse outcome into the script window.
state.appendMessage(scriptEvent{
At: now,
MsgID: msgID,
Subject: truncateDaemon(h.Subject, phpRelayBreakdownSubjectMax),
FromMismatch: sig.FromMismatch,
AdditionalSignal: sig.AdditionalSignal,
SourceIP: sig.SourceIP,
}, h.Recipients)
state.recordActive(msgID, now)
if p.policies == nil || !p.policies.IsProxyIP(sig.SourceIP) {
p.eng.ips.appendMessage(sig.SourceIP, sig.ScriptKey, now, h.Subject, h.Recipients)
}
if p.rebuilding.Load() {
return
}
findings := p.eng.evaluatePaths(sig.ScriptKey, sig.SourceIP, h.EnvelopeUser, now)
for _, f := range findings {
p.alerter(f)
}
}
// msgIDFromPath returns the msgID portion of a /path/<msgID>-H file.
func msgIDFromPath(path string) string {
base := filepath.Base(path)
if !strings.HasSuffix(base, "-H") {
return ""
}
return strings.TrimSuffix(base, "-H")
}
// runRecoveryScan walks every -H file under spoolRoot/*/, sorts by mtime
// (oldest first), invokes onFile up to maxFiles. Returns the number scanned
// and whether the cap was hit.
//
//nolint:unused // consumed by daemon wiring (Task O2)
func runRecoveryScan(spoolRoot string, maxFiles int, onFile func(string)) (int, bool) {
type entry struct {
path string
mod time.Time
}
var entries []entry
subs, err := os.ReadDir(spoolRoot)
if err != nil {
return 0, false
}
for _, sub := range subs {
if !sub.IsDir() {
continue
}
subPath := filepath.Join(spoolRoot, sub.Name())
files, err := os.ReadDir(subPath)
if err != nil {
continue
}
for _, f := range files {
if !strings.HasSuffix(f.Name(), "-H") {
continue
}
full := filepath.Join(subPath, f.Name())
fi, err := os.Stat(full)
if err != nil {
continue
}
entries = append(entries, entry{path: full, mod: fi.ModTime()})
}
}
sort.Slice(entries, func(i, j int) bool { return entries[i].mod.Before(entries[j].mod) })
truncated := false
if len(entries) > maxFiles {
entries = entries[:maxFiles]
truncated = true
}
for _, e := range entries {
onFile(e.path)
}
return len(entries), truncated
}
// phpRelayStartupWalkMax bounds how many recent -H files the startup walker
// replays so a stuffed queue cannot stall daemon startup. The live watcher and
// the Path 2b history scan backstop anything beyond the cap.
const phpRelayStartupWalkMax = 20000
// maxDetectionWindow returns the widest window any php_relay path evaluates
// over. The startup spool walker skips -H files older than this: such mail can
// no longer contribute to any current detection and only costs parse time.
// Lives in this linux-only file because the walker is its sole caller.
func (e *evaluator) maxDetectionWindow() time.Duration {
maxMin := 60 // Path 2 absolute-volume window is hardcoded at 60 min
relayCfg := e.config().EmailProtection.PHPRelay
if v := relayCfg.RateWindowMin; v > maxMin {
maxMin = v
}
if v := relayCfg.FanoutWindowMin; v > maxMin {
maxMin = v
}
return time.Duration(maxMin) * time.Minute
}
// runStartupSpoolWalker walks the currently-queued -H files through the pipeline
// in REBUILD mode, then performs one re-evaluation pass over the reconstructed
// scriptStates. Findings are emitted ONLY in the re-evaluation pass, so the
// rebuild itself never produces duplicate findings for the same in-queue mail.
//
// Each message is attributed to its -H file ModTime, so mail queued over days is
// not compressed into one instant (which used to fabricate Path 1/2/4 bursts and
// mass-freeze the legit queue). Files older than the widest detection window are
// skipped, and the walk is bounded to the newest phpRelayStartupWalkMax files so
// a stuffed queue cannot stall startup.
func runStartupSpoolWalker(spoolRoot string, p *spoolPipeline) {
p.SetRebuilding(true)
now := time.Now()
cutoff := now.Add(-p.eng.maxDetectionWindow())
type walkEntry struct {
path string
mod time.Time
}
var entries []walkEntry
if subs, err := os.ReadDir(spoolRoot); err == nil {
for _, sub := range subs {
if !sub.IsDir() {
continue
}
subPath := filepath.Join(spoolRoot, sub.Name())
files, ferr := os.ReadDir(subPath)
if ferr != nil {
continue
}
for _, f := range files {
if !strings.HasSuffix(f.Name(), "-H") {
continue
}
full := filepath.Join(subPath, f.Name())
fi, serr := os.Stat(full)
if serr != nil {
continue
}
mod := fi.ModTime()
if mod.After(now) {
mod = now
}
if mod.Before(cutoff) {
continue // skip mail older than any detection window
}
entries = append(entries, walkEntry{path: full, mod: mod})
}
}
}
// Newest first, then cap: keep the most recent (most relevant) messages and
// bound startup work on a stuffed queue.
sort.Slice(entries, func(i, j int) bool { return entries[i].mod.After(entries[j].mod) })
if len(entries) > phpRelayStartupWalkMax {
entries = entries[:phpRelayStartupWalkMax]
}
// Replay chronologically so bounded per-script/per-account rings retain the
// newest events after their own caps are applied.
sort.Slice(entries, func(i, j int) bool { return entries[i].mod.Before(entries[j].mod) })
for _, e := range entries {
p.onFileAt(e.path, e.mod)
}
p.SetRebuilding(false)
// Re-evaluation pass.
snap := p.eng.scripts.Snapshot()
for k, s := range snap {
// We don't have per-script source IP in the snapshot; pass empty
// sourceIP. Path 4 (HTTP-IP fanout) is keyed off perIPWindow which
// was already populated in OnFile, so an empty SourceIP here just
// means the per-script Path 4 finding doesn't carry an IP -- the
// window itself still triggers correctly via direct OnFile calls
// during normal operation.
cpuser := ""
// Best-effort cpuser: read it from any active msgID's index entry.
if p.msgIndex != nil {
if ids, _ := s.snapshotActiveMsgs(); len(ids) > 0 {
if e, ok := p.msgIndex.Get(ids[0]); ok {
cpuser = e.CPUser
}
}
}
for _, f := range p.eng.evaluatePaths(k, "", cpuser, now) {
p.alerter(f)
}
}
}
// spoolSupervisor wraps a goroutine that may panic. After maxRestarts
// consecutive panics, it stops trying and invokes OnFailed (used to emit
// a Critical finding email_php_relay_watcher_failed).
type spoolSupervisor struct {
fn func(ctx context.Context)
maxRestarts int
OnFailed func()
}
//nolint:unused // consumed by daemon wiring (Task O2)
func newSpoolSupervisor(fn func(ctx context.Context), maxRestarts int) *spoolSupervisor {
return &spoolSupervisor{fn: fn, maxRestarts: maxRestarts, OnFailed: func() {}}
}
func (s *spoolSupervisor) Run(ctx context.Context) {
backoff := 100 * time.Millisecond
for attempt := 0; attempt <= s.maxRestarts; attempt++ {
select {
case <-ctx.Done():
return
default:
}
func() {
defer func() {
if r := recover(); r != nil {
_ = r // Panic recovered; loop will sleep + retry.
}
}()
s.fn(ctx)
}()
if ctx.Err() != nil {
return
}
if attempt == s.maxRestarts {
s.OnFailed()
return
}
select {
case <-ctx.Done():
return
case <-time.After(backoff):
}
if backoff < 5*time.Second {
backoff *= 2
}
}
}
package daemon
import (
"errors"
"fmt"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
)
// phpRelayUnavailableError explains why no PHP-relay controller exists.
//
// Every handler used to answer "phprelay controller not wired (Phase O2)",
// which reads as "this feature has not been written yet". It has been: the
// wiring runs from startPHPRelay, and startControlListener runs before it, so
// the controller is attached whenever it is built. What stops it being built
// is the platform gate -- a non-cPanel host, or the feature switched off (its
// default, and absent entirely from a config that predates the key).
//
// Reporting the gate that actually closed lets the operator act. Naming a
// phase number sends them to read source they do not have.
func phpRelayUnavailableError(cfg *config.Config, isCPanel bool) error {
if !isCPanel {
return errors.New("php_relay guard is inactive: it supports cPanel hosts only")
}
if cfg == nil || !cfg.EmailProtection.PHPRelay.Enabled {
return fmt.Errorf("php_relay guard is disabled: set email_protection.php_relay.enabled: true in csm.yaml, then %s", rehashAndRestartHint)
}
// Enabled on a supported platform and still absent: startup failed. Do
// not blame a setting that is already correct.
return errors.New("php_relay guard is enabled but did not start; check the daemon log for php_relay errors at startup")
}
// rehashAndRestartHint is the standard follow-up after a csm.yaml edit. The
// config hash is part of the integrity baseline, so a restart without a
// rehash makes the daemon refuse to start.
const rehashAndRestartHint = "run `csm rehash && systemctl restart csm`"
// phpRelayUnavailable is the handler-side wrapper: handlers hold a daemon,
// not a platform verdict.
func (c *ControlListener) phpRelayUnavailable() error {
var cfg *config.Config
if c.d != nil {
cfg = c.d.cfg
}
return phpRelayUnavailableError(cfg, platform.Detect().IsCPanel())
}
//go:build linux
package daemon
import (
"context"
"fmt"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/emailspool"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/store"
)
// phpRelayAuditPath is the production audit file. Linux-only because the
// writer is only attached from startPHPRelayLinux; tests pass t.TempDir
// paths directly to eximAuditWriterAt.
const phpRelayAuditPath = "/var/log/csm/php_relay_audit.jsonl"
type phpRelayPaths struct {
cpanelConfig string
spool string
historyLog string
auditLog string
}
// startPHPRelayLinux completes the PHP-relay wiring after the platform
// gate (in startPHPRelay) has confirmed cPanel + located exim. It is
// split out from daemon.go so the heavy linux-only types stay in this
// file; the cross-platform stub in php_relay_wiring_other.go keeps the
// darwin build clean.
//
// The whole pipeline is built up exactly once at daemon start; SIGHUP
// only reloads policies (handled in daemon.go where d.policies is
// already wired). Callers must not invoke this twice.
func startPHPRelayLinux(d *Daemon) {
startPHPRelayLinuxAt(d, phpRelayPaths{
cpanelConfig: "/var/cpanel/cpanel.config",
spool: "/var/spool/exim/input",
historyLog: "/var/log/exim_mainlog",
auditLog: phpRelayAuditPath,
})
}
// Explicit paths let startup tests exercise the real workers without touching
// the host's cPanel configuration, mail spool or audit log.
func startPHPRelayLinuxAt(d *Daemon, paths phpRelayPaths) {
// 1. cPanel hourly limit + Path 2b derivation.
limit, status := readCpanelHourlyLimit(paths.cpanelConfig)
switch status {
case cpanelLimitMissing, cpanelLimitUnparsable:
emitPHPRelayFinding(d, alert.Warning, "email_php_relay_cpanel_limit_unreadable",
"cpanel.config maxemailsperhour unreadable; assuming 100")
}
// 2. Policies (suspicious mailers, HTTP-proxy ranges). LoadPolicies
// returns a usable Policies even on partial-failure so we always
// have something to install on the daemon for SIGHUP reloads.
pol, _ := emailspool.LoadPolicies(d.cfg.EmailProtection.PHPRelay.PoliciesDir)
d.policies = pol
// 3. Window state (per-script, per-IP, per-account).
psw := newPerScriptWindow()
pip := newPerIPWindow(64)
pacct := newPerAccountWindow(5000)
// 4. Evaluator. Wires the cPanel-derived effective account limit so
// Path 2b activates as soon as the first message arrives.
prMetrics := newPHPRelayMetrics()
eng := newEvaluator(psw, pip, pacct, d.cfg, prMetrics)
eng.SetPolicies(pol)
_, enabled, capped := deriveEffectiveAccountLimit(d.cfg, limit, status)
if !enabled {
emitPHPRelayFinding(d, alert.Warning, "email_php_relay_path2b_disabled",
"Path 2b disabled: cPanel limit off and no operator override")
}
if capped {
emitPHPRelayFinding(d, alert.Warning, "email_php_relay_account_volume_capped",
"operator AccountVolumePerHour capped to 95% of cPanel hourly limit")
}
eng.SetAccountLimitSource(limit, status)
SetPHPRelayEvaluator(eng)
// 5. msgIDIndex + persister. Bbolt access is via store.Global() --
// the daemon does not hold a *store.DB directly; the global handle
// is the same singleton the rest of the codebase uses (sigWatcher,
// retention, etc.).
bdb := store.Global()
persister := newMsgIndexPersister(bdb, 4096, 100*time.Millisecond)
persister.SetErrorCallback(func(f alert.Finding) {
alert.TryEnqueue(d.alertCh, f)
})
persister.SetMetrics(prMetrics)
d.registerQueueSource("phprelay.index", persister)
persister.Start()
d.phpRelayShutdown = append(d.phpRelayShutdown, persister.Stop)
idx := newMsgIDIndex(persister, 200_000)
// Let queue-completion log lines reap activeMsgs so delivered mail is not
// re-frozen nor charged against the freeze rate-limit budget.
eng.SetMsgIndex(idx)
// 6. ignoreList with bbolt-backed restore.
ignores := newIgnoreList()
ignores.SetStore(bdb)
_ = ignores.Restore()
// 7. cpanel user domains resolver (Path 4 helper).
domains := newUserDomainsResolver()
// 8. Controller (constructed before the freezer so DryRunFn can
// thread the runtime/bbolt/yaml precedence into freeze decisions).
runner := defaultRunner{}
auditor := newStructuredAuditor(eximAuditWriterAt(d.cfg, paths.auditLog))
controller := &PHPRelayController{
eng: eng,
msgIndex: idx,
ignores: ignores,
actionDryRun: &runtimeBool{},
db: bdb,
runner: runner,
eximBin: eximBinary,
auditor: auditor,
enabled: true,
platform: "cpanel",
}
if d.controlListener != nil {
d.controlListener.phprelay = controller
}
// 9. Spool pipeline (Flow A) + autoFreezer (post-emit hook).
pipeline := newSpoolPipeline(eng, domains, pol, idx, ignores, func(f alert.Finding) {
alert.TryEnqueue(d.alertCh, f)
})
freezer := newAutoFreezer(psw, d.cfg, paths.spool, eximBinary,
runner, auditor, prMetrics, controller.DryRunFn())
d.autoFreezer = freezer
// 10. Startup walker BEFORE the watcher to rebuild script state for
// messages already on the spool when the daemon starts.
runStartupSpoolWalker(paths.spool, pipeline)
var previousWatcher *spoolWatcher
watcherFn := func(ctx context.Context) {
w, err := newSpoolWatcher(paths.spool, pipeline.OnFile)
if err != nil {
d.MarkWatcher("phprelay", false)
emitPHPRelayFinding(d, alert.Critical, "email_php_relay_watcher_failed", err.Error())
return
}
if previousWatcher != nil {
w.inheritQueueHealth(previousWatcher)
}
previousWatcher = w
d.registerQueueSource("phprelay", w)
d.MarkWatcher("phprelay", true)
w.SetMetrics(prMetrics)
w.SetOverflowHandler(func() {
emitPHPRelayFinding(d, alert.Critical, "email_php_relay_inotify_overflow",
"inotify queue overflow; running bounded recovery scan")
const phpRelayOverflowScanMax = 1000
n, truncated := runRecoveryScan(paths.spool, phpRelayOverflowScanMax, pipeline.OnFile)
if truncated {
emitPHPRelayFinding(d, alert.Critical, "email_php_relay_overflow_scan_truncated",
fmt.Sprintf("overflow recovery capped at %d files; older messages skipped (Path 2b backstops)", phpRelayOverflowScanMax))
}
emitPHPRelayFinding(d, alert.Warning, "email_php_relay_inotify_overflow_recovered",
fmt.Sprintf("recovery scan processed %d -H files", n))
})
w.Run(ctx)
}
sup := newSpoolSupervisor(watcherFn, 5)
sup.OnFailed = func() {
d.MarkWatcher("phprelay", false)
emitPHPRelayFinding(d, alert.Critical, "email_php_relay_watcher_failed", "supervisor exhausted restarts")
}
ctx := stopChContext(d)
d.wg.Add(1)
obs.Go("php-relay-supervisor", func() {
defer d.wg.Done()
sup.Run(ctx)
})
// 11. Retrospective Path 2b scan over exim_mainlog so account-volume
// alerts fire on the first hour boundary even after a daemon
// restart that lost in-memory state. Threaded through ctx so a
// large mainlog on a busy host cannot outlive shutdown.
d.wg.Add(1)
obs.Go("php-relay-history-scan", func() {
defer d.wg.Done()
ScanEximHistoryForPHPRelayAccountVolume(ctx, paths.historyLog, eng, time.Now(), func(f alert.Finding) {
alert.TryEnqueue(d.alertCh, f)
})
})
// 12. Flow E maintenance ticker.
d.wg.Add(1)
obs.Go("php-relay-flow-e", func() {
defer d.wg.Done()
runPHPRelayFlowE(d, ctx, psw, pip, pacct, idx, persister, ignores, prMetrics)
})
}
// runPHPRelayFlowE drives Phase E maintenance (TTL sweeps + metric
// gauges). Single source of truth for php_relay TTLs; runs until ctx is
// cancelled (which happens when d.stopCh closes).
func runPHPRelayFlowE(
d *Daemon,
ctx context.Context,
psw *perScriptWindow,
pip *perIPWindow,
pacct *perAccountWindow,
idx *msgIDIndex,
persister *msgIndexPersister,
ignores *ignoreList,
m *phpRelayMetrics,
) {
minTicker := time.NewTicker(1 * time.Minute)
fiveMinTicker := time.NewTicker(5 * time.Minute)
defer minTicker.Stop()
defer fiveMinTicker.Stop()
for {
select {
case <-ctx.Done():
return
case <-minTicker.C:
now := time.Now()
_ = idx.SweepMemory(now.Add(-4 * time.Hour))
if m != nil {
m.MsgIDIndexSize.With("memory").Set(float64(idx.Len()))
}
if _, err := persister.SweepBolt(now.Add(-25 * time.Hour)); err != nil {
emitPHPRelayFinding(d, alert.Warning, "email_php_relay_sweep_failed", err.Error())
}
if _, err := ignores.SweepBolt(now); err != nil {
emitPHPRelayFinding(d, alert.Warning, "email_php_relay_sweep_failed", err.Error())
}
ignores.SweepExpired(now)
case <-fiveMinTicker.C:
now := time.Now()
cutoff25h := now.Add(-25 * time.Hour)
psw.PruneActiveMsgs(cutoff25h)
psw.SweepIdle(cutoff25h)
pip.SweepIdle(now.Add(-1 * time.Hour))
pacct.SweepIdle(now.Add(-24 * time.Hour))
}
}
}
// Nonblocking delivery keeps a full findings channel from stopping mail
// supervision; the shared queue tracker retains any delivery loss.
func emitPHPRelayFinding(d *Daemon, sev alert.Severity, check, msg string) {
alert.TryEnqueue(d.alertCh, alert.Finding{
Severity: sev,
Check: check,
Message: msg,
Timestamp: time.Now(),
})
}
// stopChContext bridges the daemon's stopCh (chan struct{}) to a
// context.Context for components that take a ctx (spool watcher, Flow E
// ticker). The returned context is cancelled when stopCh closes.
func stopChContext(d *Daemon) context.Context {
ctx, cancel := context.WithCancel(context.Background())
go func() {
<-d.stopCh
cancel()
}()
return ctx
}
package daemon
import (
"os"
"github.com/pidginhost/csm/internal/phpshield"
)
const phpShieldScriptPath = phpshield.ScriptPath
const phpShieldMissingScriptWarning = "php_shield is enabled but not installed; protection is inactive until reinstalled; " +
"run `csm enable --php-shield`, restart PHP or the web server (e.g. systemctl restart lsws || apachectl graceful), " +
"then restart csm; or set php_shield.enabled: false"
var phpShieldStat = os.Stat
// phpShieldInstalled reports whether the PHP shield script is deployed on disk.
func phpShieldInstalled() bool {
_, err := phpShieldStat(phpShieldScriptPath)
return err == nil
}
// phpShieldWatchDecision decides, from the config flag and whether the shield
// script is installed, whether the daemon should receive Shield event
// datagrams and whether it should warn that the Shield is enabled but absent.
//
// When php_shield.enabled is true but the shield was never installed (or an
// upgrade wiped /opt/csm), opening the event socket would retry forever. Instead
// we warn once with a remediation hint, so the
// misprovision is surfaced rather than masked or spammed.
func phpShieldWatchDecision(enabled, scriptExists bool) (watch, warnNotInstalled bool) {
if !enabled {
return false, false
}
if scriptExists {
return true, false
}
return false, true
}
package daemon
import (
"fmt"
"os"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/phptaintworker"
)
// phpTaintWorkerTimeout is shorter than the check-side outer deadline. The
// supervisor must have time to kill and reap a stuck parser before the shared
// deep-content walk gives up on the request.
const phpTaintWorkerTimeout = 20 * time.Second
func (d *Daemon) initPHPTaintAnalyzer() error {
sup, err := phptaintworker.NewSupervisor(phptaintworker.SupervisorConfig{
Command: d.binaryPath,
Args: []string{"phptaint-worker"},
Env: os.Environ(),
Timeout: phpTaintWorkerTimeout,
Log: func(format string, args ...any) {
fmt.Fprintf(os.Stderr, "[%s] phptaint-worker: "+format+"\n", append([]any{ts()}, args...)...)
},
})
if err != nil {
return err
}
d.phpTaintSup = sup
d.registerQueueSource("php_taint", sup)
checks.SetPHPTaintAnalyzer(sup)
return nil
}
func (d *Daemon) stopPHPTaintAnalyzer() {
checks.SetPHPTaintAnalyzer(nil)
if d.phpTaintSup == nil {
return
}
_ = d.phpTaintSup.Stop()
d.phpTaintSup = nil
}
//go:build linux
package daemon
import (
"fmt"
"math"
"os"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/platform"
)
// procRootDir is the procfs mount used by the ancestry walker. A var so
// tests can point it at a synthetic tree.
var procRootDir = "/proc"
// maxAncestryDepth bounds the PPid walk. Real package-manager chains are
// short (cpio <- weak-modules <- dnf <- systemd); a cap keeps a hostile or
// corrupt PPid chain from turning event handling into an unbounded scan.
const maxAncestryDepth = 16
// Injection points for tests; defaults are the real implementations.
var (
tmpExecPkgWindow = checks.PkgManagerRecentlyActive
tmpExecAncestry = ancestryEvidenceFor
tmpExecDemote = demoteTmpExec
)
// ancestryEvidence is what one walk of a process chain proved. The two
// signals are kept apart because they are corroborated differently: a comm
// can be set by the process itself and needs a package transaction in the
// same window to mean anything, while an exe under a panel root cannot be
// faked and stands on its own.
type ancestryEvidence struct {
packageManager bool
panelTool bool
}
func (e ancestryEvidence) none() bool { return !e.packageManager && !e.panelTool }
// procAncestryEvidence walks pid's PPid chain through procfs and reports what
// the chain proves. Any read or parse failure ends the walk with what was
// collected so far: a vanished process cannot prove provenance, so a chain
// that proved nothing leaves the alert at its original severity.
func procAncestryEvidence(pid int32, panelRoots []string) ancestryEvidence {
var ev ancestryEvidence
for depth := 0; depth < maxAncestryDepth && pid > 1; depth++ {
dir := fmt.Sprintf("%s/%d", procRootDir, pid)
// #nosec G304 -- procfs pseudo-files under a fixed root; pid is from the fanotify event.
comm, err := os.ReadFile(dir + "/comm")
if err != nil {
return ev
}
if isPackageManagerComm(strings.TrimSpace(string(comm))) {
ev.packageManager = true
}
if !ev.panelTool && procExeIsPanelTool(dir, panelRoots) {
ev.panelTool = true
}
if ev.packageManager && ev.panelTool {
return ev
}
// #nosec G304 -- procfs pseudo-files under a fixed root; pid is from the fanotify event.
status, err := os.ReadFile(dir + "/status")
if err != nil {
return ev
}
ppid := int32(-1)
for _, line := range strings.Split(string(status), "\n") {
if strings.HasPrefix(line, "PPid:") {
n, convErr := strconv.ParseInt(strings.TrimSpace(strings.TrimPrefix(line, "PPid:")), 10, 32)
if convErr != nil {
return ev
}
ppid = int32(n)
break
}
}
pid = ppid
}
return ev
}
// procExeIsPanelTool reports whether the process at procDir is running a
// binary the control panel installed. The path is tested before the file is
// stat()ed because an ancestor's exe can point anywhere.
func procExeIsPanelTool(procDir string, roots []string) bool {
if len(roots) == 0 {
return false
}
exe, err := os.Readlink(procDir + "/exe")
if err != nil || !exeInPanelRoot(exe, roots) {
return false
}
mode, uid, err := exeStat(exe)
if err != nil {
return false
}
return panelToolExeTrusted(exe, roots, mode, uid)
}
// panelToolRoots returns the detected panel's own tool directories. Detection
// is cached, so this is cheap enough to call per event. A var so tests can
// describe a cPanel host without one.
var panelToolRoots = func() []string { return platform.Detect().PanelToolRoots() }
// wireAncestryProvenance installs the ancestry hook the checks package uses to
// rescore sensitive-file findings. The /proc walk needs no BPF, so this runs
// on every Linux host; wireAncestryCache only adds a cache that survives
// process exit.
func wireAncestryProvenance() { checks.AncestryProvenance = ancestryProvenanceReason }
// ancestryEvidenceFor reports what pid's process chain proves. The BPF
// processctx cache is preferred because it survives process exit; hosts
// without BPF fall back to a live /proc walk, which is racy for short-lived
// writers but fails closed into the original severity.
func ancestryEvidenceFor(pid int32) ancestryEvidence {
if pid <= 0 {
return ancestryEvidence{}
}
roots := panelToolRoots()
if probe := cachedAncestryEvidence; probe != nil {
if ev := probe(uint32(pid), roots); !ev.none() {
return ev
}
}
return procAncestryEvidence(pid, roots)
}
// cachedAncestryEvidence reads the same evidence out of the BPF processctx
// cache. Nil on hosts built or running without BPF.
var cachedAncestryEvidence func(pid uint32, panelRoots []string) ancestryEvidence
const (
pkgAncestryReason = "ancestor is package manager"
panelToolAncestryReason = "ancestor is control panel maintenance"
)
// ancestryProvenanceReason names the trusted component behind pid, or "" when
// the chain proves nothing. Panel tooling is reported first: its evidence is
// an executable an unprivileged user cannot place, while a package-manager
// comm is only a name.
func ancestryProvenanceReason(pid uint32) string {
if pid == 0 || pid > math.MaxInt32 {
return ""
}
ev := ancestryEvidenceFor(int32(pid))
switch {
case ev.panelTool:
return panelToolAncestryReason
case ev.packageManager:
return pkgAncestryReason
default:
return ""
}
}
// demoteTmpExec decides whether an executable_in_tmp_realtime finding is
// demoted from Critical to Warning. The file must be root-owned in every
// case -- a non-root attacker can never qualify -- and then one of two
// provenance arms has to hold:
//
// - the writer descends from the control panel's own tooling, proven by an
// ancestor's executable path,
// - or it descends from a package manager AND a package-manager log was
// touched within the provenance window. The comm behind that arm is
// attacker-settable, so it is never the sole gate.
//
// The finding is rescored, never suppressed, so the evidence trail survives
// even if the heuristic is wrong.
func demoteTmpExec(uid uint32, pid int32, now time.Time) (bool, string) {
if uid != 0 || pid <= 0 {
return false, ""
}
// Both cheap gates first. The walk costs up to maxAncestryDepth procfs
// reads and runs for every root-owned executable written under a temp
// root, so it must not run when neither arm could change the verdict:
// no package transaction in the window, and no panel root for an exe to
// resolve inside.
window := tmpExecPkgWindow(now)
panelRooted := len(panelToolRoots()) > 0
if !window && !panelRooted {
return false, ""
}
ev := tmpExecAncestry(pid)
if panelRooted && ev.panelTool {
return true, "control panel maintenance ancestry"
}
if window && ev.packageManager {
return true, "package manager ancestry during active package window"
}
return false, ""
}
package daemon
import (
"fmt"
"os"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/platform"
)
// PlatformOverridesFrom maps the operator's web_server: block to platform
// overrides. WebServer is only set when a type was given, so an empty type
// leaves the detected webserver alone.
func PlatformOverridesFrom(cfg *config.Config) platform.Overrides {
var wsOverride *platform.WebServer
if t := cfg.WebServer.Type; t != "" {
ws := platform.WebServer(t)
wsOverride = &ws
}
return platform.Overrides{
WebServer: wsOverride,
ApacheConfigDir: cfg.WebServer.ConfigDir,
AccessLogPaths: cfg.WebServer.AccessLogs,
ErrorLogPaths: cfg.WebServer.ErrorLogs,
ModSecAuditLogPaths: cfg.WebServer.ModSecAudits,
DomlogGlobs: cfg.WebServer.DomlogGlobs,
}
}
// InstallPlatformOverrides installs the config-supplied platform overrides.
// It must run before anything in the process calls platform.Detect: a
// detection cached without them silently discards the operator's remedy for
// a wrong probe, so the daemon's log watchers attach to the wrong files. A
// lost override is reported loudly; the daemon keeps running on the probe.
func InstallPlatformOverrides(cfg *config.Config) bool {
if platform.SetOverrides(PlatformOverridesFrom(cfg)) {
return true
}
msg := "platform overrides from web_server: were ignored because platform detection ran first; log watchers follow the probe, not the config"
fmt.Fprintf(os.Stderr, "[ERROR] %s\n", msg)
obs.CaptureMsg("platform", msg)
return false
}
package daemon
import (
"net"
"net/http"
"net/http/pprof"
"runtime"
"strings"
"sync"
"time"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/obs"
)
const (
// mutexProfileFraction samples one in every N mutex contention events.
// Sampling limits stack collection overhead on busy hosts.
mutexProfileFraction = 100
// blockProfileRate samples one blocking event per this many nanoseconds
// spent blocked, so a goroutine parked for 10 microseconds is recorded
// with probability one. Shorter waits are sampled with proportionally lower
// probability, not excluded. Even unsampled waits incur timing overhead.
blockProfileRate = 10_000
)
// The runtime rates are process-wide. A listener stopping must not switch off
// another listener that is still serving profiles.
var contentionProfiles struct {
sync.Mutex
listeners int
}
// enableContentionProfiles turns on the sampling the mutex and block profiles
// depend on for the lifetime of successfully bound listeners.
func enableContentionProfiles() {
contentionProfiles.Lock()
defer contentionProfiles.Unlock()
if contentionProfiles.listeners == 0 {
runtime.SetMutexProfileFraction(mutexProfileFraction)
runtime.SetBlockProfileRate(blockProfileRate)
}
contentionProfiles.listeners++
}
func disableContentionProfiles() {
contentionProfiles.Lock()
defer contentionProfiles.Unlock()
contentionProfiles.listeners--
if contentionProfiles.listeners == 0 {
runtime.SetMutexProfileFraction(0)
runtime.SetBlockProfileRate(0)
}
}
// pprofListen is replaceable in tests so listener lifecycle coverage does not
// depend on the test sandbox allowing real sockets.
var pprofListen = net.Listen
// isLoopbackPprofAddr reports whether addr (host:port) binds only to a loopback
// interface. The pprof server exposes process internals and lets a caller
// trigger CPU/heap dumps, so it must never listen off-box. An empty or wildcard
// host ("" / ":6060") is rejected because it would bind all interfaces.
func isLoopbackPprofAddr(addr string) bool {
host, _, err := net.SplitHostPort(strings.TrimSpace(addr))
if err != nil {
return false
}
host = strings.TrimSpace(host)
if host == "" {
return false
}
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func newPprofMux() *http.ServeMux {
mux := http.NewServeMux()
mux.HandleFunc("/debug/pprof/", pprof.Index)
mux.HandleFunc("/debug/pprof/cmdline", pprof.Cmdline)
mux.HandleFunc("/debug/pprof/profile", pprof.Profile)
mux.HandleFunc("/debug/pprof/symbol", pprof.Symbol)
mux.HandleFunc("/debug/pprof/trace", pprof.Trace)
return mux
}
// startPprofListener starts a net/http/pprof server on addr, but only when addr
// is a loopback bind. A dedicated mux keeps the pprof handlers off
// http.DefaultServeMux so nothing else can serve them. The listener closes on
// daemon shutdown. It reports whether a listener was actually started.
func (d *Daemon) startPprofListener(addr string) bool {
if !isLoopbackPprofAddr(addr) {
csmlog.Warn("debug.pprof_listen ignored: not a loopback bind; pprof exposes process internals and must use 127.0.0.1/::1/localhost",
"addr", addr)
return false
}
ln, err := pprofListen("tcp", strings.TrimSpace(addr))
if err != nil {
csmlog.Warn("debug.pprof_listen ignored: listener failed", "addr", addr, "err", err)
return false
}
// Sampling starts with the listener, not at daemon start: a host with no
// pprof bind pays nothing for profiles nobody can fetch.
enableContentionProfiles()
srv := &http.Server{
Addr: addr,
Handler: newPprofMux(),
ReadHeaderTimeout: 10 * time.Second,
}
done := make(chan struct{})
d.wg.Add(2)
obs.Go("pprof-listener", func() {
defer d.wg.Done()
defer close(done)
csmlog.Info("pprof debug listener started", "addr", ln.Addr().String())
if err := srv.Serve(ln); err != nil && err != http.ErrServerClosed {
csmlog.Warn("pprof listener stopped", "err", err)
}
})
obs.Go("pprof-shutdown", func() {
defer d.wg.Done()
defer disableContentionProfiles()
select {
case <-d.stopCh:
_ = srv.Close()
case <-done:
}
})
return true
}
package daemon
import (
"math"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/processctx"
)
const (
processCtxCacheCap = 16384
processCtxCacheTTL = 30 * time.Minute
processCtxEnrichWorkers = 2
processCtxEnrichQueueCap = 1024
processCtxProcRoot = "/proc"
processCtxProcReadDeadline = 10 * time.Millisecond
)
var (
processCtxOnce sync.Once
processCtxCache *processctx.Cache
processCtxEnr *processctx.Enricher
processCtxPublished atomic.Pointer[processctx.Enricher]
processCtxRegistry = metrics.Default
)
var processCtxReadStartedAt = defaultProcessCtxReadStartedAt
func defaultProcessCtxReadStartedAt(pid int) (time.Time, bool) {
return processctx.NewProcReader(processCtxProcRoot, processCtxProcReadDeadline).ReadStartedAt(pid)
}
func processCtxStartedAt(pid int) time.Time {
t, ok := processCtxReadStartedAt(pid)
if !ok {
return time.Time{}
}
return t
}
// ProcessCtx returns the daemon-wide process-context cache and enricher,
// constructing them on first call and registering metrics on the default
// registry. Safe for concurrent callers.
func ProcessCtx() (*processctx.Cache, *processctx.Enricher) {
processCtxOnce.Do(func() {
processCtxCache = processctx.NewCache(processCtxCacheCap, processCtxCacheTTL)
reader := processctx.NewProcReader(processCtxProcRoot, processCtxProcReadDeadline)
processCtxEnr = processctx.NewEnricher(processCtxCache, reader, processctx.EnricherConfig{
Workers: processCtxEnrichWorkers,
QueueCap: processCtxEnrichQueueCap,
Resolver: daemonProcessIdentityResolver{},
})
processctx.RegisterMetrics(processCtxRegistry(), processCtxCache, processCtxEnr)
processCtxPublished.Store(processCtxEnr)
processCtxEnr.Start()
wireAncestryCache(processCtxCache)
})
return processCtxCache, processCtxEnr
}
func stopProcessCtx() {
if enr := processCtxPublished.Load(); enr != nil {
enr.Stop()
}
}
type daemonProcessIdentityResolver struct{}
func (daemonProcessIdentityResolver) Resolve(uid int) (string, string) {
if uid < 0 || uid > math.MaxUint32 {
return "", ""
}
user := checks.LookupUser(uint32(uid))
account := resolveLocalAccountForUID(uid, user)
return user, account
}
func resolveLocalAccountForUID(uid int, user string) string {
// First phase: for normal hosted account UIDs, the username is the account
// on cPanel and plain Linux fallback hosts. Later phases can replace this
// helper with a platform-backed account enumerator without changing the
// processctx package.
if uid >= 1000 && user != "" && !strings.HasPrefix(user, "uid:") {
return user
}
return ""
}
// resetProcessCtxForTest is a test seam. Callers in tests must run with
// t.Setenv or similar isolation; production code never invokes this.
func resetProcessCtxForTest() {
if processCtxEnr != nil {
processCtxEnr.Stop()
}
processCtxPublished.Store(nil)
processCtxOnce = sync.Once{}
processCtxCache = nil
processCtxEnr = nil
processCtxRegistry = metrics.NewRegistry
processCtxReadStartedAt = defaultProcessCtxReadStartedAt
}
package daemon
import (
"errors"
"github.com/pidginhost/csm/internal/firewall"
)
func isProtectedIPRefusal(err error) bool {
return errors.Is(err, firewall.ErrIPProtected)
}
package daemon
import (
"sync"
"time"
)
const purgeSuppressionWindow = 60 * time.Second
// purgeTracker correlates password purge events with subsequent
// stale-session 401 errors to suppress false-positive alerts.
//
// When a cPanel user changes their password, all existing sessions are
// invalidated. Any in-flight browser requests (AJAX polls, notifications,
// etc.) will return 401 - these are expected side effects, not attacks.
//
// Flow: login (NEW) records IP→account, PURGE records account→time,
// 401 handler checks IP→account→purgeTime to decide suppression.
var purgeTracker = &purgeState{
purges: make(map[string]time.Time),
sessions: make(map[string]string),
}
type purgeState struct {
mu sync.Mutex
purges map[string]time.Time // account → last purge time
sessions map[string]string // IP → last known account
}
// recordLogin tracks which account an IP most recently logged into.
func (ps *purgeState) recordLogin(ip, account string) {
ps.mu.Lock()
defer ps.mu.Unlock()
ps.sessions[ip] = account
ps.cleanupLocked()
}
// recordPurge records a password purge event for an account.
func (ps *purgeState) recordPurge(account string) {
ps.mu.Lock()
defer ps.mu.Unlock()
ps.purges[account] = time.Now()
}
// isPostPurge401 returns true if the IP's 401 is likely a stale session
// artifact from a recent password change (within the suppression window).
func (ps *purgeState) isPostPurge401(ip string) bool {
ps.mu.Lock()
defer ps.mu.Unlock()
account, ok := ps.sessions[ip]
if !ok {
return false
}
purgeTime, ok := ps.purges[account]
if !ok {
return false
}
return time.Since(purgeTime) < purgeSuppressionWindow
}
// cleanupLocked removes stale entries. Caller must hold ps.mu.
func (ps *purgeState) cleanupLocked() {
cutoff := time.Now().Add(-2 * purgeSuppressionWindow)
for k, t := range ps.purges {
if t.Before(cutoff) {
delete(ps.purges, k)
}
}
// Cap sessions map to prevent unbounded growth (keep most recent 500)
if len(ps.sessions) > 500 {
ps.sessions = make(map[string]string)
}
}
package daemon
import (
"fmt"
"maps"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/queuehealth"
)
type queueSource interface {
QueueStatuses(time.Time) map[string]queuehealth.Status
}
func (d *Daemon) registerBackendQueues(prefix string, backend bpf.Backend) {
if source, ok := backend.(queueSource); ok {
d.registerQueueSource(prefix, source)
}
}
func (d *Daemon) registerQueueSource(prefix string, source queueSource) {
d.queueSourcesMu.Lock()
defer d.queueSourcesMu.Unlock()
if d.queueSources == nil {
d.queueSources = make(map[string]queueSource)
}
d.queueSources[prefix] = source
}
func (d *Daemon) registeredQueueStatuses(now time.Time) map[string]queuehealth.Status {
d.queueSourcesMu.RLock()
sources := maps.Clone(d.queueSources)
d.queueSourcesMu.RUnlock()
statuses := make(map[string]queuehealth.Status)
for prefix, source := range sources {
for name, status := range source.QueueStatuses(now) {
statuses[prefix+"."+name] = status
}
}
return statuses
}
func (d *Daemon) setFileMonitor(fm *FileMonitor) {
d.fileMonitorMu.Lock()
d.fileMonitor = fm
d.fileMonitorMu.Unlock()
}
func (d *Daemon) getFileMonitor() *FileMonitor {
d.fileMonitorMu.RLock()
defer d.fileMonitorMu.RUnlock()
return d.fileMonitor
}
// FanotifyActive reports whether the fanotify file monitor is running.
func (d *Daemon) FanotifyActive() bool {
return d.getFileMonitor() != nil
}
// LogWatcherCount returns the number of running log watchers, including
// those a retry started after a missing log appeared.
func (d *Daemon) LogWatcherCount() int {
d.logWatchersMu.Lock()
defer d.logWatchersMu.Unlock()
return len(d.logWatchers)
}
func (d *Daemon) monitorQueueHealth() {
defer d.wg.Done()
ticker := time.NewTicker(5 * time.Second)
defer ticker.Stop()
var reporter queuehealth.Reporter
for {
select {
case <-d.stopCh:
return
case now := <-ticker.C:
d.reportQueueHealth(now, &reporter)
}
}
}
func (d *Daemon) reportQueueHealth(now time.Time, reporter *queuehealth.Reporter) {
events := reporter.Events(now, d.queueStatuses(now))
if len(events) == 0 {
return
}
findings := make([]alert.Finding, 0, len(events))
for _, event := range events {
check, message := "protection_queue_degraded", "Protection work is delayed or being dropped"
if event.Recovered {
check, message = "protection_queue_recovered", "Protection queue recovered"
}
s := event.Current
findings = append(findings, alert.Finding{
Check: check, Severity: alert.Warning, Message: message,
DedupKey: event.Name, Timestamp: now,
Details: fmt.Sprintf("queue=%s reason=%s %s", event.Name, s.Reason, s.Evidence()),
})
}
// This loop owns delivery: the failing ingest channel cannot carry its
// own alarm. History and passive observers receive transitions without
// running automatic response or advancing the last completed scan time.
d.store.AppendHistory(findings)
if err := alert.Dispatch(d.currentCfg(), findings); err != nil {
csmlog.Warn("queue health notification failed", "err", err)
}
}
//go:build linux
package daemon
import (
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
func (fm *FileMonitor) initQueueHealth() {
fm.queueHealthOnce.Do(func() {
fm.analyzerHealth = queuehealth.New(cap(fm.analyzerCh), time.Minute)
fm.reconcileHealth = queuehealth.New(reconcileDirCap, time.Minute)
fm.kernelQueueHealth = queuehealth.New(0, time.Minute)
fm.kernelQueue = newNotificationQueue(fanotifyDescriptor(fm.fd), fm.kernelQueueHealth, queuehealth.New(1, time.Minute))
})
}
func (fm *FileMonitor) queueStatuses(now time.Time) map[string]queuehealth.Status {
fm.initQueueHealth()
kernel, reader := fm.kernelQueue.snapshot(time.Now)
reconcile := fm.reconcileHealth.Snapshot(now)
reconcile.DepthUnit = "directories"
statuses := map[string]queuehealth.Status{
"fanotify.analyzer": fm.analyzerHealth.Snapshot(now),
"fanotify.kernel": kernel,
"fanotify.reader": reader,
"fanotify.reconcile": reconcile,
"fanotify.staged_packages": fm.stagedPackages().snapshot(now),
}
if fm.dropper != nil {
statuses["fanotify.dropper"], statuses["fanotify.dropper_findings"] = fm.dropper.tr.queueStatuses(now)
}
return statuses
}
func (sw *SpoolWatcher) initQueueHealth() {
sw.queueHealthOnce.Do(func() {
sw.scannerHealth = queuehealth.New(cap(sw.scanCh), time.Minute)
sw.kernelQueueHealth = queuehealth.New(0, time.Minute)
sw.kernelQueue = newNotificationQueue(fanotifyDescriptor(sw.fd), sw.kernelQueueHealth, queuehealth.New(1, time.Minute))
})
}
// A restarted watcher belongs to the same daemon lifetime. Carry its work
// accounting forward so repeated crashes cannot clear the recent loss window.
func (sw *SpoolWatcher) inheritQueueHealth(previous *SpoolWatcher) {
previous.initQueueHealth()
sw.queueHealthOnce.Do(func() {
sw.scannerHealth = previous.scannerHealth
sw.kernelQueueHealth = previous.kernelQueueHealth
sw.kernelQueue = newNotificationQueue(fanotifyDescriptor(sw.fd), sw.kernelQueueHealth, previous.kernelQueue.batches)
})
}
func (sw *SpoolWatcher) queueStatuses(now time.Time) map[string]queuehealth.Status {
sw.initQueueHealth()
kernel, reader := sw.kernelQueue.snapshot(time.Now)
return map[string]queuehealth.Status{
"spool.scanner": sw.scannerHealth.Snapshot(now),
"spool.kernel": kernel,
"spool.reader": reader,
}
}
package daemon
import (
"fmt"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/yara"
)
// reportRealtimeRuleCoverage raises a finding when a realtime engine is
// running with no rules. Both scanners treat a missing or empty rules
// directory as a successful load of nothing, so a mistyped rules_dir or an
// empty rule sync left every file write "scanned" against zero rules while
// the daemon looked healthy. yaraActive is false when no YARA backend is
// installed or its rules failed to compile; those states raise their own
// findings and are not repeated here.
func (d *Daemon) reportRealtimeRuleCoverage(yamlRules, yaraRules int, yaraActive bool) {
var empty []string
if yamlRules == 0 {
empty = append(empty, "YAML")
}
if yaraActive && yaraRules == 0 {
empty = append(empty, "YARA")
}
if len(empty) == 0 {
d.realtimeRulesMu.Lock()
d.realtimeRulesState = ""
d.realtimeRulesMu.Unlock()
return
}
rulesDir := "the configured rules directory"
if cfg := d.currentCfg(); cfg != nil && cfg.Signatures.RulesDir != "" {
rulesDir = cfg.Signatures.RulesDir
}
state := strings.Join(empty, ",") + "\x00" + rulesDir
d.realtimeRulesMu.Lock()
defer d.realtimeRulesMu.Unlock()
if d.realtimeRulesState == state {
return
}
if d.emitYaraFinding(alert.High, "realtime_rules_missing",
fmt.Sprintf("Real-time file scanning has no %s rules loaded from %s; every file write is scanned against nothing until rules are installed and reloaded.",
strings.Join(empty, " or "), rulesDir)) {
d.realtimeRulesState = state
}
}
// reportRealtimeRuleCoverageNow reads the live engines and reports on them.
func (d *Daemon) reportRealtimeRuleCoverageNow() {
yaraRules, yaraActive := 0, false
if b := yara.Active(); b != nil {
yaraActive = true
yaraRules = b.RuleCount()
}
d.reportRealtimeRuleCoverage(yamlRuleCount(), yaraRules, yaraActive)
}
// yamlRuleCount returns the number of YAML rules the global scanner holds;
// no scanner at all counts as zero rules.
func yamlRuleCount() int {
if s := signatures.Global(); s != nil {
return s.RuleCount()
}
return 0
}
package daemon
import "github.com/pidginhost/csm/internal/checks"
// reconcileReputationWhitelist pushes the live config's reputation.whitelist
// into the running threat database. The field is tagged safe for hot reload,
// so a SIGHUP that only changes it reports success; without this push the
// threat database kept the startup list until a full restart. Called from the
// reload success path, mirroring reconcileVerifiedBots.
func (d *Daemon) reconcileReputationWhitelist() {
cfg := d.activeOrStartupCfg()
if cfg == nil {
return
}
if db := checks.GetThreatDB(); db != nil {
db.SetConfigWhitelist(cfg.Reputation.Whitelist)
}
}
package daemon
import (
"bytes"
"github.com/pidginhost/csm/internal/checks"
)
// signalEagerReconcile fires a non-blocking notification on sig the first
// time count reaches threshold. Cross-platform helper extracted so the
// trigger can be unit-tested from a non-linux test file.
//
// - sig is a buffered cap-1 channel. The send is non-blocking (default
// branch) so a stalled receiver never wedges the caller.
// - The trigger fires only on the exact threshold (not >=). A long
// burst above threshold within the same window must not refire
// after the signal has been drained; the next window's first count
// reaching threshold rearms it once the receiver has reset counters.
// - A nil sig is a no-op (some unit tests construct partial structs
// that omit it).
func signalEagerReconcile(sig chan struct{}, count, threshold int64) {
if sig == nil {
return
}
if count != threshold {
return
}
select {
case sig <- struct{}{}:
default:
}
}
// Recognisers that suppress the lowest-tier "anomalous PHP location"
// warning for two specific shapes, without skipping content scanning:
//
// 1. Files inside cPanel's pkgacct/restorepkg staging tree. cPanel
// extracts the user backup as root into /home/cpanelpkgrestore.TMP.
// work.<id>/ for inspection, then re-extracts it under the user
// identity into /home/<account>/. Both extractions raise fanotify
// events; the user-context one carries the real signal, so the
// staging-side warning is a duplicate. The signature/YARA scanners
// still run on the staging file - only the path-only warning is
// dropped.
//
// 2. WP-Optimize probe files at wp-content/uploads/wpo/*. The plugin
// writes tiny <?php files to test whether the host honours certain
// Apache/Nginx directives (Server-Signature, mod_headers, mod_rewrite).
// They contain no input handling and no execution primitives; the
// anomalous-location warning is noise on every site running this
// plugin. As above, the signature/YARA scanners still run.
// looksLikeCpanelRestoreStaging delegates to the shared recogniser in
// internal/checks/sitedetect.go. The deep-scan path uses the same
// helper so realtime and scheduled scans agree on which files are
// duplicates of the user-context extraction.
func looksLikeCpanelRestoreStaging(path string) bool {
return checks.LooksLikeCpanelRestoreStaging(path)
}
// wpOptimizeProbeMaxSize bounds the size of files the recogniser will
// accept. WP-Optimize probes are header()/echo one-liners; anything
// larger fails the shape gate and falls through to the standard
// anomalous-location warning.
const wpOptimizeProbeMaxSize = 512
// wpOptimizeProbeDangerous is the deny list of byte sequences that, if
// present, disqualify a file from being treated as a WP-Optimize probe.
// Probes never use PHP superglobals or execution primitives; an attacker
// payload that does (the only realistic way to abuse a 512-byte file in
// /uploads/wpo/) trips this gate and continues to the standard alert.
//
// Tokens are matched case-insensitively against the file body. They are
// kept separate from the signature scanner above this recogniser so the
// gate stays valid even if a future signature update changes coverage.
var wpOptimizeProbeDangerous = [][]byte{
[]byte("$_"), // any PHP superglobal: $_POST, $_GET, $_REQUEST, $_COOKIE, $_SERVER...
[]byte("ev" + "al"), // split to keep the source-tree security hook happy
[]byte("ass" + "ert"),
[]byte("include"),
[]byte("require"),
[]byte("sys" + "tem"),
[]byte("p" + "assthru"),
[]byte("sh" + "ell_exec"),
[]byte("po" + "pen"),
[]byte("proc_open"),
[]byte("e" + "xec"), // plain exec() and any *exec* variant
[]byte("base64"), // any encoder/decoder pair
[]byte("phpinfo"), // information disclosure
[]byte("create_func"), // create_function deprecated lambda primitive
[]byte("file_get"), // file_get_contents (file disclosure)
[]byte("file_put"), // file_put_contents (write primitive)
[]byte("fwrite"),
[]byte("readfile"),
[]byte("`"), // backtick command substitution
}
// looksLikeWPOptimizeProbe is the realtime, content-aware check.
// It applies the shared path-only gate from internal/checks/sitedetect.go
// (path under /uploads/wpo/, basename test.php, plugin installed) and
// then adds two content-shape gates the deep-scan path cannot apply:
//
// - File body fits in wpOptimizeProbeMaxSize bytes.
// - File body contains none of wpOptimizeProbeDangerous.
//
// All gates together prevent a webshell hidden under /uploads/wpo/test.php
// from silencing the realtime warning: any payload large or interesting
// enough to be useful trips one of the content gates. The
// signature/YARA scanners run before this recogniser, so any existing
// rule still fires on its own pipeline regardless of suppression here.
func looksLikeWPOptimizeProbe(path string, content []byte) bool {
if !checks.LooksLikeWPOptimizeProbeByPath(path) {
return false
}
if len(content) > wpOptimizeProbeMaxSize {
return false
}
lower := bytes.ToLower(content)
for _, danger := range wpOptimizeProbeDangerous {
if bytes.Contains(lower, danger) {
return false
}
}
return true
}
package daemon
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/store"
)
// RetentionResult reports how many entries each sweep removed in a single
// RunRetentionOnce invocation.
type RetentionResult struct {
History int
AttackEvents int
Reputation int
FirewallActions int
Errors []error
}
// Deleted returns the total number of entries deleted across all sweeps.
func (r RetentionResult) Deleted() int {
return r.History + r.AttackEvents + r.Reputation + r.FirewallActions
}
// bucketMap describes which config knob drives which bucket sweep. Keeping
// this explicit rather than inferring it from field names keeps the mapping
// documented next to the orchestrator.
//
// Retention bucket mapping:
// - HistoryDays → `history` bucket (the finding archive every scan appends to)
// - FindingsDays → `attacks:events` bucket (per-attack event trail feeding scoring)
// - ReputationDays → `reputation` bucket (AbuseIPDB / local lookup cache, keyed by IP)
// - HistoryDays → `fw:actions` journal (proven firewall outcomes; this is
// also how far back a durable action can be undone)
//
// Blocked IPs are deliberately NOT on a TTL: `fw:blocked` is pruned when
// an operator or auto-response unblocks an IP, and temp-ban expiry is
// already handled by LoadFirewallState.
// RunRetentionOnce performs one sweep cycle over the retention-managed
// buckets. Safe to call when retention is disabled or when inputs are
// nil — it no-ops and returns a zero result.
//
// Setting a bucket's *Days value to zero means "don't sweep this bucket
// on this cycle"; negative values are treated as zero (validation catches
// them at config load).
func RunRetentionOnce(db *store.DB, cfg *config.Config, now time.Time) RetentionResult {
var result RetentionResult
if db == nil || cfg == nil || !cfg.Retention.Enabled {
return result
}
if cfg.Retention.HistoryDays > 0 {
cutoff := now.Add(-time.Duration(cfg.Retention.HistoryDays) * 24 * time.Hour)
n, err := db.SweepHistoryOlderThan(cutoff)
if err != nil {
result.Errors = append(result.Errors, err)
}
result.History = n
}
if cfg.Retention.HistoryDays > 0 {
cutoff := now.Add(-time.Duration(cfg.Retention.HistoryDays) * 24 * time.Hour)
n, err := db.SweepFirewallActionsOlderThan(cutoff)
if err != nil {
result.Errors = append(result.Errors, err)
}
result.FirewallActions = n
}
if cfg.Retention.FindingsDays > 0 {
cutoff := now.Add(-time.Duration(cfg.Retention.FindingsDays) * 24 * time.Hour)
n, err := db.SweepAttackEventsOlderThan(cutoff)
if err != nil {
result.Errors = append(result.Errors, err)
}
result.AttackEvents = n
}
if cfg.Retention.ReputationDays > 0 {
cutoff := now.Add(-time.Duration(cfg.Retention.ReputationDays) * 24 * time.Hour)
n, err := db.SweepReputationOlderThan(cutoff)
if err != nil {
result.Errors = append(result.Errors, err)
}
result.Reputation = n
}
return result
}
// retentionSweepDurationOnce guards /metrics registration of the
// retention-cycle counter so repeated daemon starts in a test binary are
// idempotent.
var retentionSweepDurationOnce sync.Once
var retentionSweepCounter *metrics.Counter
var retentionDeletedCounter *metrics.Counter
func registerRetentionMetrics() {
retentionSweepDurationOnce.Do(func() {
retentionSweepCounter = metrics.NewCounter(
"csm_retention_sweeps_total",
"Number of retention sweep cycles completed since daemon start.",
)
metrics.MustRegister("csm_retention_sweeps_total", retentionSweepCounter)
retentionDeletedCounter = metrics.NewCounter(
"csm_retention_deleted_total",
"Number of bucket entries deleted by the retention sweep.",
)
metrics.MustRegister("csm_retention_deleted_total", retentionDeletedCounter)
})
}
// retentionScanner is the daemon goroutine that drives RunRetentionOnce
// on the configured SweepInterval. Started from Run() only when
// cfg.Retention.Enabled is true; absent that, the sweep is dormant and
// no timer fires.
//
// Compaction is NOT triggered from here: reclaiming space safely requires
// exclusive access to the bbolt file, which only startup and `csm store
// compact` with the daemon stopped have. This goroutine instead emits an info
// log when the startup compaction rule says a restart would reclaim space.
func (d *Daemon) retentionScanner() {
defer d.wg.Done()
registerRetentionMetrics()
// First sweep happens after a short settle period so a restart storm
// does not hammer bbolt. Subsequent sweeps use the full interval.
settle := 5 * time.Minute
timer := time.NewTimer(settle)
defer timer.Stop()
for {
select {
case <-d.stopCh:
return
case <-timer.C:
d.runRetentionTick()
timer.Reset(d.retentionInterval())
}
}
}
// retentionInterval parses Retention.SweepInterval from the live config,
// falling back to 24h if the duration is malformed or non-positive. The
// live config is re-read each tick so SIGHUP can adjust cadence without
// restart (the retention struct itself stays hotreload:"restart", but the
// ticker can pick up edits on the next cycle).
func (d *Daemon) retentionInterval() time.Duration {
cfg := d.currentCfg()
if cfg == nil {
return 24 * time.Hour
}
ival, err := time.ParseDuration(cfg.Retention.SweepInterval)
if err != nil || ival <= 0 {
return 24 * time.Hour
}
return ival
}
// runRetentionTick runs one sweep + size check and emits metrics/logs.
func (d *Daemon) runRetentionTick() {
cfg := d.currentCfg()
db := store.Global()
result := RunRetentionOnce(db, cfg, time.Now())
retentionSweepCounter.Inc()
if n := result.Deleted(); n > 0 {
retentionDeletedCounter.Add(float64(n))
csmlog.Info("retention sweep",
"deleted_total", n,
"history", result.History,
"attacks_events", result.AttackEvents,
"reputation", result.Reputation,
)
}
for _, err := range result.Errors {
csmlog.Warn("retention sweep bucket error", "err", err)
}
// Compaction hint: the daemon auto-compacts at the next startup
// (maybeCompactStateAtStartup), and `csm store compact` does it now with
// the daemon stopped. Say so only when that startup check would act.
if size, free, due := compactionHintDue(db, cfg); due {
csmlog.Info("retention: state db is mostly free space; it will be auto-compacted on the next restart (or run `csm store compact` now with the daemon stopped)",
"size_bytes", size,
"free_bytes", free,
)
}
}
// compactionHintDue applies the startup compaction rule to the live db. A
// large file whose pages are still in use is not compacted, so it gets no hint.
func compactionHintDue(db *store.DB, cfg *config.Config) (size, free int64, due bool) {
if cfg == nil || db == nil {
return 0, 0, false
}
size, err := db.Size()
if err != nil {
return 0, 0, false
}
free, err = db.FreeBytes()
if err != nil {
return 0, 0, false
}
return size, free, store.CompactionDue(size, free, cfg.Retention.CompactMinSizeMB, cfg.Retention.CompactFillRatio)
}
package daemon
import (
"context"
"errors"
"fmt"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// scanJobQueueDepth is the maximum number of jobs that may wait in the queue
// while the single worker is busy. Enqueue returns an error when full.
const scanJobQueueDepth = 8
// maxScanJobFindingsPerJob caps how many findings a single job persists. A
// server-wide scan can emit tens of thousands of findings; storing them all
// bloats the state db and no operator triages that many per job. Findings past
// the cap are counted (FindingCount) but not stored, and the record's
// FindingsTruncated flag lets the UI show "showing first N of M". Declared as a
// var so tests can pin a small cap. Production never mutates it.
var maxScanJobFindingsPerJob = 10000
// scanJobFindingFlushBatch bounds how many findings are written per bbolt
// transaction. Batching amortizes the per-commit fsync while keeping any single
// transaction (and its memory) bounded on a huge scan.
const scanJobFindingFlushBatch = 1000
// maxRetainedScanJobFindings bounds the cumulative finding rows kept across all
// retained jobs, independent of the job-count retention. Without it, retaining
// the newest N jobs still lets finding rows grow without bound when scans are
// finding-heavy.
const maxRetainedScanJobFindings = 50000
// scanJobIDCounter is a process-wide monotonic counter used to make job IDs
// unique when two Enqueue calls occur within the same millisecond.
var scanJobIDCounter atomic.Int64
// newScanJobID returns a lexically time-sortable, collision-free job ID.
// Format: "sj-<unix-ms-hex>-<counter-hex>".
// Lexical sort on the hex timestamp gives newest-first ordering consistent
// with the store's ListScanJobs sort key.
func newScanJobID() string {
ms := time.Now().UnixMilli()
seq := scanJobIDCounter.Add(1)
return fmt.Sprintf("sj-%016x-%08x", ms, seq)
}
// scanJobRunner is the function the worker calls to perform the actual scan.
// It is a field on ScanJobManager so tests can substitute a blocking fixture
// without any sleep races.
type scanJobRunner func(ctx context.Context, cfg *config.Config, st *state.Store, target string, opts checks.AccountScanOptions) []alert.Finding
// accountEnumerator lists cPanel accounts for scope="all" jobs.
// It is a field on ScanJobManager so tests can substitute a fake without
// touching the filesystem.
type accountEnumerator func(cfg *config.Config) ([]string, error)
// scanJobRequest is what Enqueue pushes into the work channel.
type scanJobRequest struct {
work *scanJobWork
id string
opts checks.AccountScanOptions
quarantine bool // when true, annotateQuarantine runs on each finding
cancelCtx context.Context // per-job cancellable context
cancelFn context.CancelFunc // allows Cancel() to stop the runner
// remediated records the successful disposition per file path within this
// job, so a second check flagging the same file reports the action already
// taken instead of attempting it again. The map is shared by every copy of
// the request; only the worker goroutine touches it. It can contain at most
// maxScanJobFindingsPerJob entries because
// annotateQuarantine runs only for findings admitted under that cap.
remediated map[string]scanJobRemediation
}
type scanJobRemediation struct {
status string
detail string
}
// ScanJobManager runs full-scan jobs as background work in the daemon.
// A single worker goroutine drains a bounded channel of job requests so
// jobs are strictly serialised; Phase 1 does not fan out across accounts.
//
// Lifecycle: the caller owns a sync.WaitGroup slot (d.wg). Stop() cancels
// all in-flight work and blocks until the worker goroutine exits.
type ScanJobManager struct {
st *state.Store
cfg *config.Config
db scanJobStore // bbolt handle; resolved from store.Global() at construction
health *scanJobHealth
// runAccountScan is the runner called for each job. Tests replace this
// with a fixture to control blocking / return values without sleep races.
runAccountScan scanJobRunner
// enumerateAccounts lists accounts for scope="all" jobs. Tests replace it.
enumerateAccounts accountEnumerator
// quarantineFile performs the pure file quarantine for a single finding.
// Defaults to checks.QuarantineFindingFile; tests replace it to operate on
// a temp dir without touching real /home or /opt/csm.
quarantineFile func(f alert.Finding) (checks.RemediationResult, bool)
workCh chan scanJobRequest // bounded; Enqueue returns an error when full
stopCh chan struct{} // closed by Stop()
// cancelMu guards stopped and cancelFns across enqueue, cancel, run,
// drain, and stop. Enqueue holds it until the work item is buffered so
// Stop cannot close the worker before seeing the new job context.
cancelMu sync.Mutex
stopped bool
cancelFns map[string]context.CancelFunc
// workerDone is closed when the worker goroutine exits.
workerDone chan struct{}
// stopOnce ensures Stop() is idempotent: a second call is a no-op rather
// than a panic on double-close of stopCh.
stopOnce sync.Once
}
// NewScanJobManager creates a ScanJobManager and reconciles any job left in
// state "queued" or "running" from a previous daemon crash to "error" with reason
// "daemon_restarted". Returns an error if the global bbolt handle is nil.
func NewScanJobManager(st *state.Store, cfg *config.Config) (*ScanJobManager, error) {
db := store.Global()
if db == nil {
return nil, errors.New("scan-job manager: global bbolt store is nil")
}
m := &ScanJobManager{
st: st,
cfg: cfg,
db: db,
health: newScanJobHealth(),
runAccountScan: defaultScanRunner,
enumerateAccounts: checks.EnumerateScanAccounts,
quarantineFile: checks.QuarantineFindingFile,
workCh: make(chan scanJobRequest, scanJobQueueDepth),
stopCh: make(chan struct{}),
cancelFns: make(map[string]context.CancelFunc),
workerDone: make(chan struct{}),
}
if err := m.reconcileInterruptedJobs(); err != nil {
return nil, fmt.Errorf("scan-job manager: reconcile: %w", err)
}
// The worker goroutine is tracked in its own WaitGroup so Stop() can
// drain it independently of the daemon's wg. The daemon wires this
// manager via d.wg + obs.Go when it calls startScanJobManager().
obs.Go("scan-job-worker", m.worker)
return m, nil
}
// defaultScanRunner delegates to checks.RunAccountScanWithOptions.
func defaultScanRunner(ctx context.Context, cfg *config.Config, st *state.Store, target string, opts checks.AccountScanOptions) []alert.Finding {
return checks.RunAccountScanWithOptions(ctx, cfg, st, target, opts)
}
// reconcileInterruptedJobs marks unfinished persisted jobs as "error" with
// reason "daemon_restarted" because their in-memory requests cannot be resumed.
// Called once on construction before the worker
// starts, so no concurrent writes race with this read-modify-write pass.
func (m *ScanJobManager) reconcileInterruptedJobs() error {
jobs, err := m.db.ListScanJobs()
if err != nil {
return err
}
for _, rec := range jobs {
if rec.State != "running" && rec.State != "queued" {
continue
}
rec.State = "error"
rec.Error = "daemon_restarted"
rec.Finished = time.Now().UTC()
if putErr := m.db.PutScanJob(rec); putErr != nil {
return putErr
}
m.health.losses.Lose(time.Now(), 1)
csmlog.Warn("scan job reconciled after restart", "job_id", rec.ID)
}
return nil
}
// Enqueue creates a new queued job record and pushes it into the work channel.
// Returns the job ID. Returns an error when the queue is full.
// A quarantine flag is recorded on the job record for use by a later phase;
// no live auto-response is wired in Phase 1.
func (m *ScanJobManager) Enqueue(scope, target string, opts checks.AccountScanOptions, quarantine bool) (string, error) {
admission := m.health.admission.Begin(time.Now())
completed := false
defer func() {
if completed {
admission.Finish(time.Now())
} else {
admission.Reject(time.Now())
}
}()
m.cancelMu.Lock()
defer m.cancelMu.Unlock()
admission.Start(time.Now())
if m.stopped {
completed = true
return "", errors.New("scan-job manager stopped")
}
id := newScanJobID()
// Options map for storage (human-readable; not used by the worker).
optsMap := map[string]any{
"max_files": opts.MaxFiles,
"force_content": opts.ForceContent,
"force_file_index": opts.ForceFileIndex,
"respect_ignores": opts.RespectIgnores,
"max_file_bytes": opts.MaxFileBytes,
"quarantine": quarantine,
}
rec := store.ScanJobRecord{
ID: id,
Scope: scope,
Target: target,
State: "queued",
Created: time.Now().UTC(),
Options: optsMap,
}
if err := m.db.PutScanJob(rec); err != nil {
return "", fmt.Errorf("scan-job enqueue: persist: %w", err)
}
// Build a per-job cancellable context. Cancel is called either by Cancel()
// or by Stop(), which cancels all live job contexts before waiting for the
// worker goroutine to exit.
jobCtx, jobCancel := context.WithCancel(context.Background())
jobCtx, progress := checks.WithCheckDispatchProgress(jobCtx)
m.cancelFns[id] = jobCancel
req := scanJobRequest{
work: m.health.begin(progress.Snapshot),
id: id,
opts: opts,
quarantine: quarantine,
remediated: make(map[string]scanJobRemediation),
cancelCtx: jobCtx,
cancelFn: jobCancel,
}
select {
case m.workCh <- req:
completed = true
return id, nil
default:
req.work.finish(false)
completed = true
// Queue full -- remove the persisted record and cancel the context.
jobCancel()
delete(m.cancelFns, id)
// Mark the just-persisted record as error so it does not look queued.
rec.State = "error"
rec.Error = "queue_full"
rec.Finished = time.Now().UTC()
_ = m.db.PutScanJob(rec)
return "", errors.New("scan-job queue is full")
}
}
// Cancel cancels the job with the given ID. If the job is still queued the
// worker will transition it to "canceled" without running the scanner. If the
// job is running, its context is canceled and the worker persists any findings
// the runner returns after ctx.Done() before setting state "canceled".
// Returns an error when the ID is unknown or already in a terminal state.
func (m *ScanJobManager) Cancel(id string) error {
m.cancelMu.Lock()
fn, ok := m.cancelFns[id]
m.cancelMu.Unlock()
if !ok {
return fmt.Errorf("scan-job cancel: unknown or already terminal job %q", id)
}
fn()
return nil
}
// ListJobs returns all scan job records ordered newest-first.
// It is the encapsulated accessor for handleScanStatus and the WebUI Phase 2b
// handler so neither needs to reach the manager's private db field.
func (m *ScanJobManager) ListJobs() ([]store.ScanJobRecord, error) {
return m.db.ListScanJobs()
}
// ListFindings returns a paginated slice of findings for the given job ID plus
// the total number of findings stored for that job. offset and limit follow the
// usual page semantics (limit=0 returns all findings).
// It is the encapsulated accessor for handleScanReport and the WebUI Phase 2b
// handler so neither needs to reach the manager's private db field.
func (m *ScanJobManager) ListFindings(id string, offset, limit int) ([]alert.Finding, int, error) {
return m.db.ListScanJobFindings(id, offset, limit)
}
// Progress returns the current record for the given job ID.
// ok is false when the ID is not found in the store.
func (m *ScanJobManager) Progress(id string) (store.ScanJobRecord, bool) {
rec, ok, err := m.db.GetScanJob(id)
if err != nil || !ok {
return store.ScanJobRecord{}, false
}
return rec, true
}
// Stop cancels all in-flight or queued work and waits for the worker goroutine
// to exit. After Stop returns the manager must not be used. Stop is idempotent:
// a sync.Once guards the close of stopCh so repeated calls are safe.
// Canceling all live job contexts before blocking on workerDone ensures that
// any scan currently executing inside runAccountScan returns promptly; the
// runner honors ctx and the job lands in terminal state "canceled" with its
// partial findings retained.
func (m *ScanJobManager) Stop() {
m.stopOnce.Do(func() {
m.cancelMu.Lock()
m.stopped = true
close(m.stopCh)
for _, fn := range m.cancelFns {
fn()
}
m.cancelMu.Unlock()
})
<-m.workerDone
}
// worker is the single serialised goroutine that processes scan jobs.
// It holds no references to closed resources after Stop() returns because:
// 1. Stop() closes stopCh and cancels all live job contexts before blocking
// on workerDone; any in-flight scan honors ctx.Done() and returns promptly
// with partial findings, landing in terminal state "canceled".
// 2. The select below exits as soon as stopCh fires, and all remaining
// queued jobs are drained as "canceled" -- none write to the store after
// Close() because the daemon closes the store only after Stop() returns.
func (m *ScanJobManager) worker() {
defer m.workerExited()
for {
select {
case <-m.stopCh:
// Drain the queue: mark any pending jobs canceled without running.
m.drainQueueOnStop()
return
case req := <-m.workCh:
req.work.run(func() { m.runJob(req) })
}
}
}
func (m *ScanJobManager) workerExited() {
// A panic or Goexit can bypass the normal drain. Refuse new admissions
// under the same lock used to publish them before releasing queued work.
m.cancelMu.Lock()
defer m.cancelMu.Unlock()
defer close(m.workerDone)
m.stopped = true
for id, cancel := range m.cancelFns {
cancel()
delete(m.cancelFns, id)
}
for {
select {
case req := <-m.workCh:
req.work.finish(false)
default:
return
}
}
}
// drainQueueOnStop consumes all remaining items in workCh after stopCh closes
// and marks each as "canceled". Nothing writes to the store after this returns.
func (m *ScanJobManager) drainQueueOnStop() {
for {
select {
case req := <-m.workCh:
req.work.run(func() {
req.cancelFn()
m.cancelMu.Lock()
delete(m.cancelFns, req.id)
m.cancelMu.Unlock()
m.setTerminal(req, "canceled", "")
})
default:
return
}
}
}
// runJob executes a single scan job to completion and persists the result.
// The job context (req.cancelCtx) may be canceled externally by Cancel() or
// by Stop() -> drainQueueOnStop. This function handles both cases:
// - If the context is already canceled when we check, skip the scan entirely.
// - After the runner returns, inspect ctx.Err() to choose the terminal state.
func (m *ScanJobManager) runJob(req scanJobRequest) {
db := req.trackedStore(m.db)
defer func() {
// Always remove the cancel function entry when the job is done.
m.cancelMu.Lock()
delete(m.cancelFns, req.id)
m.cancelMu.Unlock()
req.cancelFn()
}()
// If the job was canceled before the worker got to it (e.g. Cancel()
// called while it was queued), skip the scan and go straight to terminal.
if req.cancelCtx.Err() != nil {
m.setTerminal(req, "canceled", "")
return
}
// Transition to "running".
rec, ok, err := db.GetScanJob(req.id)
if err != nil || !ok {
csmlog.Warn("scan job missing at run time", "job_id", req.id)
return
}
rec.State = "running"
rec.Started = time.Now().UTC()
if putErr := db.PutScanJob(rec); putErr != nil {
csmlog.Warn("scan job state update failed", "job_id", req.id, "err", putErr)
return
}
cfg := m.cfg
if cfg == nil {
cfg = &config.Config{}
}
// Branch on scope: "all" runs every account; anything else (default
// "account") runs the single target as before.
var findingCount, findingsStored int
var truncated bool
var jobState string
if rec.Scope == "all" {
findingCount, findingsStored, truncated, jobState = m.runAllAccounts(req, rec, cfg)
} else {
// Run the scan. The runner blocks until complete or ctx is canceled.
req.work.progressed()
findings := m.runAccountScan(req.cancelCtx, cfg, m.st, rec.Target, req.opts)
req.work.progressed()
// Persist findings (batched, capped), even if canceled mid-scan, so the
// "cancel keeps partial" guarantee holds. Quarantine is applied only to
// findings that will be stored; truncated findings must not trigger
// unaudited file moves.
findingsStored, truncated = m.persistFindings(req, 0, findings, func(f alert.Finding) alert.Finding {
return m.annotateQuarantine(req, f)
})
findingCount = len(findings)
jobState = "done"
if req.cancelCtx.Err() != nil {
jobState = "canceled"
}
}
// Refresh the record before writing the terminal state so FilesScanned
// and FindingCount reflect any in-progress updates (future phases may
// update these mid-scan via callbacks; for now we set them from findings).
rec2, ok2, err2 := db.GetScanJob(req.id)
if err2 != nil || !ok2 {
// Fall back to the snapshot we already have.
rec2 = rec
}
rec2.State = jobState
rec2.Finished = time.Now().UTC()
rec2.FindingCount = findingCount
rec2.FindingsStored = findingsStored
rec2.FindingsTruncated = truncated
rec2.CurrentAccount = "" // clear transient progress field on completion
if putErr := db.PutScanJob(rec2); putErr != nil {
csmlog.Warn("scan job terminal state failed", "job_id", req.id, "err", putErr)
}
// Prune oldest jobs to keep the store within the configured retention,
// bounding both the job count and the cumulative finding rows.
retention := cfg.Thresholds.ScanJobRetention
if retention <= 0 {
retention = 20
}
if _, pruneErr := db.PruneScanJobs(retention, maxRetainedScanJobFindings); pruneErr != nil {
csmlog.Warn("scan job prune failed", "err", pruneErr)
}
}
// persistFindings writes findings for a job in batched transactions, honoring
// the per-job findings cap. stored is how many findings the job has already
// persisted (so the cap spans every account of a scope="all" job). prepare is
// applied only to the findings that fit under the cap. It returns how many
// findings from this slice were persisted and whether the cap or a write error
// dropped any of them.
func (m *ScanJobManager) persistFindings(req scanJobRequest, stored int, findings []alert.Finding, prepare func(alert.Finding) alert.Finding) (written int, truncated bool) {
db := req.trackedStore(m.db)
jobID := req.id
room := len(findings)
if maxScanJobFindingsPerJob > 0 {
remaining := maxScanJobFindingsPerJob - stored
if remaining <= 0 {
return 0, len(findings) > 0
}
if room > remaining {
room = remaining
truncated = true
}
}
for off := 0; off < room; off += scanJobFindingFlushBatch {
end := off + scanJobFindingFlushBatch
if end > room {
end = room
}
batch := findings[off:end]
if prepare != nil {
batch = append([]alert.Finding(nil), batch...)
for i := range batch {
req.work.progressed()
batch[i] = prepare(batch[i])
req.work.progressed()
}
}
if err := db.AppendScanJobFindings(jobID, stored+off, batch); err != nil {
csmlog.Warn("scan job finding batch persist failed", "job_id", jobID, "seq", stored+off, "err", err)
return written, true
}
written += len(batch)
}
return written, truncated
}
// runAllAccounts executes the scan body for a scope="all" job: enumerates
// accounts, scans each with panic isolation, persists findings incrementally,
// and returns (findingCount, terminalState). Terminal state is "error" if
// enumeration fails, "canceled" if the context is done after the loop,
// "done" otherwise.
//
// For cancel mid-iteration: the check at the top of the per-account loop
// detects a canceled context before starting the next account, so partial
// findings already persisted from completed accounts are retained.
func (m *ScanJobManager) runAllAccounts(req scanJobRequest, rec store.ScanJobRecord, cfg *config.Config) (findingCount, findingsStored int, truncated bool, jobState string) {
db := req.trackedStore(m.db)
req.work.progressed()
accounts, err := m.enumerateAccounts(cfg)
req.work.progressed()
if err != nil {
req.work.fail()
// Persist error state immediately so the outer runJob terminal write
// picks up the correct state (it will overwrite State/Finished again,
// but we set Error here via a direct PutScanJob call first).
recErr, ok2, getErr := db.GetScanJob(req.id)
if getErr == nil && ok2 {
recErr.State = "error"
recErr.Error = err.Error()
recErr.Finished = time.Now().UTC()
_ = db.PutScanJob(recErr)
}
return 0, 0, false, "error"
}
// Set AccountsTotal and persist immediately so progress is visible.
rec.AccountsTotal = len(accounts)
if putErr := db.PutScanJob(rec); putErr != nil {
csmlog.Warn("scan job accounts_total persist failed", "job_id", req.id, "err", putErr)
}
findingCount = 0
findingsStored = 0
for _, account := range accounts {
// Cancel check: stop before starting next account.
if req.cancelCtx.Err() != nil {
break
}
// Update progress: current account being scanned.
rec.CurrentAccount = account
if putErr := db.PutScanJob(rec); putErr != nil {
csmlog.Warn("scan job progress persist failed", "job_id", req.id, "err", putErr)
}
// Run the account scan with panic isolation so one bad account cannot
// abort the entire server-wide job.
req.work.progressed()
findings := m.runAccountScanIsolated(req, cfg, account)
req.work.progressed()
// Attribute each finding to the account when the check did not already
// set TenantID, then persist the account's findings as one batch. The
// per-job cap spans accounts via the running findingsStored counter.
now := time.Now().UTC()
for i := range findings {
if findings[i].TenantID == "" {
findings[i].TenantID = account
}
if findings[i].Timestamp.IsZero() {
findings[i].Timestamp = now
}
}
written, trunc := m.persistFindings(req, findingsStored, findings, func(f alert.Finding) alert.Finding {
return m.annotateQuarantine(req, f)
})
findingsStored += written
truncated = truncated || trunc
findingCount += len(findings)
rec.AccountsDone++
rec.FindingCount = findingCount
rec.FindingsStored = findingsStored
rec.FindingsTruncated = truncated
if putErr := db.PutScanJob(rec); putErr != nil {
csmlog.Warn("scan job accounts_done persist failed", "job_id", req.id, "err", putErr)
}
}
// Clear the transient CurrentAccount field before terminal state is written.
rec.CurrentAccount = ""
jobState = "done"
if req.cancelCtx.Err() != nil {
jobState = "canceled"
}
return findingCount, findingsStored, truncated, jobState
}
// runAccountScanIsolated calls m.runAccountScan inside a recover() so a panic
// in any account's scanner is converted to a synthetic account_scan_error
// finding rather than crashing the worker goroutine and aborting the job.
func (m *ScanJobManager) runAccountScanIsolated(req scanJobRequest, cfg *config.Config, account string) (findings []alert.Finding) {
defer func() {
if r := recover(); r != nil {
req.work.fail()
findings = []alert.Finding{{
Severity: alert.High,
Check: "account_scan_error",
Message: fmt.Sprintf("scan failed for account %s", account),
Details: fmt.Sprintf("%v", r),
TenantID: account,
Timestamp: time.Now().UTC(),
}}
}
}()
return m.runAccountScan(req.cancelCtx, cfg, m.st, account, req.opts)
}
// setTerminal writes a terminal state for a job without running any scan.
// Used to transition queued-but-canceled jobs during drainQueueOnStop.
func (m *ScanJobManager) setTerminal(req scanJobRequest, jobState, errMsg string) {
db := req.trackedStore(m.db)
id := req.id
rec, ok, err := db.GetScanJob(id)
if err != nil || !ok {
return
}
rec.State = jobState
rec.Error = errMsg
rec.Finished = time.Now().UTC()
_ = db.PutScanJob(rec)
}
// annotateQuarantine runs the full-scan file remediation on f when the job's
// quarantine flag is set, and stamps the finding's RemediationStatus /
// RemediationDetail fields with the outcome. When the flag is false
// (report-only) f is returned unchanged so existing consumers see no JSON diff.
//
// Eligible standalone malware files are quarantined, while eligible WordPress
// core, plugin and theme files are cleaned in place. Ineligible findings such
// as process kills, DB cleanup and htaccess edits are marked "left_for_review".
// This method never calls alert.Dispatch, state.AppendHistory, or
// StoreLatestScanFindings; remediation is not an alert.
func (m *ScanJobManager) annotateQuarantine(req scanJobRequest, f alert.Finding) alert.Finding {
if !req.quarantine {
return f
}
// A cancelled job keeps its partial findings but stops acting on them:
// the operator withdrew consent, and findings produced during teardown
// are the least verified of the batch.
if req.cancelCtx != nil && req.cancelCtx.Err() != nil {
return f
}
if prior, done := req.remediated[f.FilePath]; done && f.FilePath != "" {
f.RemediationStatus = prior.status
f.RemediationDetail = prior.detail
return f
}
result, eligible := m.quarantineFile(f)
if !eligible {
f.RemediationStatus = "left_for_review"
return f
}
if result.Success {
f.RemediationStatus = result.RemediationStatus
if f.RemediationStatus == "" {
f.RemediationStatus = "quarantined"
}
f.RemediationDetail = result.Action
if req.remediated != nil {
req.remediated[f.FilePath] = scanJobRemediation{
status: f.RemediationStatus,
detail: result.Action,
}
}
} else {
req.work.fail()
f.RemediationStatus = "failed"
f.RemediationDetail = result.Error
}
return f
}
// startScanJobManager initialises the ScanJobManager and wires it into the
// daemon's lifecycle. Call this from Daemon.Run() after store.Global() is set
// and before startControlListener() so the control socket can immediately
// accept job submissions.
func (d *Daemon) startScanJobManager() (*ScanJobManager, error) {
m, err := NewScanJobManager(d.store, d.cfg)
if err != nil {
return nil, err
}
d.registerQueueSource("scans", m)
// Track the manager's lifetime in the daemon wait-group using the same
// obs.Go + defer d.wg.Done() pattern every other background worker uses.
// NewScanJobManager already started the worker goroutine via its own
// obs.Go call; this entry just blocks d.wg.Wait() until that goroutine
// finishes draining. Stop() must be called before d.wg.Wait() so that
// workerDone is already closed by the time we reach here.
d.wg.Add(1)
obs.Go("scan-job-manager", func() {
defer d.wg.Done()
<-m.workerDone
})
return m, nil
}
package daemon
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/queuehealth"
)
const scanJobControlBudget = time.Minute
type scanJobHealth struct {
mu sync.Mutex
pending map[*scanJobWork]struct{}
losses *queuehealth.Tracker
admission *queuehealth.Tracker
progress time.Time
fullSince time.Time
}
type scanJobWork struct {
owner *scanJobHealth
checks func(time.Time) checks.DispatchProgressSnapshot
started time.Time
progress time.Time
failed bool
}
func newScanJobHealth() *scanJobHealth {
return &scanJobHealth{
pending: make(map[*scanJobWork]struct{}),
losses: queuehealth.New(0, scanJobControlBudget),
admission: queuehealth.New(0, scanJobControlBudget),
}
}
func (h *scanJobHealth) begin(progress func(time.Time) checks.DispatchProgressSnapshot) *scanJobWork {
h.mu.Lock()
defer h.mu.Unlock()
now := time.Now()
if len(h.pending) == 0 {
h.progress = now
}
w := &scanJobWork{owner: h, checks: progress, progress: now}
h.pending[w] = struct{}{}
h.updateFull(now)
return w
}
func (h *scanJobHealth) updateFull(now time.Time) {
waiting := 0
for w := range h.pending {
if w.started.IsZero() {
waiting++
}
}
if waiting >= scanJobQueueDepth {
if h.fullSince.IsZero() {
h.fullSince = now
}
} else {
h.fullSince = time.Time{}
}
}
func (w *scanJobWork) progressed() {
w.owner.mu.Lock()
defer w.owner.mu.Unlock()
w.progress = time.Now()
w.owner.progress = w.progress
}
func (w *scanJobWork) fail() {
w.owner.mu.Lock()
defer w.owner.mu.Unlock()
w.failLocked(time.Now())
}
func (w *scanJobWork) failLocked(now time.Time) {
if !w.failed {
w.failed = true
w.owner.losses.Lose(now, 1)
}
}
func (w *scanJobWork) finish(completed bool) {
h := w.owner
h.mu.Lock()
defer h.mu.Unlock()
now := time.Now()
if !completed {
w.failLocked(now)
}
delete(h.pending, w)
if !w.started.IsZero() {
h.progress = now
}
h.updateFull(now)
}
func (w *scanJobWork) run(fn func()) {
h := w.owner
h.mu.Lock()
w.started = time.Now()
w.progress = w.started
h.progress = w.started
h.updateFull(w.started)
h.mu.Unlock()
completed := false
defer func() { w.finish(completed) }()
fn()
completed = true
}
func (h *scanJobHealth) snapshot(now time.Time) queuehealth.Status {
h.mu.Lock()
defer h.mu.Unlock()
s := h.losses.Snapshot(now)
s.Capacity = scanJobQueueDepth
s.LagBasis = "consumer_progress"
progress := h.progress
var overdue bool
for w := range h.pending {
if w.started.IsZero() {
s.Depth++
continue
}
s.InFlight++
s.ProcessingSeconds = max(s.ProcessingSeconds, now.Sub(w.started).Seconds())
child := w.checks(now)
latest := w.progress
if child.LastProgress.After(latest) {
latest = child.LastProgress
}
if latest.After(progress) {
progress = latest
}
if child.Active {
overdue = overdue || child.Overdue
} else {
overdue = overdue || now.Sub(latest) >= scanJobControlBudget
}
}
if s.Depth > 0 {
s.LagSeconds = max(0, now.Sub(progress).Seconds())
}
switch {
case overdue:
s.Reason = "processing_lag"
case s.Depth > 0 && s.InFlight == 0 && now.Sub(progress) >= scanJobControlBudget:
s.Reason = "backlog_lag"
case !h.fullSince.IsZero() && now.Sub(h.fullSince) >= 30*time.Second:
s.Reason = "queue_full"
}
if s.Reason != "" {
s.Status = "degraded"
}
return s
}
// QueueStatuses retains work through persistence and worker cleanup. It never
// takes the cancellation lock or asks the store for progress.
func (m *ScanJobManager) QueueStatuses(now time.Time) map[string]queuehealth.Status {
admission := m.health.admission.Snapshot(now)
admission.CapacityUnavailable = true
return map[string]queuehealth.Status{"jobs": m.health.snapshot(now), "admission": admission}
}
package daemon
import (
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/store"
)
type scanJobStore interface {
PutScanJob(store.ScanJobRecord) error
GetScanJob(string) (store.ScanJobRecord, bool, error)
ListScanJobs() ([]store.ScanJobRecord, error)
AppendScanJobFindings(string, int, []alert.Finding) error
ListScanJobFindings(string, int, int) ([]alert.Finding, int, error)
PruneScanJobs(int, int) (int, error)
}
// Only worker operations advance the job clock. Status polling through the
// manager's store must not conceal stalled worker persistence.
type scanJobWorkStore struct {
scanJobStore
work *scanJobWork
}
func (req scanJobRequest) trackedStore(db scanJobStore) scanJobWorkStore {
return scanJobWorkStore{scanJobStore: db, work: req.work}
}
func (db scanJobWorkStore) PutScanJob(rec store.ScanJobRecord) error {
db.work.progressed()
err := db.scanJobStore.PutScanJob(rec)
if err != nil {
db.work.fail()
}
db.work.progressed()
return err
}
func (db scanJobWorkStore) GetScanJob(id string) (store.ScanJobRecord, bool, error) {
db.work.progressed()
rec, ok, err := db.scanJobStore.GetScanJob(id)
if err != nil || !ok {
db.work.fail()
}
db.work.progressed()
return rec, ok, err
}
func (db scanJobWorkStore) AppendScanJobFindings(id string, seq int, findings []alert.Finding) error {
db.work.progressed()
err := db.scanJobStore.AppendScanJobFindings(id, seq, findings)
if err != nil {
db.work.fail()
}
db.work.progressed()
return err
}
func (db scanJobWorkStore) PruneScanJobs(keep, maxFindings int) (int, error) {
db.work.progressed()
n, err := db.scanJobStore.PruneScanJobs(keep, maxFindings)
if err != nil {
db.work.fail()
}
db.work.progressed()
return n, err
}
//go:build !(linux && bpf)
package daemon
import (
"context"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
)
// sensitiveFileBPF is the no-tag placeholder for the BPF sensitive-file
// monitor. The real type with BPF map handles, link, and the ringbuf reader
// lives in sensitive_file_bpf.go behind //go:build linux && bpf.
type sensitiveFileBPF struct{}
func (s *sensitiveFileBPF) Mode() string { return "bpf" }
func (s *sensitiveFileBPF) EventCount() uint64 { return 0 }
func (s *sensitiveFileBPF) Run(_ context.Context) {}
func startSensitiveFileBPF(_ context.Context, _ chan<- alert.Finding, _ *config.Config) (*sensitiveFileBPF, error) {
return nil, bpf.ErrNotBuilt
}
// Code generated by bpf2go; DO NOT EDIT.
//go:build 386 || amd64
package sensitive_file_bpfprog
import (
"bytes"
_ "embed"
"fmt"
"io"
"structs"
"github.com/cilium/ebpf"
)
type SensitiveFileCsmQueueStats struct {
_ structs.HostLayout
Lost uint64
Submitted uint64
}
type SensitiveFileFileid struct {
_ structs.HostLayout
Dev uint64
Ino uint64
}
type SensitiveFileSensitiveEvent struct {
_ structs.HostLayout
Uid uint32
Pid uint32
Mask uint32
_ [4]byte
Dev uint64
Ino uint64
Comm [16]uint8
}
// Names of all BPF objects in the ELF.
//
// Used for safe lookups in a Collection or CollectionSpec.
const (
SensitiveFileMapEvents = "events"
SensitiveFileMapQueueStats = "queue_stats"
SensitiveFileMapWatched = "watched"
SensitiveFileProgCsmFilePerm = "csm_file_perm"
SensitiveFileVarUnused = "unused"
)
// LoadSensitiveFile returns the embedded CollectionSpec for SensitiveFile.
func LoadSensitiveFile() (*ebpf.CollectionSpec, error) {
reader := bytes.NewReader(_SensitiveFileBytes)
spec, err := ebpf.LoadCollectionSpecFromReader(reader)
if err != nil {
return nil, fmt.Errorf("can't load SensitiveFile: %w", err)
}
return spec, err
}
// LoadSensitiveFileObjects loads SensitiveFile and converts it into a struct.
//
// The following types are suitable as obj argument:
//
// *SensitiveFileObjects
// *SensitiveFilePrograms
// *SensitiveFileMaps
//
// See ebpf.CollectionSpec.LoadAndAssign documentation for details.
func LoadSensitiveFileObjects(obj any, opts *ebpf.CollectionOptions) error {
spec, err := LoadSensitiveFile()
if err != nil {
return err
}
return spec.LoadAndAssign(obj, opts)
}
// SensitiveFileSpecs contains maps and programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type SensitiveFileSpecs struct {
SensitiveFileProgramSpecs
SensitiveFileMapSpecs
SensitiveFileVariableSpecs
}
// SensitiveFileProgramSpecs contains programs before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type SensitiveFileProgramSpecs struct {
CsmFilePerm *ebpf.ProgramSpec `ebpf:"csm_file_perm"`
}
// SensitiveFileMapSpecs contains maps before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type SensitiveFileMapSpecs struct {
Events *ebpf.MapSpec `ebpf:"events"`
QueueStats *ebpf.MapSpec `ebpf:"queue_stats"`
Watched *ebpf.MapSpec `ebpf:"watched"`
}
// SensitiveFileVariableSpecs contains global variables before they are loaded into the kernel.
//
// It can be passed ebpf.CollectionSpec.Assign.
type SensitiveFileVariableSpecs struct {
Unused *ebpf.VariableSpec `ebpf:"unused"`
}
// SensitiveFileObjects contains all objects after they have been loaded into the kernel.
//
// It can be passed to LoadSensitiveFileObjects or ebpf.CollectionSpec.LoadAndAssign.
type SensitiveFileObjects struct {
SensitiveFilePrograms
SensitiveFileMaps
SensitiveFileVariables
}
func (o *SensitiveFileObjects) Close() error {
return _SensitiveFileClose(
&o.SensitiveFilePrograms,
&o.SensitiveFileMaps,
)
}
// SensitiveFileMaps contains all maps after they have been loaded into the kernel.
//
// It can be passed to LoadSensitiveFileObjects or ebpf.CollectionSpec.LoadAndAssign.
type SensitiveFileMaps struct {
Events *ebpf.Map `ebpf:"events"`
QueueStats *ebpf.Map `ebpf:"queue_stats"`
Watched *ebpf.Map `ebpf:"watched"`
}
func (m *SensitiveFileMaps) Close() error {
return _SensitiveFileClose(
m.Events,
m.QueueStats,
m.Watched,
)
}
// SensitiveFileVariables contains all global variables after they have been loaded into the kernel.
//
// It can be passed to LoadSensitiveFileObjects or ebpf.CollectionSpec.LoadAndAssign.
type SensitiveFileVariables struct {
Unused *ebpf.Variable `ebpf:"unused"`
}
// SensitiveFilePrograms contains all programs after they have been loaded into the kernel.
//
// It can be passed to LoadSensitiveFileObjects or ebpf.CollectionSpec.LoadAndAssign.
type SensitiveFilePrograms struct {
CsmFilePerm *ebpf.Program `ebpf:"csm_file_perm"`
}
func (p *SensitiveFilePrograms) Close() error {
return _SensitiveFileClose(
p.CsmFilePerm,
)
}
func _SensitiveFileClose(closers ...io.Closer) error {
for _, closer := range closers {
if err := closer.Close(); err != nil {
return err
}
}
return nil
}
// Do not access this directly.
//
//go:embed sensitivefile_x86_bpfel.o
var _SensitiveFileBytes []byte
//go:build linux
package daemon
import "golang.org/x/sys/unix"
// kernelDeviceID converts the device number stat(2) reports, which uses the
// glibc encoding (major<<8)|(minor&0xff)|((minor&~0xff)<<12), into the
// kernel's own dev_t (major<<20)|minor. That is what an LSM program reads from
// inode->i_sb->s_dev, so a BPF map keyed by device must use this form. The two
// encodings agree only for major 0 (tmpfs, overlay), which is why a map keyed
// with the raw stat value works in test rigs and misses every file on a real
// block device.
func kernelDeviceID(statDev uint64) uint64 {
return uint64(unix.Major(statDev))<<20 | uint64(unix.Minor(statDev))
}
package daemon
import (
"encoding/binary"
"errors"
)
// SensitiveFileEvent matches struct sensitive_event in sensitive_file.bpf.c
// byte for byte. Userspace looks up (Dev, Ino) in the in-memory mirror of
// the BPF watchset map to recover the path string at finding time.
type SensitiveFileEvent struct {
UID uint32
PID uint32
Mask uint32
Dev uint64
Ino uint64
Comm string
}
const sensitiveFileEventSize = 4 + 4 + 4 + 4 + 8 + 8 + 16
func decodeSensitiveFileEvent(b []byte) (SensitiveFileEvent, error) {
if len(b) < sensitiveFileEventSize {
return SensitiveFileEvent{}, errors.New("sensitive file event short buffer")
}
ev := SensitiveFileEvent{
UID: binary.LittleEndian.Uint32(b[0:4]),
PID: binary.LittleEndian.Uint32(b[4:8]),
Mask: binary.LittleEndian.Uint32(b[8:12]),
Dev: binary.LittleEndian.Uint64(b[16:24]),
Ino: binary.LittleEndian.Uint64(b[24:32]),
}
ev.Comm = nullTerm(b[32 : 32+16])
return ev, nil
}
package daemon
import (
"context"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/state"
)
// sensitiveFilePoller wraps checks.CheckSensitiveFiles in a goroutine.
// Used when the BPF backend is unavailable or operator-disabled. Detection
// latency equals the poll interval (default 5 minutes).
type sensitiveFilePoller struct {
cfg *config.Config
store *state.Store
alertCh chan<- alert.Finding
count atomic.Uint64
}
func newSensitiveFilePoller(cfg *config.Config, store *state.Store, alertCh chan<- alert.Finding) *sensitiveFilePoller {
return &sensitiveFilePoller{cfg: cfg, store: store, alertCh: alertCh}
}
func (p *sensitiveFilePoller) Mode() string { return "legacy" }
func (p *sensitiveFilePoller) EventCount() uint64 { return p.count.Load() }
func (p *sensitiveFilePoller) Run(ctx context.Context) {
interval := sensitiveFilePollerInterval(p.cfg)
t := time.NewTicker(interval)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
for _, f := range checks.CheckSensitiveFiles(ctx, p.cfg, p.store) {
p.count.Add(1)
if !alert.TryEnqueue(p.alertCh, f) {
csmlog.Warn("sensitive_file legacy: alert channel full, dropping finding")
}
}
}
}
}
func sensitiveFilePollerInterval(cfg *config.Config) time.Duration {
if d := cfg.Detection.SensitiveFilesPollInterval; d > 0 {
return d
}
return 5 * time.Minute
}
package daemon
import (
"context"
"errors"
"strings"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/state"
)
// StartSensitiveFileMonitor selects the active sensitive-file write monitor
// based on cfg.Detection.SensitiveFilesBackend and host capability.
//
// "auto" (default) -- try BPF, fall back to legacy hash-comparison polling.
// "bpf" -- require BPF; return nil if unavailable (no fallback).
// "legacy" -- pin legacy polling.
// "none" -- disable the live monitor (the periodic check still runs).
//
// Unknown values fall back to "auto" with a warning. The metric
// csm_bpf_backend{feature="sensitive_files", kind="..."} reflects the
// chosen path.
func StartSensitiveFileMonitor(alertCh chan<- alert.Finding, cfg *config.Config, store *state.Store) bpf.Backend {
choice := strings.ToLower(strings.TrimSpace(cfg.Detection.SensitiveFilesBackend))
if choice == "" {
choice = bpf.BackendAuto
}
switch choice {
case bpf.BackendAuto, bpf.BackendBPF, bpf.BackendLegacy, bpf.BackendNone:
default:
csmlog.Warn("sensitive_files: unknown backend choice, using auto", "value", choice)
choice = bpf.BackendAuto
}
if choice == bpf.BackendNone {
csmlog.Info("sensitive_files: disabled by config")
bpf.SetActive("sensitive_files", bpf.BackendNone)
return nil
}
var bpfErr error
if choice == bpf.BackendAuto || choice == bpf.BackendBPF {
if b, err := tryStartSensitiveFileBPFFn(context.Background(), alertCh, cfg); err == nil && b != nil {
csmlog.Info("sensitive_files", "backend", "bpf", "choice", choice)
bpf.SetActive("sensitive_files", bpf.BackendBPF)
return b
} else if err != nil {
bpfErr = err
level := "bpf-unsupported"
if errors.Is(err, bpf.ErrNotBuilt) {
level = "bpf-not-built"
}
csmlog.Info("sensitive_files: BPF unavailable", "state", level, "reason", err.Error(), "choice", choice)
if choice == bpf.BackendBPF {
csmlog.Warn("sensitive_files: backend=bpf but BPF unavailable; no live monitor", "reason", err.Error())
bpf.SetActive("sensitive_files", bpf.BackendNone)
emitBPFUnavailableFinding(alertCh, "sensitive_files", choice, "", err)
return nil
}
}
}
poller := newSensitiveFilePoller(cfg, store, alertCh)
csmlog.Info("sensitive_files", "backend", "legacy", "choice", choice)
bpf.SetActive("sensitive_files", bpf.BackendLegacy)
if bpfErr != nil {
emitBPFUnavailableFinding(alertCh, "sensitive_files", choice, bpf.BackendLegacy, bpfErr)
}
return poller
}
var tryStartSensitiveFileBPFFn = tryStartSensitiveFileBPF
func tryStartSensitiveFileBPF(ctx context.Context, ch chan<- alert.Finding, cfg *config.Config) (bpf.Backend, error) {
b, err := startSensitiveFileBPF(ctx, ch, cfg)
if err != nil {
return nil, err
}
return b, nil
}
package daemon
import (
"crypto/sha256"
"encoding/hex"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/store"
)
// Signature-update-driven retroactive rescan.
//
// CSM's signature rules update independently of the deep-tier
// scanner. Without this watcher, a fresh ruleset only catches files
// that change AFTER the update -- existing files that newly match
// stay silent until the next time they happen to be touched. Real
// attacks aren't that polite.
//
// The watcher polls cfg.Signatures.RulesDir every sigWatchInterval,
// stat()s every *.yaml / *.yml / *.yar / *.yara file, and sets the
// daemon's forceFullRescan flag whenever any tracked file's content
// changes. The next deep-tier tick reads + clears the flag and runs
// the full account tree instead of the fanotify short-list.
//
// A file is hashed only when its mtime or size moved. Package upgrades
// and the rules updater rewrite files whose content is unchanged, and a
// full rescan reads every file on the host, so a moved mtime alone is
// not a reason to arm.
//
// The per-file state is persisted in bbolt (sig_watch bucket) so a
// daemon restart does not look like "all files are new" and trigger a
// phantom rescan on first tick.
const sigWatchInterval = 60 * time.Second
var sigWatchExtensions = []string{".yaml", ".yml", ".yar", ".yara"}
var (
sigRescansTotalOnce sync.Once
sigRescansTotal *metrics.Counter
)
// observeSignatureRescan increments the operator-facing counter the
// first time the watcher arms a rescan in any process lifetime, and
// every subsequent time. Called from the deep-tier path AFTER a full
// retro-sweep completes, so the counter measures completed sweeps,
// not queued ones.
func observeSignatureRescan() {
sigRescansTotalOnce.Do(func() {
sigRescansTotal = metrics.NewCounter(
"csm_signature_rescans_total",
"Signature-update-driven full deep-tier rescans completed. Incremented when the deep-tier scheduler picks up the forceFullRescan flag set by the signature watcher and finishes a sweep against the new ruleset.",
)
metrics.MustRegister("csm_signature_rescans_total", sigRescansTotal)
})
sigRescansTotal.Inc()
}
// sigWatcher carries the watcher's loop state. The daemon owns one
// instance; the goroutine in (*Daemon).signatureWatcher drives it.
//
// cfg and store are re-resolved per tick, not captured at
// construction. Originally we cached cfg.Signatures.RulesDir and
// store.Global() into struct fields and discovered two ways that
// could go wrong: a hot-reload that changed the rules dir would
// silently keep walking the old path, and a daemon ordering quirk
// where store.Global() is nil at goroutine spawn would leave the
// watcher persistence-blind for the rest of its lifetime. Live
// resolution closes both.
type sigWatcher struct {
cfgFunc func() *config.Config
storeFunc func() *store.DB
rescanFlag *atomic.Bool
alertCh chan<- alert.Finding
interval time.Duration
hashFile func(string, os.FileInfo) (string, error)
// Initialised on first tick from store.GetSignatureFiles(); the
// in-memory map is the authoritative working copy for the loop.
last map[string]store.SignatureFileState
// persisted is false until last has been written to bbolt, and again
// after it changes, so an unchanged tick commits nothing.
persisted bool
}
// newSigWatcher constructs a watcher with production defaults.
// cfgFunc and storeFunc are called per tick so config hot-reloads
// and lazy bbolt initialisation are picked up automatically.
// Callers can override the interval after construction for tests.
func newSigWatcher(cfgFunc func() *config.Config, storeFunc func() *store.DB, flag *atomic.Bool, alertCh chan<- alert.Finding) *sigWatcher {
return &sigWatcher{
cfgFunc: cfgFunc,
storeFunc: storeFunc,
rescanFlag: flag,
alertCh: alertCh,
interval: sigWatchInterval,
hashFile: hashRulesFile,
}
}
// loadInitial pulls the persisted state into memory. Called once when
// w.last is nil. A read error here is non-fatal -- the watcher operates
// with an empty map and the next tick re-persists, so the cost of a
// transient bbolt error is at most one phantom rescan.
func (w *sigWatcher) loadInitial(sdb *store.DB) {
if sdb == nil {
w.last = map[string]store.SignatureFileState{}
return
}
got, err := sdb.GetSignatureFiles()
if err != nil {
csmlog.Warn("sig_watch: loading persisted state", "err", err)
w.last = map[string]store.SignatureFileState{}
return
}
w.last = got
w.persisted = true
}
// tick performs one walk of the rules dir and arms the rescan flag
// when any tracked file's content changed. Removed files drop out of
// the persisted map without triggering a rescan -- the spec calls out
// only a change to an existing file as a trigger.
func (w *sigWatcher) tick() {
cfg := w.cfgFunc()
if !sigWatchEnabled(cfg) {
return
}
rulesDir := cfg.Signatures.RulesDir
if rulesDir == "" {
return
}
sdb := w.storeFunc()
// Defer first-time persistence load until we have a non-nil
// store. A nil store on the first tick (race against bbolt
// open) means we operate purely in-memory; once bbolt is up,
// the next tick triggers loadInitial as if for the first time
// because last is still nil.
if w.last == nil && sdb != nil {
w.loadInitial(sdb)
}
if w.last == nil {
w.last = map[string]store.SignatureFileState{}
}
current := walkRulesDir(rulesDir)
next := make(map[string]store.SignatureFileState, len(current))
var changed []sigWatchChange
for path, file := range current {
info := file.info
old, seen := w.last[path]
if seen && old.SHA256 != "" && old.Size == info.Size() && old.Mtime.Equal(info.ModTime()) {
next[path] = old
continue
}
digest, err := w.hashFile(path, info)
if err != nil {
csmlog.Warn("sig_watch: hashing rules", "path", path, "err", err)
if !seen {
continue
}
// Keep the last good comparison point, including its stamp, so
// the next tick retries instead of accepting an unreadable update.
if old.SHA256 != "" {
next[path] = old
continue
}
// With no recorded hash, preserve the legacy mtime fallback
// even when this read cannot establish a content baseline.
}
state := store.SignatureFileState{Mtime: info.ModTime(), Size: info.Size(), SHA256: digest}
next[path] = state
switch {
case !seen:
// New file. The spec treats first-observation as a
// non-event so a fresh `update-rules` install does not
// cause a rescan when the daemon also starts cold.
case old.SHA256 != "":
if old.SHA256 != state.SHA256 {
changed = append(changed, sigWatchChange{Path: path, Old: old.Mtime, New: state.Mtime})
}
case old.Size < 0 && old.Mtime.Equal(file.legacyMtime),
old.Size == state.Size && old.Mtime.Equal(state.Mtime):
// A legacy record with the same stamp gains a hash without
// treating the first tick after upgrade as a rules change.
default:
// The stamp moved and the contents cannot be compared, so
// this is treated as a change, as it was before content
// hashes were tracked.
changed = append(changed, sigWatchChange{Path: path, Old: old.Mtime, New: state.Mtime})
}
}
if !sameSignatureState(w.last, next) {
w.persisted = false
}
w.last = next
if sdb != nil && !w.persisted {
if err := sdb.PutSignatureFiles(next); err != nil {
csmlog.Warn("sig_watch: persisting state", "err", err)
} else {
w.persisted = true
}
}
if len(changed) == 0 {
return
}
w.rescanFlag.Store(true)
for _, c := range changed {
alert.TryEnqueue(w.alertCh, alert.Finding{
Severity: alert.Warning,
Check: "signature_update_rescan_queued",
Message: fmt.Sprintf("Signature update detected, full deep rescan queued: %s", filepath.Base(c.Path)),
Details: fmt.Sprintf("File: %s\nOld mtime: %s\nNew mtime: %s", c.Path, c.Old.UTC().Format(time.RFC3339), c.New.UTC().Format(time.RFC3339)),
FilePath: c.Path,
Timestamp: time.Now(),
})
}
}
// hashRulesFile accepts a hash only while the walked file, the open file and
// the final pathname still identify the same regular file and stamp.
func hashRulesFile(path string, expected os.FileInfo) (string, error) {
if !expected.Mode().IsRegular() {
return "", fmt.Errorf("rules file is not regular")
}
// A replacement by a FIFO between walk and open must not stall the watcher.
f, err := os.OpenFile(path, os.O_RDONLY|unix.O_NONBLOCK, 0) // #nosec G304 -- path comes from walking the operator-configured rules dir and the opened identity is checked below.
if err != nil {
return "", err
}
defer func() { _ = f.Close() }()
before, err := f.Stat()
if err != nil {
return "", err
}
if !sameRulesFile(expected, before) {
return "", fmt.Errorf("rules file changed before hashing or is not regular")
}
h := sha256.New()
n, err := io.Copy(h, io.LimitReader(f, before.Size()+1))
if err != nil {
return "", err
}
after, err := f.Stat()
if err != nil {
return "", err
}
current, err := os.Stat(path)
if err != nil {
return "", err
}
if n != before.Size() || !sameRulesFile(before, after) || !sameRulesFile(after, current) {
return "", fmt.Errorf("rules file changed while hashing")
}
return hex.EncodeToString(h.Sum(nil)), nil
}
func sameRulesFile(a, b os.FileInfo) bool {
return a.Mode().IsRegular() && b.Mode().IsRegular() && os.SameFile(a, b) &&
a.Size() == b.Size() && a.ModTime().Equal(b.ModTime())
}
func sameSignatureState(a, b map[string]store.SignatureFileState) bool {
if len(a) != len(b) {
return false
}
for path, x := range a {
y, ok := b[path]
if !ok || x.Size != y.Size || x.SHA256 != y.SHA256 || !x.Mtime.Equal(y.Mtime) {
return false
}
}
return true
}
type sigWatchFile struct {
info os.FileInfo
// Old builds recorded the link's mtime for symlinks, not the target's.
legacyMtime time.Time
}
// walkRulesDir returns the target file info of every signature file under dir.
// Sub-directories are walked too -- the YARA Forge updater puts
// files under tier-named subfolders. Errors during walk (missing
// dir, EACCES on a sub-tree) are swallowed; we want one bad path
// not to crash the watcher or stop sibling traversal.
func walkRulesDir(dir string) map[string]sigWatchFile {
out := map[string]sigWatchFile{}
_ = filepath.Walk(dir, func(path string, info os.FileInfo, err error) error {
if err != nil {
// Missing dir or EACCES on a sub-tree -- ignore so the
// watcher does not crash. We deliberately swallow err
// instead of returning it; filepath.SkipDir is the
// idiomatic alternative but we want to keep walking
// siblings, not the descendants of one bad path.
return filepath.SkipDir
}
if info.IsDir() {
return nil
}
ext := strings.ToLower(filepath.Ext(path))
if !sigWatchExtMatches(ext) {
return nil
}
file := sigWatchFile{info: info, legacyMtime: info.ModTime()}
if info.Mode()&os.ModeSymlink != 0 {
if target, err := os.Stat(path); err == nil {
file.info = target
}
}
out[path] = file
return nil
})
return out
}
// sigWatchExtMatches returns true when ext is one of the file
// extensions the watcher tracks. Lower-case input expected.
func sigWatchExtMatches(ext string) bool {
for _, want := range sigWatchExtensions {
if ext == want {
return true
}
}
return false
}
// sigWatchEnabled resolves the tri-state cfg flag. Same shape as
// dbObjectScanningEnabled in the checks package: nil = on, *true =
// on, *false = off.
func sigWatchEnabled(cfg *config.Config) bool {
if cfg == nil {
return true
}
if cfg.Detection.RescanOnSignatureUpdate == nil {
return true
}
return *cfg.Detection.RescanOnSignatureUpdate
}
// sigWatchChange records one changed file for the alert detail
// message.
type sigWatchChange struct {
Path string
Old time.Time
New time.Time
}
// signatureWatcher is the daemon's signature-watch goroutine. Runs
// until d.stopCh is closed; ticks every sigWatchInterval, sets
// d.forceFullRescan when any tracked rule file's content changes.
//
// Cfg and store are accessed via getter closures (not captured
// values) so a hot-reload of signatures.rules_dir takes effect on
// the next tick and a late-initialised bbolt is picked up
// automatically.
func (d *Daemon) signatureWatcher() {
defer d.wg.Done()
w := newSigWatcher(
func() *config.Config { return d.currentCfg() },
store.Global,
&d.forceFullRescan,
d.alertCh,
)
ticker := time.NewTicker(w.interval)
defer ticker.Stop()
// Initial tick on start so the watcher converges quickly when
// the daemon comes up shortly after an `update-rules` invocation.
w.tick()
for {
select {
case <-d.stopCh:
return
case <-ticker.C:
w.tick()
}
}
}
package daemon
import (
"fmt"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
)
// smtpIPEntry tracks failed-auth timestamps and suppression state for one IP.
// slowTimes/slowAccounts hold the long-horizon failure history that catches
// attackers pacing below the fast window; slowAccounts maps each targeted
// mailbox to its most recent failure so distinct-target counting can prune.
type smtpIPEntry struct {
times []time.Time
slowTimes []time.Time
slowAccounts map[string]time.Time
// slowLastSuccess is the most recent successful SMTP auth from this IP.
// A success inside the slow window marks the source as a live legitimate
// client (e.g. an office NAT where some devices still authenticate) and
// disqualifies the slow block.
slowLastSuccess time.Time
suppressed time.Time
lastSeen time.Time
}
// evictionRank keeps recent successful clients and meaningful accumulated
// evidence behind disposable one-shot failure entries. This prevents a fresh
// one-failure-per-IP flood from evicting legitimate source history.
func (e *smtpIPEntry) evictionRank(now time.Time, window, slowWindow time.Duration) int {
if slowWindow > 0 {
cutoff := now.Add(-slowWindow)
if (!e.slowLastSuccess.IsZero() && !e.slowLastSuccess.Before(cutoff)) || countTimesAtOrAfter(e.slowTimes, cutoff) > 1 {
return evictionRankProtectedIP
}
}
if hasTimeAtOrAfter(e.times, now.Add(-window)) {
return evictionRankActiveIP
}
return evictionRankIdleIP
}
// smtpSubnetEntry tracks unique attacker IPs within a /24.
type smtpSubnetEntry struct {
ips map[string]time.Time // ip -> firstSeen in window
suppressed time.Time
lastSeen time.Time
}
// smtpAccountEntry tracks unique attacker IPs per mailbox.
type smtpAccountEntry struct {
ips map[string]time.Time
suppressed time.Time
lastSeen time.Time
}
// slowBruteMinAccounts is how many distinct mailboxes a single IP's
// long-horizon failures must target before the slow-brute signal fires. A
// misconfigured client with a stale saved password hammers one mailbox; a
// mailbox walk touches several.
const slowBruteMinAccounts = config.SlowBruteMinThreshold
// Slow-brute state is nested inside a tracked IP, so maxTracked cannot bound
// it. These caps match the largest accepted threshold and keep a high-rate
// source or a stream of unique attacker-controlled mailbox names bounded.
const (
slowBruteMaxTimesPerIP = config.SlowBruteMaxThreshold
slowBruteMaxAccountsPerIP = config.SlowBruteMaxThreshold
)
// slowBruteWalkAccounts fires the slow signal on distinct-mailbox breadth
// alone. A walk probing one or two passwords per mailbox stays under any
// failure-count floor forever (observed live: 34 failures across 26
// mailboxes), but no legitimate client fails against this many distinct
// mailboxes from one address without a single success.
const slowBruteWalkAccounts = 10
// pruneSlowAccounts drops per-mailbox last-failure records older than cutoff.
func pruneSlowAccounts(accounts map[string]time.Time, cutoff time.Time) {
for acct, ts := range accounts {
if ts.Before(cutoff) {
delete(accounts, acct)
}
}
}
// appendSlowFailure retains the newest bounded history. Validation caps the
// configured threshold at the same value, so a reachable threshold is never
// discarded. Reslicing avoids copying the full history on every event.
func appendSlowFailure(times []time.Time, ts time.Time) []time.Time {
times = append(times, ts)
if len(times) > slowBruteMaxTimesPerIP {
times = times[len(times)-slowBruteMaxTimesPerIP:]
}
return times
}
// recordSlowAccount prunes before insertion and refuses new keys after the
// per-IP cap. Three distinct live keys are sufficient for detection, so a
// capped map preserves the signal while preventing unique-name memory growth.
// The bool reports whether this event's account is represented in the map.
func recordSlowAccount(accounts map[string]time.Time, account string, now, cutoff time.Time) (map[string]time.Time, bool) {
pruneSlowAccounts(accounts, cutoff)
if account == "" {
return accounts, false
}
if accounts == nil {
accounts = make(map[string]time.Time)
}
if _, exists := accounts[account]; exists || len(accounts) < slowBruteMaxAccountsPerIP {
accounts[account] = now
return accounts, true
}
return accounts, false
}
// smtpAuthTracker aggregates dovecot auth-failure events into three
// detection signals: per-IP brute force, per-/24 password spray, and
// per-mailbox account spray.
//
// Thread-safe; Record may be called concurrently from multiple log readers.
type smtpAuthTracker struct {
mu sync.Mutex
perIPThreshold int
subnetThreshold int
accountSprayThreshold int
window time.Duration
suppression time.Duration
slowThreshold int
slowWindow time.Duration
maxTracked int
now func() time.Time
ips map[string]*smtpIPEntry
subnets map[string]*smtpSubnetEntry
accounts map[string]*smtpAccountEntry
// Diagnostic counters (guarded by mu): cumulative Record invocations and
// findings emitted. The daemon logs these so a "zero smtp_bruteforce in
// production despite thousands of auth failures" can be pinned to either
// "Record never called" or "called but never crosses threshold".
recordCalls int64
findingsEmitted int64
// backendDownFn, when set, reports whether the active socket probe currently
// sees the mail auth backend down. SMTP-AUTH (exim->dovecot) fails the same
// way during a cpdoveauthd outage, so suppress brute/subnet auto-block then.
backendDownFn func() bool
}
// SetBackendDownCheck installs the active-probe callback the tracker consults to
// learn whether the mail auth backend is down. When it returns true, brute-force
// and subnet auto-block are suppressed. Set once at startup before log readers
// begin.
func (t *smtpAuthTracker) SetBackendDownCheck(fn func() bool) {
t.mu.Lock()
defer t.mu.Unlock()
t.backendDownFn = fn
}
// newSMTPAuthTracker constructs a tracker. `now` is injected so tests can
// use deterministic clocks; pass `time.Now` in production.
func newSMTPAuthTracker(
perIPThreshold int,
subnetThreshold int,
accountSprayThreshold int,
window time.Duration,
suppression time.Duration,
slowThreshold int,
slowWindow time.Duration,
maxTracked int,
now func() time.Time,
) *smtpAuthTracker {
if now == nil {
now = time.Now
}
return &smtpAuthTracker{
perIPThreshold: perIPThreshold,
subnetThreshold: subnetThreshold,
accountSprayThreshold: accountSprayThreshold,
window: window,
suppression: suppression,
slowThreshold: slowThreshold,
slowWindow: slowWindow,
maxTracked: maxTracked,
now: now,
ips: make(map[string]*smtpIPEntry),
subnets: make(map[string]*smtpSubnetEntry),
accounts: make(map[string]*smtpAccountEntry),
}
}
// Size returns the total number of tracked entities (IPs + subnets + accounts).
func (t *smtpAuthTracker) Size() int {
t.mu.Lock()
defer t.mu.Unlock()
return len(t.ips) + len(t.subnets) + len(t.accounts)
}
// Record processes one dovecot auth-failure observation. Returns zero or more
// findings that callers should append to their finding slice.
//
// ip MUST be non-private, non-loopback, and non-infra — callers enforce this
// before invoking Record.
func (t *smtpAuthTracker) Record(ip, account string) []alert.Finding {
if ip == "" {
return nil
}
t.mu.Lock()
defer t.mu.Unlock()
t.recordCalls++
now := t.now()
cutoff := now.Add(-t.window)
var findings []alert.Finding
// During a mail-auth-backend outage exim->dovecot SMTP AUTH fails for every
// user regardless of password, so the failure-rate signals would mass-block
// legitimate senders. Suppress them while the probe reports the backend down.
degraded := t.backendDownFn != nil && t.backendDownFn()
// --- Per-IP tracker ---
e, ok := t.ips[ip]
if !ok {
e = &smtpIPEntry{}
t.ips[ip] = e
}
e.times = pruneTimes(e.times, cutoff)
e.times = append(e.times, now)
e.lastSeen = now
if t.perIPThreshold > 0 && len(e.times) >= t.perIPThreshold && !now.Before(e.suppressed) && !degraded {
e.suppressed = now.Add(t.suppression)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "smtp_bruteforce",
Message: fmt.Sprintf("SMTP brute force from %s: %d failed auths in %v",
ip, len(e.times), t.window),
Details: "Real-time detection of dovecot_login auth failures",
Timestamp: now,
SourceIP: ip,
})
}
// --- Long-horizon slow-brute tracker ---
// Catches attackers pacing below the fast window (e.g. one failure every
// few minutes for hours). Requiring several distinct target mailboxes
// separates a mailbox walk from a misconfigured client retrying one stale
// saved password, which fails against a single mailbox no matter how long
// it runs.
if t.slowThreshold > 0 && t.slowWindow > 0 {
if degraded {
// Backend failures say nothing about credentials. Discard the
// long-lived evidence so outage traffic cannot trigger a delayed
// block as soon as the backend recovers.
e.slowTimes = nil
e.slowAccounts = nil
} else {
slowCutoff := now.Add(-t.slowWindow)
e.slowTimes = pruneTimes(e.slowTimes, slowCutoff)
e.slowTimes = appendSlowFailure(e.slowTimes, now)
e.slowAccounts, _ = recordSlowAccount(e.slowAccounts, normalizeMailAuthAccount(account), now, slowCutoff)
if e.slowLastSuccess.Before(slowCutoff) {
e.slowLastSuccess = time.Time{}
}
if (len(e.slowTimes) >= t.slowThreshold || len(e.slowAccounts) >= slowBruteWalkAccounts) &&
len(e.slowAccounts) >= slowBruteMinAccounts &&
e.slowLastSuccess.IsZero() &&
!now.Before(e.suppressed) {
e.suppressed = now.Add(t.suppression)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "smtp_bruteforce",
Message: fmt.Sprintf("SMTP brute force from %s: %d failed auths across %d mailboxes in %v",
ip, len(e.slowTimes), len(e.slowAccounts), t.slowWindow),
Details: "Long-horizon detection of paced dovecot_login auth failures that stay below the fast per-IP window",
Timestamp: now,
SourceIP: ip,
})
}
}
}
// --- Per-/24 subnet tracker (IPv4 only) ---
if prefix := extractPrefix24Daemon(ip); prefix != "" {
s, ok := t.subnets[prefix]
if !ok {
s = &smtpSubnetEntry{ips: make(map[string]time.Time)}
t.subnets[prefix] = s
}
pruneSubnetIPs(s, cutoff)
s.ips[ip] = now
s.lastSeen = now
if t.subnetThreshold > 0 && len(s.ips) >= t.subnetThreshold && !now.Before(s.suppressed) && !degraded {
s.suppressed = now.Add(t.suppression)
cidr := prefix + ".0/24"
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "smtp_subnet_spray",
Message: fmt.Sprintf("SMTP password spray from %s.0/24: %d unique IPs in %v",
prefix, len(s.ips), t.window),
Details: "Real-time detection of dovecot_login auth failures from many IPs in one /24",
Timestamp: now,
SourceIP: cidr,
})
}
}
// --- Per-account spray tracker ---
// Keyed on the trimmed, lower-cased mailbox: "User@X.RO", "user@x.ro"
// and a padded spelling are one target, and a spray across those
// variants must reach the distinct-source threshold as one account.
// (Only the tracker key is folded; the local part keeps its case
// everywhere it is displayed.)
account = strings.ToLower(strings.TrimSpace(account))
if account != "" {
a, ok := t.accounts[account]
if !ok {
a = &smtpAccountEntry{ips: make(map[string]time.Time)}
t.accounts[account] = a
}
pruneAccountIPs(a, cutoff)
a.ips[ip] = now
a.lastSeen = now
if t.accountSprayThreshold > 0 && len(a.ips) >= t.accountSprayThreshold && !now.Before(a.suppressed) {
a.suppressed = now.Add(t.suppression)
_, acctDomain := alert.SplitEmail(account)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "smtp_account_spray",
Message: fmt.Sprintf("SMTP password spray targeting %s: %d unique IPs in %v",
account, len(a.ips), t.window),
Details: "Distributed login attempts across many IPs against one mailbox (visibility only — no auto-block).",
Timestamp: now,
SourceIP: ip,
Domain: acctDomain,
Mailbox: account,
})
}
}
t.enforceMaxTracked(ip)
t.findingsEmitted += int64(len(findings))
return findings
}
// RecordSuccess notes a successful SMTP authentication from ip, disqualifying
// the source from the slow-brute block for as long as the success stays inside
// the slow window. Callers filter infra/private/loopback IPs first.
func (t *smtpAuthTracker) RecordSuccess(ip string) {
if ip == "" {
return
}
t.mu.Lock()
defer t.mu.Unlock()
// An explicitly disabled slow detector must not retain successful-client
// state for every authenticated SMTP delivery on the host.
if t.slowThreshold <= 0 || t.slowWindow <= 0 {
return
}
e, ok := t.ips[ip]
if !ok {
e = &smtpIPEntry{}
t.ips[ip] = e
}
now := t.now()
e.slowLastSuccess = now
e.lastSeen = now
t.enforceMaxTracked(ip)
}
// Stats returns cumulative Record invocations and findings emitted since
// startup. Used by the daemon's periodic diagnostic log.
func (t *smtpAuthTracker) Stats() (calls, emits int64) {
t.mu.Lock()
defer t.mu.Unlock()
return t.recordCalls, t.findingsEmitted
}
// pruneTimes drops timestamps older than cutoff. Reuses the backing array.
func pruneTimes(times []time.Time, cutoff time.Time) []time.Time {
recent := times[:0]
for _, ts := range times {
if !ts.Before(cutoff) {
recent = append(recent, ts)
}
}
return recent
}
// extractPrefix24Daemon returns the first three octets of an IPv4 address as
// "a.b.c", or "" if the input isn't an IPv4 address in dotted-quad form.
func extractPrefix24Daemon(ip string) string {
parts := 0
end := 0
for i := 0; i < len(ip); i++ {
if ip[i] == '.' {
parts++
if parts == 3 {
end = i
break
}
}
}
if parts != 3 {
return ""
}
// Reject IPv6 mapped or containing colons.
for i := 0; i < end; i++ {
if ip[i] == ':' {
return ""
}
}
return ip[:end]
}
// pruneSubnetIPs drops per-/24 IP entries whose last-seen is older than cutoff.
func pruneSubnetIPs(s *smtpSubnetEntry, cutoff time.Time) {
for ip, ts := range s.ips {
if ts.Before(cutoff) {
delete(s.ips, ip)
}
}
}
// pruneAccountIPs drops per-account IP entries whose last-seen is older than cutoff.
func pruneAccountIPs(a *smtpAccountEntry, cutoff time.Time) {
for ip, ts := range a.ips {
if ts.Before(cutoff) {
delete(a.ips, ip)
}
}
}
// Purge removes entries with no recent activity (older than window + suppression).
// Called from a background goroutine every minute.
func (t *smtpAuthTracker) Purge() {
t.mu.Lock()
defer t.mu.Unlock()
now := t.now()
activityCutoff := now.Add(-(t.window + t.suppression))
slowCutoff := now.Add(-t.slowWindow)
for k, e := range t.ips {
e.times = pruneTimes(e.times, now.Add(-t.window))
e.slowTimes = pruneTimes(e.slowTimes, slowCutoff)
pruneSlowAccounts(e.slowAccounts, slowCutoff)
if e.slowLastSuccess.Before(slowCutoff) {
e.slowLastSuccess = time.Time{}
}
if len(e.times) == 0 && len(e.slowTimes) == 0 &&
e.slowLastSuccess.IsZero() && !e.lastSeen.After(activityCutoff) {
delete(t.ips, k)
}
}
for k, s := range t.subnets {
pruneSubnetIPs(s, now.Add(-t.window))
if len(s.ips) == 0 && !s.lastSeen.After(activityCutoff) {
delete(t.subnets, k)
}
}
for k, a := range t.accounts {
pruneAccountIPs(a, now.Add(-t.window))
if len(a.ips) == 0 && !a.lastSeen.After(activityCutoff) {
delete(t.accounts, k)
}
}
}
// enforceMaxTracked evicts the least-recently-seen entries until the total
// number of tracked entities (IPs + subnets + accounts) is <= maxTracked.
// Caller must hold t.mu.
// keepIP is never evicted; see the mail tracker for why the entry a caller
// just wrote must survive its own bookkeeping.
func (t *smtpAuthTracker) enforceMaxTracked(keepIP string) {
total := len(t.ips) + len(t.subnets) + len(t.accounts)
if total <= t.maxTracked {
return
}
now := t.now()
// Ranked like the mail tracker: attacker-chosen account and subnet keys
// go first, idle sources next, sources with in-window failures or
// slow-brute evidence last, so a flood of unique mailbox names cannot
// evict the evidence that governs a source's auto-block decision.
type victim struct {
kind string // "ip" | "subnet" | "account"
key string
seen time.Time
rank int
}
victims := make([]victim, 0, total)
for k, v := range t.ips {
if k == keepIP {
continue
}
victims = append(victims, victim{"ip", k, v.lastSeen, v.evictionRank(now, t.window, t.slowWindow)})
}
for k, v := range t.subnets {
victims = append(victims, victim{"subnet", k, v.lastSeen, evictionRankSubnet})
}
for k, v := range t.accounts {
victims = append(victims, victim{"account", k, v.lastSeen, evictionRankAccount})
}
sort.Slice(victims, func(i, j int) bool {
if victims[i].rank != victims[j].rank {
return victims[i].rank < victims[j].rank
}
return victims[i].seen.Before(victims[j].seen)
})
// Evict to 95% of cap so subsequent inserts don't re-trigger the sort.
target := t.maxTracked * 95 / 100
for i := 0; i < len(victims); i++ {
if len(t.ips)+len(t.subnets)+len(t.accounts) <= target {
break
}
v := victims[i]
switch v.kind {
case "ip":
delete(t.ips, v.key)
case "subnet":
delete(t.subnets, v.key)
case "account":
delete(t.accounts, v.key)
}
}
}
package daemon
import (
"fmt"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/eximlog"
)
// smtpProbeBlockExpiryString returns the configured block expiry string when
// auto-response will actually block the source IP for an `smtp_probe_abuse`
// finding (auto_response.enabled AND block_ips true AND dry_run false), or ""
// otherwise.
// The returned string is what the operator put in csm.yaml ("24h", "12h"...)
// so the alert text matches the value they configured rather than Go's
// canonical Duration formatting.
func smtpProbeBlockExpiryString() string {
cfg := config.Active()
if cfg == nil {
return ""
}
if !cfg.AutoResponse.Enabled || !cfg.AutoResponse.BlockIPs || cfg.AutoResponseDryRunEnabled() {
return ""
}
if cfg.AutoResponse.BlockExpiry == "" {
return "24h"
}
return cfg.AutoResponse.BlockExpiry
}
// smtpProbeEntry records connect timestamps and suppression for one IP.
type smtpProbeEntry struct {
times []time.Time
suppressed time.Time
lastSeen time.Time
}
// smtpProbeTracker counts raw SMTP connect events per source IP and emits an
// `smtp_probe_abuse` finding when an IP exceeds the threshold inside the
// rolling window.
//
// This is the connection-rate complement to smtpAuthTracker: scanners that
// probe-and-disconnect (no AUTH attempt) never trigger the auth tracker, so
// they need their own signal. The thresholds are deliberately set well above
// any legitimate MUA usage; Thunderbird/iPhone bursts of 10-15 parallel
// sessions per send fall comfortably under, scanner storms with hundreds of
// connect/min are caught.
type smtpProbeTracker struct {
mu sync.Mutex
threshold int
window time.Duration
suppression time.Duration
maxTracked int
now func() time.Time
// expiryStrFn returns the operator-visible block expiry (e.g. "24h") when
// live auto-blocking is enabled, or "" when no auto-block will run. Read
// at finding time so a SIGHUP reload of auto_response.* is reflected in
// the next emitted finding's Details.
expiryStrFn func() string
ips map[string]*smtpProbeEntry
}
func newSMTPProbeTracker(threshold int, window, suppression time.Duration, maxTracked int, now func() time.Time, expiryStrFn func() string) *smtpProbeTracker {
if now == nil {
now = time.Now
}
return &smtpProbeTracker{
threshold: threshold,
window: window,
suppression: suppression,
maxTracked: maxTracked,
now: now,
expiryStrFn: expiryStrFn,
ips: make(map[string]*smtpProbeEntry),
}
}
// Size returns the number of tracked source IPs.
func (t *smtpProbeTracker) Size() int {
t.mu.Lock()
defer t.mu.Unlock()
return len(t.ips)
}
// Record observes one SMTP connect event. ip MUST be non-private,
// non-loopback, and non-infra. Callers enforce this before invoking Record.
// Returns zero or one finding (no per-call multiplication).
func (t *smtpProbeTracker) Record(ip string) []alert.Finding {
if ip == "" {
return nil
}
t.mu.Lock()
defer t.mu.Unlock()
// threshold is read under the lock: it is mutable via SetThresholds on a
// SIGHUP reload, so a lockless fast-path read here would race the swap.
if t.threshold <= 0 {
return nil
}
now := t.now()
cutoff := now.Add(-t.window)
e, ok := t.ips[ip]
if !ok {
e = &smtpProbeEntry{}
t.ips[ip] = e
}
e.times = pruneTimes(e.times, cutoff)
e.times = append(e.times, now)
e.lastSeen = now
var findings []alert.Finding
if len(e.times) >= t.threshold && !now.Before(e.suppressed) {
e.suppressed = now.Add(t.suppression)
// The Details message is computed here, before AutoBlockIPs runs in
// dispatchBatch. We can only report the *intent* (scheduled for
// auto-block) - the actual outcome (blocked / rate-limited / already
// blocked / challenged) is published by the companion `auto_block`
// finding emitted by checks.AutoBlockIPs in the same batch.
details := "Sustained SMTP connect rate above the configured threshold. Likely scanner / dictionary probe;"
if t.expiryStrFn != nil {
if exp := t.expiryStrFn(); exp != "" {
details += fmt.Sprintf(" scheduled for auto-block (%s).", exp)
} else {
details += " consider manual block."
}
} else {
details += " consider manual block."
}
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "smtp_probe_abuse",
Message: fmt.Sprintf("SMTP probe abuse from %s: %d connections in %v",
ip, len(e.times), t.window),
Details: details,
Timestamp: now,
SourceIP: ip,
})
}
t.enforceMaxTracked()
return findings
}
// Purge removes IPs with no activity since (window + suppression) ago.
// Called periodically to prevent unbounded growth.
func (t *smtpProbeTracker) Purge() {
t.mu.Lock()
defer t.mu.Unlock()
now := t.now()
activityCutoff := now.Add(-(t.window + t.suppression))
for k, e := range t.ips {
e.times = pruneTimes(e.times, now.Add(-t.window))
if len(e.times) == 0 && !e.lastSeen.After(activityCutoff) {
delete(t.ips, k)
}
}
}
// enforceMaxTracked evicts the least-recently-seen IPs to keep memory bounded.
// Caller must hold t.mu.
func (t *smtpProbeTracker) enforceMaxTracked() {
if t.maxTracked <= 0 || len(t.ips) <= t.maxTracked {
return
}
type victim struct {
key string
seen time.Time
}
victims := make([]victim, 0, len(t.ips))
for k, e := range t.ips {
victims = append(victims, victim{k, e.lastSeen})
}
sort.Slice(victims, func(i, j int) bool { return victims[i].seen.Before(victims[j].seen) })
target := t.maxTracked * 95 / 100
for i := 0; i < len(victims) && len(t.ips) > target; i++ {
delete(t.ips, victims[i].key)
}
}
// parseEximSMTPConnectIP extracts the connecting source IP from an exim
// mainlog "SMTP connection from ..." line. Returns "" when the line is not
// a connect event.
//
// Exim formats vary:
//
// SMTP connection from [1.2.3.4]:65417 (TCP/IP connection count = 7)
// SMTP connection from (helo.example.com) [1.2.3.4]:43018 lost D=5s
// SMTP connection from ([helo-as-ip]) [1.2.3.4]:38294 lost D=15s
// SMTP connection from ([192.168.0.94]) [1.2.3.4]:64547 D=5s closed by QUIT
//
// Client attribution is delegated to eximlog so HELO address literals and
// malformed bracketed tokens follow the same rules as every other consumer.
func parseEximSMTPConnectIP(line string) string {
const marker = "SMTP connection from "
markerStart := strings.Index(line, marker)
if markerStart < 0 {
return ""
}
if hStart, ok := eximlog.HFieldStart(line); ok && hStart-len(" H=") < markerStart {
return ""
}
return eximlog.ClientIP(line)
}
package daemon
import (
"os"
"path/filepath"
"sort"
)
// eximSplitSpoolDirName reports whether name is one of the hash
// subdirectories Exim creates under input/ when split_spool_directory is
// on: a single character from [0-9A-Za-z], taken from the message ID.
func eximSplitSpoolDirName(name string) bool {
if len(name) != 1 {
return false
}
c := name[0]
return (c >= '0' && c <= '9') || (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z')
}
// spoolMarkTargets returns the directories a spool watcher has to mark to
// see every message file under root: the root itself (flat spool) and every
// split-spool hash subdirectory present right now, sorted. A missing root
// yields nil. Exim creates hash directories lazily, so callers re-run this
// periodically and mark whatever is new.
func spoolMarkTargets(root string) []string {
entries, err := os.ReadDir(root)
if err != nil {
return nil
}
targets := []string{root}
for _, e := range entries {
if e.IsDir() && eximSplitSpoolDirName(e.Name()) {
targets = append(targets, filepath.Join(root, e.Name()))
}
}
sort.Strings(targets[1:])
return targets
}
//go:build linux
package daemon
import (
"errors"
"fmt"
"os"
"path/filepath"
"runtime/debug"
"strings"
"sync"
"sync/atomic"
"syscall"
"time"
"unsafe"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailav"
"github.com/pidginhost/csm/internal/metrics"
emime "github.com/pidginhost/csm/internal/mime"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/queuehealth"
)
// spoolQueueOverflowTotal counts FAN_Q_OVERFLOW records on the spool watcher.
// Registered once per process; hand-built test watchers that skip
// registerSpoolMetrics leave it nil, so callers must nil-check.
var (
spoolQueueOverflowTotal *metrics.Counter
spoolMetricsOnce sync.Once
)
func registerSpoolMetrics() {
spoolMetricsOnce.Do(func() {
spoolQueueOverflowTotal = metrics.NewCounter(
"csm_spool_fanotify_queue_overflow_total",
"FAN_Q_OVERFLOW events on the Exim spool watcher: the kernel queue filled and dropped spool open events. In permission mode the kernel allows the overflowing opens, so mail may be delivered without an AV scan. Any sustained growth needs investigation.",
)
metrics.MustRegister("csm_spool_fanotify_queue_overflow_total", spoolQueueOverflowTotal)
})
}
// fanotify constants for permission events (not in Go stdlib).
const (
FAN_CLASS_CONTENT = 0x00000004
FAN_OPEN_PERM = 0x00010000
FAN_ALLOW = 0x01
FAN_DENY = 0x02
FAN_EVENT_ON_CHILD = 0x08000000
)
// fanotifyResponse is the struct written back to the fanotify fd
// to allow or deny a permission event.
type fanotifyResponse struct {
Fd int32
Response uint32
}
const responseSize = int(unsafe.Sizeof(fanotifyResponse{}))
// SpoolWatcher monitors Exim spool directories for new messages using a
// dedicated fanotify instance with permission events (FAN_OPEN_PERM).
// It is completely separate from the FileMonitor.
type SpoolWatcher struct {
fd int
cfg *config.Config
alertCh chan<- alert.Finding
orchestrator *emailav.Orchestrator
quarantine *emailav.Quarantine
permissionMode bool // true if using FAN_OPEN_PERM, false if fallback to FAN_CLOSE_WRITE
// eventMask is the fanotify mask every spool directory is marked with.
// spoolRoots are the Exim input directories found at start; marked
// records each pathname seen so periodic rescans can report newly found
// split-spool hash subdirectories without growing beyond the fixed Exim
// hash alphabet. Every rescan still refreshes the kernel mark because an
// unlink drops the inode-bound mark while a recreated path stays recorded.
eventMask uint64
spoolRoots []string
markedMu sync.Mutex
marked map[string]struct{}
// emailAVTempDir is the staging directory CreateTemp uses for
// extracted attachments. Established once at watcher construction
// (0700, daemon-owned) so an unprivileged local uid cannot race
// the scanner via /tmp symlink swaps. Empty is only for hand-built
// test watchers; production construction requires state_path.
emailAVTempDir string
// selfPID is os.Getpid(). fanotify permission events generated by CSM's
// own reads of the spool -D file (during MIME parsing) must be allowed
// without re-scanning, or the scan re-enters fanotify and deadlocks every
// worker. See dispatchEvent.
selfPID int
scanCh chan spoolEvent
pipeFds [2]int
stopOnce sync.Once
drainOnce sync.Once
stopCh chan struct{}
wg sync.WaitGroup
pipeClosed int32 // atomic
fdClosed int32 // atomic - guards sw.fd against double-close
runActive int32 // atomic - Run owns scanCh shutdown while set
degradedMu sync.Mutex
lastDegradedAt time.Time
encryptedMu sync.Mutex
lastEncryptedAt time.Time
panicMu sync.Mutex
lastPanicAt time.Time
// queueOverflows counts FAN_Q_OVERFLOW records. In permission mode a kernel
// queue overflow means opens were let through without a scan verdict, so
// mail may have been delivered unscanned. overflowMu rate-limits the
// operator finding so a storm does not flood the alert channel.
queueOverflows int64 // atomic
overflowMu sync.Mutex
lastOverflowAt time.Time
// holds tracks hold-budget expiries and decides when to stop holding mail.
holds holdWatchdog
queueHealthOnce sync.Once
scannerHealth *queuehealth.Tracker
kernelQueueHealth *queuehealth.Tracker
kernelQueue *notificationQueue
}
type spoolEvent struct {
queueTicket queuehealth.Ticket
path string
fd int // fanotify event fd (for permission response)
pid int32
needResp bool // true if permission event requiring response
// guard bounds how long this open stays suspended. Nil for events built
// outside dispatchEvent, which answer directly from needResp.
guard *holdGuard
}
// finish hands the kernel a verdict for this event, once, and closes the
// event fd. The guard may already have answered when the scan outran the
// hold budget; then this verdict only affects quarantine, not delivery.
func (evt spoolEvent) finish(sw *SpoolWatcher, response uint32) {
switch {
case evt.guard != nil:
evt.guard.finish(response)
return
case evt.needResp:
spoolWriteResponse(sw, int32(evt.fd), response) // #nosec G115 -- POSIX fd fits in int32.
}
_ = unix.Close(evt.fd)
}
// NewSpoolWatcher creates a dedicated fanotify instance for Exim spool scanning.
// Attempts FAN_CLASS_CONTENT with FAN_OPEN_PERM first; falls back to
// FAN_CLASS_NOTIF with FAN_CLOSE_WRITE if permission events are unavailable.
func NewSpoolWatcher(cfg *config.Config, alertCh chan<- alert.Finding, orch *emailav.Orchestrator, quar *emailav.Quarantine) (*SpoolWatcher, error) {
emailAVTempDir, err := resolveEmailAVTempDir(cfg)
if err != nil {
return nil, err
}
registerSpoolMetrics()
sw := &SpoolWatcher{
cfg: cfg,
alertCh: alertCh,
orchestrator: orch,
quarantine: quar,
emailAVTempDir: emailAVTempDir,
selfPID: os.Getpid(),
scanCh: make(chan spoolEvent, 256),
stopCh: make(chan struct{}),
}
// Try permission-capable class first
fd, err := unix.FanotifyInit(FAN_CLASS_CONTENT|FAN_CLOEXEC|FAN_NONBLOCK, unix.O_RDONLY)
if err == nil {
sw.fd = fd
sw.permissionMode = true
fmt.Fprintf(os.Stderr, "[%s] spool watcher: permission events enabled (FAN_OPEN_PERM)\n", ts())
} else {
// Fallback to notification-only
fd, err = unix.FanotifyInit(FAN_CLASS_NOTIF|FAN_CLOEXEC|FAN_NONBLOCK, unix.O_RDONLY)
if err != nil {
return nil, fmt.Errorf("fanotify_init: %w (neither permission nor notification mode available)", err)
}
sw.fd = fd
sw.permissionMode = false
fmt.Fprintf(os.Stderr, "[%s] spool watcher: WARNING - permission events unavailable, using notification mode (small delivery race window possible)\n", ts())
if sw.cfg.EmailAV.FailMode == "tempfail" {
fmt.Fprintf(os.Stderr, "[%s] spool watcher: WARNING - fail_mode=tempfail requested but cannot be honoured without FAN_OPEN_PERM; operating fail-open\n", ts())
}
}
// Mark spool directories
spoolDirs := []string{"/var/spool/exim/input", "/var/spool/exim4/input"}
var eventMask uint64
if sw.permissionMode {
eventMask = FAN_OPEN_PERM | FAN_EVENT_ON_CHILD
} else {
eventMask = FAN_CLOSE_WRITE | FAN_EVENT_ON_CHILD
}
sw.eventMask = eventMask
sw.marked = make(map[string]struct{})
marked := 0
for _, dir := range spoolDirs {
if _, err := os.Stat(dir); err != nil {
continue
}
sw.spoolRoots = append(sw.spoolRoots, dir)
marked += sw.markSpoolTargets(dir)
}
if marked == 0 {
_ = unix.Close(sw.fd)
return nil, fmt.Errorf("no Exim spool directories found to watch")
}
// Create pipe for stop signaling
if err := unix.Pipe2(sw.pipeFds[:], unix.O_NONBLOCK|unix.O_CLOEXEC); err != nil {
_ = unix.Close(sw.fd)
return nil, fmt.Errorf("creating pipe: %w", err)
}
return sw, nil
}
// spoolRescanInterval bounds how long a split-spool hash directory Exim
// created after start stays unwatched.
var spoolRescanInterval = time.Minute
// markSpoolTargets marks root and every split-spool hash subdirectory under
// it, and returns how many pathnames were marked for the first time.
// FAN_EVENT_ON_CHILD on a directory mark covers only its
// direct children, so on a split spool (the cPanel default) the -D files,
// which live one level down, are only seen through the subdirectory marks.
// Uses FAN_MARK_ADD (not FAN_MARK_MOUNT) to scope to the directory.
func (sw *SpoolWatcher) markSpoolTargets(root string) int {
added := 0
for _, dir := range spoolMarkTargets(root) {
sw.markedMu.Lock()
_, done := sw.marked[dir]
sw.markedMu.Unlock()
if err := unix.FanotifyMark(sw.fd, FAN_MARK_ADD, sw.eventMask, -1, dir); err != nil {
fmt.Fprintf(os.Stderr, "[%s] spool watcher: cannot watch %s: %v\n", ts(), dir, err)
continue
}
if !done {
sw.markedMu.Lock()
sw.marked[dir] = struct{}{}
sw.markedMu.Unlock()
added++
fmt.Fprintf(os.Stderr, "[%s] spool watcher: watching %s\n", ts(), dir)
}
}
return added
}
// rescanSpoolDirs marks split-spool hash directories that appeared since
// the last pass.
func (sw *SpoolWatcher) rescanSpoolDirs() {
for _, root := range sw.spoolRoots {
sw.markSpoolTargets(root)
}
}
// Run starts the event loop and scanner workers. Blocks until Stop() is called.
func (sw *SpoolWatcher) Run() {
// Event loop. Set up epoll before starting any workers so a setup
// failure has nothing to unwind beyond drainAndClose: workers started
// first would park on scanCh forever and hang daemon shutdown.
epfd, err := unix.EpollCreate1(unix.EPOLL_CLOEXEC)
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] spool watcher: epoll_create: %v\n", ts(), err)
sw.drainAndClose()
return
}
defer func() { _ = unix.Close(epfd) }()
// #nosec G115 -- POSIX fd fits in int32 (rlimit ~1024). Same for all fd→int32 in this file.
if err := unix.EpollCtl(epfd, unix.EPOLL_CTL_ADD, sw.fd, &unix.EpollEvent{Events: unix.EPOLLIN, Fd: int32(sw.fd)}); err != nil {
fmt.Fprintf(os.Stderr, "[%s] spool watcher: epoll_ctl(fanotify fd): %v\n", ts(), err)
sw.drainAndClose()
return
}
// #nosec G115 -- POSIX fd fits in int32.
if err := unix.EpollCtl(epfd, unix.EPOLL_CTL_ADD, sw.pipeFds[0], &unix.EpollEvent{Events: unix.EPOLLIN, Fd: int32(sw.pipeFds[0])}); err != nil {
fmt.Fprintf(os.Stderr, "[%s] spool watcher: epoll_ctl(pipe fd): %v\n", ts(), err)
sw.drainAndClose()
return
}
// Start scanner workers
atomic.StoreInt32(&sw.runActive, 1)
defer atomic.StoreInt32(&sw.runActive, 0)
concurrency := sw.cfg.EmailAV.ScanConcurrency
if concurrency < 1 {
concurrency = 4
}
for i := 0; i < concurrency; i++ {
sw.wg.Add(1)
obs.Go("spool-scanner", sw.scanWorker)
}
events := make([]unix.EpollEvent, 16)
buf := make([]byte, 4096)
lastRescan := time.Now()
for {
select {
case <-sw.stopCh:
sw.drainAndClose()
return
default:
}
if time.Since(lastRescan) >= spoolRescanInterval {
sw.rescanSpoolDirs()
lastRescan = time.Now()
}
n, err := unix.EpollWait(epfd, events, 500)
if err != nil {
if err == unix.EINTR {
continue
}
select {
case <-sw.stopCh:
sw.drainAndClose()
return
default:
continue
}
}
for i := 0; i < n; i++ {
// #nosec G115 -- POSIX fd fits in int32.
if events[i].Fd == int32(sw.pipeFds[0]) {
sw.drainAndClose()
return
}
// #nosec G115 -- POSIX fd fits in int32.
if events[i].Fd == int32(sw.fd) {
sw.readEvents(buf)
}
}
}
}
func (sw *SpoolWatcher) readEvents(buf []byte) {
sw.initQueueHealth()
for {
n, err := sw.kernelQueue.read(buf, sw.parseEvents)
if err != nil || n < metadataSize {
return
}
}
}
// parseEvents walks a raw fanotify buffer and dispatches each record. A
// FAN_Q_OVERFLOW record (Fd == FAN_NOFD, FAN_Q_OVERFLOW bit set) is handled
// explicitly so a kernel queue overflow is not silently swallowed.
func (sw *SpoolWatcher) parseEvents(buf []byte) {
offset := 0
for offset+metadataSize <= len(buf) {
// #nosec G103 -- fanotify delivers a packed binary stream;
// reinterpretation is required and bounded by metadataSize above.
meta := (*fanotifyEventMetadata)(unsafe.Pointer(&buf[offset]))
if meta.EventLen < uint32(metadataSize) || int(meta.EventLen) > len(buf)-offset {
break
}
if meta.Mask&unix.FAN_Q_OVERFLOW != 0 {
sw.handleQueueOverflow()
} else if meta.Fd >= 0 {
sw.dispatchEvent(meta.Fd, meta.Pid)
}
offset += int(meta.EventLen)
}
}
// handleQueueOverflow reacts to a FAN_Q_OVERFLOW record on the spool watcher.
// In permission mode the kernel resolves opens it cannot queue by allowing
// them, so an overflow means some inbound messages were delivered without an
// AV scan. Count it and emit a rate-limited Warning that says so.
func (sw *SpoolWatcher) handleQueueOverflow() {
sw.initQueueHealth()
sw.kernelQueueHealth.Lose(time.Now(), 1)
atomic.AddInt64(&sw.queueOverflows, 1)
if spoolQueueOverflowTotal != nil {
spoolQueueOverflowTotal.Inc()
}
sw.overflowMu.Lock()
if !sw.lastOverflowAt.IsZero() && time.Since(sw.lastOverflowAt) < time.Minute {
sw.overflowMu.Unlock()
return
}
sw.lastOverflowAt = time.Now()
sw.overflowMu.Unlock()
sw.emitFinding("email_av_queue_overflow", alert.Warning,
"Spool fanotify queue overflowed: the kernel dropped spool open events. In permission mode overflowing opens are allowed through, so one or more messages may have been delivered unscanned during the storm. Investigate the mail spike and consider a manual rescan of recent messages.")
}
// dispatchEvent decides what to do with a single fanotify event: allow our own
// reads immediately, allow non -D opens, or hand the -D open to a scan worker.
// It owns the event fd: every path either enqueues it (worker closes it) or
// responds and closes it here.
func (sw *SpoolWatcher) dispatchEvent(fd int32, pid int32) {
// MAIL-02: never scan CSM's own reads. When a scan worker opens the -D
// file to parse it, that open re-enters this fanotify instance as a
// FAN_OPEN_PERM event for our own pid. With the scan blocked waiting on a
// verdict, self-recursion deadlocks the workers and stalls Exim delivery
// server-wide until a restart. fanotify(7) documents excluding getpid()
// for exactly this reason. Allow immediately, without scanning.
if int(pid) == sw.selfPID {
if sw.permissionMode {
sw.writeResponse(fd, FAN_ALLOW)
}
_ = unix.Close(int(fd))
return
}
path, err := os.Readlink(fmt.Sprintf("/proc/self/fd/%d", fd))
if err != nil {
// Must respond even on error paths (permission mode)
if sw.permissionMode {
sw.writeResponse(fd, FAN_ALLOW)
}
_ = unix.Close(int(fd))
return
}
// Only process *-D files (message body)
if !strings.HasSuffix(path, "-D") {
if sw.permissionMode {
sw.writeResponse(fd, FAN_ALLOW)
}
_ = unix.Close(int(fd))
return
}
// While the scanner is behind, release opens instead of queueing them.
if sw.dispatchBypass(fd, sw.permissionMode) {
return
}
// Hand to a scan worker. The kernel keeps the opener suspended meanwhile,
// so the guard's budget starts here: time spent queued counts against the
// same deadline as time spent scanning.
sw.initQueueHealth()
evt := spoolEvent{
queueTicket: sw.scannerHealth.Begin(time.Now()),
path: path,
fd: int(fd),
pid: pid,
needResp: sw.permissionMode,
guard: sw.newHoldGuard(fd, sw.permissionMode),
}
select {
case sw.scanCh <- evt:
// Worker will handle response and fd close
case <-sw.stopCh:
// Shutting down - allow and close
evt.queueTicket.Reject(time.Now())
evt.finish(sw, FAN_ALLOW)
default:
// Never wait in the reader: later events in this batch (including our
// own parser opens) are also suspended and do not yet have timers.
evt.queueTicket.Reject(time.Now())
evt.guard.expire()
evt.finish(sw, FAN_ALLOW)
}
}
// scanWorker processes spool events: MIME parse, scan, quarantine/allow.
func (sw *SpoolWatcher) scanWorker() {
defer sw.wg.Done()
// drainAndClose closes scanCh after shutdown. Keep draining queued events
// until then: dispatchEvent may enqueue a permission fd while Stop is racing,
// and the worker owns the response/close path for every enqueued fd.
for {
select {
case evt, ok := <-sw.scanCh:
if !ok {
return
}
sw.handleSpoolEventSafe(evt)
case <-sw.stopCh:
if atomic.LoadInt32(&sw.runActive) != 0 {
for evt := range sw.scanCh {
sw.handleSpoolEventSafe(evt)
}
return
}
for {
select {
case evt, ok := <-sw.scanCh:
if !ok {
return
}
sw.handleSpoolEventSafe(evt)
default:
return
}
}
}
}
}
// spoolEventHandler processes one queued spool event. Var so tests can
// substitute a panicking handler.
var spoolEventHandler = (*SpoolWatcher).handleSpoolEvent
// handleSpoolEventSafe contains panics and guarantees event cleanup even if a
// handler fails before installing its own defer. Shared guard ownership makes
// this fallback safe after a handler has already answered and closed the fd.
// Re-raising would restart the daemon, and Exim would redeliver the same
// message into the same panic.
func (sw *SpoolWatcher) handleSpoolEventSafe(evt spoolEvent) {
if evt.guard == nil {
evt.guard = &holdGuard{sw: sw, fd: int32(evt.fd), needResp: evt.needResp} // #nosec G115 -- POSIX fd fits in int32.
}
defer func() {
evt.finish(sw, FAN_ALLOW)
if r := recover(); r != nil {
sw.reportScannerPanic(evt.path, r)
}
}()
work := queuehealth.Work[spoolEvent]{Value: evt, Ticket: evt.queueTicket}
work.Process(func(queued spoolEvent) { spoolEventHandler(sw, queued) })
}
// reportScannerPanic logs the panic with its stack, forwards it to
// observability and raises a critical finding at most once per ten minutes.
func (sw *SpoolWatcher) reportScannerPanic(path string, r interface{}) {
obs.CaptureMsg("spool-scanner", fmt.Sprintf("panic scanning %s: %v", path, r))
fmt.Fprintf(os.Stderr, "[%s] spool watcher: recovered panic scanning %s: %v\n%s", ts(), path, r, debug.Stack())
sw.panicMu.Lock()
defer sw.panicMu.Unlock()
if !sw.lastPanicAt.IsZero() && time.Since(sw.lastPanicAt) < 10*time.Minute {
return
}
sw.lastPanicAt = time.Now()
sw.emitFinding("email_av_scanner_panic", alert.Critical,
fmt.Sprintf("Email AV scanner panicked on %s and let it through unscanned: %v", filepath.Base(path), r))
}
func (sw *SpoolWatcher) handleSpoolEvent(evt spoolEvent) {
// CRITICAL: deferred FAN_ALLOW - every code path must allow by default.
// Only overridden to FAN_DENY when policy requires deferral or quarantine.
response := uint32(FAN_ALLOW)
defer func() {
evt.finish(sw, response)
// finish joins the winning verdict; checking before it could miss a
// timer that wins concurrently, or mistake a tempfail for delivery.
released := evt.guard != nil && evt.guard.timedOut() && evt.guard.timeoutVerdict == FAN_ALLOW
if released && response == FAN_DENY {
sw.emitFinding("email_av_late_verdict", alert.Warning,
fmt.Sprintf("Message %s was allowed before its scan finished (hold budget %s) and the scan then asked to stop delivery. A copy may already have been delivered; see the scan and quarantine findings for the final outcome.",
strings.TrimSuffix(filepath.Base(evt.path), "-D"), spoolHoldBudget))
}
}()
// Derive message ID: strip -D suffix and directory
base := filepath.Base(evt.path)
msgID := strings.TrimSuffix(base, "-D")
spoolDir := filepath.Dir(evt.path)
headerPath := filepath.Join(spoolDir, msgID+"-H")
bodyPath := evt.path
// MIME parse - fail-open on error
limits := emime.Limits{
MaxAttachmentSize: sw.cfg.EmailAV.MaxAttachmentSize,
MaxArchiveDepth: sw.cfg.EmailAV.MaxArchiveDepth,
MaxArchiveFiles: sw.cfg.EmailAV.MaxArchiveFiles,
MaxExtractionSize: sw.cfg.EmailAV.MaxExtractionSize,
TempDir: sw.emailAVTempDir,
}
tempfail := sw.cfg.EmailAV.FailMode == "tempfail" && evt.needResp
extraction, err := emime.ParseSpoolMessage(headerPath, bodyPath, limits)
if err != nil {
// MAIL-03: Exim opens the -D body file before it writes the matching
// -H header file (the -H lands only after the message is accepted).
// fanotify delivers that reception-time open to us while the -H is
// legitimately absent. Allow silently: emitting a parse-error Warning
// here would fire once per inbound message, and deferring (tempfail)
// would defer 100% of inbound mail. Any ENOENT on a spool file is safe
// to allow -- a message whose spool files are gone cannot be scanned.
if errors.Is(err, os.ErrNotExist) {
return
}
fmt.Fprintf(os.Stderr, "[%s] spool watcher: MIME parse error for %s: %v\n", ts(), msgID, err)
sw.emitFinding("email_av_parse_error", alert.Warning, fmt.Sprintf("MIME parse failed for message %s: %v", msgID, err))
if tempfail {
response = FAN_DENY // tempfail: Exim retries later
}
return
}
// Clean up temp files when done
defer func() {
for _, p := range extraction.Parts {
os.Remove(p.TempPath)
}
}()
sw.emitEncryptedArchiveWarning(msgID, extraction.EncryptedEntries, extraction.EncryptedEntriesOmitted)
if extraction.Partial {
partialResult := &emailav.ScanResult{PartialExtraction: true}
if shouldTempfailEmailDelivery(tempfail, partialResult, nil) {
response = FAN_DENY
sw.emitDegradedWarning(fmt.Sprintf("Incomplete email attachment extraction for message %s (%s) - delivery deferred (tempfail mode)", msgID, partialExtractionReason(extraction)))
return
}
sw.emitDegradedWarning(fmt.Sprintf("Incomplete email attachment extraction for message %s (%s) - delivery allowed (fail-open mode)", msgID, partialExtractionReason(extraction)))
}
if len(extraction.Parts) == 0 {
return // No attachments to scan - allow
}
// Scan
result := sw.orchestrator.ScanParts(msgID, extraction.Parts, extraction.Partial)
// Emit degraded/timeout findings for operator visibility
if result.AllEnginesDown {
if shouldTempfailEmailDelivery(tempfail, result, nil) {
response = FAN_DENY // tempfail: defer delivery until engines recover
sw.emitDegradedWarning(fmt.Sprintf("All AV engines unavailable - message %s deferred (tempfail mode)", msgID))
return
}
sw.emitDegradedWarning(fmt.Sprintf("All AV engines unavailable - message %s delivered unscanned", msgID))
}
if len(result.TimedOutEngines) > 0 {
sw.emitFinding("email_av_timeout", alert.Warning,
fmt.Sprintf("Scan timeout for message %s on engine(s): %s", msgID, strings.Join(result.TimedOutEngines, ", ")))
if shouldTempfailEmailDelivery(tempfail, result, nil) {
sw.emitDegradedWarning(fmt.Sprintf("Incomplete AV scan - message %s deferred after engine timeout", msgID))
response = FAN_DENY
return
}
}
if len(result.ErroredEngines) > 0 {
sw.emitFinding("email_av_scan_error", alert.Warning,
fmt.Sprintf("Scan error for message %s on engine(s): %s", msgID, strings.Join(result.ErroredEngines, ", ")))
if shouldTempfailEmailDelivery(tempfail, result, nil) {
sw.emitDegradedWarning(fmt.Sprintf("Incomplete AV scan - message %s deferred after engine error", msgID))
response = FAN_DENY
return
}
}
if !result.Infected {
return // Clean - allow
}
// Infected - attempt quarantine
env := emailav.QuarantineEnvelope{
From: extraction.From,
To: extraction.To,
Subject: extraction.Subject,
Direction: extraction.Direction,
}
if sw.cfg.EmailAV.QuarantineInfected {
if err := sw.quarantineSpoolEvent(evt, msgID, spoolDir, result, env); err != nil {
fmt.Fprintf(os.Stderr, "[%s] spool watcher: quarantine failed for %s: %v\n", ts(), msgID, err)
sw.emitFinding("email_av_quarantine_error", alert.Warning,
fmt.Sprintf("Quarantine failed for infected message %s: %v", msgID, err))
if shouldTempfailEmailDelivery(tempfail, result, err) {
sw.emitDegradedWarning(fmt.Sprintf("Quarantine failed for infected message %s - delivery deferred (tempfail mode)", msgID))
response = FAN_DENY
}
} else {
// Quarantine succeeded - deny the open so Exim can't deliver
response = FAN_DENY
}
}
// Emit alert finding
sigNames := make([]string, len(result.Findings))
for i, f := range result.Findings {
sigNames[i] = fmt.Sprintf("%s(%s)", f.Signature, f.Engine)
}
msg := fmt.Sprintf("Malware detected in %s email from %s to %s: %s [subject: %s]",
extraction.Direction, extraction.From, strings.Join(extraction.To, ","),
strings.Join(sigNames, ", "), extraction.Subject)
sw.emitFinding("email_malware", alert.Critical, msg)
}
func partialExtractionReason(extraction *emime.ExtractionResult) string {
if extraction.PartialReason != "" {
return extraction.PartialReason
}
return "partial extraction"
}
func (sw *SpoolWatcher) writeResponse(fd int32, response uint32) {
resp := fanotifyResponse{Fd: fd, Response: response}
// #nosec G103 -- serializing the fanotify response struct for the
// kernel write; unsafe cast to a byte slice of the exact struct size.
respBytes := (*[responseSize]byte)(unsafe.Pointer(&resp))[:]
sw.initQueueHealth()
_, err := sw.kernelQueue.write(respBytes)
if err != nil {
// The kernel holds blocked processes until a response is written or
// the fanotify fd is closed. A failed write means the fd is broken -
// close it to release ALL pending permission events (fail-open),
// then signal the event loop to exit so the daemon can restart us.
fmt.Fprintf(os.Stderr, "[%s] spool watcher: FATAL - fanotify response write failed: %v - closing fd to release pending events\n", ts(), err)
sw.closeFd()
sw.Stop()
}
}
// closeFd closes the fanotify fd exactly once, even if called from multiple paths.
func (sw *SpoolWatcher) closeFd() {
if atomic.CompareAndSwapInt32(&sw.fdClosed, 0, 1) {
sw.initQueueHealth()
_ = sw.kernelQueue.close()
}
}
func (sw *SpoolWatcher) emitFinding(check string, severity alert.Severity, message string) bool {
if alert.TryEnqueue(sw.alertCh, alert.Finding{
Severity: severity,
Check: check,
Message: message,
Timestamp: time.Now(),
}) {
return true
} else {
// Alert channel full - drop
return false
}
}
// encryptedArchiveAlertInterval rate-limits the encrypted-archive finding.
// Password-protected archives are a steady part of ordinary business mail, so
// this is a recurring condition rather than an incident; a per-minute limit
// would put over a thousand findings a day in front of the operator.
const encryptedArchiveAlertInterval = time.Hour
// emitEncryptedArchiveWarning reports attachments delivered without being
// scanned because their archive entries are encrypted. This is deliberately
// not an email_av_degraded finding: nothing is degraded, and no retry will
// make the content readable. Keeping it separate lets an operator set policy
// on unscannable mail without losing the signal that scanning itself broke.
func (sw *SpoolWatcher) emitEncryptedArchiveWarning(msgID string, entries []emime.EncryptedArchiveEntry, omitted int) {
if len(entries) == 0 && omitted == 0 {
return
}
sw.encryptedMu.Lock()
defer sw.encryptedMu.Unlock()
if time.Since(sw.lastEncryptedAt) < encryptedArchiveAlertInterval {
return
}
named := make([]string, 0, len(entries))
for _, e := range entries {
named = append(named, fmt.Sprintf("%s in %s", e.Filename, e.ArchiveName))
}
if omitted > 0 {
named = append(named, fmt.Sprintf("%d additional encrypted member(s) (names omitted)", omitted))
}
if sw.emitFinding("email_av_encrypted_archive", alert.Warning,
fmt.Sprintf("Encrypted archive attachment could not be scanned for message %s: %s",
msgID, strings.Join(named, ", "))) {
// Only delivered warnings consume the allowance; queue pressure must
// not hide this condition for an hour after the queue recovers.
sw.lastEncryptedAt = time.Now()
}
}
// emitDegradedWarning emits an email_av_degraded finding, rate-limited to
// once per minute to avoid flooding the alert channel when clamd is down.
func (sw *SpoolWatcher) emitDegradedWarning(message string) {
sw.degradedMu.Lock()
if time.Since(sw.lastDegradedAt) < time.Minute {
sw.degradedMu.Unlock()
return
}
sw.lastDegradedAt = time.Now()
sw.degradedMu.Unlock()
sw.emitFinding("email_av_degraded", alert.Warning, message)
}
func (sw *SpoolWatcher) drainAndClose() {
sw.drainOnce.Do(func() {
close(sw.scanCh)
sw.wg.Wait()
sw.closeFd()
if atomic.CompareAndSwapInt32(&sw.pipeClosed, 0, 1) {
_ = unix.Close(sw.pipeFds[0])
_ = unix.Close(sw.pipeFds[1])
}
})
}
// PermissionMode returns true if using FAN_OPEN_PERM, false if FAN_CLOSE_WRITE fallback.
func (sw *SpoolWatcher) PermissionMode() bool {
return sw.permissionMode
}
// resolveEmailAVTempDir returns the private directory CreateTemp should use
// for extracted email parts.
func resolveEmailAVTempDir(cfg *config.Config) (string, error) {
if cfg == nil {
return "", fmt.Errorf("email AV temp dir: nil config")
}
if cfg.StatePath == "" {
return "", fmt.Errorf("email AV temp dir: state_path is empty")
}
dir := filepath.Join(cfg.StatePath, "emailav-tmp")
if err := os.MkdirAll(dir, 0o700); err != nil {
return "", fmt.Errorf("creating email AV temp dir %s: %w", dir, err)
}
if err := secureEmailAVTempDir(dir); err != nil {
return "", err
}
return dir, nil
}
func secureEmailAVTempDir(dir string) error {
info, err := os.Lstat(dir)
if err != nil {
return fmt.Errorf("checking email AV temp dir %s: %w", dir, err)
}
if info.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("email AV temp dir %s is a symlink", dir)
}
if !info.IsDir() {
return fmt.Errorf("email AV temp dir %s is not a directory", dir)
}
uid, gid, ok := fileOwner(info)
if ok {
euid := os.Geteuid()
if uid != euid {
if euid != 0 {
return fmt.Errorf("email AV temp dir %s is owned by uid %d, want uid %d", dir, uid, euid)
}
if chownErr := os.Chown(dir, euid, os.Getegid()); chownErr != nil {
return fmt.Errorf("owning email AV temp dir %s: %w", dir, chownErr)
}
} else if gid != os.Getegid() && euid == 0 {
if chownErr := os.Chown(dir, euid, os.Getegid()); chownErr != nil {
return fmt.Errorf("owning email AV temp dir %s: %w", dir, chownErr)
}
}
}
// #nosec G302 -- 0700 is the intended directory mode: daemon-only
// execute+read+write so unprivileged uids cannot enumerate or race
// staged email attachments. gosec's <=0600 rule does not distinguish
// directories from regular files.
if chmodErr := os.Chmod(dir, 0o700); chmodErr != nil {
return fmt.Errorf("chmod email AV temp dir %s: %w", dir, chmodErr)
}
info, err = os.Lstat(dir)
if err != nil {
return fmt.Errorf("checking email AV temp dir %s after chmod: %w", dir, err)
}
if info.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("email AV temp dir %s became a symlink", dir)
}
if !info.IsDir() {
return fmt.Errorf("email AV temp dir %s is not a directory", dir)
}
if info.Mode().Perm() != 0o700 {
return fmt.Errorf("email AV temp dir %s mode is %o, want 700", dir, info.Mode().Perm())
}
if uid, _, ok := fileOwner(info); ok && uid != os.Geteuid() {
return fmt.Errorf("email AV temp dir %s is owned by uid %d, want uid %d", dir, uid, os.Geteuid())
}
return nil
}
func fileOwner(info os.FileInfo) (uid, gid int, ok bool) {
stat, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return 0, 0, false
}
return int(stat.Uid), int(stat.Gid), true
}
// Stop signals the event loop to exit.
func (sw *SpoolWatcher) Stop() {
sw.stopOnce.Do(func() {
close(sw.stopCh)
if atomic.LoadInt32(&sw.pipeClosed) == 0 {
_, _ = unix.Write(sw.pipeFds[1], []byte{0})
}
sw.closeFd()
})
}
//go:build linux
package daemon
import (
"fmt"
"os"
"sync"
"sync/atomic"
"time"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/alert"
)
// spoolHoldBudget bounds how long one mail open may stay suspended waiting for
// a scan verdict. The kernel keeps the opening process frozen until we answer,
// so an unbounded wait turns slow scanning into frozen mail: a saturated
// scanner once held 167 Exim processes and stalled outbound mail host-wide.
// The scan continues after the deadline and can still quarantine.
var spoolHoldBudget = 5 * time.Second
const (
// spoolHoldExpiryWindow and spoolHoldExpiryThreshold define the burst that
// means the scanner cannot keep up with mail rather than one slow message.
spoolHoldExpiryWindow = time.Minute
spoolHoldExpiryThreshold = 10
// spoolBypassCooldown is how long opens are released unheld after a burst.
spoolBypassCooldown = 5 * time.Minute
)
// spoolWriteResponse is the seam to the kernel verdict write. Var so tests can
// observe verdicts without a live fanotify descriptor.
var spoolWriteResponse = (*SpoolWatcher).writeResponse
// holdGuard owns the verdict for one suspended open. Whoever gets there first
// wins -- the scan, or the budget timer -- and the kernel is answered once.
type holdGuard struct {
sw *SpoolWatcher
fd int32
needResp bool
once sync.Once
closeOnce sync.Once
timer *time.Timer
expired atomic.Bool
timeoutVerdict uint32 // read only after respond has joined once
}
// newHoldGuard arms the budget timer for a suspended open. Outside permission
// mode nothing is suspended, so there is no verdict to write and no timer.
func (sw *SpoolWatcher) newHoldGuard(fd int32, needResp bool) *holdGuard {
g := &holdGuard{sw: sw, fd: fd, needResp: needResp && sw.permissionMode}
if g.needResp {
g.timer = time.AfterFunc(spoolHoldBudget, g.expire)
}
return g
}
// timeoutResponse is the verdict used when the budget runs out. fail_mode
// tempfail means the operator chose never to deliver unscanned mail, so defer
// and let Exim retry; otherwise release the message and keep scanning.
func (sw *SpoolWatcher) timeoutResponse() uint32 {
if sw.cfg.EmailAV.FailMode == "tempfail" {
return FAN_DENY
}
return FAN_ALLOW
}
func (g *holdGuard) expire() {
if !g.needResp {
return
}
g.once.Do(func() {
g.timeoutVerdict = g.sw.timeoutResponse()
g.expired.Store(true)
spoolWriteResponse(g.sw, g.fd, g.timeoutVerdict)
// Keep all callback work inside once so finishing an event also joins
// its timer before the watcher or test hooks can be torn down.
g.sw.noteHoldExpiry(time.Now())
})
}
// respond hands the kernel the scan's verdict, unless the budget already
// answered for this event.
func (g *holdGuard) respond(response uint32) {
if g.timer != nil {
g.timer.Stop()
}
if !g.needResp {
return
}
g.once.Do(func() {
spoolWriteResponse(g.sw, g.fd, response)
})
}
func (g *holdGuard) finish(response uint32) {
g.closeOnce.Do(func() {
// once waits for a timer's in-progress write before the fd can be
// closed and recycled. Panic cleanup may finish the same event again.
g.respond(response)
_ = unix.Close(int(g.fd))
})
}
// timedOut reports whether the budget answered before the scan did.
func (g *holdGuard) timedOut() bool { return g.expired.Load() }
// holdWatchdog tracks budget expiries. A burst of them means scanning cannot
// keep up with mail, and holding every further message for the full budget
// only spreads the delay, so the watcher stops holding for a cooldown.
type holdWatchdog struct {
mu sync.Mutex
expiries []time.Time
bypassUntil time.Time
}
// recordExpiry notes one expiry and reports whether it started a bypass.
func (w *holdWatchdog) recordExpiry(now time.Time) bool {
w.mu.Lock()
defer w.mu.Unlock()
cutoff := now.Add(-spoolHoldExpiryWindow)
kept := w.expiries[:0]
for _, at := range w.expiries {
if at.After(cutoff) {
kept = append(kept, at)
}
}
kept = append(kept, now)
w.expiries = kept
if len(w.expiries) < spoolHoldExpiryThreshold || now.Before(w.bypassUntil) {
return false
}
w.bypassUntil = now.Add(spoolBypassCooldown)
w.expiries = w.expiries[:0]
return true
}
func (w *holdWatchdog) bypassing(now time.Time) bool {
w.mu.Lock()
defer w.mu.Unlock()
return now.Before(w.bypassUntil)
}
// noteHoldExpiry records an exhausted deadline or admission capacity and
// reports the first transition into bypass.
func (sw *SpoolWatcher) noteHoldExpiry(now time.Time) {
if !sw.holds.recordExpiry(now) {
return
}
action := "allowed without scanning"
if sw.timeoutResponse() == FAN_DENY {
action = "deferred without scanning (tempfail mode)"
}
fmt.Fprintf(os.Stderr, "[%s] spool watcher: scan capacity or hold budget repeatedly exhausted - mail %s for %s\n", ts(), action, spoolBypassCooldown)
sw.emitFinding("email_av_hold_bypass", alert.Critical,
fmt.Sprintf("Email AV scanning fell behind mail delivery: scan capacity or the %s hold budget was exhausted %d times within %s. New messages are %s for the next %s; these messages are not queued for a later scan. Scans already running continue. Investigate scanner load.",
spoolHoldBudget, spoolHoldExpiryThreshold, spoolHoldExpiryWindow, action, spoolBypassCooldown))
}
// dispatchBypass answers an open immediately while bypassing, without
// queueing it. Queueing is what left Exim waiting behind a saturated scanner.
// Reports whether it handled the event.
func (sw *SpoolWatcher) dispatchBypass(fd int32, needResp bool) bool {
if !sw.holds.bypassing(time.Now()) {
return false
}
if needResp && sw.permissionMode {
spoolWriteResponse(sw, fd, sw.timeoutResponse())
}
_ = unix.Close(int(fd))
sw.initQueueHealth()
sw.scannerHealth.Lose(time.Now(), 1)
return true
}
package daemon
import "github.com/pidginhost/csm/internal/emailav"
func shouldTempfailEmailDelivery(tempfail bool, result *emailav.ScanResult, quarantineErr error) bool {
if !tempfail {
return false
}
if quarantineErr != nil {
return true
}
if result == nil {
return false
}
return result.PartialExtraction ||
result.AllEnginesDown ||
len(result.TimedOutEngines) > 0 ||
len(result.ErroredEngines) > 0
}
//go:build linux
package daemon
import (
"fmt"
"io"
"os"
"path/filepath"
"golang.org/x/sys/unix"
"github.com/pidginhost/csm/internal/emailav"
)
// A verdict deadline can let Exim start delivery while we scan. Take its body
// lock before touching either spool file, even if the deadline has not fired
// yet: it may fire during the move. An OFD lock also excludes other workers
// and survives unrelated parser closes, unlike process-wide F_SETLK locks.
func (sw *SpoolWatcher) quarantineSpoolEvent(evt spoolEvent, msgID, spoolDir string, result *emailav.ScanResult, env emailav.QuarantineEnvelope) error {
fd, err := unix.Open(evt.path, unix.O_RDWR|unix.O_NOFOLLOW|unix.O_CLOEXEC|unix.O_NONBLOCK, 0)
if err != nil {
return fmt.Errorf("opening spool body for quarantine: %w", err)
}
defer func() { _ = unix.Close(fd) }()
lock := unix.Flock_t{Type: unix.F_WRLCK, Whence: int16(io.SeekStart)}
// #nosec G115 -- successful unix.Open returned a nonnegative POSIX fd.
if err := unix.FcntlFlock(uintptr(fd), unix.F_OFD_SETLK, &lock); err != nil {
return fmt.Errorf("locking spool body for quarantine: %w", err)
}
var original, current unix.Stat_t
if err := unix.Fstat(evt.fd, &original); err != nil {
return fmt.Errorf("checking event body for quarantine: %w", err)
}
if err := unix.Lstat(evt.path, ¤t); err != nil {
return fmt.Errorf("checking spool body for quarantine: %w", err)
}
var locked unix.Stat_t
if err := unix.Fstat(fd, &locked); err != nil {
return fmt.Errorf("checking locked body for quarantine: %w", err)
}
if current.Mode&unix.S_IFMT != unix.S_IFREG || original.Dev != current.Dev || original.Ino != current.Ino || locked.Dev != current.Dev || locked.Ino != current.Ino {
return fmt.Errorf("spool body changed before quarantine")
}
// A delivery journal contains recipients already delivered. The quarantine
// format stores only H/D; separating a surviving journal would lose that
// state on release and could send duplicates. Leave it for Exim recovery.
if _, err := os.Lstat(filepath.Join(spoolDir, msgID+"-J")); !os.IsNotExist(err) {
if err != nil {
return fmt.Errorf("checking delivery journal: %w", err)
}
return fmt.Errorf("spool delivery journal requires Exim recovery before quarantine")
}
if _, err := os.Lstat(filepath.Join(spoolDir, msgID+"-H")); err != nil {
return fmt.Errorf("checking spool header for quarantine: %w", err)
}
return sw.quarantine.QuarantineMessage(msgID, spoolDir, result, env)
}
package daemon
import (
"context"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/queuehealth"
"github.com/pidginhost/csm/internal/verdict"
)
// verdictAskFunc asks the operator's callback about one destination.
type verdictAskFunc func(ctx context.Context, req verdict.Request) (verdict.Response, error)
type verdictEnricherOpts struct {
Ask verdictAskFunc
Workers int
Queue int
TTL time.Duration
}
type verdictEntry struct {
tenantID string
verdict string
note string
at time.Time
}
type verdictKey struct {
ip string
reason string
severity string
}
type verdictJob struct {
key verdictKey
ticket queuehealth.Ticket
}
// verdictEnricher annotates findings with the operator's verdict without ever
// standing between the BPF ring buffer and the alert channel.
//
// The callback used to run inline in the consumer loop: one denied connection
// could hold the reader for the callback's whole timeout while a 256-slot ring
// overflowed, so an optional annotation cost real security events. Enrichment
// now happens on a bounded worker pool behind a short-lived cache. A finding is
// dispatched immediately either way; what a cache miss loses is the annotation
// on that first event, not the event.
type verdictEnricher struct {
ask verdictAskFunc
ttl time.Duration
jobs chan verdictJob
workers int
cacheCap int
wg sync.WaitGroup
mu sync.Mutex
cache map[verdictKey]verdictEntry
inFlight map[verdictKey]bool
dropped atomic.Int64
queueStats *queuehealth.Tracker
ctx context.Context
stopped bool
stopOnce sync.Once
}
func newVerdictEnricher(opts verdictEnricherOpts) *verdictEnricher {
if opts.Workers <= 0 {
opts.Workers = 2
}
if opts.Workers > 4 {
opts.Workers = 4
}
if opts.Queue <= 0 {
opts.Queue = 64
}
if opts.TTL <= 0 {
opts.TTL = time.Minute
}
return &verdictEnricher{
ask: opts.Ask,
ttl: opts.TTL,
jobs: make(chan verdictJob, opts.Queue),
workers: opts.Workers,
cacheCap: opts.Queue,
cache: make(map[verdictKey]verdictEntry),
inFlight: make(map[verdictKey]bool),
queueStats: queuehealth.New(opts.Queue, time.Minute),
}
}
func (e *verdictEnricher) start(ctx context.Context) {
e.mu.Lock()
e.ctx = ctx
e.mu.Unlock()
for i := 0; i < e.workers; i++ {
e.wg.Add(1)
go func() {
defer e.wg.Done()
e.work(ctx)
}()
}
}
func (e *verdictEnricher) wait() {
e.stopOnce.Do(func() {
e.mu.Lock()
e.stopped = true
close(e.jobs)
e.mu.Unlock()
e.wg.Wait()
e.mu.Lock()
defer e.mu.Unlock()
for job := range e.jobs {
delete(e.inFlight, job.key)
job.ticket.Reject(time.Now())
e.dropped.Add(1)
}
})
}
// QueueStatuses reports annotation work. The row is advisory: the finding is
// dispatched with or without a verdict, so a stalled lookup delays context
// rather than protection.
func (e *verdictEnricher) QueueStatuses(now time.Time) map[string]queuehealth.Status {
status := e.queueStats.Snapshot(now)
status.Advisory = true
return map[string]queuehealth.Status{"verdict": status}
}
func (e *verdictEnricher) droppedEnrichments() int64 { return e.dropped.Load() }
// annotate applies a cached verdict to f and reports whether it could. On a
// miss it queues the lookup and returns immediately: the caller dispatches the
// finding now and later events for the same destination carry the answer.
func (e *verdictEnricher) annotate(f *alert.Finding, ip, reason, severity string) bool {
if e == nil || f == nil || e.ask == nil {
return false
}
key := verdictKey{ip: ip, reason: reason, severity: severity}
e.mu.Lock()
entry, ok := e.cache[key]
fresh := ok && time.Since(entry.at) < e.ttl
if ok && !fresh {
delete(e.cache, key)
}
if fresh {
e.mu.Unlock()
applyVerdictEntry(f, entry)
return true
}
defer e.mu.Unlock()
if e.stopped || (e.ctx != nil && e.ctx.Err() != nil) {
e.queueStats.Lose(time.Now(), 1)
e.dropped.Add(1)
return false
}
if e.inFlight[key] {
return false
}
ticket := e.queueStats.Begin(time.Now())
e.inFlight[key] = true
select {
case e.jobs <- verdictJob{key: key, ticket: ticket}:
default:
// Saturated: the annotation is what we give up, never the finding.
delete(e.inFlight, key)
ticket.Reject(time.Now())
e.dropped.Add(1)
}
return false
}
func (e *verdictEnricher) work(ctx context.Context) {
for {
if ctx.Err() != nil {
return
}
select {
case <-ctx.Done():
return
case job, ok := <-e.jobs:
if !ok {
return
}
if err := e.processJob(ctx, job); err != nil {
csmlog.Warn("bpf enforcement verdict callback failed", "err", err, "dst", job.key.ip)
}
}
}
}
func (e *verdictEnricher) processJob(ctx context.Context, job verdictJob) error {
job.ticket.Start(time.Now())
completed := false
defer func() {
e.mu.Lock()
defer e.mu.Unlock()
delete(e.inFlight, job.key)
if completed {
job.ticket.Finish(time.Now())
} else {
job.ticket.Reject(time.Now())
e.dropped.Add(1)
}
}()
resp, err := e.ask(ctx, verdict.Request{
IP: job.key.ip,
Reason: job.key.reason,
Severity: job.key.severity,
Source: "bpf_enforcement",
})
if err != nil {
// A failure is not an answer: the next event must be able to retry.
return err
}
e.mu.Lock()
now := time.Now()
e.pruneCacheLocked(now)
e.cache[job.key] = verdictEntry{
tenantID: resp.TenantID,
verdict: resp.Verdict,
note: resp.Note,
at: now,
}
e.mu.Unlock()
completed = true
return nil
}
func (e *verdictEnricher) pruneCacheLocked(now time.Time) {
for key, entry := range e.cache {
if now.Sub(entry.at) >= e.ttl {
delete(e.cache, key)
}
}
if len(e.cache) < e.cacheCap {
return
}
var oldestKey verdictKey
var oldestAt time.Time
for key, entry := range e.cache {
if oldestAt.IsZero() || entry.at.Before(oldestAt) {
oldestKey = key
oldestAt = entry.at
}
}
delete(e.cache, oldestKey)
}
func applyVerdictEntry(f *alert.Finding, entry verdictEntry) {
if entry.tenantID != "" && f.TenantID == "" {
f.TenantID = entry.tenantID
}
if entry.verdict != "" {
appendFindingDetail(f, "Verdict callback: "+entry.verdict)
}
if entry.tenantID != "" {
appendFindingDetail(f, "Verdict tenant: "+entry.tenantID)
}
if entry.note != "" {
appendFindingDetail(f, "Verdict note: "+entry.note)
}
}
package daemon
import (
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/store"
"github.com/pidginhost/csm/internal/threatintel"
)
// verifiedBotEntries converts the operator-configured reputation.verified_bots
// list into the threatintel registry shape.
func verifiedBotEntries(cfg *config.Config) []threatintel.BotEntry {
if cfg == nil {
return nil
}
out := make([]threatintel.BotEntry, 0, len(cfg.Reputation.VerifiedBots))
for _, b := range cfg.Reputation.VerifiedBots {
out = append(out, threatintel.BotEntry{
Name: b.Name,
UASubstrings: b.UASubstrings,
RDNSSuffixes: b.RDNSSuffixes,
IPRanges: b.IPRanges,
})
}
return out
}
// reconcileVerifiedBots re-applies reputation.verified_bots after a SIGHUP so
// operators can add or change good bots without a restart. Re-stamping the
// PTR-verdict cache with the new list drops cached verdicts when the list
// changed, so a previously-spoofed IP is re-checked under the new suffixes
// instead of staying pinned for the cache TTL.
func (d *Daemon) reconcileVerifiedBots() {
cfg := d.activeOrStartupCfg()
entries := verifiedBotEntries(cfg)
threatintel.SetOperatorBots(entries)
if !cfg.BotVerifyEnabled() {
return
}
db := store.Global()
if db == nil {
return
}
ver := threatintel.OperatorBotsCacheVersion(threatintel.LogicVersion, entries)
if dropped, err := db.EnsureBotVerifyLogicVersion(ver); err == nil && dropped {
csmlog.Info("bot-verify cache dropped after verified_bots change")
}
if d.botVerifier == nil {
d.startBotVerifier(db, entries)
return
}
d.botVerifier.SetOperatorEntries(entries)
}
func (d *Daemon) startBotVerifier(db *store.DB, entries []threatintel.BotEntry) {
bv := threatintel.NewAsyncBotVerifier(db.PutBotVerify, db)
bv.SetOperatorEntries(entries)
d.registerQueueSource("bot_verification", bv)
d.botVerifier = bv
d.wg.Add(1)
obs.Go("bot-verify", func() {
defer d.wg.Done()
bv.Run(d.stopCh)
})
checks.SetBotVerifier(bv, db.GetBotVerify)
}
package daemon
import (
"bufio"
"bytes"
"context"
"errors"
"fmt"
"io"
"os"
"os/exec"
"strings"
"sync"
"syscall"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/eximlog"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/store"
)
const (
recentOutgoingMailHoldWindow = 2 * time.Hour
logWatcherMaxLineBytes = 256 * 1024
logWatcherOffsetMarkerBytes = 256
)
// LogLineHandler parses a log line and returns findings (if any).
type LogLineHandler func(line string, cfg *config.Config) []alert.Finding
// LogWatcher tails a log file using inotify and processes new lines.
type LogWatcher struct {
path string
cfg *config.Config
handler LogLineHandler
alertCh chan<- alert.Finding
file *os.File
offset int64
fileID logFileID
marker []byte
closeOnce sync.Once
}
type logFileID struct {
dev uint64
ino uint64
known bool
}
func fileID(info os.FileInfo) logFileID {
if info == nil {
return logFileID{}
}
if st, ok := info.Sys().(*syscall.Stat_t); ok {
return logFileID{
dev: uint64(st.Dev), // #nosec G115 -- device IDs are non-negative on supported Unix hosts
ino: uint64(st.Ino), // #nosec G115 -- inode numbers are non-negative
known: st.Dev != 0 || st.Ino != 0,
}
}
return logFileID{}
}
func (id logFileID) same(other logFileID) bool {
return id.known && other.known && id.dev == other.dev && id.ino == other.ino
}
func readOffsetMarker(f *os.File, offset int64) ([]byte, bool, error) {
if f == nil || offset <= 0 {
return nil, true, nil
}
n := int64(logWatcherOffsetMarkerBytes)
if offset < n {
n = offset
}
buf := make([]byte, n)
read, err := f.ReadAt(buf, offset-n)
if err != nil && !errors.Is(err, io.EOF) {
return nil, false, err
}
return buf[:read], read == len(buf), nil
}
// NewLogWatcher creates a watcher for a log file.
func NewLogWatcher(path string, cfg *config.Config, handler LogLineHandler, alertCh chan<- alert.Finding) (*LogWatcher, error) {
// #nosec G304 -- path is operator-configured log path from csm.yaml.
f, err := os.Open(path)
if err != nil {
return nil, err
}
// Seek to end - only process new lines
offset, err := f.Seek(0, io.SeekEnd)
if err != nil {
_ = f.Close()
return nil, err
}
info, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, err
}
marker, markerOK, err := readOffsetMarker(f, offset)
if err != nil {
_ = f.Close()
return nil, fmt.Errorf("read offset marker for %s: %w", path, err)
}
if !markerOK {
// The file rotated or shrank between Seek and ReadAt. Start without a
// marker rather than fail: a constructor error would disable this
// watcher until daemon restart, while readNewLines treats the saved
// position as untrusted and reads from the beginning on the next tick.
marker = nil
}
return &LogWatcher{
path: path,
cfg: cfg,
handler: handler,
alertCh: alertCh,
file: f,
offset: offset,
fileID: fileID(info),
marker: marker,
}, nil
}
// currentCfg returns the live daemon config so SIGHUP changes to thresholds,
// infra_ips, trusted_countries, and suppression settings reach the log-line
// handlers without a restart. Falls back to the startup snapshot before the
// first hot-reload publishes an active config.
func (w *LogWatcher) currentCfg() *config.Config {
if cfg := config.Active(); cfg != nil {
return cfg
}
return w.cfg
}
// Run starts watching the log file. Uses polling (every 2 seconds) instead of
// inotify to avoid complexity with log rotation. Simple, reliable, low overhead.
func (w *LogWatcher) Run(stopCh <-chan struct{}) {
// Run owns the file for its lifetime. Closing it here, rather than from a
// separate shutdown goroutine, keeps w.file single-threaded: a concurrent
// Stop() close used to race readNewLines/reopen, and the freed fd could be
// reused by another goroutine mid-Stat/Read.
defer w.closeFile()
ticker := time.NewTicker(2 * time.Second)
defer ticker.Stop()
// Also reopen the file every 5 minutes to handle log rotation
reopenTicker := time.NewTicker(5 * time.Minute)
defer reopenTicker.Stop()
for {
select {
case <-stopCh:
return
case <-reopenTicker.C:
w.reopen()
case <-ticker.C:
w.readNewLines()
}
}
}
// Stop closes the watcher's file. Safe to call when Run was never started
// (unit tests open a watcher just to drive readNewLines directly). When Run is
// active it closes the file itself on stopCh, so the daemon must not call Stop
// concurrently with a running Run.
func (w *LogWatcher) Stop() {
w.closeFile()
}
func (w *LogWatcher) closeFile() {
w.closeOnce.Do(func() {
if w.file != nil {
_ = w.file.Close()
}
})
}
func (w *LogWatcher) readNewLines() {
if w.file == nil {
w.reopen()
if w.file == nil {
return
}
}
info, err := w.file.Stat()
if err != nil {
w.reopen()
return
}
// File was truncated or rotated (smaller than our offset)
if info.Size() < w.offset {
w.reopen()
return
}
if !w.offsetMarkerMatches(w.file) {
w.offset = 0
w.marker = nil
}
// No new data
if info.Size() == w.offset {
return
}
// Seek to where we left off
_, err = w.file.Seek(w.offset, io.SeekStart)
if err != nil {
return
}
reader := bufio.NewReaderSize(w.file, 64*1024)
committedOffset := w.offset
for {
rawLine, truncated, readErr := readBoundedWatcherLine(reader, logWatcherMaxLineBytes)
if len(rawLine) > 0 && readErr != nil {
break
}
if readErr == nil {
if current, seekErr := w.file.Seek(0, io.SeekCurrent); seekErr == nil {
committedOffset = current - int64(reader.Buffered())
}
}
if truncated {
fmt.Fprintf(os.Stderr, "[%s] Warning: skipped oversized log line from %s at %d bytes\n", ts(), w.path, logWatcherMaxLineBytes)
}
if len(rawLine) > 0 && !truncated {
line := trimWatcherLineEnding(rawLine)
if line == "" {
if readErr != nil {
break
}
continue
}
findings := w.handler(line, w.currentCfg())
for _, f := range findings {
if f.Timestamp.IsZero() {
f.Timestamp = time.Now()
}
if !alert.TryEnqueue(w.alertCh, f) {
if f.Check == "exim_frozen_realtime" {
releaseEximFrozenDedup(line)
}
// Channel full - drop (backpressure)
fmt.Fprintf(os.Stderr, "[%s] Warning: alert channel full, dropping finding from %s\n", ts(), w.path)
}
}
}
if readErr != nil {
break
}
}
w.offset = committedOffset
w.refreshOffsetMarker()
}
func readBoundedWatcherLine(r *bufio.Reader, maxBytes int) (string, bool, error) {
var b strings.Builder
truncated := false
for {
chunk, err := r.ReadSlice('\n')
if len(chunk) > 0 {
switch {
case truncated:
case b.Len()+len(chunk) <= maxBytes:
b.Write(chunk)
default:
if room := maxBytes - b.Len(); room > 0 {
b.Write(chunk[:room])
}
truncated = true
}
}
if errors.Is(err, bufio.ErrBufferFull) {
continue
}
return b.String(), truncated, err
}
}
func trimWatcherLineEnding(line string) string {
line = strings.TrimSuffix(line, "\n")
return strings.TrimSuffix(line, "\r")
}
func (w *LogWatcher) reopen() {
if w.file != nil {
_ = w.file.Close()
// Drop the closed handle so a failed open below doesn't leave a dead
// fd behind for the next readNewLines to Stat in a loop.
w.file = nil
}
f, err := os.Open(w.path)
if err != nil {
return
}
info, err := f.Stat()
if err != nil {
_ = f.Close()
return
}
w.file = f
id := fileID(info)
switch {
case w.fileID.known && id.known && !w.fileID.same(id):
// Rotated by rename+create: new file, read from the start regardless
// of its size.
w.offset = 0
case info.Size() < w.offset:
// Truncated in place (copytruncate rotation).
w.offset = 0
case !w.offsetMarkerMatches(f):
// The saved offset now points into different content. This catches a
// truncate-and-regrow between polling ticks and cheap inode reuse.
w.offset = 0
}
if w.offset == 0 {
w.marker = nil
} else {
w.refreshOffsetMarker()
}
// Same file, size >= offset, matching marker: keep w.offset so lines
// written since the last read tick are not skipped. readNewLines seeks
// before every read.
w.fileID = id
}
func (w *LogWatcher) offsetMarkerMatches(f *os.File) bool {
if w.offset == 0 {
return true
}
if len(w.marker) == 0 {
return false
}
marker, ok, err := readOffsetMarker(f, w.offset)
return err == nil && ok && bytes.Equal(marker, w.marker)
}
func (w *LogWatcher) refreshOffsetMarker() {
marker, ok, err := readOffsetMarker(w.file, w.offset)
if err != nil || !ok {
w.marker = nil
return
}
w.marker = marker
}
// --- Log line handlers ---
func parseSessionLogLine(line string, cfg *config.Config) []alert.Finding {
var findings []alert.Finding
// cPanel login from non-infra IP - only alert on direct form login,
// not API-created sessions (from portal create_user_session)
if strings.Contains(line, "[cpaneld]") && strings.Contains(line, " NEW ") {
// Track IP→account for purge correlation (before any filtering)
if loginIP, loginAccount := parseCpanelSessionLogin(line); loginIP != "" && loginAccount != "" {
purgeTracker.recordLogin(loginIP, loginAccount)
}
switch {
case cfg.Suppressions.SuppressCpanelLogin:
// Skip all cPanel login alerts
case strings.Contains(line, "method=create_user_session") ||
strings.Contains(line, "method=create_session") ||
strings.Contains(line, "create_user_session"):
// Portal-created session - no alert
default:
ip, account := parseCpanelSessionLogin(line)
if ip != "" && account != "" && !isInfraIPDaemon(ip, cfg.InfraIPs) &&
!isTrustedCountry(ip, cfg.Suppressions.TrustedCountries) {
// WARNING severity - logins are useful for audit trail but
// not paging-level. Multi-IP correlation and brute-force
// stay at CRITICAL/HIGH via their own checks.
method := "unknown"
if strings.Contains(line, "method=handle_form_login") {
method = "direct form login"
} else if idx := strings.Index(line, "method="); idx >= 0 {
rest := line[idx+7:]
if comma := strings.IndexAny(rest, ",\n "); comma > 0 {
method = rest[:comma]
}
}
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "cpanel_login_realtime",
Message: fmt.Sprintf("cPanel direct login from non-infra IP: %s (account: %s, method: %s)", ip, account, method),
Details: truncateDaemon(line, 300),
SourceIP: ip,
TenantID: account,
})
}
}
}
// Password purge
if strings.Contains(line, "PURGE") && strings.Contains(line, "password_change") {
account := parsePurgeDaemon(line)
if account != "" {
purgeTracker.recordPurge(account)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "cpanel_password_purge_realtime",
Message: fmt.Sprintf("cPanel password purge for: %s", account),
TenantID: account,
})
}
}
return findings
}
// parseSecureLogLine reports accepted SSH logins from the authentication log.
// The scheduled ssh_logins check reads the same file and will meet this line
// again, so the finding comes from the shared builder: identical findings
// collapse in the state store, while a login this watcher never saw is still
// reported by the scan.
func parseSecureLogLine(line string, cfg *config.Config) []alert.Finding {
if f, ok := checks.SSHAcceptedLoginFinding(line, cfg); ok {
return []alert.Finding{f}
}
return nil
}
func parseEximLogLine(line string, cfg *config.Config) []alert.Finding {
var findings []alert.Finding
// 1. Frozen bounces - spam indicator. The capitalized form is the initial
// freeze event ("Frozen (delivery error message)"); the lowercase form is
// exim re-logging "Message is frozen" on every queue run. Dedup by queue
// ID so one stuck message alerts once, and unfreeze events never do.
if eximFrozenShouldAlert(line, time.Now()) {
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "exim_frozen_realtime",
Message: "Exim frozen message detected",
Details: truncateDaemon(line, 200),
})
}
// 2. Outgoing mail hold - account is held by cPanel.
// Format: "Sender office@example.com has an outgoing mail hold"
// or: "Domain example.org has an outgoing mail hold"
//
// Exim emits this rejection from the enforce_mail_permissions router on
// EVERY queued-message retry while the hold is active, so re-applying the
// hold here creates a feedback loop: an operator who clears a
// false-positive hold (e.g. caused by external transit defers like the
// 2026-05-11 Microsoft edge outage) sees CSM re-set the hold within
// seconds because old queued messages keep retrying. cPanel's
// TailWatch::Eximstats is the authoritative source for setting
// the hold. CSM records the hold so later retry-limit noise from
// the held domain is not promoted to a fresh spam outbreak.
permissionText := mailPermissionLogText(line)
if strings.Contains(permissionText, "outgoing mail hold") {
sender := extractMailHoldSender(permissionText)
domain := extractDomainFromEmail(sender)
if domain == "" {
domain = sender // may already be a bare domain
}
if domain != "" {
recordRecentOutgoingMailHold(domain)
RecordCompromisedDomain(domain)
}
// Alert only once per domain per hour
dedupKey := "email_hold:" + domain
if db := store.Global(); db != nil {
lastAlert := db.GetMetaString(dedupKey)
if lastAlert != "" && !isDedupExpired(lastAlert, 1*time.Hour) {
// Already alerted for this domain recently - skip finding
} else {
_ = db.SetMetaString(dedupKey, time.Now().Format(time.RFC3339))
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_compromised_account",
Message: fmt.Sprintf("Account %s is on cPanel outgoing mail hold", sender),
Details: truncateDaemon(line, 300),
Mailbox: mailboxOnly(sender),
Domain: domain,
TenantID: checks.MailOwner(domain),
})
}
}
}
// 3. Max defers/failures exceeded.
//
// cPanel's TailWatch::Eximstats has already throttled the domain by the
// time exim emits this line from enforce_mail_permissions, so the line is
// not independent evidence of an outbound spam blast: the same governor
// trips on inbound junk, full mailboxes, and forwarder bounces. Escalate
// to a compromise (CRITICAL + auto-hold) only when CSM's own
// authenticated-send rate window corroborates a real outbound blast for
// the domain. Otherwise report a deliverability event and leave the hold
// to cPanel, so an operator who clears a false-positive hold is not
// immediately re-held.
if strings.Contains(permissionText, "max defers and failures per hour") {
domain := extractEximDomain(permissionText)
if recentOutgoingMailHold(domain) {
return findings
}
if domainHasOutboundBlast(domain, cfg) {
held := maybeHoldOutgoingMail(cfg, domain)
if held {
recordRecentOutgoingMailHold(domain)
}
message := fmt.Sprintf("Spam outbreak: %s exceeded max defers/failures with high outbound volume", domain)
if held {
message += " - outgoing mail auto-suspended"
}
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_spam_outbreak",
Message: message,
Details: truncateDaemon(line, 300),
Domain: domain,
TenantID: checks.MailOwner(domain),
})
if domain != "" {
RecordCompromisedDomain(domain)
}
} else {
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_defer_fail_governor",
Message: fmt.Sprintf("%s hit the cPanel defer/fail governor; no outbound spam volume observed", domain),
Details: truncateDaemon(line, 300),
Domain: domain,
})
}
}
// 4. SMTP credentials leaked in subject - compromised account
// Pattern: T="...host:port,user@domain,PASSWORD..." in the subject field
if strings.Contains(line, " <= ") && strings.Contains(line, "T=\"") {
subject := extractEximSubject(line)
subjectLower := strings.ToLower(subject)
// Detect credential patterns: host:port,user,password or
// SMTP credentials in subject (common in credential stuffing attacks)
if (strings.Contains(subject, ":587,") || strings.Contains(subject, ":465,") ||
strings.Contains(subject, ":25,")) &&
strings.Contains(subject, "@") {
sender := extractEximSender(line)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_credential_leak",
Message: fmt.Sprintf("SMTP credentials leaked in email subject from %s", sender),
Details: fmt.Sprintf("The email subject contains what appears to be SMTP credentials (host:port,user,password). This account is likely compromised by a bulk mail service.\nSubject: %s", truncateDaemon(subject, 100)),
Mailbox: mailboxOnly(sender),
Domain: extractDomainFromEmail(sender),
TenantID: mailAccountOwner(extractAuthUser(line)),
})
}
// Also detect common spam subject patterns
if strings.Contains(subjectLower, "password") && strings.Contains(subjectLower, "smtp") {
sender := extractEximSender(line)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_credential_leak",
Message: fmt.Sprintf("Suspicious email subject with SMTP/password keywords from %s", sender),
Details: truncateDaemon(line, 300),
Mailbox: mailboxOnly(sender),
Domain: extractDomainFromEmail(sender),
TenantID: mailAccountOwner(extractAuthUser(line)),
})
}
}
// 5. Authentication from known bulk mail services
if strings.Contains(line, " <= ") && strings.Contains(line, "A=dovecot_") {
knownSpamServices := []string{
"truelist.io", "sendinblue.com", "mailspree.co",
"bulkmailer.", "massmailsoftware.", "sendblaster.",
}
lineLower := strings.ToLower(line)
for _, service := range knownSpamServices {
if strings.Contains(lineLower, service) {
sender := extractEximSender(line)
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_compromised_account",
Message: fmt.Sprintf("Compromised email account %s authenticated from bulk mail service %s", sender, service),
Details: truncateDaemon(line, 300),
Mailbox: mailboxOnly(sender),
Domain: extractDomainFromEmail(sender),
TenantID: mailAccountOwner(extractAuthUser(line)),
})
break
}
}
}
// 6. Dovecot auth failure - brute force indicator
// Format: "dovecot_login authenticator failed for (HELO) [IP]:port: 535 ... (set_id=user@domain)"
if strings.Contains(line, "authenticator failed") && strings.Contains(line, "dovecot") {
ip := eximlog.ClientIP(line)
account := extractSetID(line)
msg := "Email authentication failure"
if account != "" {
msg += " for " + account
}
if ip != "" {
msg += " from " + ip
}
// cPanel-local mailboxes log set_id as a bare local part with no
// "@domain"; treating it as a Mailbox would leave the structured
// field empty (mailboxOnly drops bare names) and force the
// correlator to fall back to SourceIP, splitting one targeted
// account across many attacker IPs. Route the bare form to
// TenantID so the incident groups by account.
mailbox, domain, tenant := splitMailAccount(account)
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_auth_failure_realtime",
Message: msg,
Details: truncateDaemon(line, 300),
SourceIP: ip,
Mailbox: mailbox,
Domain: domain,
TenantID: tenant,
})
}
// 7. DKIM signing failures
if dkimDomain := parseDKIMFailureDomain(line); dkimDomain != "" {
dedupKey := "dkim_fail:" + dkimDomain
if db := store.Global(); db != nil {
lastAlert := db.GetMetaString(dedupKey)
if lastAlert == "" || isDedupExpired(lastAlert, 24*time.Hour) {
_ = db.SetMetaString(dedupKey, time.Now().Format(time.RFC3339))
findings = append(findings, alert.Finding{
Severity: alert.Warning,
Check: "email_dkim_failure",
Message: fmt.Sprintf("DKIM signing failed for %s - check key file and DNS TXT record", dkimDomain),
Details: truncateDaemon(line, 300),
Timestamp: time.Now(),
Domain: dkimDomain,
})
}
}
}
// 8. SPF/DMARC outbound rejections
if spfDomain, spfReason := parseSPFDMARCRejection(line); spfDomain != "" {
dedupKey := "spf_reject:" + spfDomain
if db := store.Global(); db != nil {
lastAlert := db.GetMetaString(dedupKey)
if lastAlert == "" || isDedupExpired(lastAlert, 24*time.Hour) {
_ = db.SetMetaString(dedupKey, time.Now().Format(time.RFC3339))
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_spf_rejection",
Message: fmt.Sprintf("Outbound mail from %s rejected due to SPF/DMARC failure", spfDomain),
Details: fmt.Sprintf("Reason: %s\n%s", spfReason, truncateDaemon(line, 200)),
Timestamp: time.Now(),
Domain: spfDomain,
})
}
}
}
// 9. Outbound rate limiting for authenticated users
if strings.Contains(line, " <= ") && strings.Contains(line, "A=dovecot_") {
authUser := extractAuthUser(line)
if authUser != "" {
rateFindings := checkEmailRate(authUser, cfg)
findings = append(findings, rateFindings...)
}
}
// 10. Cloud-relay credential abuse (multiple authenticated sends from
// distinct cloud-provider IPs for the same mailbox).
// Use the AUTH identity (A=dovecot_*:<user>), not the envelope-from:
// the envelope-from can be forged by the attacker, while the AUTH
// identity is the credential actually being abused and the one we
// must lock out.
for _, f := range parseCloudRelayFinding(line, cfg) {
handleCloudRelayCredentialAbuse(cfg, extractAuthUser(line))
findings = append(findings, f)
}
if eng := PHPRelayEvaluator(); eng != nil {
findings = append(findings, eng.parsePHPRelayAccountVolume(line, time.Now())...)
}
return findings
}
// extractEximSender extracts the sender address from an exim log line.
// Format: "... <= sender@domain.com H=..."
func extractEximSender(line string) string {
idx := strings.Index(line, " <= ")
if idx < 0 {
return ""
}
rest := line[idx+4:]
fields := strings.Fields(rest)
if len(fields) > 0 {
return fields[0]
}
return ""
}
// extractEximDomain extracts a domain from an exim log line mentioning
// "Domain X has exceeded".
func extractEximDomain(line string) string {
idx := strings.Index(line, "Domain ")
if idx < 0 {
return ""
}
rest := line[idx+7:]
if sp := strings.IndexByte(rest, ' '); sp > 0 {
return rest[:sp]
}
return rest
}
// extractEximSubject extracts the subject from T="..." in an exim log line.
func extractEximSubject(line string) string {
idx := strings.Index(line, "T=\"")
if idx < 0 {
return ""
}
rest := line[idx+3:]
end := strings.Index(rest, "\"")
if end < 0 {
return rest
}
return rest[:end]
}
// --- Helpers (avoid import cycle with checks package) ---
func parseCpanelSessionLogin(line string) (ip, account string) {
idx := strings.Index(line, "[cpaneld]")
if idx < 0 {
return "", ""
}
rest := strings.TrimSpace(line[idx+len("[cpaneld]"):])
fields := strings.Fields(rest)
if len(fields) < 3 {
return "", ""
}
ip = fields[0]
for i, f := range fields {
if f == "NEW" && i+1 < len(fields) {
parts := strings.SplitN(fields[i+1], ":", 2)
if len(parts) >= 1 {
account = parts[0]
}
break
}
}
return ip, account
}
func parsePurgeDaemon(line string) string {
idx := strings.Index(line, "PURGE")
if idx < 0 {
return ""
}
rest := strings.TrimSpace(line[idx+len("PURGE"):])
fields := strings.Fields(rest)
if len(fields) < 1 {
return ""
}
parts := strings.SplitN(fields[0], ":", 2)
if len(parts) >= 1 {
return parts[0]
}
return ""
}
// isInfraIPDaemon is the realtime path's infra test. It delegates to the
// checks package so the two cannot drift: an earlier copy here lacked the
// Cloudflare branch, so realtime findings carried edge addresses that the
// scan path filters, and those were exported and blocked as attackers.
func isInfraIPDaemon(ip string, infraNets []string) bool {
return checks.IsInfraIP(ip, infraNets)
}
// mergeInfraIPs combines top-level infra IPs with firewall-specific ones,
// deduplicating entries. This allows the firewall to include additional CIDRs
// (e.g. server's own range) that need port access but shouldn't suppress alerts.
func mergeInfraIPs(topLevel, fwSpecific []string) []string {
return firewall.MergeInfraIPs(topLevel, fwSpecific)
}
// outgoingMailHoldUsersPath is the cPanel file listing users currently
// under OUTGOING_MAIL_HOLD. Read-only; cPanel/WHM owns mutation. var
// (not const) so tests can point it at a fixture.
var outgoingMailHoldUsersPath = "/etc/outgoing_mail_hold_users"
// whmapi1HoldExec invokes `whmapi1 hold_outgoing_email user=<user>`.
// Declared as var so tests can replace it without spawning whmapi1.
var whmapi1HoldExec = func(user string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
// #nosec G204 -- whmapi1 is the cPanel API binary; user is a cPanel
// account name resolved from /etc/userdomains (cpanel-managed).
return exec.CommandContext(ctx, "whmapi1", "hold_outgoing_email", "user="+user).CombinedOutput()
}
// userOnOutgoingMailHold reports whether the cPanel user already appears
// in /etc/outgoing_mail_hold_users. Used to short-circuit redundant
// whmapi1 calls when exim re-emits "exceeded max defers/failures" every
// retry hour while the hold is already active.
func userOnOutgoingMailHold(user string) bool {
if user == "" {
return false
}
f, err := os.Open(outgoingMailHoldUsersPath)
if err != nil {
return false
}
defer func() { _ = f.Close() }()
scanner := bufio.NewScanner(f)
for scanner.Scan() {
if strings.TrimSpace(scanner.Text()) == user {
return true
}
}
return false
}
// maybeHoldOutgoingMail applies an outgoing-mail hold only when auto-response
// is enabled and not in dry-run. Holding a customer's outbound mail is a
// customer-impacting action, so it honours the same master switch and dry-run
// safety default as IP blocking and quarantine; an operator evaluating CSM in
// monitor mode must never have mail held out from under them. It returns true
// only when a hold was actually applied (or already active), so callers can
// keep their hold-dedup bookkeeping accurate.
func maybeHoldOutgoingMail(cfg *config.Config, domainOrEmail string) bool {
if cfg == nil || !cfg.AutoResponse.Enabled || cfg.AutoResponseDryRunEnabled() {
fmt.Fprintf(os.Stderr, "[%s] auto-suspend: would hold outgoing mail for %s (auto_response disabled or dry-run)\n",
time.Now().Format("2006-01-02 15:04:05"), domainOrEmail)
return false
}
return autoSuspendOutgoingMail(domainOrEmail)
}
// autoSuspendOutgoingMail calls whmapi1 to hold outgoing mail for the cPanel
// account that owns the given domain or email address. It returns true when
// the hold is applied or already active. Declared as var so tests can swap in
// a recorder without spawning whmapi1.
var autoSuspendOutgoingMail = autoSuspendOutgoingMailReal
func autoSuspendOutgoingMailReal(domainOrEmail string) bool {
if domainOrEmail == "" {
return false
}
// Extract domain from email if needed
domain := domainOrEmail
if atIdx := strings.LastIndexByte(domain, '@'); atIdx >= 0 {
domain = domain[atIdx+1:]
}
// Look up cPanel username for this domain
user := lookupCPanelUser(domain)
if user == "" {
fmt.Fprintf(os.Stderr, "[%s] auto-suspend: could not find cPanel user for domain %s\n",
time.Now().Format("2006-01-02 15:04:05"), domain)
return false
}
// Skip if cPanel already lists this user as held. Re-issuing the
// hold has no operational effect, but on a sustained exim retry
// loop (queued bounces keep the defer/fail ratio above threshold
// hour after hour) the redundant whmapi1 calls produce a stream
// of "AUTO-SUSPEND" log lines that look like a fresh incident.
if userOnOutgoingMailHold(user) {
return true
}
out, err := whmapi1HoldExec(user)
if err != nil {
fmt.Fprintf(os.Stderr, "[%s] auto-suspend: whmapi1 hold_outgoing_email failed for %s: %v\n%s\n",
time.Now().Format("2006-01-02 15:04:05"), user, err, string(out))
return false
}
fmt.Fprintf(os.Stderr, "[%s] AUTO-SUSPEND: outgoing mail held for cPanel user %s (domain: %s)\n",
time.Now().Format("2006-01-02 15:04:05"), user, domain)
return true
}
// userdomainsPath is the cPanel domain→user map file. var (not const)
// so tests can point it at a fixture under t.TempDir(). Production must
// not mutate at runtime.
var userdomainsPath = "/etc/userdomains"
// lookupCPanelUser finds the cPanel username that owns a domain.
// Reads userdomainsPath which maps "domain: user" per line.
func lookupCPanelUser(domain string) string {
f, err := os.Open(userdomainsPath)
if err != nil {
return ""
}
defer func() { _ = f.Close() }()
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := scanner.Text()
parts := strings.SplitN(line, ":", 2)
if len(parts) != 2 {
continue
}
d := strings.TrimSpace(parts[0])
u := strings.TrimSpace(parts[1])
if strings.EqualFold(d, domain) {
return u
}
}
return ""
}
// extractMailHoldSender extracts the account/domain from outgoing mail hold messages.
//
// Two formats:
//
// "Sender user@domain has an outgoing mail hold" -> "user@domain"
// "Domain example.com has an outgoing mail hold" -> "example.com"
func extractMailHoldSender(line string) string {
// Try "Sender user@domain" first
if idx := strings.Index(line, "Sender "); idx >= 0 {
rest := line[idx+7:]
if sp := strings.IndexByte(rest, ' '); sp > 0 {
return rest[:sp]
}
return rest
}
// Try "Domain example.com" format
if idx := strings.Index(line, "Domain "); idx >= 0 {
rest := line[idx+7:]
if sp := strings.IndexByte(rest, ' '); sp > 0 {
return rest[:sp]
}
return rest
}
return ""
}
func recordRecentOutgoingMailHold(domain string) {
if domain == "" {
return
}
db := store.Global()
if db == nil {
return
}
_ = db.SetMetaString("email_hold_seen:"+domain, time.Now().Format(time.RFC3339))
}
func recentOutgoingMailHold(domain string) bool {
if domain == "" {
return false
}
db := store.Global()
if db == nil {
return false
}
stored := db.GetMetaString("email_hold_seen:" + domain)
return stored != "" && !isDedupExpired(stored, recentOutgoingMailHoldWindow)
}
// extractSetID extracts the account from "(set_id=user@domain)" or "(set_id=user)" in exim logs.
func extractSetID(line string) string {
const prefix = "set_id="
idx := strings.Index(line, prefix)
if idx < 0 {
return ""
}
rest := line[idx+len(prefix):]
end := strings.IndexAny(rest, ")\n ")
if end < 0 {
return rest
}
return rest[:end]
}
func truncateDaemon(s string, maxLen int) string {
if len(s) <= maxLen {
return s
}
return s[:maxLen] + "..."
}
// parseDKIMFailureDomain extracts domain from "DKIM: signing failed for {domain}"
func parseDKIMFailureDomain(line string) string {
const prefix = "DKIM: signing failed for "
idx := strings.Index(line, prefix)
if idx < 0 {
return ""
}
rest := line[idx+len(prefix):]
end := strings.IndexAny(rest, ": \t\n")
if end < 0 {
return rest
}
return rest[:end]
}
// parseSPFDMARCRejection extracts SENDER domain and rejection reason from
// exim ** permanent failure lines. Sender comes from <envelope_sender>.
func parseSPFDMARCRejection(line string) (senderDomain, reason string) {
starIdx := strings.Index(line, " ** ")
if starIdx < 0 {
return "", ""
}
// Extract envelope sender from <sender@domain> - search AFTER the **
// marker to avoid matching earlier <> fields (e.g. H=<hostname>).
rest := line[starIdx:]
ltIdx := strings.Index(rest, "<")
gtIdx := strings.Index(rest, ">")
if ltIdx < 0 || gtIdx < 0 || gtIdx <= ltIdx+1 {
return "", ""
}
sender := rest[ltIdx+1 : gtIdx]
atIdx := strings.LastIndexByte(sender, '@')
if atIdx < 0 || atIdx >= len(sender)-1 {
return "", ""
}
domain := sender[atIdx+1:]
// Extract rejection reason after last " : "
colonIdx := strings.LastIndex(line, " : ")
if colonIdx < 0 {
return "", ""
}
reason = strings.TrimSpace(line[colonIdx+3:])
if !isSPFDMARCRelated(reason) {
return "", ""
}
if len(reason) > 200 {
reason = reason[:200]
}
return domain, reason
}
// isSPFDMARCRelated checks if a rejection reason is SPF/DMARC related.
// Generic 5.7.1 alone is NOT sufficient - requires explicit auth keywords.
func isSPFDMARCRelated(reason string) bool {
if reason == "" {
return false
}
lower := strings.ToLower(reason)
for _, kw := range []string{"spf", "dmarc", "dkim"} {
if strings.Contains(lower, kw) {
return true
}
}
for _, code := range []string{"5.7.23", "5.7.25", "5.7.26"} {
if strings.Contains(lower, code) {
return true
}
}
if strings.Contains(lower, "5.7.1") {
if strings.Contains(lower, "authentication") || strings.Contains(lower, "ptr record") ||
strings.Contains(lower, "sender policy") || strings.Contains(lower, "alignment") {
return true
}
}
return false
}
// isDedupExpired checks if a stored RFC3339 timestamp is older than the given duration.
func isDedupExpired(stored string, window time.Duration) bool {
t, err := time.Parse(time.RFC3339, stored)
if err != nil {
return true
}
return time.Since(t) > window
}
// --- Outbound email rate limiting ---
// rateWindow tracks send timestamps for a single authenticated user.
type rateWindow struct {
mu sync.Mutex
times []time.Time
alerted string // last threshold level alerted ("warn" or "crit") - prevents repeated alerts per window
}
// add appends a timestamp to the window.
func (rw *rateWindow) add(t time.Time) {
rw.times = append(rw.times, t)
}
// countInWindow returns the number of timestamps within the window duration.
// Caller must hold rw.mu.
func (rw *rateWindow) countInWindow(now time.Time, window time.Duration) int {
cutoff := now.Add(-window)
count := 0
for _, t := range rw.times {
if t.After(cutoff) {
count++
}
}
return count
}
// prune removes timestamps older than the window duration and resets the
// alerted flag when the count drops below thresholds. Caller must hold rw.mu.
func (rw *rateWindow) prune(now time.Time, window time.Duration) {
cutoff := now.Add(-window)
kept := rw.times[:0]
for _, t := range rw.times {
if t.After(cutoff) {
kept = append(kept, t)
}
}
rw.times = kept
}
// emailRateWindows tracks per-user send rate windows.
var emailRateWindows sync.Map // map[string]*rateWindow
// extractAuthUser shares the Exim identity parser with scheduled mail checks.
func extractAuthUser(line string) string {
return eximlog.AuthenticatedUser(line)
}
// isHighVolumeSender checks if a user is in the high-volume senders allowlist.
func isHighVolumeSender(user string, allowlist []string) bool {
for _, allowed := range allowlist {
if strings.EqualFold(user, allowed) {
return true
}
}
return false
}
// extractDomainFromEmail returns the domain part of an email address.
func extractDomainFromEmail(email string) string {
idx := strings.LastIndexByte(email, '@')
if idx < 0 || idx >= len(email)-1 {
return ""
}
return email[idx+1:]
}
// mailboxOnly returns the input only when it looks like a full mailbox
// (contains '@'); otherwise returns "". Used by realtime emit sites that
// receive either "user@domain" or a bare domain — the bare domain belongs
// in the Domain field, not Mailbox, so the correlator does not collapse
// distinct mailboxes onto a domain key.
func mailboxOnly(s string) string {
if strings.IndexByte(s, '@') < 0 {
return ""
}
return s
}
// splitMailAccount classifies an authenticated mail account string into
// the three correlation fields. A full mailbox ("user@domain") routes to
// Mailbox + Domain. A bare local part (cPanel-style, no '@') routes to
// TenantID so the incident correlator groups by account, not by attacker
// SourceIP. An empty input returns three empty strings.
func splitMailAccount(account string) (mailbox, domain, tenant string) {
if account == "" {
return "", "", ""
}
if strings.IndexByte(account, '@') < 0 {
return "", "", account
}
return account, extractDomainFromEmail(account), ""
}
// hasRecentCompromisedFinding checks if there's a recent email_compromised_account
// or email_spam_outbreak finding for the given domain (suppresses rate alerts).
func hasRecentCompromisedFinding(domain string) bool {
emailRateSuppressed.mu.Lock()
defer emailRateSuppressed.mu.Unlock()
if ts, ok := emailRateSuppressed.domains[domain]; ok {
if time.Since(ts) < time.Hour {
return true
}
delete(emailRateSuppressed.domains, domain)
}
return false
}
// emailRateSuppressed tracks domains with recent compromised/spam findings.
var emailRateSuppressed = struct {
mu sync.Mutex
domains map[string]time.Time
}{domains: make(map[string]time.Time)}
// RecordCompromisedDomain marks a domain as having a recent compromised finding.
// Called from parseEximLogLine when email_compromised_account or email_spam_outbreak fires.
func RecordCompromisedDomain(domain string) {
emailRateSuppressed.mu.Lock()
defer emailRateSuppressed.mu.Unlock()
emailRateSuppressed.domains[domain] = time.Now()
}
// domainHasOutboundBlast reports whether authenticated senders under the given
// domain have produced enough outbound volume within the rate window to
// corroborate an actual spam outbreak. A cPanel defer/fail governor trip alone
// is not such evidence. Returns false when rate thresholds are unconfigured, so
// an operator who has not tuned the rate window never auto-holds on a bare
// governor line.
func domainHasOutboundBlast(domain string, cfg *config.Config) bool {
if domain == "" || cfg == nil {
return false
}
threshold := cfg.EmailProtection.RateWarnThreshold
windowDur := time.Duration(cfg.EmailProtection.RateWindowMin) * time.Minute
if threshold <= 0 || windowDur <= 0 {
return false
}
now := time.Now()
total := 0
emailRateWindows.Range(func(key, val any) bool {
user, ok := key.(string)
if !ok || !strings.EqualFold(extractDomainFromEmail(user), domain) {
return true
}
rw, ok := val.(*rateWindow)
if !ok {
return true
}
rw.mu.Lock()
total += rw.countInWindow(now, windowDur)
rw.mu.Unlock()
return total < threshold // stop iterating once corroborated
})
return total >= threshold
}
// checkEmailRate processes an outbound email for rate limiting.
// Returns findings if thresholds are exceeded.
func checkEmailRate(user string, cfg *config.Config) (findings []alert.Finding) {
// Registered before the unlock defer so owner I/O runs after it.
defer func() { stampMailAccountOwner(findings, user) }()
// Guard: skip if thresholds are zero (misconfigured or disabled)
if cfg.EmailProtection.RateWarnThreshold <= 0 || cfg.EmailProtection.RateCritThreshold <= 0 {
return nil
}
if isHighVolumeSender(user, cfg.EmailProtection.HighVolumeSenders) {
return nil
}
// Load or create rate window for this user
val, _ := emailRateWindows.LoadOrStore(user, &rateWindow{})
rw := val.(*rateWindow)
now := time.Now()
windowDur := time.Duration(cfg.EmailProtection.RateWindowMin) * time.Minute
rw.mu.Lock()
defer rw.mu.Unlock()
// Check domain suppression BEFORE adding to window - prevents
// phantom rate inflation for suppressed domains.
domain := extractDomainFromEmail(user)
if domain != "" && hasRecentCompromisedFinding(domain) {
return nil
}
rw.add(now)
count := rw.countInWindow(now, windowDur)
// Reset alerted state when count drops below warn threshold -
// allows re-alerting on the next burst after the window slides.
if count < cfg.EmailProtection.RateWarnThreshold {
rw.alerted = ""
}
mailbox, domain, _ := splitMailAccount(user)
if count >= cfg.EmailProtection.RateCritThreshold {
if rw.alerted != "crit" {
rw.alerted = "crit"
findings = append(findings, alert.Finding{
Severity: alert.Critical,
Check: "email_rate_critical",
Message: fmt.Sprintf("Email rate CRITICAL: %s sent %d messages in %d minutes (threshold: %d)", user, count, cfg.EmailProtection.RateWindowMin, cfg.EmailProtection.RateCritThreshold),
Details: fmt.Sprintf("User: %s\nMessages in window: %d\nWindow: %d minutes\nThreshold: %d", user, count, cfg.EmailProtection.RateWindowMin, cfg.EmailProtection.RateCritThreshold),
Mailbox: mailbox,
Domain: domain,
})
}
} else if count >= cfg.EmailProtection.RateWarnThreshold {
if rw.alerted != "warn" && rw.alerted != "crit" {
rw.alerted = "warn"
findings = append(findings, alert.Finding{
Severity: alert.High,
Check: "email_rate_warning",
Message: fmt.Sprintf("Email rate WARNING: %s sent %d messages in %d minutes (threshold: %d)", user, count, cfg.EmailProtection.RateWindowMin, cfg.EmailProtection.RateWarnThreshold),
Details: fmt.Sprintf("User: %s\nMessages in window: %d\nWindow: %d minutes\nThreshold: %d", user, count, cfg.EmailProtection.RateWindowMin, cfg.EmailProtection.RateWarnThreshold),
Mailbox: mailbox,
Domain: domain,
})
}
}
return findings
}
// StartEmailRateEviction starts a background goroutine that prunes expired
// rate windows every 10 minutes. Same pattern as StartModSecEviction.
func StartEmailRateEviction(stopCh <-chan struct{}) {
obs.Go("email-rate-eviction", func() {
ticker := time.NewTicker(10 * time.Minute)
defer ticker.Stop()
for {
select {
case <-stopCh:
return
case now := <-ticker.C:
evictEmailRateWindows(now)
}
}
})
}
// evictEmailRateWindows prunes all per-user rate windows and deletes empty entries.
func evictEmailRateWindows(now time.Time) {
// Use a generous 60-minute eviction window to avoid premature deletion.
// The actual rate window is checked during rate evaluation.
evictWindow := 60 * time.Minute
emailRateWindows.Range(func(key, val any) bool {
rw := val.(*rateWindow)
rw.mu.Lock()
rw.prune(now, evictWindow)
empty := len(rw.times) == 0
if empty {
rw.alerted = ""
}
rw.mu.Unlock()
if empty {
emailRateWindows.Delete(key)
}
return true
})
// Also prune the suppressed domains map
emailRateSuppressed.mu.Lock()
for domain, ts := range emailRateSuppressed.domains {
if time.Since(ts) > time.Hour {
delete(emailRateSuppressed.domains, domain)
}
}
emailRateSuppressed.mu.Unlock()
}
package daemon
import "regexp"
// looksLikePHPWebshell returns true when PHP file content exhibits the
// canonical realtime-detectable webshell shapes:
// 1. A request superglobal ($_GET / $_POST / $_REQUEST / $_COOKIE /
// php://input) flowing into a code-execution primitive
// (eval / assert / system / passthru / exec / shell_exec / proc_open
// / popen / create_function), with optional decoder layers
// (base64_decode / gzinflate / str_rot13).
// 2. eval/assert wrapping a decoder of an arbitrary base64 / gzinflate
// blob (the obfuscated-payload primitive).
//
// Returns false on legitimate code that uses dangerous functions in
// non-attack contexts (Pear Text_Diff/Engine/shell.php's shell_exec call
// to Unix `diff`, TinyMCE charmap.php's static glyph data array).
func looksLikePHPWebshell(data []byte) bool {
if len(data) == 0 {
return false
}
// No inner byte cap. The fanotify callers own the read window and
// record when that window is hit. A duplicate cap here would silently
// shorten any caller that legitimately passed a larger buffer, hiding
// truncation from that caller without adding protection beyond the
// upstream reader limit.
for _, re := range webshellContentRegexes {
if re.Match(data) {
return true
}
}
return requestVariableFlowsToDangerousFunction(data)
}
// webshellContentRegexes are the realtime-detection-grade content patterns
// for looksLikePHPWebshell. Compiled once at package init.
var webshellContentRegexes = []*regexp.Regexp{
// Request superglobal directly piped into a code-execution primitive
// in the same expression (with optional decoder layers).
regexp.MustCompile(`(?i)\b(?:eval|assert|system|passthru|exec|shell_exec|proc_open|popen|create_function)\s*\(\s*(?:gzinflate\s*\(\s*|str_rot13\s*\(\s*|base64_decode\s*\(\s*|@\s*)*\s*\$_(?:GET|POST|REQUEST|COOKIE|FILES|SERVER)\b`),
// php://input piped into eval/assert/system in the same expression.
regexp.MustCompile(`(?i)\b(?:eval|assert|system|passthru|exec|shell_exec|proc_open|popen|create_function)\s*\(\s*(?:gzinflate\s*\(\s*|str_rot13\s*\(\s*|base64_decode\s*\(\s*|@\s*)*\s*file_get_contents\s*\(\s*['"]php://input`),
// eval/assert wrapping a base64/gzinflate/str_rot13 decoder of a
// long literal blob (obfuscated-payload primitive).
regexp.MustCompile(`(?i)\b(?:eval|assert)\s*\(\s*(?:gzinflate\s*\(\s*|str_rot13\s*\(\s*)?base64_decode\s*\(\s*['"][A-Za-z0-9+/=]{40,}`),
}
var (
requestAssignmentRegex = regexp.MustCompile(`(?is)\$([A-Za-z_][A-Za-z0-9_]*)\s*=\s*(?:@?\s*)?(?:\$_(?:GET|POST|REQUEST|COOKIE|FILES|SERVER)\b(?:\s*\[[^\]]{0,200}\])?|file_get_contents\s*\(\s*['"]php://input['"]\s*\))[^;]{0,200};`)
dangerousVariableCallRegex = regexp.MustCompile(`(?is)\b(?:eval|assert|system|passthru|exec|shell_exec|proc_open|popen|create_function)\s*\(\s*(?:@?\s*)?\$([A-Za-z_][A-Za-z0-9_]*)\b`)
)
func requestVariableFlowsToDangerousFunction(data []byte) bool {
matches := requestAssignmentRegex.FindAllSubmatchIndex(data, -1)
for _, match := range matches {
if len(match) < 4 {
continue
}
varName := string(data[match[2]:match[3]])
windowEnd := match[1] + 800
if windowEnd > len(data) {
windowEnd = len(data)
}
for _, call := range dangerousVariableCallRegex.FindAllSubmatch(data[match[1]:windowEnd], -1) {
if len(call) == 2 && string(call[1]) == varName {
return true
}
}
}
return false
}
package daemon
import (
"context"
"fmt"
"os"
"sync"
"sync/atomic"
"syscall"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/yara"
"github.com/pidginhost/csm/internal/yaraworker"
)
// yaraMetricsOnce guards registration of the yara-worker restart
// counter hook so a baseline re-run or a second daemon instance in the
// same test binary does not panic with "duplicate registration".
var yaraMetricsOnce sync.Once
// yaraWorkerOn reports whether the daemon should run YARA-X in a
// supervised child process. The field is a *bool tri-state: nil means
// "use system default" (true, per ROADMAP item 2 follow-up), *true is
// explicit opt-in, *false is explicit opt-out.
//
// A nil cfg falls back to false so a pathological caller does not
// accidentally spin up a worker; production never passes nil.
func yaraWorkerOn(cfg *config.Config) bool {
if cfg == nil {
return false
}
if cfg.Signatures.YaraWorkerEnabled == nil {
return true
}
return *cfg.Signatures.YaraWorkerEnabled
}
// yaraWorkerWatcher names the YARA-X worker in the watcher status map.
const yaraWorkerWatcher = "yara_worker"
// initYaraBackend wires up either the out-of-process YARA-X supervisor
// (default, per ROADMAP item 2 follow-up) or the in-process scanner
// (when config.Signatures.YaraWorkerEnabled is explicitly *false).
// Both paths register themselves as the yara package's active backend
// so existing callers work unchanged.
//
// In worker mode, yara.Init is deliberately NOT called: the rule
// compile happens inside the child process, not the daemon. This is
// the point of the feature (ROADMAP item 2) — a cgo crash while
// compiling or scanning stays contained to the child. Matches carry
// string-valued rule metadata (see yara.Match.Meta / yaraipc.Match.Meta)
// so the emailav YARA-X adapter works identically under both backends
// — severity no longer needs the in-process *yara_x.Rules object.
func (d *Daemon) initYaraBackend() error {
if !yaraWorkerOn(d.cfg) {
if yaraScanner := yara.Init(d.cfg.Signatures.RulesDir, d.cfg.Signatures.DisabledRules...); yaraScanner != nil {
fmt.Fprintf(os.Stderr, "[%s] YARA-X scanner active: %d rule file(s)\n", ts(), yaraScanner.RuleCount())
}
return nil
}
sup, err := yaraworker.NewSupervisor(d.yaraSupervisorConfig())
if err != nil {
return fmt.Errorf("creating yara-worker supervisor: %w", err)
}
// Own the supervisor before the first Start attempt. If startup fails and
// the retry goroutine is still trying to bring it online when the daemon
// stops, stopYaraBackend can cancel/kill that in-flight Start instead of
// leaving an orphaned worker attempt outside shutdown ownership.
d.yaraSup = sup
if err := sup.Start(context.Background()); err != nil {
// Every YARA scan is off until the worker starts, so report it where
// doctor and the status API look for components that failed to attach.
d.MarkWatcher(yaraWorkerWatcher, false)
// A boot-time start failure must not disable YARA for the daemon's
// whole lifetime. Retry in the background with backoff and raise a
// finding once the failure looks persistent, mirroring the
// post-start crash path (onYaraWorkerRestart). The daemon keeps
// running; scanning comes online once the worker starts.
obs.Go("yara-init-retry", func() { d.retryYaraStart(sup) })
return fmt.Errorf("starting yara-worker (retrying in background): %w", err)
}
d.activateYaraBackend(sup)
return nil
}
// activateYaraBackend installs a started supervisor as the active YARA
// backend, wires its restart metric, and surfaces a still-broken rule compile
// as a finding so a worker that is up but scanning nothing is visible instead
// of masquerading as a healthy zero-rule host.
func (d *Daemon) activateYaraBackend(sup *yaraworker.Supervisor) {
yara.SetActive(sup)
// Start launches supervision before returning. Serialize boot readiness
// with crash reporting so activation cannot erase an early crash; only
// the stable callback may restore health after an exit.
d.yaraCrashMu.Lock()
if sup.RestartCount() == 0 {
d.MarkWatcher(yaraWorkerWatcher, true)
}
d.yaraCrashMu.Unlock()
// Expose the supervisor's cumulative restart count to Prometheus.
// Registered once per process; subsequent calls re-point nothing
// (the closure captures `sup`, and a second daemon.Run in the
// same process would need to arrange for the metric to follow).
yaraMetricsOnce.Do(func() {
metrics.RegisterCounterFunc(
"csm_yara_worker_restarts_total",
"Number of times the YARA-X worker subprocess has been restarted by its supervisor.",
func() float64 { return float64(sup.RestartCount()) },
)
})
fmt.Fprintf(os.Stderr, "[%s] %s\n", ts(), yaraWorkerStatusLine(sup.RuleCount(), sup.ChildPID()))
d.reportYaraCompileStatus(sup.CompileError())
d.reportRealtimeRuleCoverage(yamlRuleCount(), sup.RuleCount(), sup.CompileError() == "")
}
// yaraSupervisorConfig describes the worker the daemon runs and how its
// lifecycle is reported.
func (d *Daemon) yaraSupervisorConfig() yaraworker.SupervisorConfig {
return yaraworker.SupervisorConfig{
BinaryPath: d.binaryPath,
SocketPath: yaraworker.DefaultSocketPath(),
RulesDir: d.cfg.Signatures.RulesDir,
ConfigFile: d.cfg.ConfigFile,
ConfigDir: d.cfg.ConfigDir,
DisabledRules: d.cfg.Signatures.DisabledRules,
StartTimeout: 10 * time.Second,
MinRestartInterval: time.Second,
MaxRestartInterval: 60 * time.Second,
StableDuration: 30 * time.Second,
ClientTimeout: 30 * time.Second,
OnRestart: d.onYaraWorkerRestart,
OnStable: d.onYaraWorkerStable,
Logf: func(format string, args ...any) {
fmt.Fprintf(os.Stderr, "[%s] yara-worker: "+format+"\n", append([]any{ts()}, args...)...)
},
}
}
// yaraWorkerStatusLine describes the worker's startup state. A worker that
// compiled zero rules is running but matches nothing, so it must not be
// called active: that wording let a host report a healthy scanner while every
// scan silently checked against an empty rule set.
func yaraWorkerStatusLine(ruleCount int, childPID int) string {
if ruleCount == 0 {
return fmt.Sprintf("YARA-X worker started with 0 rules compiled - scanning nothing until rules load (pid=%d)", childPID)
}
return fmt.Sprintf("YARA-X worker active: %d rule(s) compiled in child process (pid=%d)", ruleCount, childPID)
}
// retryYaraStart re-attempts a failed worker start with capped exponential
// backoff until it succeeds or the daemon stops, raising one finding once the
// failure is persistent so the outage is visible.
func (d *Daemon) retryYaraStart(sup *yaraworker.Supervisor) {
ok := retryStartWithStopContext(
func(ctx context.Context) error { return sup.Start(ctx) },
d.stopCh,
time.Second, 60*time.Second, 3,
func(attempt int, err error) {
d.emitYaraFinding(alert.Critical, "yara_backend_unavailable",
fmt.Sprintf("YARA-X worker has failed to start %d times (%v); real-time malware scanning is disabled while the supervisor keeps retrying.", attempt, err))
},
)
if !ok {
return
}
fmt.Fprintf(os.Stderr, "[%s] yara-worker started after retrying a failed boot\n", ts())
d.activateYaraBackend(sup)
}
func retryStartWithStopContext(start func(context.Context) error, stop <-chan struct{}, minBackoff, maxBackoff time.Duration, alertAfter int, onPersistent func(attempt int, err error)) bool {
ctx, cancel := context.WithCancel(context.Background())
retryDone := make(chan struct{})
obs.Go("yara-init-retry-stop", func() {
select {
case <-stop:
cancel()
case <-retryDone:
}
})
ok := retryStart(
func() error { return start(ctx) },
stop,
minBackoff, maxBackoff, alertAfter,
onPersistent,
)
close(retryDone)
if !ok {
cancel()
}
return ok
}
// retryStart calls start with capped exponential backoff until it returns nil,
// stop is signalled, or forever. onPersistent, if set, fires exactly once
// after alertAfter consecutive failures so a persistent outage raises a single
// finding. Returns true if start eventually succeeded.
func retryStart(start func() error, stop <-chan struct{}, minBackoff, maxBackoff time.Duration, alertAfter int, onPersistent func(attempt int, err error)) bool {
backoff := minBackoff
for attempt := 1; ; attempt++ {
select {
case <-time.After(backoff):
case <-stop:
return false
}
if stopRequested(stop) {
return false
}
if err := start(); err == nil {
return !stopRequested(stop)
} else if attempt == alertAfter && onPersistent != nil {
if stopRequested(stop) {
return false
}
onPersistent(attempt, err)
}
backoff *= 2
if backoff > maxBackoff {
backoff = maxBackoff
}
}
}
func stopRequested(stop <-chan struct{}) bool {
if stop == nil {
return false
}
select {
case <-stop:
return true
default:
return false
}
}
// reportYaraCompileStatus raises a finding when the worker is alive but its
// rules failed to compile. A clean compile ("") is silent.
func (d *Daemon) reportYaraCompileStatus(compileErr string) {
if compileErr == "" {
return
}
d.emitYaraFinding(alert.Critical, "yara_worker_compile_failed",
fmt.Sprintf("YARA-X worker is up but its rules failed to compile: %s. Real-time malware scanning is disabled until the rules are fixed and reloaded.", compileErr))
}
// emitYaraFinding pushes a finding without blocking; a saturated alert channel
// is accounted for by the daemon's general drop counter.
func (d *Daemon) emitYaraFinding(sev alert.Severity, check, msg string) bool {
finding := alert.Finding{Severity: sev, Check: check, Message: msg, Timestamp: time.Now()}
if alert.TryEnqueue(d.alertCh, finding) {
return true
} else {
atomic.AddInt64(&d.droppedAlerts, 1)
return false
}
}
// stopYaraBackend is called during daemon shutdown. Safe to call when
// worker mode is off.
func (d *Daemon) stopYaraBackend() {
if d.yaraSup == nil {
return
}
if err := d.yaraSup.Stop(); err != nil {
fmt.Fprintf(os.Stderr, "[%s] yara-worker stop: %v\n", ts(), err)
}
yara.SetActive(nil)
}
// onYaraWorkerStable records a worker that stayed up after starting or
// restarting as attached again.
func (d *Daemon) onYaraWorkerStable() {
select {
case <-d.stopCh:
return
default:
}
d.MarkWatcher(yaraWorkerWatcher, true)
}
// onYaraWorkerRestart is called once per unplanned worker exit. Emits
// a Critical finding the first time, then one every minute after to
// avoid spamming alerts while a broken rule package is in place.
func (d *Daemon) onYaraWorkerRestart(exitCode int, sig syscall.Signal, ranFor time.Duration) {
// The supervisor suppresses this callback once its own context is
// cancelled, but that happens in stopYaraBackend, late in shutdown.
// stopCh closes at the top of shutdown and covers that window. The unit's
// KillMode=mixed lets the daemon begin shutdown before stopping workers;
// a worker can still exit independently during the drain.
//
// A nil stopCh (zero-value Daemon, several tests) blocks forever on
// receive, so the default arm is what keeps those reporting.
select {
case <-d.stopCh:
return
default:
}
now := time.Now()
d.yaraCrashMu.Lock()
// Pair with activation's readiness publication. Scanning stays failed
// until a restarted worker proves it stays up.
d.MarkWatcher(yaraWorkerWatcher, false)
last := d.yaraLastCrashAlert
d.yaraLastCrashAlert = now
d.yaraCrashMu.Unlock()
fmt.Fprintf(os.Stderr, "[%s] yara-worker exited code=%d signal=%v ran=%s\n",
ts(), exitCode, sig, ranFor.Round(time.Millisecond))
if !last.IsZero() && now.Sub(last) < time.Minute {
return
}
finding := alert.Finding{
Severity: alert.Critical,
Check: "yara_worker_crashed",
Timestamp: now,
Message: fmt.Sprintf("YARA-X worker crashed (exit=%d signal=%v after %s); YARA scanning is offline; supervisor will attempt a restart. Worker health remains failed until a restarted worker stays up for 30s.", exitCode, sig, ranFor.Round(time.Millisecond)),
}
alert.TryEnqueue(d.alertCh, finding)
}
package emailav
import (
"encoding/binary"
"fmt"
"io"
"net"
"os"
"strings"
"time"
)
// ClamdScanner scans files via the clamd Unix socket using INSTREAM.
type ClamdScanner struct {
socketPath string
}
// NewClamdScanner creates a ClamdScanner that connects to clamd at the given socket path.
func NewClamdScanner(socketPath string) *ClamdScanner {
return &ClamdScanner{socketPath: socketPath}
}
func (s *ClamdScanner) Name() string { return "clamav" }
// Available checks if clamd is reachable by attempting a connection.
func (s *ClamdScanner) Available() bool {
conn, err := net.DialTimeout("unix", s.socketPath, 2*time.Second)
if err != nil {
return false
}
_ = conn.Close()
return true
}
// Scan sends the file at path to clamd via INSTREAM and returns the verdict.
func (s *ClamdScanner) Scan(path string) (Verdict, error) {
// #nosec G304 -- path is mail queue file path from mail scanner walk.
f, err := os.Open(path)
if err != nil {
return Verdict{}, fmt.Errorf("opening file: %w", err)
}
defer f.Close()
conn, err := net.DialTimeout("unix", s.socketPath, 5*time.Second)
if err != nil {
return Verdict{}, fmt.Errorf("connecting to clamd: %w", err)
}
defer func() { _ = conn.Close() }()
err = conn.SetDeadline(time.Now().Add(30 * time.Second))
if err != nil {
return Verdict{}, fmt.Errorf("setting deadline: %w", err)
}
// Send INSTREAM command
_, err = conn.Write([]byte("nINSTREAM\n"))
if err != nil {
return Verdict{}, fmt.Errorf("sending INSTREAM: %w", err)
}
// Stream file content in chunks
buf := make([]byte, 8192)
lenBuf := make([]byte, 4)
for {
n, readErr := f.Read(buf)
if n > 0 {
// #nosec G115 -- n is bounded by len(buf)=8192; fits in uint32.
binary.BigEndian.PutUint32(lenBuf, uint32(n))
_, err = conn.Write(lenBuf)
if err != nil {
return Verdict{}, fmt.Errorf("sending chunk length: %w", err)
}
_, err = conn.Write(buf[:n])
if err != nil {
return Verdict{}, fmt.Errorf("sending chunk data: %w", err)
}
}
if readErr == io.EOF {
break
}
if readErr != nil {
return Verdict{}, fmt.Errorf("reading file: %w", readErr)
}
}
// Send terminator (4 zero bytes)
binary.BigEndian.PutUint32(lenBuf, 0)
_, err = conn.Write(lenBuf)
if err != nil {
return Verdict{}, fmt.Errorf("sending terminator: %w", err)
}
// Read response
resp := make([]byte, 4096)
n, err := conn.Read(resp)
if err != nil && err != io.EOF {
return Verdict{}, fmt.Errorf("reading response: %w", err)
}
return parseClamdResponse(string(resp[:n]))
}
// parseClamdResponse parses a clamd INSTREAM response line.
// "stream: OK\n" → clean
// "stream: Win.Trojan.Agent-123 FOUND\n" → infected
func parseClamdResponse(resp string) (Verdict, error) {
resp = strings.TrimSpace(resp)
if strings.HasSuffix(resp, "OK") {
return Verdict{Infected: false}, nil
}
if strings.HasSuffix(resp, "FOUND") {
// Extract signature: "stream: <sig> FOUND"
resp = strings.TrimPrefix(resp, "stream: ")
sig := strings.TrimSuffix(resp, " FOUND")
return Verdict{
Infected: true,
Signature: sig,
Severity: "critical",
}, nil
}
return Verdict{}, fmt.Errorf("unexpected clamd response: %q", resp)
}
package emailav
import (
"context"
"fmt"
"os"
"sync"
"time"
emime "github.com/pidginhost/csm/internal/mime"
"github.com/pidginhost/csm/internal/obs"
)
// Orchestrator runs multiple scanners in parallel against extracted email parts.
type Orchestrator struct {
scanners []Scanner
scanTimeout time.Duration
health *scanQueue
}
// NewOrchestrator creates an orchestrator with the given scanners and per-scan timeout.
func NewOrchestrator(scanners []Scanner, scanTimeout time.Duration) *Orchestrator {
return &Orchestrator{
scanners: scanners,
scanTimeout: scanTimeout,
health: newScanQueue(),
}
}
// ScanParts scans all extracted parts with all available engines.
// Fail-open: unavailable engines, timeouts, and errors are recorded but do not
// mark the message as infected.
func (o *Orchestrator) ScanParts(messageID string, parts []emime.ExtractedPart, partial bool) *ScanResult {
result := &ScanResult{
MessageID: messageID,
ScannedAt: time.Now(),
PartialExtraction: partial,
}
// Determine which engines are available
var available []Scanner
for _, s := range o.scanners {
if s.Available() {
available = append(available, s)
result.EnginesUsed = append(result.EnginesUsed, s.Name())
} else {
result.FailedEngines = append(result.FailedEngines, s.Name())
fmt.Fprintf(os.Stderr, "[emailav] engine %s unavailable\n", s.Name())
}
}
if len(available) == 0 {
// fail-open: no engines available - rate-limit the warning
result.AllEnginesDown = true
return result
}
// Scan each part with all available engines
for _, part := range parts {
findings, timedOut, errored := o.scanPart(part, available)
result.Findings = append(result.Findings, findings...)
result.TimedOutEngines = append(result.TimedOutEngines, timedOut...)
result.ErroredEngines = append(result.ErroredEngines, errored...)
}
result.Infected = len(result.Findings) > 0
return result
}
// scanPart scans a single part with all available engines concurrently.
// Returns findings and lists of engine names that timed out or errored.
func (o *Orchestrator) scanPart(part emime.ExtractedPart, scanners []Scanner) ([]Finding, []string, []string) {
ctx, cancel := context.WithTimeout(context.Background(), o.scanTimeout)
defer cancel()
results := make(chan engineScanResult, len(scanners))
var wg sync.WaitGroup
work := make([]*scanWork, 0, len(scanners))
defer func() {
for _, w := range work {
w.finishDelivery(false)
}
}()
for _, s := range scanners {
work = append(work, o.startScan(ctx, s, part.TempPath, results, &wg))
}
// Close results channel when all scans complete
obs.SafeGo("emailav-drain", func() {
wg.Wait()
close(results)
})
var findings []Finding
var timedOut []string
var errored []string
for r := range results {
r.work.received()
if r.err != nil {
fmt.Fprintf(os.Stderr, "[emailav] %s scan error on %s: %v\n", r.engine, part.Filename, r.err)
if r.timedOut {
timedOut = append(timedOut, r.engine)
} else {
errored = append(errored, r.engine)
}
r.work.finishDelivery(true)
continue // fail-open
}
if r.verdict.Infected {
f := Finding{
Filename: part.Filename,
Engine: r.engine,
Signature: r.verdict.Signature,
Severity: r.verdict.Severity,
}
if part.Nested {
f.Filename = part.ArchiveName + "/" + part.Filename
}
findings = append(findings, f)
}
r.work.finishDelivery(true)
}
return findings, timedOut, errored
}
type engineScanResult struct {
engine string
verdict Verdict
err error
timedOut bool
work *scanWork
}
func (o *Orchestrator) startScan(ctx context.Context, scanner Scanner, path string, results chan<- engineScanResult, wg *sync.WaitGroup) *scanWork {
deadline, _ := ctx.Deadline()
w := o.health.begin(deadline)
wg.Add(1)
obs.SafeGo("emailav-scan", func() {
defer wg.Done()
published := false
defer func() {
if !published {
w.finishDelivery(false)
}
}()
done := make(chan engineScanResult, 1)
obs.SafeGo("emailav-engine", func() {
success := false
defer func() { w.finishEngine(success) }()
w.start()
v, err := scanner.Scan(path)
done <- engineScanResult{engine: scanner.Name(), verdict: v, err: err, work: w}
success = err == nil
})
var r engineScanResult
select {
case r = <-done:
case <-ctx.Done():
r = engineScanResult{engine: scanner.Name(), err: fmt.Errorf("scan timeout"), timedOut: true, work: w}
}
w.publish(r.err != nil)
results <- r
published = true
})
return w
}
package emailav
import (
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"syscall"
"time"
"github.com/pidginhost/csm/internal/quarantinefs"
"github.com/pidginhost/csm/internal/safepath"
)
type movedFile struct {
src string
dst string
}
// vars so tests can force EXDEV and source swaps without depending on
// the host filesystem layout.
var (
moveFileRename = renameSpoolFile
moveFileAfterCrossDeviceCopy = func(string, string) error { return nil }
moveFileSyncDir = quarantinefs.SyncDir
)
// QuarantineEnvelope holds the email envelope info for quarantine metadata.
type QuarantineEnvelope struct {
From string
To []string
Subject string
Direction string
}
// QuarantineMetadata is the JSON sidecar written with quarantined messages.
type QuarantineMetadata struct {
MessageID string `json:"message_id"`
Direction string `json:"direction"`
From string `json:"from"`
To []string `json:"to"`
Subject string `json:"subject"`
QuarantinedAt time.Time `json:"quarantined_at"`
OriginalSpoolDir string `json:"original_spool_dir"`
Findings []Finding `json:"findings"`
PartialScan bool `json:"partial_scan"`
EnginesUsed []string `json:"engines_used"`
}
// Quarantine manages the per-message email quarantine directory.
type Quarantine struct {
baseDir string // e.g. /opt/csm/quarantine/email
allowedSpoolDirs []string
}
// NewQuarantine creates a quarantine manager for the given base directory.
func NewQuarantine(baseDir string) *Quarantine {
return &Quarantine{
baseDir: baseDir,
allowedSpoolDirs: []string{"/var/spool/exim/input", "/var/spool/exim4/input"},
}
}
// QuarantineMessage moves spool files into a per-message quarantine directory
// and writes metadata.json.
func (q *Quarantine) QuarantineMessage(msgID, spoolDir string, result *ScanResult, env QuarantineEnvelope) error {
msgID, err := cleanQuarantineMessageID(msgID)
if err != nil {
return err
}
msgDir := filepath.Join(q.baseDir, msgID)
meta := QuarantineMetadata{
MessageID: msgID,
Direction: env.Direction,
From: env.From,
To: env.To,
Subject: env.Subject,
QuarantinedAt: time.Now(),
OriginalSpoolDir: spoolDir,
Findings: result.Findings,
PartialScan: result.PartialExtraction,
EnginesUsed: result.EnginesUsed,
}
metaData, err := json.MarshalIndent(meta, "", " ")
if err != nil {
return fmt.Errorf("marshaling metadata: %w", err)
}
if err := quarantinefs.EnsureDir(q.baseDir, 0700); err != nil {
return fmt.Errorf("creating quarantine root: %w", err)
}
if err := os.Mkdir(msgDir, 0700); err != nil {
return fmt.Errorf("creating quarantine dir: %w", err)
}
if err := quarantinefs.SyncDir(q.baseDir); err != nil {
return fmt.Errorf("syncing quarantine root: %w", err)
}
metaPath := filepath.Join(msgDir, "metadata.json")
if err := quarantinefs.WriteExclusive(metaPath, bytes.NewReader(metaData), 0600); err != nil {
_ = os.RemoveAll(msgDir)
return fmt.Errorf("writing metadata before moving spool files: %w", err)
}
var moved []movedFile
for _, suffix := range []string{"-H", "-D"} {
src := filepath.Join(spoolDir, msgID+suffix)
dst := filepath.Join(msgDir, msgID+suffix)
if err := moveFile(src, dst); err != nil {
var partial *partialSpoolMove
if errors.As(err, &partial) {
return fmt.Errorf("quarantine partly applied; recovery files and metadata retained at %s: %w", msgDir, err)
}
if !os.IsNotExist(err) {
rollbackErr := rollbackMovedFiles(moved)
if rollbackErr != nil {
return fmt.Errorf("moving spool file %s: %w (rollback failed, recovery files retained at %s: %v)", suffix, err, msgDir, rollbackErr)
}
_ = os.RemoveAll(msgDir)
return fmt.Errorf("moving spool file %s: %w", suffix, err)
}
continue
}
moved = append(moved, movedFile{src: src, dst: dst})
}
if len(moved) == 0 {
_ = os.RemoveAll(msgDir)
return fmt.Errorf("no spool files found for %s", msgID)
}
return nil
}
// ListMessages returns all quarantined email messages.
func (q *Quarantine) ListMessages() ([]QuarantineMetadata, error) {
entries, err := os.ReadDir(q.baseDir)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, fmt.Errorf("reading quarantine dir: %w", err)
}
var msgs []QuarantineMetadata
for _, entry := range entries {
if !entry.IsDir() {
continue
}
meta, err := q.readMetadata(entry.Name())
if err != nil {
continue
}
msgs = append(msgs, *meta)
}
return msgs, nil
}
// GetMessage returns the metadata for a single quarantined message.
func (q *Quarantine) GetMessage(msgID string) (*QuarantineMetadata, error) {
msgID, err := cleanQuarantineMessageID(msgID)
if err != nil {
return nil, err
}
return q.readMetadata(msgID)
}
// ReleaseMessage moves spool files back to the original spool directory
// and removes the quarantine directory.
func (q *Quarantine) ReleaseMessage(msgID string) error {
msgID, err := cleanQuarantineMessageID(msgID)
if err != nil {
return err
}
meta, err := q.readMetadata(msgID)
if err != nil {
return fmt.Errorf("reading metadata: %w", err)
}
spoolDir, err := q.validateReleaseSpoolDir(meta.OriginalSpoolDir)
if err != nil {
return err
}
msgDir := filepath.Join(q.baseDir, msgID)
// The queue header makes the message visible to Exim; restore its body first.
for _, suffix := range []string{"-D", "-H"} {
src := filepath.Join(msgDir, msgID+suffix)
dst := filepath.Join(spoolDir, msgID+suffix)
if err := moveFile(src, dst); err != nil {
// If source doesn't exist, skip (partial quarantine)
if os.IsNotExist(err) {
continue
}
return fmt.Errorf("moving %s back to spool: %w", suffix, err)
}
}
if err := os.RemoveAll(msgDir); err != nil {
return fmt.Errorf("message released but quarantine cleanup failed: %w", err)
}
return quarantinefs.SyncDir(q.baseDir)
}
// DeleteMessage permanently removes a quarantined message. Unlike
// ReleaseMessage it does not require metadata: deleting an orphaned or
// partially-written entry is a legitimate cleanup operation.
func (q *Quarantine) DeleteMessage(msgID string) error {
msgID, err := cleanQuarantineMessageID(msgID)
if err != nil {
return err
}
msgDir := filepath.Join(q.baseDir, msgID)
return os.RemoveAll(msgDir)
}
// CleanExpired removes quarantine directories older than maxAge.
// Returns the number of directories cleaned.
func (q *Quarantine) CleanExpired(maxAge time.Duration) (int, error) {
entries, err := os.ReadDir(q.baseDir)
if err != nil {
if os.IsNotExist(err) {
return 0, nil
}
return 0, err
}
cleaned := 0
cutoff := time.Now().Add(-maxAge)
for _, entry := range entries {
if !entry.IsDir() {
continue
}
meta, err := q.readMetadata(entry.Name())
if err != nil {
continue
}
if meta.QuarantinedAt.Before(cutoff) {
os.RemoveAll(filepath.Join(q.baseDir, entry.Name()))
cleaned++
}
}
return cleaned, nil
}
func (q *Quarantine) readMetadata(msgID string) (*QuarantineMetadata, error) {
msgID, err := cleanQuarantineMessageID(msgID)
if err != nil {
return nil, err
}
metaPath := filepath.Join(q.baseDir, msgID, "metadata.json")
// #nosec G304 -- msgID is restricted to a single path segment under baseDir.
data, err := os.ReadFile(metaPath)
if err != nil {
return nil, err
}
var meta QuarantineMetadata
if err := json.Unmarshal(data, &meta); err != nil {
return nil, err
}
return &meta, nil
}
func cleanQuarantineMessageID(msgID string) (string, error) {
if msgID == "" {
return "", fmt.Errorf("message id is required")
}
if msgID == "." || msgID == ".." || strings.ContainsAny(msgID, `/\`) {
return "", fmt.Errorf("invalid message id")
}
if filepath.Base(msgID) != msgID {
return "", fmt.Errorf("invalid message id")
}
return msgID, nil
}
// moveFile renames src to dst, falling back to a fd-bound copy when
// rename fails (cross-device or other EXDEV-style errors). Callers
// construct dst by filepath.Join under the quarantine base dir
// (config-owned) plus a validated single-segment identifier.
//
// The cross-device path opens src with O_NOFOLLOW, copies content
// from the open fd into a freshly-created dst, and unlinks src by the
// path only after verifying the path still names the same inode. An
// attacker who swaps src for a symlink or replaces the file
// mid-copy is rejected.
func moveFile(src, dst string) error {
if err := quarantinefs.SyncFilePath(src); err != nil {
return err
}
if err := moveFileRename(src, dst); err == nil {
return syncSpoolMove(src, dst)
} else if !errors.Is(err, syscall.EXDEV) {
// Non-EXDEV rename errors are not a cross-device condition;
// surface them directly so the caller does not silently
// fall back to copy semantics for permission or path errors.
return err
}
// #nosec G304 -- src is mail queue path from scanner walk;
// O_NOFOLLOW plus the identity check below catch symlink swaps.
fd, err := os.OpenFile(src, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return fmt.Errorf("opening cross-device source: %w", err)
}
defer fd.Close()
srcInfo, err := fd.Stat()
if err != nil {
return fmt.Errorf("stat cross-device source: %w", err)
}
if !srcInfo.Mode().IsRegular() {
return fmt.Errorf("refusing cross-device copy of non-regular %s", src)
}
// #nosec G304 G306 -- dst is constructed under the quarantine baseDir;
// 0600 keeps the copy private.
dstFile, err := os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o600)
if err != nil {
return fmt.Errorf("creating cross-device destination: %w", err)
}
removeDst := true
defer func() {
if removeDst {
_ = os.Remove(dst)
}
}()
if _, copyErr := io.Copy(dstFile, fd); copyErr != nil {
_ = dstFile.Close()
return fmt.Errorf("cross-device copy body: %w", copyErr)
}
stat := srcInfo.Sys().(*syscall.Stat_t)
if ownerErr := dstFile.Chown(int(stat.Uid), int(stat.Gid)); ownerErr != nil {
_ = dstFile.Close()
return fmt.Errorf("cross-device ownership: %w", ownerErr)
}
// Chown can clear special mode bits, so permissions follow ownership.
if modeErr := dstFile.Chmod(srcInfo.Mode()); modeErr != nil {
_ = dstFile.Close()
return fmt.Errorf("cross-device permissions: %w", modeErr)
}
if timeErr := safepath.SetModTime(dstFile, srcInfo.ModTime()); timeErr != nil {
_ = dstFile.Close()
return fmt.Errorf("cross-device modification time: %w", timeErr)
}
if syncErr := dstFile.Sync(); syncErr != nil {
_ = dstFile.Close()
return fmt.Errorf("cross-device sync: %w", syncErr)
}
if closeErr := dstFile.Close(); closeErr != nil {
return fmt.Errorf("cross-device close: %w", closeErr)
}
if hookErr := moveFileAfterCrossDeviceCopy(src, dst); hookErr != nil {
return fmt.Errorf("cross-device post-copy check: %w", hookErr)
}
pathInfo, err := os.Lstat(src)
if err != nil {
return fmt.Errorf("stat cross-device source path: %w", err)
}
if !sameUnixInode(srcInfo, pathInfo) {
return fmt.Errorf("cross-device source %s changed during copy", src)
}
if err := moveFileSyncDir(filepath.Dir(dst)); err != nil {
return fmt.Errorf("syncing cross-device destination: %w", err)
}
if err := os.Remove(src); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("removing cross-device source: %w", err)
}
removeDst = false
if err := moveFileSyncDir(filepath.Dir(src)); err != nil {
return &partialSpoolMove{src: src, dst: dst, cause: err}
}
return nil
}
type partialSpoolMove struct {
src, dst string
cause error
}
func (e *partialSpoolMove) Error() string {
return fmt.Sprintf("moved %s to %s, but directory sync failed; inspect both paths before retrying: %v", e.src, e.dst, e.cause)
}
func (e *partialSpoolMove) Unwrap() error { return e.cause }
func syncSpoolMove(src, dst string) error {
for _, dir := range []string{filepath.Dir(dst), filepath.Dir(src)} {
if err := moveFileSyncDir(dir); err != nil {
return &partialSpoolMove{src: src, dst: dst, cause: err}
}
}
return nil
}
func renameSpoolFile(src, dst string) error {
source, err := safepath.OpenDir(filepath.Dir(src))
if err != nil {
return err
}
defer func() { _ = source.Close() }()
destination, err := safepath.OpenDir(filepath.Dir(dst))
if err != nil {
return err
}
defer func() { _ = destination.Close() }()
return source.RenameTo(filepath.Base(src), destination, filepath.Base(dst))
}
// sameUnixInode compares two FileInfos via underlying syscall.Stat_t so
// the cross-device move can verify the source path still names the
// same inode it opened. Falls back to false if either Sys() does not
// expose Stat_t (e.g. non-Linux builds).
func sameUnixInode(a, b os.FileInfo) bool {
if a == nil || b == nil {
return false
}
if os.SameFile(a, b) {
return true
}
as, aok := a.Sys().(*syscall.Stat_t)
bs, bok := b.Sys().(*syscall.Stat_t)
return aok && bok && as.Dev == bs.Dev && as.Ino == bs.Ino
}
func rollbackMovedFiles(moved []movedFile) error {
for i := len(moved) - 1; i >= 0; i-- {
if err := moveFile(moved[i].dst, moved[i].src); err != nil {
return err
}
}
return nil
}
func (q *Quarantine) validateReleaseSpoolDir(spoolDir string) (string, error) {
cleanDir := filepath.Clean(spoolDir)
if cleanDir == "" || !filepath.IsAbs(cleanDir) {
return "", fmt.Errorf("invalid original spool directory")
}
// Fail closed when the supplied path will not resolve. The previous
// behaviour fell back to the unresolved literal, which let a path
// that does not currently exist on disk still match an allowed-list
// entry that also does not exist (e.g. on a host where only one of
// the default exim spool defaults is installed). A release that
// targets a path we cannot resolve cannot have its identity
// confirmed, so it must not proceed.
resolvedDir, err := filepath.EvalSymlinks(cleanDir)
if err != nil {
return "", fmt.Errorf("resolving original spool directory %q: %w", cleanDir, err)
}
for _, allowed := range q.allowedSpoolDirs {
cleanAllowed := filepath.Clean(allowed)
// Allowed entries are operator-config trusted defaults; the
// shipped list intentionally covers both exim and exim4, so a
// host that only has one of them must silently skip the
// non-installed entry rather than falling back to the literal
// (which would silently widen the trust boundary). The input
// path was already required to resolve above, so a missing
// allowed entry can never alias the resolved input.
resolvedAllowed, err := filepath.EvalSymlinks(cleanAllowed)
if err != nil {
continue
}
if resolvedDir == resolvedAllowed {
return resolvedDir, nil
}
}
return "", fmt.Errorf("original spool directory is not trusted: %s", cleanDir)
}
package emailav
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type scanQueue struct {
mu sync.Mutex
pending map[*scanWork]struct{}
losses *queuehealth.Tracker
}
type scanWork struct {
queue *scanQueue
queued, deadline time.Time
started, returned, ready, receiving time.Time
engineDone, deliveryDone, failed bool
}
func newScanQueue() *scanQueue {
return &scanQueue{pending: make(map[*scanWork]struct{}), losses: queuehealth.New(0, time.Minute)}
}
func (q *scanQueue) begin(deadline time.Time) *scanWork {
w := &scanWork{queue: q, queued: time.Now(), deadline: deadline}
q.mu.Lock()
q.pending[w] = struct{}{}
q.mu.Unlock()
return w
}
func (w *scanWork) start() {
w.queue.mu.Lock()
w.started = time.Now()
w.queue.mu.Unlock()
}
func (w *scanWork) finishEngine(success bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
w.engineDone = true
w.returned = time.Now()
if !success {
w.failLocked()
}
w.releaseLocked()
}
func (w *scanWork) publish(failed bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
w.ready = time.Now()
if failed {
w.failLocked()
}
}
func (w *scanWork) received() {
w.queue.mu.Lock()
w.receiving = time.Now()
w.queue.mu.Unlock()
}
func (w *scanWork) finishDelivery(success bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
if w.deliveryDone {
return
}
w.deliveryDone = true
if !success {
w.failLocked()
}
w.releaseLocked()
}
func (w *scanWork) failLocked() {
if !w.failed {
w.failed = true
w.queue.losses.Lose(time.Now(), 1)
}
}
func (w *scanWork) releaseLocked() {
// Timeout releases result delivery while Scan may still run. Conversely,
// a returned engine can leave work in either buffered result handoff.
if w.engineDone && w.deliveryDone {
delete(w.queue.pending, w)
}
}
// QueueStatuses reads memory without probing an engine or taking its locks.
func (o *Orchestrator) QueueStatuses(now time.Time) map[string]queuehealth.Status {
q := o.health
q.mu.Lock()
defer q.mu.Unlock()
s := q.losses.Snapshot(now)
// Concurrent ScanParts calls and engines outliving their callers have no
// fixed global admission limit, despite each result channel being bounded.
s.CapacityUnavailable = true
var waitingLate, runningLate bool
for w := range q.pending {
origin, deadline, running := w.queued, w.deadline, false
switch {
case !w.engineDone:
if !w.started.IsZero() {
origin, running = w.started, true
}
case !w.receiving.IsZero():
origin, running = w.receiving, true
deadline = origin.Add(time.Minute)
default:
origin = w.returned
if !w.ready.IsZero() && w.ready.Before(origin) {
origin = w.ready
}
deadline = origin.Add(time.Minute)
}
if running {
s.InFlight++
s.ProcessingSeconds = max(s.ProcessingSeconds, now.Sub(origin).Seconds())
runningLate = runningLate || !now.Before(deadline)
} else {
s.Depth++
s.LagSeconds = max(s.LagSeconds, now.Sub(origin).Seconds())
waitingLate = waitingLate || !now.Before(deadline)
}
}
switch {
case waitingLate:
s.Reason = "backlog_lag"
case runningLate:
s.Reason = "processing_lag"
}
if s.Reason != "" {
s.Status = "degraded"
}
return map[string]queuehealth.Status{"scans": s}
}
//go:build !yara
package emailav
import "github.com/pidginhost/csm/internal/yara"
// YaraXScanner is a no-op stub when YARA-X is not compiled in. The
// constructor still accepts a yara.Backend so callers can pass
// yara.Active() uniformly across build tags.
type YaraXScanner struct{}
// NewYaraXScanner returns a scanner that is never available.
func NewYaraXScanner(_ yara.Backend) *YaraXScanner {
return &YaraXScanner{}
}
// NewActiveYaraXScanner returns the same no-op scanner in non-YARA builds.
func NewActiveYaraXScanner() *YaraXScanner {
return &YaraXScanner{}
}
func (s *YaraXScanner) Name() string { return "yara-x" }
func (s *YaraXScanner) Available() bool { return false }
func (s *YaraXScanner) Scan(_ string) (Verdict, error) {
return Verdict{}, nil
}
package emailspool
import (
"bufio"
"errors"
"fmt"
"io"
"os"
"strconv"
"strings"
"golang.org/x/net/idna"
)
// ExtractDomain returns the lowercased, IDN-normalised domain portion of an
// RFC 5322 address or display-name form. Returns "" on parse failure.
//
// Quoted local parts ("a@b"@example.com) are handled by treating the address
// as the substring after the LAST unquoted '@'.
func ExtractDomain(s string) string {
s = strings.TrimSpace(s)
if s == "" {
return ""
}
if i := strings.LastIndex(s, "<"); i >= 0 {
if j := strings.LastIndex(s, ">"); j > i {
s = s[i+1 : j]
}
}
at := lastUnquotedAt(s)
if at < 0 {
return ""
}
domain := strings.TrimSpace(s[at+1:])
domain = strings.ToLower(domain)
// On IDNA failure keep the raw lowercased domain rather than returning "".
// Callers gate mismatch checks on a non-empty domain, so an empty return
// would let a malformed attacker-controlled label bypass those checks.
if ascii, err := idna.ToASCII(domain); err == nil {
domain = ascii
}
return domain
}
// lastUnquotedAt returns the index of the rightmost '@' character that is not
// inside a double-quoted segment, or -1 if none.
func lastUnquotedAt(s string) int {
inQuote := false
last := -1
for i := 0; i < len(s); i++ {
c := s[i]
switch c {
case '\\':
i++ // skip the next byte (quoted-pair)
case '"':
inQuote = !inQuote
case '@':
if !inQuote {
last = i
}
}
}
return last
}
// IsSubdomainOrEqual reports whether candidate is base or a subdomain of base.
// Both inputs are case-insensitive. Empty inputs return false.
func IsSubdomainOrEqual(candidate, base string) bool {
if candidate == "" || base == "" {
return false
}
c := strings.ToLower(strings.TrimSuffix(candidate, "."))
b := strings.ToLower(strings.TrimSuffix(base, "."))
if c == b {
return true
}
return strings.HasSuffix(c, "."+b)
}
// MaxSpoolHeaderBytes bounds how much of an Exim -H file we read.
// 32 KiB covers rich messages with DKIM, ARC, and folded MIME headers
// without unbounded memory.
const MaxSpoolHeaderBytes = 32 * 1024
// Headers is the parsed envelope + interesting RFC 5322 fields from a
// cPanel-Exim spool -H file. EnvelopeUser comes from the file's line 2
// (the local UID under which Exim accepted the message); the RFC 5322
// fields come from the message's mail headers section. Empty string means
// "header absent" -- callers should not synthesise defaults from absence.
type Headers struct {
EnvelopeUser string
EnvelopeUID int
From string
ReplyTo string
Subject string
XPHPScript string
XMailer string
UserAgent string
MessageID string
// Recipients holds the envelope recipient addresses when the Exim -H
// recipient block can be unambiguously located; empty when the block is
// absent or its shape cannot be validated. Consumers that gate on
// recipient diversity must treat empty as "unknown" and fail open.
Recipients []string
}
// ParseHeaders reads the given Exim -H file and returns a Headers.
// Returns the parse error if the file cannot be opened or is structurally
// invalid; otherwise missing individual headers leave the corresponding
// Headers field empty.
func ParseHeaders(path string) (Headers, error) {
// #nosec G304 -- path is supplied by the daemon's Exim spool
// watcher and resolved from /var/spool/exim/input/, an
// operator-trusted directory enumerated by the spool walker.
f, err := os.Open(path)
if err != nil {
return Headers{}, err
}
defer f.Close()
h, err := ParseHeadersReader(f)
if err != nil {
return Headers{}, fmt.Errorf("parse %s: %w", path, err)
}
return h, nil
}
// ParseHeadersReader is the io.Reader form of ParseHeaders for callers that
// already have the spool bytes in memory or behind a custom seam (e.g. the
// checks package's osFS abstraction). It applies the same Exim -H parsing
// rules as ParseHeaders -- envelope preamble, blank-line separator,
// "NNNX " prefixed RFC 5322 headers -- and is bounded by MaxSpoolHeaderBytes
// per token; oversize input returns bufio.ErrTooLong.
func ParseHeadersReader(r io.Reader) (Headers, error) {
var h Headers
// Per-line memory is bounded by the scanner's max buffer
// (MaxSpoolHeaderBytes); a token larger than that returns
// bufio.ErrTooLong. We deliberately do NOT wrap r in an io.LimitReader:
// when LimitReader returns EOF mid-token, bufio.Scanner emits the
// partial token without error and oversize spool files are silently
// truncated. That hides the failure from operators.
sc := bufio.NewScanner(r)
sc.Buffer(make([]byte, 0, 8192), MaxSpoolHeaderBytes)
// Line 1: msgID-H (we ignore the value; presence is enough)
if !sc.Scan() {
return Headers{}, errors.New("empty spool file")
}
// Line 2: "<user> <uid> <gid>"
if !sc.Scan() {
return Headers{}, errors.New("missing envelope user line")
}
fields := splitFields(sc.Text())
if len(fields) < 2 {
return Headers{}, fmt.Errorf("malformed envelope user line: %q", sc.Text())
}
h.EnvelopeUser = fields[0]
if uid, err := strconv.Atoi(fields[1]); err == nil {
h.EnvelopeUID = uid
}
// Skip remaining envelope metadata until the blank line that separates
// it from the RFC 5322 header section. Exim's -H format places mail
// headers AFTER a blank line that follows the recipient list; recipients
// are preceded by a numeric count line.
inHeaders := false
var preamble []string
lastHeader := ""
skippingDeleted := false
for sc.Scan() {
line := sc.Text()
if !inHeaders {
if line == "" {
inHeaders = true
continue
}
preamble = append(preamble, line)
continue
}
// RFC 5322 header section. Each header line in Exim's -H format
// starts with "<count><flag> <name>: <value>" where count is a
// variable-width decimal byte count for the full stored header
// (including folded continuation lines) and flag is a single Exim
// marker ('T', 'F', 'R', '*', space, etc.).
if isFoldedHeaderLine(line) {
if !skippingDeleted && lastHeader != "" {
appendEximHeaderContinuation(&h, lastHeader, line)
}
continue
}
lastHeader = ""
skippingDeleted = false
name, value, deleted, ok := parseEximHeaderLine(line)
if deleted {
skippingDeleted = true
continue
}
if !ok {
continue
}
if canonical, handled := setEximHeaderValue(&h, name, value); handled {
lastHeader = canonical
}
}
if err := sc.Err(); err != nil {
return Headers{}, fmt.Errorf("scan spool: %w", err)
}
if !inHeaders {
return Headers{}, errors.New("missing header section separator")
}
h.Recipients = extractEximRecipients(preamble)
return h, nil
}
// extractEximRecipients recovers the envelope recipient addresses from the
// Exim -H preamble (the lines between the envelope-user line and the blank
// header separator). Exim writes the recipient block last: a line holding only
// the recipient count N, then exactly N recipient lines that run to the end of
// the preamble. Anchoring on "index + 1 + N == len(preamble)" plus an
// address-shape check on every claimed recipient locates the block without a
// full grammar and tolerates option/ACL lines above it. Any ambiguity, including
// a malformed anchored candidate, returns nil so callers treat recipients as
// unknown and fail open.
func extractEximRecipients(preamble []string) []string {
var candidate []string
ambiguous := false
for i, line := range preamble {
n, ok := parseBareUint(line)
if !ok || n < 1 || i+1+n != len(preamble) {
continue
}
rcpts := make([]string, 0, n)
valid := true
for _, r := range preamble[i+1:] {
addr := firstField(r)
if addr == "" || !strings.Contains(addr, "@") {
valid = false
break
}
rcpts = append(rcpts, addr)
}
if !valid {
ambiguous = true
continue
}
if candidate != nil {
ambiguous = true
continue
}
candidate = rcpts
}
if ambiguous {
return nil
}
return candidate
}
// parseBareUint reports whether s is a single non-negative integer token with
// no other characters. The length cap rejects timestamps and other long
// numeric option values that are not recipient counts.
func parseBareUint(s string) (int, bool) {
s = strings.TrimSpace(s)
if s == "" || len(s) > 9 {
return 0, false
}
for i := 0; i < len(s); i++ {
if s[i] < '0' || s[i] > '9' {
return 0, false
}
}
n, err := strconv.Atoi(s)
if err != nil {
return 0, false
}
return n, true
}
func firstField(s string) string {
s = strings.TrimSpace(s)
if i := strings.IndexAny(s, " \t"); i >= 0 {
s = s[:i]
}
return s
}
// parseEximHeaderLine returns the header name, value, deletion marker, and
// prefix-valid marker for an Exim -H header line of the form
// "NNNF Header-Name: value". The prefix is a variable-width decimal length
// (Exim writes the byte count of the full stored header, including folded
// continuation lines), followed by a flag byte ('T', 'F', 'R', '*', space,
// etc.) and a space separator; the whole prefix is stripped before splitting
// on the colon.
func parseEximHeaderLine(line string) (name string, value string, deleted bool, ok bool) {
i := 0
for i < len(line) && isDigit(line[i]) {
i++
}
if i == 0 || i+1 >= len(line) {
return "", "", false, false
}
// After the digit run: a flag byte (letter, or space for unflagged
// headers such as X-PHP-Script in real cPanel-Exim spool output) then a
// single space separator.
if !isEximHeaderFlag(line[i]) || line[i+1] != ' ' {
return "", "", false, false
}
if line[i] == '*' {
return "", "", true, true
}
rest := line[i+2:]
colon := indexByte(rest, ':')
if colon < 0 {
return "", "", false, true
}
name = rest[:colon]
if colon+1 < len(rest) {
value = trimLeadingSpace(rest[colon+1:])
}
return name, value, false, true
}
func isFoldedHeaderLine(line string) bool {
return len(line) > 0 && (line[0] == ' ' || line[0] == '\t')
}
func isEximHeaderFlag(b byte) bool {
return b == ' ' || b == '*' || isLetter(b)
}
func setEximHeaderValue(h *Headers, name, value string) (string, bool) {
switch strings.ToLower(name) {
case "from":
h.From = value
return "from", true
case "reply-to":
h.ReplyTo = value
return "reply-to", true
case "subject":
h.Subject = value
return "subject", true
case "x-php-script":
h.XPHPScript = value
return "x-php-script", true
case "x-mailer":
h.XMailer = value
return "x-mailer", true
case "user-agent":
h.UserAgent = value
return "user-agent", true
case "message-id":
h.MessageID = value
return "message-id", true
default:
return "", false
}
}
func appendEximHeaderContinuation(h *Headers, name, line string) {
value := trimLeadingSpace(line)
switch name {
case "from":
h.From = appendFoldedHeaderValue(h.From, value)
case "reply-to":
h.ReplyTo = appendFoldedHeaderValue(h.ReplyTo, value)
case "subject":
h.Subject = appendFoldedHeaderValue(h.Subject, value)
case "x-php-script":
h.XPHPScript = appendFoldedHeaderValue(h.XPHPScript, value)
case "x-mailer":
h.XMailer = appendFoldedHeaderValue(h.XMailer, value)
case "user-agent":
h.UserAgent = appendFoldedHeaderValue(h.UserAgent, value)
case "message-id":
h.MessageID = appendFoldedHeaderValue(h.MessageID, value)
}
}
func appendFoldedHeaderValue(current, continuation string) string {
if current == "" {
return continuation
}
if continuation == "" {
return current
}
return current + " " + continuation
}
func isDigit(b byte) bool { return b >= '0' && b <= '9' }
func isLetter(b byte) bool { return (b >= 'a' && b <= 'z') || (b >= 'A' && b <= 'Z') }
func indexByte(s string, b byte) int {
for i := 0; i < len(s); i++ {
if s[i] == b {
return i
}
}
return -1
}
func trimLeadingSpace(s string) string {
i := 0
for i < len(s) && (s[i] == ' ' || s[i] == '\t') {
i++
}
return s[i:]
}
func splitFields(s string) []string {
var out []string
start := -1
for i := 0; i < len(s); i++ {
if s[i] == ' ' || s[i] == '\t' {
if start >= 0 {
out = append(out, s[start:i])
start = -1
}
} else if start < 0 {
start = i
}
}
if start >= 0 {
out = append(out, s[start:])
}
return out
}
package emailspool
import (
"fmt"
"net"
"os"
"path/filepath"
"strings"
"sync"
"gopkg.in/yaml.v3"
)
// Policies is the loaded contents of policies/email/*.yaml. Used by Stage 1
// (mailer classes, http_proxy_ranges) and Stage 2 (policy_blocks). Each
// data file has its own schema; failing to parse one file does not abort
// the whole load -- each category degrades independently.
type Policies struct {
mu sync.RWMutex
suspiciousMail []string
safeMail []string
proxyNets []*net.IPNet
// selfIPs are the host's own interface addresses. PHP-relay Path 4
// (HTTP-IP fanout) treats them as proxy-equivalent so WordPress cron
// and any other local loopback-to-public traffic does not page on
// "one IP triggered N scripts". Populated by RefreshSelfIPs.
selfIPs []net.IP
}
// hostIPsFunc returns the set of non-loopback host IPs to treat as self.
// Package-level so tests can inject deterministic addresses.
var hostIPsFunc = enumerateHostIPs
// enumerateHostIPs reads every IPv4/IPv6 address bound to a non-loopback
// interface. Errors from net.InterfaceAddrs (extremely rare; only on syscall
// failure) drop us to an empty list; loopback fallback in IsProxyIP still
// covers 127/8 and ::1.
func enumerateHostIPs() []net.IP {
addrs, err := net.InterfaceAddrs()
if err != nil {
return nil
}
out := make([]net.IP, 0, len(addrs))
for _, a := range addrs {
var ip net.IP
switch v := a.(type) {
case *net.IPNet:
ip = v.IP
case *net.IPAddr:
ip = v.IP
}
if ip == nil || ip.IsLoopback() || ip.IsLinkLocalUnicast() {
continue
}
out = append(out, ip)
}
return out
}
// RefreshSelfIPs re-enumerates host addresses and replaces the cached set.
// Safe to call concurrently with IsProxyIP. LoadPolicies / Reload call this
// automatically so SIGHUP picks up cPanel alias IP additions without code
// changes to callers; exposed for tests and the daemon Flow E ticker.
func (p *Policies) RefreshSelfIPs() {
ips := hostIPsFunc()
p.mu.Lock()
p.selfIPs = ips
p.mu.Unlock()
}
type mailerClassesYAML struct {
Version int `yaml:"version"`
Suspicious []string `yaml:"suspicious"`
Safe []string `yaml:"safe"`
}
//nolint:unused // consumed by E2/E3
type httpProxyRangesYAML struct {
Version int `yaml:"version"`
CIDRs []string `yaml:"cidrs"`
}
// LoadPolicies reads all known policy files in dir. Missing files are
// treated as "no entries"; corrupt files return an error per file but do
// not abort the load. The returned Policies is safe for concurrent use.
func LoadPolicies(dir string) (*Policies, error) {
p := &Policies{}
if err := p.load(dir); err != nil {
return p, err
}
return p, nil
}
func (p *Policies) load(dir string) error {
p.mu.Lock()
// Refresh self IPs alongside file-backed policies so SIGHUP picks up
// any cPanel alias IP additions that landed since startup.
p.selfIPs = hostIPsFunc()
defer p.mu.Unlock()
var firstErr error
setErr := func(err error) {
if firstErr == nil {
firstErr = err
}
}
// mailer_classes.yaml
// #nosec G304 -- dir is the operator-supplied policy directory; filename is a fixed literal under it.
if data, err := os.ReadFile(filepath.Join(dir, "mailer_classes.yaml")); err == nil {
var raw mailerClassesYAML
if uerr := yaml.Unmarshal(data, &raw); uerr != nil {
setErr(fmt.Errorf("parse mailer_classes.yaml: %w", uerr))
} else {
p.suspiciousMail = lowerList(raw.Suspicious)
p.safeMail = lowerList(raw.Safe)
}
} else if !os.IsNotExist(err) {
setErr(fmt.Errorf("read mailer_classes.yaml: %w", err))
}
// http_proxy_ranges.yaml
// #nosec G304 -- dir is the operator-supplied policy directory; filename is a fixed literal under it.
if data, err := os.ReadFile(filepath.Join(dir, "http_proxy_ranges.yaml")); err == nil {
var raw httpProxyRangesYAML
if uerr := yaml.Unmarshal(data, &raw); uerr != nil {
setErr(fmt.Errorf("parse http_proxy_ranges.yaml: %w", uerr))
} else {
nets := make([]*net.IPNet, 0, len(raw.CIDRs))
for _, c := range raw.CIDRs {
_, n, perr := net.ParseCIDR(strings.TrimSpace(c))
if perr != nil {
setErr(fmt.Errorf("invalid CIDR %q in http_proxy_ranges.yaml: %w", c, perr))
continue
}
nets = append(nets, n)
}
p.proxyNets = nets
}
} else if !os.IsNotExist(err) {
setErr(fmt.Errorf("read http_proxy_ranges.yaml: %w", err))
}
return firstErr
}
// Reload refreshes from dir while keeping previous values for any category
// whose new file is corrupt (existing reload-error contract). Used by SIGHUP.
//
//nolint:unused // consumed by E3
func (p *Policies) Reload(dir string) error {
return p.load(dir)
}
// MailerSuspicious reports whether x-mailer header matches any suspicious
// substring. Substrings, not exact match.
func (p *Policies) MailerSuspicious(xMailer string) bool {
if xMailer == "" {
return false
}
p.mu.RLock()
defer p.mu.RUnlock()
low := strings.ToLower(xMailer)
for _, s := range p.suspiciousMail {
if strings.Contains(low, s) {
return true
}
}
return false
}
// MailerSafe reports whether x-mailer matches any safe substring.
func (p *Policies) MailerSafe(xMailer string) bool {
if xMailer == "" {
return false
}
p.mu.RLock()
defer p.mu.RUnlock()
low := strings.ToLower(xMailer)
for _, s := range p.safeMail {
if strings.Contains(low, s) {
return true
}
}
return false
}
// IsProxyIP reports whether ip falls within any configured CDN/proxy CIDR,
// is one of the host's own interface addresses, or is a loopback address.
// Used by Path 4 to skip fanout counting for IPs that are CDN front IPs,
// the local host (WordPress cron, panel-internal callbacks), or 127/::1.
func (p *Policies) IsProxyIP(ip string) bool {
if ip == "" {
return false
}
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
if parsed.IsLoopback() {
return true
}
p.mu.RLock()
defer p.mu.RUnlock()
for _, n := range p.proxyNets {
if n.Contains(parsed) {
return true
}
}
for _, self := range p.selfIPs {
if self.Equal(parsed) {
return true
}
}
return false
}
func lowerList(in []string) []string {
out := make([]string, 0, len(in))
for _, s := range in {
s = strings.ToLower(strings.TrimSpace(s))
if s != "" {
out = append(out, s)
}
}
return out
}
// Package eximlog extracts the connecting client from Exim main log lines.
//
// Exim renders the peer as `hostname (HELO) [IP]:port`, prefixed with `H=`
// in most records but bare in a few (authenticator failures, TLS errors).
// Everything before the bracketed address is attacker-influenced. The HELO
// may be an RFC 5321 address literal such as `[203.0.113.9]`, and servers that
// accept junk HELO values can log delimiter-like text there. The message
// Subject can carry brackets too. Every consumer that turns a log line into
// an IP for blocking or reputation scoring must therefore go through this
// package rather than grabbing the first bracketed token.
package eximlog
import (
"net"
"strings"
)
// ClientIP returns the connecting client's IP from an Exim log line, or ""
// when the line carries none. It reads either a real H= field or one of the
// known records that Exim writes through host_and_ident without an H= prefix.
func ClientIP(line string) string {
hStart, hasHField := HFieldStart(line)
bareStart, bareMarker := unprefixedClientStart(line)
if bareMarker >= 0 && (!hasHField || bareMarker < hFieldMarkerStart(line, hStart)) {
return hostAndIdentClientIP(line[bareStart:])
}
if hasHField {
return HFieldClientIP(line[hStart:])
}
return ""
}
func hFieldMarkerStart(line string, valueStart int) int {
if strings.HasPrefix(line, "H=") {
return 0
}
return valueStart - len(" H=")
}
// unprefixedClientStart recognizes the Exim records that render
// host_and_ident(FALSE) directly. A marker inside T= is message data, not a
// peer field. markerStart is returned separately so ClientIP can prefer a
// genuine H= field that occurs earlier on the line.
func unprefixedClientStart(line string) (start, markerStart int) {
markerStart = -1
offset, fields := peerFields(line)
t := strings.Index(fields, " T=")
for _, marker := range []string{
"authenticator failed for ",
"TLS error on connection from ",
"SMTP connection from ",
} {
idx := strings.Index(fields, marker)
if idx < 0 || (t >= 0 && t < idx) {
continue
}
if markerStart < 0 || offset+idx < markerStart {
markerStart = offset + idx
start = offset + idx + len(marker)
}
}
return start, markerStart
}
// peerFields bounds host-field searches to reception metadata on arrivals.
// The sender, authentication and post-size fields can contain client data.
// Other records are searched from the start.
func peerFields(line string) (int, string) {
accept := strings.Index(line, " <= ")
if accept < 0 || !arrivalPrefix(line[:accept]) {
return 0, line
}
sender := accept + len(" <= ")
offset := sender + envelopeEnd(line[sender:])
fields := line[offset:]
for _, boundary := range []string{" A=", " S="} {
if end := strings.Index(fields, boundary); end >= 0 {
fields = fields[:end]
}
}
return offset, fields
}
// HFieldStart returns the offset just past the H= marker and true when the
// line carries a real H= field. An H= that appears after T= is inside the
// Subject and is ignored, as is one in an arrival's message data.
func HFieldStart(line string) (int, bool) {
if strings.HasPrefix(line, "H=") {
return len("H="), true
}
offset, fields := peerFields(line)
if h := strings.Index(fields, " H="); h >= 0 {
if t := strings.Index(fields, " T="); t >= 0 && t < h {
return 0, false
}
return offset + h + len(" H="), true
}
return 0, false
}
// HFieldClientIP returns the connecting address inside an H= value.
func HFieldClientIP(s string) string {
ip, _ := HFieldClientIPAndEnd(s)
return ip
}
// HFieldClientIPAndEnd returns the connecting address and the byte offset
// immediately after its closing bracket. The offset lets callers discard the
// entire attacker-controlled H= value before parsing later Exim fields. A
// candidate must be followed by a real H= boundary, and a second plausible
// candidate makes the field ambiguous instead of letting junk HELO text win.
func HFieldClientIPAndEnd(s string) (string, int) {
// Remote ident follows the peer. A U= marker before a candidate means
// greeting delimiters hid the real peer and ident boundary. Text after
// the peer is not checked this way: subjects, addresses and login names
// placed there by any sender must not remove the connecting address.
identStart := strings.Index(s, " U=")
parenDepth := 0
quoted := false
client := ""
clientEnd := 0
for i := 0; i < len(s); i++ {
if quoted {
if s[i] == '\\' && i+1 < len(s) {
i++
continue
}
if s[i] == '"' {
quoted = false
}
continue
}
// Authentication and later message fields can contain client data.
// Their address literals and delimiters do not describe the peer.
if parenDepth == 0 && (strings.HasPrefix(s[i:], " A=") || strings.HasPrefix(s[i:], " S=")) {
break
}
switch s[i] {
case '(':
parenDepth++
case ')':
if parenDepth == 0 {
return "", 0
}
parenDepth--
case '"':
if parenDepth == 0 {
quoted = true
}
case '[':
end := strings.IndexByte(s[i+1:], ']')
if end < 0 {
return "", 0
}
if parenDepth == 0 && !interfaceAddressAt(s, i) {
candidate := s[i+1 : i+1+end]
after := s[i+1+end+1:]
if net.ParseIP(candidate) != nil && hFieldClientIPTerminated(after) {
if client != "" || (identStart >= 0 && identStart < i) {
return "", 0
}
client = candidate
clientEnd = i + end + 2
}
}
i += end + 1
}
}
if parenDepth != 0 || quoted {
return "", 0
}
return client, clientEnd
}
func beginsNextField(s string) bool {
if len(s) == 0 || (s[0] != ' ' && s[0] != '\t') {
return false
}
rest := strings.TrimLeft(s, " \t")
if strings.HasPrefix(rest, "for ") {
return true
}
eq := strings.IndexByte(rest, '=')
if eq <= 0 || eq > 3 {
return false
}
for i := 0; i < eq; i++ {
c := rest[i]
if (c < 'A' || c > 'Z') && (c < 'a' || c > 'z') {
return false
}
}
return true
}
func hFieldClientIPTerminated(s string) bool {
rest := withoutLoggedPort(s)
if after, ok := strings.CutPrefix(rest, " TFO"); ok {
rest = strings.TrimPrefix(after, "*")
}
return rest == "" || beginsNextField(rest) ||
strings.HasPrefix(rest, " authenticator failed") ||
strings.HasPrefix(rest, " rejected RCPT")
}
func hostAndIdentClientIPTerminated(s string) bool {
rest := withoutLoggedPort(s)
if rest == "" || beginsNextField(rest) || strings.HasPrefix(rest, ": ") {
return true
}
for _, suffix := range []string{" (", " lost", " D=", " closed"} {
if strings.HasPrefix(rest, suffix) {
return true
}
}
return false
}
func withoutLoggedPort(s string) string {
if len(s) < 2 || s[0] != ':' || s[1] < '0' || s[1] > '9' {
return s
}
i := 2
for i < len(s) && s[i] >= '0' && s[i] <= '9' {
i++
}
return s[i:]
}
// hostAndIdentClientIP returns the connecting address from Exim's unprefixed
// host_and_ident output. Exim encloses the HELO in parentheses before the
// client, so an address literal inside that group is attacker text. More than
// one plausible peer or malformed parentheses are rejected; otherwise a junk
// HELO could make CSM block an address supplied by the peer.
func hostAndIdentClientIP(s string) string {
// host_and_ident writes remote U= after the peer, as in an H= record.
identStart := strings.Index(s, " U=")
parenDepth := 0
quoted := false
client := ""
for i := 0; i < len(s); i++ {
if quoted {
if s[i] == '\\' && i+1 < len(s) {
i++
continue
}
if s[i] == '"' {
quoted = false
}
continue
}
switch s[i] {
case '(':
parenDepth++
case ')':
if parenDepth == 0 {
return ""
}
parenDepth--
case '"':
if parenDepth == 0 {
quoted = true
}
case '[':
end := strings.IndexByte(s[i+1:], ']')
if end < 0 {
return ""
}
if parenDepth == 0 && !interfaceAddressAt(s, i) {
candidate := s[i+1 : i+1+end]
after := s[i+1+end+1:]
if net.ParseIP(candidate) != nil && hostAndIdentClientIPTerminated(after) {
if client != "" || (identStart >= 0 && identStart < i) {
return ""
}
// Failure details, including the attempted login name,
// follow the peer and are not host information.
if failureDetailsFollow(after) {
return candidate
}
client = candidate
}
}
i += end + 1
}
}
if parenDepth != 0 || quoted {
return ""
}
return client
}
// failureDetailsFollow reports whether s, the text after a peer's closing
// bracket, holds only the logged port, local interface and connection ID
// before the ": " that starts failure details. Remote ident text can contain
// that separator, so a U= field never qualifies.
func failureDetailsFollow(s string) bool {
rest := withoutLoggedPort(s)
if strings.HasPrefix(rest, " I=[") {
end := strings.IndexByte(rest, ']')
if end < 0 || net.ParseIP(rest[len(" I=["):end]) == nil {
return false
}
rest = withoutLoggedPort(rest[end+1:])
}
if strings.HasPrefix(rest, " Ci=") {
digits := rest[len(" Ci="):]
n := 0
for n < len(digits) && digits[n] >= '0' && digits[n] <= '9' {
n++
}
if n == 0 {
return false
}
rest = digits[n:]
}
return strings.HasPrefix(rest, ": ")
}
func interfaceAddressAt(s string, bracket int) bool {
return bracket >= len(" I=") && s[bracket-len(" I="):bracket] == " I="
}
package eximlog
import (
"strings"
"time"
)
// AuthenticatedUser reads the Dovecot authentication identity on an Exim
// arrival record. The envelope sender, HELO and quoted fields are not proof
// of authentication. The complete mainlog prefix must identify an arrival.
func AuthenticatedUser(line string) string {
fields, ok := submissionFields(line)
if !ok {
return ""
}
return fields.auth
}
// Submitter returns an authenticated identity, or the local Exim caller on
// a non-network P=local arrival. Records containing a remote U= identity
// are ambiguous and cannot prove a submitter. The caller must resolve the
// returned identity against the local inventory before assigning a tenant.
func Submitter(line string) string {
fields, ok := submissionFields(line)
if !ok {
return ""
}
if fields.auth != "" {
return fields.auth
}
if !fields.remote && !fields.authSeen && fields.protocol == "local" {
return fields.user
}
return ""
}
type submitFields struct {
auth, user, protocol string
remote, authSeen bool
}
func submissionFields(line string) (submitFields, bool) {
var out submitFields
accept := strings.Index(line, " <= ")
if accept < 0 || !arrivalPrefix(line[:accept]) {
return out, false
}
rest := line[accept+len(" <= "):]
end := envelopeEnd(rest)
if end <= 0 || end == len(rest) {
return out, false
}
rest = rest[end:]
seen := map[string]bool{}
metadata:
for rest != "" {
rest = strings.TrimLeft(rest, " \t\r\n")
// Exim writes submission metadata before the message size. Later
// fields contain message data, including addr-spec message IDs with
// quoted words that can resemble authentication or local-user fields.
if rest == "" || strings.HasPrefix(rest, "S=") || strings.HasPrefix(rest, "T=") || strings.HasPrefix(rest, "for ") {
break
}
// Only a top-level host field establishes a network submission.
// Searching ahead would mistake H= inside message data for metadata.
if strings.HasPrefix(rest, "H=") {
_, clientEnd := HFieldClientIPAndEnd(rest[len("H="):])
if out.remote || clientEnd == 0 {
return submitFields{}, false
}
out.remote = true
rest = rest[len("H=")+clientEnd:]
continue
}
end := fieldEnd(rest)
field := rest[:end]
key, value, hasValue := strings.Cut(field, "=")
if hasValue && (key == "A" || key == "U" || key == "P") {
if seen[key] {
return submitFields{}, false
}
seen[key] = true
switch key {
case "A":
out.authSeen = true
authenticator, identity, ok := strings.Cut(value, ":")
if ok && (authenticator == "dovecot_login" || authenticator == "dovecot_plain") {
// smtp_mailauth can append an envelope identity after the
// authenticated identity. Only the second A= item is trusted.
auth, _, mailauth := strings.Cut(identity, ":")
out.auth = auth
if strings.Count(out.auth, "@") > 1 || strings.HasPrefix(out.auth, "@") ||
strings.HasSuffix(out.auth, "@") || strings.ContainsAny(out.auth, " \t\r\n\"\\") {
return submitFields{}, false
}
if mailauth {
// The optional envelope is client-supplied xtext,
// not additional submission metadata.
break metadata
}
}
case "U":
out.user = value
case "P":
out.protocol = value
}
}
rest = rest[end:]
}
// Exim appends remote RFC 1413 ident text without quoting spaces. It
// can therefore imitate all following metadata, including A= and S=.
// Neither an apparent authentication field nor its order proves an
// identity on such a record. Local U= still comes from the server.
if out.remote && seen["U"] {
return submitFields{}, false
}
return out, true
}
// An arrival must start with a date, time, optional zone/PID and message ID.
// A <= fragment in a delivery reply or subject cannot supply an identity.
func arrivalPrefix(prefix string) bool {
fields := strings.Fields(prefix)
if len(fields) < 3 {
return false
}
if _, err := time.Parse("2006-01-02 15:04:05", fields[0]+" "+fields[1]); err != nil {
return false
}
fields = fields[2:]
if _, err := time.Parse("-0700", fields[0]); err == nil {
fields = fields[1:]
}
if len(fields) > 0 && strings.HasPrefix(fields[0], "[") {
pid := fields[0]
if len(pid) < 3 || pid[len(pid)-1] != ']' || strings.Trim(pid[1:len(pid)-1], "0123456789") != "" {
return false
}
fields = fields[1:]
}
if len(fields) != 1 {
return false
}
id := strings.Split(fields[0], "-")
if len(id) != 3 {
return false
}
for _, part := range id {
if part == "" || strings.Trim(part, "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ") != "" {
return false
}
}
return true
}
func envelopeEnd(s string) int {
quoted := false
for i := 0; i < len(s); i++ {
switch {
case quoted && s[i] == '\\' && i+1 < len(s):
i++
case s[i] == '"':
quoted = !quoted
case !quoted && (s[i] == ' ' || s[i] == '\t' || s[i] == '\n'):
return i
}
}
return len(s)
}
// fieldEnd consumes quoted values together so their contents cannot
// masquerade as authentication or local-user metadata. Quotes can begin
// within a value, including an optional mailbox appended to an A= field.
func fieldEnd(s string) int {
eq := strings.IndexByte(s, '=')
var quote byte
for i := 0; i < len(s); i++ {
if quote != 0 {
switch {
case s[i] == '\\' && i+1 < len(s):
i++
case s[i] == quote:
quote = 0
}
continue
}
switch {
case s[i] == '"' || (s[i] == '\'' && eq >= 0 && i == eq+1):
quote = s[i]
case s[i] == ' ' || s[i] == '\t' || s[i] == '\r' || s[i] == '\n':
return i
}
}
return len(s)
}
package firewall
import (
"errors"
"time"
"github.com/pidginhost/csm/internal/actionlog"
)
// recordFirewallAction mirrors successful legacy entries. Paths with multiple
// outcomes record at their public operation boundary instead.
func recordFirewallAction(action, ip, reason, source string, duration time.Duration) {
recordFirewallResult(action, ip, reason, source, duration, actionlog.Applied, nil)
}
func firewallRecord(action, ip, reason, source string, duration time.Duration) actionlog.Record {
op, actor := firewallActionOp(action, source)
rec := actionlog.Record{Op: op, Action: action, Actor: actor, Target: ip, Reason: reason, Result: actionlog.Applied}
if duration > 0 {
rec.ActorDetail = "expires in " + duration.String()
}
// A block/allow reversal needs the prior state. Creating an allow is not
// an undo: it also suppresses future automatic blocks for this address.
return rec
}
func recordFirewallResult(action, ip, reason, source string, duration time.Duration, result actionlog.Result, err error) {
recordFirewallFindingResult(action, ip, reason, source, duration, result, err, "")
}
func recordFirewallFindingResult(action, ip, reason, source string, duration time.Duration, result actionlog.Result, err error, findingID string) {
rec := firewallRecord(action, ip, reason, source, duration)
rec.FindingID = findingID
rec.Result = result
if err != nil {
rec.Error = err.Error()
rec.Result = actionlog.Failed
if errors.Is(err, ErrIPProtected) {
rec.Result = actionlog.Refused
}
}
actionlog.Write(rec)
}
func recordFirewallFailure(action, ip, reason, source string, duration time.Duration, err error) {
if err != nil {
recordFirewallResult(action, ip, reason, source, duration, actionlog.Failed, err)
}
}
func recordBlockOutcome(ip, reason string, duration time.Duration, outcome BlockOutcome, err error, manual bool, findingID string) {
if !manual && outcome == BlockOutcomeNoop && err == nil {
return
}
rec := firewallRecord("block", ip, reason, InferProvenance("block", reason), duration)
rec.FindingID = findingID
rec.Op = "respond.block_ip"
if manual {
rec.Op = "operate.manual_firewall"
}
switch outcome {
case BlockOutcomeDryRun:
rec.Result = actionlog.DryRun
case BlockOutcomeAllowed, BlockOutcomeAllowlisted:
rec.Result = actionlog.Refused
}
if err != nil {
rec.Result = actionlog.Failed
rec.Error = err.Error()
if errors.Is(err, ErrIPProtected) {
rec.Result = actionlog.Refused
}
}
actionlog.Write(rec)
}
func firewallActionOp(action, source string) (string, actionlog.Actor) {
actor := actionlog.DefaultActor()
switch source {
case SourceCLI:
actor = actionlog.CLI
case SourceWebUI:
actor = actionlog.WebUI
}
switch action {
case "apply", "restart":
return "integrate.firewall_ruleset", actor
}
if source == SourceCLI || source == SourceWebUI || actor == actionlog.CLI {
return "operate.manual_firewall", actor
}
switch action {
case "flush", "unblock", "remove_allow", "allow_port", "remove_port_allow", "unblock_subnet":
return "operate.manual_firewall", actor
}
return "respond.block_ip", actor
}
//go:build linux
package firewall
import (
"fmt"
"time"
"github.com/google/nftables"
)
// Persist intent before touching the kernel so a restart converges to the
// requested policy. A failed atomic kernel batch restores the previous intent.
// The caller holds e.mu across both writes and the kernel transaction.
func (e *Engine) commitAllowedRemovals(prior, next FirewallState, removeIPs []string, req ActionRequest) error {
if e.lifecycle != nil {
return e.runDurableLocked(req, nil, next)
}
if err := e.persistFirewallIntent(prior, next); err != nil {
return fmt.Errorf("persisting allow removal: %w", err)
}
if err := e.deleteAllowedElements(removeIPs, next.Allowed); err != nil {
if restoreErr := e.saveState(&prior); restoreErr != nil {
return fmt.Errorf("partial failure: %w (state restore failed: %w)", err, restoreErr)
}
return err
}
return nil
}
func (e *Engine) deleteAllowedElements(ips []string, remaining []AllowedEntry) error {
if len(ips) == 0 {
return nil
}
conn := e.newMutationConn()
sets := make(map[*nftables.Set]bool)
for _, ip := range ips {
set, key, err := e.resolveIPSet(ip, e.setAllowed, e.setAllowed6)
if err != nil || set == nil {
continue
}
if err := conn.SetDeleteElements(set, []nftables.SetElement{{Key: key}}); err != nil {
return fmt.Errorf("removing allow for %s: %w", ip, err)
}
sets[set] = true
}
if err := conn.Flush(); err != nil {
if !isNftNotFound(err) {
return fmt.Errorf("removing allows: %w", err)
}
// An element already absent aborts the whole delete batch. Rebuild the
// affected sets atomically so one missing element cannot wedge cleanup.
return e.replaceAllowedSets(sets, remaining)
}
return nil
}
func (e *Engine) replaceAllowedSets(sets map[*nftables.Set]bool, entries []AllowedEntry) error {
conn := e.newMutationConn()
elements := make(map[*nftables.Set][]nftables.SetElement)
seen := make(map[string]bool)
now := time.Now()
for _, entry := range entries {
if !entry.ExpiresAt.IsZero() && !now.Before(entry.ExpiresAt) {
continue
}
ip, ok := canonicalIPKey(entry.IP)
if !ok || seen[ip] {
continue
}
set, key, err := e.resolveIPSet(ip, e.setAllowed, e.setAllowed6)
if err != nil || !sets[set] {
continue
}
seen[ip] = true
elements[set] = append(elements[set], nftables.SetElement{Key: key})
}
for set := range sets {
conn.FlushSet(set)
if err := addElementsChunked(conn, set, elements[set]); err != nil {
return err
}
}
if err := conn.Flush(); err != nil {
return fmt.Errorf("rebuilding allowed sets: %w", err)
}
return nil
}
// A separate non-lasting connection discards queued messages on any error.
// Its constructor cannot fail because it does not dial until Flush.
func (e *Engine) newMutationConn() *nftables.Conn {
conn, _ := nftables.New(nftables.WithSockOptions(applyNFTSocketBuffer), nftables.WithNetNSFd(e.conn.NetNS), nftables.WithTestDial(e.conn.TestDial))
return conn
}
package firewall
import (
"bufio"
"encoding/json"
"log"
"os"
"path/filepath"
"time"
)
const maxAuditFileSize = 10 * 1024 * 1024 // 10 MB
// AuditEntry records a firewall modification for compliance and forensics.
type AuditEntry struct {
Timestamp time.Time `json:"timestamp"`
Action string `json:"action"` // block, unblock, allow, remove_allow, flush, apply
IP string `json:"ip,omitempty"`
Reason string `json:"reason,omitempty"`
Source string `json:"source,omitempty"`
Duration string `json:"duration,omitempty"`
}
// AppendAudit writes an audit entry to the JSONL audit log.
// Rotates the log when it exceeds 10 MB.
func AppendAudit(statePath, action, ip, reason, source string, duration time.Duration) {
if source == "" {
source = InferProvenance(action, reason)
}
recordFirewallAction(action, ip, reason, source, duration)
appendAudit(statePath, action, ip, reason, source, duration)
}
func appendAudit(statePath, action, ip, reason, source string, duration time.Duration) {
if source == "" {
source = InferProvenance(action, reason)
}
entry := AuditEntry{
Timestamp: time.Now(),
Action: action,
IP: ip,
Reason: reason,
Source: source,
}
if duration > 0 {
entry.Duration = duration.String()
}
path := filepath.Join(statePath, "audit.jsonl")
data, err := json.Marshal(entry)
if err != nil {
return
}
data = append(data, '\n')
// Rotate if file exceeds max size
if info, statErr := os.Stat(path); statErr == nil && info.Size() > maxAuditFileSize {
_ = os.Rename(path, path+".1")
}
// #nosec G304 -- path is filepath.Join under operator-configured statePath.
f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600)
if err != nil {
// Without this log, perm/disk-full/inode-exhaust modes drop the
// audit entry with no operator-visible signal -- the Write/Close
// branches below already log, so a silent Open path was the only
// remaining hole in the audit pipeline.
log.Printf("firewall: audit open failed for %s: %v", path, err)
return
}
if _, writeErr := f.Write(data); writeErr != nil {
_ = f.Close()
log.Printf("firewall: audit write failed for %s: %v", path, writeErr)
return
}
// Close error on a writable file is the disk-full / fsync signal --
// without it, a dropped audit entry leaves no record anywhere.
if closeErr := f.Close(); closeErr != nil {
log.Printf("firewall: audit close failed for %s: %v", path, closeErr)
}
}
// ReadAuditLog returns the last N audit entries from the log.
func ReadAuditLog(statePath string, limit int) []AuditEntry {
path := filepath.Join(statePath, "firewall", "audit.jsonl")
// #nosec G304 -- filepath.Join under operator-configured statePath.
f, err := os.Open(path)
if err != nil {
return nil
}
defer f.Close()
var all []AuditEntry
scanner := bufio.NewScanner(f)
for scanner.Scan() {
var entry AuditEntry
if json.Unmarshal(scanner.Bytes(), &entry) == nil {
all = append(all, entry)
}
}
if limit > 0 && len(all) > limit {
all = all[len(all)-limit:]
}
return all
}
//go:build linux
package firewall
import "time"
// BlockIPForUndo captures both sides of a forced operator block under the
// mutation lock. A later CLI decision must never become part of this undo.
func (e *Engine) BlockIPForUndo(ip, reason string, ttl time.Duration) (before, after *BlockedEntry, resultErr error) {
defer func() {
if e.shouldLegacyOutcome(resultErr) {
recordBlockOutcome(ip, reason, ttl, BlockOutcomeLive, resultErr, true, "")
}
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return nil, nil, err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
state := e.loadStateFile()
if e.stateReadErr != nil {
return nil, nil, e.stateReadErr
}
if entry, exists := blockedStateEntry(state, ip); exists {
before = &entry
}
_, err = e.blockIPRequestLocked(ip, reason, ttl, false, false,
ActionRequest{Operation: "block", Target: ip, Reason: reason, TTL: ttl}, nil, true)
if err != nil {
return nil, nil, err
}
state = e.loadStateFile()
entry, _ := blockedStateEntry(state, ip)
return before, &entry, nil
}
// UnblockIPForUndo captures the removed block in the same critical section.
func (e *Engine) UnblockIPForUndo(ip string) (before *BlockedEntry, resultErr error) {
defer func() { e.legacyFirewallFailure("unblock", ip, "", "", 0, resultErr) }()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return nil, err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
state := e.loadStateFile()
if e.stateReadErr != nil {
return nil, e.stateReadErr
}
if entry, exists := blockedStateEntry(state, ip); exists {
before = &entry
}
if err := e.unblockIPLocked(ip); err != nil {
return nil, err
}
return before, nil
}
// RestoreBlockIfUnchanged applies a saved lifetime only while the action's
// firewall snapshot still matches. The comparison and mutation share e.mu.
func (e *Engine) RestoreBlockIfUnchanged(ip string, expected, prior *BlockedEntry) (resultErr error) {
var action, reason string
var duration time.Duration
defer func() {
switch action {
case "block":
if e.shouldLegacyOutcome(resultErr) {
recordBlockOutcome(ip, reason, duration, BlockOutcomeLive, resultErr, true, "")
}
case "unblock":
e.legacyFirewallFailure("unblock", ip, "", "", 0, resultErr)
}
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
state := e.loadStateFile()
if e.stateReadErr != nil {
return e.stateReadErr
}
current, exists := blockedStateEntry(state, ip)
if exists != (expected != nil) || (exists && !SameBlockedEntry(current, *expected)) {
return ErrBlockChanged
}
if !exists {
live, err := e.isBlockedLiveLocked(ip)
if err != nil {
return err
}
if live {
return ErrBlockChanged
}
}
if prior != nil {
ttl := time.Duration(0)
if !prior.ExpiresAt.IsZero() {
ttl = time.Until(prior.ExpiresAt)
if ttl <= 0 {
prior = nil
}
}
if prior != nil {
action, reason, duration = "block", prior.Reason, ttl
_, err := e.blockIPRequestLocked(ip, prior.Reason, ttl, false, false,
ActionRequest{Operation: "block", Target: ip, Reason: prior.Reason, TTL: ttl}, nil, false)
return err
}
}
if !exists {
return nil
}
action = "unblock"
return e.unblockIPLocked(ip)
}
package firewall
import "net"
// LiveBlockedSnapshot is a point-in-time membership view of the kernel's
// blocked IP sets, taken by dumping each configured family set once.
//
// HasV4 / HasV6 record which families the snapshot actually covers. A set can
// be absent -- firewall.ipv6 disabled leaves the v6 set nil -- and an
// uncovered family must never read as "not blocked", or a reconcile pass
// would prune every tracked block of that family on its first cycle.
type LiveBlockedSnapshot struct {
V4 map[string]struct{}
V6 map[string]struct{}
HasV4 bool
HasV6 bool
}
// Contains reports whether ip is in the live blocked set. known is false when
// the snapshot does not cover the IP's address family, in which case the
// caller must keep its cached answer rather than treat ip as unblocked.
//
// Family selection mirrors the per-IP path: an IPv4-mapped IPv6 address
// resolves against the v4 set, because that is the set the block path keys it
// into. A malformed IP is reported as definitively absent, matching
// IsBlockedLive.
func (s LiveBlockedSnapshot) Contains(ip string) (blocked, known bool) {
parsed := net.ParseIP(ip)
if parsed == nil {
return false, true
}
if v4 := parsed.To4(); v4 != nil {
if !s.HasV4 {
return false, false
}
_, ok := s.V4[v4.String()]
return ok, true
}
if !s.HasV6 {
return false, false
}
_, ok := s.V6[parsed.To16().String()]
return ok, true
}
//go:build linux
package firewall
import (
"fmt"
"net"
"os"
"github.com/google/nftables"
)
// UpdateCloudflareSet flushes and repopulates the Cloudflare nftables sets.
func (e *Engine) UpdateCloudflareSet(ipv4, ipv6 []string) error {
e.mu.Lock()
defer e.mu.Unlock()
if err := e.lifecycleReadyLocked(); err != nil {
return err
}
// Both sets are created together by createSets, but ConnectExisting
// loads them independently, so one can be present while the other is
// nil. Require both -- FlushSet/SetAddElements panic on a nil set.
if e.setCFWhitelist == nil || e.setCFWhitelist6 == nil {
return fmt.Errorf("cf_whitelist sets not initialized")
}
// Flush existing entries
e.conn.FlushSet(e.setCFWhitelist)
e.conn.FlushSet(e.setCFWhitelist6)
// Populate IPv4 CIDRs
var elems4 []nftables.SetElement
for _, cidr := range ipv4 {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
continue
}
start := network.IP.To4()
end := lastIPInRange(network)
if start == nil || end == nil {
continue
}
elems4 = appendIntervalSetElements(elems4, start, end)
}
if len(elems4) > 0 {
if err := e.conn.SetAddElements(e.setCFWhitelist, normalizeIntervalElements(elems4)); err != nil {
return fmt.Errorf("adding CF IPv4 elements: %w", err)
}
}
// Populate IPv6 CIDRs
var elems6 []nftables.SetElement
for _, cidr := range ipv6 {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
continue
}
if network.IP.To4() != nil {
continue
}
start := network.IP.To16()
end := lastIPInRange(network)
if start == nil || end == nil {
continue
}
elems6 = appendIntervalSetElements(elems6, start, end)
}
if len(elems6) > 0 {
if err := e.conn.SetAddElements(e.setCFWhitelist6, normalizeIntervalElements(elems6)); err != nil {
return fmt.Errorf("adding CF IPv6 elements: %w", err)
}
}
if err := e.conn.Flush(); err != nil {
return fmt.Errorf("flushing CF whitelist: %w", err)
}
fmt.Fprintf(os.Stderr, "firewall: cloudflare whitelist updated: %d IPv4, %d IPv6 CIDRs\n",
len(ipv4), len(ipv6))
return nil
}
// CloudflareIPs returns the currently configured Cloudflare CIDRs from the cached state.
func (e *Engine) CloudflareIPs() (ipv4, ipv6 []string) {
return LoadCFState(e.statePath)
}
// CloudflareCovers reports whether ip is inside a cached Cloudflare range,
// meaning a block of it leaves TCP 80/443 reachable (the CF accept precedes
// the blocked drop in the input chain).
func (e *Engine) CloudflareCovers(ip string) bool {
v4, v6 := LoadCFState(e.statePath)
return CloudflareRangesCover(v4, v6, ip)
}
package firewall
import (
"bufio"
"context"
"errors"
"fmt"
"net"
"net/http"
"os"
"strings"
"time"
"github.com/pidginhost/csm/internal/atomicio"
)
const (
cfIPv4URL = "https://www.cloudflare.com/ips-v4"
cfIPv6URL = "https://www.cloudflare.com/ips-v6"
// CloudflareCoverageWarning is shared by every operator-facing block
// response so API, CLI, and finding wording cannot drift apart.
CloudflareCoverageWarning = "IP is inside a Cloudflare allow range; ports 80/443 from it are still accepted"
)
// FetchCloudflareIPs downloads the current Cloudflare IP ranges, honouring
// ctx. A caller shutting down cancels it rather than waiting out the HTTP
// timeout: on a host that cannot reach cloudflare.com -- an egress-restricted
// server among them -- that wait is the full timeout, twice.
func FetchCloudflareIPs(ctx context.Context) (ipv4, ipv6 []string, err error) {
client := &http.Client{Timeout: 30 * time.Second}
return fetchCloudflareIPs(ctx, client)
}
func fetchCloudflareIPs(ctx context.Context, client *http.Client) (ipv4, ipv6 []string, err error) {
ipv4, ipv4Err := fetchCIDRList(ctx, client, cfIPv4URL)
if ipv4Err != nil {
ipv4Err = fmt.Errorf("fetching CF IPv4: %w", ipv4Err)
}
if ctx.Err() != nil {
// Cancelled between the two fetches: do not start the second.
return ipv4, nil, errors.Join(ipv4Err, ctx.Err())
}
ipv6, ipv6Err := fetchCIDRList(ctx, client, cfIPv6URL)
if ipv6Err != nil {
ipv6Err = fmt.Errorf("fetching CF IPv6: %w", ipv6Err)
}
return ipv4, ipv6, errors.Join(ipv4Err, ipv6Err)
}
// fetchCIDRList fetches a URL and parses one CIDR per line.
func fetchCIDRList(ctx context.Context, client *http.Client, url string) ([]string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
resp, err := client.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("HTTP %d from %s", resp.StatusCode, url)
}
scanner := bufio.NewScanner(resp.Body)
cidrs := parseCloudflareResponse(scanner)
if err := scanner.Err(); err != nil {
return nil, fmt.Errorf("reading %s: %w", url, err)
}
// A 200 with no ranges is an interstitial, a proxy page or a truncated
// body, never an empty Cloudflare list. Publishing it would flush every
// Cloudflare guard, so the previous list must stay in force.
if len(cidrs) == 0 {
return nil, fmt.Errorf("no CIDR ranges in response from %s", url)
}
return cidrs, nil
}
// parseCloudflareResponse parses lines from a scanner, returning valid CIDRs.
// Skips blank lines, comments, and invalid entries.
func parseCloudflareResponse(scanner *bufio.Scanner) []string {
var cidrs []string
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
_, _, err := net.ParseCIDR(line)
if err != nil {
continue
}
cidrs = append(cidrs, line)
}
return cidrs
}
// SaveCFState persists the Cloudflare CIDRs for status and coverage checks.
func SaveCFState(statePath string, ipv4, ipv6 []string, refreshed time.Time) error {
path := statePath
if !strings.HasSuffix(path, "/firewall") {
path += "/firewall"
}
if err := os.MkdirAll(path, 0700); err != nil {
return fmt.Errorf("creating Cloudflare state directory: %w", err)
}
var sb strings.Builder
fmt.Fprintf(&sb, "# refreshed: %s\n", refreshed.Format(time.RFC3339))
sb.WriteString("# ipv4\n")
for _, cidr := range ipv4 {
sb.WriteString(cidr)
sb.WriteByte('\n')
}
sb.WriteString("# ipv6\n")
for _, cidr := range ipv6 {
sb.WriteString(cidr)
sb.WriteByte('\n')
}
file := path + "/cf_whitelist.txt"
if err := atomicio.AtomicWrite(file, 0600, []byte(sb.String())); err != nil {
return fmt.Errorf("writing Cloudflare state: %w", err)
}
return nil
}
// CloudflareRangesCover reports whether ip falls inside any of the given
// Cloudflare CIDR ranges. The input chain accepts Cloudflare edges on TCP
// 80/443 before the blocked-IP drop, so blocking a covered IP does not stop
// its web traffic; callers use this to warn the operator at block time.
// Malformed CIDRs are skipped so one bad cache line cannot hide coverage.
func CloudflareRangesCover(v4, v6 []string, ip string) bool {
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
ranges := v4
if parsed.To4() == nil {
ranges = v6
}
for _, cidr := range ranges {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
continue
}
if network.Contains(parsed) {
return true
}
}
return false
}
// LoadCFState reads the cached Cloudflare CIDRs.
func LoadCFState(statePath string) (ipv4, ipv6 []string) {
path := statePath
if !strings.HasSuffix(path, "/firewall") {
path += "/firewall"
}
file := path + "/cf_whitelist.txt"
// #nosec G304 -- fixed filename under operator-configured statePath.
f, err := os.Open(file)
if err != nil {
return nil, nil
}
defer f.Close()
scanner := bufio.NewScanner(f)
section := "ipv4"
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "# ipv4" {
section = "ipv4"
continue
}
if line == "# ipv6" {
section = "ipv6"
continue
}
if line == "" || strings.HasPrefix(line, "#") {
continue
}
switch section {
case "ipv4":
ipv4 = append(ipv4, line)
case "ipv6":
ipv6 = append(ipv6, line)
}
}
return ipv4, ipv6
}
// LoadCFRefreshTime reads the last CF refresh time from state.
func LoadCFRefreshTime(statePath string) time.Time {
path := statePath
if !strings.HasSuffix(path, "/firewall") {
path += "/firewall"
}
file := path + "/cf_whitelist.txt"
// #nosec G304 -- fixed filename under operator-configured statePath.
f, err := os.Open(file)
if err != nil {
return time.Time{}
}
defer f.Close()
scanner := bufio.NewScanner(f)
if scanner.Scan() {
line := scanner.Text()
if strings.HasPrefix(line, "# refreshed: ") {
ts := strings.TrimPrefix(line, "# refreshed: ")
if t, err := time.Parse(time.RFC3339, ts); err == nil {
return t
}
}
}
return time.Time{}
}
package firewall
import (
"errors"
"fmt"
"net"
"strings"
)
// FirewallConfig defines the nftables firewall configuration.
type FirewallConfig struct {
Enabled bool `yaml:"enabled"`
// Open ports (IPv4)
TCPIn []int `yaml:"tcp_in"`
TCPOut []int `yaml:"tcp_out"`
UDPIn []int `yaml:"udp_in"`
UDPOut []int `yaml:"udp_out"`
// IPv6 - enable dual-stack filtering
IPv6 bool `yaml:"ipv6"`
TCP6In []int `yaml:"tcp6_in"` // if empty, uses tcp_in
TCP6Out []int `yaml:"tcp6_out"` // if empty, uses tcp_out
UDP6In []int `yaml:"udp6_in"` // if empty, uses udp_in
UDP6Out []int `yaml:"udp6_out"` // if empty, uses udp_out
// Ports restricted to infra IPs only
RestrictedTCP []int `yaml:"restricted_tcp"`
// RequiredTCPOut declares outbound TCP ports a service on this host needs.
// It is checked, never merged: validation warns when an effective outbound
// family policy omits one, so an integration can state its
// requirement in its own conf.d fragment and have `csm doctor` verify it
// against the effective policy.
RequiredTCPOut []int `yaml:"required_tcp_out"`
// TCPOutAllow permits outbound TCP to a specific destination on a port
// range. tcp_out is []int and cannot express a range, so a host acting as
// a client for a range-using protocol (passive FTP data channels) has no
// way to state its need without opening every high port to the internet.
// Empty means no destination-scoped egress, which is the prior behaviour.
TCPOutAllow []OutAllowRule `yaml:"tcp_out_allow"`
// Passive FTP range
PassiveFTPStart int `yaml:"passive_ftp_start"`
PassiveFTPEnd int `yaml:"passive_ftp_end"`
// Infra IPs (CIDR notation)
InfraIPs []string `yaml:"infra_ips"`
// Rate limiting (per-source nftables meters). SYN/conn-rate/UDP are
// dual-stack (IPv6 keyed per /64); ConnLimit is IPv4-only (per-source
// ct count uses nf_conncount, whose GC has a kernel UAF on the el8 kernel).
ConnRateLimit int `yaml:"conn_rate_limit"` // new connections per minute per source (IPv6 per /64)
SYNFloodProtection bool `yaml:"syn_flood_protection"`
ConnLimit int `yaml:"conn_limit"` // max concurrent connections per IPv4 source, IPv4 only (0 = disabled)
// Per-port flood protection - per-source rate limit per port and IP family.
PortFlood []PortFloodRule `yaml:"port_flood"`
// UDP flood protection - per-source rate limit on UDP packets (IPv6 per /64)
UDPFlood bool `yaml:"udp_flood"`
UDPFloodRate int `yaml:"udp_flood_rate"` // packets per second
UDPFloodBurst int `yaml:"udp_flood_burst"` // burst allowance
// Country blocking
CountryBlock []string `yaml:"country_block"` // ISO country codes
CountryDBPath string `yaml:"country_db_path"`
// Ports to drop silently without logging (reduces log noise from scanners)
DropNoLog []int `yaml:"drop_nolog"`
// Max blocked IPs (prevents memory exhaustion, 0 = unlimited)
DenyIPLimit int `yaml:"deny_ip_limit"`
DenyTempIPLimit int `yaml:"deny_temp_ip_limit"`
// Outbound SMTP restriction - block outgoing mail except from allowed users
SMTPBlock bool `yaml:"smtp_block"`
SMTPAllowUsers []string `yaml:"smtp_allow_users"` // usernames allowed to send
SMTPPorts []int `yaml:"smtp_ports"`
// Dynamic DNS - resolve hostnames to IPs, update allowed set periodically
DynDNSHosts []string `yaml:"dyndns_hosts"`
// Logging
LogDropped bool `yaml:"log_dropped"`
LogRate int `yaml:"log_rate"` // log entries per minute
// DoS-exempt ranges - source CIDRs excluded from connection/mail-port
// meters and auto-subnet escalation. Intended for shared-source ranges
// such as carrier CGNAT blocks and well-known mail-provider egress.
DOSExemptRanges []string `yaml:"dos_exempt_ranges"`
DOSExemptKnownMailProviders *bool `yaml:"dos_exempt_known_mail_providers"`
}
// MergeInfraIPs returns the effective firewall infra list in configuration
// order. Equivalent IP spellings share one entry because nftables enforces the
// parsed address or network, not the operator's original text.
func MergeInfraIPs(topLevel, firewallSpecific []string) []string {
const infraIPHintCap = 1 << 16
hint := len(topLevel) + len(firewallSpecific)
if hint < 0 || hint > infraIPHintCap {
hint = infraIPHintCap
}
seen := make(map[string]struct{}, hint)
merged := make([]string, 0, hint)
appendEntries := func(entries []string) {
for _, raw := range entries {
entry := strings.TrimSpace(raw)
if entry == "" {
continue
}
key := canonicalInfraIPKey(entry)
if _, ok := seen[key]; ok {
continue
}
seen[key] = struct{}{}
merged = append(merged, entry)
}
}
appendEntries(topLevel)
appendEntries(firewallSpecific)
if len(merged) == 0 {
return nil
}
return merged
}
func canonicalInfraIPKey(entry string) string {
if _, network, err := net.ParseCIDR(entry); err == nil {
return "network:" + network.String()
}
if ip := net.ParseIP(entry); ip != nil {
bits := net.IPv6len * 8
if ip.To4() != nil {
bits = net.IPv4len * 8
}
return "network:" + (&net.IPNet{IP: ip, Mask: net.CIDRMask(bits, bits)}).String()
}
return "host:" + strings.ToLower(strings.TrimSuffix(entry, "."))
}
func effectiveIPv6Ports(ipv4, ipv6 []int) []int {
if len(ipv6) > 0 {
return ipv6
}
return ipv4
}
// restrictedOutputNeedsIPv4Bypass reports whether the restricted output chain
// would drop all IPv4 egress. effectiveIPv6Ports mirrors IPv4 out-ports onto
// IPv6 but never the reverse, so configuring only IPv6 egress ports puts the
// shared inet chain into DROP mode with no IPv4 accept rules. Callers invoke
// this only on the restricted path and, when true, must accept IPv4 wholesale
// after the SMTP guard so smtp_block still applies to IPv4.
func (c *FirewallConfig) restrictedOutputNeedsIPv4Bypass() bool {
return len(c.TCPOut) == 0 && len(c.UDPOut) == 0
}
// ExemptKnownMailProviders reports whether bundled mail-provider ranges are
// included in the DoS-exempt set. Defaults to true when unset.
func (c *FirewallConfig) ExemptKnownMailProviders() bool {
if c.DOSExemptKnownMailProviders == nil {
return true
}
return *c.DOSExemptKnownMailProviders
}
// ParseOutAllowDst accepts a bare IP or a CIDR and returns it as a network.
// A bare address becomes a host route; mapped IPv4 prefixes use IPv4 widths.
func ParseOutAllowDst(dst string) (*net.IPNet, error) {
if dst == "" {
return nil, errors.New("dst is required")
}
if _, network, err := net.ParseCIDR(dst); err == nil {
if ip4 := network.IP.To4(); ip4 != nil {
// ParseCIDR retains a 128-bit mask for mapped IPv4 prefixes.
// The engine loads only four bytes, so normalize both together.
network.IP = ip4
network.Mask = network.Mask[len(network.Mask)-net.IPv4len:]
}
return network, nil
}
ip := net.ParseIP(dst)
if ip == nil {
return nil, fmt.Errorf("dst %q is neither an IP address nor a CIDR", dst)
}
if ip4 := ip.To4(); ip4 != nil {
return &net.IPNet{IP: ip4, Mask: net.CIDRMask(32, 32)}, nil
}
return &net.IPNet{IP: ip.To16(), Mask: net.CIDRMask(128, 128)}, nil
}
// OutAllowRule permits outbound TCP to one destination on a port range.
// Dst is an IP or CIDR; 0.0.0.0/0 and ::/0 mean any destination in that family.
type OutAllowRule struct {
Dst string `yaml:"dst"`
PortStart int `yaml:"port_start"`
PortEnd int `yaml:"port_end"`
}
// PortFloodRule defines per-port connection rate limiting.
type PortFloodRule struct {
Port int `yaml:"port"`
Proto string `yaml:"proto"` // "tcp" or "udp"
Hits int `yaml:"hits"` // max new connections
Seconds int `yaml:"seconds"` // time window in seconds
}
// DefaultConfig returns a sensible default firewall configuration
// matching a typical cPanel server.
func DefaultConfig() *FirewallConfig {
return &FirewallConfig{
Enabled: false,
// SSH (22) is intentionally absent. Many cPanel hosts move sshd to
// 2087 or other alt ports; operators who keep sshd on 22 must add
// it explicitly. TCP 853 enables DNS-over-TLS; UDP 853 enables
// DNS-over-QUIC.
TCPIn: []int{
20, 21, 25, 26, 53, 80, 110, 143, 443, 465, 587,
853, 993, 995, 2077, 2078, 2079, 2080, 2082, 2083,
2091, 2095, 2096,
},
TCPOut: []int{
20, 21, 25, 26, 37, 43, 53, 80, 110, 113, 443,
465, 587, 853, 873, 993, 995, 2082, 2083, 2086, 2087,
2089, 2195, 2325, 2703,
},
UDPIn: []int{53, 443, 853},
// 6277/24441 are DCC/Pyzor network checks used by SpamAssassin.
// Without them outbound spam-scoring queries silently fail.
UDPOut: []int{53, 113, 123, 443, 853, 873, 6277, 24441},
RestrictedTCP: []int{2086, 2087, 2325, 9443},
TCPOutAllow: nil,
PassiveFTPStart: 49152,
PassiveFTPEnd: 65534,
// 200 new connections per minute per source (IPv6 aggregated per /64).
// Sized to tolerate shared CGNAT egress (residential ISPs / mobile
// carriers) where hundreds of subscribers share a single public
// address; the original 30/min default produced spurious drops on
// such ranges.
ConnRateLimit: 200,
SYNFloodProtection: true,
// 400 concurrent connections per IPv4 source. Sized for power users
// (multi-tab webmail + IMAP IDLE on multiple devices + Thunderbird
// parallel send + HTTPS browsing) plus headroom for shared CGNAT
// egress IPs. Operators with very heavy IDLE-style workloads can
// raise it further.
ConnLimit: 400,
// 600 hits / 300 s = 120 new connections per minute per source IP.
// Sized to tolerate normal MUA bursts (Thunderbird/iPhone/Outlook each
// open 5-15 parallel sessions when sending one email or syncing IMAP
// after suspend) while still catching true single-IP floods. Detection
// of low-and-slow scanners belongs to userspace, not this rule.
PortFlood: []PortFloodRule{
{Port: 25, Proto: "tcp", Hits: 600, Seconds: 300},
{Port: 465, Proto: "tcp", Hits: 600, Seconds: 300},
{Port: 587, Proto: "tcp", Hits: 600, Seconds: 300},
},
UDPFlood: true,
UDPFloodRate: 100,
UDPFloodBurst: 500,
DropNoLog: []int{23, 67, 68, 111, 113, 135, 136, 137, 138, 139, 445, 500, 513, 520},
DenyIPLimit: 3000,
DenyTempIPLimit: 500,
SMTPBlock: false,
SMTPPorts: []int{25, 465, 587},
LogDropped: true,
LogRate: 5,
}
}
package firewall
import "path/filepath"
// CountryDBDir returns the directory holding the per-country CIDR files:
// country_db_path when set, otherwise the geoip directory under the state
// path, which is where `csm firewall update-geoip` writes. Every consumer
// (engine, CLI update and lookup) resolves through here so an empty setting
// cannot leave country blocking armed on one surface and inert on another.
func CountryDBDir(cfg *FirewallConfig, statePath string) string {
if cfg != nil && cfg.CountryDBPath != "" {
return cfg.CountryDBPath
}
return filepath.Join(statePath, "geoip")
}
// countryBlockingActive reports whether the operator asked for country
// blocking. The DB directory always resolves, so the code list alone decides.
func countryBlockingActive(cfg *FirewallConfig) bool {
return cfg != nil && len(cfg.CountryBlock) > 0
}
package firewall
import (
"context"
"fmt"
"net"
"os"
"sort"
"sync"
"time"
)
// hostHealth tracks per-host resolution state for the re-resolve guard.
type hostHealth struct {
lastSuccess time.Time
firstFailure time.Time
findingEmitted bool
}
// DynDNSResolver periodically resolves hostnames and updates the firewall allowed set.
type DynDNSResolver struct {
mu sync.Mutex
hosts []string
resolved map[string][]string // hostname -> all resolved IPs
// infraHosts is the subset of hosts that also feed the engine's
// infra-IP guard. When a host appears here, every successful
// resolution additionally calls engine.UpdateInfraResolved so the
// hostname stays blockable-refusable even when the IP behind it
// rotates. Hosts not in this map only feed the allowed-IPs set.
infraHosts map[string]struct{}
engine interface {
AllowIP(ip string, reason string) error
RemoveAllowIPBySource(ip string, source string) error
}
// infraEngine receives infra-mode updates. Optional - separate
// interface so the AllowIP/RemoveAllowIPBySource consumers don't
// have to grow when an operator does not declare any infra hosts.
infraEngine interface {
UpdateInfraResolved(host string, ips []string)
DropInfraResolved(host string)
}
// lookupFn is the context-bound DNS lookup function. Defaults to
// net.DefaultResolver.LookupHost so callers can bound a stuck
// resolver by cancelling the parent context. Tests may replace this
// with a stub.
lookupFn func(ctx context.Context, host string) ([]string, error)
// Guard fields for the DNS re-resolve guard (Task 3).
muGuard sync.RWMutex
hostHealth map[string]*hostHealth
gracePeriod time.Duration
unresolvable map[string]struct{}
findingSink func(name string)
}
// NewDynDNSResolver creates a resolver for the given hostnames.
func NewDynDNSResolver(hosts []string, engine interface {
AllowIP(ip string, reason string) error
RemoveAllowIPBySource(ip string, source string) error
}) *DynDNSResolver {
r := &DynDNSResolver{
hosts: hosts,
resolved: make(map[string][]string),
engine: engine,
hostHealth: make(map[string]*hostHealth),
unresolvable: make(map[string]struct{}),
gracePeriod: 10 * time.Minute,
}
r.lookupFn = net.DefaultResolver.LookupHost
return r
}
// AddHost appends a hostname to the resolver's host list.
// It is safe to call concurrently. Used in tests to add hosts after construction.
func (d *DynDNSResolver) AddHost(host string) {
d.mu.Lock()
defer d.mu.Unlock()
d.hosts = append(d.hosts, host)
}
// RegisterInfraHost marks host as an infra hostname. Every subsequent
// successful resolution will, in addition to AllowIP, call
// engine.UpdateInfraResolved so the resolved IPs feed the infra-block
// guard. Idempotent. Wire an infra engine via SetInfraEngine before
// the resolver's first tick to make this effective.
func (d *DynDNSResolver) RegisterInfraHost(host string) {
d.mu.Lock()
defer d.mu.Unlock()
if d.infraHosts == nil {
d.infraHosts = make(map[string]struct{})
}
d.infraHosts[host] = struct{}{}
}
// SetInfraEngine wires the engine that receives infra-mode resolution
// updates. Setting it to nil disables infra routing without affecting
// the regular allowed-IPs path.
func (d *DynDNSResolver) SetInfraEngine(eng interface {
UpdateInfraResolved(host string, ips []string)
DropInfraResolved(host string)
}) {
d.mu.Lock()
defer d.mu.Unlock()
d.infraEngine = eng
}
// markLastSuccess seeds the lastSuccess timestamp for a host to now.
// Used in tests to simulate a prior successful resolution without a real
// DNS lookup.
func (d *DynDNSResolver) markLastSuccess(host string) {
d.muGuard.Lock()
defer d.muGuard.Unlock()
hh := d.hostHealth[host]
if hh == nil {
hh = &hostHealth{}
d.hostHealth[host] = hh
}
hh.lastSuccess = time.Now()
}
// SetFindingSink installs the callback invoked when a host has been
// unresolvable for longer than gracePeriod. Called from the daemon at
// startup, after the alert pipeline is wired.
func (d *DynDNSResolver) SetFindingSink(sink func(host string)) {
d.muGuard.Lock()
defer d.muGuard.Unlock()
d.findingSink = sink
}
// UnresolvableHosts lists infra_ips hostnames currently failing to resolve
// beyond the grace period.
func (d *DynDNSResolver) UnresolvableHosts() []string {
d.muGuard.RLock()
defer d.muGuard.RUnlock()
out := make([]string, 0, len(d.unresolvable))
for h := range d.unresolvable {
out = append(out, h)
}
sort.Strings(out)
return out
}
// Run starts the periodic resolver. Blocks until stopCh is closed.
func (d *DynDNSResolver) Run(stopCh <-chan struct{}) {
parent, cancelParent := context.WithCancel(context.Background())
defer cancelParent()
done := make(chan struct{})
go func() {
select {
case <-stopCh:
cancelParent()
case <-done:
}
}()
defer close(done)
select {
case <-stopCh:
cancelParent()
return
default:
}
d.runTick(parent)
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for {
select {
case <-parent.Done():
return
case <-ticker.C:
d.runTick(parent)
}
}
}
// runTick is the periodic Run helper. It bounds a single resolution
// cycle so a stuck DNS server cannot hold the ticker beyond its budget
// and stack late ticks. The per-tick budget is purposely larger than
// the default 30-second per-resolver timeout but smaller than the
// 5-minute ticker period.
func (d *DynDNSResolver) runTick(parent context.Context) {
ctx, cancel := context.WithTimeout(parent, dyndnsTickBudget)
defer cancel()
d.tickOnce(ctx)
}
const dyndnsTickBudget = 60 * time.Second
// tickOnce performs one resolution cycle over all registered hosts.
// The periodic Run loop calls this; tests can call it directly. The
// context bounds the whole cycle so a single stuck DNS server cannot
// hold the tick for the resolver's implicit timeout (~30s) and pile
// late ticks on top of the configured 5-minute period.
func (d *DynDNSResolver) tickOnce(ctx context.Context) {
d.mu.Lock()
hosts := make([]string, len(d.hosts))
copy(hosts, d.hosts)
d.mu.Unlock()
for _, host := range hosts {
if ctx.Err() != nil {
return
}
d.resolveHost(ctx, host)
}
}
// resolveAll is retained for backward compatibility with existing tests.
func (d *DynDNSResolver) resolveAll() {
d.tickOnce(context.Background())
}
func (d *DynDNSResolver) resolveHost(ctx context.Context, host string) {
newIPs, err := d.lookupFn(ctx, host)
if err != nil || len(newIPs) == 0 {
if ctx.Err() == context.Canceled {
return
}
fmt.Fprintf(os.Stderr, "dyndns: failed to resolve %s: %v\n", host, err)
d.updateGuardFailure(host)
return
}
sort.Strings(newIPs)
d.mu.Lock()
oldIPs := d.resolved[host]
d.mu.Unlock()
oldSet := make(map[string]bool)
for _, ip := range oldIPs {
oldSet[ip] = true
}
newSet := make(map[string]bool)
var infraIPs []string
for _, ip := range newIPs {
newSet[ip] = true
if parsed := net.ParseIP(ip); parsed != nil {
infraIPs = append(infraIPs, parsed.String())
}
}
var successIPs []string
// Keep failed removals tracked so the next DNS refresh retries them.
for _, ip := range oldIPs {
if !newSet[ip] {
if err := d.engine.RemoveAllowIPBySource(ip, SourceDynDNS); err != nil {
successIPs = append(successIPs, ip)
fmt.Fprintf(os.Stderr, "dyndns: error removing %s (%s): %v\n", ip, host, err)
continue
}
fmt.Fprintf(os.Stderr, "dyndns: %s removed %s (no longer resolves)\n", host, ip)
}
}
// Add new IPs
reason := fmt.Sprintf("dyndns: %s", host)
for _, ip := range newIPs {
if oldSet[ip] {
successIPs = append(successIPs, ip) // already allowed
continue
}
if err := d.engine.AllowIP(ip, reason); err != nil {
fmt.Fprintf(os.Stderr, "dyndns: error allowing %s (%s): %v\n", ip, host, err)
continue
}
successIPs = append(successIPs, ip)
fmt.Fprintf(os.Stderr, "dyndns: %s resolved to %s (added)\n", host, ip)
}
// Retain active allows and old addresses whose removal still needs retrying.
d.mu.Lock()
d.resolved[host] = successIPs
isInfra := false
if _, ok := d.infraHosts[host]; ok {
isInfra = true
}
infraEngine := d.infraEngine
d.mu.Unlock()
// Infra mode feeds the block guard from DNS itself, not from the
// allow-list mutation result. A transient nftables write failure
// should not make a resolved management hostname blockable.
if isInfra && infraEngine != nil {
infraEngine.UpdateInfraResolved(host, infraIPs)
}
// Successful resolution: clear any guard state.
d.updateGuardSuccess(host)
}
// updateGuardSuccess marks a host as successfully resolved in the guard state.
func (d *DynDNSResolver) updateGuardSuccess(host string) {
d.muGuard.Lock()
hh := d.hostHealth[host]
if hh == nil {
hh = &hostHealth{}
d.hostHealth[host] = hh
}
hh.lastSuccess = time.Now()
hh.firstFailure = time.Time{}
if _, was := d.unresolvable[host]; was {
delete(d.unresolvable, host)
fmt.Fprintf(os.Stderr, "dyndns: %s recovered (resolution succeeded)\n", host)
}
hh.findingEmitted = false
d.muGuard.Unlock()
}
// updateGuardFailure updates guard state after a failed resolution for a host.
// If the host has been unresolvable beyond gracePeriod and no finding has been
// emitted yet, it marks the host unresolvable and invokes the finding sink.
func (d *DynDNSResolver) updateGuardFailure(host string) {
d.muGuard.Lock()
hh := d.hostHealth[host]
if hh == nil {
hh = &hostHealth{}
d.hostHealth[host] = hh
}
now := time.Now()
if hh.firstFailure.IsZero() {
hh.firstFailure = now
}
since := hh.lastSuccess
if since.IsZero() {
since = hh.firstFailure
}
var sinkToCall func(string)
if now.Sub(since) > d.gracePeriod &&
!hh.findingEmitted {
d.unresolvable[host] = struct{}{}
hh.findingEmitted = true
sinkToCall = d.findingSink
}
d.muGuard.Unlock()
// Invoke sink outside the lock to prevent deadlock if sink calls back in.
if sinkToCall != nil {
sinkToCall(host)
}
}
package firewall
import "net"
// EffectiveDOSExemptNets returns the union of operator ranges and (when enabled)
// provider ranges, split into IPv4 and IPv6 *net.IPNet slices. providerNets may
// be nil. Invalid operator entries are skipped (validation already rejected them
// at load; this is defense in depth).
func EffectiveDOSExemptNets(cfg *FirewallConfig, providerNets []*net.IPNet) (v4, v6 []*net.IPNet) {
var all []*net.IPNet
if cfg != nil {
for _, s := range cfg.DOSExemptRanges {
n := parseExemptNet(s)
if n != nil {
all = append(all, n)
}
}
}
// nil cfg treats ExemptKnownMailProviders as true (safe default).
if cfg == nil || cfg.ExemptKnownMailProviders() {
all = append(all, providerNets...)
}
// Split and copy. To4() returns non-nil for IPv4, including
// IPv4-mapped-in-IPv6 addresses, so it correctly classifies both.
// Copies ensure callers cannot mutate the returned slices back into
// any shared state held by the caller or this package.
for _, n := range all {
if n == nil {
continue
}
ip := make(net.IP, len(n.IP))
copy(ip, n.IP)
mask := make(net.IPMask, len(n.Mask))
copy(mask, n.Mask)
c := &net.IPNet{IP: ip, Mask: mask}
if n.IP.To4() != nil {
v4 = append(v4, c)
} else {
v6 = append(v6, c)
}
}
return v4, v6
}
// parseExemptNet parses s as a CIDR or a bare IP (becomes /32 or /128).
// Returns nil for invalid input.
func parseExemptNet(s string) *net.IPNet {
if _, n, err := net.ParseCIDR(s); err == nil {
return n
}
ip := net.ParseIP(s)
if ip == nil {
return nil
}
if ip4 := ip.To4(); ip4 != nil {
return &net.IPNet{IP: ip4, Mask: net.CIDRMask(32, 32)}
}
return &net.IPNet{IP: ip.To16(), Mask: net.CIDRMask(128, 128)}
}
//go:build linux
package firewall
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"strings"
"sync"
"syscall"
"time"
"github.com/google/nftables"
"github.com/google/nftables/binaryutil"
"github.com/google/nftables/expr"
"github.com/mdlayher/netlink"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/atomicio"
)
// Engine manages the nftables firewall ruleset.
// Manages the nftables ruleset via netlink.
type Engine struct {
lifecycle *Lifecycle
stateRevision uint64
stateReadErr error
mu sync.Mutex
conn *nftables.Conn
// listTables, when set, replaces conn.ListTables (tests inject a
// failing listing).
listTables func() ([]*nftables.Table, error)
cfg *FirewallConfig
// Captured only after our own successful ruleset transaction. Config-file
// hashes cannot authenticate the rules currently installed in the kernel.
appliedRuleset string
readRuleset func() (string, error)
// dryRunRecorder is called by BlockIP when auto_response.dry_run is
// active. Set by SetDryRunRecorder after construction so the firewall
// package does not import internal/store (which would be a cycle).
dryRunRecorder func(ip, reason string, timeout time.Duration)
// dryRunEnabled reports whether auto_response.dry_run is active.
// Set by the daemon so this package does not import internal/config.
dryRunEnabled func() bool
// verdictAsker, when set, is consulted after local block validation and
// before the dry-run gate. The callback returns (verdict, tenantID, note,
// error). Verdict "allow" short-circuits the block; "block" or empty
// proceeds with the default flow.
// Errors are fail-open: the daemon proceeds with the default block and
// logs the failure. Daemon owns the underlying verdict.Client; this
// package stays free of the internal/verdict import.
verdictAsker func(ctx context.Context, ip, reason string) (verdict, tenantID, note string, err error)
// softAllowFn, when set, reports whether an IP belongs to a verified-bot
// range (built-in/auto-updated crawler snapshots plus operator
// reputation.verified_bots IP-range entries). The auto-block path consults
// it so a published crawler the auto-blocker would otherwise re-add to
// blocked_ips is left out. Operator `firewall deny` (BlockIPForce) bypasses
// this gate. Set by the daemon so this package does not import
// internal/threatintel.
softAllowFn func(ip string) bool
// shutdownCtx, when set, scopes the lifetime of any in-flight verdict
// callback to daemon shutdown. Without it, BlockIPOutcome used
// context.Background() and a wedged panel callback kept the
// auto-block caller waiting for the full http.Client.Timeout during
// graceful restart. Nil falls back to context.Background() so unit
// tests that build the Engine literal without a daemon keep working.
shutdownCtx context.Context
table *nftables.Table
chainIn *nftables.Chain
chainOut *nftables.Chain
setBlocked *nftables.Set
setBlockedNet *nftables.Set
setAllowed *nftables.Set
setInfra *nftables.Set
setCountry *nftables.Set
// Cloudflare IP whitelist sets (interval for CIDR matching)
setCFWhitelist *nftables.Set // IPv4
setCFWhitelist6 *nftables.Set // IPv6
// DoS-exempt ranges (interval sets for CIDR matching).
// Sources in these sets bypass per-IP rate-limit / conn-limit / port-flood
// rules. setDOSExempt6 is nil when IPv6 is disabled.
setDOSExempt *nftables.Set
setDOSExempt6 *nftables.Set
// dosExemptProviderNets holds the mail-provider overlay pushed by the
// daemon. Written by SetDOSExemptProviderNets and RefreshDOSExemptSets;
// read by createSets() and RefreshDOSExemptSets. Always guarded by mu.
dosExemptProviderNets []*net.IPNet
// IPv6 sets (nil if IPv6 disabled)
setBlocked6 *nftables.Set
setBlockedNet6 *nftables.Set
setAllowed6 *nftables.Set
setInfra6 *nftables.Set
setCountry6 *nftables.Set
// Meters for per-IP rate limiting
meterSYN *nftables.Set
meterConn *nftables.Set
meterUDP *nftables.Set
meterConnlim *nftables.Set
// IPv6 siblings of the rate meters (keyed on the /64 prefix). There is
// deliberately no v6 connlimit: per-source ct count uses nf_conncount,
// whose GC has a kernel use-after-free on the el8 kernel, so concurrent
// connection limiting stays IPv4-only.
meterSYN6 *nftables.Set
meterConn6 *nftables.Set
meterUDP6 *nftables.Set
meterPortFlood4 map[string]*nftables.Set
meterPortFlood6 map[string]*nftables.Set
statePath string
// Cached parsed firewall state. Population is lazy: the first call
// to loadStateFile under e.mu reads + parses state.json once, then
// subsequent calls return a deep copy from the in-memory cache as
// long as the on-disk metadata key is unchanged.
//
// Before this cache existed, every single mutator (BlockIP,
// AllowIP, saveBlockedEntry, etc.) reloaded and re-parsed the full
// 325 KiB state.json from disk, and every IsBlocked / IsAllowed
// call did a linear scan over the parsed slices. On a busy
// production host that meant ~72 state.json opens per second
// steady state, which showed up as the dominant CPU hot spot in
// roadmap audit 7.1.
//
// All four fields are written only while e.mu is held. The index
// maps are rebuilt every time stateCache is repopulated so
// O(1) lookups stay coherent with the cached slices. blockedIPIndex
// stores the slice position for each canonical IP.
stateCache *FirewallState
stateCacheKey stateFileCacheKey
blockedIPIndex map[string]int
allowedIPIndex map[string]struct{}
portAllowedIndex map[string]struct{}
blockedCIDRIndex map[string]struct{}
liveBlockLookup func(set *nftables.Set, key []byte) (bool, error)
// liveBlockedDump, when non-nil, replaces the netlink dump of one
// blocked set. Tests inject it to avoid a real nft connection.
liveBlockedDump func(set *nftables.Set) ([]nftables.SetElement, error)
// liveBlockCounts, when non-nil, returns the live (perm, temp)
// element counts across blocked v4 + v6 sets. Tests inject this
// to avoid spinning up a real nft connection. Nil falls back to
// GetSetElements queries.
liveBlockCounts func() (perm, temp int, err error)
// infraResolved maps a hostname declared under cfg.InfraIPs to its
// last successfully-resolved set of IPs. blockIPTarget refuses to
// block any of these so a transient DNS pause cannot let an
// attacker block CSM's own panel hostname into a lockout. Mutated
// under e.mu by UpdateInfraResolved / DropInfraResolved from the
// DynDNS resolver.
infraResolved map[string]map[string]struct{}
// localAddrs caches the host's own non-loopback interface addresses.
// The block guard refuses to block any of these regardless of
// cfg.InfraIPs so a misconfigured config or a stray scan that loops
// back to the daemon cannot brick the host. Refreshed lazily under
// e.mu when localAddrsExpiresAt has elapsed.
localAddrs map[string]struct{}
localAddrsExpiresAt time.Time
// localAddrsLookup, when non-nil, replaces net.InterfaceAddrs() for
// tests. Returning an error leaves the cache untouched.
localAddrsLookup func() ([]string, error)
}
type stateFileCacheKey struct {
modTime time.Time
changeTime time.Time
size int64
dev uint64
ino uint64
}
func stateFileCacheKeyFromInfo(info os.FileInfo) stateFileCacheKey {
key := stateFileCacheKey{
modTime: info.ModTime(),
size: info.Size(),
}
if stat, ok := info.Sys().(*syscall.Stat_t); ok {
key.changeTime = time.Unix(stat.Ctim.Sec, stat.Ctim.Nsec)
key.dev = uint64(stat.Dev)
key.ino = uint64(stat.Ino)
}
return key
}
func (k stateFileCacheKey) matches(info os.FileInfo) bool {
other := stateFileCacheKeyFromInfo(info)
return k.size == other.size &&
k.dev == other.dev &&
k.ino == other.ino &&
k.modTime.Equal(other.modTime) &&
k.changeTime.Equal(other.changeTime)
}
// BlockedEntry represents a blocked IP with metadata.
type BlockedEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
BlockedAt time.Time `json:"blocked_at"`
ExpiresAt time.Time `json:"expires_at"` // zero = permanent
}
// AllowedEntry represents an allowed IP with metadata.
type AllowedEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
Port int `json:"port,omitempty"` // 0 = all ports
ExpiresAt time.Time `json:"expires_at,omitempty"` // zero = permanent
}
// SubnetEntry represents a blocked CIDR range.
type SubnetEntry struct {
CIDR string `json:"cidr"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
BlockedAt time.Time `json:"blocked_at"`
ExpiresAt time.Time `json:"expires_at,omitempty"`
}
// PortAllowEntry represents a port-specific IP allow (e.g. tcp|in|d=PORT|s=IP).
type PortAllowEntry struct {
IP string `json:"ip"`
Port int `json:"port"`
Proto string `json:"proto"` // "tcp" or "udp"
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
}
// FirewallState is persisted to disk for restore on restart.
type FirewallState struct {
Blocked []BlockedEntry `json:"blocked"`
BlockedNet []SubnetEntry `json:"blocked_nets"`
Allowed []AllowedEntry `json:"allowed"`
PortAllowed []PortAllowEntry `json:"port_allowed"`
}
// portU16 converts an operator-configured int port to the uint16 nftables
// expects. Returns 0 for out-of-range values; 0 is an unroutable TCP/UDP port
// so a misconfigured rule fails closed (no traffic matches) rather than
// wrapping silently to a valid-but-wrong port.
func portU16(p int) uint16 {
if p < 0 || p > 65535 {
return 0
}
// #nosec G115 -- bounded above.
return uint16(p)
}
// nftSocketBufferBytes is the netlink socket read/write buffer CSM requests for
// its nftables connections. The default OS netlink buffer
// (net.core.rmem_default, typically ~208 KiB) cannot hold the kernel ack stream
// for a full ruleset Apply on a host with a large blocklist; the apply then
// fails with ENOBUFS ("netlink receive: recvmsg: no buffer space available")
// and the firewall is left unmanaged. SetReadBuffer uses SO_RCVBUFFORCE under
// CAP_NET_ADMIN, so the larger size holds regardless of net.core.rmem_max.
const nftSocketBufferBytes = 8 << 20 // 8 MiB
// applyNFTSocketBuffer enlarges a netlink connection's receive and send
// buffers. It is best-effort: a resize failure is ignored so the connection
// still works with the default buffer (sufficient for small rulesets), and it
// returns nil so it never aborts the dial inside WithSockOptions.
func applyNFTSocketBuffer(nc *netlink.Conn) error {
_ = nc.SetReadBuffer(nftSocketBufferBytes)
_ = nc.SetWriteBuffer(nftSocketBufferBytes)
return nil
}
// newNFTConn opens an nftables netlink connection with enlarged socket buffers.
// WithSockOptions runs on every (re)dial, including the transient per-Flush
// sockets of a non-lasting connection, so each Apply uses the larger buffer.
func newNFTConn() (*nftables.Conn, error) {
return nftables.New(nftables.WithSockOptions(applyNFTSocketBuffer))
}
// NewEngine creates a new nftables firewall engine.
func NewEngine(cfg *FirewallConfig, statePath string) (*Engine, error) {
conn, err := newNFTConn()
if err != nil {
return nil, fmt.Errorf("nftables connection: %w", err)
}
e := &Engine{
conn: conn,
cfg: cfg,
statePath: filepath.Join(statePath, "firewall"),
}
_ = os.MkdirAll(e.statePath, 0700)
return e, nil
}
// ConnectExisting connects to an already-running CSM firewall.
// Used by CLI commands to modify the live ruleset without reapplying all rules.
func ConnectExisting(cfg *FirewallConfig, statePath string) (*Engine, error) {
conn, err := newNFTConn()
if err != nil {
return nil, fmt.Errorf("nftables connection: %w", err)
}
// Find existing CSM table
tables, err := conn.ListTables()
if err != nil {
return nil, fmt.Errorf("listing tables: %w", err)
}
var table *nftables.Table
for _, t := range tables {
if t.Name == "csm" && t.Family == nftables.TableFamilyINet {
table = t
break
}
}
if table == nil {
return nil, fmt.Errorf("CSM firewall not running (table 'csm' not found) - run 'csm firewall restart' first")
}
setBlocked, err := conn.GetSetByName(table, "blocked_ips")
if err != nil {
return nil, fmt.Errorf("blocked_ips set not found: %w", err)
}
setBlockedNet, err := conn.GetSetByName(table, "blocked_nets")
if err != nil {
return nil, fmt.Errorf("blocked_nets set not found: %w", err)
}
setAllowed, err := conn.GetSetByName(table, "allowed_ips")
if err != nil {
return nil, fmt.Errorf("allowed_ips set not found: %w", err)
}
setInfra, err := conn.GetSetByName(table, "infra_ips")
if err != nil {
return nil, fmt.Errorf("infra_ips set not found: %w", err)
}
setDOSExempt, err := conn.GetSetByName(table, "dos_exempt_nets")
if err != nil {
return nil, fmt.Errorf("dos_exempt_nets set not found: %w", err)
}
e := &Engine{
conn: conn,
cfg: cfg,
table: table,
setBlocked: setBlocked,
setBlockedNet: setBlockedNet,
setAllowed: setAllowed,
setInfra: setInfra,
setDOSExempt: setDOSExempt,
statePath: filepath.Join(statePath, "firewall"),
}
// Try to find Cloudflare whitelist sets (optional)
if s, err := conn.GetSetByName(table, "cf_whitelist"); err == nil {
e.setCFWhitelist = s
}
if s, err := conn.GetSetByName(table, "cf_whitelist6"); err == nil {
e.setCFWhitelist6 = s
}
// Try to find IPv6 sets (optional - may not exist if IPv6 disabled)
if s, err := conn.GetSetByName(table, "blocked_ips6"); err == nil {
e.setBlocked6 = s
}
if s, err := conn.GetSetByName(table, "blocked_nets6"); err == nil {
e.setBlockedNet6 = s
}
if s, err := conn.GetSetByName(table, "allowed_ips6"); err == nil {
e.setAllowed6 = s
}
if s, err := conn.GetSetByName(table, "infra_ips6"); err == nil {
e.setInfra6 = s
}
if s, err := conn.GetSetByName(table, "dos_exempt_nets6"); err == nil {
e.setDOSExempt6 = s
}
return e, nil
}
// SetDryRunRecorder installs a callback that is invoked by BlockIP whenever
// auto_response.dry_run is active. The daemon calls this after construction
// to wire in store.RecordDryRunBlock without creating an import cycle between
// internal/firewall and internal/store.
// errTableListing marks an Apply aborted because the kernel table listing
// failed; nothing was changed.
var errTableListing = errors.New("listing nftables tables")
// listTablesFn returns the table lister: the injected seam when a test set
// one, else the live connection.
func (e *Engine) listTablesFn() func() ([]*nftables.Table, error) {
if e.listTables != nil {
return e.listTables
}
return e.conn.ListTables
}
// SetConfig replaces the ruleset input the next Apply builds from. The
// engine was constructed with a value copy of the firewall block taken at
// daemon start, so every re-apply path (`csm firewall restart`,
// `apply-confirmed`, SIGHUP) rebuilt the old rules while reporting success
// and the operator's csm.yaml edit landed unprotected at the next restart.
// A nil cfg is ignored.
func (e *Engine) SetConfig(cfg *FirewallConfig) {
if cfg == nil {
return
}
e.mu.Lock()
defer e.mu.Unlock()
e.cfg = cfg
}
// Config returns the ruleset input the engine currently applies.
func (e *Engine) Config() *FirewallConfig {
e.mu.Lock()
defer e.mu.Unlock()
return e.cfg
}
func (e *Engine) SetDryRunRecorder(fn func(ip, reason string, timeout time.Duration)) {
e.mu.Lock()
e.dryRunRecorder = fn
e.mu.Unlock()
}
// SetDryRunEnabledFunc installs the callback BlockIP uses to decide whether
// auto_response.dry_run should intercept an automatic block. Nil means live.
func (e *Engine) SetDryRunEnabledFunc(fn func() bool) {
e.mu.Lock()
e.dryRunEnabled = fn
e.mu.Unlock()
}
// SetVerdictAsker installs the verdict callback the daemon constructs at
// startup. Nil disables the verdict callback (the gate skips entirely).
func (e *Engine) SetVerdictAsker(fn func(ctx context.Context, ip, reason string) (string, string, string, error)) {
e.mu.Lock()
e.verdictAsker = fn
e.mu.Unlock()
}
func (e *Engine) verdictAskerFn() func(ctx context.Context, ip, reason string) (string, string, string, error) {
e.mu.Lock()
fn := e.verdictAsker
e.mu.Unlock()
return fn
}
// SetSoftAllowChecker installs the callback the auto-block path consults to
// decide whether an IP belongs to a verified-bot range. Nil disables the
// verified-bot side of the soft-allow gate (operator allowed_ips is still
// honoured). The daemon wires this to threatintel so the firewall package
// stays free of that import.
func (e *Engine) SetSoftAllowChecker(fn func(ip string) bool) {
e.mu.Lock()
e.softAllowFn = fn
e.mu.Unlock()
}
// autoBlockSoftAllowed reports whether an automatic block of ip must be
// skipped because ip is on a soft-allow list: an operator full-IP allow,
// an operator port-specific allow, or a verified-bot range. Only the
// auto-block path (BlockIPOutcome) consults this; BlockIPForce bypasses it so
// an explicit operator deny still wins.
func (e *Engine) autoBlockSoftAllowed(ip string) bool {
if e.operatorSoftAllowed(ip) {
return true
}
return e.autoBlockVerifiedRange(ip)
}
func (e *Engine) autoBlockVerifiedRange(ip string) bool {
e.mu.Lock()
fn := e.softAllowFn
e.mu.Unlock()
return fn != nil && fn(ip)
}
func (e *Engine) logAutoBlockSoftAllowed(ip string) {
fmt.Fprintf(os.Stderr, "[%s] auto-block: %s is allowlisted or a verified bot - not blocking\n",
time.Now().Format("2006-01-02 15:04:05"), ip)
}
// SetShutdownContext installs a context whose cancellation aborts any
// in-flight verdict callback. The daemon ties this to its stopCh so a
// graceful shutdown does not have to wait for an unresponsive panel
// callback to return.
func (e *Engine) SetShutdownContext(ctx context.Context) {
e.mu.Lock()
e.shutdownCtx = ctx
e.mu.Unlock()
}
func (e *Engine) verdictContext() context.Context {
e.mu.Lock()
ctx := e.shutdownCtx
e.mu.Unlock()
if ctx == nil {
return context.Background()
}
return ctx
}
// SetDOSExemptProviderNets stores the mail-provider IP ranges used by the
// dos_exempt_nets set. The daemon calls this before Apply() and on each
// provider refresh. Nil is valid when the daemon has no provider data yet;
// createSets() will produce an empty set in that case.
func (e *Engine) SetDOSExemptProviderNets(nets []*net.IPNet) {
e.mu.Lock()
e.dosExemptProviderNets = cloneIPNetSlice(nets)
e.mu.Unlock()
}
func cloneIPNetSlice(nets []*net.IPNet) []*net.IPNet {
if len(nets) == 0 {
return nil
}
out := make([]*net.IPNet, len(nets))
for i, n := range nets {
if n == nil {
continue
}
ip := make(net.IP, len(n.IP))
copy(ip, n.IP)
mask := make(net.IPMask, len(n.Mask))
copy(mask, n.Mask)
out[i] = &net.IPNet{IP: ip, Mask: mask}
}
return out
}
// dosExemptIntervalElems builds nftables interval set elements from nets.
// Callers must pass nets from a single IP family (all IPv4 or all IPv6).
func dosExemptIntervalElems(nets []*net.IPNet) []nftables.SetElement {
if len(nets) == 0 {
return nil
}
var elems []nftables.SetElement
for _, n := range nets {
if n == nil {
continue
}
var start net.IP
if n.IP.To4() != nil {
start = n.IP.To4()
} else {
start = n.IP.To16()
}
end := lastIPInRange(n)
if start == nil || end == nil {
continue
}
elems = appendIntervalSetElements(elems, start, end)
}
return normalizeIntervalElements(elems)
}
// RefreshDOSExemptSets repopulates the dos_exempt_nets[6] interval sets with
// a new provider overlay in a single batched kernel transaction. If the kernel
// batch fails the previous set contents remain active and dosExemptProviderNets
// is not updated, preserving the last-known good state.
func (e *Engine) RefreshDOSExemptSets(providerNets []*net.IPNet) error {
e.mu.Lock()
defer e.mu.Unlock()
if err := e.lifecycleReadyLocked(); err != nil {
return err
}
if e.setDOSExempt == nil {
return fmt.Errorf("dos_exempt_nets set not initialized")
}
v4, v6 := EffectiveDOSExemptNets(e.cfg, providerNets)
// Queue flush + repopulate for IPv4. Nothing is sent to the kernel yet.
e.conn.FlushSet(e.setDOSExempt)
if elems4 := dosExemptIntervalElems(v4); len(elems4) > 0 {
if err := e.conn.SetAddElements(e.setDOSExempt, elems4); err != nil {
return fmt.Errorf("queuing dos_exempt_nets elements: %w", err)
}
}
// Queue flush + repopulate for IPv6 when the set exists.
if e.setDOSExempt6 != nil {
e.conn.FlushSet(e.setDOSExempt6)
if elems6 := dosExemptIntervalElems(v6); len(elems6) > 0 {
if err := e.conn.SetAddElements(e.setDOSExempt6, elems6); err != nil {
return fmt.Errorf("queuing dos_exempt_nets6 elements: %w", err)
}
}
}
// Atomic commit. On failure the kernel keeps the previous elements.
if err := e.conn.Flush(); err != nil {
return fmt.Errorf("refreshing dos_exempt sets: %w", err)
}
// Update overlay only after the kernel confirmed the transaction.
e.dosExemptProviderNets = cloneIPNetSlice(providerNets)
return nil
}
// dosExemptV4Lookup returns a two-expression sequence that loads the IPv4
// source address (network-header offset 12, 4 bytes) into reg and performs an
// inverted set lookup against dos_exempt_nets. The caller must place the result
// after an IPv4 family guard and before the meter expressions. The Ct and
// Payload expressions that follow reload reg independently, so register reuse
// is safe.
func (e *Engine) dosExemptV4Lookup(reg uint32) []expr.Any {
return []expr.Any{
&expr.Payload{DestRegister: reg, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Lookup{SourceRegister: reg, SetName: e.setDOSExempt.Name, SetID: e.setDOSExempt.ID, Invert: true},
}
}
// dosExemptV6Lookup is the IPv6 analogue of dosExemptV4Lookup, loading the
// IPv6 source address (network-header offset 8, 16 bytes) against
// dos_exempt_nets6. Use only in rules that carry IPv6 traffic.
func (e *Engine) dosExemptV6Lookup(reg uint32) []expr.Any {
return []expr.Any{
&expr.Payload{DestRegister: reg, Base: expr.PayloadBaseNetworkHeader, Offset: 8, Len: 16},
&expr.Lookup{SourceRegister: reg, SetName: e.setDOSExempt6.Name, SetID: e.setDOSExempt6.ID, Invert: true},
}
}
// Apply builds and atomically applies the complete nftables ruleset.
// All operations (delete old table + create new table/rules +
// populate persisted block/allow entries) are batched into a single
// netlink transaction. If the flush fails, the kernel keeps whatever
// ruleset was running before - the server is never left without a
// firewall. Equally important: the new ruleset never appears with
// EMPTY blocked sets between the table-swap and the persisted-state
// load; an attacker IP from state.json is blocked from the moment
// the new table becomes the live one.
func (e *Engine) Apply() (resultErr error) {
defer func() {
if e.shouldLegacyOutcome(resultErr) {
recordFirewallResult("apply", "csm", "", "", 0, actionlog.Applied, resultErr)
}
}()
e.mu.Lock()
defer e.mu.Unlock()
if err := e.lifecycleReadyLocked(); err != nil {
return err
}
if e.lifecycle != nil {
return e.applyDurableLocked()
}
return e.applyRulesetLocked("")
}
func (e *Engine) applyRulesetLocked(marker string) error {
// Compute the elements to seed each set from persisted state
// BEFORE touching nft. Pure computation; if state.json is
// missing or malformed the slices stay empty and the new table
// still applies (no firewall regression on a fresh install).
initial := e.computeInitialBlockStateLocked()
// Check if existing CSM table needs replacing.
// If so, include the delete in the same atomic batch as the new table.
// A listing that fails must abort: AddTable is create-only, so without
// the delete every rule below would be appended a second time to the
// live chains (doubled meters, halved rate limits).
tables, err := e.listTablesFn()()
if err != nil {
return fmt.Errorf("%w: %v", errTableListing, err)
}
for _, t := range tables {
if t.Name == "csm" && t.Family == nftables.TableFamilyINet {
e.conn.DelTable(t)
break
}
}
// Create table - all operations below are batched, nothing is sent until Flush()
e.table = e.conn.AddTable(&nftables.Table{
Name: "csm",
Family: nftables.TableFamilyINet,
})
// Create IP sets
if err := e.createSets(); err != nil {
return fmt.Errorf("creating sets: %w", err)
}
// Create chains and rules
if err := e.createInputChain(); err != nil {
return fmt.Errorf("creating input chain: %w", err)
}
if err := e.createOutputChain(); err != nil {
return fmt.Errorf("creating output chain: %w", err)
}
// Queue initial set elements from persisted state into the same
// netlink batch as the table+set+chain creation above. Without
// this, Apply previously Flushed an empty-set ruleset, then a
// separate loadState() Flush populated the sets - leaving a
// brief window where the new table existed without the
// persisted blocks.
if err := e.queueInitialBlockStateLocked(initial); err != nil {
return fmt.Errorf("queueing initial firewall state: %w", err)
}
if marker != "" {
e.conn.AddChain(&nftables.Chain{Name: marker, Table: e.table})
}
// Apply atomically - if this fails, nftables keeps whatever was running before
if err := e.conn.Flush(); err != nil {
return fmt.Errorf("applying ruleset: %w", err)
}
// A failed capture leaves monitoring visibly unbaselined; it must neither
// undo an applied firewall nor bless a later, possibly external edit.
e.appliedRuleset, _ = e.readRulesetLocked()
return nil
}
// createSets creates the nftables named sets for IP management.
func (e *Engine) createSets() error {
// Blocked IPs set. HasTimeout enables per-element timeouts for temporary
// blocks; there must be NO set-level default timeout: the kernel assigns
// the default to every element added without its own timeout (the netlink
// library omits zero element timeouts), so a default would silently expire
// permanent blocks (deny, promote) while state.json still says blocked.
e.setBlocked = &nftables.Set{
Table: e.table,
Name: "blocked_ips",
KeyType: nftables.TypeIPAddr,
HasTimeout: true,
}
if err := e.conn.AddSet(e.setBlocked, nil); err != nil {
return fmt.Errorf("blocked set: %w", err)
}
// Blocked subnets set (interval for CIDR ranges, permanent)
e.setBlockedNet = &nftables.Set{
Table: e.table,
Name: "blocked_nets",
KeyType: nftables.TypeIPAddr,
Interval: true,
}
if err := e.conn.AddSet(e.setBlockedNet, nil); err != nil {
return fmt.Errorf("blocked nets set: %w", err)
}
// Allowed IPs set
e.setAllowed = &nftables.Set{
Table: e.table,
Name: "allowed_ips",
KeyType: nftables.TypeIPAddr,
}
if err := e.conn.AddSet(e.setAllowed, nil); err != nil {
return fmt.Errorf("allowed set: %w", err)
}
// Infra IPs set (interval for CIDR support) - split IPv4 and IPv6
e.setInfra = &nftables.Set{
Table: e.table,
Name: "infra_ips",
KeyType: nftables.TypeIPAddr,
Interval: true,
}
var infraElements []nftables.SetElement
var infra6Elements []nftables.SetElement
for _, cidr := range e.cfg.InfraIPs {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
ip := net.ParseIP(cidr)
if ip == nil {
continue
}
if ip4 := ip.To4(); ip4 != nil {
infraElements = appendIntervalSetElements(infraElements, ip4, ip4)
} else if e.cfg.IPv6 {
ip16 := ip.To16()
infra6Elements = appendIntervalSetElements(infra6Elements, ip16, ip16)
}
continue
}
if network.IP.To4() != nil {
start := network.IP.To4()
end := lastIPInRange(network)
if start != nil && end != nil {
infraElements = appendIntervalSetElements(infraElements, start, end)
}
} else if e.cfg.IPv6 {
start := network.IP.To16()
end := lastIPInRange(network)
if start != nil && end != nil {
infra6Elements = appendIntervalSetElements(infra6Elements, start, end)
}
}
}
if err := e.conn.AddSet(e.setInfra, normalizeIntervalElements(infraElements)); err != nil {
return fmt.Errorf("infra set: %w", err)
}
// Cloudflare IP whitelist sets (interval for CIDR ranges, accept on 80/443 only)
e.setCFWhitelist = &nftables.Set{
Table: e.table,
Name: "cf_whitelist",
KeyType: nftables.TypeIPAddr,
Interval: true,
}
if err := e.conn.AddSet(e.setCFWhitelist, nil); err != nil {
return fmt.Errorf("cf_whitelist set: %w", err)
}
e.setCFWhitelist6 = &nftables.Set{
Table: e.table,
Name: "cf_whitelist6",
KeyType: nftables.TypeIP6Addr,
Interval: true,
}
if err := e.conn.AddSet(e.setCFWhitelist6, nil); err != nil {
return fmt.Errorf("cf_whitelist6 set: %w", err)
}
// Country-blocked IPs set (interval for CIDR ranges)
if countryBlockingActive(e.cfg) {
e.setCountry = &nftables.Set{
Table: e.table,
Name: "country_blocked",
KeyType: nftables.TypeIPAddr,
Interval: true,
}
var countryElements []nftables.SetElement
for _, code := range e.cfg.CountryBlock {
countryElements = append(countryElements, loadCountryCIDRs(CountryDBDir(e.cfg, e.statePath), code)...)
}
if err := e.conn.AddSet(e.setCountry, normalizeIntervalElements(countryElements)); err != nil {
fmt.Fprintf(os.Stderr, "firewall: warning creating country set: %v\n", err)
e.setCountry = nil
} else if len(countryElements) > 0 {
fmt.Fprintf(os.Stderr, "firewall: loaded %d country block ranges for %v\n",
len(countryElements)/2, e.cfg.CountryBlock)
}
}
// IPv6 sets
if e.cfg.IPv6 {
// No set-level default timeout; see the blocked_ips comment above.
e.setBlocked6 = &nftables.Set{
Table: e.table, Name: "blocked_ips6",
KeyType: nftables.TypeIP6Addr, HasTimeout: true,
}
if err := e.conn.AddSet(e.setBlocked6, nil); err != nil {
return fmt.Errorf("blocked6 set: %w", err)
}
e.setBlockedNet6 = &nftables.Set{
Table: e.table, Name: "blocked_nets6",
KeyType: nftables.TypeIP6Addr, Interval: true,
}
if err := e.conn.AddSet(e.setBlockedNet6, nil); err != nil {
return fmt.Errorf("blocked_nets6 set: %w", err)
}
e.setAllowed6 = &nftables.Set{
Table: e.table, Name: "allowed_ips6",
KeyType: nftables.TypeIP6Addr,
}
if err := e.conn.AddSet(e.setAllowed6, nil); err != nil {
return fmt.Errorf("allowed6 set: %w", err)
}
e.setInfra6 = &nftables.Set{
Table: e.table, Name: "infra_ips6",
KeyType: nftables.TypeIP6Addr, Interval: true,
}
if err := e.conn.AddSet(e.setInfra6, normalizeIntervalElements(infra6Elements)); err != nil {
return fmt.Errorf("infra6 set: %w", err)
}
// IPv6 country-block set, mirroring the IPv4 set above. Without it an
// attacker on an IPv6 address from a blocked country was never dropped.
if countryBlockingActive(e.cfg) {
e.setCountry6 = &nftables.Set{
Table: e.table, Name: "country_blocked6",
KeyType: nftables.TypeIP6Addr, Interval: true,
}
var country6Elements []nftables.SetElement
for _, code := range e.cfg.CountryBlock {
country6Elements = append(country6Elements, loadCountryCIDRs6(CountryDBDir(e.cfg, e.statePath), code)...)
}
if err := e.conn.AddSet(e.setCountry6, normalizeIntervalElements(country6Elements)); err != nil {
fmt.Fprintf(os.Stderr, "firewall: warning creating IPv6 country set: %v\n", err)
e.setCountry6 = nil
} else if len(country6Elements) > 0 {
fmt.Fprintf(os.Stderr, "firewall: loaded %d IPv6 country block ranges for %v\n",
len(country6Elements)/2, e.cfg.CountryBlock)
}
}
}
// DoS-exempt sets: sources here bypass connection meters and mail-port
// flood rules. Populated from operator dos_exempt_ranges + provider
// overlay (pushed by the daemon before Apply).
v4Exempt, v6Exempt := EffectiveDOSExemptNets(e.cfg, e.dosExemptProviderNets)
e.setDOSExempt = &nftables.Set{
Table: e.table,
Name: "dos_exempt_nets",
KeyType: nftables.TypeIPAddr,
Interval: true,
}
if err := e.conn.AddSet(e.setDOSExempt, dosExemptIntervalElems(v4Exempt)); err != nil {
return fmt.Errorf("dos_exempt_nets set: %w", err)
}
if e.cfg.IPv6 {
e.setDOSExempt6 = &nftables.Set{
Table: e.table,
Name: "dos_exempt_nets6",
KeyType: nftables.TypeIP6Addr,
Interval: true,
}
if err := e.conn.AddSet(e.setDOSExempt6, dosExemptIntervalElems(v6Exempt)); err != nil {
return fmt.Errorf("dos_exempt_nets6 set: %w", err)
}
}
// Meter sets for per-IP rate limiting (dynamic sets). The IPv6 siblings key
// on the /64 prefix; see ipv6MeterSet.
if e.cfg.SYNFloodProtection {
e.meterSYN = &nftables.Set{
Table: e.table, Name: "meter_syn", KeyType: nftables.TypeIPAddr,
Dynamic: true, HasTimeout: true, Timeout: time.Minute,
}
_ = e.conn.AddSet(e.meterSYN, nil)
if e.cfg.IPv6 {
e.meterSYN6 = e.ipv6MeterSet("meter_syn6")
}
}
if e.cfg.ConnRateLimit > 0 {
e.meterConn = &nftables.Set{
Table: e.table, Name: "meter_conn", KeyType: nftables.TypeIPAddr,
Dynamic: true, HasTimeout: true, Timeout: time.Minute,
}
_ = e.conn.AddSet(e.meterConn, nil)
if e.cfg.IPv6 {
e.meterConn6 = e.ipv6MeterSet("meter_conn6")
}
}
if e.cfg.UDPFlood && e.cfg.UDPFloodRate > 0 {
e.meterUDP = &nftables.Set{
Table: e.table, Name: "meter_udp", KeyType: nftables.TypeIPAddr,
Dynamic: true, HasTimeout: true, Timeout: time.Minute,
}
_ = e.conn.AddSet(e.meterUDP, nil)
if e.cfg.IPv6 {
e.meterUDP6 = e.ipv6MeterSet("meter_udp6")
}
}
if e.cfg.ConnLimit > 0 {
e.meterConnlim = &nftables.Set{
Table: e.table, Name: "meter_connlimit", KeyType: nftables.TypeIPAddr,
Dynamic: true,
}
_ = e.conn.AddSet(e.meterConnlim, nil)
}
if portFloodNeedsMeter(e.cfg.PortFlood) {
e.meterPortFlood4 = make(map[string]*nftables.Set)
if e.cfg.IPv6 {
e.meterPortFlood6 = make(map[string]*nftables.Set)
}
for _, pf := range e.cfg.PortFlood {
if !usablePortFloodRule(pf) {
continue
}
name4 := portFloodMeterName(pf, portFloodIPv4)
if _, ok := e.meterPortFlood4[name4]; !ok {
set := &nftables.Set{
Table: e.table, Name: name4, KeyType: nftables.TypeIPAddr,
Dynamic: true, HasTimeout: true, Timeout: time.Minute,
}
_ = e.conn.AddSet(set, nil)
e.meterPortFlood4[name4] = set
}
if e.cfg.IPv6 {
name6 := portFloodMeterName(pf, portFloodIPv6)
if _, ok := e.meterPortFlood6[name6]; !ok {
set := &nftables.Set{
Table: e.table, Name: name6, KeyType: nftables.TypeIP6Addr,
Dynamic: true, HasTimeout: true, Timeout: time.Minute,
}
_ = e.conn.AddSet(set, nil)
e.meterPortFlood6[name6] = set
}
}
}
}
return nil
}
// portFloodNeedsMeter reports whether any port_flood rule has a usable rate,
// so the meter set is only created when at least one rule will reference it.
func portFloodNeedsMeter(rules []PortFloodRule) bool {
for _, pf := range rules {
if usablePortFloodRule(pf) {
return true
}
}
return false
}
func usablePortFloodRule(pf PortFloodRule) bool {
return pf.Hits > 0 && pf.Seconds > 0 && pf.Port > 0
}
// createInputChain builds the input filter chain with proper rule ordering.
func (e *Engine) createInputChain() error {
policy := nftables.ChainPolicyDrop
e.chainIn = e.conn.AddChain(&nftables.Chain{
Name: "input",
Table: e.table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookInput,
Priority: nftables.ChainPriorityFilter,
Policy: &policy,
})
// Rule 1: Allow loopback
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyIIFNAME, Register: 1},
&expr.Cmp{
Op: expr.CmpOpEq,
Register: 1,
Data: []byte("lo\x00"),
},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
if !e.cfg.IPv6 {
e.conn.AddRule(&nftables.Rule{Table: e.table, Chain: e.chainIn, Exprs: familyBypassRuleExprs(10)})
}
// Rule 2: Drop INVALID conntrack state (malformed packets)
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitINVALID),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(0)},
&expr.Verdict{Kind: expr.VerdictDrop},
},
})
// Rule 3: Allow infra IPs FIRST - infra must NEVER be blocked, even accidentally
e.addSetMatchRule(e.setInfra, expr.VerdictAccept)
e.addSetMatchRuleV6(e.setInfra6, expr.VerdictAccept)
// Rule 4: Cloudflare IP whitelist - accept on TCP 80/443 only.
// CF IPs can still be blocked on other ports (unlike infra).
e.addCFWhitelistRule(e.setCFWhitelist, false)
e.addCFWhitelistRule(e.setCFWhitelist6, true)
// Rule 5: Drop blocked IPs before established/related so active
// keep-alive connections do not bypass a new block.
e.addSetMatchRule(e.setBlocked, expr.VerdictDrop)
e.addSetMatchRuleV6(e.setBlocked6, expr.VerdictDrop)
// Rule 6: Drop blocked subnets (interval set for CIDR ranges)
e.addSetMatchRule(e.setBlockedNet, expr.VerdictDrop)
e.addSetMatchRuleV6(e.setBlockedNet6, expr.VerdictDrop)
// Rule 7: Allow established/related connections after block checks.
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1,
DestRegister: 1,
Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitESTABLISHED | expr.CtStateBitRELATED),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{
Op: expr.CmpOpNeq,
Register: 1,
Data: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
// Rule 8: Allow explicitly allowed IPs
e.addSetMatchRule(e.setAllowed, expr.VerdictAccept)
e.addSetMatchRuleV6(e.setAllowed6, expr.VerdictAccept)
// Rule 9: Port-specific allows (IP+port, e.g. MySQL access for specific IPs)
state := e.loadStateFile()
for _, pa := range state.PortAllowed {
exprs := buildPortAllowExprs(pa, e.cfg.IPv6)
if exprs == nil {
continue
}
e.conn.AddRule(&nftables.Rule{Table: e.table, Chain: e.chainIn, Exprs: exprs})
}
// ICMPv4 echo-request (type 8)
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{1}}, // ICMPv4
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{8}}, // echo-request
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
// ICMPv6 echo-request (type 128) + neighbor discovery (types 133-137, required for IPv6)
if e.cfg.IPv6 {
for _, icmp6Type := range []byte{128, 133, 134, 135, 136, 137} {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{58}}, // ICMPv6
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{icmp6Type}},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
}
}
// Per-IP SYN flood protection via meter
if e.cfg.SYNFloodProtection && e.meterSYN != nil {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.synFloodRuleExprs(),
})
if e.meterSYN6 != nil {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.synFloodRuleExprs6(),
})
}
}
// Per-IP new connection rate limit via meter
if e.cfg.ConnRateLimit > 0 && e.meterConn != nil {
// #nosec G115 -- ConnRateLimit is an operator-configured int (typical 10–1000);
// /2 is non-negative and well below uint32 max.
burst := uint32(e.cfg.ConnRateLimit / 2)
if burst < 5 {
burst = 5
}
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.connMeterRuleExprs(uint64(e.cfg.ConnRateLimit), burst),
})
if e.meterConn6 != nil {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.connMeterRuleExprs6(uint64(e.cfg.ConnRateLimit), burst),
})
}
}
// Per-IP concurrent connection limit (CONNLIMIT)
if e.cfg.ConnLimit > 0 && e.meterConnlim != nil {
// #nosec G115 -- ConnLimit is operator-configured non-negative int; fits in uint32.
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.connlimitRuleExprs(uint32(e.cfg.ConnLimit)),
})
}
// Rule 9: Country block - drop traffic from blocked countries
if e.setCountry != nil {
e.addSetMatchRule(e.setCountry, expr.VerdictDrop)
}
if e.setCountry6 != nil {
e.addSetMatchRuleV6(e.setCountry6, expr.VerdictDrop)
}
// Per-port flood protection - rate-limit new connections per source IP.
// Each rule has separate IPv4 and IPv6 meters so ports and families do not
// consume each other's token buckets.
for _, pf := range e.cfg.PortFlood {
items := []struct {
family portFloodIPFamily
meter *nftables.Set
}{
{family: portFloodIPv4, meter: e.meterPortFlood4[portFloodMeterName(pf, portFloodIPv4)]},
}
if e.cfg.IPv6 {
items = append(items, struct {
family portFloodIPFamily
meter *nftables.Set
}{family: portFloodIPv6, meter: e.meterPortFlood6[portFloodMeterName(pf, portFloodIPv6)]})
}
for _, item := range items {
exprs := e.portFloodRuleExprs(pf, item.meter, item.family)
if exprs == nil {
continue
}
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: exprs,
})
}
}
// Per-IP UDP flood protection via meter
if e.cfg.UDPFlood && e.cfg.UDPFloodRate > 0 && e.meterUDP != nil {
// #nosec G115 -- UDPFloodBurst is operator-configured non-negative int.
burst := uint32(e.cfg.UDPFloodBurst)
if burst < 10 {
burst = 10
}
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.udpFloodRuleExprs(uint64(e.cfg.UDPFloodRate), burst),
})
if e.meterUDP6 != nil {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: e.udpFloodRuleExprs6(uint64(e.cfg.UDPFloodRate), burst),
})
}
}
// Build restricted port set - these are only reachable via infra IPs (rule 4)
restricted := make(map[int]bool)
for _, p := range e.cfg.RestrictedTCP {
restricted[p] = true
}
// Open TCP ports (public) - restricted ports excluded
for _, port := range e.cfg.TCPIn {
if restricted[port] {
continue
}
e.addPortAcceptRule(port, true, 2)
}
// Open UDP ports (public)
for _, port := range e.cfg.UDPIn {
e.addPortAcceptRule(port, false, 2)
}
if e.cfg.IPv6 {
for _, port := range effectiveIPv6Ports(e.cfg.TCPIn, e.cfg.TCP6In) {
if !restricted[port] {
e.addPortAcceptRule(port, true, 10)
}
}
for _, port := range effectiveIPv6Ports(e.cfg.UDPIn, e.cfg.UDP6In) {
e.addPortAcceptRule(port, false, 10)
}
}
// Passive FTP range
if e.cfg.PassiveFTPStart > 0 && e.cfg.PassiveFTPEnd > 0 {
e.addPortRangeAcceptRule(e.cfg.PassiveFTPStart, e.cfg.PassiveFTPEnd, true, 2)
if e.cfg.IPv6 {
e.addPortRangeAcceptRule(e.cfg.PassiveFTPStart, e.cfg.PassiveFTPEnd, true, 10)
}
}
// Silent drop for commonly-scanned ports (no logging)
for _, port := range e.cfg.DropNoLog {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}}, // TCP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(port))},
&expr.Verdict{Kind: expr.VerdictDrop},
},
})
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{17}}, // UDP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(port))},
&expr.Verdict{Kind: expr.VerdictDrop},
},
})
}
// Rate-limited log for remaining dropped packets
if e.cfg.LogDropped {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Limit{
Type: expr.LimitTypePkts,
Rate: uint64(max(e.cfg.LogRate, 1)),
Unit: expr.LimitTimeMinute,
Burst: 5,
},
&expr.Log{Key: 1, Data: []byte("CSM-DROP: ")},
},
})
}
// Default policy is DROP - anything not matched above is dropped
return nil
}
// createOutputChain builds the output filter chain.
// Restricts outbound to configured ports only (prevents C2 on non-standard ports).
func (e *Engine) createOutputChain() error {
tcp6Out := effectiveIPv6Ports(e.cfg.TCPOut, e.cfg.TCP6Out)
udp6Out := effectiveIPv6Ports(e.cfg.UDPOut, e.cfg.UDP6Out)
if len(e.cfg.TCPOut) == 0 && len(e.cfg.UDPOut) == 0 && (!e.cfg.IPv6 || len(tcp6Out) == 0 && len(udp6Out) == 0) {
// No outbound restrictions configured - accept all
policy := nftables.ChainPolicyAccept
e.chainOut = e.conn.AddChain(&nftables.Chain{
Name: "output",
Table: e.table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookOutput,
Priority: nftables.ChainPriorityFilter,
Policy: &policy,
})
return nil
}
// Outbound filtering enabled
policy := nftables.ChainPolicyDrop
e.chainOut = e.conn.AddChain(&nftables.Chain{
Name: "output",
Table: e.table,
Type: nftables.ChainTypeFilter,
Hooknum: nftables.ChainHookOutput,
Priority: nftables.ChainPriorityFilter,
Policy: &policy,
})
// Allow established/related outbound
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1,
DestRegister: 1,
Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitESTABLISHED | expr.CtStateBitRELATED),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(0)},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
// Allow loopback outbound
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyOIFNAME, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte("lo\x00")},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
if !e.cfg.IPv6 {
e.conn.AddRule(&nftables.Rule{Table: e.table, Chain: e.chainOut, Exprs: familyBypassRuleExprs(10)})
}
// SMTP block - restrict outbound mail to allowed users only.
// resolveSMTPAllowedUIDs unconditionally includes root and mailnull
// (exim's queue runner); without mailnull on the allow list queued
// mail is silently dropped while CSM still reports healthy.
smtpBlocked := make(map[int]bool)
if e.cfg.SMTPBlock && len(e.cfg.SMTPPorts) > 0 {
allowedUIDs := resolveSMTPAllowedUIDs(e.cfg.SMTPAllowUsers)
for _, port := range e.cfg.SMTPPorts {
smtpBlocked[port] = true
// Accept from each allowed UID
for _, uid := range allowedUIDs {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(port))},
&expr.Meta{Key: expr.MetaKeySKUID, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(uid)},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
}
// Drop SMTP from everyone else
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(port))},
&expr.Verdict{Kind: expr.VerdictDrop},
},
})
}
}
// The operator restricted egress for another family (e.g. IPv6 only) but
// configured no IPv4 out-ports. Accept IPv4 egress wholesale rather than let
// the chain's DROP policy sever it. Placed after the SMTP guard above so
// smtp_block still governs IPv4 mail.
if e.cfg.restrictedOutputNeedsIPv4Bypass() {
e.conn.AddRule(&nftables.Rule{Table: e.table, Chain: e.chainOut, Exprs: familyBypassRuleExprs(2)})
}
// Destination-scoped outbound allows. Deliberately emitted AFTER the
// smtp_block guard above: nftables takes the first verdict, so the drop on
// each smtp port is already decided and no range here can bypass it. The
// config validator rejects an overlapping range as well; rule order must
// not be the only guard between a config key and outbound mail.
for _, r := range e.cfg.TCPOutAllow {
exprs := buildOutAllowExprs(r, e.cfg.IPv6)
if exprs == nil {
continue
}
e.conn.AddRule(&nftables.Rule{Table: e.table, Chain: e.chainOut, Exprs: exprs})
}
// Allow configured outbound TCP ports (skip SMTP-blocked ports - handled above)
for _, port := range e.cfg.TCPOut {
if smtpBlocked[port] {
continue
}
e.addOutboundPortRule(port, true, 2)
}
// Allow configured outbound UDP ports
for _, port := range e.cfg.UDPOut {
e.addOutboundPortRule(port, false, 2)
}
if e.cfg.IPv6 {
for _, port := range tcp6Out {
if !smtpBlocked[port] {
e.addOutboundPortRule(port, true, 10)
}
}
for _, port := range udp6Out {
e.addOutboundPortRule(port, false, 10)
}
}
// Allow only safe ICMP outbound (echo-reply + echo-request, block dest-unreachable)
// Blocking ICMP type 3 (dest-unreachable) prevents leaking closed port info to scanners
for _, icmpType := range []byte{0, 8} { // 0=echo-reply, 8=echo-request
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{1}}, // ICMP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{icmpType}},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
}
// ICMPv6 outbound - allow echo-reply (129) + echo-request (128) + ND (133-137)
if e.cfg.IPv6 {
for _, icmp6Type := range []byte{128, 129, 133, 134, 135, 136, 137} {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{58}}, // ICMPv6
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 0, Len: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{icmp6Type}},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
}
}
// REJECT outbound TCP with RST (faster failure than silent DROP)
// UDP still silently drops via chain policy.
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}},
&expr.Reject{Type: 1}, // NFT_REJECT_TCP_RST
},
})
return nil
}
// --- Helper methods ---
// resolveIPSet returns the appropriate set and key bytes for an IP address.
// Falls back to IPv6 set if the IP is not IPv4.
func (e *Engine) resolveIPSet(ip string, set4, set6 *nftables.Set) (*nftables.Set, []byte, error) {
parsed := net.ParseIP(ip)
if parsed == nil {
return nil, nil, fmt.Errorf("invalid IP: %s", ip)
}
if ip4 := parsed.To4(); ip4 != nil {
return set4, ip4, nil
}
if set6 == nil {
return nil, nil, fmt.Errorf("IPv6 not enabled in firewall config: %s", ip)
}
return set6, parsed.To16(), nil
}
func canonicalFirewallIP(ip string) (string, error) {
parsed := net.ParseIP(ip)
if parsed == nil {
return "", fmt.Errorf("invalid IP: %s", ip)
}
return parsed.String(), nil
}
func canonicalIPKey(ip string) (string, bool) {
parsed := net.ParseIP(ip)
if parsed == nil {
return "", false
}
return parsed.String(), true
}
func stateIPKey(ip string) string {
if key, ok := canonicalIPKey(ip); ok {
return key
}
return ip
}
func sameIPString(a, b string) bool {
if a == b {
return true
}
aKey, aOK := canonicalIPKey(a)
if !aOK {
return false
}
bKey, bOK := canonicalIPKey(b)
return bOK && aKey == bKey
}
func preferLongerLivedBlock(current, candidate BlockedEntry) BlockedEntry {
if current.ExpiresAt.IsZero() {
return current
}
if candidate.ExpiresAt.IsZero() || candidate.ExpiresAt.After(current.ExpiresAt) {
return candidate
}
return current
}
func preferLongerLivedAllow(current, candidate AllowedEntry) AllowedEntry {
if current.ExpiresAt.IsZero() {
return current
}
if candidate.ExpiresAt.IsZero() || candidate.ExpiresAt.After(current.ExpiresAt) {
return candidate
}
return current
}
func normalizeFirewallStateIPs(state *FirewallState) {
blockedPos := make(map[string]int, len(state.Blocked))
blocked := state.Blocked[:0]
for _, entry := range state.Blocked {
key, ok := canonicalIPKey(entry.IP)
if !ok {
blocked = append(blocked, entry)
continue
}
entry.IP = key
if i, exists := blockedPos[key]; exists {
blocked[i] = preferLongerLivedBlock(blocked[i], entry)
continue
}
blockedPos[key] = len(blocked)
blocked = append(blocked, entry)
}
state.Blocked = blocked
type allowedKey struct {
ip string
source string
}
allowedPos := make(map[allowedKey]int, len(state.Allowed))
allowed := state.Allowed[:0]
for _, entry := range state.Allowed {
key, ok := canonicalIPKey(entry.IP)
if !ok {
allowed = append(allowed, entry)
continue
}
entry.IP = key
dedupKey := allowedKey{ip: key, source: entry.Source}
if i, exists := allowedPos[dedupKey]; exists {
allowed[i] = preferLongerLivedAllow(allowed[i], entry)
continue
}
allowedPos[dedupKey] = len(allowed)
allowed = append(allowed, entry)
}
state.Allowed = allowed
type portAllowKey struct {
ip string
port int
proto string
}
portAllowedPos := make(map[portAllowKey]struct{}, len(state.PortAllowed))
portAllowed := state.PortAllowed[:0]
for _, entry := range state.PortAllowed {
key, ok := canonicalIPKey(entry.IP)
if !ok {
portAllowed = append(portAllowed, entry)
continue
}
entry.IP = key
dedupKey := portAllowKey{ip: key, port: entry.Port, proto: entry.Proto}
if _, exists := portAllowedPos[dedupKey]; exists {
continue
}
portAllowedPos[dedupKey] = struct{}{}
portAllowed = append(portAllowed, entry)
}
state.PortAllowed = portAllowed
}
// addSetMatchRule adds an IPv4 source-IP set match rule on the input chain.
func (e *Engine) addSetMatchRule(set *nftables.Set, verdict expr.VerdictKind) {
exprs := buildSetMatchRuleExprs(set, verdict, 2, 12, 4)
if exprs == nil {
return
}
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: exprs,
})
}
// addSetMatchRuleV6 adds an IPv6 source-IP set match rule on the input chain.
func (e *Engine) addSetMatchRuleV6(set *nftables.Set, verdict expr.VerdictKind) {
exprs := buildSetMatchRuleExprs(set, verdict, 10, 8, 16)
if exprs == nil {
return
}
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: exprs,
})
}
func buildSetMatchRuleExprs(set *nftables.Set, verdict expr.VerdictKind, nfproto byte, sourceOffset, sourceLen uint32) []expr.Any {
if set == nil {
return nil
}
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{nfproto}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: sourceOffset, Len: sourceLen},
&expr.Lookup{SourceRegister: 1, SetName: set.Name, SetID: set.ID},
&expr.Verdict{Kind: verdict},
}
}
// ipv4NFProtoGuard returns the NFPROTO==IPV4 family guard prepended to every
// per-IP meter rule. The meters live in the dual-stack inet table and key on a
// raw IPv4 network-header source load (offset 12, len 4); on an IPv6 packet that
// offset reads bytes 4..7 of the 16-byte IPv6 source and writes a fake IPv4 key
// into the IPv4-typed meter set. The guard makes the meter rules run on IPv4
// only. Register 1 is reloaded by the expression that follows, so the guard does
// not disturb later register use.
func ipv4NFProtoGuard() []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{2}}, // NFPROTO_IPV4
}
}
// ipv6NFProtoGuard is the NFPROTO==IPV6 sibling of ipv4NFProtoGuard, prepended
// to the IPv6 meter rules so they only run on IPv6 packets.
func ipv6NFProtoGuard() []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{10}}, // NFPROTO_IPV6
}
}
func buildPortAllowExprs(pa PortAllowEntry, ipv6Enabled bool) []expr.Any {
parsed := net.ParseIP(pa.IP)
if parsed == nil || pa.Port <= 0 || pa.Port > 65535 {
return nil
}
proto := byte(6)
if pa.Proto == "udp" {
proto = 17
}
common := func(ip net.IP, offset, length uint32) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: offset, Len: length},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: ip},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(pa.Port))},
&expr.Verdict{Kind: expr.VerdictAccept},
}
}
if ip4 := parsed.To4(); ip4 != nil {
return append(ipv4NFProtoGuard(), common(ip4, 12, 4)...)
}
if !ipv6Enabled {
return nil
}
return append(ipv6NFProtoGuard(), common(parsed.To16(), 8, 16)...)
}
// ipv6Saddr64Key loads the IPv6 source address into register 1 and masks it to
// its /64 prefix. IPv6 hosts routinely own a whole /64 and rotate addresses
// within it for free, so the flood meters key on the /64 rather than the full
// /128, which a rotating source would trivially evade.
func ipv6Saddr64Key() []expr.Any {
return []expr.Any{
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 8, Len: 16},
&expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: 16,
Mask: []byte{
0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff,
0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
},
Xor: make([]byte, 16),
},
}
}
// ipv6MeterSet builds a dynamic IPv6 meter set keyed on the /64 prefix, with the
// same one-minute element timeout as the v4 meters.
func (e *Engine) ipv6MeterSet(name string) *nftables.Set {
set := &nftables.Set{
Table: e.table, Name: name, KeyType: nftables.TypeIP6Addr,
Dynamic: true, HasTimeout: true, Timeout: time.Minute,
}
_ = e.conn.AddSet(set, nil)
return set
}
// synFloodRuleExprs6 is the IPv6 sibling of synFloodRuleExprs: same TCP-SYN
// match, keyed on the /64 source prefix into the v6 meter.
func (e *Engine) synFloodRuleExprs6() []expr.Any {
exprs := append(ipv6NFProtoGuard(), []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}}, // TCP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 13, Len: 1},
&expr.Bitwise{SourceRegister: 1, DestRegister: 1, Len: 1, Mask: []byte{0x12}, Xor: []byte{0x00}},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0x02}}, // SYN only
}...)
exprs = append(exprs, ipv6Saddr64Key()...)
return append(exprs,
&expr.Dynset{
SrcRegKey: 1, SetName: e.meterSYN6.Name, SetID: e.meterSYN6.ID, Operation: 1,
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: 25, Unit: expr.LimitTimeSecond, Burst: 100, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
)
}
// udpFloodRuleExprs6 is the IPv6 sibling of udpFloodRuleExprs.
func (e *Engine) udpFloodRuleExprs6(rate uint64, burst uint32) []expr.Any {
exprs := append(ipv6NFProtoGuard(), []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{17}}, // UDP
}...)
exprs = append(exprs, ipv6Saddr64Key()...)
return append(exprs,
&expr.Dynset{
SrcRegKey: 1, SetName: e.meterUDP6.Name, SetID: e.meterUDP6.ID, Operation: 1,
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: rate, Unit: expr.LimitTimeSecond, Burst: burst, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
)
}
// connMeterRuleExprs6 is the IPv6 sibling of connMeterRuleExprs. The exempt
// lookup matches the full v6 source (exempt ranges are CIDRs); the meter keys on
// the /64 prefix.
func (e *Engine) connMeterRuleExprs6(rate uint64, burst uint32) []expr.Any {
core := []expr.Any{
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitNEW),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(0)},
}
core = append(core, ipv6Saddr64Key()...)
core = append(core,
&expr.Dynset{
SrcRegKey: 1, SetName: e.meterConn6.Name, SetID: e.meterConn6.ID, Operation: 1,
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: rate, Unit: expr.LimitTimeMinute, Burst: burst, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
)
if e.setDOSExempt6 != nil {
core = append(e.dosExemptV6Lookup(1), core...)
}
return append(ipv6NFProtoGuard(), core...)
}
// synFloodRuleExprs builds the expression list for the per-IP SYN flood
// rate-limit rule. The rule is intentionally NOT exempt-aware: SYN flood
// protection targets TCP half-open storms and is applied without a dos_exempt
// guard (unlike connection meters).
func (e *Engine) synFloodRuleExprs() []expr.Any {
return append(ipv4NFProtoGuard(), []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}}, // TCP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 13, Len: 1},
&expr.Bitwise{SourceRegister: 1, DestRegister: 1, Len: 1, Mask: []byte{0x12}, Xor: []byte{0x00}},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{0x02}}, // SYN only
// Load source IP for per-IP metering
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Dynset{
SrcRegKey: 1,
SetName: e.meterSYN.Name,
SetID: e.meterSYN.ID,
Operation: 1, // NFT_DYNSET_OP_UPDATE
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: 25, Unit: expr.LimitTimeSecond, Burst: 100, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
}...)
}
// udpFloodRuleExprs builds the expression list for the per-IP UDP flood
// rate-limit rule. Like synFloodRuleExprs, this rule carries no dos_exempt
// guard.
func (e *Engine) udpFloodRuleExprs(rate uint64, burst uint32) []expr.Any {
return append(ipv4NFProtoGuard(), []expr.Any{
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{17}}, // UDP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Dynset{
SrcRegKey: 1,
SetName: e.meterUDP.Name,
SetID: e.meterUDP.ID,
Operation: 1,
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: rate, Unit: expr.LimitTimeSecond, Burst: burst, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
}...)
}
// connMeterRuleExprs builds the expression list for the per-IP new-connection
// rate-limit rule. When setDOSExempt is non-nil, an inverted source-set lookup
// is prepended so sources in dos_exempt_nets bypass the rule entirely.
//
// Register-reuse safety: dosExemptV4Lookup writes saddr into reg 1. The
// following Ct and Payload exprs reload reg 1 with ct-state and saddr
// respectively, so the exempt check can never interfere with the meter logic.
func (e *Engine) connMeterRuleExprs(rate uint64, burst uint32) []expr.Any {
exprs := []expr.Any{
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitNEW),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(0)},
// Load source IP for per-IP metering
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Dynset{
SrcRegKey: 1,
SetName: e.meterConn.Name,
SetID: e.meterConn.ID,
Operation: 1,
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: rate, Unit: expr.LimitTimeMinute, Burst: burst, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
}
if e.setDOSExempt != nil {
exprs = append(e.dosExemptV4Lookup(1), exprs...)
}
// The IPv4 family guard must precede the exempt lookup, which also loads an
// IPv4 network-header source and would otherwise misread IPv6 packets.
return append(ipv4NFProtoGuard(), exprs...)
}
// connlimitRuleExprs builds the expression list for the per-IP concurrent
// connection limit rule. When setDOSExempt is non-nil, an inverted source-set
// lookup is prepended so sources in dos_exempt_nets bypass the rule.
func (e *Engine) connlimitRuleExprs(limit uint32) []expr.Any {
exprs := []expr.Any{
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitNEW),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(0)},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
&expr.Dynset{
SrcRegKey: 1,
SetName: e.meterConnlim.Name,
SetID: e.meterConnlim.ID,
Operation: 1,
Exprs: []expr.Any{
&expr.Connlimit{Count: limit, Flags: 1}, // 1 = over
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
}
if e.setDOSExempt != nil {
exprs = append(e.dosExemptV4Lookup(1), exprs...)
}
// The IPv4 family guard must precede the exempt lookup, which also loads an
// IPv4 network-header source and would otherwise misread IPv6 packets.
return append(ipv4NFProtoGuard(), exprs...)
}
// addCFWhitelistRule adds an accept rule for Cloudflare IPs on TCP ports 80 and 443.
// Equivalent to: ip saddr @cf_whitelist tcp dport {80, 443} accept
func (e *Engine) addCFWhitelistRule(set *nftables.Set, ipv6 bool) {
if set == nil {
return
}
for _, port := range []uint16{80, 443} {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: cfWhitelistRuleExprs(set, ipv6, port),
})
}
}
// cfWhitelistRuleExprs builds the expression list for one Cloudflare whitelist
// accept rule: <family> saddr @set tcp dport <port> accept. The NFPROTO guard
// must precede the raw network-header source load: without it the v4 load also
// runs on IPv6 packets, where offset 12 reads bytes 4..7 of the IPv6 source,
// so a colliding IPv6 address would be accepted ahead of the block rules.
func cfWhitelistRuleExprs(set *nftables.Set, ipv6 bool, port uint16) []expr.Any {
var exprs []expr.Any
if ipv6 {
exprs = append(ipv6NFProtoGuard(),
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 8, Len: 16},
)
} else {
exprs = append(ipv4NFProtoGuard(),
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: 12, Len: 4},
)
}
return append(exprs,
&expr.Lookup{SourceRegister: 1, SetName: set.Name, SetID: set.ID},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}}, // TCP
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(port)},
&expr.Verdict{Kind: expr.VerdictAccept},
)
}
func buildFamilyPortRuleExprs(nfproto byte, port int, tcp bool) []expr.Any {
proto := byte(6) // TCP
if !tcp {
proto = 17 // UDP
}
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{nfproto}},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}},
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(port))},
&expr.Verdict{Kind: expr.VerdictAccept},
}
}
// buildOutAllowExprs builds one destination-scoped outbound accept rule:
// nfproto guard, TCP, destination address, destination port range, accept.
//
// The neighbouring buildPortAllowExprs matches the SOURCE address because it
// serves the input chain. This rule serves the output chain and matches the
// DESTINATION, at offset 16 for IPv4 and 24 for IPv6.
//
// Returns nil when the rule cannot be expressed -- an unparseable destination,
// or an IPv6 destination while dual-stack filtering is off. Callers must emit
// nothing rather than a rule matching more than the operator asked for.
func buildOutAllowExprs(r OutAllowRule, ipv6Enabled bool) []expr.Any {
network, err := ParseOutAllowDst(r.Dst)
if err != nil {
return nil
}
if !validOutAllowRange(r) {
return nil
}
var (
exprs []expr.Any
addr net.IP
mask net.IPMask
offset uint32
length uint32
)
if ip4 := network.IP.To4(); ip4 != nil {
exprs, addr, mask, offset, length = ipv4NFProtoGuard(), ip4, network.Mask, 16, 4
} else {
if !ipv6Enabled {
return nil
}
exprs, addr, mask, offset, length = ipv6NFProtoGuard(), network.IP.To16(), network.Mask, 24, 16
}
exprs = append(exprs,
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{6}}, // TCP
)
// A zero-length prefix means every destination, and the honest encoding of
// that is no address match at all rather than a compare against nothing.
if ones, _ := mask.Size(); ones > 0 {
exprs = append(exprs, &expr.Payload{
DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: offset, Len: length,
})
if ones < int(length)*8 {
exprs = append(exprs, &expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: length,
Mask: mask, Xor: make([]byte, length),
})
}
exprs = append(exprs, &expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: addr.Mask(mask)})
}
return append(exprs,
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpGte, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(r.PortStart))},
&expr.Cmp{Op: expr.CmpOpLte, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(r.PortEnd))},
&expr.Verdict{Kind: expr.VerdictAccept},
)
}
// validOutAllowRange mirrors the config validator so a rule that slipped past
// validation still emits nothing rather than a range that matches everything.
func validOutAllowRange(r OutAllowRule) bool {
return r.PortStart >= 1 && r.PortStart <= 65535 &&
r.PortEnd >= 1 && r.PortEnd <= 65535 &&
r.PortStart <= r.PortEnd
}
func familyBypassRuleExprs(nfproto byte) []expr.Any {
return []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{nfproto}},
&expr.Verdict{Kind: expr.VerdictAccept},
}
}
func (e *Engine) addPortAcceptRule(port int, tcp bool, nfproto byte) {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: buildFamilyPortRuleExprs(nfproto, port, tcp),
})
}
func (e *Engine) addPortRangeAcceptRule(startPort, endPort int, tcp bool, nfproto byte) {
proto := byte(6)
if !tcp {
proto = 17
}
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainIn,
Exprs: []expr.Any{
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{nfproto}},
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}},
// Load dest port once, check range
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpGte, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(startPort))},
&expr.Cmp{Op: expr.CmpOpLte, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(endPort))},
&expr.Verdict{Kind: expr.VerdictAccept},
},
})
}
func (e *Engine) addOutboundPortRule(port int, tcp bool, nfproto byte) {
e.conn.AddRule(&nftables.Rule{
Table: e.table,
Chain: e.chainOut,
Exprs: buildFamilyPortRuleExprs(nfproto, port, tcp),
})
}
// --- Public API ---
// BlockIP adds an IP to the blocked set with optional timeout.
// timeout 0 = permanent block.
//
// Thin wrapper over BlockIPOutcome that discards the outcome. Existing
// callers that only need success/error semantics keep working; auto-
// response callers should use BlockIPOutcome so they can suppress local
// side effects (state mutation, AUTO-BLOCK alert) when the kernel was
// not actually touched.
func (e *Engine) BlockIP(ip string, reason string, timeout time.Duration) error {
_, err := e.BlockIPOutcome(ip, reason, timeout)
return err
}
// BlockIPOutcome is the AUTO-RESPONSE entry point. It performs the same
// guards, verdict-callback consultation, and dry-run gating as BlockIP,
// but additionally reports which path was taken via BlockOutcome so the
// caller can decide whether to record local state. See the BlockOutcome
// godoc for the meaning of each return value.
//
// Operator-initiated commands (csm firewall block, Web UI manual block) must
// call BlockIPForce instead, which skips the dry-run gate unconditionally.
func (e *Engine) BlockIPOutcome(ip string, reason string, timeout time.Duration) (outcome BlockOutcome, resultErr error) {
return e.BlockIPOutcomeWithFindingID(ip, reason, timeout, "")
}
// BlockIPOutcomeWithFindingID preserves the originating audit identity across
// every automatic outcome without changing the block policy or result.
func (e *Engine) BlockIPOutcomeWithFindingID(ip, reason string, timeout time.Duration, findingID string) (outcome BlockOutcome, resultErr error) {
return e.BlockIPRequest(ActionRequest{Operation: "block", Target: ip, Reason: reason, TTL: timeout, FindingID: findingID, Automatic: true, Actor: string(actionlog.Daemon)}, nil)
}
func (e *Engine) blockIPOutcomeRequest(req ActionRequest, budget *ScanAdmission) (outcome BlockOutcome, resultErr error) {
ip, reason, timeout, findingID := req.Target, req.Reason, req.TTL, req.FindingID
defer func() {
if e.shouldLegacyOutcome(resultErr) || (resultErr == nil && outcome != BlockOutcomeLive) {
recordBlockOutcome(ip, reason, timeout, outcome, resultErr, false, findingID)
}
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return BlockOutcomeNoop, err
}
ip = canonical
// Soft-allow gate: an automatic block must never re-add an operator
// full-IP/port allow or a verified-bot range to blocked_ips. The nftables
// input chain drops @blocked_ips before it accepts operator allows, so the
// only safe fix is to keep the IP out of the blocked set. Checked before
// the local validation guards so a soft-allowed IP is skipped cleanly even
// when the deny limit is full. Operator `firewall deny` uses BlockIPForce,
// which never reaches this path, so an explicit deny still overrides a
// soft-allow.
if e.autoBlockSoftAllowed(ip) {
e.logAutoBlockSoftAllowed(ip)
return BlockOutcomeAllowlisted, nil
}
// Local safety checks always run before consulting the external callback.
// The callback can downgrade a block decision, but it cannot bypass
// malformed-IP, IPv6-disabled, infra-IP, or block-limit guards.
alreadyBlocked, err := e.validateBlockIP(ip, timeout, true)
if err != nil {
return BlockOutcomeNoop, err
}
if alreadyBlocked {
return BlockOutcomeNoop, nil
}
// Verdict gate: consult the panel after local validation and before the
// dry-run gate so that an "allow" verdict short-circuits everything.
// Fail-open: errors proceed with the default block. Nil callback skips
// the gate entirely.
if asker := e.verdictAskerFn(); asker != nil {
v, tenant, note, verr := asker(e.verdictContext(), ip, reason)
switch {
case verr != nil:
fmt.Fprintf(os.Stderr, "[%s] verdict callback failed for %s: %v - proceeding with default block\n",
time.Now().Format("2006-01-02 15:04:05"), ip, verr)
case v == "allow":
fmt.Fprintf(os.Stderr, "[%s] verdict callback returned allow for %s (tenant=%q note=%q) - not blocking\n",
time.Now().Format("2006-01-02 15:04:05"), ip, tenant, note)
return BlockOutcomeAllowed, nil
case tenant != "" || note != "":
fmt.Fprintf(os.Stderr, "[%s] verdict callback returned block for %s (tenant=%q note=%q) - proceeding with default block\n",
time.Now().Format("2006-01-02 15:04:05"), ip, tenant, note)
}
// "block" / empty / error -> proceed with default flow.
}
if e.autoBlockSoftAllowed(ip) {
e.logAutoBlockSoftAllowed(ip)
return BlockOutcomeAllowlisted, nil
}
// Dry-run gate: the daemon callback reads the current daemon config at
// call time so a SIGHUP takes effect without a daemon restart. Nil
// callback means live.
if e.autoResponseDryRunEnabled() {
fmt.Fprintf(os.Stderr, "[%s] auto_response dry_run: would have blocked %s (%s)\n",
time.Now().Format("2006-01-02 15:04:05"), ip, reason)
e.recordDryRunBlock(ip, reason, timeout)
return BlockOutcomeDryRun, nil
}
if e.autoBlockVerifiedRange(ip) {
e.logAutoBlockSoftAllowed(ip)
return BlockOutcomeAllowlisted, nil
}
lockedOutcome, err := e.blockIPLockedRequest(ip, reason, timeout, true, true, req, budget)
if err != nil {
return lockedOutcome, err
}
if lockedOutcome == BlockOutcomeAllowlisted {
e.logAutoBlockSoftAllowed(ip)
return BlockOutcomeAllowlisted, nil
}
return lockedOutcome, nil
}
func (e *Engine) autoResponseDryRunEnabled() bool {
e.mu.Lock()
fn := e.dryRunEnabled
e.mu.Unlock()
return fn != nil && fn()
}
// BlockIPForce adds an IP to the blocked set unconditionally, bypassing the
// auto_response.dry_run gate. Use this for operator-initiated commands (CLI,
// Web UI manual block) where the operator has explicitly decided to block.
func (e *Engine) BlockIPForce(ip string, reason string, timeout time.Duration) (resultErr error) {
defer func() {
if e.shouldLegacyOutcome(resultErr) {
recordBlockOutcome(ip, reason, timeout, BlockOutcomeLive, resultErr, true, "")
}
}()
return e.blockIPLocked(ip, reason, timeout, false)
}
// BlockIPForcePreserveLifetime is the Web UI timed-block path. The guard
// and replacement share the engine lock, including concurrent CLI promotions.
func (e *Engine) BlockIPForcePreserveLifetime(ip, reason string, timeout time.Duration) (resultErr error) {
defer func() {
if e.shouldLegacyOutcome(resultErr) {
recordBlockOutcome(ip, reason, timeout, BlockOutcomeLive, resultErr, true, "")
}
}()
_, err := e.blockIPLockedRequestGuarded(ip, reason, timeout, false, false,
ActionRequest{Operation: "block", Target: ip, Reason: reason, TTL: timeout}, nil, true)
return err
}
// PromoteToPermanentBlock upgrades an existing temporary block on ip to a
// permanent one: it clears the kernel timeout by deleting the timed element
// and re-adding it without a timeout, and zeroes ExpiresAt in state. The
// ordinary block path cannot do this during PermBlock escalation because it
// skips an already-blocked IP, so the kernel timeout would otherwise expire
// the block the operator wanted made permanent. Returns an error if the IP is
// not currently blocked (nothing to promote).
func (e *Engine) PromoteToPermanentBlock(ip, reason string) (resultErr error) {
return e.PromoteToPermanentBlockWithFindingID(ip, reason, "")
}
// PromoteToPermanentBlockWithFindingID ties escalation to its triggering
// observation while preserving the existing permanent-block transaction.
func (e *Engine) PromoteToPermanentBlockWithFindingID(ip, reason, findingID string) (resultErr error) {
source := InferProvenance("permblock", reason)
defer func() {
if e.shouldLegacyOutcome(resultErr) {
recordFirewallFindingResult("permblock", ip, reason, source, 0, actionlog.Applied, resultErr, findingID)
}
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
// This path deliberately bypasses blockIPLocked, so it does not inherit
// that path's address guard. An entry that reached the set before the
// guard existed must not be made permanent here.
if isUnblockableAddress(ip) {
return ipProtectedErrorf("refusing to block non-routable address: %s", ip)
}
e.mu.Lock()
defer e.mu.Unlock()
targetSet, key, err := e.resolveIPSet(ip, e.setBlocked, e.setBlocked6)
if err != nil {
return err
}
priorState := e.loadStateFile()
entry := BlockedEntry{IP: ip, Reason: reason, Source: SourceSystem, BlockedAt: time.Now()}
found := false
wasTemporary := false
for _, b := range priorState.Blocked {
if sameIPString(b.IP, ip) {
found = true
wasTemporary = !b.ExpiresAt.IsZero()
entry.BlockedAt = b.BlockedAt
if b.Source != "" {
entry.Source = b.Source
}
break
}
}
if !found {
return fmt.Errorf("cannot promote %s: not currently blocked", ip)
}
if e.cfg != nil && wasTemporary && e.cfg.DenyIPLimit > 0 {
perm, _, ok := e.livePermTempCountsLocked(priorState)
if !ok {
perm = countPermanentBlockedEntries(priorState)
}
if perm >= e.cfg.DenyIPLimit {
return fmt.Errorf("permanent deny limit reached (%d)", e.cfg.DenyIPLimit)
}
}
// ExpiresAt left zero: permanent.
nextState := copyFirewallState(priorState)
upsertBlockedEntryInState(&nextState, entry)
if e.lifecycle != nil {
return e.runDurableLocked(ActionRequest{Operation: "promote", Target: ip, Reason: reason, FindingID: findingID, Automatic: findingID != ""}, nil, nextState)
}
if err := e.persistFirewallIntent(priorState, nextState); err != nil {
return fmt.Errorf("persisting permanent promotion for %s: %w", ip, err)
}
// Delete the timed element and re-add it without a timeout in one
// transaction, so the address is never unblocked in between.
if err := e.conn.SetDeleteElements(targetSet, []nftables.SetElement{{Key: key}}); err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return fmt.Errorf("promoting %s: delete timed element: %w (state restore failed: %v)", ip, err, restoreErr)
}
return fmt.Errorf("promoting %s: delete timed element: %w", ip, err)
}
if err := e.conn.SetAddElements(targetSet, []nftables.SetElement{{Key: key}}); err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return fmt.Errorf("promoting %s: re-add permanent element: %w (state restore failed: %v)", ip, err, restoreErr)
}
return fmt.Errorf("promoting %s: re-add permanent element: %w", ip, err)
}
if err := e.conn.Flush(); err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return fmt.Errorf("promoting %s: flush: %w (state restore failed: %v)", ip, err, restoreErr)
}
return fmt.Errorf("promoting %s: flush: %w", ip, err)
}
source = entry.Source
e.legacyFileAuditLocked("permblock", ip, reason, entry.Source, 0)
return nil
}
// recordDryRunBlock persists a dry-run record through the daemon-installed
// recorder so operators can review the count via /api/v1/status.
// No-op when no recorder is installed.
func (e *Engine) recordDryRunBlock(ip, reason string, timeout time.Duration) {
e.mu.Lock()
recorder := e.dryRunRecorder
e.mu.Unlock()
if recorder != nil {
recorder(ip, reason, timeout)
}
}
// blockIPLocked is the real implementation called by both BlockIP and BlockIPForce.
func (e *Engine) blockIPLocked(ip string, reason string, timeout time.Duration, skipExisting bool) error {
_, err := e.blockIPLockedMaybeSoftAllowed(ip, reason, timeout, skipExisting, false)
return err
}
func (e *Engine) blockIPLockedMaybeSoftAllowed(ip string, reason string, timeout time.Duration, skipExisting bool, enforceSoftAllow bool) (BlockOutcome, error) {
return e.blockIPLockedRequest(ip, reason, timeout, skipExisting, enforceSoftAllow, ActionRequest{Operation: "block", Target: ip, Reason: reason, TTL: timeout, Automatic: enforceSoftAllow}, nil)
}
func (e *Engine) blockIPLockedRequest(ip string, reason string, timeout time.Duration, skipExisting bool, enforceSoftAllow bool, req ActionRequest, budget *ScanAdmission) (BlockOutcome, error) {
return e.blockIPLockedRequestGuarded(ip, reason, timeout, skipExisting, enforceSoftAllow, req, budget, false)
}
func (e *Engine) blockIPLockedRequestGuarded(ip string, reason string, timeout time.Duration, skipExisting bool, enforceSoftAllow bool, req ActionRequest, budget *ScanAdmission, preserveLifetime bool) (BlockOutcome, error) {
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return BlockOutcomeNoop, err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
return e.blockIPRequestLocked(ip, reason, timeout, skipExisting, enforceSoftAllow, req, budget, preserveLifetime)
}
// blockIPRequestLocked runs with e.mu held, including snapshot-checked undo.
func (e *Engine) blockIPRequestLocked(ip string, reason string, timeout time.Duration, skipExisting bool, enforceSoftAllow bool, req ActionRequest, budget *ScanAdmission, preserveLifetime bool) (BlockOutcome, error) {
req.Target = ip
req = normalizeActionRequest(req)
if found, replayErr := e.replayActionLocked(req); found || replayErr != nil {
return BlockOutcomeNoop, replayErr
}
if readyErr := e.lifecycleReadyLocked(); readyErr != nil {
return BlockOutcomeNoop, readyErr
}
priorState := e.loadStateFile()
if e.lifecycle != nil && e.stateReadErr != nil {
return BlockOutcomeNoop, e.stateReadErr
}
if preserveLifetime && timeout > 0 {
if e.stateReadErr != nil {
return BlockOutcomeNoop, e.stateReadErr
}
if entry, found := blockedStateEntry(priorState, ip); found {
if entry.ExpiresAt.IsZero() {
return BlockOutcomeNoop, ErrPermanentBlock
}
if time.Until(entry.ExpiresAt) > timeout {
return BlockOutcomeNoop, ErrLongerBlock
}
} else {
set, key, err := e.resolveIPSet(ip, e.setBlocked, e.setBlocked6)
if err != nil {
return BlockOutcomeNoop, err
}
live, temporary, classified := e.liveBlockElementKindLocked(set, key)
if !classified {
return BlockOutcomeNoop, fmt.Errorf("cannot verify existing block lifetime for %s", ip)
}
if live && !temporary {
return BlockOutcomeNoop, ErrPermanentBlock
}
}
}
// Safety, capacity, eviction and the resulting state share one snapshot.
if enforceSoftAllow && (ipSetIndexContains(e.allowedIPIndex, ip) || ipSetIndexContains(e.portAllowedIndex, ip)) {
return BlockOutcomeAllowlisted, nil
}
targetSet, key, alreadyBlocked, evictTempIP, err := e.blockIPTargetFromState(ip, timeout, skipExisting, priorState)
if err != nil {
return BlockOutcomeNoop, err
}
if alreadyBlocked {
return BlockOutcomeNoop, nil
}
// Persist to state BEFORE adding the kernel element. state.json is the
// seed source on the next Apply, so a crash after this point but before the
// kernel add leaves a durable record that Apply re-applies. The previous
// ordering (kernel first, then state) left a window where a process kill
// produced a permanent kernel block with no state row: it never expired
// (timeout 0) yet a later state-seeded Apply would silently drop it.
// Zero ExpiresAt means permanent.
entry := BlockedEntry{
IP: ip,
Reason: reason,
Source: InferProvenance("block", reason),
BlockedAt: time.Now(),
}
if timeout > 0 {
entry.ExpiresAt = time.Now().Add(timeout)
}
var evictSet *nftables.Set
var evictKey []byte
if evictTempIP != "" {
evictSet, evictKey, err = e.resolveIPSet(evictTempIP, e.setBlocked, e.setBlocked6)
if err != nil {
return BlockOutcomeNoop, fmt.Errorf("resolving temp block eviction target %s: %w", evictTempIP, err)
}
}
// A forced block over an address that is already blocked changes the
// timeout (deny over an auto-block, tempban over a permanent deny).
// nf_tables treats NEWSETELEM without NLM_F_EXCL on an existing key as
// an acknowledged no-op that keeps the old timeout, so the element must be
// deleted and re-added in the same batch, as PromoteToPermanentBlock does;
// otherwise state.json and every CSM surface disagree with the kernel.
replaceExisting := !skipExisting && firewallStateHasBlocked(priorState, ip)
if !skipExisting && !replaceExisting {
// State is normally written before the kernel batch, but older builds
// and partial recovery can leave a live kernel element without a state
// row. A plain add would be acknowledged without replacing its timeout,
// recreating the state/kernel drift this path is meant to repair.
live, liveErr := e.isBlockedLiveLocked(ip)
if liveErr != nil {
return BlockOutcomeNoop, fmt.Errorf("checking existing block for %s: %w", ip, liveErr)
}
replaceExisting = live
}
nextState := copyFirewallState(priorState)
if evictTempIP != "" {
removeBlockedIPFromState(&nextState, evictTempIP)
}
upsertBlockedEntryInState(&nextState, entry)
if e.lifecycle != nil {
if err := e.runDurableLocked(req, budget, nextState); err != nil {
if errors.Is(err, ErrActionAuditPending) {
return BlockOutcomeLive, err
}
return BlockOutcomeNoop, err
}
return BlockOutcomeLive, nil
}
if err := e.persistFirewallIntent(priorState, nextState); err != nil {
return BlockOutcomeNoop, fmt.Errorf("persisting block for %s: %w", ip, err)
}
elem := []nftables.SetElement{{Key: key, Timeout: timeout}}
if replaceExisting {
if err := e.conn.SetDeleteElements(targetSet, []nftables.SetElement{{Key: key}}); err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return BlockOutcomeNoop, fmt.Errorf("replacing blocked element: %w (state restore failed: %v)", err, restoreErr)
}
return BlockOutcomeNoop, fmt.Errorf("replacing blocked element: %w", err)
}
}
if err := e.conn.SetAddElements(targetSet, elem); err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return BlockOutcomeNoop, fmt.Errorf("adding to blocked set: %w (state restore failed: %v)", err, restoreErr)
}
return BlockOutcomeNoop, fmt.Errorf("adding to blocked set: %w", err)
}
if evictTempIP != "" {
if err := e.conn.SetDeleteElements(evictSet, []nftables.SetElement{{Key: evictKey}}); err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return BlockOutcomeNoop, fmt.Errorf("evicting temp block %s: %w (state restore failed: %v)", evictTempIP, err, restoreErr)
}
return BlockOutcomeNoop, fmt.Errorf("evicting temp block %s: %w", evictTempIP, err)
}
}
if err := e.conn.Flush(); err != nil {
if (replaceExisting || evictSet != nil) && isNftNotFound(err) {
// state.json said blocked but the kernel had already expired the
// replacement or eviction element, so the whole batch was rejected
// on that delete. Retry without whichever stale delete failed.
err = e.retryBlockAddAfterMissingElement(targetSet, elem, evictSet, evictKey, replaceExisting)
}
if err != nil {
if restoreErr := e.restoreBlockStateAfterFailureLocked(priorState, ip); restoreErr != nil {
return BlockOutcomeNoop, fmt.Errorf("flushing: %w (state restore failed: %v)", err, restoreErr)
}
return BlockOutcomeNoop, fmt.Errorf("flushing: %w", err)
}
}
if evictTempIP != "" {
e.legacyActionAuditLocked("evict_temp", evictTempIP, "temp deny limit reached; evicted soonest-expiring entry", SourceSystem, 0)
}
e.legacyFileAuditLocked("block", ip, reason, entry.Source, timeout)
return BlockOutcomeLive, nil
}
// retryBlockAddAfterMissingElement re-queues a block batch without the
// delete of an element the kernel no longer holds. Caller must hold e.mu.
func (e *Engine) retryBlockAddAfterMissingElement(targetSet *nftables.Set, elem []nftables.SetElement, evictSet *nftables.Set, evictKey []byte, replaceExisting bool) error {
queueAdd := func() error {
if err := e.conn.SetAddElements(targetSet, elem); err != nil {
return fmt.Errorf("retry adding to blocked set: %w", err)
}
return nil
}
flush := func(stage string) error {
if err := e.conn.Flush(); err != nil {
return fmt.Errorf("%s: %w", stage, err)
}
return nil
}
if !replaceExisting {
// The only delete in the rejected batch was the expired eviction
// victim. Re-add the new target without repeating that stale delete.
if err := queueAdd(); err != nil {
return err
}
return flush("retry flushing block without eviction")
}
// First assume the replacement target expired but the eviction victim is
// still live. Adding the target and evicting the victim remains atomic.
if err := queueAdd(); err != nil {
return err
}
if evictSet != nil {
if err := e.conn.SetDeleteElements(evictSet, []nftables.SetElement{{Key: evictKey}}); err != nil {
return fmt.Errorf("retry evicting temp block: %w", err)
}
}
if err := flush("retry flushing block after missing replacement"); err == nil {
return nil
} else if !isNftNotFound(err) {
return err
}
// The eviction victim was the stale element instead. Replace the target
// without repeating the eviction delete, keeping a live target blocked for
// the whole atomic transaction.
if err := e.conn.SetDeleteElements(targetSet, []nftables.SetElement{{Key: elem[0].Key}}); err != nil {
return fmt.Errorf("retry deleting replacement target: %w", err)
}
if err := queueAdd(); err != nil {
return err
}
if err := flush("retry flushing block without eviction"); err == nil {
return nil
} else if !isNftNotFound(err) {
return err
}
// Both deletes raced expired elements. No live target remains to preserve,
// so an add-only transaction is the final bounded retry.
if err := queueAdd(); err != nil {
return err
}
return flush("retry flushing add-only block")
}
func (e *Engine) validateBlockIP(ip string, timeout time.Duration, skipExisting bool) (bool, error) {
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return false, err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
_, _, alreadyBlocked, _, err := e.blockIPTarget(ip, timeout, skipExisting)
return alreadyBlocked, err
}
func (e *Engine) blockIPTarget(ip string, timeout time.Duration, skipExisting bool) (*nftables.Set, []byte, bool, string, error) {
state := e.loadStateFile()
if e.lifecycle != nil && e.stateReadErr != nil {
return nil, nil, false, "", e.stateReadErr
}
return e.blockIPTargetFromState(ip, timeout, skipExisting, state)
}
func (e *Engine) blockIPTargetFromState(ip string, timeout time.Duration, skipExisting bool, st FirewallState) (*nftables.Set, []byte, bool, string, error) {
parsed := net.ParseIP(ip)
if parsed == nil {
return nil, nil, false, "", fmt.Errorf("invalid IP: %s", ip)
}
ip = parsed.String()
// SAFETY: never block infra IPs - prevents admin lockout.
// Runs before resolveIPSet so that an IPv6-disabled config (set6 == nil)
// cannot bypass the infra guard when the caller passes the canonical
// IPv6 form of a listed infra address.
for _, cidr := range e.cfg.InfraIPs {
_, network, cidrErr := net.ParseCIDR(cidr)
if cidrErr != nil {
if infraIP := net.ParseIP(cidr); infraIP != nil && infraIP.String() == parsed.String() {
return nil, nil, false, "", ipProtectedErrorf("refusing to block infra IP: %s", ip)
}
continue
}
if network.Contains(parsed) {
return nil, nil, false, "", ipProtectedErrorf("refusing to block infra IP: %s (in %s)", ip, cidr)
}
}
// Hostnames in cfg.InfraIPs are resolved by the DynDNS loop and
// pushed in via UpdateInfraResolved. Without this check a hostname
// listed as infra would only be honoured when the operator also
// pinned the IP, so a moving panel IP would silently drop out of
// the lockout guard.
if host, ok := e.infraIPResolvedHostLocked(ip); ok {
return nil, nil, false, "", ipProtectedErrorf("refusing to block infra IP: %s (resolved from %s)", ip, host)
}
// Daemon's own interface addresses are always off-limits. Without
// this guard a stray request from the host to itself (cron, panel
// callback, internal probe) could trigger an auto-block that
// firewalls every customer hosted on the same IP.
if e.isLocalAddrLocked(ip) {
return nil, nil, false, "", ipProtectedErrorf("refusing to block local host IP: %s (own interface address)", ip)
}
// Loopback and link-local are excluded from the interface set above, so
// they need their own refusal: the panel's own proxied requests appear as
// loopback, and a WAF denial count for it must not become a block.
if isUnblockableAddress(ip) {
return nil, nil, false, "", ipProtectedErrorf("refusing to block non-routable address: %s", ip)
}
targetSet, key, err := e.resolveIPSet(ip, e.setBlocked, e.setBlocked6)
if err != nil {
return nil, nil, false, "", err
}
cachedBlockMissingLive := false
if skipExisting && firewallStateHasBlocked(st, ip) {
liveBlocked, liveErr := e.isBlockedLiveLocked(ip)
if liveErr != nil || liveBlocked {
// Treat a probe error as "still blocked" so we never demote a
// cached block on transient netlink trouble. Returning nil
// here is intentional and the conservative posture.
return targetSet, key, true, "", nil //nolint:nilerr // intentional fail-safe on netlink probe error
}
cachedBlockMissingLive = true
}
// Enforce deny IP limits. Prefer counts from the live nft set
// so an entry the kernel already expired no longer counts
// against the cap. Fall back to the cached state.json count
// (existing behaviour) if the live query is unavailable.
if e.cfg.DenyIPLimit > 0 || e.cfg.DenyTempIPLimit > 0 {
perm, temp, ok := e.livePermTempCountsLocked(st)
stateEntry, replacingStateEntry := blockedStateEntry(st, ip)
if ok && !skipExisting {
if replacingStateEntry {
// Replacing an element already occupying a limit slot is not a new
// block. Only subtract it when the kernel confirms the state entry
// is still live; stale state must not hide some other live element.
if live, liveErr := e.isBlockedLiveLocked(ip); liveErr == nil && live {
if stateEntry.ExpiresAt.IsZero() && perm > 0 {
perm--
} else if !stateEntry.ExpiresAt.IsZero() && temp > 0 {
temp--
}
}
} else if live, temporary, classified := e.liveBlockElementKindLocked(targetSet, key); classified && live {
// Recovery may leave a kernel element with no state row. It still
// occupies one of the counted slots, so replacing it must release
// that same slot before enforcing the limit.
if temporary && temp > 0 {
temp--
} else if !temporary && perm > 0 {
perm--
}
}
}
if !ok {
excludeIP := ""
if cachedBlockMissingLive || (!skipExisting && replacingStateEntry) {
excludeIP = ip
}
perm, temp = blockedStatePermTempCounts(st, excludeIP)
}
if timeout == 0 && e.cfg.DenyIPLimit > 0 && perm >= e.cfg.DenyIPLimit {
return nil, nil, false, "", fmt.Errorf("permanent deny limit reached (%d)", e.cfg.DenyIPLimit)
}
if timeout > 0 && e.cfg.DenyTempIPLimit > 0 && temp >= e.cfg.DenyTempIPLimit {
// Validation must stay read-only because verdict and dry-run gates run
// after it. The live block path applies this eviction target.
victim, ok := soonestExpiringTempIP(st, ip)
if !ok {
return nil, nil, false, "", fmt.Errorf("temporary deny limit reached (%d) and no temp entry to evict", e.cfg.DenyTempIPLimit)
}
if _, _, err := e.resolveIPSet(victim, e.setBlocked, e.setBlocked6); err != nil {
return nil, nil, false, "", fmt.Errorf("temporary deny limit reached (%d) and eviction target %s is unusable: %w", e.cfg.DenyTempIPLimit, victim, err)
}
return targetSet, key, false, victim, nil
}
}
return targetSet, key, false, "", nil
}
// soonestExpiringTempIP returns the IP of the temporary block closest to
// expiry, skipping permanent blocks and excludeIP. Pure helper so the
// eviction policy is unit-testable without nftables.
func soonestExpiringTempIP(st FirewallState, excludeIP string) (string, bool) {
var best string
var bestExp time.Time
found := false
for _, b := range st.Blocked {
if sameIPString(b.IP, excludeIP) || b.ExpiresAt.IsZero() {
continue
}
if !found || b.ExpiresAt.Before(bestExp) {
best, bestExp, found = b.IP, b.ExpiresAt, true
}
}
return best, found
}
func firewallStateHasBlocked(state FirewallState, ip string) bool {
_, ok := blockedStateEntry(state, ip)
return ok
}
func blockedStateEntry(state FirewallState, ip string) (BlockedEntry, bool) {
for _, entry := range state.Blocked {
if sameIPString(entry.IP, ip) {
return entry, true
}
}
return BlockedEntry{}, false
}
func countPermanentBlockedEntries(state FirewallState) int {
perm, _ := blockedStatePermTempCounts(state, "")
return perm
}
func blockedStatePermTempCounts(state FirewallState, excludeIP string) (perm, temp int) {
seen := make(map[string]bool, len(state.Blocked))
for _, entry := range state.Blocked {
key, ok := canonicalIPKey(entry.IP)
if !ok {
continue
}
if excludeIP != "" && sameIPString(key, excludeIP) {
continue
}
isTemp := !entry.ExpiresAt.IsZero()
if priorTemp, exists := seen[key]; exists {
if priorTemp && !isTemp {
seen[key] = false
}
continue
}
seen[key] = isTemp
}
for _, isTemp := range seen {
if isTemp {
temp++
} else {
perm++
}
}
return perm, temp
}
// UnblockIP removes an IP from the blocked set and state.
func (e *Engine) UnblockIP(ip string) (resultErr error) {
defer func() { e.legacyFirewallFailure("unblock", ip, "", "", 0, resultErr) }()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
return e.unblockIPLocked(ip)
}
func (e *Engine) unblockIPLocked(ip string) error {
targetSet, key, err := e.resolveIPSet(ip, e.setBlocked, e.setBlocked6)
if err != nil {
return err
}
// Remove from state BEFORE the kernel delete: a crash between the two
// must converge to the operator's intent (unblocked) on the next Apply,
// not silently re-block the IP from a stale state row. On kernel failure
// the prior state is restored so the entry stays visible for a retry.
priorState := e.loadStateFile()
nextState := copyFirewallState(priorState)
removeBlockedIPFromState(&nextState, ip)
if e.lifecycle != nil {
return e.runDurableLocked(ActionRequest{Operation: "unblock", Target: ip}, nil, nextState)
}
if err := e.persistFirewallIntent(priorState, nextState); err != nil {
return fmt.Errorf("persisting unblock for %s: %w", ip, err)
}
if err := e.conn.SetDeleteElements(targetSet, []nftables.SetElement{{Key: key}}); err != nil {
if restoreErr := e.saveState(&priorState); restoreErr != nil {
return fmt.Errorf("removing from blocked set: %w (state restore failed: %v)", err, restoreErr)
}
return fmt.Errorf("removing from blocked set: %w", err)
}
if err := e.conn.Flush(); err != nil {
if isNftNotFound(err) {
// The delete target already disappeared from nft, so the
// persisted unblock is the only remaining state to keep.
e.legacyActionAuditLocked("unblock", ip, "", "", 0)
return nil
}
if restoreErr := e.saveState(&priorState); restoreErr != nil {
return fmt.Errorf("flushing: %w (state restore failed: %v)", err, restoreErr)
}
return fmt.Errorf("flushing: %w", err)
}
e.legacyActionAuditLocked("unblock", ip, "", "", 0)
return nil
}
// IsBlocked returns true if the IP is currently in the engine's blocked state.
// Uses the persisted state file (which is cleaned of expired entries on load).
//
// The lookup is O(1) via the blockedIPIndex map populated from the
// cached state. Linear scans over the parsed slice are gone -- on
// hosts with hundreds of persisted blocks the scan was the dominant
// cost of every connection-handler IsBlocked check.
func (e *Engine) IsBlocked(ip string) bool {
e.mu.Lock()
defer e.mu.Unlock()
e.ensureStateCacheLocked()
if _, ok := e.blockedIPIndex[ip]; ok {
return true
}
// The index is keyed by canonical form (blockIPLocked stores
// net.ParseIP(ip).String()). A caller passing a non-canonical form -- an
// IPv4-mapped IPv6 like "::ffff:1.2.3.4" from a dual-stack listener --
// would otherwise miss the block. Retry with the canonical form.
if parsed := net.ParseIP(ip); parsed != nil {
if canon := parsed.String(); canon != ip {
_, ok := e.blockedIPIndex[canon]
return ok
}
}
return false
}
// IsAllowed reports whether ip is on the operator allowed_ips set, read from
// the in-memory cache built from state.json. Mirrors IsBlocked, including the
// canonical-form retry for IPv4-mapped IPv6 callers.
func (e *Engine) IsAllowed(ip string) bool {
e.mu.Lock()
defer e.mu.Unlock()
e.ensureStateCacheLocked()
return ipSetIndexContains(e.allowedIPIndex, ip)
}
func (e *Engine) operatorSoftAllowed(ip string) bool {
e.mu.Lock()
defer e.mu.Unlock()
return e.operatorSoftAllowedLocked(ip)
}
func (e *Engine) operatorSoftAllowedLocked(ip string) bool {
e.ensureStateCacheLocked()
return ipSetIndexContains(e.allowedIPIndex, ip) || ipSetIndexContains(e.portAllowedIndex, ip)
}
func ipSetIndexContains(index map[string]struct{}, ip string) bool {
if _, ok := index[ip]; ok {
return true
}
if parsed := net.ParseIP(ip); parsed != nil {
if canon := parsed.String(); canon != ip {
_, ok := index[canon]
return ok
}
}
return false
}
// IsBlockedLive queries the live nftables set, not the in-memory cache built
// from state.json. The cache can drift from the kernel when nft auto-expires
// entries faster than CSM rewrites state.json, or when an out-of-band flush
// happens. Reconcile loops should consult this method so the local tracker
// shrinks in lock-step with the kernel; per-packet hot paths should stay on
// IsBlocked since this issues a netlink RTT.
//
// Malformed IPs are reported as absent. Netlink and engine-initialization
// failures are returned so callers can keep their cached answer instead of
// deleting local state on a transient lookup failure.
func (e *Engine) IsBlockedLive(ip string) (bool, error) {
e.mu.Lock()
defer e.mu.Unlock()
return e.isBlockedLiveLocked(ip)
}
// livePermTempCountsLocked returns the count of permanent and temporary
// entries across the blocked v4 + v6 nft sets. Used by blockIPTarget so
// the deny limits trip against the kernel's actual state instead of
// stale state.json entries that the kernel already expired. Live keys
// that still exist in CSM state are classified from state because nft
// timeout attributes can reflect inherited/default set behaviour rather
// than the operator's block intent. Out-of-state live keys fall back to
// the kernel expiration attributes.
//
// Must be called with e.mu held; blockIPTarget already holds the lock
// at the only call site.
func (e *Engine) livePermTempCountsLocked(state FirewallState) (perm, temp int, ok bool) {
if e.liveBlockCounts != nil {
p, t, err := e.liveBlockCounts()
if err != nil {
return 0, 0, false
}
return p, t, true
}
if e.conn == nil {
return 0, 0, false
}
stateTempByIP := blockedStateTempByIP(state)
gotAny := false
for _, set := range []*nftables.Set{e.setBlocked, e.setBlocked6} {
if set == nil {
continue
}
elements, err := e.conn.GetSetElements(set)
if err != nil {
return 0, 0, false
}
gotAny = true
p, t := countLiveBlockElements(elements, stateTempByIP)
perm += p
temp += t
}
if !gotAny {
return 0, 0, false
}
return perm, temp, true
}
func blockedStateTempByIP(state FirewallState) map[string]bool {
byIP := make(map[string]bool, len(state.Blocked)*2)
for _, entry := range state.Blocked {
if entry.IP == "" {
continue
}
temp := !entry.ExpiresAt.IsZero()
byIP[entry.IP] = temp
if parsed := net.ParseIP(entry.IP); parsed != nil {
byIP[parsed.String()] = temp
}
}
return byIP
}
func countLiveBlockElements(elements []nftables.SetElement, stateTempByIP map[string]bool) (perm, temp int) {
for _, el := range elements {
if ip, ok := setElementIPString(el.Key); ok {
if stateTemp, found := stateTempByIP[ip]; found {
if stateTemp {
temp++
} else {
perm++
}
continue
}
}
if el.Timeout > 0 || el.Expires > 0 {
temp++
} else {
perm++
}
}
return perm, temp
}
func (e *Engine) liveBlockElementKindLocked(set *nftables.Set, key []byte) (live, temporary, classified bool) {
if set == nil || (e.liveBlockedDump == nil && e.conn == nil) {
return false, false, false
}
elements, err := e.dumpBlockedSetLocked(set)
if err != nil {
return false, false, false
}
for _, el := range elements {
if bytes.Equal(el.Key, key) {
return true, el.Timeout > 0 || el.Expires > 0, true
}
}
return false, false, true
}
func setElementIPString(key []byte) (string, bool) {
switch len(key) {
case net.IPv4len:
return net.IP(key).String(), true
case net.IPv6len:
return net.IP(key).String(), true
default:
return "", false
}
}
func (e *Engine) isBlockedLiveLocked(ip string) (bool, error) {
parsed := net.ParseIP(ip)
if parsed == nil {
return false, nil
}
var (
set *nftables.Set
key []byte
)
if ip4 := parsed.To4(); ip4 != nil {
set = e.setBlocked
key = ip4
} else {
set = e.setBlocked6
key = parsed.To16()
}
if set == nil {
return false, fmt.Errorf("blocked set unavailable for %s", ip)
}
if e.liveBlockLookup != nil {
return e.liveBlockLookup(set, key)
}
if e.conn == nil {
return false, fmt.Errorf("nftables connection unavailable")
}
elements, err := e.conn.GetSetElements(set)
if err != nil {
return false, fmt.Errorf("listing blocked set: %w", err)
}
for _, el := range elements {
if bytes.Equal(el.Key, key) {
return true, nil
}
}
return false, nil
}
// LiveBlockedSet dumps the blocked v4 and v6 sets once and returns a
// membership snapshot. IsBlockedLive issues a full set dump per call, so a
// reconcile pass over N tracked IPs cost N dumps of the entire set; callers
// testing more than one IP should snapshot instead.
//
// A nil set is left uncovered rather than reported empty, and so is one whose
// dump failed: the returned snapshot always describes whichever families did
// answer, with the failures joined into the error. Callers keep their cached
// view for uncovered families instead of erasing local state on a transient
// netlink error, while a family that dumped cleanly stays reconciled.
func (e *Engine) LiveBlockedSet() (LiveBlockedSnapshot, error) {
e.mu.Lock()
defer e.mu.Unlock()
if e.setBlocked == nil && e.setBlocked6 == nil {
return LiveBlockedSnapshot{}, fmt.Errorf("blocked sets unavailable")
}
if e.liveBlockedDump == nil && e.conn == nil {
return LiveBlockedSnapshot{}, fmt.Errorf("nftables connection unavailable")
}
var snap LiveBlockedSnapshot
var dumpErrs []error
for _, target := range []struct {
set *nftables.Set
members *map[string]struct{}
covered *bool
}{
{e.setBlocked, &snap.V4, &snap.HasV4},
{e.setBlocked6, &snap.V6, &snap.HasV6},
} {
if target.set == nil {
continue
}
elements, err := e.dumpBlockedSetLocked(target.set)
if err != nil {
// One family failing does not invalidate the other. Leaving this
// family uncovered keeps its callers on their cached answer while
// the family that did dump stays reconciled against the kernel.
dumpErrs = append(dumpErrs, fmt.Errorf("listing blocked set %s: %w", target.set.Name, err))
continue
}
members := make(map[string]struct{}, len(elements))
for _, el := range elements {
// Keys come from the kernel; anything that is not an address
// width cannot be matched against and is not ours to interpret.
if len(el.Key) != net.IPv4len && len(el.Key) != net.IPv6len {
continue
}
members[net.IP(el.Key).String()] = struct{}{}
}
*target.members = members
*target.covered = true
}
return snap, errors.Join(dumpErrs...)
}
func (e *Engine) dumpBlockedSetLocked(set *nftables.Set) ([]nftables.SetElement, error) {
if e.liveBlockedDump != nil {
return e.liveBlockedDump(set)
}
return e.conn.GetSetElements(set)
}
// AllowIP adds an IP to the allowed set and persists it.
// If the IP is currently blocked, the block is removed first.
func (e *Engine) AllowIP(ip string, reason string) error {
return e.allowIP(ip, reason, 0, "allow")
}
// TempAllowIP adds a temporary allow with expiry. Uses the same allowed set
// but tracks expiry in state - CleanExpiredAllows removes them periodically.
func (e *Engine) TempAllowIP(ip string, reason string, timeout time.Duration) error {
return e.allowIP(ip, reason, timeout, "temp_allow")
}
func (e *Engine) allowIP(ip string, reason string, timeout time.Duration, action string) (resultErr error) {
defer func() {
e.legacyFirewallFailure(action, ip, reason, InferProvenance(action, reason), timeout, resultErr)
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
blockedSet, blockedKey, _ := e.resolveIPSet(ip, e.setBlocked, e.setBlocked6)
allowedSet, allowedKey, err := e.resolveIPSet(ip, e.setAllowed, e.setAllowed6)
if err != nil {
return err
}
if allowedSet == nil {
return fmt.Errorf("allowed set unavailable for %s", ip)
}
entry := AllowedEntry{IP: ip, Reason: reason, Source: InferProvenance(action, reason)}
if action == "temp_allow" && timeout > 0 {
entry.ExpiresAt = time.Now().Add(timeout)
}
priorState := e.loadStateFile()
nextState := copyFirewallState(priorState)
removeBlockedIPFromState(&nextState, ip)
upsertAllowedEntryInState(&nextState, entry)
if e.lifecycle != nil {
return e.runDurableLocked(ActionRequest{Operation: action, Target: ip, Reason: reason, TTL: timeout}, nil, nextState)
}
if err := e.persistFirewallIntent(priorState, nextState); err != nil {
return fmt.Errorf("persisting %s for %s: %w", action, ip, err)
}
conn := e.newMutationConn()
if blockedSet != nil {
if err := conn.SetDeleteElements(blockedSet, []nftables.SetElement{{Key: blockedKey}}); err != nil {
if restoreErr := e.saveState(&priorState); restoreErr != nil {
return fmt.Errorf("removing from blocked set: %w (state restore failed: %v)", err, restoreErr)
}
logNftSetOpErr(action+" remove from blocked", ip, err)
return fmt.Errorf("removing from blocked set: %w", err)
}
}
if err := conn.SetAddElements(allowedSet, []nftables.SetElement{{Key: allowedKey}}); err != nil {
if restoreErr := e.saveState(&priorState); restoreErr != nil {
return fmt.Errorf("adding to allowed set: %w (state restore failed: %v)", err, restoreErr)
}
return fmt.Errorf("adding to allowed set: %w", err)
}
if err := conn.Flush(); err != nil {
retryErr := e.retryAllowAfterBenignFlushError(allowedSet, allowedKey, err)
if retryErr == nil {
e.legacyActionAuditLocked(action, ip, reason, entry.Source, timeout)
return nil
}
if restoreErr := e.saveState(&priorState); restoreErr != nil {
return fmt.Errorf("flushing: %w (state restore failed: %v)", err, restoreErr)
}
if isNftNotFound(err) {
return fmt.Errorf("flushing: %w (retry failed: %w)", err, retryErr)
}
return fmt.Errorf("flushing: %w", err)
}
e.legacyActionAuditLocked(action, ip, reason, entry.Source, timeout)
return nil
}
func (e *Engine) retryAllowAfterBenignFlushError(allowedSet *nftables.Set, allowedKey []byte, flushErr error) error {
if !isNftNotFound(flushErr) {
return flushErr
}
// Only the blocked delete can be absent when the allowed set exists.
// Retrying the add alone keeps the retry atomic and avoids a committed
// delete followed by a failed add.
conn := e.newMutationConn()
if err := conn.SetAddElements(allowedSet, []nftables.SetElement{{Key: allowedKey}}); err != nil {
return fmt.Errorf("retry adding to allowed set: %w", err)
}
if err := conn.Flush(); err != nil {
return fmt.Errorf("retry flushing allowed add: %w", err)
}
return nil
}
// CleanExpiredAllows removes expired temporary allows from the set and state.
// An IP is only removed from nftables if no non-expired entries remain for it.
// Called periodically by the daemon.
func (e *Engine) CleanExpiredAllows() int {
e.mu.Lock()
defer e.mu.Unlock()
prior, ok := e.loadStateFileRawLocked()
if !ok {
return 0
}
next := copyFirewallState(prior)
next.Allowed = nil
var expired []AllowedEntry
expiredIPs := make(map[string]bool)
now := time.Now()
for _, entry := range prior.Allowed {
if !entry.ExpiresAt.IsZero() && !now.Before(entry.ExpiresAt) {
expired = append(expired, entry)
expiredIPs[stateIPKey(entry.IP)] = true
} else {
next.Allowed = append(next.Allowed, entry)
}
}
if len(expired) == 0 {
return 0
}
for _, entry := range next.Allowed {
delete(expiredIPs, stateIPKey(entry.IP))
}
var removeIPs []string
for ip := range expiredIPs {
removeIPs = append(removeIPs, ip)
}
if err := e.commitAllowedRemovals(prior, next, removeIPs, ActionRequest{Operation: "temp_allow_expired", Target: "*", Actor: "daemon", Source: SourceSystem}); err != nil {
for _, entry := range expired {
if e.lifecycle == nil {
recordFirewallFailure("temp_allow_expired", entry.IP, "", SourceSystem, 0, err)
}
}
fmt.Fprintf(os.Stderr, "firewall: expired-allow cleanup failed: %v\n", err)
return 0
}
for _, entry := range expired {
e.legacyActionAuditLocked("temp_allow_expired", entry.IP, "", SourceSystem, 0)
}
return len(expired)
}
// CleanExpiredSubnets removes expired temporary subnet blocks from nftables and state.
func (e *Engine) CleanExpiredSubnets() int {
e.mu.Lock()
defer e.mu.Unlock()
prior, ok := e.loadStateFileRawLocked()
if !ok {
return 0
}
next := copyFirewallState(prior)
var active, expired []SubnetEntry
now := time.Now()
for _, entry := range prior.BlockedNet {
if !entry.ExpiresAt.IsZero() && !now.Before(entry.ExpiresAt) {
expired = append(expired, entry)
} else {
active = append(active, entry)
}
}
if len(expired) == 0 {
return 0
}
next.BlockedNet = active
if err := e.updateSubnetMutation(prior, next, ActionRequest{Operation: "temp_subnet_expired", Target: "*", Actor: "daemon", Source: SourceSystem}, nil); err != nil {
for _, entry := range expired {
if e.lifecycle == nil {
recordFirewallFailure("temp_subnet_expired", entry.CIDR, "", SourceSystem, 0, err)
}
}
fmt.Fprintf(os.Stderr, "firewall: expired-subnet cleanup failed: %v\n", err)
return 0
}
for _, entry := range expired {
e.legacyActionAuditLocked("temp_subnet_expired", entry.CIDR, "", SourceSystem, 0)
}
return len(expired)
}
// RemoveAllowIP removes an IP from the allowed set and state.
func (e *Engine) RemoveAllowIP(ip string) (resultErr error) {
defer func() { e.legacyFirewallFailure("remove_allow", ip, "", "", 0, resultErr) }()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
prior := e.loadStateFile()
next := copyFirewallState(prior)
next.Allowed = nil
for _, entry := range prior.Allowed {
if !sameIPString(entry.IP, ip) {
next.Allowed = append(next.Allowed, entry)
}
}
if err := e.commitAllowedRemovals(prior, next, []string{ip}, ActionRequest{Operation: "remove_allow", Target: ip}); err != nil {
return err
}
e.legacyActionAuditLocked("remove_allow", ip, "", "", 0)
return nil
}
// RemoveAllowIPBySource removes only allow entries from a specific source.
// The IP is only removed from the nftables set if no other sources remain.
func (e *Engine) RemoveAllowIPBySource(ip, source string) (resultErr error) {
defer func() { e.legacyFirewallFailure("remove_allow", ip, "", source, 0, resultErr) }()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
prior := e.loadStateFile()
next := copyFirewallState(prior)
_, ipGone := removeAllowedSourceFromState(&next, ip, source)
var removeIPs []string
if ipGone {
removeIPs = []string{ip}
}
if err := e.commitAllowedRemovals(prior, next, removeIPs, ActionRequest{Operation: "remove_allow", Target: ip, Reason: "source: " + source, Source: source}); err != nil {
return err
}
e.legacyActionAuditLocked("remove_allow", ip, "source: "+source, source, 0)
return nil
}
// AllowIPPort adds a port-specific IP allow. The rule is persisted to state
// and applied on the next Apply(). For immediate effect, call Apply() after.
func (e *Engine) AllowIPPort(ip string, port int, proto string, reason string) (resultErr error) {
defer func() {
e.legacyFirewallFailure("allow_port", fmt.Sprintf("%s:%d/%s", ip, port, proto), reason, InferProvenance("allow_port", reason), 0, resultErr)
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
if port < 1 || port > 65535 {
return fmt.Errorf("invalid port: %d", port)
}
if proto != "tcp" && proto != "udp" {
proto = "tcp"
}
e.mu.Lock()
defer e.mu.Unlock()
st := e.loadStateFile()
// Deduplicate
for _, existing := range st.PortAllowed {
if sameIPString(existing.IP, ip) && existing.Port == port && existing.Proto == proto {
return e.savePortPolicyLocked(&st, ActionRequest{Operation: "configure_port_allow", Target: fmt.Sprintf("%s:%d/%s", ip, port, proto), Reason: reason}) // Confirm durability even after a prior uncertain write.
}
}
st.PortAllowed = append(st.PortAllowed, PortAllowEntry{
IP: ip, Port: port, Proto: proto, Reason: reason, Source: InferProvenance("allow_port", reason),
})
if err := e.savePortPolicyLocked(&st, ActionRequest{Operation: "configure_port_allow", Target: fmt.Sprintf("%s:%d/%s", ip, port, proto), Reason: reason}); err != nil {
return fmt.Errorf("persisting port allow change: %w", err)
}
e.legacyActionAuditLocked("allow_port", fmt.Sprintf("%s:%d/%s", ip, port, proto), reason, InferProvenance("allow_port", reason), 0)
return nil
}
// RemoveAllowIPPort removes a port-specific IP allow from state.
func (e *Engine) RemoveAllowIPPort(ip string, port int, proto string) (resultErr error) {
defer func() {
e.legacyFirewallFailure("remove_port_allow", fmt.Sprintf("%s:%d/%s", ip, port, proto), "", "", 0, resultErr)
}()
canonical, err := canonicalFirewallIP(ip)
if err != nil {
return err
}
ip = canonical
e.mu.Lock()
defer e.mu.Unlock()
st := e.loadStateFile()
var remaining []PortAllowEntry
found := false
for _, entry := range st.PortAllowed {
if sameIPString(entry.IP, ip) && entry.Port == port && entry.Proto == proto {
found = true
continue
}
remaining = append(remaining, entry)
}
if !found {
return fmt.Errorf("port allow not found: %s:%d/%s", ip, port, proto)
}
st.PortAllowed = remaining
if err := e.savePortPolicyLocked(&st, ActionRequest{Operation: "remove_port_allow", Target: fmt.Sprintf("%s:%d/%s", ip, port, proto)}); err != nil {
return fmt.Errorf("persisting port allow change: %w", err)
}
e.legacyActionAuditLocked("remove_port_allow", fmt.Sprintf("%s:%d/%s", ip, port, proto), "", "", 0)
return nil
}
// FlushBlocked removes all IPs from the blocked set and clears persisted state.
func (e *Engine) FlushBlocked() (resultErr error) {
defer func() { e.legacyFirewallFailure("flush", "", "", "", 0, resultErr) }()
e.mu.Lock()
defer e.mu.Unlock()
// Persist the operator's intent before changing nftables. If the process
// dies between these steps, the next Apply must converge to an empty block
// set instead of restoring every flushed IP from stale state.
priorState := e.loadStateFile()
nextState := copyFirewallState(priorState)
nextState.Blocked = nil
if e.lifecycle != nil {
return e.runDurableLocked(ActionRequest{Operation: "flush", Target: "*"}, nil, nextState)
}
if err := e.persistFirewallIntent(priorState, nextState); err != nil {
return fmt.Errorf("persisting flush: %w", err)
}
e.conn.FlushSet(e.setBlocked)
if e.setBlocked6 != nil {
e.conn.FlushSet(e.setBlocked6)
}
if err := e.conn.Flush(); err != nil {
if restoreErr := e.saveState(&priorState); restoreErr != nil {
return fmt.Errorf("flushing blocked set: %w (state restore failed: %v)", err, restoreErr)
}
return fmt.Errorf("flushing blocked set: %w", err)
}
e.legacyActionAuditLocked("flush", "", fmt.Sprintf("cleared %d entries", len(priorState.Blocked)), SourceSystem, 0)
return nil
}
// subnetSafetyGuardLocked refuses a CIDR block that would firewall traffic the
// daemon must keep reachable: infra IPs or ranges, DNS-resolved infra hosts,
// local interface addresses, full-IP allows, port-specific allows, or the
// default route. Subnet blocks cover many addresses, and the output chain has
// no infra carve-out, so an unsafe subnet can lock out operators or kill the
// daemon's own egress.
// Must be called with e.mu held.
func (e *Engine) subnetSafetyGuardLocked(network *net.IPNet) error {
state := e.loadStateFile()
if e.lifecycle != nil && e.stateReadErr != nil {
return e.stateReadErr
}
return e.subnetSafetyGuardStateLocked(network, state)
}
func (e *Engine) subnetSafetyGuardStateLocked(network *net.IPNet, state FirewallState) error {
if ones, _ := network.Mask.Size(); ones == 0 {
return ipProtectedErrorf("refusing to block default route: %s", network.String())
}
// An unspecified host is not a usable target, but its containing range
// can be: operators may block 0.0.0.0/8 as a bogon range.
if ones, bits := network.Mask.Size(); ones == bits && network.IP.IsUnspecified() {
return ipProtectedErrorf("refusing to block non-routable range: %s", network.String())
}
// Checking only the first address misses larger ranges covering a local
// scope. Interface enumeration intentionally omits these scopes.
for _, protected := range protectedLocalRanges {
if network.Contains(protected.IP) || protected.Contains(network.IP) {
return ipProtectedErrorf("refusing to block subnet %s: overlaps protected range %s", network, protected)
}
}
for _, raw := range e.cfg.InfraIPs {
if _, infraNet, cidrErr := net.ParseCIDR(raw); cidrErr == nil {
if network.Contains(infraNet.IP) || infraNet.Contains(network.IP) {
return ipProtectedErrorf("refusing to block subnet %s: overlaps infra range %s", network.String(), raw)
}
continue
}
if infraIP := net.ParseIP(raw); infraIP != nil && network.Contains(infraIP) {
return ipProtectedErrorf("refusing to block subnet %s: contains infra IP %s", network.String(), raw)
}
}
for host, set := range e.infraResolved {
for key := range set {
if ip := net.ParseIP(key); ip != nil && network.Contains(ip) {
return ipProtectedErrorf("refusing to block subnet %s: contains infra IP %s (resolved from %s)", network.String(), key, host)
}
}
}
e.refreshLocalAddrsLocked()
for key := range e.localAddrs {
if ip := net.ParseIP(key); ip != nil && network.Contains(ip) {
return ipProtectedErrorf("refusing to block subnet %s: contains local host IP %s", network.String(), key)
}
}
for _, entry := range state.Allowed {
if ip := net.ParseIP(entry.IP); ip != nil && network.Contains(ip) {
return ipProtectedErrorf("refusing to block subnet %s: contains allowed IP %s", network.String(), entry.IP)
}
}
for _, entry := range state.PortAllowed {
if ip := net.ParseIP(entry.IP); ip != nil && network.Contains(ip) {
return ipProtectedErrorf("refusing to block subnet %s: contains port-allowed IP %s", network.String(), entry.IP)
}
}
return nil
}
var protectedLocalRanges = func() []*net.IPNet {
ranges := []*net.IPNet{
{IP: net.IPv4(127, 0, 0, 0), Mask: net.CIDRMask(8, 32)},
{IP: net.IPv4(169, 254, 0, 0), Mask: net.CIDRMask(16, 32)},
{IP: net.IPv4(224, 0, 0, 0), Mask: net.CIDRMask(24, 32)},
{IP: net.ParseIP("::1"), Mask: net.CIDRMask(128, 128)},
{IP: net.ParseIP("fe80::"), Mask: net.CIDRMask(10, 128)},
}
// Multicast flags vary independently of the link-local scope nibble.
for flags := byte(0); flags < 16; flags++ {
ip := make(net.IP, net.IPv6len)
ip[0], ip[1] = 0xff, flags<<4|2
ranges = append(ranges, &net.IPNet{IP: ip, Mask: net.CIDRMask(16, 128)})
}
return ranges
}()
func (e *Engine) subnetBlockPlanLocked(cidr string, state FirewallState) (*net.IPNet, bool, error) {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
return nil, false, fmt.Errorf("invalid CIDR: %s", cidr)
}
if err := e.subnetSafetyGuardStateLocked(network, state); err != nil {
return nil, false, err
}
for _, entry := range state.BlockedNet {
if entry.CIDR == network.String() {
return network, true, nil
}
}
targetSet, _, _ := e.resolveSubnetSet(network)
if targetSet == nil {
return nil, false, fmt.Errorf("no matching set for %s (IPv6 disabled?)", cidr)
}
return network, false, nil
}
// ValidateSubnetBlock runs the same safety and capability checks as
// BlockSubnet without changing persisted or kernel firewall state.
func (e *Engine) ValidateSubnetBlock(cidr string) error {
e.mu.Lock()
defer e.mu.Unlock()
state := e.loadStateFile()
if e.lifecycle != nil && e.stateReadErr != nil {
return e.stateReadErr
}
_, _, err := e.subnetBlockPlanLocked(cidr, state)
return err
}
// BlockSubnet adds a CIDR range to the blocked subnets set (IPv4 or IPv6).
// timeout 0 = permanent block.
func (e *Engine) BlockSubnet(cidr string, reason string, timeout time.Duration) (resultErr error) {
return e.BlockSubnetWithFindingID(cidr, reason, timeout, "")
}
// BlockSubnetWithFindingID records a causal finding for an automatic subnet
// decision. Manual and maintenance callers use BlockSubnet without one.
func (e *Engine) BlockSubnetWithFindingID(cidr, reason string, timeout time.Duration, findingID string) (resultErr error) {
return e.BlockSubnetRequest(ActionRequest{Operation: "block_subnet", Target: cidr, Reason: reason, TTL: timeout, FindingID: findingID, Automatic: findingID != ""}, nil)
}
func (e *Engine) BlockSubnetRequest(req ActionRequest, budget *ScanAdmission) (resultErr error) {
cidr, reason, timeout, findingID := req.Target, req.Reason, req.TTL, req.FindingID
defer func() {
if e.shouldLegacyOutcome(resultErr) {
result, auditErr := actionlog.Applied, resultErr
if errors.Is(resultErr, ErrActionDryRun) {
result, auditErr = actionlog.DryRun, nil
}
recordFirewallFindingResult("block_subnet", cidr, reason, InferProvenance("block_subnet", reason), timeout, result, auditErr, findingID)
}
}()
_, network, parseErr := net.ParseCIDR(req.Target)
if parseErr != nil {
return fmt.Errorf("invalid CIDR: %s", req.Target)
}
req.Operation = "block_subnet"
req.Target = network.String()
req = normalizeActionRequest(req)
dryRun := req.Automatic && e.autoResponseDryRunEnabled()
e.mu.Lock()
defer e.mu.Unlock()
if found, replayErr := e.replayActionLocked(req); found || replayErr != nil {
return replayErr
}
if dryRun {
return ErrActionDryRun
}
if readyErr := e.lifecycleReadyLocked(); readyErr != nil {
return readyErr
}
priorState := e.loadStateFile()
if e.lifecycle != nil && e.stateReadErr != nil {
return e.stateReadErr
}
network, alreadyBlocked, err := e.subnetBlockPlanLocked(cidr, priorState)
if err != nil {
return err
}
cidr = network.String()
if alreadyBlocked {
if e.lifecycle != nil {
return nil
}
// A prior write may be visible despite a durability error, before the
// kernel changed. Retrying must reconcile that saved intent as well.
if err := e.updateSubnetMutation(priorState, priorState, req, budget); err != nil {
return err
}
e.legacyFileAuditLocked("block_subnet", network.String(), reason, InferProvenance("block_subnet", reason), timeout)
return nil
}
entry := SubnetEntry{
CIDR: network.String(),
Reason: reason,
Source: InferProvenance("block_subnet", reason),
BlockedAt: time.Now(),
}
if timeout > 0 {
entry.ExpiresAt = time.Now().Add(timeout)
}
// Persist to state BEFORE the kernel add, mirroring blockIPLocked: state
// seeds the next Apply, so a crash between the kernel add and a later
// state write would otherwise leave a kernel block that silently
// disappears on restart. On kernel failure the prior state is restored.
nextState := copyFirewallState(priorState)
addSubnetEntryIfMissingInState(&nextState, entry)
if err := e.updateSubnetMutation(priorState, nextState, req, budget); err != nil {
return err
}
e.legacyFileAuditLocked("block_subnet", network.String(), reason, entry.Source, timeout)
return nil
}
// IsSubnetBlocked returns true if the CIDR is present in the persisted subnet block state.
func (e *Engine) IsSubnetBlocked(cidr string) bool {
_, network, err := net.ParseCIDR(cidr)
if err != nil {
return false
}
e.mu.Lock()
defer e.mu.Unlock()
return e.isSubnetBlockedStateLocked(network.String())
}
// BlockedSubnetCovering reports the blocked CIDR (if any) that contains ip.
// The input-chain drops blocked_nets before the allowed_ips accept, so an
// allow on an IP inside a blocked subnet has no effect: the subnet drop still
// fires. Callers surface this so an operator is not told an IP is reachable
// when a subnet rule still blocks it. The subnet block stays authoritative by
// design (see subnetSafetyGuardLocked); this only reports, it does not unblock.
func (e *Engine) BlockedSubnetCovering(ip string) (string, bool) {
e.mu.Lock()
defer e.mu.Unlock()
return subnetCovering(e.loadStateFile().BlockedNet, ip)
}
// UnblockSubnet removes a CIDR range from the blocked subnets set (IPv4 or IPv6).
func (e *Engine) UnblockSubnet(cidr string) (resultErr error) {
defer func() { e.legacyFirewallFailure("unblock_subnet", cidr, "", "", 0, resultErr) }()
e.mu.Lock()
defer e.mu.Unlock()
_, network, err := net.ParseCIDR(cidr)
if err != nil {
return fmt.Errorf("invalid CIDR: %s", cidr)
}
targetSet, _, _ := e.resolveSubnetSet(network)
if targetSet == nil {
return fmt.Errorf("no matching set for %s (IPv6 disabled?)", cidr)
}
prior := e.loadStateFile()
next := copyFirewallState(prior)
remaining := next.BlockedNet[:0]
for _, entry := range next.BlockedNet {
if entry.CIDR != network.String() {
remaining = append(remaining, entry)
}
}
next.BlockedNet = remaining
if err := e.updateSubnetMutation(prior, next, ActionRequest{Operation: "unblock_subnet", Target: network.String()}, nil); err != nil {
return err
}
e.legacyActionAuditLocked("unblock_subnet", network.String(), "", "", 0)
return nil
}
// resolveSubnetSet returns the correct blocked_nets set and start/end keys
// for a CIDR. Returns (nil, nil, nil) when IPv6 is disabled for a v6 CIDR,
// or when lastIPInRange cannot produce an interval end (malformed
// net.IPNet whose IP is neither 4 nor 16 bytes) -- callers already check
// for a nil set, so this also short-circuits the degenerate case where
// nextIP(nil) would feed an empty Key to the kernel.
func (e *Engine) resolveSubnetSet(network *net.IPNet) (*nftables.Set, net.IP, net.IP) {
end := lastIPInRange(network)
if end == nil {
return nil, nil, nil
}
if start := network.IP.To4(); start != nil {
return e.setBlockedNet, start, end
}
if e.setBlockedNet6 != nil {
return e.setBlockedNet6, network.IP.To16(), end
}
return nil, nil, nil
}
func intervalSetElements(start, end net.IP) []nftables.SetElement {
if start == nil || end == nil {
return nil
}
return appendIntervalSetElements(nil, start, end)
}
func appendIntervalSetElements(dst []nftables.SetElement, start, end net.IP) []nftables.SetElement {
if start == nil || end == nil {
return dst
}
endMarker, ok := nextIntervalKey(end)
if !ok {
// nftables represents an interval through the last address with an
// open end. Wrapping its exclusive end to zero would change its meaning.
return append(dst, nftables.SetElement{Key: start})
}
return append(dst,
nftables.SetElement{Key: start},
nftables.SetElement{Key: endMarker, IntervalEnd: true},
)
}
// UpdateInfraResolved records the IP set last resolved for an infra
// hostname. Replaces any previous entry for that host so the resolver's
// per-tick refresh leaves no stale ghost IPs. Pass an empty ips slice
// to remove the host entirely (e.g. when DNS stopped resolving).
func (e *Engine) UpdateInfraResolved(host string, ips []string) {
e.mu.Lock()
defer e.mu.Unlock()
if host == "" {
return
}
if e.infraResolved == nil {
e.infraResolved = make(map[string]map[string]struct{})
}
if len(ips) == 0 {
delete(e.infraResolved, host)
return
}
set := make(map[string]struct{}, len(ips))
for _, ip := range ips {
// Normalise via net.ParseIP so canonical form is stored
// (collapses IPv6 forms and drops malformed values).
if parsed := net.ParseIP(ip); parsed != nil {
set[parsed.String()] = struct{}{}
}
}
if len(set) == 0 {
delete(e.infraResolved, host)
return
}
e.infraResolved[host] = set
}
// DropInfraResolved clears all resolved IPs for a host. Equivalent to
// UpdateInfraResolved(host, nil); separate name surfaces operator
// intent at call sites that purposefully retire a hostname.
func (e *Engine) DropInfraResolved(host string) {
e.UpdateInfraResolved(host, nil)
}
// infraIPResolvedHostLocked reports whether ip matches any IP recorded
// for any tracked infra hostname. Must be called with e.mu held; the
// existing blockIPTarget path already does so. The lookup normalizes ip
// to the same canonical form the storage path applies (net.ParseIP
// collapses IPv6 and rewrites ::ffff:1.2.3.4 to 1.2.3.4), so a caller
// passing the IPv4-mapped or uncanonical form still hits the guard.
func (e *Engine) infraIPResolvedHostLocked(ip string) (string, bool) {
if e.infraResolved == nil {
return "", false
}
parsed := net.ParseIP(ip)
if parsed == nil {
return "", false
}
key := parsed.String()
for host, set := range e.infraResolved {
if _, ok := set[key]; ok {
return host, true
}
}
return "", false
}
// localAddrsCacheTTL bounds how stale the host-own-IP set can get before
// the next block call rebuilds it. The trade-off: a newly assigned local
// address could be auto-blocked for up to this window if a flagged source
// happens to share that address. Local-address changes (operator running
// `ip addr add`) are rare, so 60s keeps both the FP window and the steady
// per-block syscall cost negligible. Refreshing every miss instead would
// pin the cache to permanently-fresh under any scan storm.
const localAddrsCacheTTL = 60 * time.Second
// refreshLocalAddrsLocked rebuilds the cache of host-own interface
// addresses when the TTL has expired. Must be called with e.mu held.
// Failure leaves the previous cache in place so a transient netlink
// hiccup cannot demote the lockout guard.
func (e *Engine) refreshLocalAddrsLocked() {
if e.localAddrs != nil && !e.localAddrsExpiresAt.IsZero() && time.Now().Before(e.localAddrsExpiresAt) {
return
}
var ips []string
if e.localAddrsLookup != nil {
got, err := e.localAddrsLookup()
if err != nil {
return
}
ips = got
} else {
addrs, err := net.InterfaceAddrs()
if err != nil {
return
}
for _, addr := range addrs {
ipnet, ok := addr.(*net.IPNet)
if !ok {
continue
}
ips = append(ips, ipnet.IP.String())
}
}
set := make(map[string]struct{}, len(ips))
for _, raw := range ips {
key, ok := localAddrGuardKey(raw)
if !ok {
continue
}
set[key] = struct{}{}
}
e.localAddrs = set
e.localAddrsExpiresAt = time.Now().Add(localAddrsCacheTTL)
}
// isUnblockableAddress reports addresses that can never legitimately be
// blocked, independent of what the interface enumeration returned.
//
// The local-address guard is built from localAddrGuardKey, which drops
// loopback and link-local, so those were the one class the guard did not
// cover -- and loopback is precisely what a control panel's own proxied
// requests appear as. A block of 127.0.0.1 was therefore accepted rather than
// refused; nothing broke only because the input chain accepts "iifname lo"
// before reaching the blocked set, leaving an entry that looked effective
// while doing nothing.
func isUnblockableAddress(ip string) bool {
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
return parsed.IsLoopback() ||
parsed.IsUnspecified() ||
parsed.IsLinkLocalUnicast() ||
parsed.IsLinkLocalMulticast()
}
func localAddrGuardKey(raw string) (string, bool) {
parsed := net.ParseIP(raw)
if parsed == nil {
return "", false
}
if isUnblockableAddress(raw) {
return "", false
}
return parsed.String(), true
}
// isLocalAddrLocked reports whether ip is one of the daemon's own host
// addresses. Must be called with e.mu held.
func (e *Engine) isLocalAddrLocked(ip string) bool {
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
e.refreshLocalAddrsLocked()
if len(e.localAddrs) == 0 {
return false
}
_, ok := e.localAddrs[parsed.String()]
return ok
}
// BlockedCount returns the number of live blocked IP entries the engine
// is enforcing. Sourced from the same state file Status() uses, so
// `/api/v1/status` and `csm firewall status` agree on the number. Expired
// entries are pruned by loadStateFile before being counted.
func (e *Engine) BlockedCount() int {
e.mu.Lock()
defer e.mu.Unlock()
s := e.loadStateFile()
return countBlockedRules(s.Blocked, e.cfg != nil && e.cfg.IPv6)
}
// RuleCounts returns the cardinality of every firewall rule category from
// the engine state file with expired temp bans pruned. Callers needing a
// live count (e.g. Prometheus gauges) must use this rather than the bbolt
// store, which holds only the migration-time snapshot.
func (e *Engine) RuleCounts() RuleCounts {
e.mu.Lock()
defer e.mu.Unlock()
s := e.loadStateFile()
return countRuleEntries(s, e.cfg != nil && e.cfg.IPv6)
}
// BlockedSubnets returns a snapshot of the active persisted subnet blocks.
// The returned slice is a fresh copy: loadStateFile reuses the warm shared
// state cache, so the slice it returns can alias internal engine state. Copy
// it here so a caller mutating the result cannot corrupt the cache. SubnetEntry
// is a value type (strings + time.Time, no reference fields), so a slice copy
// is a sufficient deep copy.
func (e *Engine) BlockedSubnets() []SubnetEntry {
e.mu.Lock()
defer e.mu.Unlock()
s := e.loadStateFile()
return append([]SubnetEntry(nil), s.BlockedNet...)
}
// Status returns current firewall statistics.
//
// Takes e.mu so the cached state can be read coherently. Before the
// cache existed loadStateFile was lock-free because every call did its
// own ReadFile + Unmarshal; now that loadStateFile mutates the shared
// cache + index, the lock is required.
func (e *Engine) Status() map[string]interface{} {
e.mu.Lock()
defer e.mu.Unlock()
state := e.loadStateFile()
return map[string]interface{}{
"enabled": e.cfg.Enabled,
"tcp_in": e.cfg.TCPIn,
"tcp_out": e.cfg.TCPOut,
"udp_in": e.cfg.UDPIn,
"udp_out": e.cfg.UDPOut,
"infra_ips": e.cfg.InfraIPs,
"blocked": len(state.Blocked),
"allowed": len(state.Allowed),
"log_dropped": e.cfg.LogDropped,
}
}
// --- State persistence ---
// initialBlockState is the pre-computed pool of nft set elements to
// seed a freshly-built csm table from persisted state. Populated by
// computeInitialBlockStateLocked and consumed by
// queueInitialBlockStateLocked.
type initialBlockState struct {
blocked4, blocked6 []nftables.SetElement
allowed4, allowed6 []nftables.SetElement
blockedNet4, blockedNet6 []nftables.SetElement
}
// computeInitialBlockStateLocked reads state.json and returns the
// nft elements needed to repopulate the blocked / allowed / blocked-
// net sets. Pure computation; does not touch nft. Safe to call from
// Apply before any AddTable / AddSet so the result can be queued
// into the same atomic netlink batch.
func (e *Engine) computeInitialBlockStateLocked() initialBlockState {
state := e.loadStateFile()
if e.lifecycle != nil {
now := time.Now()
sets := desiredActionSets(state, e.cfg.IPv6, now)
return initialBlockState{blocked4: actionElements(sets["blocked_ips"], now), blocked6: actionElements(sets["blocked_ips6"], now), allowed4: actionElements(sets["allowed_ips"], now), allowed6: actionElements(sets["allowed_ips6"], now), blockedNet4: actionElements(sets["blocked_nets"], now), blockedNet6: actionElements(sets["blocked_nets6"], now)}
}
now := time.Now()
var ibs initialBlockState
restoredBlocked := make(map[string]bool)
for _, entry := range state.Blocked {
if !entry.ExpiresAt.IsZero() && now.After(entry.ExpiresAt) {
continue
}
key, ok := canonicalIPKey(entry.IP)
if !ok {
continue
}
if restoredBlocked[key] {
continue
}
parsed := net.ParseIP(entry.IP)
timeout := time.Duration(0)
if !entry.ExpiresAt.IsZero() {
timeout = time.Until(entry.ExpiresAt)
}
if ip4 := parsed.To4(); ip4 != nil {
ibs.blocked4 = append(ibs.blocked4, nftables.SetElement{Key: ip4, Timeout: timeout})
} else if e.cfg.IPv6 {
ibs.blocked6 = append(ibs.blocked6, nftables.SetElement{Key: parsed.To16(), Timeout: timeout})
}
restoredBlocked[key] = true
}
restoredAllowed := make(map[string]bool)
for _, entry := range state.Allowed {
if !entry.ExpiresAt.IsZero() && now.After(entry.ExpiresAt) {
continue
}
key, ok := canonicalIPKey(entry.IP)
if !ok {
continue
}
if restoredAllowed[key] {
continue
}
parsed := net.ParseIP(entry.IP)
if ip4 := parsed.To4(); ip4 != nil {
ibs.allowed4 = append(ibs.allowed4, nftables.SetElement{Key: ip4})
} else if e.cfg.IPv6 {
ibs.allowed6 = append(ibs.allowed6, nftables.SetElement{Key: parsed.To16()})
}
restoredAllowed[key] = true
}
ibs.blockedNet4, ibs.blockedNet6 = subnetIntervalElements(state.BlockedNet, e.cfg.IPv6, now)
return ibs
}
// queueInitialBlockStateLocked queues the previously-computed
// elements into the still-pending Apply netlink batch. Apply Flushes
// the whole batch as one transaction. No Flush here.
func (e *Engine) queueInitialBlockStateLocked(ibs initialBlockState) error {
if err := e.addElementsChunked(e.setBlocked, ibs.blocked4); err != nil {
return err
}
if e.setBlocked6 != nil {
if err := e.addElementsChunked(e.setBlocked6, ibs.blocked6); err != nil {
return err
}
}
if err := e.addElementsChunked(e.setAllowed, ibs.allowed4); err != nil {
return err
}
if e.setAllowed6 != nil {
if err := e.addElementsChunked(e.setAllowed6, ibs.allowed6); err != nil {
return err
}
}
if err := e.addElementsChunked(e.setBlockedNet, ibs.blockedNet4); err != nil {
return err
}
if e.setBlockedNet6 != nil {
if err := e.addElementsChunked(e.setBlockedNet6, ibs.blockedNet6); err != nil {
return err
}
}
return nil
}
// addElementsChunked issues SetAddElements in fixed-size chunks. A single
// SetAddElements call encodes all elements into one netlink message whose
// size scales linearly with len(elems); the kernel's netlink socket rmem
// (typically 208 KB, tunable via net.core.rmem_max) caps how big that
// message can be before the receive path refuses it with ENOBUFS. At
// ~28 bytes per element worst-case a 1000-element chunk is ~28 KB, well
// under the default rmem and comfortably below any realistic rmem_max.
// The batch size must stay even so merged interval boundary pairs never
// split across chunks. Only the final range can have an open upper end.
func (e *Engine) addElementsChunked(s *nftables.Set, elems []nftables.SetElement) error {
return addElementsChunked(e.conn, s, elems)
}
func addElementsChunked(conn *nftables.Conn, s *nftables.Set, elems []nftables.SetElement) error {
const chunk = 1000
for i := 0; i < len(elems); i += chunk {
end := i + chunk
if end > len(elems) {
end = len(elems)
}
if err := conn.SetAddElements(s, elems[i:end]); err != nil {
op := fmt.Sprintf("add elements to set %q chunk %d-%d", s.Name, i, end)
logNftSetOpErr(op, "initial restore", err)
return fmt.Errorf("adding initial elements to set %q chunk %d-%d: %w", s.Name, i, end, err)
}
}
return nil
}
// logNftSetOpErr keeps operator-visible nft error logs grep-friendly.
func logNftSetOpErr(op, target string, err error) {
fmt.Fprintf(os.Stderr, "firewall: nft %s for %s failed: %v\n", op, target, err)
}
// loadStateFile returns a deep copy of the cached firewall state with
// expired entries pruned. The on-disk state.json is re-read only when
// the file metadata key differs from the cached key (or the cache is empty).
//
// All callers must hold e.mu. The returned value is safe to mutate
// without affecting the cache; mutators write back via saveState which
// rebuilds the cache from the passed-in struct.
//
// Hot read paths (IsBlocked, IsSubnetBlocked, IsAllowed) bypass this
// allocation by consulting the index maps directly.
func (e *Engine) loadStateFile() FirewallState {
e.ensureStateCacheLocked()
if e.stateCache == nil {
return FirewallState{}
}
s := e.stateCache
return FirewallState{
Blocked: append([]BlockedEntry(nil), s.Blocked...),
BlockedNet: append([]SubnetEntry(nil), s.BlockedNet...),
Allowed: append([]AllowedEntry(nil), s.Allowed...),
PortAllowed: append([]PortAllowEntry(nil), s.PortAllowed...),
}
}
// loadStateFileRawLocked reads state.json without pruning expired entries.
// Expiry cleanup needs the stale rows so it can remove matching nftables
// elements before writing the active state back.
func (e *Engine) loadStateFileRawLocked() (FirewallState, bool) {
if e.lifecycle != nil {
state, _, err := e.readCommittedStateLocked()
return state, err == nil
}
stateFile := filepath.Join(e.statePath, "state.json")
data, err := os.ReadFile(stateFile) // #nosec G304 -- filepath.Join under operator-configured statePath.
if err != nil {
if os.IsNotExist(err) {
return FirewallState{}, true
}
return FirewallState{}, false
}
var state FirewallState
if err := json.Unmarshal(data, &state); err != nil {
return FirewallState{}, false
}
return state, true
}
// ensureStateCacheLocked populates or refreshes e.stateCache. Cheap on
// cache-hit (one stat). On cache-miss it does the full ReadFile +
// json.Unmarshal that the pre-cache implementation did on every call.
//
// Expired entries are pruned in-place after load so IsBlocked and the
// other index-backed lookups never report a stale block. Pruning
// updates the index maps via rebuildIndexLocked.
func (e *Engine) ensureStateCacheLocked() {
if e.lifecycle != nil {
state, _, err := e.readCommittedStateLocked()
if err == nil {
e.installCommittedCache(state)
}
return
}
stateFile := filepath.Join(e.statePath, "state.json")
info, statErr := os.Stat(stateFile)
if statErr == nil && e.stateCache != nil && e.stateCacheKey.matches(info) {
e.applyExpiryLocked()
return
}
var fresh FirewallState
switch {
case statErr == nil:
// #nosec G304 -- filepath.Join under operator-configured statePath.
data, readErr := os.ReadFile(stateFile)
if readErr != nil {
e.keepPriorStateCacheLocked()
return
}
if err := json.Unmarshal(data, &fresh); err != nil {
e.keepPriorStateCacheLocked()
return
}
e.stateCacheKey = stateFileCacheKeyFromInfo(info)
case os.IsNotExist(statErr):
e.stateCacheKey = stateFileCacheKey{}
case e.stateCache != nil:
// Transient stat error (permission, EIO). Keep the prior
// cache rather than dropping all blocks.
e.applyExpiryLocked()
return
}
normalizeFirewallStateIPs(&fresh)
e.stateCache = &fresh
e.applyExpiryLocked()
e.rebuildIndexLocked()
}
func (e *Engine) keepPriorStateCacheLocked() {
if e.stateCache != nil {
e.applyExpiryLocked()
return
}
e.stateCache = &FirewallState{}
e.stateCacheKey = stateFileCacheKey{}
e.rebuildIndexLocked()
}
// applyExpiryLocked prunes expired entries from the cached state in
// place. Returns whether anything changed; the index maps are rebuilt
// when something did.
func (e *Engine) applyExpiryLocked() {
if e.stateCache == nil {
return
}
now := time.Now()
s := e.stateCache
changed := false
if pruned, dropped := pruneBlocked(s.Blocked, now); dropped {
s.Blocked = pruned
changed = true
}
if pruned, dropped := pruneBlockedNet(s.BlockedNet, now); dropped {
s.BlockedNet = pruned
changed = true
}
if pruned, dropped := pruneAllowed(s.Allowed, now); dropped {
s.Allowed = pruned
changed = true
}
if changed {
e.rebuildIndexLocked()
}
}
// rebuildIndexLocked refreshes the three lookup maps from the cached
// state. Must be called any time e.stateCache is mutated.
func (e *Engine) rebuildIndexLocked() {
if e.stateCache == nil {
e.blockedIPIndex = nil
e.allowedIPIndex = nil
e.portAllowedIndex = nil
e.blockedCIDRIndex = nil
return
}
s := e.stateCache
blocked := make(map[string]int, len(s.Blocked))
for i, entry := range s.Blocked {
blocked[entry.IP] = i
if key, ok := canonicalIPKey(entry.IP); ok {
blocked[key] = i
}
}
e.blockedIPIndex = blocked
allowed := make(map[string]struct{}, len(s.Allowed))
for _, entry := range s.Allowed {
allowed[entry.IP] = struct{}{}
if key, ok := canonicalIPKey(entry.IP); ok {
allowed[key] = struct{}{}
}
}
e.allowedIPIndex = allowed
portAllowed := make(map[string]struct{}, len(s.PortAllowed))
for _, entry := range s.PortAllowed {
portAllowed[entry.IP] = struct{}{}
if key, ok := canonicalIPKey(entry.IP); ok {
portAllowed[key] = struct{}{}
}
}
e.portAllowedIndex = portAllowed
subnets := make(map[string]struct{}, len(s.BlockedNet))
for _, entry := range s.BlockedNet {
subnets[entry.CIDR] = struct{}{}
}
e.blockedCIDRIndex = subnets
}
// pruneBlocked returns the list with expired blocked-IP entries removed
// and a flag indicating whether anything was dropped. Returns the
// original slice when no entries expired so we avoid pointless
// allocations on the steady-state hot path.
func pruneBlocked(in []BlockedEntry, now time.Time) ([]BlockedEntry, bool) {
expired := 0
for _, entry := range in {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
expired++
}
}
if expired == 0 {
return in, false
}
out := make([]BlockedEntry, 0, len(in)-expired)
for _, entry := range in {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
continue
}
out = append(out, entry)
}
return out, true
}
func pruneBlockedNet(in []SubnetEntry, now time.Time) ([]SubnetEntry, bool) {
expired := 0
for _, entry := range in {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
expired++
}
}
if expired == 0 {
return in, false
}
out := make([]SubnetEntry, 0, len(in)-expired)
for _, entry := range in {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
continue
}
out = append(out, entry)
}
return out, true
}
func pruneAllowed(in []AllowedEntry, now time.Time) ([]AllowedEntry, bool) {
expired := 0
for _, entry := range in {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
expired++
}
}
if expired == 0 {
return in, false
}
out := make([]AllowedEntry, 0, len(in)-expired)
for _, entry := range in {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
continue
}
out = append(out, entry)
}
return out, true
}
var writeFirewallStateJSON = atomicio.AtomicWriteJSON
// saveState writes the firewall state to disk atomically (write to .tmp,
// rename into place) and rebuilds the in-memory cache to reflect the
// just-written snapshot. Callers must hold e.mu.
//
// The cache rebuild deep-copies the input slices so a caller that
// keeps mutating the local FirewallState after saveState returns cannot
// corrupt the cache.
func (e *Engine) saveState(s *FirewallState) error {
if e.lifecycle != nil {
return e.runDurableLocked(ActionRequest{Operation: "state", Target: "*"}, nil, *s)
}
state := copyFirewallState(*s)
normalizeFirewallStateIPs(&state)
path := filepath.Join(e.statePath, "state.json")
if err := writeFirewallStateJSON(path, 0o600, &state); err != nil {
if firewallStateFileMatches(path, 0o600, &state) {
e.setStateCacheLocked(path, &state)
return &stateDurabilityError{cause: err}
}
fmt.Fprintf(os.Stderr, "firewall: persist state.json failed: %v\n", err)
e.clearStateCacheLocked()
return err
}
e.setStateCacheLocked(path, &state)
return nil
}
func firewallStateFileMatches(path string, perm os.FileMode, s *FirewallState) bool {
info, err := os.Stat(path)
if err != nil || info.Mode().Perm() != perm {
return false
}
want, err := json.MarshalIndent(s, "", " ")
if err != nil {
return false
}
// #nosec G304 -- path is the engine-owned state file path.
got, err := os.ReadFile(path)
if err != nil {
return false
}
return bytes.Equal(got, want)
}
func (e *Engine) setStateCacheLocked(path string, s *FirewallState) {
var cacheKey stateFileCacheKey
if info, statErr := os.Stat(path); statErr == nil {
cacheKey = stateFileCacheKeyFromInfo(info)
}
state := copyFirewallState(*s)
normalizeFirewallStateIPs(&state)
e.stateCache = &FirewallState{
Blocked: append([]BlockedEntry(nil), state.Blocked...),
BlockedNet: append([]SubnetEntry(nil), state.BlockedNet...),
Allowed: append([]AllowedEntry(nil), state.Allowed...),
PortAllowed: append([]PortAllowEntry(nil), state.PortAllowed...),
}
e.stateCacheKey = cacheKey
e.rebuildIndexLocked()
}
func (e *Engine) clearStateCacheLocked() {
e.stateCache = nil
e.stateCacheKey = stateFileCacheKey{}
e.rebuildIndexLocked()
}
func copyFirewallState(s FirewallState) FirewallState {
return FirewallState{
Blocked: append([]BlockedEntry(nil), s.Blocked...),
BlockedNet: append([]SubnetEntry(nil), s.BlockedNet...),
Allowed: append([]AllowedEntry(nil), s.Allowed...),
PortAllowed: append([]PortAllowEntry(nil), s.PortAllowed...),
}
}
func upsertBlockedEntryInState(state *FirewallState, entry BlockedEntry) {
if key, ok := canonicalIPKey(entry.IP); ok {
entry.IP = key
}
out := state.Blocked[:0]
written := false
for _, existing := range state.Blocked {
if sameIPString(existing.IP, entry.IP) {
if !written {
out = append(out, entry)
written = true
}
continue
}
out = append(out, existing)
}
if !written {
out = append(out, entry)
}
state.Blocked = out
}
func removeBlockedIPFromState(state *FirewallState, ip string) {
remaining := state.Blocked[:0]
for _, entry := range state.Blocked {
if sameIPString(entry.IP, ip) {
continue
}
remaining = append(remaining, entry)
}
state.Blocked = remaining
}
func upsertAllowedEntryInState(state *FirewallState, entry AllowedEntry) {
if key, ok := canonicalIPKey(entry.IP); ok {
entry.IP = key
}
out := state.Allowed[:0]
written := false
for _, existing := range state.Allowed {
if sameIPString(existing.IP, entry.IP) && existing.Source == entry.Source {
if !written {
out = append(out, entry)
written = true
}
continue
}
out = append(out, existing)
}
if !written {
out = append(out, entry)
}
state.Allowed = out
}
func addSubnetEntryIfMissingInState(state *FirewallState, entry SubnetEntry) bool {
for _, existing := range state.BlockedNet {
if existing.CIDR == entry.CIDR {
return false
}
}
state.BlockedNet = append(state.BlockedNet, entry)
return true
}
func isNftNotFound(err error) bool {
return errors.Is(err, syscall.ENOENT)
}
func (e *Engine) restoreBlockStateAfterFailureLocked(state FirewallState, ip string) error {
if err := e.saveState(&state); err != nil {
fmt.Fprintf(os.Stderr, "firewall: restore state after failed block for %s failed: %v\n", ip, err)
return err
}
return nil
}
func (e *Engine) saveBlockedEntry(entry BlockedEntry) error {
if entry.Source == "" {
entry.Source = InferProvenance("block", entry.Reason)
}
if key, ok := canonicalIPKey(entry.IP); ok {
entry.IP = key
}
state := e.loadStateFile()
upsertBlockedEntryInState(&state, entry)
return e.saveState(&state)
}
func (e *Engine) removeBlockedState(ip string) error {
state := e.loadStateFile()
var remaining []BlockedEntry
for _, entry := range state.Blocked {
if !sameIPString(entry.IP, ip) {
remaining = append(remaining, entry)
}
}
state.Blocked = remaining
return e.saveState(&state)
}
func (e *Engine) saveAllowedEntry(entry AllowedEntry) error {
if entry.Source == "" {
entry.Source = InferProvenance("allow", entry.Reason)
}
if key, ok := canonicalIPKey(entry.IP); ok {
entry.IP = key
}
state := e.loadStateFile()
upsertAllowedEntryInState(&state, entry)
return e.saveState(&state)
}
func (e *Engine) removeAllowedState(ip string) error {
state := e.loadStateFile()
var remaining []AllowedEntry
for _, entry := range state.Allowed {
if !sameIPString(entry.IP, ip) {
remaining = append(remaining, entry)
}
}
state.Allowed = remaining
return e.saveState(&state)
}
// removeAllowedStateBySource reports whether the final source was removed.
func (e *Engine) removeAllowedStateBySource(ip, source string) (bool, error) {
state := e.loadStateFile()
found, ipGone := removeAllowedSourceFromState(&state, ip, source)
if !found {
return false, nil
}
if err := e.saveState(&state); err != nil {
return false, err
}
return ipGone, nil
}
func removeAllowedSourceFromState(state *FirewallState, ip, source string) (found, ipGone bool) {
remaining := state.Allowed[:0]
ipGone = true
for _, entry := range state.Allowed {
if sameIPString(entry.IP, ip) && entry.Source == source {
found = true
continue
}
remaining = append(remaining, entry)
if sameIPString(entry.IP, ip) {
ipGone = false
}
}
state.Allowed = remaining
return found, ipGone
}
func (e *Engine) saveSubnetEntry(entry SubnetEntry) error {
if entry.Source == "" {
entry.Source = InferProvenance("block_subnet", entry.Reason)
}
state := e.loadStateFile()
if !addSubnetEntryIfMissingInState(&state, entry) {
return nil
}
return e.saveState(&state)
}
func (e *Engine) isSubnetBlockedStateLocked(cidr string) bool {
e.ensureStateCacheLocked()
_, ok := e.blockedCIDRIndex[cidr]
return ok
}
func (e *Engine) removeSubnetState(cidr string) error {
state := e.loadStateFile()
var remaining []SubnetEntry
for _, entry := range state.BlockedNet {
if entry.CIDR != cidr {
remaining = append(remaining, entry)
}
}
state.BlockedNet = remaining
return e.saveState(&state)
}
// IP helpers (nextIP, lastIPInRange, fileExistsFirewall) moved to ip_helpers.go (no build tag).
// loadCountryCIDRs reads CIDR ranges from a country file.
// Expected format: one CIDR per line in {dbPath}/{CODE}.cidr
func loadCountryCIDRs(dbPath, countryCode string) []nftables.SetElement {
file := filepath.Join(dbPath, strings.ToUpper(countryCode)+".cidr")
// #nosec G304 -- filepath.Join under operator-configured GeoIP dbPath.
data, err := os.ReadFile(file)
if err != nil {
return nil
}
var elements []nftables.SetElement
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
_, network, err := net.ParseCIDR(line)
if err != nil {
continue
}
start := network.IP.To4()
end := lastIPInRange(network)
if start != nil && end != nil {
elements = appendIntervalSetElements(elements, start, end)
}
}
return elements
}
// loadCountryCIDRs6 reads IPv6 CIDR ranges from {dbPath}/{CODE}.cidr6 and
// builds interval set elements. v4 CIDRs in the file (if any) are skipped so
// only IPv6 ranges reach the IPv6 set.
func loadCountryCIDRs6(dbPath, countryCode string) []nftables.SetElement {
file := filepath.Join(dbPath, strings.ToUpper(countryCode)+".cidr6")
// #nosec G304 -- filepath.Join under operator-configured GeoIP dbPath.
data, err := os.ReadFile(file)
if err != nil {
return nil
}
var elements []nftables.SetElement
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
_, network, err := net.ParseCIDR(line)
if err != nil {
continue
}
if network.IP.To4() != nil {
continue // not an IPv6 range
}
start := network.IP.To16()
end := lastIPInRange(network)
if start != nil && end != nil {
elements = appendIntervalSetElements(elements, start, end.To16())
}
}
return elements
}
//go:build linux
package firewall
import (
"bytes"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"net"
"reflect"
"strings"
"time"
"github.com/google/nftables"
"github.com/pidginhost/csm/internal/actionlog"
)
// AttachLifecycle is a construction-time injection, not a migration. The owner
// must already have initialized the complete store and own the daemon lock.
func (e *Engine) AttachLifecycle(l *Lifecycle) error {
e.mu.Lock()
defer e.mu.Unlock()
if l == nil || l.Store == nil || e.lifecycle != nil {
return errors.New("invalid firewall lifecycle attachment")
}
state, revision, err := l.Store.ReadFirewallState()
if err != nil {
return err
}
if l.Audit == nil {
l.Audit = writeFirewallActionAudit
}
e.lifecycle = l
e.stateRevision = revision
e.installCommittedCache(state)
return nil
}
func writeFirewallActionAudit(a FirewallAction) error {
result := actionlog.Result(a.Phase)
return actionlog.WriteDurable(actionlog.Record{ActionID: a.Request.ID, ActionVersion: a.AuditVersion, IncidentID: a.Request.IncidentID, UndoOf: a.Request.UndoOf, Timestamp: a.UpdatedAt, Op: firewallAuditOperation(a.Request), Action: a.Request.Operation, Actor: actionlog.Actor(a.Request.Actor), ActorDetail: a.Request.ActorDetail, FindingID: a.Request.FindingID, Target: a.Request.Target, Reason: a.Request.Reason, Result: result, Error: a.Detail})
}
func (e *Engine) lifecycleEnabled() bool { e.mu.Lock(); defer e.mu.Unlock(); return e.lifecycle != nil }
func (e *Engine) installCommittedCache(state FirewallState) {
state = copyFirewallState(state)
e.stateCache = &state
e.applyExpiryLocked()
e.rebuildIndexLocked()
}
func (e *Engine) lifecycleReadyLocked() error {
if e.lifecycle == nil {
return nil
}
pending, err := e.lifecycle.Store.PendingFirewallActions()
if err != nil {
return err
}
if len(pending) > 0 {
return fmt.Errorf("%w: %s", ErrActionUnknown, pending[0].Request.ID)
}
state, _, err := e.readCommittedStateLocked()
if err != nil {
return err
}
e.installCommittedCache(state)
return nil
}
func (e *Engine) RecoverActions() error {
e.mu.Lock()
defer e.mu.Unlock()
if e.lifecycle == nil {
return errors.New("firewall lifecycle unavailable")
}
recoverErr := e.lifecycle.Recover(engineActionKernel{e})
state, _, readErr := e.readCommittedStateLocked()
if readErr == nil {
e.installCommittedCache(state)
}
return errors.Join(recoverErr, readErr)
}
// PendingActions reports the actions recovery could not settle. While any
// exists the engine refuses new mutations, so this is what an operator needs
// to see before deciding anything.
func (e *Engine) PendingActions() ([]FirewallAction, error) {
e.mu.Lock()
defer e.mu.Unlock()
if e.lifecycle == nil {
return nil, errors.New("firewall lifecycle unavailable")
}
return e.lifecycle.Store.PendingFirewallActions()
}
// ResolveAction records an outcome an operator established by hand for an
// action the kernel cannot prove. The committed cache follows the outcome,
// whether it came from the operator or from the kernel proving it after all.
func (e *Engine) ResolveAction(id, outcome, detail string) (FirewallAction, error) {
e.mu.Lock()
defer e.mu.Unlock()
if e.lifecycle == nil {
return FirewallAction{}, errors.New("firewall lifecycle unavailable")
}
a, err := e.lifecycle.Resolve(id, outcome, detail, engineActionKernel{e})
state, _, readErr := e.readCommittedStateLocked()
if readErr == nil {
e.installCommittedCache(state)
}
return a, errors.Join(durableActionOutcome(a, err), readErr)
}
func (e *Engine) UndoAction(req ActionRequest) (FirewallAction, error) {
e.mu.Lock()
defer e.mu.Unlock()
if e.lifecycle == nil {
return FirewallAction{}, errors.New("firewall lifecycle unavailable")
}
a, err := e.lifecycle.Undo(req, engineActionKernel{e})
err = durableActionOutcome(a, err)
if a.Phase == "verified" {
current, _, readErr := e.readCommittedStateLocked()
if readErr != nil {
return a, errors.Join(err, readErr)
}
e.installCommittedCache(current)
}
return a, err
}
func (e *Engine) runDurableLocked(req ActionRequest, budget *ScanAdmission, next FirewallState) error {
if e.stateReadErr != nil {
return e.stateReadErr
}
expectedRevision := e.stateRevision
if err := e.lifecycleReadyLocked(); err != nil {
return err
}
prior, revision, err := e.readCommittedStateLocked()
if err != nil {
return err
}
if revision != expectedRevision {
return fmt.Errorf("%w: planning snapshot changed", ErrStateConflict)
}
if req.ID == "" {
req.ID = rand.Text()
}
req = normalizeActionRequest(req)
plan := FirewallAction{Request: req, Before: prior, After: next, Revision: revision, CreatedAt: time.Now(), Budget: budget}
if prepareErr := e.prepareActionKernel(&plan); prepareErr != nil {
return prepareErr
}
a, err := e.lifecycle.Execute(plan, engineActionKernel{e})
if a.Phase == "verified" {
e.installCommittedCache(a.After)
}
return durableActionOutcome(a, err)
}
func stateFingerprint(value any) string {
raw, _ := json.Marshal(value)
sum := sha256.Sum256(raw)
return "csm:" + hex.EncodeToString(sum[:])
}
func (e *Engine) actionSet(name string) *nftables.Set {
for _, set := range []*nftables.Set{e.setBlocked, e.setBlocked6, e.setAllowed, e.setAllowed6, e.setBlockedNet, e.setBlockedNet6} {
if set != nil && set.Name == name {
return set
}
}
return nil
}
func actionSetNames(ipv6 bool) []string {
names := []string{"blocked_ips", "allowed_ips", "blocked_nets"}
if ipv6 {
names = append(names, "blocked_ips6", "allowed_ips6", "blocked_nets6")
}
return names
}
func desiredActionSets(state FirewallState, ipv6 bool, now time.Time) map[string][]ActionElement {
out := make(map[string][]ActionElement)
for _, name := range actionSetNames(ipv6) {
out[name] = nil
}
for _, b := range state.Blocked {
if !b.ExpiresAt.IsZero() && !b.ExpiresAt.After(now) {
continue
}
ip := net.ParseIP(b.IP)
if ip == nil {
continue
}
name := "blocked_ips"
key := ip.To4()
if key == nil {
if !ipv6 {
continue
}
name += "6"
key = ip.To16()
}
out[name] = append(out[name], ActionElement{Key: key, Comment: stateFingerprint(b), ExpiresAt: b.ExpiresAt})
}
allowed := make(map[string][]AllowedEntry)
for _, a := range state.Allowed {
if a.ExpiresAt.IsZero() || a.ExpiresAt.After(now) {
allowed[a.IP] = append(allowed[a.IP], a)
}
}
for address, rows := range allowed {
ip := net.ParseIP(address)
if ip == nil {
continue
}
name := "allowed_ips"
key := ip.To4()
if key == nil {
if !ipv6 {
continue
}
name += "6"
key = ip.To16()
}
out[name] = append(out[name], ActionElement{Key: key, Comment: stateFingerprint(rows)})
}
var activeSubnets []SubnetEntry
for _, entry := range state.BlockedNet {
if entry.ExpiresAt.IsZero() || entry.ExpiresAt.After(now) {
activeSubnets = append(activeSubnets, entry)
}
}
v4, v6 := subnetIntervalElements(activeSubnets, ipv6, now)
for i, elems := range [][]nftables.SetElement{v4, v6} {
name := "blocked_nets"
if i == 1 {
name += "6"
}
for _, elem := range elems {
comment := ""
// Interval-end sentinels cannot carry element userdata. The start
// marker identifies the complete union, including source expiry.
if !elem.IntervalEnd {
comment = stateFingerprint(activeSubnets)
}
out[name] = append(out[name], ActionElement{Key: elem.Key, End: elem.IntervalEnd, Comment: comment})
}
}
return out
}
func (e *Engine) prepareActionKernel(a *FirewallAction) error {
now := time.Now()
a.CreatedAt = now
desired := desiredActionSets(a.After, e.setBlocked6 != nil, now)
expected := desiredActionSets(a.Before, e.setBlocked6 != nil, now)
unpruned := desiredActionSets(a.Before, e.setBlocked6 != nil, time.Time{})
conn, connErr := newLifecycleConn(e)
if connErr != nil {
return connErr
}
names := requiredActionSets(*a)
a.KernelBefore = nil
a.KernelAfter = nil
for _, name := range names {
actual, readErr := readActionSet(conn, name)
if readErr != nil {
return readErr
}
elems := actual.elements
if !actual.exists && e.actionSet(name) != nil {
return fmt.Errorf("%w: set %s disappeared", ErrActionUnknown, name)
}
if !matchActionElements(expected[name], elems, now) && !matchActionElements(unpruned[name], elems, now) {
return fmt.Errorf("%w: live set %s differs from committed state", ErrActionUnknown, name)
}
before := ActionSet{Name: name, Exists: actual.exists}
now := time.Now()
for _, elem := range elems {
entry := ActionElement{Key: bytes.Clone(elem.Key), End: elem.IntervalEnd, Comment: elem.Comment}
if elem.Timeout > 0 {
entry.ExpiresAt = now.Add(elem.Expires)
for _, known := range expected[name] {
if bytes.Equal(known.Key, entry.Key) && known.Comment == entry.Comment {
entry.ExpiresAt = known.ExpiresAt
break
}
}
}
before.Elements = append(before.Elements, entry)
}
a.KernelBefore = append(a.KernelBefore, before)
a.KernelAfter = append(a.KernelAfter, ActionSet{Name: name, Exists: actual.exists, Elements: desired[name]})
}
return nil
}
type engineActionKernel struct{ e *Engine }
func (k engineActionKernel) ApplyFirewallAction(a FirewallAction) error {
if a.Request.Operation == "apply" {
conn, err := newLifecycleConn(k.e)
if err != nil {
return err
}
priorConn := k.e.conn
k.e.conn = conn
defer func() { k.e.conn = priorConn }()
return k.e.applyRulesetLocked(a.Ruleset.Marker)
}
if err := k.applyActionSets(a, true); err != nil {
if !isNftNotFound(err) {
return err
}
// The kernel expires timed elements on its own, so a delete can name an
// element that is already gone. That batch changed nothing, so rewrite
// the complete set instead of reporting an uncertain outcome.
return k.applyActionSets(a, false)
}
return nil
}
// applyActionSets writes the intended effect of one action. A whole-set rewrite
// costs one message per retained element, which on a busy host is most of the
// work a single block does, so unchanged elements are left alone where the set
// allows it.
func (k engineActionKernel) applyActionSets(a FirewallAction, delta bool) error {
conn, connErr := newLifecycleConn(k.e)
if connErr != nil {
return connErr
}
now := time.Now()
for i, state := range a.KernelAfter {
if !state.Exists {
continue
}
set := k.e.actionSet(state.Name)
if set == nil {
return fmt.Errorf("firewall recovery set unavailable: %s", state.Name)
}
if !delta || i >= len(a.KernelBefore) || !deltaApplicable(a.KernelBefore[i], state) {
conn.FlushSet(set)
if err := addElementsChunked(conn, set, actionElements(state.Elements, now)); err != nil {
return err
}
continue
}
add, remove := actionElementDelta(a.KernelBefore[i], state, now)
// Each element list has a uint16 netlink attribute length. Bound
// deletion messages too, while keeping all chunks in one transaction.
for offset := 0; offset < len(remove); offset += 1000 {
end := min(offset+1000, len(remove))
if err := conn.SetDeleteElements(set, remove[offset:end]); err != nil {
return err
}
}
if err := addElementsChunked(conn, set, add); err != nil {
return err
}
}
return conn.Flush()
}
// deltaApplicable reports whether a set can be changed element by element.
// Interval sets carry paired start and end markers whose union changes shape
// when any member changes, so those are always rewritten whole.
func deltaApplicable(before, after ActionSet) bool {
if !before.Exists || before.Name != after.Name {
return false
}
for _, set := range []ActionSet{before, after} {
for _, elem := range set.Elements {
if elem.End {
return false
}
}
}
return true
}
func actionElementKey(elem ActionElement) string {
return fmt.Sprintf("%x/%t", elem.Key, elem.End)
}
// actionElementDelta returns the elements to add and to remove so the live set
// matches the intended effect. An element whose comment or expiry changed is
// removed and re-added in the same batch, which nftables applies atomically.
func actionElementDelta(before, after ActionSet, now time.Time) (add, remove []nftables.SetElement) {
live := make(map[string]ActionElement, len(before.Elements))
for _, elem := range before.Elements {
live[actionElementKey(elem)] = elem
}
intended := make(map[string]ActionElement, len(after.Elements))
for _, elem := range actionElements(after.Elements, now) {
entry := ActionElement{Key: elem.Key, End: elem.IntervalEnd, Comment: elem.Comment}
if elem.Timeout > 0 {
entry.ExpiresAt = now.Add(elem.Timeout)
}
key := actionElementKey(entry)
intended[key] = entry
prior, held := live[key]
if held && prior.Comment == entry.Comment && prior.ExpiresAt.Equal(entry.ExpiresAt) {
continue
}
if held {
remove = append(remove, nftables.SetElement{Key: prior.Key, IntervalEnd: prior.End})
}
add = append(add, elem)
}
for _, elem := range before.Elements {
if _, wanted := intended[actionElementKey(elem)]; wanted {
continue
}
remove = append(remove, nftables.SetElement{Key: elem.Key, IntervalEnd: elem.End})
}
return add, remove
}
func actionElements(entries []ActionElement, now time.Time) []nftables.SetElement {
var out []nftables.SetElement
for _, entry := range entries {
timeout := time.Duration(0)
if !entry.ExpiresAt.IsZero() {
timeout = entry.ExpiresAt.Sub(now)
if timeout < time.Millisecond {
continue
}
}
out = append(out, nftables.SetElement{Key: entry.Key, IntervalEnd: entry.End, Comment: entry.Comment, Timeout: timeout})
}
return out
}
func matchActionElements(expected []ActionElement, actual []nftables.SetElement, at time.Time) bool {
remaining := make(map[string]ActionElement)
for _, entry := range expected {
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(at) {
continue
}
remaining[fmt.Sprintf("%x/%t", entry.Key, entry.End)] = entry
}
for _, elem := range actual {
key := fmt.Sprintf("%x/%t", elem.Key, elem.IntervalEnd)
want, ok := remaining[key]
if !ok || want.Comment != elem.Comment {
return false
}
if want.ExpiresAt.IsZero() {
if elem.Timeout != 0 {
return false
}
} else {
if elem.Timeout <= 0 {
return false
}
delta := at.Add(elem.Expires).Sub(want.ExpiresAt)
if delta > 2*time.Second || delta < -2*time.Second {
return false
}
}
delete(remaining, key)
}
return len(remaining) == 0
}
func (k engineActionKernel) ObserveFirewallAction(a FirewallAction) (ActionObservation, error) {
if err := k.ValidateEvidence(a); err != nil {
return ActionObservation{}, err
}
var generation uint32
if a.Request.Operation == "apply" {
var genErr error
generation, genErr = k.e.actionGeneration()
if genErr != nil {
return ActionObservation{}, genErr
}
}
result := ActionObservation{Before: true, After: true}
for i, before := range a.KernelBefore {
after := a.KernelAfter[i]
if before.Name != after.Name {
return ActionObservation{}, errors.New("mismatched firewall kernel evidence")
}
conn, connErr := newLifecycleConn(k.e)
if connErr != nil {
return ActionObservation{}, connErr
}
started := time.Now()
actual, err := readActionSet(conn, before.Name)
if err != nil {
return ActionObservation{}, err
}
elems := actual.elements
if actual.exists != before.Exists {
result.Before = false
}
if actual.exists != after.Exists {
result.After = false
}
if time.Since(started) > 2*time.Second {
return ActionObservation{}, errors.New("firewall verification took too long")
}
result.Before = result.Before && matchActionElements(before.Elements, elems, started)
result.After = result.After && matchActionElements(after.Elements, elems, started)
}
if a.Request.Operation == "apply" {
rules, err := k.e.rulesetEvidenceMatches(a, generation)
if err != nil {
return ActionObservation{}, err
}
result.Before = result.Before && rules.Before
result.After = result.After && rules.After
}
return result, nil
}
// BlockIPRequest carries durable request identity and source-specific admission
// policy through the existing automatic and manual safety gates.
func (e *Engine) BlockIPRequest(req ActionRequest, budget *ScanAdmission) (BlockOutcome, error) {
canonical, err := canonicalFirewallIP(req.Target)
if err != nil {
recordBlockOutcome(req.Target, req.Reason, req.TTL, BlockOutcomeNoop, err, !req.Automatic, req.FindingID)
return BlockOutcomeNoop, err
}
req.Target = canonical
req.Operation = "block"
req = normalizeActionRequest(req)
e.mu.Lock()
if found, replayErr := e.replayActionLocked(req); found || replayErr != nil {
e.mu.Unlock()
return BlockOutcomeNoop, replayErr
}
readyErr := e.lifecycleReadyLocked()
e.mu.Unlock()
if readyErr != nil {
return BlockOutcomeNoop, readyErr
}
if req.Automatic {
return e.blockIPOutcomeRequest(req, budget)
}
return e.blockIPLockedRequest(req.Target, req.Reason, req.TTL, false, false, req, budget)
}
func firewallAuditOperation(req ActionRequest) string {
if req.Automatic {
return "respond.block_ip"
}
return "operate.manual_firewall"
}
func (e *Engine) DurableActionsEnabled() bool { return e.lifecycleEnabled() }
func (e *Engine) FirewallScanBudget(window string) (int, error) {
e.mu.Lock()
defer e.mu.Unlock()
if e.lifecycle == nil {
return 0, errors.New("durable firewall budget unavailable")
}
return e.lifecycle.Store.ReadFirewallScanBudget(window)
}
type liveActionSet struct {
exists bool
elements []nftables.SetElement
}
func readActionSet(conn *nftables.Conn, name string) (liveActionSet, error) {
set, err := conn.GetSetByName(&nftables.Table{Name: "csm", Family: nftables.TableFamilyINet}, name)
if err != nil {
if isNftNotFound(err) {
return liveActionSet{}, nil
}
return liveActionSet{}, err
}
elems, err := conn.GetSetElements(set)
return liveActionSet{exists: true, elements: elems}, err
}
func (e *Engine) applyDurableLocked() error {
state, revision, err := e.readCommittedStateLocked()
if err != nil {
return err
}
req := ActionRequest{ID: rand.Text(), Operation: "apply", Target: "csm", Actor: string(actionlog.DefaultActor()), Source: SourceSystem}
plan := FirewallAction{Request: req, Before: state, After: state, Revision: revision, CreatedAt: time.Now()}
if evidenceErr := e.prepareRulesetEvidence(&plan); evidenceErr != nil {
return evidenceErr
}
plan.CreatedAt = time.Now()
desired := desiredActionSets(state, e.cfg.IPv6, plan.CreatedAt)
conn, err := newLifecycleConn(e)
if err != nil {
return err
}
for _, name := range actionSetNames(true) {
actual, readErr := readActionSet(conn, name)
if readErr != nil {
return readErr
}
before := ActionSet{Name: name, Exists: actual.exists}
now := time.Now()
for _, elem := range actual.elements {
entry := ActionElement{Key: bytes.Clone(elem.Key), End: elem.IntervalEnd, Comment: elem.Comment}
if elem.Timeout > 0 {
entry.ExpiresAt = now.Add(elem.Expires)
}
before.Elements = append(before.Elements, entry)
}
elements, exists := desired[name]
plan.KernelBefore = append(plan.KernelBefore, before)
plan.KernelAfter = append(plan.KernelAfter, ActionSet{Name: name, Exists: exists, Elements: elements})
}
generation, genErr := e.actionGeneration()
if genErr != nil {
return genErr
}
if generation != plan.Ruleset.Generation {
return fmt.Errorf("%w: ruleset changed during planning", ErrActionUnknown)
}
a, executeErr := e.lifecycle.Execute(plan, engineActionKernel{e})
return durableActionOutcome(a, executeErr)
}
func requiredActionSets(a FirewallAction) []string {
if a.Request.Operation == "apply" {
return actionSetNames(true)
}
// A removal must prove the target set even when committed state already
// omits it. Otherwise an untracked live element is reported as removed.
blocked := a.Request.Operation == "unblock" || a.Request.Operation == "flush"
allowed := a.Request.Operation == "remove_allow"
subnets := a.Request.Operation == "unblock_subnet"
var names []string
if blocked || !reflect.DeepEqual(a.Before.Blocked, a.After.Blocked) {
names = append(names, "blocked_ips", "blocked_ips6")
}
if allowed || !reflect.DeepEqual(a.Before.Allowed, a.After.Allowed) {
names = append(names, "allowed_ips", "allowed_ips6")
}
if subnets || !reflect.DeepEqual(a.Before.BlockedNet, a.After.BlockedNet) {
names = append(names, "blocked_nets", "blocked_nets6")
}
return names
}
func (k engineActionKernel) ValidateEvidence(a FirewallAction) error {
names := requiredActionSets(a)
if len(names) != len(a.KernelBefore) || len(names) != len(a.KernelAfter) {
return errors.New("incomplete firewall kernel evidence")
}
expected := desiredActionSets(a.After, true, a.CreatedAt)
for i, name := range names {
before, after := a.KernelBefore[i], a.KernelAfter[i]
if before.Name != name || after.Name != name {
return errors.New("mismatched firewall kernel evidence")
}
if !strings.HasSuffix(name, "6") && !after.Exists {
return errors.New("missing IPv4 firewall set")
}
for _, set := range []ActionSet{before, after} {
if !set.Exists && len(set.Elements) != 0 {
return errors.New("absent set has elements")
}
keys := make(map[string]bool)
for _, elem := range set.Elements {
size := 4
if strings.HasSuffix(name, "6") {
size = 16
}
key := fmt.Sprintf("%x/%t", elem.Key, elem.End)
if len(elem.Key) != size || keys[key] {
return errors.New("invalid or duplicate firewall element")
}
keys[key] = true
}
}
if after.Exists && !sameActionElements(after.Elements, expected[name]) {
return errors.New("kernel evidence differs from intended effect")
}
}
return nil
}
func sameActionElements(a, b []ActionElement) bool {
if len(a) != len(b) {
return false
}
entries := make(map[string]ActionElement, len(a))
for _, entry := range a {
entries[fmt.Sprintf("%x/%t", entry.Key, entry.End)] = entry
}
for _, entry := range b {
prior, ok := entries[fmt.Sprintf("%x/%t", entry.Key, entry.End)]
if !ok || prior.Comment != entry.Comment || !prior.ExpiresAt.Equal(entry.ExpiresAt) {
return false
}
}
return true
}
func (k engineActionKernel) PrepareFirewallUndo(plan *FirewallAction) error {
return k.e.prepareActionKernel(plan)
}
// The caller holds e.mu. Durable paths deliver their versioned journal event.
func (e *Engine) legacyActionAuditLocked(action, target, reason, source string, ttl time.Duration) {
if e.lifecycle == nil {
AppendAudit(e.statePath, action, target, reason, source, ttl)
}
}
func (e *Engine) legacyFileAuditLocked(action, target, reason, source string, ttl time.Duration) {
if e.lifecycle == nil {
appendAudit(e.statePath, action, target, reason, source, ttl)
}
}
// Failure defers run after the engine lock is released. Cleanup loops already
// hold that lock and use their durable outcome instead.
func (e *Engine) legacyFirewallFailure(action, target, reason, source string, ttl time.Duration, err error) {
if e.shouldLegacyOutcome(err) {
recordFirewallFailure(action, target, reason, source, ttl, err)
}
}
func normalizeActionRequest(req ActionRequest) ActionRequest {
if req.Source == "" {
req.Source = InferProvenance(req.Operation, req.Reason)
}
if req.Actor == "" {
_, actor := firewallActionOp(req.Operation, req.Source)
req.Actor = string(actor)
}
return req
}
// A retry observes the original plan and refreshes current committed state.
// Historical After must never overwrite cache state from a later action.
func (e *Engine) replayActionLocked(req ActionRequest) (bool, error) {
if e.lifecycle == nil || req.ID == "" {
return false, nil
}
previous, err := e.lifecycle.Store.ReadFirewallAction(req.ID)
if errors.Is(err, ErrActionMissing) {
return false, nil
}
if err != nil {
return false, err
}
if previous.Request != req {
return true, ErrStateConflict
}
result, executeErr := e.lifecycle.Execute(previous, engineActionKernel{e})
if executeErr != nil && result.Phase != "verified" {
return true, durableActionOutcome(result, executeErr)
}
current, _, err := e.readCommittedStateLocked()
if err == nil {
e.installCommittedCache(current)
}
return true, errors.Join(durableActionOutcome(result, executeErr), err)
}
func (e *Engine) savePortPolicyLocked(state *FirewallState, req ActionRequest) error {
if e.lifecycle != nil {
return e.runDurableLocked(req, nil, *state)
}
return e.saveState(state)
}
// An admitted action owns its versioned outcome; rejected attempts retain the
// existing best-effort refusal record at their public operation boundary.
type durableAttemptError struct{ error }
func (err *durableAttemptError) Unwrap() error { return err.error }
func durableActionOutcome(a FirewallAction, err error) error {
if err == nil {
return nil
}
if a.Phase == "verified" {
return &durableAttemptError{errors.Join(ErrActionAuditPending, err)}
}
if a.Request.ID != "" || errors.Is(err, ErrActionUnknown) || errors.Is(err, ErrStateCommitUncertain) {
return &durableAttemptError{err}
}
return err
}
func (e *Engine) shouldLegacyOutcome(err error) bool {
if !e.lifecycleEnabled() {
return true
}
var admitted *durableAttemptError
return err != nil && !errors.As(err, &admitted)
}
// Pin mutation plans to the snapshot they read. A later successful read must
// not silently attach an older After state to a newer committed revision.
func (e *Engine) readCommittedStateLocked() (FirewallState, uint64, error) {
state, revision, err := e.lifecycle.Store.ReadFirewallState()
e.stateReadErr = err
if err == nil {
e.stateRevision = revision
}
return state, revision, err
}
package firewall
import (
"bufio"
"fmt"
"io"
"net"
"net/http"
"os"
"path/filepath"
"strings"
"time"
)
const (
geoIPBaseURL = "https://raw.githubusercontent.com/herrbischoff/country-ip-blocks/master/ipv4/"
geoIPBaseURLv6 = "https://raw.githubusercontent.com/herrbischoff/country-ip-blocks/master/ipv6/"
)
// UpdateGeoIPDB downloads country CIDR lists from a public source.
// Creates one file per country code per family: {dbPath}/{CC}.cidr (IPv4)
// and {dbPath}/{CC}.cidr6 (IPv6). IPv6 is best-effort so a country with no v6
// allocation does not fail the update. The return value is the number of CIDR
// files refreshed.
func UpdateGeoIPDB(dbPath string, countryCodes []string) (int, error) {
client := &http.Client{Timeout: 30 * time.Second}
return updateGeoIPDBWithClient(dbPath, countryCodes, client)
}
func updateGeoIPDBWithClient(dbPath string, countryCodes []string, client *http.Client) (int, error) {
if err := os.MkdirAll(dbPath, 0700); err != nil {
return 0, fmt.Errorf("creating geoip directory: %w", err)
}
updated := 0
for _, code := range countryCodes {
code = strings.ToLower(strings.TrimSpace(code))
if len(code) != 2 {
continue
}
cc := strings.ToUpper(code)
if downloadCIDRFile(client, geoIPBaseURL+code+".cidr", filepath.Join(dbPath, cc+".cidr")) {
updated++
}
if downloadCIDRFile(client, geoIPBaseURLv6+code+".cidr", filepath.Join(dbPath, cc+".cidr6")) {
updated++
}
}
return updated, nil
}
// countryCIDRMaxBytes bounds one country CIDR download.
const countryCIDRMaxBytes = 16 << 20
// countCIDRLines returns how many lines of path parse as a CIDR, skipping
// blanks and comments.
func countCIDRLines(path string) int {
// #nosec G304 -- path is the download's own temp file under dbPath.
data, err := os.ReadFile(path)
if err != nil {
return 0
}
count := 0
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if _, _, err := net.ParseCIDR(line); err == nil {
count++
}
}
return count
}
// downloadCIDRFile fetches url into outPath atomically. Returns false (and
// logs) on any HTTP, write, or invalid-payload condition so the caller can
// treat each family independently.
func downloadCIDRFile(client *http.Client, url, outPath string) bool {
resp, err := client.Get(url)
if err != nil {
fmt.Fprintf(os.Stderr, "geoip: error downloading %s: %v\n", url, err)
return false
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
fmt.Fprintf(os.Stderr, "geoip: %s returned HTTP %d\n", url, resp.StatusCode)
return false
}
tmpPath := outPath + ".tmp"
// #nosec G304 -- filepath.Join under operator-configured dbPath; code from fixed list.
f, err := os.Create(tmpPath)
if err != nil {
fmt.Fprintf(os.Stderr, "geoip: error creating %s: %v\n", tmpPath, err)
return false
}
// Bounded: a country list is well under a megabyte; an upstream that
// streams more is not serving the list.
n, copyErr := io.Copy(f, io.LimitReader(resp.Body, countryCIDRMaxBytes+1))
closeErr := f.Close()
if copyErr != nil {
_ = os.Remove(tmpPath)
fmt.Fprintf(os.Stderr, "geoip: error writing %s: %v\n", tmpPath, copyErr)
return false
}
if closeErr != nil {
_ = os.Remove(tmpPath)
fmt.Fprintf(os.Stderr, "geoip: error closing %s: %v\n", tmpPath, closeErr)
return false
}
if n > countryCIDRMaxBytes {
_ = os.Remove(tmpPath)
fmt.Fprintf(os.Stderr, "geoip: %s exceeds %d bytes, keeping previous file\n", url, countryCIDRMaxBytes)
return false
}
// Validate before install: a 200 with no parseable CIDR (an HTML
// interstitial, a moved path) must not replace the last good file, or
// the next restart builds an empty country set while the update
// reported success.
if cidrs := countCIDRLines(tmpPath); cidrs == 0 {
_ = os.Remove(tmpPath)
fmt.Fprintf(os.Stderr, "geoip: %s holds no CIDR entries (%d bytes), keeping previous file\n", url, n)
return false
}
if err := os.Rename(tmpPath, outPath); err != nil {
_ = os.Remove(tmpPath)
fmt.Fprintf(os.Stderr, "geoip: error installing %s: %v\n", outPath, err)
return false
}
fmt.Fprintf(os.Stderr, "geoip: updated %s (%d bytes)\n", outPath, n)
return true
}
// LookupIP finds which country CIDR files contain the given IP.
// Returns matching country codes.
func LookupIP(dbPath string, ip string) []string {
parsed := net.ParseIP(ip)
if parsed == nil {
return nil
}
// IPv4 (incl. v4-mapped) matches against .cidr files; IPv6 against .cidr6.
var suffix string
var needle net.IP
if ip4 := parsed.To4(); ip4 != nil {
suffix = ".cidr"
needle = ip4
} else {
suffix = ".cidr6"
needle = parsed.To16()
}
if needle == nil {
return nil
}
entries, err := os.ReadDir(dbPath)
if err != nil {
return nil
}
var matches []string
for _, entry := range entries {
if entry.IsDir() || !strings.HasSuffix(entry.Name(), suffix) {
continue
}
code := strings.TrimSuffix(entry.Name(), suffix)
if containsIP(filepath.Join(dbPath, entry.Name()), needle) {
matches = append(matches, code)
}
}
return matches
}
func containsIP(cidrFile string, ip net.IP) bool {
// #nosec G304 -- cidrFile is filepath.Join under operator-configured dbPath.
f, err := os.Open(cidrFile)
if err != nil {
return false
}
defer f.Close()
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") {
continue
}
_, network, err := net.ParseCIDR(line)
if err != nil {
continue
}
if network.Contains(ip) {
return true
}
}
return false
}
//go:build linux
package firewall
import (
"bytes"
"fmt"
"net"
"slices"
"time"
"github.com/google/nftables"
)
// normalizeIntervalElements unions independent start/end pairs of one address
// family. Stored source entries stay separate; only their kernel view is merged.
// A start without an end represents a range through the family's last address.
func normalizeIntervalElements(elements []nftables.SetElement) []nftables.SetElement {
type ipRange struct{ start, end []byte }
ranges := make([]ipRange, 0, len(elements)/2+1)
for i := 0; i < len(elements); i++ {
r := ipRange{start: elements[i].Key}
if i+1 < len(elements) && elements[i+1].IntervalEnd {
i++
r.end = bytes.Clone(elements[i].Key)
for j := len(r.end) - 1; j >= 0; j-- {
r.end[j]--
if r.end[j] != 0xff {
break
}
}
} else {
r.end = bytes.Repeat([]byte{0xff}, len(r.start))
}
ranges = append(ranges, r)
}
slices.SortFunc(ranges, func(a, b ipRange) int { return bytes.Compare(a.start, b.start) })
var out []nftables.SetElement
for i := 0; i < len(ranges); {
merged := ranges[i]
i++
for i < len(ranges) {
next, hasNext := nextIntervalKey(merged.end)
if hasNext && bytes.Compare(ranges[i].start, next) > 0 {
break
}
if bytes.Compare(ranges[i].end, merged.end) > 0 {
merged.end = ranges[i].end
}
i++
}
out = appendIntervalSetElements(out, merged.start, merged.end)
}
return out
}
// Preserve the set's key width even when an IPv6 boundary has the numeric
// shape of an IPv4-mapped address.
func nextIntervalKey(key []byte) ([]byte, bool) {
next := bytes.Clone(key)
for i := len(next) - 1; i >= 0; i-- {
next[i]++
if next[i] != 0 {
return next, true
}
}
return nil, false
}
func subnetIntervalElements(entries []SubnetEntry, ipv6 bool, now time.Time) (v4, v6 []nftables.SetElement) {
for _, entry := range entries {
if !entry.ExpiresAt.IsZero() && !now.Before(entry.ExpiresAt) {
continue
}
_, network, err := net.ParseCIDR(entry.CIDR)
if err != nil {
continue
}
// The public block path refuses default routes. Old state must obey
// that same lockout guard when it is restored at startup.
if ones, _ := network.Mask.Size(); ones == 0 {
continue
}
end := lastIPInRange(network)
if start := network.IP.To4(); start != nil {
v4 = appendIntervalSetElements(v4, start, end)
} else if ipv6 {
v6 = appendIntervalSetElements(v6, network.IP.To16(), end)
}
}
return normalizeIntervalElements(v4), normalizeIntervalElements(v6)
}
// Rebuild the union in one transaction. Deleting the original boundaries of
// an overlapping entry could otherwise remove a different source's protection.
func (e *Engine) replaceBlockedSubnetSets(entries []SubnetEntry) error {
// A pending config reload may differ from the installed ruleset. Keep
// every family that still has a live set until Apply replaces the table.
v4, v6 := subnetIntervalElements(entries, e.setBlockedNet6 != nil, time.Now())
// A non-lasting connection only dials at Flush and cannot fail here. A
// separate batch also makes any queue error discard the entire update.
conn, _ := nftables.New(nftables.WithSockOptions(applyNFTSocketBuffer), nftables.WithNetNSFd(e.conn.NetNS), nftables.WithTestDial(e.conn.TestDial))
for _, family := range []struct {
set *nftables.Set
elements []nftables.SetElement
}{{e.setBlockedNet, v4}, {e.setBlockedNet6, v6}} {
if family.set == nil {
continue
}
conn.FlushSet(family.set)
if err := addElementsChunked(conn, family.set, family.elements); err != nil {
return err
}
}
if err := conn.Flush(); err != nil {
return fmt.Errorf("replacing subnet sets: %w", err)
}
return nil
}
func (e *Engine) updateSubnetMutation(prior, next FirewallState, req ActionRequest, budget *ScanAdmission) error {
if e.lifecycle != nil {
return e.runDurableLocked(req, budget, next)
}
if err := e.persistFirewallIntent(prior, next); err != nil {
return fmt.Errorf("persisting subnet change: %w", err)
}
if err := e.replaceBlockedSubnetSets(next.BlockedNet); err != nil {
if restoreErr := e.saveState(&prior); restoreErr != nil {
return fmt.Errorf("%w (state restore failed: %v)", err, restoreErr)
}
return err
}
return nil
}
package firewall
import (
"net"
"os"
)
// subnetCovering returns the first CIDR in entries that contains ip, if any.
// Pure helper for BlockedSubnetCovering so the containment logic is testable
// without a kernel-attached engine.
func subnetCovering(entries []SubnetEntry, ip string) (string, bool) {
parsed := net.ParseIP(ip)
if parsed == nil {
return "", false
}
for _, entry := range entries {
_, network, err := net.ParseCIDR(entry.CIDR)
if err != nil {
continue
}
if network.Contains(parsed) {
return entry.CIDR, true
}
}
return "", false
}
// nextIP returns the IP address immediately following the given IP.
// When ip is the all-ones address for its family, nextIP clamps to ip
// instead of wrapping to all-zeros. nftables interval keys use a separate
// width-preserving helper because their address family is already fixed.
func nextIP(ip net.IP) net.IP {
next, _ := nextIPSafe(ip)
return next
}
// nextIPSafe returns the successor of ip plus whether the successor
// exists in the same address family.
func nextIPSafe(ip net.IP) (net.IP, bool) {
next := canonicalIPBytes(ip)
if next == nil {
return nil, false
}
for i := len(next) - 1; i >= 0; i-- {
if next[i] != 0xff {
next[i]++
return next, true
}
next[i] = 0
}
return canonicalIPBytes(ip), false
}
func canonicalIPBytes(ip net.IP) net.IP {
if ip4 := ip.To4(); ip4 != nil {
out := make(net.IP, net.IPv4len)
copy(out, ip4)
return out
}
if ip16 := ip.To16(); ip16 != nil {
out := make(net.IP, net.IPv6len)
copy(out, ip16)
return out
}
return nil
}
// lastIPInRange returns the last IP address in a CIDR range.
func lastIPInRange(network *net.IPNet) net.IP {
ip := network.IP.To4()
if ip == nil {
ip = network.IP.To16()
}
if ip == nil {
return nil
}
mask := network.Mask
if len(ip) == net.IPv4len && len(mask) == net.IPv6len {
mask = mask[net.IPv6len-net.IPv4len:]
}
last := make(net.IP, len(ip))
for i := range ip {
if i < len(mask) {
last[i] = ip[i] | ^mask[i]
} else {
last[i] = ip[i]
}
}
return last
}
// fileExistsFirewall checks if a file exists.
func fileExistsFirewall(path string) bool {
_, err := os.Stat(path)
return err == nil
}
package firewall
import (
"encoding/json"
"errors"
"fmt"
"time"
)
func sameActionState(a, b FirewallState) bool {
left, err := json.Marshal(a)
if err != nil {
return false
}
right, err := json.Marshal(b)
return err == nil && string(left) == string(right)
}
func actionResult(a FirewallAction) error {
switch a.Phase {
case "verified":
return nil
case "failed":
return fmt.Errorf("%w: %s", ErrActionFailed, a.Request.ID)
default:
return fmt.Errorf("%w: %s", ErrActionUnknown, a.Request.ID)
}
}
// Execute admits once and executes only a fresh request. Repeated IDs inspect
// the original evidence; they never renew the plan or replay a kernel batch.
func (l *Lifecycle) Execute(plan FirewallAction, kernel ActionKernel) (FirewallAction, error) {
l.mu.Lock()
defer l.mu.Unlock()
return l.execute(plan, kernel)
}
func (l *Lifecycle) execute(plan FirewallAction, kernel ActionKernel) (FirewallAction, error) {
a, fresh, err := l.Store.AdmitFirewallAction(plan)
if err != nil {
return FirewallAction{}, err
}
if !fresh {
if a.Phase != "verified" && a.Phase != "failed" {
a, err = l.reconcile(a, kernel, nil)
} else {
err = actionResult(a)
}
return a, errors.Join(err, l.deliverAudit())
}
a, err = l.Store.TransitionFirewallAction(a.Request.ID, "executing", "", time.Now())
if err != nil {
return FirewallAction{}, fmt.Errorf("%w: execution admission: %w", ErrActionUnknown, err)
}
observed, inspectErr := kernel.ObserveFirewallAction(a)
if inspectErr != nil || !observed.Before {
cause := errors.New("target changed before execution")
if inspectErr != nil {
cause = inspectErr
}
result, unknownErr := l.markUnknown(a, cause)
return result, errors.Join(unknownErr, l.deliverAudit())
}
kernelErr := kernel.ApplyFirewallAction(a)
if kernelErr == nil {
a, err = l.Store.TransitionFirewallAction(a.Request.ID, "applied", "", time.Now())
if err != nil {
return FirewallAction{}, fmt.Errorf("%w: applied outcome persistence: %w", ErrActionUnknown, err)
}
}
a, err = l.reconcile(a, kernel, kernelErr)
return a, errors.Join(err, l.deliverAudit())
}
func (l *Lifecycle) markUnknown(a FirewallAction, cause error) (FirewallAction, error) {
result, err := l.Store.TransitionFirewallAction(a.Request.ID, "unknown", cause.Error(), time.Now())
return result, errors.Join(fmt.Errorf("%w: %s: %w", ErrActionUnknown, a.Request.ID, cause), err)
}
func (l *Lifecycle) reconcile(a FirewallAction, kernel ActionKernel, kernelErr error) (FirewallAction, error) {
observed, err := kernel.ObserveFirewallAction(a)
if err != nil {
return l.markUnknown(a, err)
}
phase := ""
switch {
case observed.After && !observed.Before:
phase = "verified"
case observed.Before && !observed.After:
phase = "failed"
case observed.After && observed.Before && a.Phase == "applied":
phase = "verified"
case observed.After && observed.Before && a.Phase == "planned":
phase = "failed"
default:
return l.markUnknown(a, errors.New("kernel does not prove a complete before or after state"))
}
detail := ""
if kernelErr != nil {
detail = kernelErr.Error()
}
result, err := l.Store.TransitionFirewallAction(a.Request.ID, phase, detail, time.Now())
if err != nil {
return FirewallAction{}, fmt.Errorf("%w: outcome persistence: %w", ErrActionUnknown, err)
}
return result, actionResult(result)
}
// Recover proves incomplete outcomes without executing host mutations. A
// proven rejected action is recovered successfully, though its result is failed.
func (l *Lifecycle) Recover(kernel ActionKernel) error {
l.mu.Lock()
defer l.mu.Unlock()
actions, err := l.Store.PendingFirewallActions()
if err != nil {
return err
}
var failures []error
for _, a := range actions {
if _, err := l.reconcile(a, kernel, nil); err != nil && !errors.Is(err, ErrActionFailed) {
failures = append(failures, err)
}
}
failures = append(failures, l.deliverAudit())
return errors.Join(failures...)
}
// Resolve records the outcome an operator established by hand for an action
// the kernel could not prove. It is the only way out of an uncertain outcome,
// which otherwise refuses every later mutation. Kernel evidence still wins: a
// proven outcome is recorded instead of the asserted one, and a settled action
// is never rewritten. Detail names who decided, and reaches the audit trail.
func (l *Lifecycle) Resolve(id, outcome, detail string, kernel ActionKernel) (FirewallAction, error) {
l.mu.Lock()
defer l.mu.Unlock()
if outcome != "verified" && outcome != "failed" {
return FirewallAction{}, fmt.Errorf("firewall action outcome must be verified or failed, not %q", outcome)
}
if detail == "" {
return FirewallAction{}, errors.New("firewall action resolution must record who decided")
}
a, err := l.Store.ReadFirewallAction(id)
if err != nil {
return FirewallAction{}, err
}
if a.Phase != "unknown" {
return FirewallAction{}, fmt.Errorf("%w: action %s is %s, not uncertain", ErrStateConflict, id, a.Phase)
}
proven, reconcileErr := l.reconcile(a, kernel, nil)
if proven.Phase == "verified" || proven.Phase == "failed" {
return proven, errors.Join(ignoreActionResult(reconcileErr), l.deliverAudit())
}
// Only a durably recorded unknown permits an operator assertion. A
// failed transition can mean the kernel proved the opposite outcome;
// a storage error is not permission to replace that evidence.
if proven.Phase != "unknown" {
return proven, reconcileErr
}
result, err := l.Store.TransitionFirewallAction(id, outcome, detail, time.Now())
if err != nil {
return FirewallAction{}, fmt.Errorf("%w: operator outcome persistence: %w", ErrActionUnknown, err)
}
return result, l.deliverAudit()
}
// A resolution reports what the outcome turned out to be. A proven rejection
// is a complete answer to the operator's question, not a failure to answer it.
func ignoreActionResult(err error) error {
if errors.Is(err, ErrActionFailed) {
return nil
}
return err
}
func (l *Lifecycle) DeliverAudit() error {
l.mu.Lock()
defer l.mu.Unlock()
return l.deliverAudit()
}
func (l *Lifecycle) deliverAudit() error {
actions, err := l.Store.FirewallAuditPending()
if err != nil {
return err
}
for _, a := range actions {
if l.Audit == nil {
return errors.New("durable firewall audit writer unavailable")
}
if err := l.Audit(a); err != nil {
return fmt.Errorf("firewall audit pending for %s: %w", a.Request.ID, err)
}
if err := l.Store.AcknowledgeFirewallAudit(a.Request.ID, a.AuditVersion); err != nil {
return err
}
}
return nil
}
// Undo admits an inverse state transition, never a command string. The complete
// recorded result must still be current; this also protects eviction victims.
func (l *Lifecycle) Undo(req ActionRequest, kernel ActionKernel) (FirewallAction, error) {
l.mu.Lock()
defer l.mu.Unlock()
if req.Operation != "undo" || req.UndoOf == "" {
return FirewallAction{}, errors.New("invalid firewall undo request")
}
if existing, err := l.Store.ReadFirewallAction(req.ID); err == nil {
if existing.Request != req {
return FirewallAction{}, ErrStateConflict
}
return l.execute(existing, kernel)
} else if !errors.Is(err, ErrActionMissing) {
return FirewallAction{}, err
}
original, err := l.Store.ReadFirewallAction(req.UndoOf)
if err != nil {
return FirewallAction{}, err
}
if original.Request.Operation == "apply" {
return FirewallAction{}, errors.New("ruleset configuration undo requires the confirmed-apply rollback protocol")
}
if original.Phase != "verified" {
return FirewallAction{}, errors.New("only a verified firewall action can be undone")
}
state, revision, err := l.Store.ReadFirewallState()
if err != nil {
return FirewallAction{}, err
}
if !sameActionState(state, original.After) {
return FirewallAction{}, fmt.Errorf("%w: undo state changed", ErrStateConflict)
}
observed, err := kernel.ObserveFirewallAction(original)
if err != nil {
return FirewallAction{}, err
}
if !observed.After {
return FirewallAction{}, fmt.Errorf("%w: undo target changed", ErrStateConflict)
}
plan := FirewallAction{Request: req, Before: state, After: original.Before, Revision: revision, CreatedAt: time.Now(), KernelBefore: original.KernelAfter, KernelAfter: original.KernelBefore}
if planner, ok := kernel.(interface{ PrepareFirewallUndo(*FirewallAction) error }); ok {
if err := planner.PrepareFirewallUndo(&plan); err != nil {
return FirewallAction{}, err
}
}
return l.execute(plan, kernel)
}
//go:build linux
package firewall
import (
"encoding/binary"
"errors"
"fmt"
"math"
"github.com/google/nftables"
"github.com/mdlayher/netlink"
"github.com/mdlayher/netlink/nltest"
"golang.org/x/sys/unix"
)
// The namespace generation advances on a committed nftables transaction.
// Together with a unique inert chain in our atomic Apply batch it proves both
// that the intended batch committed and that no later ruleset edit intervened.
// Unrelated namespace edits conservatively require recovery review too.
func (e *Engine) actionGeneration() (uint32, error) {
var socket *netlink.Conn
if e.conn.TestDial != nil {
socket = nltest.Dial(e.conn.TestDial)
} else {
var err error
socket, err = netlink.Dial(unix.NETLINK_NETFILTER, &netlink.Config{NetNS: e.conn.NetNS})
if err != nil {
return 0, err
}
if err = applyLifecycleSocketDeadline(socket); err != nil {
return 0, err
}
}
defer func() { _ = socket.Close() }()
messages, err := socket.Execute(netlink.Message{Header: netlink.Header{Type: netlink.HeaderType(unix.NFNL_SUBSYS_NFTABLES<<8 | unix.NFT_MSG_GETGEN), Flags: netlink.Request}, Data: []byte{unix.AF_UNSPEC, 0, 0, 0}})
if err != nil {
return 0, err
}
for _, message := range messages {
if len(message.Data) < 4 {
continue
}
attrs, decodeErr := netlink.UnmarshalAttributes(message.Data[4:])
if decodeErr != nil {
return 0, decodeErr
}
for _, attr := range attrs {
if attr.Type&0x3fff == unix.NFTA_GEN_ID && len(attr.Data) == 4 {
return binary.BigEndian.Uint32(attr.Data), nil
}
}
}
return 0, errors.New("missing nftables generation evidence")
}
func (e *Engine) prepareRulesetEvidence(a *FirewallAction) error {
generation, err := e.actionGeneration()
if err != nil {
return err
}
if generation == math.MaxUint32 {
return errors.New("nftables generation rollover requires a new recovery baseline")
}
a.Ruleset = &ActionRuleset{Generation: generation, Marker: "csm_action_" + stateFingerprint(a.Request)[4:]}
return nil
}
func (e *Engine) rulesetEvidenceMatches(a FirewallAction, generation uint32) (ActionObservation, error) {
if a.Ruleset == nil || a.Ruleset.Marker != "csm_action_"+stateFingerprint(a.Request)[4:] || a.Ruleset.Generation == math.MaxUint32 {
return ActionObservation{}, errors.New("invalid ruleset recovery evidence")
}
conn, err := newLifecycleConn(e)
if err != nil {
return ActionObservation{}, err
}
chains, err := conn.ListChainsOfTableFamily(nftables.TableFamilyINet)
if err != nil {
return ActionObservation{}, err
}
marker := false
for _, chain := range chains {
if chain.Table.Name == "csm" && chain.Name == a.Ruleset.Marker {
marker = true
}
}
after, err := e.actionGeneration()
if err != nil {
return ActionObservation{}, err
}
if after != generation {
return ActionObservation{}, fmt.Errorf("%w: ruleset changed during verification", ErrActionUnknown)
}
return ActionObservation{Before: generation == a.Ruleset.Generation && !marker, After: generation == a.Ruleset.Generation+1 && marker}, nil
}
//go:build linux
package firewall
import (
"errors"
"fmt"
"time"
"github.com/google/nftables"
"github.com/mdlayher/netlink"
)
// Bound the actual socket exchange, including both request writes and replies.
// This option runs on every transient dial, so a reused nftables connection
// receives a fresh deadline for each operation.
func applyLifecycleSocketDeadline(connection *netlink.Conn) error {
if err := connection.SetDeadline(time.Now().Add(2 * time.Second)); err != nil {
// nftables does not close a socket when one of its socket options fails.
_ = connection.Close()
return fmt.Errorf("setting firewall lifecycle socket deadline: %w", err)
}
return nil
}
func newLifecycleConn(engine *Engine) (*nftables.Conn, error) {
if engine == nil || engine.conn == nil {
return nil, errors.New("firewall lifecycle transport unavailable")
}
options := []nftables.ConnOption{
nftables.WithNetNSFd(engine.conn.NetNS),
nftables.WithTestDial(engine.conn.TestDial),
nftables.WithSockOptions(applyNFTSocketBuffer),
}
// The explicitly injected nltest transport has no OS socket or deadline
// support. Never suppress deadline errors on a live transport.
if engine.conn.TestDial == nil {
options = append(options, nftables.WithSockOptions(applyLifecycleSocketDeadline))
}
return nftables.New(options...)
}
//go:build linux
package firewall
import (
"errors"
"fmt"
)
type stateDurabilityError struct{ cause error }
func (e *stateDurabilityError) Error() string {
return fmt.Sprintf("state file matches requested change but durability is unconfirmed: %v", e.cause)
}
func (e *stateDurabilityError) Unwrap() error { return e.cause }
// A write can fail after rename made the next state visible. Restore the
// prior intent before returning so background cleanup retains its retry rows.
// Both failures must be exposed if storage also refuses the rollback.
func (e *Engine) persistFirewallIntent(prior, next FirewallState) error {
err := e.saveState(&next)
if err == nil {
return nil
}
var uncertain *stateDurabilityError
if !errors.As(err, &uncertain) {
return err
}
if restoreErr := e.saveState(&prior); restoreErr != nil {
return fmt.Errorf("partial failure: %w (state restore failed: %w)", err, restoreErr)
}
return fmt.Errorf("intent write failed; previous state restored: %w", err)
}
//go:build linux
package firewall
import (
"fmt"
"strings"
"github.com/google/nftables"
"github.com/google/nftables/binaryutil"
"github.com/google/nftables/expr"
)
type portFloodIPFamily struct {
name string
nfproto byte
sourceOffset uint32
sourceLen uint32
}
var (
portFloodIPv4 = portFloodIPFamily{name: "v4", nfproto: 2, sourceOffset: 12, sourceLen: 4}
portFloodIPv6 = portFloodIPFamily{name: "v6", nfproto: 10, sourceOffset: 8, sourceLen: 16}
)
// buildPortFloodExprs returns the nftables expressions for one per-port
// flood-protection rule. The rule rate-limits new TCP/UDP connections to
// pf.Port per source address by updating a dynamic meter set. The caller
// supplies a family-specific meter, so IPv4 and IPv6 do not share buckets.
//
// Returning nil signals the caller to skip rule creation (zero rate, missing
// meter, or zero-port).
func buildPortFloodExprs(pf PortFloodRule, meter *nftables.Set, family portFloodIPFamily) []expr.Any {
if meter == nil || pf.Hits <= 0 || pf.Seconds <= 0 || pf.Port <= 0 {
return nil
}
proto := byte(6) // TCP
if pf.Proto == "udp" {
proto = 17
}
// hits/seconds -> packets per minute (multiply first to keep precision).
ratePerMin := uint64(pf.Hits) * 60 / uint64(pf.Seconds)
if ratePerMin < 1 {
ratePerMin = 1
}
burst := uint32(ratePerMin / 4)
if burst < 2 {
burst = 2
}
return []expr.Any{
// Restrict to the family that matches the meter key type.
&expr.Meta{Key: expr.MetaKeyNFPROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{family.nfproto}},
// Match new connections only.
&expr.Ct{Register: 1, SourceRegister: false, Key: expr.CtKeySTATE},
&expr.Bitwise{
SourceRegister: 1, DestRegister: 1, Len: 4,
Mask: binaryutil.NativeEndian.PutUint32(expr.CtStateBitNEW),
Xor: binaryutil.NativeEndian.PutUint32(0),
},
&expr.Cmp{Op: expr.CmpOpNeq, Register: 1, Data: binaryutil.NativeEndian.PutUint32(0)},
// L4 protocol filter (TCP or UDP).
&expr.Meta{Key: expr.MetaKeyL4PROTO, Register: 1},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: []byte{proto}},
// Destination port.
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseTransportHeader, Offset: 2, Len: 2},
&expr.Cmp{Op: expr.CmpOpEq, Register: 1, Data: binaryutil.BigEndian.PutUint16(portU16(pf.Port))},
// Load source address into reg1; this is the meter key.
&expr.Payload{DestRegister: 1, Base: expr.PayloadBaseNetworkHeader, Offset: family.sourceOffset, Len: family.sourceLen},
// Update meter entry for this source IP, evaluating its own token bucket.
&expr.Dynset{
SrcRegKey: 1,
SetName: meter.Name,
SetID: meter.ID,
Operation: 1, // NFT_DYNSET_OP_UPDATE
Exprs: []expr.Any{
&expr.Limit{Type: expr.LimitTypePkts, Rate: ratePerMin, Unit: expr.LimitTimeMinute, Burst: burst, Over: true},
},
},
&expr.Verdict{Kind: expr.VerdictDrop},
}
}
func (e *Engine) portFloodRuleExprs(pf PortFloodRule, meter *nftables.Set, family portFloodIPFamily) []expr.Any {
exprs := buildPortFloodExprs(pf, meter, family)
if exprs == nil || !isMailTCP(pf) {
return exprs
}
var exempt []expr.Any
switch {
case family == portFloodIPv4 && e.setDOSExempt != nil:
exempt = e.dosExemptV4Lookup(1)
case family == portFloodIPv6 && e.setDOSExempt6 != nil:
exempt = e.dosExemptV6Lookup(1)
default:
return exprs
}
guarded := make([]expr.Any, 0, len(exprs)+len(exempt))
guarded = append(guarded, exprs[:2]...)
guarded = append(guarded, exempt...)
return append(guarded, exprs[2:]...)
}
func portFloodMeterName(pf PortFloodRule, family portFloodIPFamily) string {
return fmt.Sprintf("meter_pf_%s_%d_%s", portFloodProto(pf), pf.Port, family.name)
}
func portFloodProto(pf PortFloodRule) string {
if pf.Proto == "udp" {
return "udp"
}
return "tcp"
}
// isMailTCP reports whether pf is a TCP rule for a standard mail relay or
// submission port (25, 465, 587). Only these ports receive the DoS-exempt
// inverted-lookup guard; all other TCP ports and UDP rules are left unchanged.
func isMailTCP(pf PortFloodRule) bool {
if !strings.EqualFold(pf.Proto, "tcp") {
return false
}
switch pf.Port {
case 25, 465, 587:
return true
default:
return false
}
}
//go:build linux
package firewall
import "fmt"
type ipProtectedError struct {
msg string
}
func (e ipProtectedError) Error() string {
return e.msg
}
func (e ipProtectedError) Unwrap() error {
return ErrIPProtected
}
func ipProtectedErrorf(format string, args ...any) error {
return ipProtectedError{msg: fmt.Sprintf(format, args...)}
}
package firewall
import "strings"
const (
SourceUnknown = "unknown"
SourceWebUI = "web_ui"
SourceCLI = "cli"
SourceAutoResponse = "auto_response"
SourceChallenge = "challenge"
SourceWhitelist = "whitelist"
SourceDynDNS = "dyndns"
SourceSystem = "system"
)
// InferProvenance classifies a firewall entry source from structured action/reason text.
// This keeps provenance logic centralized instead of spreading fragile string checks
// throughout the web UI and firewall call sites.
func InferProvenance(action, reason string) string {
action = strings.ToLower(strings.TrimSpace(action))
reason = strings.ToLower(strings.TrimSpace(reason))
switch {
case strings.Contains(reason, "dyndns:"):
return SourceDynDNS
case strings.Contains(reason, "passed challenge"),
strings.Contains(reason, "challenge-timeout"),
strings.Contains(reason, "challenge timeout"):
return SourceChallenge
case strings.Contains(reason, "temp whitelist"),
strings.Contains(reason, "whitelist"),
strings.Contains(reason, "bulk whitelist"),
strings.Contains(reason, "customer ip"):
return SourceWhitelist
case strings.Contains(reason, "auto-block"),
strings.Contains(reason, "permbblock"),
strings.Contains(reason, "permblock"),
strings.Contains(reason, "auto-netblock"):
return SourceAutoResponse
case strings.Contains(reason, "via cli"):
return SourceCLI
case strings.Contains(reason, "via csm web ui"),
strings.Contains(reason, "via ui"),
strings.Contains(reason, "allowed from firewall lookup"),
strings.Contains(reason, "manual block"):
return SourceWebUI
case action == "temp_allow_expired":
return SourceSystem
case action == "flush":
return SourceSystem
default:
return SourceUnknown
}
}
// Package rollback implements the firewall settings tentative-apply
// workflow: a save with a deadline that auto-reverts unless the operator
// confirms before the timer expires. The manager survives daemon restarts
// (state is persisted in bbolt) so that the apply itself can take down the
// daemon for a config reload without losing the rollback intent.
package rollback
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"os"
"os/exec"
"sync"
"time"
"github.com/pidginhost/csm/internal/actionlog"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/store"
)
// MinTimeout and MaxTimeout bound the operator-supplied window. The lower
// bound exists so a misclick cannot leave the operator one second to react;
// the upper bound caps how long a botched apply can sit on disk before
// auto-recovery kicks in.
const (
MinTimeout = 1 * time.Minute
MaxTimeout = 30 * time.Minute
DefaultTimeout = 5 * time.Minute
)
// Restarter performs a daemon restart. Production wires this to a
// systemctl exec; tests substitute a fake.
type Restarter func(ctx context.Context) error
// Status describes the current pending rollback for status APIs and the
// Web UI banner. AppliedAt and ExpiresAt are UTC; SecondsRemaining is a
// derived hint computed at call time.
type Status struct {
Pending bool `json:"pending"`
AppliedAt time.Time `json:"applied_at,omitzero"`
ExpiresAt time.Time `json:"expires_at,omitzero"`
SecondsRemaining int64 `json:"remaining_seconds,omitempty"`
AppliedBy string `json:"applied_by,omitempty"`
PrevHash string `json:"prev_hash,omitempty"`
NewHash string `json:"new_hash,omitempty"`
identity string
}
// Manager owns the active timer and serialises Apply/Confirm/Revert/Recover
// against each other. Storage and restart are injected so the manager can be
// driven in tests without a real bbolt or systemctl.
type Manager struct {
mu sync.Mutex
db *store.DB
configPath string
restart Restarter
now func() time.Time
timer *time.Timer
}
// Process-wide singleton. The daemon installs one Manager at startup; the
// Web UI handlers and the control-socket commands look it up by calling
// Global so they do not need to thread the manager through every layer.
var (
globalMu sync.Mutex
global *Manager
)
// SetGlobal installs the process-wide manager. Safe to call from
// daemon startup; subsequent SetGlobal calls overwrite.
func SetGlobal(m *Manager) {
globalMu.Lock()
global = m
globalMu.Unlock()
}
// Global returns the installed manager or nil when none has been set
// (e.g. inside CLI commands that load config but never start the
// daemon). Callers must nil-check.
func Global() *Manager {
globalMu.Lock()
defer globalMu.Unlock()
return global
}
// NewManager wires a manager. Use SystemctlRestart for the production
// restart path. now defaults to time.Now when nil.
func NewManager(db *store.DB, configPath string, restart Restarter, now func() time.Time) *Manager {
if now == nil {
now = time.Now
}
return &Manager{
db: db,
configPath: configPath,
restart: restart,
now: now,
}
}
// SystemctlRestart issues `systemctl restart csm.service`. The context
// timeout caps how long we wait for systemctl itself; the daemon restart
// it triggers is asynchronous from systemctl's perspective.
func SystemctlRestart(ctx context.Context) error {
// #nosec G204 -- fixed argv, no operator input interpolated.
out, err := exec.CommandContext(ctx, "systemctl", "restart", "csm.service").CombinedOutput()
if err != nil {
return fmt.Errorf("systemctl restart csm: %w (%s)", err, string(out))
}
return nil
}
// HashYAML returns the sha256 hex digest of yaml bytes. Used so the
// Web UI and CLI can show operators the before/after hash without the
// full file contents.
func HashYAML(data []byte) string {
sum := sha256.Sum256(data)
return "sha256:" + hex.EncodeToString(sum[:])
}
// clampTimeout returns a timeout within [MinTimeout, MaxTimeout]; it
// substitutes DefaultTimeout for zero so callers can pass 0 to mean
// "use the default."
func clampTimeout(d time.Duration) time.Duration {
if d == 0 {
return DefaultTimeout
}
if d < MinTimeout {
return MinTimeout
}
if d > MaxTimeout {
return MaxTimeout
}
return d
}
// Apply records prevYAML as the snapshot to restore on expiry, computes
// the expiry deadline, persists the rollback entry, and arms the local
// timer. The caller is responsible for holding ConfigWriteMutex across Apply
// and the newYAML write, then triggering the restart.
//
// applyBy is logged with the rollback record (e.g. token name or "cli")
// so audits can trace the source.
func (m *Manager) Apply(prevYAML, newYAML []byte, timeout time.Duration, applyBy string) (Status, error) {
m.mu.Lock()
defer m.mu.Unlock()
if existing, ok := m.db.GetFirewallRollback(); ok {
return statusFromRecord(existing, m.now()), fmt.Errorf("rollback already pending; confirm or revert first")
}
timeout = clampTimeout(timeout)
now := m.now().UTC()
rb := store.FirewallRollback{
PrevYAML: prevYAML,
PrevHash: HashYAML(prevYAML),
NewHash: HashYAML(newYAML),
AppliedAt: now,
ExpiresAt: now.Add(timeout),
AppliedBy: applyBy,
}
if err := m.db.SaveFirewallRollback(rb); err != nil {
return Status{}, fmt.Errorf("persist rollback: %w", err)
}
m.armTimerLocked(timeout)
return statusFromRecord(rb, m.now()), nil
}
// Confirm drops the pending rollback. The new config stays on disk;
// no daemon restart is required. Idempotent: confirming with no
// pending entry is a no-op.
func (m *Manager) Confirm() error {
expected := m.Status()
if !expected.Pending {
return nil
}
return m.ConfirmIfCurrent(expected)
}
// ConfirmIfCurrent confirms only the rollback represented by expected. A
// confirmation request can race another operator confirming the old record
// and applying a replacement; it must not silently clear the newer record.
// Taking the config writer lock also keeps confirmation from landing between
// tentative-apply's rollback staging and its csm.yaml write.
func (m *Manager) ConfirmIfCurrent(expected Status) error {
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
return m.confirmIfCurrent(expected)
}
// AbortApplyIfCurrent drops a staged rollback when tentative-apply could not
// write the new config. The caller must already hold ConfigWriteMutex for the
// complete staging/write transaction.
func (m *Manager) AbortApplyIfCurrent(expected Status) error {
return m.confirmIfCurrent(expected)
}
func (m *Manager) confirmIfCurrent(expected Status) error {
m.mu.Lock()
defer m.mu.Unlock()
rb, ok := m.db.GetFirewallRollback()
if !ok || !sameRollbackStatus(expected, rb) {
return fmt.Errorf("rollback changed before confirm; refresh and retry")
}
if m.timer != nil {
m.timer.Stop()
m.timer = nil
}
return m.db.ClearFirewallRollback()
}
// Revert restores the snapshot to disk and triggers a daemon restart.
// Returns an error if there is no pending rollback so callers can
// surface a clean "nothing to revert" message instead of silently
// succeeding.
func (m *Manager) Revert(ctx context.Context) error {
expected := m.Status()
if !expected.Pending {
return fmt.Errorf("no pending rollback")
}
return m.RevertIfCurrent(ctx, expected)
}
// RevertIfCurrent restores only the rollback represented by expected. The
// manager re-checks after acquiring the config writer lock because a blocked
// revert must not restore a replacement rollback created while it waited.
func (m *Manager) RevertIfCurrent(ctx context.Context, expected Status) error {
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
m.mu.Lock()
defer m.mu.Unlock()
rb, ok := m.db.GetFirewallRollback()
if !ok {
return fmt.Errorf("no pending rollback")
}
if !sameRollbackStatus(expected, rb) {
return fmt.Errorf("rollback changed before revert; refresh and retry")
}
return m.applyRevertLocked(ctx, rb)
}
// RecoverOnStartup is called once during daemon startup. If a pending
// rollback exists, the manager either fires the revert immediately
// (timer already expired while the daemon was down) or rearms a
// time.AfterFunc for the remaining window.
//
// The bool return is true when an immediate revert was performed so
// the caller can decide whether to bail out of further startup work
// while the restart it triggers takes effect.
func (m *Manager) RecoverOnStartup(ctx context.Context) (reverted bool, err error) {
m.mu.Lock()
rb, ok := m.db.GetFirewallRollback()
if !ok {
m.mu.Unlock()
return false, nil
}
now := m.now()
if !now.Before(rb.ExpiresAt) {
m.mu.Unlock()
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
m.mu.Lock()
defer m.mu.Unlock()
rb, ok = m.db.GetFirewallRollback()
if !ok {
return false, nil
}
now = m.now()
if now.Before(rb.ExpiresAt) {
m.armTimerLocked(rb.ExpiresAt.Sub(now))
return false, nil
}
err := m.applyRevertLocked(ctx, rb)
if err != nil {
return false, err
}
return true, nil
}
remaining := rb.ExpiresAt.Sub(now)
m.armTimerLocked(remaining)
m.mu.Unlock()
return false, nil
}
// Status reports the pending rollback for /api/v1/.../rollback and
// the CLI status command. Pending=false means "nothing in flight".
func (m *Manager) Status() Status {
m.mu.Lock()
defer m.mu.Unlock()
rb, ok := m.db.GetFirewallRollback()
if !ok {
return Status{}
}
return statusFromRecord(rb, m.now())
}
func (m *Manager) armTimerLocked(d time.Duration) {
if m.timer != nil {
m.timer.Stop()
}
m.timer = time.AfterFunc(d, func() {
// Build a fresh context so the AfterFunc goroutine has a
// usable deadline; the original ctx from RecoverOnStartup
// would have been cancelled by the time this fires.
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
if err := m.timerExpired(ctx); err != nil {
fmt.Fprintf(os.Stderr, "rollback: timer expiry failed: %v\n", err)
}
})
}
func (m *Manager) timerExpired(ctx context.Context) error {
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
m.mu.Lock()
defer m.mu.Unlock()
rb, ok := m.db.GetFirewallRollback()
if !ok {
return nil
}
now := m.now()
if now.Before(rb.ExpiresAt) {
m.armTimerLocked(rb.ExpiresAt.Sub(now))
return nil
}
return m.applyRevertLocked(ctx, rb)
}
// applyRevertLocked restores the snapshot bytes to the config path, clears
// the pending record, and only then triggers the daemon restart. The order
// is load-bearing: systemctl restart kills this very process, so anything
// sequenced after a successful restart never runs. Clearing after the
// restart left the record in place, and every subsequent boot found it
// expired, reverted, and restarted again - a permanent restart loop. Once
// the snapshot is on disk any later daemon start converges on the reverted
// config, so a failed restart (surfaced as an error) loses nothing but
// immediacy. A failed clear aborts before the restart for the same reason:
// restarting with the record still present is what loops.
// Caller must hold m.mu.
func (m *Manager) applyRevertLocked(ctx context.Context, rb store.FirewallRollback) (resultErr error) {
rec := actionlog.Record{Op: "operate.manual_firewall", Action: "rollback_config", Target: m.configPath, Before: actionlog.Stat(m.configPath), Result: actionlog.Failed}
recorded := false
defer func() {
if !recorded {
if resultErr != nil {
rec.Error = resultErr.Error()
}
actionlog.Write(rec)
}
}()
if m.timer != nil {
m.timer.Stop()
m.timer = nil
}
if len(rb.PrevYAML) == 0 {
return fmt.Errorf("rollback record has empty prev_yaml; cannot restore")
}
if err := integrity.WriteConfigBytesAtomic(m.configPath, rb.PrevYAML); err != nil {
return fmt.Errorf("restore previous config: %w", err)
}
rec.Result = actionlog.Applied
rec.After = actionlog.Stat(m.configPath)
if err := m.db.ClearFirewallRollback(); err != nil {
return fmt.Errorf("clear rollback before restart: %w", err)
}
// Restart may terminate this process before it returns. The config change
// has completed and must be recorded before requesting the restart.
actionlog.Write(rec)
recorded = true
if m.restart != nil {
if err := m.restart(ctx); err != nil {
return fmt.Errorf("trigger restart after revert (config already restored on disk): %w", err)
}
}
return nil
}
func statusFromRecord(rb store.FirewallRollback, now time.Time) Status {
remaining := int64(rb.ExpiresAt.Sub(now).Seconds())
if remaining < 0 {
remaining = 0
}
return Status{
Pending: true,
AppliedAt: rb.AppliedAt,
ExpiresAt: rb.ExpiresAt,
SecondsRemaining: remaining,
AppliedBy: rb.AppliedBy,
PrevHash: rb.PrevHash,
NewHash: rb.NewHash,
identity: rollbackIdentity(rb),
}
}
func sameRollbackStatus(expected Status, current store.FirewallRollback) bool {
return expected.Pending && expected.identity != "" && expected.identity == rollbackIdentity(current)
}
func rollbackIdentity(rb store.FirewallRollback) string {
payload := fmt.Sprintf("%x\n%q\n%q\n%q\n%q\n%q",
rb.PrevYAML, rb.PrevHash, rb.NewHash,
rb.AppliedAt.UTC().Format(time.RFC3339Nano),
rb.ExpiresAt.UTC().Format(time.RFC3339Nano), rb.AppliedBy)
sum := sha256.Sum256([]byte(payload))
return hex.EncodeToString(sum[:])
}
package firewall
// RuleCounts holds firewall rule cardinalities sourced from the engine
// state file, which is the authoritative store. The parallel bbolt
// fw:* buckets are written only during migration, so anything counting
// live rules must read the engine, not the store. Expired temp bans are
// excluded.
type RuleCounts struct {
Blocked int
Allowed int
Subnets int
PortAllowed int
}
// Total returns the sum across all rule categories.
func (c RuleCounts) Total() int {
return c.Blocked + c.Allowed + c.Subnets + c.PortAllowed
}
//go:build linux
package firewall
import "net"
func countRuleEntries(state FirewallState, ipv6Enabled bool) RuleCounts {
return RuleCounts{
Blocked: countBlockedRules(state.Blocked, ipv6Enabled),
Allowed: countAllowedRules(state.Allowed, ipv6Enabled),
Subnets: countSubnetRules(state.BlockedNet, ipv6Enabled),
PortAllowed: countPortAllowRules(state.PortAllowed, ipv6Enabled),
}
}
func countBlockedRules(entries []BlockedEntry, ipv6Enabled bool) int {
seen := make(map[string]struct{}, len(entries))
for _, entry := range entries {
key, ok := ruleIPKey(entry.IP, ipv6Enabled)
if !ok {
continue
}
seen[key] = struct{}{}
}
return len(seen)
}
func countAllowedRules(entries []AllowedEntry, ipv6Enabled bool) int {
seen := make(map[string]struct{}, len(entries))
for _, entry := range entries {
key, ok := ruleIPKey(entry.IP, ipv6Enabled)
if !ok {
continue
}
seen[key] = struct{}{}
}
return len(seen)
}
func countSubnetRules(entries []SubnetEntry, ipv6Enabled bool) int {
seen := make(map[string]struct{}, len(entries))
for _, entry := range entries {
_, network, err := net.ParseCIDR(entry.CIDR)
if err != nil {
continue
}
if network.IP.To4() == nil && !ipv6Enabled {
continue
}
seen[network.String()] = struct{}{}
}
return len(seen)
}
func countPortAllowRules(entries []PortAllowEntry, ipv6Enabled bool) int {
count := 0
for _, entry := range entries {
if _, ok := ruleIPKey(entry.IP, ipv6Enabled); !ok {
continue
}
count++
}
return count
}
func ruleIPKey(raw string, ipv6Enabled bool) (string, bool) {
parsed := net.ParseIP(raw)
if parsed == nil {
return "", false
}
if ip4 := parsed.To4(); ip4 != nil {
return "4:" + net.IP(ip4).String(), true
}
if !ipv6Enabled {
return "", false
}
ip16 := parsed.To16()
if ip16 == nil {
return "", false
}
return "6:" + net.IP(ip16).String(), true
}
//go:build linux
package firewall
import (
"context"
"os/exec"
"time"
)
// RulesetSnapshot returns the live rules and the baseline captured immediately
// after the last successful Apply. The engine lock keeps our own concurrent
// reconfiguration from splitting the pair. An empty baseline means capture
// failed or this engine attached to an existing table without applying it.
func (e *Engine) RulesetSnapshot() (current, applied string, err error) {
e.mu.Lock()
defer e.mu.Unlock()
current, err = e.readRulesetLocked()
return current, e.appliedRuleset, err
}
func (e *Engine) readRulesetLocked() (string, error) {
if e.readRuleset != nil {
return e.readRuleset()
}
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
// Terse/stateless output omits members, counters, and elapsed timeouts.
// Keep transient traffic state out of both the baseline and its memory cost.
out, err := exec.CommandContext(ctx, "nft", "-s", "-t", "list", "table", "inet", "csm").Output()
if err != nil {
return "", err
}
return string(out), nil
}
package firewall
import (
"fmt"
"os"
"os/user"
"strconv"
)
// smtpAllowlistLookupUser is the user-database accessor that
// resolveSMTPAllowedUIDs goes through. Production binds it to
// user.Lookup; tests can swap a fixture so unit tests do not depend on
// /etc/passwd of the host they run on.
var smtpAllowlistLookupUser = user.Lookup
// resolveSMTPAllowedUIDs returns the deduplicated list of UIDs that are
// allowed to open outbound SMTP connections when smtp_block is on.
//
// Two UIDs are unconditional:
// - 0 (root): operator commands, csm itself, system tools.
// - mailnull: exim's queue-runner runs as this user on cPanel; if it
// is dropped, queued mail never leaves the host even though cPanel
// thinks it released the hold. Silently breaking outbound mail is
// the worst possible failure mode for a security firewall, so we
// allow it unconditionally rather than rely on the operator
// remembering to list it under smtp_allow_users.
//
// `allowUsers` is the operator-supplied set; each entry is resolved
// through smtpAllowlistLookupUser. Unknown or unparseable entries are
// reported to stderr (matching the legacy behavior in createOutputChain)
// and skipped, so a typo in the YAML does not crash the firewall engine.
func resolveSMTPAllowedUIDs(allowUsers []string) []uint32 {
// Cap the size hint so a pathological config (or future caller bug)
// cannot drive a multi-gigabyte allocation; the +2 covers root and
// mailnull which are added unconditionally below.
const smtpAllowHintCap = 1 << 16
hint := len(allowUsers)
if hint > smtpAllowHintCap {
hint = smtpAllowHintCap
}
seen := make(map[uint32]struct{}, hint+2)
out := make([]uint32, 0, hint+2)
add := func(uid uint32) {
if _, ok := seen[uid]; ok {
return
}
seen[uid] = struct{}{}
out = append(out, uid)
}
add(0)
if u, err := smtpAllowlistLookupUser("mailnull"); err == nil {
if uid, parseErr := strconv.ParseUint(u.Uid, 10, 32); parseErr == nil {
add(uint32(uid))
}
}
for _, name := range allowUsers {
u, err := smtpAllowlistLookupUser(name)
if err != nil {
fmt.Fprintf(os.Stderr, "firewall: smtp_allow_users: unknown user %q\n", name)
continue
}
uid, err := strconv.ParseUint(u.Uid, 10, 32)
if err != nil {
fmt.Fprintf(os.Stderr, "firewall: smtp_allow_users: invalid uid for %s: %v\n", name, err)
continue
}
add(uint32(uid))
}
return out
}
package firewall
import (
"encoding/json"
"errors"
"os"
"path/filepath"
"time"
)
// ErrPermanentBlock requires an explicit unblock before a timed Web UI block.
var ErrPermanentBlock = errors.New("IP is permanently blocked; unblock it explicitly before applying a timed block")
// ErrLongerBlock requires an explicit unblock before shortening a timed block.
var ErrLongerBlock = errors.New("IP has a longer block; unblock it explicitly before shortening its lifetime")
// ErrBlockChanged prevents undo from replacing a later firewall decision.
var ErrBlockChanged = errors.New("block changed since the action; undo is no longer available")
// SameBlockedEntry compares snapshots without relying on time.Time locations.
func SameBlockedEntry(a, b BlockedEntry) bool {
return a.IP == b.IP && a.Reason == b.Reason && a.Source == b.Source &&
a.BlockedAt.Equal(b.BlockedAt) && a.ExpiresAt.Equal(b.ExpiresAt)
}
// LoadState reads the authoritative firewall state file directly without requiring
// a running engine. A missing state file is a valid fresh-host state and returns
// an empty FirewallState.
func LoadState(statePath string) (*FirewallState, error) {
stateFile := filepath.Join(statePath, "firewall", "state.json")
// #nosec G304 -- filepath.Join under operator-configured statePath.
data, err := os.ReadFile(stateFile)
if err != nil {
if os.IsNotExist(err) {
return &FirewallState{}, nil
}
return nil, err
}
var state FirewallState
if err := json.Unmarshal(data, &state); err != nil {
return nil, err
}
// Clean expired entries
now := time.Now()
var active []BlockedEntry
for _, entry := range state.Blocked {
if entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt) {
active = append(active, entry)
}
}
state.Blocked = active
var activeNets []SubnetEntry
for _, entry := range state.BlockedNet {
if entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt) {
activeNets = append(activeNets, entry)
}
}
state.BlockedNet = activeNets
var activeAllowed []AllowedEntry
for _, entry := range state.Allowed {
if entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt) {
activeAllowed = append(activeAllowed, entry)
}
}
state.Allowed = activeAllowed
return &state, nil
}
//go:build linux
package firewall
// Supported reports whether this build includes the nftables engine.
func Supported() bool { return true }
// Package forensic produces evidence archives for incident response.
//
// A snapshot bundles the structured outputs an operator needs to hand a
// customer after a database-layer compromise: full trigger / event /
// routine definitions, the admin user roster, active session metadata,
// and the recent-mtime list under the account's document roots. Wrapped
// in a tar+gzip with a manifest and a SHA256 sidecar.
//
// The snapshot intentionally excludes credentials. Password rotation is
// a separate runbook step; bundling new credentials with the evidence
// archive would conflate two opposing flows (hand-to-customer for
// audit vs hand-to-ops for rotation) and risk credential leakage if
// the archive is later shared casually.
package forensic
import (
"archive/tar"
"bytes"
"compress/gzip"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"syscall"
"time"
)
// SchemaTarget identifies one MySQL schema to include in the snapshot
// along with the WordPress table prefix needed to enumerate the user
// roster correctly. Discovered from wp-config.php in the account's
// document roots by the production hook.
type SchemaTarget struct {
Schema string
TablePrefix string
ConfigPath string
}
// DiscoveryAudit records why a snapshot captured the targets it did.
// It is written into the manifest so operators can validate the
// evidence boundary without inspecting the host again.
type DiscoveryAudit struct {
AccountRoot string
PrivatePathsExcluded bool
PrivateTopPaths []string
SkippedPaths []SkippedPath
}
// SkippedPath records one discovery path that was not safe or useful
// to capture.
type SkippedPath struct {
Path string
Reason string
}
// Sources lets the caller swap each I/O dependency for a test double.
// Production wiring lives in cmd/csm; tests pass deterministic
// closures.
type Sources struct {
DiscoverTargets func(account string) []SchemaTarget
DumpSchema func(schema string) ([]byte, error)
ListAdmins func(schema, tablePrefix string) ([]byte, error)
ListSessions func(schema, tablePrefix string) ([]byte, error)
ListRecentFiles func(accountRoot string, since time.Time) ([]byte, error)
}
// Snapshot is the operator-facing configuration. Account and OutPath
// are required; Sources is required when running outside the default
// production wiring (see cmd/csm for the defaults).
type Snapshot struct {
Account string
OutPath string
Timestamp time.Time
DiscoveryAudit DiscoveryAudit
Sources Sources
}
var accountNamePattern = regexp.MustCompile(`^[A-Za-z0-9_-]{1,32}$`)
var schemaNamePattern = regexp.MustCompile(`^[A-Za-z0-9_@]+$`)
var tablePrefixPattern = regexp.MustCompile(`^[A-Za-z0-9_]+$`)
// AccountNameValid keeps the account string conservative enough to use
// in archive entry names, manifest keys, and shell-free SQL queries.
func AccountNameValid(name string) bool {
return accountNamePattern.MatchString(name)
}
func schemaNameValid(name string) bool {
return schemaNamePattern.MatchString(name)
}
func tablePrefixValid(name string) bool {
return tablePrefixPattern.MatchString(name)
}
// Write builds the archive at s.OutPath and a `<out>.sha256` sidecar,
// returning the archive path and the SHA256 hex digest. Errors abort
// before any partial state is written.
func (s Snapshot) Write() (string, string, error) {
if !AccountNameValid(s.Account) {
return "", "", fmt.Errorf("invalid account name: %q", s.Account)
}
if s.OutPath == "" {
return "", "", errors.New("forensic snapshot: OutPath required")
}
if err := ValidateOutPath(s.Account, s.OutPath); err != nil {
return "", "", err
}
if s.Sources.DiscoverTargets == nil {
return "", "", errors.New("forensic snapshot: Sources.DiscoverTargets required")
}
ts := s.Timestamp
if ts.IsZero() {
ts = time.Now().UTC()
}
targets := s.Sources.DiscoverTargets(s.Account)
sort.Slice(targets, func(i, j int) bool { return targets[i].Schema < targets[j].Schema })
var buf bytes.Buffer
gw := gzip.NewWriter(&buf)
tw := tar.NewWriter(gw)
var manifestB strings.Builder
fmt.Fprintf(&manifestB, "account=%s\n", s.Account)
fmt.Fprintf(&manifestB, "timestamp=%s\n", ts.UTC().Format(time.RFC3339))
fmt.Fprintf(&manifestB, "schema_count=%d\n", len(targets))
writeDiscoveryAudit(&manifestB, s.DiscoveryAudit)
var validTargets, invalidTargets int
var dumpOK, dumpErr int
var adminsOK, adminsErr int
var sessionsOK, sessionsErr int
recentStatus := "disabled"
for i, tgt := range targets {
if !schemaNameValid(tgt.Schema) || !tablePrefixValid(tgt.TablePrefix) {
invalidTargets++
fmt.Fprintf(&manifestB, "invalid_target=%d schema=%q table_prefix=%q\n", i, tgt.Schema, tgt.TablePrefix)
name := "schema/invalid-target-" + strconv.Itoa(i) + ".err"
data := []byte(fmt.Sprintf("invalid schema target: schema=%q table_prefix=%q\n", tgt.Schema, tgt.TablePrefix))
if err := writeArchiveEntry(tw, name, data, ts); err != nil {
return "", "", err
}
continue
}
validTargets++
fmt.Fprintf(&manifestB, "schema=%s table_prefix=%s", tgt.Schema, tgt.TablePrefix)
if tgt.ConfigPath != "" {
fmt.Fprintf(&manifestB, " config_path=%q", tgt.ConfigPath)
}
fmt.Fprint(&manifestB, "\n")
// Schema dump.
if s.Sources.DumpSchema != nil {
data, err := s.Sources.DumpSchema(tgt.Schema)
name := "schema/" + tgt.Schema + "-routines.sql"
if err != nil {
dumpErr++
name += ".err"
data = []byte(err.Error() + "\n")
} else {
dumpOK++
}
if werr := writeArchiveEntry(tw, name, data, ts); werr != nil {
return "", "", werr
}
}
// Admin roster.
if s.Sources.ListAdmins != nil {
data, err := s.Sources.ListAdmins(tgt.Schema, tgt.TablePrefix)
name := "schema/" + tgt.Schema + "-admins.tsv"
if err != nil {
adminsErr++
name += ".err"
data = []byte(err.Error() + "\n")
} else {
adminsOK++
}
if werr := writeArchiveEntry(tw, name, data, ts); werr != nil {
return "", "", werr
}
}
// Sessions.
if s.Sources.ListSessions != nil {
data, err := s.Sources.ListSessions(tgt.Schema, tgt.TablePrefix)
name := "schema/" + tgt.Schema + "-sessions.tsv"
if err != nil {
sessionsErr++
name += ".err"
data = []byte(err.Error() + "\n")
} else {
sessionsOK++
}
if werr := writeArchiveEntry(tw, name, data, ts); werr != nil {
return "", "", werr
}
}
}
// Recent files.
if s.Sources.ListRecentFiles != nil {
data, err := s.Sources.ListRecentFiles("/home/"+s.Account, ts.Add(-7*24*time.Hour))
name := "files/recent-mtimes.tsv"
if err != nil {
recentStatus = "error"
name += ".err"
data = []byte(err.Error() + "\n")
} else {
recentStatus = "ok"
}
if werr := writeArchiveEntry(tw, name, data, ts); werr != nil {
return "", "", werr
}
}
fmt.Fprintf(&manifestB, "valid_target_count=%d\n", validTargets)
fmt.Fprintf(&manifestB, "invalid_target_count=%d\n", invalidTargets)
fmt.Fprintf(&manifestB, "dump_success_count=%d\n", dumpOK)
fmt.Fprintf(&manifestB, "dump_error_count=%d\n", dumpErr)
fmt.Fprintf(&manifestB, "admins_success_count=%d\n", adminsOK)
fmt.Fprintf(&manifestB, "admins_error_count=%d\n", adminsErr)
fmt.Fprintf(&manifestB, "sessions_success_count=%d\n", sessionsOK)
fmt.Fprintf(&manifestB, "sessions_error_count=%d\n", sessionsErr)
fmt.Fprintf(&manifestB, "recent_mtimes_status=%s\n", recentStatus)
// Manifest last so the schema list reflects what was actually
// processed.
if err := writeArchiveEntry(tw, "manifest.txt", []byte(manifestB.String()), ts); err != nil {
return "", "", err
}
if err := tw.Close(); err != nil {
return "", "", fmt.Errorf("closing tar: %w", err)
}
if err := gw.Close(); err != nil {
return "", "", fmt.Errorf("closing gzip: %w", err)
}
if err := writeNewFile(s.OutPath, buf.Bytes()); err != nil {
return "", "", fmt.Errorf("writing archive: %w", err)
}
sum := sha256.Sum256(buf.Bytes())
hexSum := hex.EncodeToString(sum[:])
sidecar := s.OutPath + ".sha256"
sidecarBody := fmt.Sprintf("%s %s\n", hexSum, filepath.Base(s.OutPath))
if err := writeNewFile(sidecar, []byte(sidecarBody)); err != nil {
if removeErr := os.Remove(s.OutPath); removeErr != nil {
return "", "", fmt.Errorf("writing sidecar: %w; removing incomplete archive: %v", err, removeErr)
}
return "", "", fmt.Errorf("writing sidecar: %w", err)
}
return s.OutPath, hexSum, nil
}
// writeNewFile creates path as a brand-new 0600 file and writes data to it.
// The destination is operator-chosen and often sits in a world-writable
// directory, where the compromised account can pre-plant a symlink; an
// exclusive, non-following create refuses anything already at the path
// instead of writing evidence through it.
func writeNewFile(path string, data []byte) error {
// #nosec G304 -- path is the operator's --out argument; O_EXCL|O_NOFOLLOW
// refuse a pre-existing file or symlink.
f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL|syscall.O_NOFOLLOW, 0o600)
if err != nil {
return err
}
if _, err := f.Write(data); err != nil {
_ = f.Close()
_ = os.Remove(path)
return err
}
if err := f.Close(); err != nil {
_ = os.Remove(path)
return err
}
return nil
}
func writeDiscoveryAudit(b *strings.Builder, audit DiscoveryAudit) {
if audit.AccountRoot != "" {
fmt.Fprintf(b, "discovery_root=%q\n", audit.AccountRoot)
}
if audit.PrivatePathsExcluded {
fmt.Fprintln(b, "private_paths_excluded=true")
}
if len(audit.PrivateTopPaths) > 0 {
paths := append([]string(nil), audit.PrivateTopPaths...)
sort.Strings(paths)
fmt.Fprintf(b, "private_top_excluded=%q\n", strings.Join(paths, ","))
}
for i, skipped := range audit.SkippedPaths {
fmt.Fprintf(b, "skipped_path=%d path=%q reason=%q\n", i, skipped.Path, skipped.Reason)
}
}
// ValidateOutPath rejects destinations that would land inside the
// target account's home directory. Writing forensic evidence somewhere
// the suspect user can read defeats the point.
func ValidateOutPath(account, outPath string) error {
if !AccountNameValid(account) {
return fmt.Errorf("invalid account name: %q", account)
}
abs, err := filepath.Abs(outPath)
if err != nil {
return fmt.Errorf("resolving out path: %w", err)
}
home := filepath.Clean("/home/" + account)
homes := []string{home}
if realHome, err := filepath.EvalSymlinks(home); err == nil {
homes = append(homes, realHome)
}
paths := []string{abs}
if realAbs, err := filepath.EvalSymlinks(abs); err == nil {
paths = append(paths, realAbs)
}
if realParent, err := filepath.EvalSymlinks(filepath.Dir(abs)); err == nil {
paths = append(paths, filepath.Join(realParent, filepath.Base(abs)))
}
for _, h := range homes {
for _, p := range paths {
if pathWithin(h, p) {
return fmt.Errorf("out path must not be inside /home/%s/", account)
}
}
}
return nil
}
func pathWithin(root, path string) bool {
rel, err := filepath.Rel(filepath.Clean(root), filepath.Clean(path))
if err != nil {
return false
}
return rel == "." || (rel != ".." && !strings.HasPrefix(rel, ".."+string(filepath.Separator)) && !filepath.IsAbs(rel))
}
// writeArchiveEntry adds a single file to the tar stream with a fixed
// mtime so identical inputs produce byte-identical archives. The mode
// is 0600 because forensic content is operator-only.
func writeArchiveEntry(tw *tar.Writer, name string, data []byte, ts time.Time) error {
if unsafeArchiveEntryName(name) {
return fmt.Errorf("unsafe archive entry name: %q", name)
}
hdr := &tar.Header{
Name: name,
Size: int64(len(data)),
Mode: 0o600,
ModTime: ts.UTC(),
}
if err := tw.WriteHeader(hdr); err != nil {
return fmt.Errorf("tar header %s: %w", name, err)
}
if _, err := tw.Write(data); err != nil {
return fmt.Errorf("tar body %s: %w", name, err)
}
return nil
}
func unsafeArchiveEntryName(name string) bool {
if name == "" || filepath.IsAbs(name) || strings.Contains(name, `\`) {
return true
}
clean := filepath.Clean(name)
return clean != name || clean == "." || clean == ".." || strings.HasPrefix(clean, ".."+string(filepath.Separator))
}
package geoip
// KnownEditions returns the MaxMind database editions CSM supports in the
// settings UI. The first slice is the free GeoLite2 family; the second is
// the paid GeoIP2 family. The lists are curated to what MaxMind actually
// publishes via the geoipupdate protocol; adding an edition here makes it
// selectable in the Settings → GeoIP → Database editions dropdown.
func KnownEditions() (free, commercial []string) {
free = []string{
"GeoLite2-City",
"GeoLite2-Country",
"GeoLite2-ASN",
}
commercial = []string{
"GeoIP2-City",
"GeoIP2-Country",
"GeoIP2-ISP",
"GeoIP2-Domain",
"GeoIP2-Connection-Type",
"GeoIP2-Anonymous-IP",
"GeoIP2-Enterprise",
}
return free, commercial
}
// Package geoip provides IP geolocation via MaxMind GeoLite2 databases
// and on-demand RDAP lookups for detailed ISP/org information.
package geoip
import (
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/netip"
"os"
"path/filepath"
"sync"
"time"
"github.com/oschwald/maxminddb-golang/v2"
)
// Info contains geolocation and network information for an IP.
type Info struct {
IP string `json:"ip"`
Country string `json:"country"` // ISO 3166-1 alpha-2 (e.g. "US")
CountryName string `json:"country_name"` // Full name (e.g. "United States")
City string `json:"city,omitempty"`
ASN uint `json:"asn,omitempty"` // Autonomous System Number
ASOrg string `json:"as_org,omitempty"` // AS Organization (ISP)
Network string `json:"network,omitempty"` // CIDR range
RDAPOrg string `json:"rdap_org,omitempty"` // Detailed org from RDAP (on-demand)
RDAPName string `json:"rdap_name,omitempty"` // Network name from RDAP
RDAPCountry string `json:"rdap_country,omitempty"` // Country from RDAP
}
// DB holds the MaxMind database readers.
type DB struct {
mu sync.RWMutex
cityDB *maxminddb.Reader
asnDB *maxminddb.Reader
dbDir string
rdapMu sync.Mutex
rdapTTL map[string]rdapCacheEntry
}
type rdapCacheEntry struct {
info Info
fetched time.Time
// failed marks a lookup that did not complete (transport error,
// non-200, undecodable body). Such an entry is retried after
// rdapNegativeTTL instead of standing as an empty answer for a day.
failed bool
}
// rdapBaseURL is the RDAP bootstrap endpoint; var so tests can point it at
// a local server. rdapNegativeTTL bounds how long a failed lookup is kept
// before the address is tried again.
var (
rdapBaseURL = "https://rdap.org/ip/"
rdapNegativeTTL = 10 * time.Minute
rdapPositiveTTL = 24 * time.Hour
rdapMaxResponseBytes = int64(1 << 20)
errRDAPLookupIncomplete = errors.New("rdap lookup did not complete")
)
// MaxMind GeoLite2 record structures
type cityRecord struct {
Country struct {
ISOCode string `maxminddb:"iso_code"`
Names map[string]string `maxminddb:"names"`
} `maxminddb:"country"`
City struct {
Names map[string]string `maxminddb:"names"`
} `maxminddb:"city"`
}
type asnRecord struct {
ASN uint `maxminddb:"autonomous_system_number"`
Org string `maxminddb:"autonomous_system_organization"`
}
// Open loads MaxMind databases from the given directory.
// Expects GeoLite2-City.mmdb and/or GeoLite2-ASN.mmdb.
// Returns nil if no databases found (graceful degradation).
func Open(dbDir string) *DB {
if dbDir == "" {
return nil
}
db := &DB{
dbDir: dbDir,
rdapTTL: make(map[string]rdapCacheEntry),
}
cityPath := filepath.Join(dbDir, "GeoLite2-City.mmdb")
if r, err := maxminddb.Open(cityPath); err == nil {
db.cityDB = r
fmt.Fprintf(os.Stderr, "geoip: loaded %s\n", cityPath)
}
asnPath := filepath.Join(dbDir, "GeoLite2-ASN.mmdb")
if r, err := maxminddb.Open(asnPath); err == nil {
db.asnDB = r
fmt.Fprintf(os.Stderr, "geoip: loaded %s\n", asnPath)
}
if db.cityDB == nil && db.asnDB == nil {
fmt.Fprintf(os.Stderr, "geoip: no databases found in %s (download GeoLite2-City.mmdb and GeoLite2-ASN.mmdb)\n", dbDir)
return nil
}
return db
}
// Close releases database resources.
func (db *DB) Close() {
if db == nil {
return
}
db.mu.Lock()
defer db.mu.Unlock()
if db.cityDB != nil {
_ = db.cityDB.Close()
}
if db.asnDB != nil {
_ = db.asnDB.Close()
}
}
// Reload opens new database readers and swaps them in atomically.
// Opens replacement readers first - if both fail, the old readers stay in place.
// If one succeeds and the other fails, only the successful one is swapped.
func (db *DB) Reload() error {
if db == nil {
return fmt.Errorf("geoip: cannot reload nil DB")
}
cityPath := filepath.Join(db.dbDir, "GeoLite2-City.mmdb")
asnPath := filepath.Join(db.dbDir, "GeoLite2-ASN.mmdb")
// Open replacements before taking lock
newCity, cityErr := maxminddb.Open(cityPath)
newASN, asnErr := maxminddb.Open(asnPath)
if cityErr != nil && asnErr != nil {
return fmt.Errorf("geoip: reload failed - city: %v, asn: %v", cityErr, asnErr)
}
db.mu.Lock()
if newCity != nil {
if db.cityDB != nil {
_ = db.cityDB.Close()
}
db.cityDB = newCity
}
if newASN != nil {
if db.asnDB != nil {
_ = db.asnDB.Close()
}
db.asnDB = newASN
}
db.mu.Unlock()
if cityErr != nil {
fmt.Fprintf(os.Stderr, "geoip: reload warning - city DB failed: %v\n", cityErr)
}
if asnErr != nil {
fmt.Fprintf(os.Stderr, "geoip: reload warning - ASN DB failed: %v\n", asnErr)
}
return nil
}
// OpenFresh creates a new DB from databases on disk.
// Use when no DB existed at startup and databases have since been downloaded.
// Returns nil if no databases found (same behavior as Open).
func OpenFresh(dbDir string) *DB {
return Open(dbDir)
}
// Lookup returns geolocation info for an IP from local MaxMind databases.
// Fast (microseconds), no network calls.
func (db *DB) Lookup(ip string) Info {
info := Info{IP: ip}
if db == nil {
return info
}
addr, err := netip.ParseAddr(ip)
if err != nil {
return info
}
db.mu.RLock()
defer db.mu.RUnlock()
if db.cityDB != nil {
var record cityRecord
result := db.cityDB.Lookup(addr)
if err := result.Decode(&record); err == nil {
info.Country = record.Country.ISOCode
info.CountryName = record.Country.Names["en"]
info.City = record.City.Names["en"]
if prefix := result.Prefix(); prefix.IsValid() {
info.Network = prefix.String()
}
}
}
if db.asnDB != nil {
var record asnRecord
result := db.asnDB.Lookup(addr)
if err := result.Decode(&record); err == nil {
info.ASN = record.ASN
info.ASOrg = record.Org
}
}
return info
}
// LookupWithRDAP returns geolocation info plus on-demand RDAP details.
// The RDAP lookup is cached for 24 hours.
func (db *DB) LookupWithRDAP(ip string) Info {
info := db.Lookup(ip)
// Check RDAP cache
db.rdapMu.Lock()
if cached, ok := db.rdapTTL[ip]; ok {
ttl := rdapPositiveTTL
if cached.failed {
ttl = rdapNegativeTTL
}
if time.Since(cached.fetched) < ttl {
db.rdapMu.Unlock()
info.RDAPOrg = cached.info.RDAPOrg
info.RDAPName = cached.info.RDAPName
info.RDAPCountry = cached.info.RDAPCountry
return info
}
}
db.rdapMu.Unlock()
// Fetch from RDAP. A failure is cached only briefly so an outage does
// not blank the registry data for a day.
rdapInfo, err := fetchRDAP(ip)
info.RDAPOrg = rdapInfo.RDAPOrg
info.RDAPName = rdapInfo.RDAPName
info.RDAPCountry = rdapInfo.RDAPCountry
// Cache
db.rdapMu.Lock()
db.rdapTTL[ip] = rdapCacheEntry{info: rdapInfo, fetched: time.Now(), failed: err != nil}
db.evictRDAPLocked()
db.rdapMu.Unlock()
return info
}
// maxRDAPCacheEntries hard-caps the RDAP cache. Caller holds db.rdapMu.
const maxRDAPCacheEntries = 10000
// evictRDAPLocked bounds the RDAP cache. It first drops expired entries; if the
// map is still over the cap (every entry fresh, e.g. a burst of distinct
// lookups within the 24h TTL), it evicts the oldest entries until at the cap so
// the map cannot grow without bound. Caller holds db.rdapMu.
func (db *DB) evictRDAPLocked() {
if len(db.rdapTTL) <= maxRDAPCacheEntries {
return
}
for k, v := range db.rdapTTL {
if time.Since(v.fetched) > 24*time.Hour {
delete(db.rdapTTL, k)
}
}
for len(db.rdapTTL) > maxRDAPCacheEntries {
var oldestKey string
var oldest time.Time
for k, v := range db.rdapTTL {
if oldestKey == "" || v.fetched.Before(oldest) {
oldestKey = k
oldest = v.fetched
}
}
if oldestKey == "" {
break
}
delete(db.rdapTTL, oldestKey)
}
}
// RDAP lookup - fetches from the appropriate RIR
func fetchRDAP(ip string) (Info, error) {
var info Info
url := rdapBaseURL + ip
client := &http.Client{Timeout: 5 * time.Second}
resp, err := client.Get(url)
if err != nil {
return info, fmt.Errorf("%w: %v", errRDAPLookupIncomplete, err)
}
defer resp.Body.Close()
if resp.StatusCode != 200 {
return info, fmt.Errorf("%w: HTTP %d", errRDAPLookupIncomplete, resp.StatusCode)
}
var rdap struct {
Name string `json:"name"`
Country string `json:"country"`
Handle string `json:"handle"`
Entities []struct {
VCardArray []interface{} `json:"vcardArray"`
Roles []string `json:"roles"`
} `json:"entities"`
}
// Read one byte past the cap so a complete JSON prefix with an oversized
// trailing body cannot be accepted and cached as a positive answer.
data, err := io.ReadAll(io.LimitReader(resp.Body, rdapMaxResponseBytes+1))
if err != nil {
return info, fmt.Errorf("%w: %v", errRDAPLookupIncomplete, err)
}
if int64(len(data)) > rdapMaxResponseBytes {
return info, fmt.Errorf("%w: response exceeds %d bytes", errRDAPLookupIncomplete, rdapMaxResponseBytes)
}
if err := json.Unmarshal(data, &rdap); err != nil {
return info, fmt.Errorf("%w: %v", errRDAPLookupIncomplete, err)
}
info.RDAPName = rdap.Name
info.RDAPCountry = rdap.Country
// Extract org name from entities
for _, entity := range rdap.Entities {
for _, role := range entity.Roles {
if role == "registrant" || role == "abuse" {
if len(entity.VCardArray) >= 2 {
if props, ok := entity.VCardArray[1].([]interface{}); ok {
for _, prop := range props {
if arr, ok := prop.([]interface{}); ok && len(arr) >= 4 {
if name, ok := arr[0].(string); ok && name == "fn" {
if val, ok := arr[3].(string); ok {
info.RDAPOrg = val
}
}
}
}
}
}
}
}
}
return info, nil
}
package geoip
import (
"archive/tar"
"compress/gzip"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"strings"
"time"
"github.com/oschwald/maxminddb-golang/v2"
)
const (
maxMindBaseURL = "https://download.maxmind.com/geoip/databases"
maxDownloadSize = 150 * 1024 * 1024 // 150MB
// maxExtractedSize bounds the size of the .mmdb entry we extract from
// the tar.gz. Real GeoLite2 files are ~70MB; 500MiB is generous. The
// check sits on the tar header and prevents io.Copy from writing a
// bomb-compressed entry to disk.
maxExtractedSize = 500 * 1024 * 1024
downloadTimeout = 120 * time.Second
)
// EditionResult reports the outcome of updating a single GeoLite2 edition.
type EditionResult struct {
Edition string // e.g. "GeoLite2-City"
Status string // "updated", "up_to_date", "error"
Err error // nil unless Status == "error"
}
// Update downloads GeoLite2 databases from MaxMind's direct download API.
// Returns one EditionResult per edition. Returns nil if credentials are empty.
func Update(dbDir, accountID, licenseKey string, editions []string) []EditionResult {
if accountID == "" || licenseKey == "" {
return nil
}
if err := os.MkdirAll(dbDir, 0700); err != nil {
result := make([]EditionResult, len(editions))
for i, ed := range editions {
result[i] = EditionResult{Edition: ed, Status: "error", Err: fmt.Errorf("creating directory: %w", err)}
}
return result
}
client := &http.Client{Timeout: downloadTimeout}
results := make([]EditionResult, len(editions))
for i, edition := range editions {
results[i] = updateEdition(client, dbDir, accountID, licenseKey, edition)
}
return results
}
func updateEdition(client *http.Client, dbDir, accountID, licenseKey, edition string) EditionResult {
return updateEditionWithURL(client, dbDir, accountID, licenseKey, edition, maxMindBaseURL)
}
func updateEditionWithURL(client *http.Client, dbDir, accountID, licenseKey, edition, baseURL string) EditionResult {
url := fmt.Sprintf("%s/%s/download?suffix=tar.gz", baseURL, edition)
markerPath := filepath.Join(dbDir, ".last-modified-"+edition)
// Read stored Last-Modified
storedLM := ""
// #nosec G304 -- filepath.Join under operator-configured dbDir.
if data, err := os.ReadFile(markerPath); err == nil {
storedLM = strings.TrimSpace(string(data))
}
// HEAD request to check Last-Modified
headReq, err := http.NewRequest("HEAD", url, nil)
if err != nil {
return EditionResult{Edition: edition, Status: "error", Err: err}
}
headReq.SetBasicAuth(accountID, licenseKey)
headResp, err := client.Do(headReq)
if err != nil {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("HEAD request: %w", err)}
}
headResp.Body.Close()
if headResp.StatusCode == 401 {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("invalid MaxMind credentials")}
}
if headResp.StatusCode == 429 {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("rate limited by MaxMind")}
}
if headResp.StatusCode != 200 {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("HEAD returned HTTP %d", headResp.StatusCode)}
}
remoteLM := headResp.Header.Get("Last-Modified")
if storedLM != "" && remoteLM != "" && storedLM == remoteLM {
return EditionResult{Edition: edition, Status: "up_to_date"}
}
// GET request to download
getReq, err := http.NewRequest("GET", url, nil)
if err != nil {
return EditionResult{Edition: edition, Status: "error", Err: err}
}
getReq.SetBasicAuth(accountID, licenseKey)
getResp, err := client.Do(getReq)
if err != nil {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("download: %w", err)}
}
defer getResp.Body.Close()
if getResp.StatusCode != 200 {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("download returned HTTP %d", getResp.StatusCode)}
}
// Reject oversized responses upfront if Content-Length is known
if getResp.ContentLength > maxDownloadSize {
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("download too large: %d bytes (max %d)", getResp.ContentLength, maxDownloadSize)}
}
// LimitReader as safety net for responses without Content-Length
mmdbTmpPath := filepath.Join(dbDir, edition+".mmdb.tmp")
if err := extractMMDB(io.LimitReader(getResp.Body, maxDownloadSize), mmdbTmpPath, edition); err != nil {
os.Remove(mmdbTmpPath)
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("extract: %w", err)}
}
if err := validateMMDB(mmdbTmpPath); err != nil {
os.Remove(mmdbTmpPath)
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("validate: %w", err)}
}
// Atomic install
destPath := filepath.Join(dbDir, edition+".mmdb")
if err := os.Rename(mmdbTmpPath, destPath); err != nil {
os.Remove(mmdbTmpPath)
return EditionResult{Edition: edition, Status: "error", Err: fmt.Errorf("install: %w", err)}
}
// Save Last-Modified marker
if remoteLM != "" {
_ = os.WriteFile(markerPath, []byte(remoteLM), 0600)
}
return EditionResult{Edition: edition, Status: "updated"}
}
func validateMMDB(path string) error {
db, err := maxminddb.Open(path)
if err != nil {
return err
}
return db.Close()
}
// extractMMDB reads a tar.gz stream and extracts the .mmdb file to destPath.
// MaxMind tar.gz archives contain a single directory with the .mmdb inside,
// e.g. GeoLite2-City_20260328/GeoLite2-City.mmdb
func extractMMDB(r io.Reader, destPath, edition string) error {
gz, err := gzip.NewReader(r)
if err != nil {
return fmt.Errorf("gzip: %w", err)
}
defer func() { _ = gz.Close() }()
tr := tar.NewReader(gz)
suffix := edition + ".mmdb"
for {
header, err := tr.Next()
if err == io.EOF {
return fmt.Errorf("no %s found in archive", suffix)
}
if err != nil {
return fmt.Errorf("reading tar: %w", err)
}
if header.Typeflag != tar.TypeReg {
continue
}
if !strings.HasSuffix(header.Name, suffix) {
continue
}
if header.Size > maxExtractedSize {
return fmt.Errorf("archive entry too large: %d bytes", header.Size)
}
// #nosec G304 -- destPath is filepath.Join under operator-configured dbDir.
f, err := os.OpenFile(destPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600)
if err != nil {
return fmt.Errorf("creating %s: %w", destPath, err)
}
_, copyErr := io.Copy(f, io.LimitReader(tr, maxExtractedSize))
closeErr := f.Close()
if copyErr != nil {
return fmt.Errorf("writing mmdb: %w", copyErr)
}
if closeErr != nil {
return fmt.Errorf("closing mmdb: %w", closeErr)
}
return nil
}
}
package health
import (
"maps"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
// Provider is the contract the daemon (or a stub for tests) implements
// so the snapshot builder doesn't depend on internal/daemon directly.
type Provider interface {
Hostname() string
StartedAt() time.Time
LatestScan() time.Time
BaselineAt() time.Time
WatcherStatuses() map[string]bool
StoreHealthy() bool
StoreSizeMB() float64
SeverityCounts() map[string]int
BlocklistSize() int
IncidentsOpen() int
BPFEnforcementActive() bool
HistoryCount() int
ConfigHash() string
BinaryHash() string
DryRunBlocksCount() int
AutomationStatus() AutomationStatus
UpdateInfo() UpdateInfo
Mode() string
// CorrelationAttribution is nil until the first active-set merge.
CorrelationAttribution() *CorrelationAttribution
QueueStatuses() map[string]queuehealth.Status
}
// Build assembles a Snapshot from the provider plus the static version
// string and capability list. Safe to call from any goroutine; the
// provider's accessors are expected to be lock-protected internally.
func Build(p Provider, version string, capabilities []string) Snapshot {
started := p.StartedAt()
uptime := int64(0)
if !started.IsZero() {
uptime = int64(time.Since(started).Seconds())
}
caps := append([]string(nil), capabilities...)
var wordpress map[string]WPVerificationCounts
if wp, ok := p.(WordPressVerificationProvider); ok {
wordpress = maps.Clone(wp.WordPressVerification())
}
return Snapshot{
WordPressVerification: wordpress,
Queues: maps.Clone(p.QueueStatuses()),
Version: version,
Hostname: p.Hostname(),
StartedAt: started,
UptimeSec: uptime,
LatestScan: p.LatestScan(),
BaselineAt: p.BaselineAt(),
BlocklistSize: p.BlocklistSize(),
IncidentsOpen: p.IncidentsOpen(),
BPFEnforcementActive: p.BPFEnforcementActive(),
HistoryCount: p.HistoryCount(),
Severities: cloneIntMap(p.SeverityCounts()),
Watchers: cloneBoolMap(p.WatcherStatuses()),
StoreHealthy: p.StoreHealthy(),
StoreSizeMB: p.StoreSizeMB(),
ConfigHash: p.ConfigHash(),
BinaryHash: p.BinaryHash(),
Capabilities: caps,
DryRunBlocks: p.DryRunBlocksCount(),
Automation: p.AutomationStatus(),
Update: p.UpdateInfo(),
Mode: p.Mode(),
CorrelationAttribution: cloneCorrelationAttribution(p.CorrelationAttribution()),
}
}
func cloneCorrelationAttribution(in *CorrelationAttribution) *CorrelationAttribution {
if in == nil {
return nil
}
return &CorrelationAttribution{
Current: cloneIntMap(in.Current),
Cumulative: cloneIntMap(in.Cumulative),
ActiveSetUpdates: in.ActiveSetUpdates,
Since: in.Since,
}
}
func cloneIntMap(m map[string]int) map[string]int {
out := make(map[string]int, len(m))
for k, v := range m {
out[k] = v
}
return out
}
func cloneBoolMap(m map[string]bool) map[string]bool {
out := make(map[string]bool, len(m))
for k, v := range m {
out[k] = v
}
return out
}
package health
import (
"github.com/pidginhost/csm/internal/bpf"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/maillog"
)
// Capabilities is the static list of features this build supports. Phpanel
// reads it via /api/v1/capabilities to feature-detect without version
// sniffing. Add a string here when shipping a feature; remove when ripping
// one out. Keep the base order stable; build-tag gated capabilities are
// appended only when this binary actually supports them.
//
// BPF capability strings (`bpf-...`) are appended dynamically based on the
// running kernel's accepted BPF program types. Their presence depends on
// build tag and host kernel and is therefore not stable across deployments.
func Capabilities() []string {
caps := []string{
"confd.dropins.v1", // P1
"profile.phpanel-agent.v1", // P1
"status.json.v1", // P2
"capabilities.v1", // P2
"doctor.v1", // P2
"config.schema.v1", // P2
"sd_notify.ready", // P2
"audit.fields.tenant.v1", // P3
"webhook.phpanel.v1", // P3
"events.sse.v1", // P3
"token.scope.readonly.v1", // P3
"mail.brute.account_key.v1", // P4
"ti.source.rspamd.v1", // P4
"auto_response.dry_run.v1", // P5
"infra_ips.guard.v1", // P5
"store.backup.v1", // P5
"ti.source.upstream.v1", // P6
"verdict.callback.v1", // P7
"systemd.dropin.example.v1", // P7
"incidents.v1",
"bpf_enforcement.available.v1",
"webui.prefs.v1", // operator preferences + saved views
"webui.undo.v1", // bulk-action undo
"mail.filter.exfil.v1", // BEC mail-filter exfiltration detection
"mail.queue.composition.v1",
"mail.forward_guard.v1", // opt-in MTA-native forward-guard (hold spam/backscatter forward copies)
"detect.http_scanner_profile.v1", // URL scanner-profile detector + challenge/block action
"challenge.stats.v1",
"verified_bots.editor.v1", // operator-managed verified-bot allowlist (rDNS + IP ranges) with web editor
"status.firewall_health.v1", // status snapshot reports firewall enabled/managed state + block counts
"mode.observe.v1", // observe posture: detection and alerting without host changes
"status.queue_health.v1",
}
if firewall.Supported() {
caps = append(caps,
"firewall.rollback.v1", // timed config rollback: apply, confirm, revert, survives a restart
"firewall.dos_exempt.v1", // ranges exempt from connection-rate and flood metering
)
}
if maillog.JournalSupported() {
caps = append(caps, "mail.source.journal.v1")
}
caps = appendBPFCaps(caps)
caps = appendActiveBPFFeatures(caps)
return caps
}
// appendActiveBPFFeatures adds one capability string per BPF-backed live
// monitor that is currently running on the kernel-side path (as opposed to
// its userspace fallback). Phases 2-4 extend this with their own feature
// keys; the test for each phase asserts the expected toggling.
func appendActiveBPFFeatures(out []string) []string {
if bpf.ActiveKind("connection_tracker") == bpf.BackendBPF {
out = append(out, "bpf-connection-tracker")
}
if bpf.ActiveKind("af_alg") == bpf.BackendBPF {
out = append(out, "bpf-af-alg-live")
}
if bpf.ActiveKind("exec_monitor") == bpf.BackendBPF {
out = append(out, "bpf-exec-monitor")
}
if bpf.ActiveKind("sensitive_files") == bpf.BackendBPF {
out = append(out, "bpf-sensitive-files")
}
return out
}
// bpfCapabilities returns the cached probe result. Tests use this to assert
// that capability strings stay in sync with the shared BPF probe.
func bpfCapabilities() bpf.Capabilities { return bpf.Probe() }
// appendBPFCaps adds one capability string per BPF program type the kernel
// accepts. Phases 1-4 add a second helper alongside this one for per-feature
// "live monitor is currently running on BPF" strings.
func appendBPFCaps(out []string) []string {
caps := bpf.Probe()
if caps.LSMAttach {
out = append(out, "bpf-lsm-attach")
}
if caps.CgroupSock {
out = append(out, "bpf-cgroup-sock")
}
if caps.Tracepoint {
out = append(out, "bpf-tracepoint")
}
if caps.Ringbuf {
out = append(out, "bpf-ringbuf")
}
return out
}
package health
import (
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
// Snapshot is the unified machine-readable health view assembled from the
// running daemon (or, on a cold lookup, from on-disk state). It is the
// single source of truth for /api/v1/status, csm status --json, csm doctor,
// and the sd_notify readiness gate.
type Snapshot struct {
WordPressVerification map[string]WPVerificationCounts `json:"wordpress_verification,omitempty"`
Queues map[string]queuehealth.Status `json:"queues,omitempty"`
Version string `json:"version"`
Hostname string `json:"hostname"`
// Mode is the operator's posture: "enforce" or "observe". An observe
// host runs detection and alerting but changes no host state.
Mode string `json:"mode,omitempty"`
StartedAt time.Time `json:"started_at"`
UptimeSec int64 `json:"uptime_seconds"`
LatestScan time.Time `json:"latest_scan,omitzero"`
BaselineAt time.Time `json:"baseline_at,omitzero"`
BlocklistSize int `json:"blocklist_size"`
IncidentsOpen int `json:"incidents_open"`
BPFEnforcementActive bool `json:"bpf_enforcement_active"`
HistoryCount int `json:"history_count"`
Severities map[string]int `json:"severities"` // "critical","high","warning"
Watchers map[string]bool `json:"watchers"` // name -> attached
StoreHealthy bool `json:"store_healthy"`
StoreSizeMB float64 `json:"store_size_mb"`
ConfigHash string `json:"config_hash,omitempty"`
BinaryHash string `json:"binary_hash,omitempty"`
Capabilities []string `json:"capabilities,omitempty"`
// DryRunBlocks is the count of firewall blocks that were intercepted by
// auto_response.dry_run and logged rather than applied to nftables.
// Cleared whenever auto-response is live; dry-run mode keeps a recent
// rolling window for operator review.
DryRunBlocks int `json:"dry_run_blocks,omitempty"`
// Automation is the operator-facing safety surface for automatic action
// rollout. It groups dry-run state, challenge routing, pending firewall
// rollback, and the last recorded automation action in one stable payload.
Automation AutomationStatus `json:"automation,omitempty"`
// CorrelationAttribution reports which checks feed cross-account
// correlation findings without a hosting owner. Nil until the daemon has
// merged an active set, and on daemons that predate the block.
CorrelationAttribution *CorrelationAttribution `json:"correlation_attribution,omitempty"`
// Update reports whether a newer CSM release is available upstream.
// Populated by internal/updatecheck. Zero value means the checker has
// not yet completed a poll (very early startup) or is disabled in
// config.
Update UpdateInfo `json:"update,omitempty"`
}
// AutomationStatus summarizes the live automation safety state. It is
// intentionally compact so status clients can decide whether the host is
// observe-only, actively mutating the firewall, or waiting for operator
// confirmation after a tentative firewall apply.
type AutomationStatus struct {
AutoResponseEnabled bool `json:"auto_response_enabled"`
AutoResponseBlockIPs bool `json:"auto_response_block_ips"`
AutoResponseDryRun bool `json:"auto_response_dry_run"`
// Termination needs a kernel process handle. A kernel that cannot pin one
// leaves configured automatic killing inoperative, so the capability and
// its cause travel with the status instead of staying in the log.
ProcessKillEnabled bool `json:"process_kill_enabled"`
ProcessSignalSupported bool `json:"process_signal_supported"`
ProcessSignalError string `json:"process_signal_error,omitempty"`
DryRunBlocks int `json:"dry_run_blocks"`
ChallengeEnabled bool `json:"challenge_enabled"`
ChallengePortGateEnabled bool `json:"challenge_port_gate_enabled"`
ChallengePortGateActive bool `json:"challenge_port_gate_active"`
ChallengePending int `json:"challenge_pending"`
ChallengeEscalated int `json:"challenge_escalated"`
// FirewallEnabled reflects firewall.enabled in config. FirewallManaged is
// true only when the daemon has a live nftables engine wired. The
// combination FirewallEnabled && !FirewallManaged means the firewall is
// configured on but the daemon is NOT managing it (e.g. the engine failed
// to apply at startup) -- a condition monitoring should alert on.
FirewallEnabled bool `json:"firewall_enabled"`
FirewallManaged bool `json:"firewall_managed"`
FirewallStartupError string `json:"firewall_startup_error,omitempty"`
FirewallBlockedIPs int `json:"firewall_blocked_ips"`
FirewallBlockedSubnets int `json:"firewall_blocked_subnets"`
FirewallRollbackPending bool `json:"firewall_rollback_pending"`
FirewallRollbackSecondsRemain int64 `json:"firewall_rollback_remaining_seconds,omitempty"`
LastAction *AutomationAction `json:"last_action,omitempty"`
}
// AutomationAction is the newest action-like finding CSM recorded.
type AutomationAction struct {
Check string `json:"check"`
Message string `json:"message"`
Timestamp time.Time `json:"timestamp"`
}
// CorrelationAttribution is the operator-facing view of cross-account
// correlation attribution. Current is the per-check count of qualifying
// findings in the latest-state active set that carry no hosting owner and
// are inside the correlation window at its most recent merge. Unstamped
// legacy rows also count. A later merge clears attributed or expired rows.
// Cumulative sums every unattributed row reported since the daemon started,
// across active-set merges and per-batch derivations, so a producer that
// recovered stays visible as having failed. Kept as its own type so
// internal/health does not import internal/checks.
type CorrelationAttribution struct {
Current map[string]int `json:"current"`
Cumulative map[string]int `json:"cumulative"`
ActiveSetUpdates int `json:"active_set_updates"`
Since time.Time `json:"since"`
}
// UpdateInfo mirrors updatecheck.Info for the health snapshot. Kept
// as a separate type so internal/health does not import
// internal/updatecheck and create a cycle.
type UpdateInfo struct {
LatestVersion string `json:"latest_version,omitempty"`
Available bool `json:"available,omitempty"`
Source string `json:"source,omitempty"`
CheckedAt time.Time `json:"checked_at,omitempty"`
Err string `json:"err,omitempty"`
}
// TotalFindings returns the sum across all severity buckets.
func (s Snapshot) TotalFindings() int {
total := 0
for _, v := range s.Severities {
total += v
}
return total
}
// AllWatchersAttached reports whether every registered watcher is attached.
// An empty Watchers map returns false (we never claim ready before probing).
func (s Snapshot) AllWatchersAttached() bool {
if len(s.Watchers) == 0 {
return false
}
for _, attached := range s.Watchers {
if !attached {
return false
}
}
return true
}
// OverallStatus collapses the snapshot into one of: "ok", "degraded", "down".
// - "down" if the snapshot was zero-valued (never assembled)
// - "degraded" if a watcher is detached, the store is unhealthy, an enabled
// firewall is unmanaged, enabled termination has no safe kernel path,
// or a protection queue is degraded. Advisory queues carry best-effort
// work and are reported without changing the host status.
// - "ok" otherwise
func (s Snapshot) OverallStatus() string {
if s.StartedAt.IsZero() && len(s.Watchers) == 0 {
return "down"
}
if !s.StoreHealthy || !s.AllWatchersAttached() || s.Automation.FirewallEnabled && !s.Automation.FirewallManaged ||
s.Automation.ProcessKillEnabled && !s.Automation.ProcessSignalSupported {
return "degraded"
}
for _, q := range s.Queues {
if q.Status == "degraded" && !q.Advisory {
return "degraded"
}
}
return "ok"
}
package incident
import (
"fmt"
"sort"
"strings"
"time"
)
// BulkStatusFilter selects stale active incidents for a bounded operator
// transition. The caller must supply at least one age guard and a positive
// limit so a broad filter cannot accidentally close every incident.
type BulkStatusFilter struct {
FromStatuses []Status
To Status
OlderThan time.Duration
LastSeenBefore time.Time
Kind Kind
Domain string
Account string
Mailbox string
Limit int
DryRun bool
Details string
Now time.Time
}
// BulkStatusItem is a small preview row for bulk incident status changes.
type BulkStatusItem struct {
ID string `json:"id"`
Kind string `json:"kind"`
Status string `json:"status"`
NewStatus string `json:"new_status"`
Severity string `json:"severity"`
Domain string `json:"domain,omitempty"`
Account string `json:"account,omitempty"`
Mailbox string `json:"mailbox,omitempty"`
CreatedAt time.Time `json:"created_at"`
LastSeenAt time.Time `json:"last_seen_at"`
}
// BulkStatusResult reports how many incidents matched the filter and how
// many were changed. Items is capped by BulkStatusFilter.Limit.
type BulkStatusResult struct {
Matched int
Updated int
Items []BulkStatusItem
}
// BulkSetStatus previews or applies one closing transition to matching
// incidents. Matching and mutation happen under the correlator lock so a
// fresh finding cannot update LastSeen between filter evaluation and close.
func (c *Correlator) BulkSetStatus(filter BulkStatusFilter) (BulkStatusResult, error) {
if filter.To != StatusResolved && filter.To != StatusDismissed {
return BulkStatusResult{}, fmt.Errorf("incident: bulk target status must be resolved or dismissed")
}
if filter.OlderThan <= 0 && filter.LastSeenBefore.IsZero() {
return BulkStatusResult{}, fmt.Errorf("incident: bulk status requires older-than or last-seen-before")
}
if filter.Limit <= 0 {
return BulkStatusResult{}, fmt.Errorf("incident: bulk status requires a positive limit")
}
now := filter.Now
if now.IsZero() {
now = c.now()
}
statusSet := make(map[Status]struct{}, len(filter.FromStatuses))
for _, status := range filter.FromStatuses {
if !validStatus(status) {
return BulkStatusResult{}, fmt.Errorf("incident: invalid status %q", status)
}
statusSet[status] = struct{}{}
}
if len(statusSet) == 0 {
return BulkStatusResult{}, fmt.Errorf("incident: bulk status requires a source status")
}
var persist []*queuedPersist
result := BulkStatusResult{Items: make([]BulkStatusItem, 0, filter.Limit)}
c.mu.Lock()
matched := make([]*Incident, 0, len(c.incidents))
for _, inc := range c.incidents {
if bulkStatusMatches(inc, filter, statusSet, now) {
matched = append(matched, inc)
}
}
sort.Slice(matched, func(i, j int) bool {
if !matched[i].UpdatedAt.Equal(matched[j].UpdatedAt) {
return matched[i].UpdatedAt.Before(matched[j].UpdatedAt)
}
return matched[i].ID < matched[j].ID
})
result.Matched = len(matched)
for _, inc := range matched {
if len(result.Items) >= filter.Limit {
break
}
from := inc.Status
result.Items = append(result.Items, bulkStatusItem(inc, filter.To))
if filter.DryRun {
continue
}
inc.Status = filter.To
inc.UpdatedAt = now
inc.ClosedAt = now
inc.ClosedBy = "operator"
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_status_changed",
Result: "ok",
Details: string(from) + " -> " + string(filter.To) + ": " + filter.Details,
})
c.counters.statusChangedTotal.Add(1)
c.unbindLocked(inc.ID)
if c.spray != nil {
c.spray.UnbindIncident(inc.ID)
}
if req, ok := c.queuePersistLocked(*inc); ok {
persist = append(persist, req)
}
result.Updated++
}
c.mu.Unlock()
c.runQueuedPersists(persist)
return result, nil
}
func bulkStatusMatches(inc *Incident, filter BulkStatusFilter, statusSet map[Status]struct{}, now time.Time) bool {
if _, ok := statusSet[inc.Status]; !ok {
return false
}
if filter.OlderThan > 0 {
cutoff := now.Add(-filter.OlderThan)
if inc.UpdatedAt.After(cutoff) {
return false
}
}
if !filter.LastSeenBefore.IsZero() && inc.UpdatedAt.After(filter.LastSeenBefore) {
return false
}
if filter.Kind != "" && inc.Kind != filter.Kind {
return false
}
if filter.Domain != "" && !strings.EqualFold(inc.Domain, filter.Domain) {
return false
}
if filter.Account != "" && !strings.EqualFold(inc.Account, filter.Account) {
return false
}
if filter.Mailbox != "" && !strings.EqualFold(inc.Mailbox, filter.Mailbox) {
return false
}
return true
}
func bulkStatusItem(inc *Incident, to Status) BulkStatusItem {
return BulkStatusItem{
ID: inc.ID,
Kind: string(inc.Kind),
Status: string(inc.Status),
NewStatus: string(to),
Severity: inc.Severity.String(),
Domain: inc.Domain,
Account: inc.Account,
Mailbox: inc.Mailbox,
CreatedAt: inc.CreatedAt,
LastSeenAt: inc.UpdatedAt,
}
}
package incident
import (
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"net"
"sort"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// ErrIncidentNotFound is returned when SetStatus or other lookups
// target an unknown incident id.
var ErrIncidentNotFound = errors.New("incident: not found")
// maxIncidentFindings and maxIncidentTimeline cap the per-incident
// fingerprint slice and operator-visible timeline so a long-running
// open incident with sustained low-severity traffic does not grow
// memory and persistence payloads without bound. Eviction keeps the
// first half (incident-opening context an operator needs to root-
// cause the incident) and the most recent half (so the timeline
// reflects current activity). Operators reading the timeline see a
// gap marker via the appended IncidentEvent that the cap fires.
const (
maxIncidentFindings = 5000
maxIncidentTimeline = 500
incidentFingerprintTruncatedMark = "...truncated:"
incidentTimelineTruncatedKind = "truncated"
)
// incidentMergeWindow is the time gap inside which two findings with the
// same correlation key are considered the same incident. Named constant
// per project convention; config exposure deferred until operators ask.
const incidentMergeWindow = 15 * time.Minute
// incidentPersistDebounce bounds how often a bookkeeping-only merge (a
// finding that appends fingerprint/timeline but does not change status,
// severity, or kind) rewrites the whole incident blob to the store. A busy
// incident under sustained attack would otherwise fsync a ~100-300KB blob on
// every finding; instead those writes coalesce to at most one per window.
// State-changing transitions always persist synchronously. Quiet bookkeeping
// waits for the next mutation or explicit flush; there is no periodic writer.
const incidentPersistDebounce = 5 * time.Second
// CorrelatorConfig is reserved for future tunables and the persistence
// hook used by the daemon to write incidents to bbolt.
type CorrelatorConfig struct {
// Persist receives ordered immutable snapshots; bookkeeping may coalesce.
// Implementations must be quick and idempotent. Errors are logged and
// counted without rolling back in-memory transitions. nil is memory-only.
Persist func(Incident) error
// OpenThreshold is the number of correlated findings required before
// a thresholded finding opens an incident. Critical-severity findings and
// check-specific first-hit signals always open immediately. Values <= 0
// default to 1 (open on first finding) for backwards compatibility with
// callers that expect the original behavior; the daemon explicitly
// configures 2 to suppress one-shot scanner noise.
OpenThreshold int
// SpraySuppression turns on the credential-spray super-incident
// path. When zero the detector is not constructed and OnFinding
// follows the legacy per-mailbox correlation path. Default-off.
SpraySuppression SpraySuppressionConfig
// IsWhitelisted is consulted before a source-IP finding can anchor
// incident correlation and by the spray detector to skip IPs the
// operator has marked as known-good (e.g. internal mail relays).
// nil short-circuits to "no IPs whitelisted".
IsWhitelisted func(ip string) bool
// CanSprayBlock is consulted immediately before recording a
// credential_spray block request. nil means "allowed" when
// OnSprayBlock is present. Implementations must be quick and must
// not call back into the correlator.
CanSprayBlock func() bool
// OnSprayBlock is invoked once per IP when the credential_spray
// detector decides the IP should be hard-blocked, based on
// SpraySuppression.BlockAtSeverity. The callback runs after the
// correlator mutex is released so firewall or verdict latency does
// not stall incident ingestion. nil disables the hand-off; the spray
// super-incident still opens and escalates, but no firewall action
// fires. The return value reports whether the firewall actually
// recorded the block: false means dry-run, transient failure, or an
// upstream gate refused, and the audit "credential_spray_block_requested"
// action is not appended in that case so operators cannot mistake a
// declined request for an enforced block. findingID identifies the latest
// eligible source observation, or is empty for an older unlinked timeline.
OnSprayBlock func(ip, reason string, ttl time.Duration, findingID string) bool
// AutoBlock turns on the generic incident-driven firewall hand-off
// for non-spray kinds. Independent of SpraySuppression; applies when
// an incident has exactly one unambiguous remote IP and its kind is
// allowed by AutoBlock.Kinds (empty = any). Default-zero means the
// path is dormant.
AutoBlock IncidentAutoBlockConfig
// CanIncidentBlock is consulted immediately before recording a
// generic incident block request. nil means "allowed" when
// OnIncidentBlock is present. Lets the daemon recheck
// auto_response.enabled / block_ips at decision time so SIGHUP
// edits take effect without rebuilding the correlator.
CanIncidentBlock func() bool
// OnIncidentBlock fires when the generic auto-block gate trips. The
// callback runs after the correlator mutex is released and returns
// true only when a live block request was accepted. Dry-run,
// disabled, and failed attempts must return false so the correlator
// can retry on the next finding instead of permanently latching the
// incident. nil disables the path even when AutoBlock is configured.
// findingID carries the same source attribution as OnSprayBlock.
OnIncidentBlock func(ip, reason string, ttl time.Duration, findingID string) bool
}
// IncidentAutoBlockConfig drives the generic incident-driven firewall
// hand-off independent of credential-spray suppression. Operators turn
// it on once they have validated that incident severity is trustworthy
// (the daemon does not promote a finding to High/Critical without
// either an explicit per-check signal or the correlator's threshold
// gate).
type IncidentAutoBlockConfig struct {
Enabled bool
// BlockAtSeverity is the minimum incident severity that triggers
// a firewall hand-off. "" / "high" / "critical". Comparison is
// case-insensitive. Any other value is ignored so operator typos
// cannot accidentally engage blocking.
BlockAtSeverity string
// BlockExpiry is the operator's configured auto-response block duration,
// the first rung of the escalation ladder. Zero falls back to 24h.
BlockExpiry time.Duration
// Kinds, when non-empty, restricts the auto-block path to the
// listed incident kinds. Empty means "every kind that carries one
// unambiguous remote IP". Credential_spray is implicitly excluded
// since the dedicated spray hand-off owns it.
Kinds map[Kind]bool
}
// IsZero reports whether the config is unset; the correlator treats a
// zero value as "generic auto-block disabled" without touching
// defaults.
func (c IncidentAutoBlockConfig) IsZero() bool {
return !c.Enabled && c.BlockAtSeverity == "" && len(c.Kinds) == 0
}
// counters holds the atomic tallies exposed via RegisterMetrics. Kept
// on the Correlator so a single instance owns its own metric state and
// tests can build isolated correlators without touching globals.
type counters struct {
createdTotal atomic.Uint64
severityChangedTotal atomic.Uint64
statusChangedTotal atomic.Uint64
findingsMergedTotal atomic.Uint64
compactedTotal atomic.Uint64
autoClosedTotal atomic.Uint64
autoCloseDryRunTotal atomic.Uint64
sprayOpenedTotal atomic.Uint64
spraySuppressedTotal atomic.Uint64
sprayDryRunTotal atomic.Uint64
}
// Correlator groups findings into incidents. In-memory state; the
// daemon is responsible for wiring it to a store via CorrelatorConfig.Persist.
type Correlator struct {
mu sync.Mutex
persistence *persistQueue
cfg CorrelatorConfig
incidents map[string]*Incident
byKey map[string]string
pending map[string]pendingFinding
// A false pending value keeps the in-flight slot occupied but discards
// its result after an incident closes, even if it is reopened meanwhile.
pendingSprayBlocks map[string]bool
pendingIncidentBlocks map[string]bool
openThreshold int
now func() time.Time
counters counters
spray *sprayDetector
// lastPersistAt records the wall-clock time of the most recent scheduled
// store write per incident id, used to debounce bookkeeping-only merges.
// Written under c.mu; cleared when the incident leaves c.incidents.
lastPersistAt map[string]time.Time
}
// pendingFinding is a finding seen for a key that has not yet met the
// open threshold. Stored only on the create path; merge into open
// incidents stays unconditional.
type pendingFinding struct {
finding alert.Finding
at time.Time
}
// NewCorrelator returns a ready Correlator. Nothing to start; this type
// is purely callback-driven.
func NewCorrelator(cfg CorrelatorConfig) *Correlator {
threshold := cfg.OpenThreshold
if threshold < 1 {
threshold = 1
}
c := &Correlator{
cfg: cfg,
persistence: newPersistQueue(),
incidents: map[string]*Incident{},
byKey: map[string]string{},
pending: map[string]pendingFinding{},
pendingSprayBlocks: map[string]bool{},
pendingIncidentBlocks: map[string]bool{},
openThreshold: threshold,
now: time.Now,
lastPersistAt: map[string]time.Time{},
}
c.spray = newSprayDetector(cfg.SpraySuppression, incidentMergeWindow, func() time.Time { return c.now() }, cfg.IsWhitelisted)
return c
}
// OnFinding ingests a Finding. Returns the incident id (if attributable)
// and whether a new incident was created. Unattributable findings yield
// ("", false, nil). Findings subject to the threshold whose key has fewer
// than OpenThreshold prior findings inside the merge window are stashed in the
// pending map and yield ("", false, nil) too; they will only open an incident
// if the threshold is met inside the window.
func (c *Correlator) OnFinding(f alert.Finding) (string, bool, error) {
key := KeyFor(f)
if key.IsEmpty() {
return "", false, nil
}
if key.Host == "" && f.SourceIP != "" && c.cfg.IsWhitelisted != nil && c.cfg.IsWhitelisted(f.SourceIP) {
return "", false, nil
}
var afterUnlock func()
c.mu.Lock()
defer func() {
c.mu.Unlock()
if afterUnlock != nil {
afterUnlock()
}
}()
keyStr := keyString(key)
now := c.now()
// Credential-spray super-incident path. When one source IP brute-forces
// many distinct mailboxes/accounts inside the merge window, collapse
// the per-mailbox fan-out into a single credential_spray incident
// keyed on RemoteIP. The detector returns sprayDecisionNone for
// non-spray traffic so the legacy correlation continues unchanged.
if c.spray != nil {
decision, hits := c.spray.Decide(f)
switch decision {
case sprayDecisionOpen:
sprayKey := Key{RemoteIP: f.SourceIP}
id, created := c.promoteOrCreateSprayLocked(sprayKey, f, now, hits)
c.spray.BindIncident(f.SourceIP, id)
c.counters.sprayOpenedTotal.Add(1)
if cb := c.maybeBlockSprayLocked(c.incidents[id], f.SourceIP, hits, now, "spray opened"); cb != nil {
afterUnlock = cb
}
return id, created, nil
case sprayDecisionSuppress:
id := c.spray.IncidentForIP(f.SourceIP)
inc, ok := c.incidents[id]
if ok && incidentStatusActive(inc.Status) {
// Fold the finding and apply the spray-specific escalation
// before a single persist so the escalating finding does not
// write the blob twice (mutate + escalate in one write).
transition := c.mutateWithFindingLocked(inc, f, now)
c.counters.findingsMergedTotal.Add(1)
if hits >= c.spray.cfg.SeverityEscalateAt && inc.Severity < alert.Critical {
from := inc.Severity
inc.Severity = alert.Critical
c.counters.severityChangedTotal.Add(1)
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_severity_changed",
Result: "ok",
Details: from.String() + " -> CRITICAL: spray sustained " + strconv.Itoa(hits) + " mailboxes",
})
transition = true
}
c.persistAfterMergeLocked(inc, now, true, transition)
// Re-evaluate the block gate on every merged finding. The
// configured BlockAtSeverity may have been armed AFTER the
// incident was opened or escalated, in which case the
// transition-time hook already fired without effect; the
// helper is idempotent via triggerSprayBlockLocked's
// action-presence and in-flight guards so a no-op call is
// harmless.
if cb := c.maybeBlockSprayLocked(inc, f.SourceIP, hits, now, "spray ongoing"); cb != nil {
afterUnlock = cb
}
c.counters.spraySuppressedTotal.Add(1)
return id, false, nil
}
// Bound incident vanished (purged) or is no longer active
// (operator resolved/dismissed). Clear the perIP binding so
// subsequent findings don't keep falling into the same dead
// lookup, then fall through to legacy so the finding still
// produces an incident rather than silently disappearing.
if id != "" {
c.spray.UnbindIncident(id)
}
case sprayDecisionNone:
if c.spray.cfg.DryRun && hits >= c.spray.cfg.DistinctMailboxes {
c.counters.sprayDryRunTotal.Add(1)
}
}
}
if id, ok := c.byKey[keyStr]; ok {
// A quiet interval must not erase an active incident's block ladder.
// Only closing the incident ends that episode, including after restore.
if inc, exists := c.incidents[id]; exists && (now.Sub(inc.UpdatedAt) <= incidentMergeWindow || inc.AutoBlock.Count > 0) {
c.mergeLocked(inc, f, now, true)
delete(c.pending, keyStr)
if cb := c.maybeBlockIncidentLocked(inc, now, "merge"); cb != nil {
afterUnlock = cb
}
return id, false, nil
}
// Stale binding -- the incident is older than the merge window.
// Drop the binding and fall through to create so a fresh incident
// owns the key going forward.
delete(c.byKey, keyStr)
}
// Threshold gate. Findings without first-hit semantics need OpenThreshold
// sightings inside the merge window before opening an incident.
if c.openThreshold > 1 && !opensIncidentImmediately(f) {
if pf, ok := c.pending[keyStr]; ok && now.Sub(pf.at) <= incidentMergeWindow {
delete(c.pending, keyStr)
id := c.createIncidentLocked(key, keyStr, pf.finding, pf.at)
inc := c.incidents[id]
// The second finding is what satisfied the open threshold, so it
// is part of incident creation and must be durable even if it
// lands inside the bookkeeping debounce window.
c.mergeAndPersistLocked(inc, f, now, true)
if cb := c.maybeBlockIncidentLocked(inc, now, "threshold promote"); cb != nil {
afterUnlock = cb
}
return id, true, nil
}
c.pending[keyStr] = pendingFinding{finding: f, at: now}
return "", false, nil
}
id := c.createIncidentLocked(key, keyStr, f, now)
delete(c.pending, keyStr)
if cb := c.maybeBlockIncidentLocked(c.incidents[id], now, "incident opened"); cb != nil {
afterUnlock = cb
}
return id, true, nil
}
// promoteOrCreateSprayLocked opens the credential_spray incident for key.
// The first failures from an IP usually opened an ordinary per-IP incident
// under the same RemoteIP key before the spray threshold tripped; creating
// a second incident under that key stole the key index and left the earlier
// incident Open with nothing able to merge into or close it. An active
// incident already holding the key is promoted in place, keeping its
// findings and timeline. Caller must hold c.mu.
func (c *Correlator) promoteOrCreateSprayLocked(key Key, f alert.Finding, now time.Time, hits int) (string, bool) {
if existingID, ok := c.byKey[keyString(key)]; ok {
if inc, live := c.incidents[existingID]; live && incidentStatusActive(inc.Status) {
fromKind := inc.Kind
inc.Kind = KindCredentialSpray
if inc.Severity < alert.High {
fromSeverity := inc.Severity
inc.Severity = alert.High
c.counters.severityChangedTotal.Add(1)
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_severity_changed",
Result: "ok",
Details: fromSeverity.String() + " -> HIGH: promoted to credential spray",
})
}
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "credential_spray_opened",
Result: "ok",
Details: f.SourceIP + " hit " + strconv.Itoa(hits) + " distinct mailboxes inside window; promoted from " + string(fromKind),
})
// A kind change must be durable, not left to the bookkeeping
// debounce.
c.mergeAndPersistLocked(inc, f, now, true)
return existingID, false
}
}
return c.createSprayIncidentLocked(key, f, now, hits), true
}
// createSprayIncidentLocked builds a credential_spray incident keyed on
// the source IP. Caller must hold c.mu. Severity is HIGH at trip and
// escalates to CRITICAL once the merge path observes
// SpraySuppressionConfig.SeverityEscalateAt distinct mailboxes.
func (c *Correlator) createSprayIncidentLocked(key Key, f alert.Finding, now time.Time, hits int) string {
id := newIncidentID()
sev := f.Severity
if sev < alert.High {
sev = alert.High
}
inc := &Incident{
ID: id,
Kind: KindCredentialSpray,
Status: StatusOpen,
Severity: sev,
CorrelationKey: cloneKey(key),
Findings: []string{},
Timeline: []IncidentEvent{},
Actions: []IncidentAction{{
Time: now,
Action: "credential_spray_opened",
Result: "ok",
Details: f.SourceIP + " hit " + strconv.Itoa(hits) + " distinct mailboxes inside window",
}},
CreatedAt: now,
UpdatedAt: now,
}
c.incidents[id] = inc
keyStr := keyString(key)
c.byKey[keyStr] = id
c.counters.createdTotal.Add(1)
c.mergeLocked(inc, f, now, false)
return id
}
// createIncidentLocked builds a new Incident, registers it in the maps,
// and seeds it with the given finding via mergeLocked. Caller must hold
// c.mu. mergeLocked is the single source of truth for Persist
// invocations, avoiding double-fire on create.
func (c *Correlator) createIncidentLocked(key Key, keyStr string, f alert.Finding, now time.Time) string {
id := newIncidentID()
displayMailbox, displayDomain := displayMailboxDomain(f.Mailbox, f.Domain)
if displayMailbox == "" && displayDomain == "" {
displayMailbox, displayDomain = key.Mailbox, key.Domain
}
inc := &Incident{
ID: id,
Kind: ClassifyKind(f),
Status: StatusOpen,
Severity: f.Severity,
Account: key.Account,
Domain: displayDomain,
Mailbox: displayMailbox,
CorrelationKey: cloneKey(key),
Findings: []string{},
Timeline: []IncidentEvent{},
Actions: []IncidentAction{},
CreatedAt: now,
UpdatedAt: now,
}
c.incidents[id] = inc
c.byKey[keyStr] = id
c.counters.createdTotal.Add(1)
c.mergeLocked(inc, f, now, false)
return id
}
// PendingCount returns the number of findings currently held in the
// threshold-gate pending map. Exposed for metrics and tests.
func (c *Correlator) PendingCount() int {
c.mu.Lock()
defer c.mu.Unlock()
return len(c.pending)
}
// PruneStalePending removes pending findings whose age relative to now
// exceeds the merge window. Returns the number pruned. Called by the
// daemon's retention loop so a host with sustained one-shot scanner
// traffic does not grow the pending map without bound.
func (c *Correlator) PruneStalePending(now time.Time) int {
c.mu.Lock()
defer c.mu.Unlock()
cutoff := now.Add(-incidentMergeWindow)
pruned := 0
for k, pf := range c.pending {
if pf.at.Before(cutoff) {
delete(c.pending, k)
pruned++
}
}
return pruned
}
// PruneStaleSpray clears spray-detector state whose lastSeen is older
// than the merge window. Wired into the same retention sweep as
// PruneStalePending.
func (c *Correlator) PruneStaleSpray(now time.Time) int {
c.mu.Lock()
defer c.mu.Unlock()
return c.spray.PruneStale(now)
}
// SprayTrackedIPs reports the count of source IPs currently held in
// the spray detector. Safe to call when the detector is disabled.
func (c *Correlator) SprayTrackedIPs() int {
c.mu.Lock()
defer c.mu.Unlock()
return c.spray.TrackedIPs()
}
// OpenCount returns the number of incidents in Open or Contained
// status. Used by the csm_incidents_open gauge; computed at scrape
// time so the value never drifts from in-memory state.
func (c *Correlator) OpenCount() int {
c.mu.Lock()
defer c.mu.Unlock()
n := 0
for _, inc := range c.incidents {
if inc.Status == StatusOpen || inc.Status == StatusContained {
n++
}
}
return n
}
// OpenCountsBySeverity returns Open and Contained incidents keyed by
// lowercase severity. Contained incidents still need operator attention,
// so they remain part of the dashboard's active-incident posture.
func (c *Correlator) OpenCountsBySeverity() map[string]int {
c.mu.Lock()
defer c.mu.Unlock()
counts := map[string]int{}
for _, inc := range c.incidents {
if inc.Status != StatusOpen && inc.Status != StatusContained {
continue
}
switch inc.Severity {
case alert.Critical:
counts["critical"]++
case alert.High:
counts["high"]++
default:
counts["warning"]++
}
}
return counts
}
// Get returns a snapshot of the incident by id.
func (c *Correlator) Get(id string) (Incident, bool) {
c.mu.Lock()
defer c.mu.Unlock()
inc, ok := c.incidents[id]
if !ok {
return Incident{}, false
}
return cloneIncident(*inc), true
}
// SnapshotPage returns a page of incidents matching status (empty
// string means all statuses), starting at offset, with at most limit
// items. The returned total is the number of records that match the
// filter regardless of the page bounds, so the caller can render an
// accurate "X of Y" header.
//
// limit <= 0 returns the rest of the filtered set after offset. The
// caller (web UI / phpanel) is expected to cap the page at a sane
// ceiling; this primitive only enforces correct slicing. Negative
// offset is clamped to zero.
//
// Items are deep-copied so callers may mutate the returned slice
// without affecting subsequent calls.
func (c *Correlator) SnapshotPage(status Status, offset, limit int) ([]Incident, int) {
if status == "" {
return c.SnapshotPageStatuses(nil, offset, limit)
}
return c.SnapshotPageStatuses([]Status{status}, offset, limit)
}
// SnapshotPageStatuses returns a page of incidents matching any status
// in statuses. An empty status list means all statuses. Sorting and
// slicing happen against internal pointers first; only the returned
// page is deep-copied.
func (c *Correlator) SnapshotPageStatuses(statuses []Status, offset, limit int) ([]Incident, int) {
c.mu.Lock()
defer c.mu.Unlock()
statusSet := make(map[Status]struct{}, len(statuses))
for _, st := range statuses {
if st != "" {
statusSet[st] = struct{}{}
}
}
matched := make([]*Incident, 0, len(c.incidents))
for _, inc := range c.incidents {
if len(statusSet) > 0 {
if _, ok := statusSet[inc.Status]; !ok {
continue
}
}
matched = append(matched, inc)
}
sortIncidentRefs(matched)
total := len(matched)
if offset < 0 {
offset = 0
}
if offset >= total {
return []Incident{}, total
}
end := total
if limit > 0 && offset+limit < end {
end = offset + limit
}
out := make([]Incident, 0, end-offset)
for _, inc := range matched[offset:end] {
out = append(out, cloneIncident(*inc))
}
return out, total
}
func sortIncidentRefs(refs []*Incident) {
sort.Slice(refs, func(i, j int) bool {
return incidentRefLess(refs[i], refs[j])
})
}
func incidentRefLess(a, b *Incident) bool {
if !a.UpdatedAt.Equal(b.UpdatedAt) {
return a.UpdatedAt.After(b.UpdatedAt)
}
return a.ID > b.ID
}
// Snapshot returns every incident sorted by UpdatedAt descending. Safe
// for concurrent callers; produces a deep-copy slice so the API layer
// can serialize it without coordinating with mutators.
func (c *Correlator) Snapshot() []Incident {
c.mu.Lock()
defer c.mu.Unlock()
refs := make([]*Incident, 0, len(c.incidents))
for _, inc := range c.incidents {
refs = append(refs, inc)
}
sortIncidentRefs(refs)
out := make([]Incident, 0, len(refs))
for _, inc := range refs {
out = append(out, cloneIncident(*inc))
}
return out
}
// mergeLocked folds f into inc. merged=true means this is a join into
// an existing incident (bumps findings_merged_total); merged=false means
// the caller already created the incident and is using mergeLocked only
// to seed the first finding -- in that case the create path owns the
// "did a new incident appear" tally.
func (c *Correlator) mergeLocked(inc *Incident, f alert.Finding, now time.Time, merged bool) {
c.mergeLockedWithPersistence(inc, f, now, merged, false)
}
func (c *Correlator) mergeAndPersistLocked(inc *Incident, f alert.Finding, now time.Time, merged bool) {
c.mergeLockedWithPersistence(inc, f, now, merged, true)
}
func (c *Correlator) mergeLockedWithPersistence(inc *Incident, f alert.Finding, now time.Time, merged, forcePersist bool) {
if merged {
c.counters.findingsMergedTotal.Add(1)
}
transition := c.mutateWithFindingLocked(inc, f, now)
c.persistAfterMergeLocked(inc, now, merged, transition || forcePersist)
}
// mutateWithFindingLocked folds f into inc and reports whether the fold caused
// a state-changing transition (kind change or severity escalation). It does not
// persist: the caller decides synchronous vs debounced persistence via
// persistAfterMergeLocked. Splitting mutation from persistence lets the
// credential_spray path apply its own escalation before a single persist,
// avoiding a redundant second write.
func (c *Correlator) mutateWithFindingLocked(inc *Incident, f alert.Finding, now time.Time) (transition bool) {
// Re-classify before appending so timeline-aware compound rules
// see the unchanged history; the new finding is passed in
// explicitly so its Check participates in the compound check.
priorKind := inc.Kind
maybeReclassifyKind(inc, f)
if inc.Kind != priorKind {
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_kind_changed",
Result: "ok",
Details: string(priorKind) + " -> " + string(inc.Kind),
})
transition = true
}
inc.Findings = appendCappedFingerprint(inc.Findings, f.Fingerprint())
ev := IncidentEvent{
FindingID: alert.FindingID(f),
Time: f.Timestamp,
Kind: "finding",
Check: f.Check,
Severity: f.Severity.String(),
Message: f.Message,
}
if f.Process != nil {
ev.PID = f.Process.PID
ev.UID = f.Process.UID
ev.Process = f.Process.Comm
}
if f.FilePath != "" {
ev.Path = f.FilePath
}
if f.SourceIP != "" {
ev.RemoteIP = f.SourceIP
}
inc.Timeline = appendCappedTimeline(inc.Timeline, ev)
if f.Severity > inc.Severity {
from := inc.Severity
inc.Severity = f.Severity
c.counters.severityChangedTotal.Add(1)
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_severity_changed",
Result: "ok",
Details: from.String() + " -> " + f.Severity.String(),
})
transition = true
}
inc.UpdatedAt = now
return transition
}
// persistAfterMergeLocked persists a just-merged incident. Incident creation
// (merged==false) and synchronous merge writes persist immediately so a restart
// never loses an open/escalate/kind-change/threshold-promotion. A pure
// bookkeeping merge is debounced so a sustained-attack incident does not fsync
// its whole blob per finding.
func (c *Correlator) persistAfterMergeLocked(inc *Incident, now time.Time, merged, syncPersist bool) {
if !merged || syncPersist {
c.markPersistedLocked(inc.ID, now)
c.persistLocked(*inc)
return
}
c.persistDebouncedLocked(inc, now)
}
func (c *Correlator) markPersistedLocked(id string, at time.Time) {
c.lastPersistAt[id] = at
}
// persistDebouncedLocked writes inc only when at least incidentPersistDebounce
// has elapsed since its last merge-path store write; otherwise the update stays
// in memory (durable state is still captured by the next transition or close).
// Caller holds c.mu.
func (c *Correlator) persistDebouncedLocked(inc *Incident, now time.Time) {
if last, ok := c.lastPersistAt[inc.ID]; ok && now.Sub(last) < incidentPersistDebounce {
if c.cfg.Persist != nil {
c.persistence.deferWrite(inc.ID)
}
return
}
c.markPersistedLocked(inc.ID, now)
c.persistLocked(*inc)
}
// FlushPendingPersists writes incidents whose latest bookkeeping-only merge was
// skipped by the debounce window. It is intentionally cheap when no incidents
// are dirty and is used by shutdown hooks before the store closes.
func (c *Correlator) FlushPendingPersists() int {
var persist []*queuedPersist
c.mu.Lock()
now := c.now()
for _, id := range c.persistence.deferredIDs() {
inc, ok := c.incidents[id]
if !ok {
c.persistence.discardDeferred(id)
delete(c.lastPersistAt, id)
continue
}
c.markPersistedLocked(id, now)
if req, ok := c.queuePersistLocked(*inc); ok {
persist = append(persist, req)
}
}
c.mu.Unlock()
c.runQueuedPersists(persist)
return len(persist)
}
// appendCappedFingerprint appends fp to fps and trims to
// maxIncidentFindings via first-half + last-half retention when the
// cap is crossed. Keeps the original opening signals plus the most
// recent traffic, dropping the middle.
func appendCappedFingerprint(fps []string, fp string) []string {
fps = append(fps, fp)
if len(fps) <= maxIncidentFindings {
return fps
}
real := make([]string, 0, len(fps))
elided := 0
for _, existing := range fps {
if n, ok := fingerprintTruncationCount(existing); ok {
elided += n
continue
}
real = append(real, existing)
}
if len(real) <= maxIncidentFindings {
return fingerprintsWithTruncationMarker(real, elided)
}
half := maxIncidentFindings / 2
tailLen := maxIncidentFindings - half
elided += len(real) - half - tailLen
head := append([]string(nil), real[:half]...)
tail := append([]string(nil), real[len(real)-tailLen:]...)
gap := []string{formatFingerprintTruncation(elided)}
return append(append(head, gap...), tail...)
}
// appendCappedTimeline behaves the same as appendCappedFingerprint
// for the operator-visible IncidentEvent slice; the truncation marker
// is rendered as a synthetic "truncated" event so the UI can show a
// "X events elided" row.
func appendCappedTimeline(events []IncidentEvent, ev IncidentEvent) []IncidentEvent {
events = append(events, ev)
if len(events) <= maxIncidentTimeline {
return events
}
real := make([]IncidentEvent, 0, len(events))
elided := 0
var markerTime time.Time
for _, existing := range events {
if n, ok := timelineTruncationCount(existing); ok {
elided += n
if markerTime.IsZero() || existing.Time.Before(markerTime) {
markerTime = existing.Time
}
continue
}
real = append(real, existing)
}
if len(real) <= maxIncidentTimeline {
return timelineWithTruncationMarker(real, elided, markerTime)
}
half := maxIncidentTimeline / 2
tailLen := maxIncidentTimeline - half
if markerTime.IsZero() {
markerTime = real[half].Time
}
elided += len(real) - half - tailLen
head := append([]IncidentEvent(nil), real[:half]...)
tail := append([]IncidentEvent(nil), real[len(real)-tailLen:]...)
gap := []IncidentEvent{timelineTruncationMarker(elided, markerTime)}
return append(append(head, gap...), tail...)
}
func fingerprintsWithTruncationMarker(fps []string, elided int) []string {
if elided == 0 {
return fps
}
half := len(fps) / 2
out := make([]string, 0, len(fps)+1)
out = append(out, fps[:half]...)
out = append(out, formatFingerprintTruncation(elided))
out = append(out, fps[half:]...)
return out
}
func fingerprintTruncationCount(fp string) (int, bool) {
if !strings.HasPrefix(fp, incidentFingerprintTruncatedMark) {
return 0, false
}
rest := strings.TrimPrefix(fp, incidentFingerprintTruncatedMark)
fields := strings.Fields(rest)
if len(fields) == 0 {
return 1, true
}
n, err := strconv.Atoi(fields[0])
if err != nil || n < 1 {
return 1, true
}
return n, true
}
func formatFingerprintTruncation(count int) string {
return incidentFingerprintTruncatedMark + strconv.Itoa(count) + " findings elided"
}
func timelineWithTruncationMarker(events []IncidentEvent, elided int, markerTime time.Time) []IncidentEvent {
if elided == 0 {
return events
}
if markerTime.IsZero() && len(events) > 0 {
markerTime = events[len(events)/2].Time
}
half := len(events) / 2
out := make([]IncidentEvent, 0, len(events)+1)
out = append(out, events[:half]...)
out = append(out, timelineTruncationMarker(elided, markerTime))
out = append(out, events[half:]...)
return out
}
func timelineTruncationCount(ev IncidentEvent) (int, bool) {
if ev.Kind != incidentTimelineTruncatedKind {
return 0, false
}
fields := strings.Fields(ev.Message)
if len(fields) == 0 {
return 1, true
}
n, err := strconv.Atoi(fields[0])
if err != nil || n < 1 {
return 1, true
}
return n, true
}
func timelineTruncationMarker(count int, at time.Time) IncidentEvent {
return IncidentEvent{
Time: at,
Kind: incidentTimelineTruncatedKind,
Message: strconv.Itoa(count) + " events elided to cap incident size",
}
}
// persistLocked invokes the Persist callback while temporarily releasing
// the correlator mutex so a re-entrant Persist that reads Correlator
// state does not deadlock. The caller MUST already hold c.mu; the
// deferred re-Lock keeps the "mu held on return" contract that
// mergeLocked's callers rely on.
func (c *Correlator) persistLocked(snap Incident) {
req, ok := c.queuePersistLocked(snap)
if !ok {
return
}
c.mu.Unlock()
defer c.mu.Lock()
c.runQueuedPersist(req)
}
// RecordOperatorBlock notes on the incident that an operator blocked its
// address from the incident view, and settles the automatic escalation ladder
// to match: the hand-off has no business re-requesting a block for an address
// the operator just blocked. A zero ttl is a permanent block, which the ladder
// records as never lapsing.
func (c *Correlator) RecordOperatorBlock(id, ip string, ttl time.Duration) error {
c.mu.Lock()
defer c.mu.Unlock()
inc, ok := c.incidents[id]
if !ok {
return ErrIncidentNotFound
}
ip = normalizeIncidentRemoteIP(ip)
if ip == "" || ip != incidentBlockCandidate(inc) {
return errors.New("incident: blocked address does not match incident source")
}
now := c.now()
if incidentStatusActive(inc.Status) {
count := inc.AutoBlock.Count
// Refreshing an existing block is not a recurrence after expiry.
if inc.AutoBlock.lapsed(now) {
count++
}
inc.AutoBlock = AutoBlockState{Count: count, LastAt: now}
if ttl > 0 {
inc.AutoBlock.ExpiresAt = now.Add(ttl)
}
} else {
// A closed record now carries an operator decision and must keep
// that decision for the operator retention window.
inc.UpdatedAt = now
inc.ClosedAt = now
inc.ClosedBy = "operator"
}
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "operator_block",
Result: "ok",
Details: ip + " blocked by operator " + blockDurationLabel(ttl),
})
c.markPersistedLocked(id, now)
c.persistLocked(*inc)
return nil
}
// resetAutoBlockLocked ends the ladder and invalidates any callback still
// running for this episode without releasing its concurrency guard early.
func (c *Correlator) resetAutoBlockLocked(inc *Incident) {
inc.AutoBlock = AutoBlockState{}
if _, ok := c.pendingSprayBlocks[inc.ID]; ok {
c.pendingSprayBlocks[inc.ID] = false
}
if _, ok := c.pendingIncidentBlocks[inc.ID]; ok {
c.pendingIncidentBlocks[inc.ID] = false
}
}
// SetStatus transitions an incident's status. On Resolved/Dismissed
// the incident is unbound from the active byKey index so future
// findings for the same correlation key start a fresh incident.
// Returns ErrIncidentNotFound if id is unknown.
func (c *Correlator) SetStatus(id string, status Status, details string) error {
if !validStatus(status) {
return fmt.Errorf("incident: invalid status %q", status)
}
c.mu.Lock()
defer c.mu.Unlock()
inc, ok := c.incidents[id]
if !ok {
return ErrIncidentNotFound
}
if inc.Status == status && (incidentStatusActive(status) || inc.ClosedBy == "operator") {
return nil
}
now := c.now()
from := inc.Status
inc.Status = status
inc.UpdatedAt = now
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_status_changed",
Result: "ok",
Details: string(from) + " -> " + string(status) + ": " + details,
})
c.counters.statusChangedTotal.Add(1)
if status == StatusResolved || status == StatusDismissed {
// SetStatus is an operator decision, including confirmation or
// dismissal of a record the daemon already closed.
inc.ClosedAt = now
inc.ClosedBy = "operator"
// Closing ends the episode. A later recurrence is new activity and
// starts at the bottom of the escalation ladder, rather than jumping
// to a permanent block off the back of a long-closed incident.
c.resetAutoBlockLocked(inc)
c.unbindLocked(id)
if c.spray != nil {
c.spray.UnbindIncident(id)
}
} else {
// Reverting from resolved/dismissed back to open or contained
// clears the close attribution so future closes attribute correctly.
inc.ClosedAt = time.Time{}
inc.ClosedBy = ""
c.bindLocked(inc)
}
c.markPersistedLocked(id, now)
c.persistLocked(*inc)
return nil
}
// CloseStale auto-resolves Open / Contained incidents whose UpdatedAt is
// older than the per-kind threshold in `idleThresholds`. Kinds absent
// from the map are never closed (the caller decides which kinds expire).
// dryRun=true counts decisions without mutating state, so an operator
// can validate thresholds before flipping the live switch. Returns
// (closed, dryRun, total-scanned).
//
// Closing unbinds both the incident key and any spray detector state,
// so future findings are evaluated as new activity instead of merging
// into the closed incident.
func (c *Correlator) CloseStale(now time.Time, idleThresholds map[Kind]time.Duration, dryRun bool) (closed, dryRunCount, scanned int) {
closed, dryRunCount, scanned, _ = c.CloseStaleLimited(now, idleThresholds, dryRun, 0)
return closed, dryRunCount, scanned
}
// CloseStaleLimited is CloseStale with a per-call cap on the number of
// incidents it resolves. limit <= 0 means unbounded (the CloseStale
// behaviour). When the cap is hit, more=true signals the caller that
// stale incidents remain so it can schedule a prompt follow-up sweep
// instead of waiting the full auto-close interval. Bounding the work
// keeps a large post-restart backlog from holding c.mu and bursting
// thousands of bbolt persists in a single tick. The cap applies only to
// live closes; a dry-run pass always scans the full set so its counters
// stay accurate.
func (c *Correlator) CloseStaleLimited(now time.Time, idleThresholds map[Kind]time.Duration, dryRun bool, limit int) (closed, dryRunCount, scanned int, more bool) {
if len(idleThresholds) == 0 {
return 0, 0, 0, false
}
var persist []*queuedPersist
c.mu.Lock()
for id, inc := range c.incidents {
if inc.Status != StatusOpen && inc.Status != StatusContained {
continue
}
threshold, ok := idleThresholds[inc.Kind]
if !ok || threshold <= 0 {
continue
}
idle := now.Sub(inc.UpdatedAt)
if idle <= threshold {
continue
}
if !dryRun && limit > 0 && closed >= limit {
more = true
break
}
scanned++
if dryRun {
c.counters.autoCloseDryRunTotal.Add(1)
dryRunCount++
continue
}
from := inc.Status
inc.Status = StatusResolved
inc.UpdatedAt = now
inc.ClosedAt = now
inc.ClosedBy = "auto:stale"
c.resetAutoBlockLocked(inc)
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_auto_closed",
Result: "ok",
Details: string(from) + " -> resolved: stale " + idle.Truncate(time.Second).String(),
})
c.counters.statusChangedTotal.Add(1)
c.counters.autoClosedTotal.Add(1)
c.unbindLocked(id)
if c.spray != nil {
c.spray.UnbindIncident(id)
}
c.markPersistedLocked(id, now)
if req, ok := c.queuePersistLocked(*inc); ok {
persist = append(persist, req)
}
closed++
}
c.mu.Unlock()
c.runQueuedPersists(persist)
return closed, dryRunCount, scanned, more
}
// closeIncidentLocked force-resolves an active incident with the given
// attribution. Caller holds c.mu and must queue/run the persist of *inc.
func (c *Correlator) closeIncidentLocked(inc *Incident, id string, now time.Time, by, detail string) {
c.resetAutoBlockLocked(inc)
from := inc.Status
inc.Status = StatusResolved
inc.UpdatedAt = now
inc.ClosedAt = now
inc.ClosedBy = by
inc.Actions = append(inc.Actions, IncidentAction{
Time: now,
Action: "incident_auto_closed",
Result: "ok",
Details: string(from) + " -> resolved: " + detail,
})
c.counters.statusChangedTotal.Add(1)
c.counters.autoClosedTotal.Add(1)
c.unbindLocked(id)
if c.spray != nil {
c.spray.UnbindIncident(id)
}
c.markPersistedLocked(id, now)
}
// CloseStaleByAge is a kind-agnostic safety cap. It force-closes any Open or
// Contained incident whose UpdatedAt is older than maxAge, independent of the
// operator's per-kind auto-close thresholds. Without it, disabling auto-close
// (or omitting a kind from the threshold map) lets active incidents accumulate
// without bound in memory and bbolt on a host under sustained attack. limit
// bounds closures per call (0 = unbounded); more reports a remaining backlog.
func (c *Correlator) CloseStaleByAge(now time.Time, maxAge time.Duration, limit int) (closed int, more bool) {
if maxAge <= 0 {
return 0, false
}
var persist []*queuedPersist
c.mu.Lock()
for id, inc := range c.incidents {
if inc.Status != StatusOpen && inc.Status != StatusContained {
continue
}
if now.Sub(inc.UpdatedAt) <= maxAge {
continue
}
if limit > 0 && closed >= limit {
more = true
break
}
c.closeIncidentLocked(inc, id, now, "auto:age_cap", "stale age cap "+maxAge.Truncate(time.Second).String())
if req, ok := c.queuePersistLocked(*inc); ok {
persist = append(persist, req)
}
closed++
}
c.mu.Unlock()
c.runQueuedPersists(persist)
return closed, more
}
// EnforceActiveCap bounds how many Open/Contained incidents are held in
// memory. When the active count exceeds maxActive it force-closes the oldest
// (by UpdatedAt) until the count is back at maxActive or the per-call limit is
// reached. This protects against a flood of distinct incidents arriving within
// the age-cap window. maxActive <= 0 disables the cap; limit <= 0 is
// unbounded. more reports that incidents remained over the cap after limit.
func (c *Correlator) EnforceActiveCap(now time.Time, maxActive, limit int) (closed int, more bool) {
if maxActive <= 0 {
return 0, false
}
c.mu.Lock()
type activeRef struct {
id string
updated time.Time
}
var active []activeRef
for id, inc := range c.incidents {
if inc.Status == StatusOpen || inc.Status == StatusContained {
active = append(active, activeRef{id: id, updated: inc.UpdatedAt})
}
}
if len(active) <= maxActive {
c.mu.Unlock()
return 0, false
}
sort.Slice(active, func(i, j int) bool { return active[i].updated.Before(active[j].updated) })
overflow := len(active) - maxActive
var persist []*queuedPersist
for _, a := range active {
if overflow <= 0 {
break
}
if limit > 0 && closed >= limit {
more = true
break
}
inc := c.incidents[a.id]
if inc == nil || (inc.Status != StatusOpen && inc.Status != StatusContained) {
continue
}
c.closeIncidentLocked(inc, a.id, now, "auto:active_cap", "active incident cap "+strconv.Itoa(maxActive))
if req, ok := c.queuePersistLocked(*inc); ok {
persist = append(persist, req)
}
closed++
overflow--
}
c.mu.Unlock()
c.runQueuedPersists(persist)
return closed, more
}
// validStatus reports whether s is one of the four spec-defined values.
// Guards SetStatus against arbitrary strings reaching the persisted
// timeline; control-socket and webui handlers also reject early but
// the correlator owns the type and must not trust callers.
func validStatus(s Status) bool {
switch s {
case StatusOpen, StatusContained, StatusResolved, StatusDismissed:
return true
}
return false
}
func incidentStatusActive(s Status) bool {
return s == StatusOpen || s == StatusContained
}
// IncrementCompactedTotal bumps the compaction counter by n. Called
// from the daemon-side retention scheduler after store.CompactIncidents
// removes records. Negative inputs are ignored so a buggy caller cannot
// underflow the monotonic counter.
func (c *Correlator) IncrementCompactedTotal(n int) {
if n < 0 {
return
}
c.counters.compactedTotal.Add(uint64(n))
}
// PruneClosedOlderThan removes resolved/dismissed incidents older than
// retention from the in-memory map. Store compaction removes the durable
// records; this keeps API/control snapshots from serving stale incidents
// until the next daemon restart.
func (c *Correlator) PruneClosedOlderThan(now time.Time, retention ClosedRetention) int {
c.mu.Lock()
defer c.mu.Unlock()
pruned := 0
var prunedIDs []string
prunedSet := make(map[string]struct{})
for id, inc := range c.incidents {
if !retention.Expired(*inc, now) {
continue
}
delete(c.incidents, id)
delete(c.lastPersistAt, id)
c.persistence.discardDeferred(id)
prunedSet[id] = struct{}{}
if c.spray != nil {
prunedIDs = append(prunedIDs, id)
}
pruned++
}
// Walk the active index once. Scanning it for each expired record
// stalls all incident activity during a large retention backlog.
for key, id := range c.byKey {
if _, ok := prunedSet[id]; ok {
delete(c.byKey, key)
}
}
if c.spray != nil {
c.spray.UnbindIncidents(prunedIDs)
}
return pruned
}
// Restore re-hydrates correlator state from a list previously loaded
// from the store. Open and Contained incidents are bound to the
// byKey index so a finding arriving inside the merge window joins the
// existing incident; Resolved/Dismissed incidents are loaded into the
// id map only (Get still returns them) but do NOT claim their key, so
// future findings for the same key start a fresh incident.
func (c *Correlator) Restore(incidents []Incident) {
c.mu.Lock()
defer c.mu.Unlock()
for i := range incidents {
inc := incidents[i]
c.incidents[inc.ID] = &inc
delete(c.lastPersistAt, inc.ID)
c.persistence.discardDeferred(inc.ID)
if inc.Status != StatusOpen && inc.Status != StatusContained {
continue
}
c.bindLocked(&inc)
// credential_spray super-incidents persist in bbolt but the spray
// detector's perIP map is in-memory only. Without this rehydration
// step a daemon restart while an attacker is mid-spray causes the
// detector to re-trip and open a duplicate super-incident even
// though the original is still active.
if c.spray != nil && inc.Kind == KindCredentialSpray && inc.CorrelationKey != nil {
ip := inc.CorrelationKey.RemoteIP
if ip != "" {
c.spray.Rehydrate(ip, inc.ID, inc.UpdatedAt)
}
}
}
}
func (c *Correlator) bindLocked(inc *Incident) {
key, ok := incidentKey(*inc)
if !ok {
return
}
c.byKey[keyString(key)] = inc.ID
}
func (c *Correlator) unbindLocked(id string) {
// Scan-and-delete by value rather than rebuilding the key: incidents
// can be keyed by account, domain, mailbox, process, remote IP, or a
// combination. byKey only holds active incidents so the scan is bounded.
for k, v := range c.byKey {
if v == id {
delete(c.byKey, k)
}
}
}
func incidentKey(inc Incident) (Key, bool) {
if inc.CorrelationKey != nil && !inc.CorrelationKey.IsEmpty() {
return canonicalizeKey(*inc.CorrelationKey), true
}
key := Key{Account: inc.Account, Domain: inc.Domain, Mailbox: inc.Mailbox}
key = canonicalizeKey(key)
if key.IsEmpty() {
return Key{}, false
}
return key, true
}
func cloneIncident(in Incident) Incident {
out := in
out.Findings = append([]string(nil), in.Findings...)
out.Timeline = append([]IncidentEvent(nil), in.Timeline...)
out.Actions = append([]IncidentAction(nil), in.Actions...)
if in.CorrelationKey != nil {
key := *in.CorrelationKey
out.CorrelationKey = &key
}
return out
}
func cloneKey(k Key) *Key {
if k.IsEmpty() {
return nil
}
key := k
return &key
}
// keyString serializes a Key into a stable string for the byKey map.
// All fields selected by KeyFor must be encoded so distinct findings
// (e.g. different PID-only processes or different remote IPs) do not
// collapse to the same bucket and falsely merge.
func keyString(k Key) string {
return fmt.Sprintf("%d:%s|%d:%s|%d:%s|%d:%s|%d|%d|%d:%s",
len(k.Host), k.Host,
len(k.Account), k.Account,
len(k.Mailbox), k.Mailbox,
len(k.Domain), k.Domain,
k.UID,
k.PID,
len(k.RemoteIP), k.RemoteIP,
)
}
func opensIncidentImmediately(f alert.Finding) bool {
if f.Severity >= alert.Critical {
return true
}
// A copy-forward is graded Warning when its owner cannot be distinguished
// from an attacker by rule shape alone. It still needs first-hit incident
// review; otherwise lowering alert severity silently removes the existing
// mailbox-takeover correlation path.
if f.Check == "email_filter_exfil" {
return true
}
// Reputation severity describes the observed surface, not confidence in
// the threat-intel match. Preserve the pre-grading first-hit incident path.
if f.Check == "ip_reputation" {
return true
}
return f.Severity >= alert.High && ClassifyKind(f) == KindHostIntegrityRisk
}
func newIncidentID() string {
var buf [6]byte
_, _ = rand.Read(buf[:])
return "inc_" + hex.EncodeToString(buf[:])
}
// maybeBlockSprayLocked is the single decision point for the
// credential_spray firewall hand-off. Unlike a transition-only hook,
// this runs on every spray decision (open, merge, escalate) so an
// operator who arms BlockAtSeverity AFTER an incident has already
// reached the configured severity still gets a block on the next
// matching finding. Idempotency is provided by
// triggerSprayBlockLocked's action-presence and in-flight guards, so
// calling this helper repeatedly against the same incident emits one
// live firewall call at a time.
//
// Returns the callback that the caller must invoke after releasing
// c.mu (matches the existing sprayDecisionOpen contract), or nil when
// no block is owed.
func (c *Correlator) maybeBlockSprayLocked(inc *Incident, ip string, hits int, now time.Time, reason string) func() {
if inc == nil || c.spray == nil || c.cfg.OnSprayBlock == nil {
return nil
}
if !incidentStatusActive(inc.Status) {
return nil
}
if !c.sprayBlockAllowed() {
return nil
}
if incidentAutoBlockExcludedOnly(inc) {
return nil
}
switch strings.ToLower(c.spray.cfg.BlockAtSeverity) {
case "high":
if inc.Severity < alert.High {
return nil
}
case "critical":
if inc.Severity < alert.Critical {
return nil
}
default:
return nil
}
return c.triggerSprayBlockLocked(inc, ip, hits, now, reason)
}
func (c *Correlator) sprayBlockAllowed() bool {
return c.cfg.CanSprayBlock == nil || c.cfg.CanSprayBlock()
}
// maybeBlockIncidentLocked is the decision point for the generic
// incident-driven firewall hand-off. Runs on every create / merge so
// an operator who arms AutoBlock AFTER an incident has already crossed
// the gate still gets a block on the next finding. Idempotent via the
// action-presence and in-flight guards in triggerIncidentBlockLocked.
//
// Skips credential_spray incidents -- the dedicated spray hand-off owns
// those so we avoid double-firing.
//
// Returns the callback that the caller must invoke after releasing
// c.mu, or nil when no block is owed.
func (c *Correlator) maybeBlockIncidentLocked(inc *Incident, now time.Time, why string) func() {
if inc == nil || c.cfg.OnIncidentBlock == nil || !c.cfg.AutoBlock.Enabled {
return nil
}
if !incidentStatusActive(inc.Status) {
return nil
}
if inc.Kind == KindCredentialSpray {
return nil
}
if c.sprayOwnsIncident(inc) {
return nil
}
if !c.incidentBlockAllowed() {
return nil
}
ip := incidentBlockCandidate(inc)
if ip == "" {
return nil
}
if len(c.cfg.AutoBlock.Kinds) > 0 && !c.cfg.AutoBlock.Kinds[inc.Kind] {
return nil
}
if incidentAutoBlockExcludedOnly(inc) {
return nil
}
switch strings.ToLower(c.cfg.AutoBlock.BlockAtSeverity) {
case "high":
if inc.Severity < alert.High {
return nil
}
case "critical":
if inc.Severity < alert.Critical {
return nil
}
default:
return nil
}
return c.triggerIncidentBlockLocked(inc, ip, now, why)
}
// incidentAutoBlockExcludedOnly keeps advisory incident signals visible without
// letting them become firewall evidence unless another blockable finding joins.
func incidentAutoBlockExcludedOnly(inc *Incident) bool {
seen := false
for _, ev := range inc.Timeline {
if ev.Kind == incidentTimelineTruncatedKind {
continue
}
if ev.Kind != "finding" || ev.Check == "" {
continue
}
seen = true
if !incidentEventAutoBlockExcluded(ev) {
return false
}
}
return seen
}
const establishedMailSourceMarker = "(established multi-mailbox source)"
func incidentEventAutoBlockExcluded(ev IncidentEvent) bool {
switch strings.ToLower(strings.TrimSpace(ev.Check)) {
case "cpanel_file_upload", "cpanel_file_upload_realtime",
"cpanel_login", "cpanel_login_realtime", "ftp_login", "ftp_login_realtime",
"webmail_login_realtime", "pam_login":
// Retained incidents can still carry the old severity of audit events.
return true
case "ftp_login_after_bruteforce",
"mail_bruteforce_suspected",
"modsec_classifier_gap",
"modsec_low_confidence_burst":
return true
case "mail_account_compromised":
severity := strings.ToUpper(strings.TrimSpace(ev.Severity))
if severity != "" {
return severity != alert.Critical.String()
}
// Incidents persisted before timeline events carried severity can still
// contain this exact advisory marker. Critical compromise messages never
// carry it, so they remain blockable after restore.
return strings.HasSuffix(strings.TrimSpace(ev.Message), establishedMailSourceMarker)
default:
return false
}
}
func (c *Correlator) incidentBlockAllowed() bool {
return c.cfg.CanIncidentBlock == nil || c.cfg.CanIncidentBlock()
}
// triggerIncidentBlockLocked is the per-incident emit point for the
// generic auto-block path. It marks the incident as in-flight, returns
// the deferred callback, then appends "incident_block_requested" only
// when the callback reports a live block request. Dry-run attempts are
// intentionally not latched so a later finding can retry after the
// operator disables dry-run.
func (c *Correlator) triggerIncidentBlockLocked(inc *Incident, ip string, now time.Time, why string) func() {
if inc == nil || c.cfg.OnIncidentBlock == nil {
return nil
}
if !inc.AutoBlock.lapsed(now) {
return nil
}
if _, ok := c.pendingIncidentBlocks[inc.ID]; ok {
return nil
}
if _, ok := c.pendingSprayBlocks[inc.ID]; ok {
return nil
}
c.pendingIncidentBlocks[inc.ID] = true
prior := inc.AutoBlock
incidentID := inc.ID
attempt := inc.AutoBlock.Count + 1
ttl := blockTTLForAttempt(attempt, c.cfg.AutoBlock.BlockExpiry)
reason := "incident " + string(inc.Kind) + " " + inc.Severity.String() + " (" + why + ")"
if attempt > 1 {
reason += "; block " + strconv.Itoa(attempt) + " after the previous one lapsed"
}
onBlock := c.cfg.OnIncidentBlock
findingID := incidentBlockFindingID(inc, ip)
return func() {
var live bool
callbackReturned := false
// The in-flight slot must clear even if onBlock panics. Otherwise,
// later findings keep seeing the incident as already in flight and
// skip the auto-block path for the rest of the incident lifetime.
defer func() {
c.mu.Lock()
defer c.mu.Unlock()
valid := c.pendingIncidentBlocks[incidentID]
delete(c.pendingIncidentBlocks, incidentID)
if !valid || !callbackReturned || !live {
return
}
current, ok := c.incidents[incidentID]
if !ok || current != inc || !incidentStatusActive(current.Status) || current.AutoBlock != prior {
return
}
current.AutoBlock = AutoBlockState{Count: attempt, LastAt: now}
if ttl > 0 {
current.AutoBlock.ExpiresAt = now.Add(ttl)
}
current.Actions = append(current.Actions, IncidentAction{
Time: now,
Action: "incident_block_requested",
Result: "ok",
Details: ip + " " + reason + " " + blockDurationLabel(ttl),
})
c.markPersistedLocked(incidentID, c.now())
c.persistLocked(*current)
}()
c.mu.Lock()
valid := c.pendingIncidentBlocks[incidentID] && c.incidents[incidentID] == inc && incidentStatusActive(inc.Status) && inc.AutoBlock == prior
c.mu.Unlock()
if !valid {
return
}
live = onBlock(ip, reason, ttl, findingID)
callbackReturned = true
}
}
func incidentBlockCandidate(inc *Incident) string {
if inc == nil {
return ""
}
if inc.CorrelationKey != nil {
if ip := normalizeIncidentRemoteIP(inc.CorrelationKey.RemoteIP); ip != "" {
return ip
}
}
var candidate string
for _, ev := range inc.Timeline {
if ev.Kind == incidentTimelineTruncatedKind {
return ""
}
ip := normalizeIncidentRemoteIP(ev.RemoteIP)
if ip == "" {
continue
}
if candidate == "" {
candidate = ip
continue
}
if candidate != ip {
return ""
}
}
return candidate
}
func (c *Correlator) sprayOwnsIncident(inc *Incident) bool {
if c == nil || c.spray == nil || inc == nil {
return false
}
owned := false
for _, ev := range inc.Timeline {
if ev.RemoteIP == "" {
continue
}
if !c.spray.cfg.PerCheck[ev.Check] {
return false
}
owned = true
}
return owned
}
func normalizeIncidentRemoteIP(raw string) string {
raw = strings.TrimSpace(raw)
if raw == "" {
return ""
}
if host, _, err := net.SplitHostPort(raw); err == nil {
raw = host
}
raw = strings.Trim(raw, "[]")
ip := net.ParseIP(raw)
if ip == nil || ip.IsLoopback() || ip.IsUnspecified() {
return ""
}
return ip.String()
}
// blockDurationLabel renders a block's lifetime for the incident action so an
// operator reading the timeline can see the ladder escalating.
func blockDurationLabel(ttl time.Duration) string {
if ttl <= 0 {
return "(permanent)"
}
return "(" + ttl.String() + ")"
}
func hasIncidentAction(actions []IncidentAction, action string) bool {
for _, a := range actions {
if a.Action == action {
return true
}
}
return false
}
// triggerSprayBlockLocked invokes the operator-supplied OnSprayBlock
// callback at most once while a spray block is already in flight and
// appends an audit action only when the callback reports the firewall
// actually applied the block. Caller holds c.mu. The returned callback
// must be invoked after unlocking. Idempotent after success; declined
// callbacks can retry on a later finding after the in-flight marker is
// cleared.
func (c *Correlator) triggerSprayBlockLocked(inc *Incident, ip string, hits int, now time.Time, why string) func() {
if inc == nil || c.cfg.OnSprayBlock == nil {
return nil
}
if !inc.AutoBlock.lapsed(now) {
return nil
}
if _, ok := c.pendingSprayBlocks[inc.ID]; ok {
return nil
}
if _, ok := c.pendingIncidentBlocks[inc.ID]; ok {
return nil
}
c.pendingSprayBlocks[inc.ID] = true
prior := inc.AutoBlock
attempt := inc.AutoBlock.Count + 1
ttl := blockTTLForAttempt(attempt, c.spray.cfg.BlockExpiry)
reason := "credential_spray: " + strconv.Itoa(hits) + " distinct mailboxes (" + why + ")"
if attempt > 1 {
reason += "; block " + strconv.Itoa(attempt) + " after the previous one lapsed"
}
onSprayBlock := c.cfg.OnSprayBlock
findingID := incidentBlockFindingID(inc, ip)
incidentID := inc.ID
return func() {
var live bool
callbackReturned := false
// Mirror the panic-safety guarantee from triggerIncidentBlockLocked:
// the in-flight slot must clear even if onSprayBlock panics so a
// recurring panic class does not latch the credential_spray
// incident out of the auto-block path forever.
defer func() {
c.mu.Lock()
defer c.mu.Unlock()
valid := c.pendingSprayBlocks[incidentID]
delete(c.pendingSprayBlocks, incidentID)
if !valid || !callbackReturned || !live {
return
}
current, ok := c.incidents[incidentID]
if !ok || current != inc || !incidentStatusActive(current.Status) || current.AutoBlock != prior {
return
}
current.AutoBlock = AutoBlockState{Count: attempt, LastAt: now}
if ttl > 0 {
current.AutoBlock.ExpiresAt = now.Add(ttl)
}
current.Actions = append(current.Actions, IncidentAction{
Time: now,
Action: "credential_spray_block_requested",
Result: "ok",
Details: ip + " " + reason + " " + blockDurationLabel(ttl),
})
c.markPersistedLocked(incidentID, c.now())
c.persistLocked(*current)
}()
c.mu.Lock()
valid := c.pendingSprayBlocks[incidentID] && c.incidents[incidentID] == inc && incidentStatusActive(inc.Status) && inc.AutoBlock == prior
c.mu.Unlock()
if !valid {
return
}
live = onSprayBlock(ip, reason, ttl, findingID)
callbackReturned = true
}
}
// incidentBlockFindingID selects the latest eligible observation for this
// source while the incident lock is held. Older timelines without an audit
// identity remain unlinked; display text cannot reconstruct the original ID.
func incidentBlockFindingID(inc *Incident, ip string) string {
ip = normalizeIncidentRemoteIP(ip)
if ip == "" {
return ""
}
for i := len(inc.Timeline) - 1; i >= 0; i-- {
ev := inc.Timeline[i]
if ev.Kind == "finding" && ev.FindingID != "" && normalizeIncidentRemoteIP(ev.RemoteIP) == ip && !incidentEventAutoBlockExcluded(ev) {
return ev.FindingID
}
}
return ""
}
package incident
import (
"sort"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// IncidentGroupsScanCap is the hard upper bound on matching incidents
// grouped per BuildGroups call. Rows excluded by status/kind filters do
// not consume the cap.
const IncidentGroupsScanCap = 10000
// Group is one row of the grouped incident view: a (kind, source)
// bucket plus rolled-up counters. Source identifies what the bucket is
// keyed on -- IP for credential-spray patterns, account/domain/mailbox
// when the source IP is unknown.
type Group struct {
Key string `json:"key"`
Kind Kind `json:"kind"`
SourceKind string `json:"source_kind"`
Source string `json:"source"`
IncidentCount int `json:"incident_count"`
OpenCount int `json:"open_count"`
ContainedCount int `json:"contained_count"`
ResolvedCount int `json:"resolved_count"`
DismissedCount int `json:"dismissed_count"`
SeverityMax alert.Severity `json:"-"`
SeverityLabel string `json:"severity_max"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
SampleIDs []string `json:"sample_ids"`
}
// GroupsResponse is what BuildGroups returns. /api/v1/incidents/groups
// sends Groups as its items, with the counts next to them.
type GroupsResponse struct {
Groups []Group `json:"groups"`
TotalGroups int `json:"total_groups"`
ScannedIncidents int `json:"scanned_incidents"`
Truncated bool `json:"truncated"`
}
// GroupFilter narrows what BuildGroups buckets. Empty fields mean "no
// filter on that dimension".
type GroupFilter struct {
// StatusSet, when non-empty, restricts incidents to the listed
// statuses. Pass {StatusOpen, StatusContained} to surface only
// active incidents (the default web UI mode).
StatusSet []Status
// Kind, when non-empty, restricts incidents to a specific kind.
Kind Kind
// Offset is the starting index into the sorted group list. Used by
// the UI to page through groups when TotalGroups exceeds MaxGroups.
Offset int
// MaxGroups caps the returned slice. Zero or negative means "no
// cap"; the handler still applies a sane upper bound.
MaxGroups int
}
// BuildGroups buckets the supplied incidents by (kind, source) and
// returns the rolled-up groups sorted by incident_count desc, then
// severity_max desc, then last_seen desc. SampleIDs holds up to three
// of the most recently updated members of each group so the UI can
// drill in without a follow-up call.
//
// `incidents` may be the full correlator snapshot. The function caps
// its scan at IncidentGroupsScanCap after status/kind filtering; the
// returned `truncated` flag reports whether the cap clipped matching
// incidents. Callers fed by Correlator.Snapshot() get an already-newest-
// first slice; sort stability there means the truncation drops the oldest
// matching entries first, which is what an operator wants.
func BuildGroups(incidents []Incident, filter GroupFilter) GroupsResponse {
statusAllowed := func(Status) bool { return true }
if len(filter.StatusSet) > 0 {
set := make(map[Status]struct{}, len(filter.StatusSet))
for _, s := range filter.StatusSet {
set[s] = struct{}{}
}
statusAllowed = func(s Status) bool {
_, ok := set[s]
return ok
}
}
type sampleEntry struct {
id string
updatedAt time.Time
}
type aggregator struct {
group Group
samples []sampleEntry
statusMap map[Status]int
}
bucketKey := func(kind Kind, sourceKind, source string) string {
return string(kind) + "|" + sourceKind + ":" + source
}
buckets := make(map[string]*aggregator)
scanned := 0
truncated := false
for _, inc := range incidents {
if !statusAllowed(inc.Status) {
continue
}
if filter.Kind != "" && inc.Kind != filter.Kind {
continue
}
if scanned >= IncidentGroupsScanCap {
truncated = true
break
}
scanned++
sourceKind, source := groupSource(inc)
k := bucketKey(inc.Kind, sourceKind, source)
agg, ok := buckets[k]
if !ok {
agg = &aggregator{
group: Group{
Key: k,
Kind: inc.Kind,
SourceKind: sourceKind,
Source: source,
FirstSeen: inc.CreatedAt,
LastSeen: inc.UpdatedAt,
},
statusMap: map[Status]int{},
}
buckets[k] = agg
}
agg.group.IncidentCount++
agg.statusMap[inc.Status]++
if inc.Severity > agg.group.SeverityMax {
agg.group.SeverityMax = inc.Severity
}
if inc.CreatedAt.Before(agg.group.FirstSeen) || agg.group.FirstSeen.IsZero() {
agg.group.FirstSeen = inc.CreatedAt
}
if inc.UpdatedAt.After(agg.group.LastSeen) {
agg.group.LastSeen = inc.UpdatedAt
}
agg.samples = append(agg.samples, sampleEntry{id: inc.ID, updatedAt: inc.UpdatedAt})
}
out := make([]Group, 0, len(buckets))
for _, agg := range buckets {
agg.group.OpenCount = agg.statusMap[StatusOpen]
agg.group.ContainedCount = agg.statusMap[StatusContained]
agg.group.ResolvedCount = agg.statusMap[StatusResolved]
agg.group.DismissedCount = agg.statusMap[StatusDismissed]
agg.group.SeverityLabel = agg.group.SeverityMax.String()
// Top-3 most recently updated members.
sort.SliceStable(agg.samples, func(i, j int) bool {
return agg.samples[i].updatedAt.After(agg.samples[j].updatedAt)
})
n := len(agg.samples)
if n > 3 {
n = 3
}
ids := make([]string, n)
for i := 0; i < n; i++ {
ids[i] = agg.samples[i].id
}
agg.group.SampleIDs = ids
out = append(out, agg.group)
}
sort.SliceStable(out, func(i, j int) bool {
if out[i].IncidentCount != out[j].IncidentCount {
return out[i].IncidentCount > out[j].IncidentCount
}
if out[i].SeverityMax != out[j].SeverityMax {
return out[i].SeverityMax > out[j].SeverityMax
}
return out[i].LastSeen.After(out[j].LastSeen)
})
totalGroups := len(out)
if filter.Offset > 0 {
if filter.Offset >= len(out) {
out = out[:0]
} else {
out = out[filter.Offset:]
}
}
if filter.MaxGroups > 0 && len(out) > filter.MaxGroups {
out = out[:filter.MaxGroups]
}
return GroupsResponse{
Groups: out,
TotalGroups: totalGroups,
ScannedIncidents: scanned,
Truncated: truncated,
}
}
// groupSource derives the (source_kind, source) pair the UI uses to
// label and drill into a group. Cascade order: host > remote_ip >
// account > domain > mailbox > "_unkeyed". The IP path is the most
// useful grouping for credential-spray patterns and the most common
// shape on busy hosts.
func groupSource(inc Incident) (sourceKind, source string) {
if inc.CorrelationKey != nil && inc.CorrelationKey.Host != "" {
return "host", inc.CorrelationKey.Host
}
if inc.CorrelationKey != nil && inc.CorrelationKey.RemoteIP != "" {
return "ip", inc.CorrelationKey.RemoteIP
}
if ip := timelineRemoteIP(inc); ip != "" {
return "ip", ip
}
if inc.Account != "" {
return "account", inc.Account
}
if inc.Domain != "" {
return "domain", inc.Domain
}
if inc.Mailbox != "" {
return "mailbox", inc.Mailbox
}
return "_unkeyed", ""
}
func timelineRemoteIP(inc Incident) string {
counts := make(map[string]int)
for _, ev := range inc.Timeline {
if ev.RemoteIP != "" {
counts[ev.RemoteIP]++
}
}
best := ""
bestCount := 0
for ip, count := range counts {
if count > bestCount || count == bestCount && (best == "" || ip < best) {
best = ip
bestCount = count
}
}
return best
}
package incident
import (
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// Key is the correlation key derived from a Finding. Empty fields mean
// "not provided"; the correlator uses the most specific non-empty fields.
// Host is the synthetic local-host actor for findings whose blast radius
// is the machine itself rather than one tenant, process, or remote IP.
type Key struct {
Host string `json:"host,omitempty"`
Account string `json:"account,omitempty"`
Domain string `json:"domain,omitempty"`
Mailbox string `json:"mailbox,omitempty"`
UID int `json:"uid,omitempty"`
PID int `json:"pid,omitempty"`
RemoteIP string `json:"remote_ip,omitempty"`
}
// excludedFromIncidents reports whether a finding records posture, a response,
// or degraded visibility rather than an event that needs containment.
func excludedFromIncidents(check string) bool {
switch check {
case "vulnerable_plugins", "outdated_plugins":
// An installed weakness is not evidence that the account was entered.
return true
case "auto_block", "auto_response", "auto_response_paused", "challenge_route",
"reputation_quota_exhausted", "threat_feed_stale",
"wp_core_unverified", "wp_plugin_inventory_unverified":
// CSM output must not re-enter the incident response path. In
// particular, response findings may gain attacker attribution later.
return true
default:
return false
}
}
// IsEmpty reports whether the key has nothing to correlate on. Such
// findings are emitted normally but do not join an incident.
func (k Key) IsEmpty() bool {
return k.Host == "" && k.Account == "" && k.Domain == "" && k.Mailbox == "" && k.UID == 0 && k.PID == 0 && k.RemoteIP == ""
}
// KeyFor extracts a correlation key from a Finding. Host-integrity
// findings are keyed to the local host so unattributed root/system events
// still become incidents. TenantID, Process.Account, CPUser, and a
// /home[N]/<account>/ heuristic provide account attribution. Domain and
// Mailbox come directly from the finding. Process UID/PID and SourceIP are
// fallback identities: they should not split account/domain/mailbox
// incidents.
//
// Mailbox + Domain are canonicalised so emitters that set either the
// full local@domain form or the split (Mailbox=local, Domain=site)
// form land on the same key. Without that, two findings about the
// same mailbox split into two incidents whenever the emitters use
// different conventions.
func KeyFor(f alert.Finding) Key {
if excludedFromIncidents(f.Check) {
return Key{}
}
switch ClassifyKind(f) {
case KindHostIntegrityRisk:
return Key{Host: "host"}
case KindWebAttack, KindMailboxBruteforce:
// Inbound attacks correlate on the attacker IP; the victim
// domain/account/mailbox is the target, not the key. This
// collapses one attacker's hits across many victims into a single
// incident. ClassifyKind only returns these kinds when a source IP
// is present.
return Key{RemoteIP: f.SourceIP}
}
mailbox, domain := canonicalizeMailboxDomain(f.Mailbox, f.Domain)
k := Key{
Account: f.TenantID,
Domain: domain,
Mailbox: mailbox,
}
if f.Process != nil && f.Process.Account != "" && k.Account == "" {
k.Account = f.Process.Account
}
if k.Account == "" && f.CPUser != "" {
k.Account = f.CPUser
}
if k.Account == "" {
k.Account = accountFromHomePath(f.FilePath)
}
if f.Process != nil && !hasStableActor(k) {
if f.Process.UID != 0 {
k.UID = f.Process.UID
}
if k.UID == 0 {
k.PID = f.Process.PID
}
}
if k.Account == "" && k.Domain == "" && k.Mailbox == "" && k.UID == 0 && k.PID == 0 {
k.RemoteIP = f.SourceIP
}
return k
}
func hasStableActor(k Key) bool {
return k.Account != "" || k.Domain != "" || k.Mailbox != ""
}
// canonicalizeMailboxDomain merges Mailbox+Domain into a stable
// (Mailbox, Domain) key pair regardless of which emit convention the
// caller used. Rules:
//
// - If Mailbox already contains "@", treat it as authoritative;
// drop Domain to avoid double-keying on conflicting site.
// - If Mailbox lacks "@" and Domain is set, splice them into the
// full form. Domain is then dropped from the key (it's already
// encoded in Mailbox).
// - Domain-only findings (no Mailbox) keep the domain as the key.
//
// Domain names are case-insensitive, so only the domain component is
// lower-cased. The local part is left intact.
func canonicalizeMailboxDomain(mailbox, domain string) (string, string) {
mailbox = strings.TrimSpace(mailbox)
domain = normalizeDomainForKey(domain)
if mailbox == "" {
return "", domain
}
if local, mailboxDomain := splitMailboxForKey(mailbox); mailboxDomain != "" {
return local + "@" + mailboxDomain, ""
}
if strings.Contains(mailbox, "@") {
return mailbox, ""
}
if domain == "" {
return mailbox, ""
}
return mailbox + "@" + domain, ""
}
func canonicalizeKey(k Key) Key {
k.Mailbox, k.Domain = canonicalizeMailboxDomain(k.Mailbox, k.Domain)
return k
}
func displayMailboxDomain(mailbox, domain string) (string, string) {
mailbox = strings.TrimSpace(mailbox)
domain = strings.TrimSpace(domain)
if mailbox == "" {
return "", domain
}
if local, mailboxDomain := splitMailboxForKey(mailbox); mailboxDomain != "" {
return local + "@" + mailboxDomain, mailboxDomain
}
if domain == "" {
return mailbox, ""
}
if strings.Contains(mailbox, "@") {
return mailbox, domain
}
normalizedDomain := normalizeDomainForKey(domain)
return mailbox + "@" + normalizedDomain, normalizedDomain
}
func splitMailboxForKey(mailbox string) (string, string) {
at := strings.LastIndexByte(mailbox, '@')
if at <= 0 || at == len(mailbox)-1 {
return "", ""
}
return mailbox[:at], normalizeDomainForKey(mailbox[at+1:])
}
func normalizeDomainForKey(domain string) string {
return strings.ToLower(strings.TrimSpace(domain))
}
// accountFromHomePath parses /home[N]/<account>/... paths. Returns the
// account segment or "" if the path does not match the cPanel-style home
// layout. Pure string parsing; does not walk the filesystem.
func accountFromHomePath(p string) string {
if p == "" {
return ""
}
parts := strings.SplitN(p, "/", 4)
if len(parts) < 3 {
return ""
}
if parts[0] != "" {
return ""
}
if !strings.HasPrefix(parts[1], "home") {
return ""
}
for _, ch := range parts[1][len("home"):] {
if ch < '0' || ch > '9' {
return ""
}
}
return parts[2]
}
package incident
import (
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// ClassifyKind returns the incident Kind for a Finding using rule
// precedence: host-integrity signals first, then inbound attacks keyed on the
// attacker IP (web_attack, mailbox_bruteforce), then mailbox-takeover signals,
// then ephemeral-path process exec, then remote-IP reputation, then the
// account-scoped web-compromise default.
func ClassifyKind(f alert.Finding) Kind {
check := strings.ToLower(f.Check)
hasAttacker := strings.TrimSpace(f.SourceIP) != ""
// Host integrity -- daemon/kernel-level signals whose blast radius is
// the host itself, not a single tenant or an inbound attacker. Listed
// explicitly so account-attributed checks that share a substring
// (e.g. suspicious_crontab on a per-user spool) stay in the tenant
// bucket.
if isHostIntegrityCheck(check) {
return KindHostIntegrityRisk
}
// Inbound attack keyed on the attacker IP. A defended web hit (WAF
// block, scanner probe, login brute-force) or a failed mailbox login
// names a victim domain/account/mailbox, but that is the attack
// target, not evidence the account is compromised. Classify by the
// attacker so defended traffic does not inflate the compromise /
// takeover buckets or inherit their long retention. Requires an
// attacker IP to key on; without one the finding falls through to the
// account-scoped tiers below. Genuine compromise is recognised by the
// on-disk and behavioural signals handled lower down, not by inbound
// hits.
if hasAttacker {
if isMailAuthAttackCheck(check) {
return KindMailboxBruteforce
}
if isInboundWebAttackCheck(check) {
return KindWebAttack
}
}
// Mailbox takeover -- a Mailbox attribute or a post-authentication
// mail-abuse check name. Mailbox attribution is the strongest single
// tenant signal so it wins over the account-scoped default below. Some
// authenticated-mail findings route bare cPanel-local accounts to
// TenantID instead of Mailbox; the check-name list keeps those here
// without sweeping in domain/config mail checks or PHP relay findings.
if f.Mailbox != "" {
return KindMailboxTakeover
}
if isMailboxTakeoverCheck(check) {
return KindMailboxTakeover
}
// Post-exploit process -- exe under ephemeral paths is a strong
// indicator of staged-then-executed payloads (cryptominers, reverse
// shells) regardless of which tenant owns the parent.
if f.Process != nil && (strings.HasPrefix(f.Process.Exe, "/tmp/") ||
strings.HasPrefix(f.Process.Exe, "/var/tmp/") ||
strings.HasPrefix(f.Process.Exe, "/dev/shm/")) {
return KindPostExploitProcess
}
// Remote-IP reputation / threat-score with no victim attribution -- an
// attacking source IP flagged by reputation, not a compromised tenant.
// Kept on the strict remote-IP-keyed gate: a reputation hit tied to an
// account is about that account, not an anonymous attacker.
if isRemoteIPThreatCheck(check) && isRemoteIPKeyed(f) {
return KindWebAttack
}
// Default -- account-scoped web compromise. Most CSM findings are
// tenant-attributed web/PHP issues, so this fallback matches the
// modal incident shape operators see.
return KindWebAccountCompromise
}
// isRemoteIPKeyed reports whether a finding would correlate on its source
// IP alone -- i.e. it has no account (tenant, cPanel user, process
// account, or /home/<account>/ path), no domain, no mailbox, and no
// stable process actor (UID/PID). It mirrors KeyFor's RemoteIP fallback
// without calling KeyFor (KeyFor depends on ClassifyKind, so the reverse
// call would recurse). Callers reach this only after the host-integrity,
// mailbox, and post-exploit-process tiers have been ruled out.
func isRemoteIPKeyed(f alert.Finding) bool {
if strings.TrimSpace(f.SourceIP) == "" {
return false
}
mailbox, domain := canonicalizeMailboxDomain(f.Mailbox, f.Domain)
if mailbox != "" || domain != "" {
return false
}
if f.TenantID != "" || f.CPUser != "" {
return false
}
if f.Process != nil && f.Process.Account != "" {
return false
}
if accountFromHomePath(f.FilePath) != "" {
return false
}
if f.Process != nil && (f.Process.UID != 0 || f.Process.PID != 0) {
return false
}
return true
}
// isRemoteIPThreatCheck covers remote-IP reputation / threat-score signals
// that flag an attacking source IP rather than a compromised tenant. Like
// inbound web attacks these correlate on the source IP alone, so they belong
// in web_attack with attacker-grade retention, not the account-compromise
// bucket. Add future remote-IP reputation checks here.
func isRemoteIPThreatCheck(check string) bool {
switch strings.ToLower(strings.TrimSpace(check)) {
case "ip_reputation", "local_threat_score":
return true
default:
return false
}
}
// isMailAuthAttackCheck covers failed mail-authentication and pre-auth mail
// probe signals. These are attacker attempts keyed on the source IP, not
// evidence that the targeted mailbox was taken over.
func isMailAuthAttackCheck(check string) bool {
switch strings.ToLower(strings.TrimSpace(check)) {
case "email_auth_failure_realtime",
"mail_account_spray",
"mail_bruteforce",
"mail_bruteforce_suspected",
"mail_subnet_spray",
"smtp_account_spray",
"smtp_bruteforce",
"smtp_probe_abuse",
"smtp_subnet_spray":
return true
default:
return false
}
}
func isInboundWebAttackCheck(check string) bool {
check = strings.ToLower(strings.TrimSpace(check))
if strings.HasPrefix(check, "http_") || strings.HasPrefix(check, "modsec_") {
return true
}
switch check {
case "admin_panel_bruteforce",
"wp_login_bruteforce",
"wp_user_enumeration",
"xmlrpc_abuse",
"waf_attack_blocked",
"api_auth_failure",
"api_auth_failure_realtime",
"webmail_bruteforce",
"webmail_login_realtime",
"whm_login_realtime",
"whm_unauth_scripts_realtime":
return true
default:
return false
}
}
// hostIntegrityChecks lists check names whose scope is the host itself
// (kernel modules, system daemon configs, root-owned credential stores)
// rather than a single tenant. Findings matching one of these jump
// straight to KindHostIntegrityRisk so incident severity reflects the
// blast radius.
var hostIntegrityChecks = map[string]bool{
"bulk_password_change": true,
"sensitive_file_modified": true,
"fake_kernel_thread": true,
"integrity": true,
"shadow_change": true,
"sshd_config_change": true,
"root_password_change": true,
"uid0_account": true,
"suid_binary": true,
"bad_asn_outbound": true,
"kernel_module": true,
"crontab_change": true,
"crond_change": true,
"firewall_ipv6_unmanaged": true,
"mail_auth_backend_degraded": true,
}
func isHostIntegrityCheck(check string) bool {
return hostIntegrityChecks[strings.ToLower(strings.TrimSpace(check))]
}
func isMailboxTakeoverCheck(check string) bool {
switch check {
case "email_cloud_relay_abuse",
"email_compromised_account",
"email_credential_leak",
"email_rate_critical",
"email_rate_warning",
"email_spam_outbreak",
"email_suspicious_geo",
"mail_account_compromised",
"mail_per_account":
return true
default:
return false
}
}
package incident
import "github.com/pidginhost/csm/internal/metrics"
// RegisterMetrics binds the correlator's counters to reg. Production
// callers should pass metrics.Default(); tests pass metrics.NewRegistry()
// to keep registration isolated.
func RegisterMetrics(reg *metrics.Registry, c *Correlator) {
reg.RegisterGaugeFunc(
"csm_incidents_open",
"Open and Contained incidents currently in correlator state.",
func() float64 { return float64(c.OpenCount()) },
)
reg.RegisterCounterFunc(
"csm_incidents_created_total",
"Total incidents created by the correlator.",
func() float64 { return float64(c.counters.createdTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_incidents_severity_changed_total",
"Incident severity escalations (severity does not downgrade, so this is monotonic).",
func() float64 { return float64(c.counters.severityChangedTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_incidents_status_changed_total",
"Incident status transitions (open/contained/resolved/dismissed).",
func() float64 { return float64(c.counters.statusChangedTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_incidents_findings_merged_total",
"Findings merged into an existing incident (not counted on incident create).",
func() float64 { return float64(c.counters.findingsMergedTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_incidents_compacted_total",
"Incidents pruned by retention compaction (resolved/dismissed beyond TTL).",
func() float64 { return float64(c.counters.compactedTotal.Load()) },
)
reg.RegisterGaugeFunc(
"csm_incidents_pending",
"Findings held in the threshold gate, awaiting a second correlated finding before opening an incident.",
func() float64 { return float64(c.PendingCount()) },
)
reg.RegisterCounterFunc(
"csm_incidents_auto_closed_total",
"Open or contained incidents auto-resolved after exceeding their per-kind idle threshold.",
func() float64 { return float64(c.counters.autoClosedTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_incidents_auto_close_dry_run_total",
"Auto-close decisions counted while dry_run was active (state unchanged).",
func() float64 { return float64(c.counters.autoCloseDryRunTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_credential_spray_opened_total",
"Credential-spray super-incidents opened (one source IP brute-forcing many mailboxes).",
func() float64 { return float64(c.counters.sprayOpenedTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_credential_spray_suppressed_mailbox_takeover_total",
"Per-mailbox incidents suppressed because a credential_spray incident already covers the source IP.",
func() float64 { return float64(c.counters.spraySuppressedTotal.Load()) },
)
reg.RegisterCounterFunc(
"csm_credential_spray_dry_run_total",
"Spray decisions counted while dry_run was active (routing unchanged).",
func() float64 { return float64(c.counters.sprayDryRunTotal.Load()) },
)
reg.RegisterGaugeFunc(
"csm_credential_spray_tracked_ips",
"Source IPs currently held in the spray detector's per-IP map.",
func() float64 { return float64(c.SprayTrackedIPs()) },
)
}
package incident
import (
"sync"
"time"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/queuehealth"
)
type queuedPersist struct {
previous <-chan struct{}
done chan struct{}
snap Incident
persist func(Incident) error
at time.Time
running, failed bool
}
type persistQueue struct {
mu sync.Mutex
tail chan struct{}
waiting map[*queuedPersist]struct{}
active *queuedPersist
idleSince time.Time
// Deferred bookkeeping waits for the next mutation or explicit flush.
// Its age is visibility, not a promised delivery deadline.
deferred map[string]time.Time
waitingLoss, activeLoss *queuehealth.Tracker
}
func newPersistQueue() *persistQueue {
tail := make(chan struct{})
close(tail)
return &persistQueue{
tail: tail,
waiting: make(map[*queuedPersist]struct{}),
deferred: make(map[string]time.Time),
waitingLoss: queuehealth.New(0, time.Minute),
activeLoss: queuehealth.New(1, time.Minute),
}
}
func (q *persistQueue) deferWrite(id string) {
q.mu.Lock()
defer q.mu.Unlock()
if _, pending := q.deferred[id]; !pending {
q.deferred[id] = time.Now()
}
}
func (q *persistQueue) discardDeferred(id string) {
q.mu.Lock()
delete(q.deferred, id)
q.mu.Unlock()
}
func (q *persistQueue) deferredIDs() []string {
q.mu.Lock()
defer q.mu.Unlock()
ids := make([]string, 0, len(q.deferred))
for id := range q.deferred {
ids = append(ids, id)
}
return ids
}
// queuePersistLocked reserves this write's place in mutation order while
// c.mu is still held. The returned callback must run after c.mu is released.
func (c *Correlator) queuePersistLocked(snap Incident) (*queuedPersist, bool) {
persist := c.cfg.Persist
q := c.persistence
if persist == nil {
q.discardDeferred(snap.ID)
return nil, false
}
req := &queuedPersist{snap: cloneIncident(snap), persist: persist, done: make(chan struct{}), at: time.Now()}
q.mu.Lock()
defer q.mu.Unlock()
// A full immutable snapshot supersedes deferred bookkeeping atomically
// with publication, so health cannot lose the owner during the transfer.
delete(q.deferred, snap.ID)
if q.active == nil && len(q.waiting) == 0 {
q.idleSince = req.at
}
req.previous = q.tail
q.tail = req.done
q.waiting[req] = struct{}{}
return req, true
}
func (q *persistQueue) loseLocked(req *queuedPersist) {
if req.failed {
return
}
req.failed = true
tracker := q.waitingLoss
if req.running {
tracker = q.activeLoss
}
tracker.Lose(time.Now(), 1)
}
func (q *persistQueue) finish(req *queuedPersist, completed bool) {
q.mu.Lock()
defer q.mu.Unlock()
if !completed {
q.loseLocked(req)
}
if req.running {
q.active = nil
q.idleSince = time.Now()
} else {
delete(q.waiting, req)
}
// Publish completion with owner removal; a successor cannot become active
// while health still attributes the slot to its predecessor.
close(req.done)
}
func (c *Correlator) runQueuedPersist(req *queuedPersist) {
q := c.persistence
<-req.previous
completed := false
defer func() { q.finish(req, completed) }()
q.mu.Lock()
delete(q.waiting, req)
q.active = req
req.running = true
req.at = time.Now()
q.mu.Unlock()
if err := req.persist(req.snap); err != nil {
q.mu.Lock()
q.loseLocked(req)
q.mu.Unlock()
// The in-memory transition has already advanced. Count failed durable
// work before logging, which can itself wait on an output writer.
csmlog.Warn("incident persist failed", "id", req.snap.ID, "kind", string(req.snap.Kind), "status", string(req.snap.Status), "err", err)
}
completed = true
}
func (c *Correlator) runQueuedPersists(batch []*queuedPersist) {
next := 0
defer func() {
// The batch reserved contiguous ordering links under c.mu. The active
// callback has finished its cleanup before this defer, so abandoning
// its unstarted tail cannot overtake an executing predecessor.
for _, req := range batch[next:] {
c.persistence.finish(req, false)
}
}()
for next < len(batch) {
req := batch[next]
next++
c.runQueuedPersist(req)
}
}
// QueueStatuses reads only queue memory, independently of correlator state and
// persistence callbacks. Waiting writes use progress of the shared writer.
func (c *Correlator) QueueStatuses(now time.Time) map[string]queuehealth.Status {
q := c.persistence
q.mu.Lock()
defer q.mu.Unlock()
waiting := q.waitingLoss.Snapshot(now)
waiting.CapacityUnavailable = true
active := q.activeLoss.Snapshot(now)
stalled := q.active == nil && !q.idleSince.IsZero() && now.Sub(q.idleSince) >= time.Minute
if req := q.active; req != nil {
active.InFlight = 1
active.ProcessingSeconds = max(0, now.Sub(req.at).Seconds())
if active.ProcessingSeconds >= time.Minute.Seconds() {
active.Status, active.Reason = "degraded", "processing_lag"
stalled = true
}
}
for req := range q.waiting {
waiting.Depth++
waiting.LagSeconds = max(waiting.LagSeconds, now.Sub(req.at).Seconds())
if stalled {
waiting.Status, waiting.Reason = "degraded", "backlog_lag"
}
}
deferred := queuehealth.Status{Status: "ok", CapacityUnavailable: true, LagBasis: "deferred_checkpoint"}
for _, at := range q.deferred {
deferred.Depth++
deferred.LagSeconds = max(deferred.LagSeconds, now.Sub(at).Seconds())
}
return map[string]queuehealth.Status{"persist.waiting": waiting, "persist.active": active, "persist.deferred": deferred}
}
package incident
import (
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// kindRank orders Kind values from weakest to strongest so the merge
// path can upgrade an incident's Kind when a stronger pattern appears
// later, but never downgrades. Higher number = stronger / more
// operator-attention.
var kindRank = map[Kind]int{
KindWebAttack: 0,
KindMailboxBruteforce: 0,
KindWebAccountCompromise: 1,
KindMailboxTakeover: 2,
KindPostExploitProcess: 3,
KindHostIntegrityRisk: 4,
KindCredentialSpray: 3,
KindHostTakeover: 5,
}
var compoundPostExploitWebChecks = map[string]struct{}{
"webshell": {},
"webshell_realtime": {},
"webshell_content_realtime": {},
"new_webshell_file": {},
"obfuscated_php": {},
"obfuscated_php_realtime": {},
"php_shield_webshell": {},
}
var compoundPostExploitNetworkChecks = map[string]struct{}{
"c2_connection": {},
"backdoor_port": {},
"backdoor_port_outbound": {},
}
// compoundHostPrivescUID0Checks, compoundHostPrivescSUIDChecks, and
// compoundHostPrivescBadASNChecks are the three host-takeover legs the
// takeover rule correlates: a new uid-0 account, a planted suid binary, and
// an outbound connection to a bad/unexpected ASN.
var compoundHostPrivescUID0Checks = map[string]struct{}{
"uid0_account": {},
}
var compoundHostPrivescSUIDChecks = map[string]struct{}{
"suid_binary": {},
}
var compoundHostPrivescBadASNChecks = map[string]struct{}{
"bad_asn_outbound": {},
}
// allCompoundFlagsSet reports whether every compound signal is already
// recorded, so the timeline hydrate loop can stop early.
func allCompoundFlagsSet(f CompoundFlags) bool {
return f.Webshell && f.C2 && f.UID0 && f.SUID && f.BadASNOutbound
}
// hostTakeoverLegs counts how many of the three distinct host-takeover legs
// an incident has observed. Two or more legs escalate to KindHostTakeover.
func hostTakeoverLegs(f CompoundFlags) int {
n := 0
if f.UID0 {
n++
}
if f.SUID {
n++
}
if f.BadASNOutbound {
n++
}
return n
}
// maybeReclassifyKind upgrades inc.Kind in place when the new finding
// classifies as a stronger Kind, or when the incident's sticky
// CompoundFlags plus the new finding cover a compound pattern that the
// per-finding classifier cannot see. Compound rules at this time:
// webshell + outbound C2 connection -> PostExploitProcess;
// uid0_account + suid_binary -> HostTakeover. Idempotent: calling with
// weaker findings is a no-op.
//
// CompoundFlags are mutated here so callers do not need a separate
// pass; they survive timeline trimming so an early webshell still
// arms the rule when a much later C2 finding arrives.
func maybeReclassifyKind(inc *Incident, f alert.Finding) {
if inc == nil {
return
}
// Hydrate sticky flags from the current timeline so incidents
// restored from bbolt (predating sticky flags) or built directly
// in tests still arm the compound rule. Timeline scan is bounded
// by maxIncidentTimeline so the cost is constant.
hydrateCompoundFlagsFromTimeline(&inc.CompoundFlags, inc.Timeline)
updateCompoundFlags(&inc.CompoundFlags, f.Check)
if newKind := ClassifyKind(f); kindRank[newKind] > kindRank[inc.Kind] && canPromoteKind(inc, f) {
inc.Kind = newKind
}
if kindRank[KindPostExploitProcess] > kindRank[inc.Kind] && inc.CompoundFlags.Webshell && inc.CompoundFlags.C2 && canPromoteKind(inc, f) {
inc.Kind = KindPostExploitProcess
}
// Host takeover: any two of the three distinct legs (new uid-0 account,
// planted suid binary, outbound connection to a bad ASN) on the same
// host inside the window.
if kindRank[KindHostTakeover] > kindRank[inc.Kind] && hostTakeoverLegs(inc.CompoundFlags) >= 2 && canPromoteKind(inc, f) {
inc.Kind = KindHostTakeover
}
}
func canPromoteKind(inc *Incident, f alert.Finding) bool {
if !remoteIPOnlyIncidentKey(inc.CorrelationKey) {
return true
}
k := KeyFor(f)
return !remoteIPOnlyKey(k)
}
func remoteIPOnlyIncidentKey(k *Key) bool {
if k == nil {
return false
}
return remoteIPOnlyKey(*k)
}
func remoteIPOnlyKey(k Key) bool {
return k.RemoteIP != "" && k.Host == "" && k.Account == "" && k.Domain == "" &&
k.Mailbox == "" && k.UID == 0 && k.PID == 0
}
// hydrateCompoundFlagsFromTimeline OR-merges timeline-derived signals
// into flags. Used as a one-shot migration for legacy/persisted
// incidents that have webshell or C2 events in their timeline but no
// CompoundFlags yet. Trimmed timelines may miss events, but anything
// still present remains a valid signal.
func hydrateCompoundFlagsFromTimeline(flags *CompoundFlags, events []IncidentEvent) {
if flags == nil || allCompoundFlagsSet(*flags) {
return
}
for _, ev := range events {
updateCompoundFlags(flags, ev.Check)
if allCompoundFlagsSet(*flags) {
return
}
}
}
// updateCompoundFlags sets sticky compound flags based on a Finding's
// check name. Once true a flag stays true so reclassify decisions are not
// silently disarmed by later trimming.
func updateCompoundFlags(flags *CompoundFlags, check string) {
if flags == nil {
return
}
check = strings.ToLower(strings.TrimSpace(check))
if _, ok := compoundPostExploitWebChecks[check]; ok {
flags.Webshell = true
}
if _, ok := compoundPostExploitNetworkChecks[check]; ok {
flags.C2 = true
}
if _, ok := compoundHostPrivescUID0Checks[check]; ok {
flags.UID0 = true
}
if _, ok := compoundHostPrivescSUIDChecks[check]; ok {
flags.SUID = true
}
if _, ok := compoundHostPrivescBadASNChecks[check]; ok {
flags.BadASNOutbound = true
}
}
package incident
import (
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// SpraySuppressionConfig is the operator-tunable knob set for the
// credential-spray detector. Defaults are conservative: enabled=false and
// dry_run=true so the detector ships dark, increments counters, and only
// changes incident routing once an operator opts in.
type SpraySuppressionConfig struct {
Enabled bool
DryRun bool
DistinctMailboxes int
SeverityEscalateAt int
PerCheck map[string]bool
MaxTrackedIPs int
// BlockAtSeverity gates the firewall hand-off. Values:
// "" / unset - detection-only (legacy behavior).
// "high" - block on incident open (DistinctMailboxes trip).
// "critical" - block on severity escalation (SeverityEscalateAt).
// Comparison is case-insensitive. Any other value is ignored so an
// operator typo cannot accidentally engage blocking.
BlockAtSeverity string
// BlockExpiry is the operator's configured auto-response block duration,
// used as the first rung of the escalation ladder. Zero falls back to 24h.
BlockExpiry time.Duration
}
// IsZero reports whether the config is unset; the correlator treats a
// zero value as "spray detection disabled" without touching defaults.
func (c SpraySuppressionConfig) IsZero() bool {
return !c.Enabled && !c.DryRun && c.DistinctMailboxes == 0 && c.SeverityEscalateAt == 0 && len(c.PerCheck) == 0 && c.MaxTrackedIPs == 0 && c.BlockAtSeverity == ""
}
// sprayDecision is the verdict the detector returns to OnFinding for a
// candidate spray-class finding.
type sprayDecision int
const (
// sprayDecisionNone means the finding is not spray traffic; the
// caller should run the normal per-mailbox correlation path.
sprayDecisionNone sprayDecision = iota
// sprayDecisionOpen means this finding tipped the IP over the
// distinct-mailbox threshold; the caller must open a new
// credential_spray incident keyed on RemoteIP and report the id back
// to the detector via BindIncident.
sprayDecisionOpen
// sprayDecisionSuppress means an active spray incident already
// exists for the IP and the finding should be attached to it
// instead of opening a per-mailbox incident.
sprayDecisionSuppress
)
// ipSprayState tracks the distinct-mailbox set hit by a single source IP
// inside the merge window, plus the bound spray incident id once the
// threshold trips.
type ipSprayState struct {
mailboxes map[string]struct{}
firstSeen time.Time
lastSeen time.Time
// incident is the bound credential_spray incident id once the
// threshold has tripped. Empty until then.
incident string
}
// sprayDetector keeps a per-IP sliding window of distinct mailboxes hit
// by spray-class checks and decides whether a finding should open a new
// credential_spray super-incident, attach to an existing one, or fall
// through to the normal per-mailbox correlator path.
//
// Concurrency: the correlator's mutex serializes all calls into the
// detector. The detector itself does not take a separate lock so the
// "single mutex protects all correlator state" invariant holds.
type sprayDetector struct {
cfg SpraySuppressionConfig
window time.Duration
now func() time.Time
isWhitelisted func(string) bool
// perIP is the live state map. Bounded by cfg.MaxTrackedIPs; entries
// that fall outside the window during Decide are pruned in place,
// and once the map exceeds the cap the oldest-by-lastSeen entry is
// evicted before insert.
perIP map[string]*ipSprayState
}
// newSprayDetector returns a detector wired to the supplied config.
// Returns nil when cfg is zero-valued so the correlator can no-op the
// fast path on hosts that have not opted in.
func newSprayDetector(cfg SpraySuppressionConfig, window time.Duration, now func() time.Time, isWhitelisted func(string) bool) *sprayDetector {
if cfg.IsZero() {
return nil
}
if cfg.MaxTrackedIPs <= 0 {
cfg.MaxTrackedIPs = 10000
}
if cfg.DistinctMailboxes <= 0 {
cfg.DistinctMailboxes = 10
}
if cfg.SeverityEscalateAt <= cfg.DistinctMailboxes {
// Default to 5x threshold so a CRITICAL bump only fires on
// genuinely sustained sprays, not on the immediate trip.
cfg.SeverityEscalateAt = cfg.DistinctMailboxes * 5
}
if isWhitelisted == nil {
isWhitelisted = func(string) bool { return false }
}
if now == nil {
now = time.Now
}
return &sprayDetector{
cfg: cfg,
window: window,
now: now,
isWhitelisted: isWhitelisted,
perIP: make(map[string]*ipSprayState),
}
}
// Decide consumes a spray-candidate finding and returns the decision
// the correlator should apply. The detector mutates internal state
// (records the mailbox, tracks the lastSeen timestamp) regardless of
// the dry_run flag so counters and audit logs reflect what the live
// path would have done; only the returned decision is gated by
// dry_run. Caller must hold the correlator mutex.
//
// hitCount is the number of distinct mailboxes the IP has hit inside
// the window after this finding is recorded. The caller uses it to
// decide severity for sprayDecisionOpen and to escalate severity on
// merge-into-existing-spray.
func (d *sprayDetector) Decide(f alert.Finding) (decision sprayDecision, hitCount int) {
if d == nil || !d.cfg.Enabled && !d.cfg.DryRun {
return sprayDecisionNone, 0
}
if f.SourceIP == "" {
return sprayDecisionNone, 0
}
if !d.cfg.PerCheck[f.Check] {
return sprayDecisionNone, 0
}
if d.isWhitelisted(f.SourceIP) {
return sprayDecisionNone, 0
}
targets := sprayTargets(f)
if len(targets) == 0 {
return sprayDecisionNone, 0
}
now := d.now()
state, ok := d.perIP[f.SourceIP]
if ok {
// Window expiration: a state whose lastSeen fell outside the
// window is stale; reset before recording the new hit so a
// fresh attack does not inherit cold mailbox counts. Bound
// entries stay until the correlator closes or unbinds their
// incident; otherwise a quiet but still-open super-incident can
// lose suppression and duplicate on the next finding.
if state.incident == "" && now.Sub(state.lastSeen) > d.window {
state = nil
delete(d.perIP, f.SourceIP)
}
}
if state == nil {
// Eviction: keep the live set bounded. Drop the oldest-by-lastSeen
// entry before inserting if we are at the cap. O(N) scan; cheap
// at the configured cap (10k) and only runs at insert time.
if len(d.perIP) >= d.cfg.MaxTrackedIPs {
d.evictOldestLocked()
}
state = &ipSprayState{
mailboxes: make(map[string]struct{}),
firstSeen: now,
}
d.perIP[f.SourceIP] = state
}
for _, target := range targets {
state.mailboxes[target] = struct{}{}
}
state.lastSeen = now
hitCount = len(state.mailboxes)
// Bound to existing spray incident? Continue suppressing.
if state.incident != "" {
if d.cfg.DryRun {
return sprayDecisionNone, hitCount
}
return sprayDecisionSuppress, hitCount
}
if hitCount < d.cfg.DistinctMailboxes {
return sprayDecisionNone, hitCount
}
// Threshold tripped. Live mode opens a new spray incident; dry_run
// only counts the decision so the operator can observe the workload
// without changing routing.
if d.cfg.DryRun {
return sprayDecisionNone, hitCount
}
return sprayDecisionOpen, hitCount
}
// sprayTargets returns the identity dimension used for the distinct-target set:
// mailbox, tenant id, cPanel user, aggregate auth targets, then message text.
func sprayTargets(f alert.Finding) []string {
target, _ := canonicalizeMailboxDomain(f.Mailbox, f.Domain)
if target != "" {
return []string{target}
}
if f.TenantID != "" {
return []string{f.TenantID}
}
if f.CPUser != "" {
return []string{f.CPUser}
}
if targets := cleanSprayTargets(f.SprayTargets); len(targets) > 0 {
return targets
}
message := strings.TrimSpace(f.Message)
if message != "" {
return []string{message}
}
return nil
}
func cleanSprayTargets(targets []string) []string {
if len(targets) == 0 {
return nil
}
seen := make(map[string]struct{}, len(targets))
out := make([]string, 0, len(targets))
for _, target := range targets {
target = strings.TrimSpace(target)
if target == "" {
continue
}
if _, ok := seen[target]; ok {
continue
}
seen[target] = struct{}{}
out = append(out, target)
}
return out
}
// BindIncident records the spray incident id the correlator just
// created in response to sprayDecisionOpen. Subsequent findings from
// the same IP return sprayDecisionSuppress while the incident stays
// bound.
// Caller holds the correlator mutex.
func (d *sprayDetector) BindIncident(ip, id string) {
if d == nil {
return
}
if state, ok := d.perIP[ip]; ok {
state.incident = id
}
}
// Rehydrate seeds the perIP map at daemon startup so an open
// credential_spray incident restored from bbolt continues to suppress
// new per-mailbox fan-out instead of allowing a duplicate super-incident
// to open. The seeded state carries no per-mailbox set: the operator
// already saw the trip on the open incident, and the suppress path only
// reads state.incident. lastSeen is set so metrics and future unbound
// expiry have a stable anchor if the incident later closes.
// Caller holds the correlator mutex.
func (d *sprayDetector) Rehydrate(ip, id string, lastSeen time.Time) {
if d == nil || ip == "" || id == "" {
return
}
d.perIP[ip] = &ipSprayState{
mailboxes: make(map[string]struct{}),
firstSeen: lastSeen,
lastSeen: lastSeen,
incident: id,
}
}
// IncidentForIP returns the bound spray incident id for ip, or "" if no
// spray is currently bound. Used by Decide's suppress path to tell the
// caller which incident to merge the finding into.
func (d *sprayDetector) IncidentForIP(ip string) string {
if d == nil {
return ""
}
if state, ok := d.perIP[ip]; ok {
return state.incident
}
return ""
}
// UnbindIncident drops detector state for every IP currently bound to
// id, so a finding from one of those IPs cannot reach the closed or
// missing incident through the suppress merge path. Dropping the state
// also resets the distinct-mailbox threshold for future activity after
// an operator closes or dismisses the incident. Caller holds the
// correlator mutex. Linear scan over perIP, same shape as the byKey
// unbind in correlator.unbindLocked: the active set is bounded by
// MaxTrackedIPs.
func (d *sprayDetector) UnbindIncident(id string) {
if d == nil || id == "" {
return
}
for ip, state := range d.perIP {
if state.incident == id {
delete(d.perIP, ip)
}
}
}
// UnbindIncidents drops detector state for every IP bound to one of the
// supplied incident ids. Caller holds the correlator mutex.
func (d *sprayDetector) UnbindIncidents(ids []string) {
if d == nil || len(ids) == 0 {
return
}
if len(ids) == 1 {
d.UnbindIncident(ids[0])
return
}
idSet := make(map[string]struct{}, len(ids))
for _, id := range ids {
if id != "" {
idSet[id] = struct{}{}
}
}
if len(idSet) == 0 {
return
}
for ip, state := range d.perIP {
if _, ok := idSet[state.incident]; ok {
delete(d.perIP, ip)
}
}
}
// PruneStale clears entries whose lastSeen is older than the window.
// Called by the daemon retention loop alongside PruneStalePending so
// the detector does not grow without bound on hosts with churning
// attacker IPs. Entries still bound to an open spray incident are
// kept regardless of age -- evicting them lets a new spray finding
// open a duplicate per-mailbox incident while the super-incident is
// still active in the correlator.
func (d *sprayDetector) PruneStale(now time.Time) int {
if d == nil {
return 0
}
pruned := 0
for ip, state := range d.perIP {
if state.incident != "" {
continue
}
if now.Sub(state.lastSeen) > d.window {
delete(d.perIP, ip)
pruned++
}
}
return pruned
}
// TrackedIPs returns the count of source IPs currently held in the
// detector. Surfaced via the csm_credential_spray_tracked_ips gauge.
func (d *sprayDetector) TrackedIPs() int {
if d == nil {
return 0
}
return len(d.perIP)
}
// evictOldestLocked drops the perIP entry with the smallest lastSeen.
// Caller holds the correlator mutex. Linear scan over MaxTrackedIPs.
func (d *sprayDetector) evictOldestLocked() {
var oldestIP string
var oldestAt time.Time
first := true
for ip, state := range d.perIP {
if first || state.lastSeen.Before(oldestAt) {
oldestIP = ip
oldestAt = state.lastSeen
first = false
}
}
if oldestIP != "" {
delete(d.perIP, oldestIP)
}
}
// Counters for the credential-spray decisions live on the Correlator's
// existing `counters` struct (see correlator.go). Keeping them there
// means there is exactly one place that owns counter mutations, which
// keeps the metrics story coherent and avoids surprising operators
// who already grep for `csm_incidents_*`.
// Package incident groups related security findings into a single
// "story" with a timeline. Original findings are not mutated or
// suppressed; the Incident is layered on top so operators read one
// escalating object instead of stitching findings together by hand.
package incident
import (
"encoding/json"
"fmt"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// Status is the lifecycle position of an incident.
type Status string
const (
StatusOpen Status = "open"
StatusContained Status = "contained"
StatusResolved Status = "resolved"
StatusDismissed Status = "dismissed"
)
// ClosedRetention is how long resolved and dismissed incidents are kept after
// their last update. Operator applies to incidents an operator closed and to
// rows closed before close attribution existed; Auto applies to incidents the
// daemon closed on its own, which carry no operator decision worth keeping as
// long.
type ClosedRetention struct {
Operator time.Duration
Auto time.Duration
}
// Expired reports whether a closed incident has outlived its retention.
// Active incidents never expire here.
func (r ClosedRetention) Expired(inc Incident, now time.Time) bool {
if inc.Status != StatusResolved && inc.Status != StatusDismissed {
return false
}
keep := r.Operator
updated := inc.UpdatedAt
if strings.HasPrefix(inc.ClosedBy, closedByAutoPrefix) {
keep = r.Auto
// Older operator writers left automatic attribution on closed
// records. Preserve those decisions during the first upgraded sweep;
// a later automatic closure starts a new retention episode.
for i := len(inc.Actions) - 1; i >= 0; i-- {
action := inc.Actions[i]
if action.Action == "incident_auto_closed" {
break
}
if (action.Action == "incident_status_changed" || action.Action == "operator_block") && !action.Time.Before(inc.ClosedAt) {
keep = r.Operator
if action.Time.After(updated) {
updated = action.Time
}
}
}
}
return updated.Before(now.Add(-keep))
}
// closedByAutoPrefix marks ClosedBy values the daemon writes when it closes
// an incident itself ("auto:stale", "auto:age_cap", "auto:active_cap").
const closedByAutoPrefix = "auto:"
// Kind is the high-level taxonomy a correlator assigns at create time.
// Stable strings; downstream tooling pins on these.
type Kind string
const (
KindWebAccountCompromise Kind = "web_account_compromise"
// KindWebAttack is an inbound web attack: a WAF hit, scanner probe,
// or login brute-force from a remote source, plus remote-IP
// reputation/threat-score signals. When such a finding names a victim
// domain or account, that is the attack target, not evidence the
// account is compromised, so these correlate on the attacker source IP
// and get a short attacker-grade retention. Keeping them out of
// web_account_compromise stops defended inbound traffic from inflating
// the account-compromise count and the 7-day review window. Genuine
// compromise is recognised by on-disk and behavioural signals
// (webshell, suspicious PHP, post-exploit process), not by inbound hits.
KindWebAttack Kind = "web_attack"
// KindMailboxBruteforce is a failed-authentication brute-force attempt
// or pre-auth mail probe against one or more mailboxes from a remote
// source. A failed login is an attack attempt, not a takeover, so it
// correlates on the attacker source with short attacker-grade retention.
// Post-authentication abuse (outbound spam, cloud relay,
// compromised-account, suspicious geo) stays in mailbox_takeover.
KindMailboxBruteforce Kind = "mailbox_bruteforce"
KindMailboxTakeover Kind = "mailbox_takeover"
KindPostExploitProcess Kind = "post_exploit_process"
KindHostIntegrityRisk Kind = "host_integrity_risk"
// KindCredentialSpray collapses a single source IP that is brute-forcing
// many distinct mailboxes/accounts inside the merge window into one
// super-incident keyed on the source IP. Prevents the per-mailbox fan-out
// that turns one attacker into thousands of mailbox_bruteforce incidents.
KindCredentialSpray Kind = "credential_spray" // #nosec G101 -- taxonomy label, not a secret
// KindHostTakeover is the compound escalation when more than one
// host-privilege-escalation leg (a new uid-0 account, a planted suid
// binary, or bad-ASN outbound connection) is seen for the same host
// inside the merge window. It ranks above KindHostIntegrityRisk so a
// confirmed multi-leg takeover stands out from a single host-integrity
// finding.
KindHostTakeover Kind = "host_takeover"
)
// Incident is the wire shape every consumer (API, control socket,
// audit propagation) sees. omitempty fields are absent from JSON when
// zero so consumers ignore optional context cleanly.
type Incident struct {
ID string `json:"id"`
Kind Kind `json:"kind"`
Status Status `json:"status"`
Severity alert.Severity `json:"severity"`
Account string `json:"account,omitempty"`
Domain string `json:"domain,omitempty"`
Mailbox string `json:"mailbox,omitempty"`
CorrelationKey *Key `json:"correlation_key,omitempty"`
Summary string `json:"summary,omitempty"`
Confidence int `json:"confidence,omitempty"`
Findings []string `json:"findings,omitempty"`
Timeline []IncidentEvent `json:"timeline,omitempty"`
Actions []IncidentAction `json:"actions,omitempty"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
// ClosedAt records the latest closure or operator decision on a closed
// incident. Reopening clears it.
ClosedAt time.Time `json:"closed_at,omitzero"`
// ClosedBy is "operator" for manual decisions on closed incidents and
// "auto:<reason>" for daemon closures. Empty for active or legacy rows.
ClosedBy string `json:"closed_by,omitempty"`
// CompoundFlags carries sticky bits used by the timeline-aware
// reclassifier. Once set, they survive timeline trimming so an
// early webshell or C2 signal still drives the compound rule when
// the matching counterpart arrives much later.
CompoundFlags CompoundFlags `json:"compound_flags,omitzero"`
// AutoBlock records what the automatic firewall hand-off already did for
// this incident. The block it applies expires; without this the marker
// saying "already blocked" did not, so an attack that outlasted its
// expiry was never blocked again.
AutoBlock AutoBlockState `json:"auto_block,omitzero"`
}
// AutoBlockState is the escalation ladder's memory for one incident. Count is
// how many blocks the hand-off has requested, ExpiresAt when the most recent
// one lapses, and a zero ExpiresAt with a nonzero Count means that block is
// permanent and nothing re-requests it. Reset when the incident leaves an
// active status, so a later recurrence starts from the bottom of the ladder.
type AutoBlockState struct {
Count int `json:"count,omitempty"`
ExpiresAt time.Time `json:"expires_at,omitzero"`
LastAt time.Time `json:"last_at,omitzero"`
}
// lapsed reports whether the hand-off may request another block: never
// blocked, or the last block has expired. A permanent block never lapses.
func (s AutoBlockState) lapsed(now time.Time) bool {
if s.Count == 0 {
return true
}
if s.ExpiresAt.IsZero() {
return false
}
return !now.Before(s.ExpiresAt)
}
// blockTTLForAttempt escalates the hand-off: the first block uses the
// operator's configured expiry, the second a week, and any later one is
// permanent (the firewall reads a zero timeout as permanent). An attacker who
// outlasts one expiry pays more each time, while a single false positive
// still ages out on its own.
func blockTTLForAttempt(attempt int, configured time.Duration) time.Duration {
switch {
case attempt <= 1:
if configured <= 0 {
return 24 * time.Hour
}
return configured
case attempt == 2:
return 7 * 24 * time.Hour
default:
return 0
}
}
// CompoundFlags records the union of compound-pattern signals an
// Incident has ever observed. Fields are sticky once true; they are
// not derived from the (possibly trimmed) timeline so reclassify is
// not silently disarmed by head+tail eviction.
type CompoundFlags struct {
Webshell bool `json:"webshell,omitempty"`
C2 bool `json:"c2,omitempty"`
// UID0, SUID, and BadASNOutbound record the three host-takeover legs:
// a new uid-0 account, a planted suid binary, and an outbound connection
// to a bad/unexpected ASN. When any two are set on one incident the
// reclassifier escalates to KindHostTakeover.
UID0 bool `json:"uid0,omitempty"`
SUID bool `json:"suid,omitempty"`
BadASNOutbound bool `json:"bad_asn_outbound,omitempty"`
}
// MarshalJSON renders Severity as its uppercase string form
// ("HIGH", "CRITICAL", "WARNING") instead of the underlying int.
// alert.Severity is an int enum, so default marshaling would emit
// numbers; consumers (web UI, control socket, audit propagation)
// expect the same human-readable token already produced by
// audit_sink and webhook dispatch.
func (i Incident) MarshalJSON() ([]byte, error) {
type wireIncident Incident
return json.Marshal(struct {
wireIncident
Severity string `json:"severity"`
}{
wireIncident: wireIncident(i),
Severity: i.Severity.String(),
})
}
// UnmarshalJSON decodes the wire shape produced by MarshalJSON. Severity
// is read from its string form ("WARNING"/"HIGH"/"CRITICAL") and converted
// back to alert.Severity. Unknown strings return an error so SIEM-side
// schema drift is loud, not silent.
func (i *Incident) UnmarshalJSON(data []byte) error {
type wireIncident Incident
aux := struct {
*wireIncident
Severity string `json:"severity"`
}{wireIncident: (*wireIncident)(i)}
if err := json.Unmarshal(data, &aux); err != nil {
return err
}
switch aux.Severity {
case "":
// allow the zero-severity case for partial decodes (tests, partial
// JSON snippets in the API). Severity stays at zero value (Warning).
case "WARNING":
i.Severity = alert.Warning
case "HIGH":
i.Severity = alert.High
case "CRITICAL":
i.Severity = alert.Critical
default:
return fmt.Errorf("incident: unknown severity %q", aux.Severity)
}
return nil
}
// IncidentEvent is one entry in an incident's timeline. Built from a
// Finding when it joins the incident; carries enough context to
// render the timeline without re-reading the original record.
type IncidentEvent struct {
Time time.Time `json:"time"`
Kind string `json:"kind"`
Check string `json:"check,omitempty"`
Severity string `json:"severity,omitempty"`
Message string `json:"message"`
FindingID string `json:"finding_id,omitempty"`
PID int `json:"pid,omitempty"`
UID int `json:"uid,omitempty"`
Process string `json:"process,omitempty"`
Path string `json:"path,omitempty"`
RemoteIP string `json:"remote_ip,omitempty"`
}
// IncidentAction is an automated or operator action that touched the
// incident. Appended to the timeline; surfaced separately so dashboards
// can filter by what the system did vs what it observed.
type IncidentAction struct {
Time time.Time `json:"time"`
Action string `json:"action"`
Result string `json:"result"`
Details string `json:"details,omitempty"`
}
package webserver
import (
"context"
"fmt"
"os/exec"
"time"
"github.com/pidginhost/csm/internal/platform"
)
// apacheHandler covers cPanel + plain Apache. cPanel ships `apachectl`
// pointing at the EasyApache binary; plain Apache on Debian/Ubuntu has
// `apache2ctl`. The selector picks the right one at construction.
type apacheHandler struct {
snippetPath string
ctlBinary string // "apachectl" or "apache2ctl"
reloadAction []string
cmdRunner cmdRunner
}
func newApacheHandler(info platform.Info, r cmdRunner) *apacheHandler {
h := &apacheHandler{cmdRunner: r}
switch {
case info.IsCPanel():
// cPanel always runs EasyApache, conf.d is the canonical drop-in.
h.snippetPath = "/etc/apache2/conf.d/csm-challenge.conf"
h.ctlBinary = "apachectl"
h.reloadAction = []string{"apachectl", "graceful"}
case info.IsDebianFamily():
// Debian/Ubuntu use apache2 + conf-enabled.
h.snippetPath = "/etc/apache2/conf-enabled/csm-challenge.conf"
h.ctlBinary = "apache2ctl"
h.reloadAction = []string{"systemctl", "reload", "apache2"}
default:
// RHEL family without cPanel: httpd + /etc/httpd/conf.d.
h.snippetPath = "/etc/httpd/conf.d/csm-challenge.conf"
h.ctlBinary = "apachectl"
h.reloadAction = []string{"systemctl", "reload", "httpd"}
}
return h
}
func (h *apacheHandler) Kind() string { return "apache" }
func (h *apacheHandler) SnippetPath() string { return h.snippetPath }
func (h *apacheHandler) Template() string { return apacheTemplate }
func (h *apacheHandler) PostInstallInstructions() string { return "" }
func (h *apacheHandler) Validate() error {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
out, err := h.cmdRunner.Run(ctx, h.ctlBinary, "configtest")
if err != nil {
return fmt.Errorf("apache configtest failed: %v\n%s", err, out)
}
return nil
}
func (h *apacheHandler) Reload() error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
out, err := h.cmdRunner.Run(ctx, h.reloadAction[0], h.reloadAction[1:]...)
if err != nil {
return fmt.Errorf("apache reload failed: %v\n%s", err, out)
}
return nil
}
// cmdRunner is the injection seam tests use to mock exec. The real
// implementation in realCmdRunner just shells out via os/exec.
type cmdRunner interface {
Run(ctx context.Context, name string, args ...string) ([]byte, error)
}
type realCmdRunner struct{}
func (realCmdRunner) Run(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- name + args come from the installer's static handler
// definitions (apachectl/apache2ctl/nginx/lswsctrl + verbs). No
// user-controlled strings reach this path.
return exec.CommandContext(ctx, name, args...).CombinedOutput()
}
// Package webserver auto-installs the CSM challenge webserver glue
// (Apache / LSWS / Nginx) with a write-validate-reload-or-revert flow.
// The operator runs `csm webserver-integration {install|upgrade|...}`;
// the package picks the right handler for the host and never reloads
// the webserver with a snippet that does not pass configtest.
package webserver
import (
"bytes"
"errors"
"fmt"
"io"
"net"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"text/template"
"github.com/pidginhost/csm/internal/challenge"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
)
// templateHeader is prepended to every rendered snippet so the
// installer can read the version back without parsing the body. The
// header is a single comment line whose value uniquely identifies
// "CSM owns this file". Manual edits that wipe or change the header
// trip ErrManualEdits at upgrade time and the file is left untouched.
const templateHeaderPrefix = "# csm-managed-version: "
// RenderConfig contains the daemon settings that have to be baked
// into webserver snippets. Values come from csm.yaml at install time
// where they are operator-configurable.
type RenderConfig struct {
ChallengeMapPath string
ChallengeNginxMapPath string
ChallengeListenAddr string
// ChallengePublicURL is the fully-qualified redirect target the
// webserver snippet emits. Required on install; empty value blocks
// the installer with a clear error. Operator typically points it at
// the host's TLS-valid control-panel domain on a CSM-owned port,
// e.g. https://server.example.com:8439/challenge.
ChallengePublicURL string
}
type templateData struct {
ChallengeMapPath string
ChallengeNginxMapPath string
ChallengePublicURL string
}
// ErrMissingPublicURL is returned by Install when challenge.public_url
// is empty. The webserver snippet has no fallback redirect target so
// the integration cannot function without it.
var ErrMissingPublicURL = errors.New("webserver integration: challenge.public_url is not set")
// ErrInvalidPublicURL is returned when challenge.public_url cannot be
// used as the direct browser redirect target.
var ErrInvalidPublicURL = errors.New("webserver integration: challenge.public_url must be an absolute http(s) URL ending in /challenge")
// ErrLoopbackPublicURL is returned when the installed snippets would
// redirect browsers to a listener that only accepts loopback traffic.
var ErrLoopbackPublicURL = errors.New("webserver integration: challenge.listen_addr must be non-loopback when challenge.public_url redirects are enabled")
// Result is the structured outcome of an installer run. JSON-friendly
// shape so the CLI can render either human text or `--json` output.
type Result struct {
Action string `json:"action"` // "install" | "upgrade" | "remove" | "status" | "validate"
Status string `json:"status"` // "ok" | "no-op" | "skipped" | "fail"
Webserver string `json:"webserver"` // detected handler kind, "" if none
SnippetPath string `json:"snippet_path"` // "" if no handler
OnDiskVer int `json:"on_disk_version,omitempty"`
ShippedVer int `json:"shipped_version,omitempty"`
Message string `json:"message,omitempty"`
// FollowUp is the operator-facing setup the integration could not
// automate on this stack. Empty when install is fully automatic.
// The CLI prints it after the result block on successful install.
FollowUp string `json:"follow_up,omitempty"`
}
// Installer drives the install / upgrade / status / remove flow. All
// I/O goes through injected hooks so unit tests can run on darwin or
// against a temp tree without touching real webserver paths.
type Installer struct {
Handler Handler
Config RenderConfig
MkdirAll func(path string, mode os.FileMode) error
WriteAt func(path string, data []byte, mode os.FileMode) error
ReadAt func(path string) ([]byte, error)
StatAt func(path string) (os.FileInfo, error)
RemoveAt func(path string) error
Stderr io.Writer
}
// New returns an Installer wired for live operation: real filesystem
// reads, atomic writes, real exec runner. The handler is auto-selected
// from platform.Detect(); pass info to override for tests.
func New(info platform.Info, cfg *config.Config) (*Installer, error) {
h, err := pickHandler(info, realCmdRunner{})
if err != nil {
return nil, err
}
return &Installer{
Handler: h,
Config: renderConfigFrom(cfg),
MkdirAll: os.MkdirAll,
WriteAt: atomicWrite,
ReadAt: os.ReadFile,
StatAt: os.Stat,
RemoveAt: os.Remove,
Stderr: os.Stderr,
}, nil
}
// Install writes the snippet for the first time (or overwrites a
// stale one) with the safe rollback flow:
//
// 1. Stash existing bytes (or note absence).
// 2. Write new snippet atomically.
// 3. Validate via the webserver's own configtest.
// 4. On pass: reload + done.
// 5. On fail: restore previous bytes + return error.
//
// Reload failure after a passing configtest is restored the same way,
// then a recovery reload is attempted so the host returns to the
// last-known-good state.
func (i *Installer) Install() (Result, error) {
res := Result{
Action: "install",
Webserver: i.Handler.Kind(),
SnippetPath: i.Handler.SnippetPath(),
ShippedVer: TemplateVersion,
}
if err := validateChallengePublicURL(i.Config); err != nil {
res.Status = "fail"
res.Message = err.Error()
return res, err
}
prevBytes, prevExists, prevVer, err := i.readSnippet()
if err != nil && !errors.Is(err, os.ErrNotExist) {
res.Status = "fail"
res.Message = err.Error()
return res, err
}
res.OnDiskVer = prevVer
if prevExists && prevVer == 0 {
res.Status = "fail"
res.Message = ErrManualEdits.Error() + ": " + i.Handler.SnippetPath()
return res, ErrManualEdits
}
if mapErr := i.ensureChallengeMapFiles(); mapErr != nil {
res.Status = "fail"
res.Message = "runtime files: " + mapErr.Error()
return res, mapErr
}
rendered, rerr := i.renderTemplate()
if rerr != nil {
res.Status = "fail"
res.Message = "render: " + rerr.Error()
return res, rerr
}
if prevExists && bytes.Equal(prevBytes, rendered) {
res.Status = "no-op"
res.Message = "snippet already current"
return res, nil
}
if err := i.WriteAt(i.Handler.SnippetPath(), rendered, 0o644); err != nil {
res.Status = "fail"
res.Message = "write: " + err.Error()
return res, err
}
if verr := i.Handler.Validate(); verr != nil {
i.restore(prevBytes, prevExists)
res.Status = "fail"
res.Message = "configtest: " + verr.Error()
return res, verr
}
if rerr := i.Handler.Reload(); rerr != nil {
i.restore(prevBytes, prevExists)
// Best-effort recovery reload. Even if it fails, the file is
// already back to the last-good content.
_ = i.Handler.Reload()
res.Status = "fail"
res.Message = "reload: " + rerr.Error() + " (rolled back)"
return res, rerr
}
res.Status = "ok"
if prevExists {
res.Message = fmt.Sprintf("snippet upgraded v%d -> v%d", prevVer, TemplateVersion)
} else {
res.Message = fmt.Sprintf("snippet installed (v%d)", TemplateVersion)
}
res.FollowUp = i.Handler.PostInstallInstructions()
return res, nil
}
// Upgrade is an alias for Install with a more honest CLI verb. The
// underlying flow is the same: idempotent install + version compare.
func (i *Installer) Upgrade() (Result, error) {
res, err := i.Install()
res.Action = "upgrade"
return res, err
}
// Status returns the current integration state without writing
// anything. Used by post-upgrade hooks and operator-facing diagnostic
// commands to detect drift.
func (i *Installer) Status() (Result, error) {
res := Result{
Action: "status",
Webserver: i.Handler.Kind(),
SnippetPath: i.Handler.SnippetPath(),
ShippedVer: TemplateVersion,
}
_, exists, ver, err := i.readSnippet()
if err != nil && !errors.Is(err, os.ErrNotExist) {
res.Status = "fail"
res.Message = err.Error()
return res, err
}
res.OnDiskVer = ver
res.Status, res.Message = classifyStatus(exists, ver, TemplateVersion)
return res, nil
}
// classifyStatus is the pure-logic version compare extracted so the
// stale / modified / ok branches can be unit-tested without depending
// on the current TemplateVersion constant.
func classifyStatus(exists bool, onDisk, shipped int) (status, message string) {
switch {
case !exists:
return "missing", "no snippet installed; run `csm webserver-integration install`"
case onDisk == 0:
return "modified", ErrManualEdits.Error()
case onDisk < shipped:
return "stale", fmt.Sprintf("on-disk v%d < shipped v%d; run `csm webserver-integration upgrade`", onDisk, shipped)
default:
return "ok", fmt.Sprintf("snippet at v%d", onDisk)
}
}
// Remove deletes the snippet, runs configtest to confirm the webserver
// is happy without it, and reloads. Mirrors the install rollback
// discipline: if removing the file makes configtest fail, restore the
// original and exit non-zero.
func (i *Installer) Remove() (Result, error) {
res := Result{
Action: "remove",
Webserver: i.Handler.Kind(),
SnippetPath: i.Handler.SnippetPath(),
ShippedVer: TemplateVersion,
}
prevBytes, prevExists, prevVer, err := i.readSnippet()
if err != nil && !errors.Is(err, os.ErrNotExist) {
res.Status = "fail"
res.Message = err.Error()
return res, err
}
res.OnDiskVer = prevVer
if !prevExists {
res.Status = "no-op"
res.Message = "snippet not present"
return res, nil
}
if prevVer == 0 {
res.Status = "fail"
res.Message = ErrManualEdits.Error() + ": refusing to delete an operator-edited file"
return res, ErrManualEdits
}
if err := i.RemoveAt(i.Handler.SnippetPath()); err != nil {
res.Status = "fail"
res.Message = "delete: " + err.Error()
return res, err
}
if verr := i.Handler.Validate(); verr != nil {
i.restore(prevBytes, prevExists)
res.Status = "fail"
res.Message = "configtest after remove: " + verr.Error()
return res, verr
}
if rerr := i.Handler.Reload(); rerr != nil {
i.restore(prevBytes, prevExists)
_ = i.Handler.Reload()
res.Status = "fail"
res.Message = "reload: " + rerr.Error() + " (rolled back)"
return res, rerr
}
res.Status = "ok"
res.Message = "snippet removed"
return res, nil
}
// Validate is a dry-run that exercises the webserver's own configtest
// against the current on-disk state. No writes, no reload.
func (i *Installer) Validate() (Result, error) {
res := Result{
Action: "validate",
Webserver: i.Handler.Kind(),
SnippetPath: i.Handler.SnippetPath(),
ShippedVer: TemplateVersion,
}
if err := i.Handler.Validate(); err != nil {
res.Status = "fail"
res.Message = err.Error()
return res, err
}
res.Status = "ok"
res.Message = "configtest passed"
return res, nil
}
// renderTemplate prefixes the handler's body with the version marker
// the installer reads back at status/upgrade time.
func (i *Installer) renderTemplate() ([]byte, error) {
tpl, err := template.New(i.Handler.Kind()).Parse(i.Handler.Template())
if err != nil {
return nil, err
}
var b strings.Builder
b.WriteString(templateHeaderPrefix)
b.WriteString(strconv.Itoa(TemplateVersion))
b.WriteByte('\n')
if err := tpl.Execute(&b, i.templateData()); err != nil {
return nil, err
}
return []byte(b.String()), nil
}
func renderConfigFrom(cfg *config.Config) RenderConfig {
rc := RenderConfig{
ChallengeMapPath: challenge.DefaultMapPath,
ChallengeNginxMapPath: challenge.DefaultNginxMapPath,
ChallengeListenAddr: "127.0.0.1",
}
if cfg == nil {
return rc
}
rc.ChallengeListenAddr = strings.TrimSpace(cfg.Challenge.ListenAddr)
if rc.ChallengeListenAddr == "" {
rc.ChallengeListenAddr = "127.0.0.1"
}
rc.ChallengePublicURL = strings.TrimSpace(cfg.Challenge.PublicURL)
return rc
}
// RenderConfigFromConfig exposes the daemon-to-template projection for
// diagnostics that need to validate the same inputs before install.
func RenderConfigFromConfig(cfg *config.Config) RenderConfig {
return renderConfigFrom(cfg)
}
func (i *Installer) templateData() templateData {
mapPath := strings.TrimSpace(i.Config.ChallengeMapPath)
if mapPath == "" {
mapPath = challenge.DefaultMapPath
}
nginxMapPath := strings.TrimSpace(i.Config.ChallengeNginxMapPath)
if nginxMapPath == "" {
nginxMapPath = challenge.DefaultNginxMapPath
}
return templateData{
ChallengeMapPath: mapPath,
ChallengeNginxMapPath: nginxMapPath,
ChallengePublicURL: strings.TrimSpace(i.Config.ChallengePublicURL),
}
}
func (i *Installer) ensureChallengeMapFiles() error {
data := i.templateData()
mkdirAll := i.MkdirAll
if mkdirAll == nil {
mkdirAll = os.MkdirAll
}
for _, f := range []struct {
path string
body []byte
}{
{path: data.ChallengeMapPath, body: []byte("# Generated by CSM.\n")},
{path: data.ChallengeNginxMapPath, body: []byte("# Generated by CSM.\n")},
} {
if err := mkdirAll(filepath.Dir(f.path), 0o755); err != nil {
return err
}
statAt := i.StatAt
if statAt == nil {
statAt = os.Stat
}
if _, err := statAt(f.path); err == nil {
continue
} else if !errors.Is(err, os.ErrNotExist) {
return err
}
if err := i.WriteAt(f.path, f.body, 0o644); err != nil {
return err
}
}
return nil
}
func validateChallengePublicURL(rc RenderConfig) error {
raw := strings.TrimSpace(rc.ChallengePublicURL)
if raw == "" {
return ErrMissingPublicURL
}
u, err := url.Parse(raw)
if err != nil || !u.IsAbs() || u.Host == "" {
return fmt.Errorf("%w: %q", ErrInvalidPublicURL, raw)
}
if u.Scheme != "http" && u.Scheme != "https" {
return fmt.Errorf("%w: unsupported scheme %q", ErrInvalidPublicURL, u.Scheme)
}
if u.User != nil || u.RawQuery != "" || u.Fragment != "" || u.Path != "/challenge" {
return fmt.Errorf("%w: %q", ErrInvalidPublicURL, raw)
}
if isLoopbackHost(u.Hostname()) {
return fmt.Errorf("%w: host %q is loopback", ErrInvalidPublicURL, u.Hostname())
}
if isLoopbackListenAddr(rc.ChallengeListenAddr) {
return ErrLoopbackPublicURL
}
return nil
}
// ValidateChallengePublicURL applies the same direct-redirect validation used
// by Install so diagnostics and tests do not drift from installer behavior.
func ValidateChallengePublicURL(rc RenderConfig) error {
return validateChallengePublicURL(rc)
}
func isLoopbackListenAddr(addr string) bool {
addr = strings.TrimSpace(addr)
if addr == "" {
return true
}
host := addr
if h, _, err := net.SplitHostPort(addr); err == nil {
host = h
}
host = strings.Trim(host, "[]")
if host == "" {
return false
}
return isLoopbackHost(host)
}
func isLoopbackHost(host string) bool {
host = strings.TrimSpace(strings.Trim(host, "[]"))
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
// readSnippet parses the on-disk snippet header to recover the
// embedded version. Returns the raw bytes, a presence flag, and the
// parsed version (zero when the file exists but lacks the marker).
func (i *Installer) readSnippet() ([]byte, bool, int, error) {
data, err := i.ReadAt(i.Handler.SnippetPath())
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil, false, 0, err
}
return nil, false, 0, err
}
ver := parseHeaderVersion(data)
return data, true, ver, nil
}
// restore writes the previous bytes back (or removes the new file if
// there were none) after a failed validate/reload step. Best-effort:
// I/O errors here are reported via stderr but do not change the
// installer's return code, because the caller already knows the
// original error.
func (i *Installer) restore(prevBytes []byte, prevExists bool) {
if !prevExists {
if err := i.RemoveAt(i.Handler.SnippetPath()); err != nil && !errors.Is(err, os.ErrNotExist) {
fmt.Fprintf(i.Stderr, "webserver integration: rollback delete failed: %v\n", err)
}
return
}
if err := i.WriteAt(i.Handler.SnippetPath(), prevBytes, 0o644); err != nil {
fmt.Fprintf(i.Stderr, "webserver integration: rollback write failed: %v\n", err)
}
}
func parseHeaderVersion(data []byte) int {
scanner := bytes.SplitN(data, []byte("\n"), 2)
if len(scanner) == 0 {
return 0
}
line := strings.TrimSpace(string(scanner[0]))
if !strings.HasPrefix(line, templateHeaderPrefix) {
return 0
}
rest := strings.TrimSpace(strings.TrimPrefix(line, templateHeaderPrefix))
v, err := strconv.Atoi(rest)
if err != nil {
return 0
}
return v
}
func pickHandler(info platform.Info, r cmdRunner) (Handler, error) {
switch info.WebServer {
case platform.WSApache:
return newApacheHandler(info, r), nil
case platform.WSLiteSpeed:
return newLSWSHandler(info, r), nil
case platform.WSNginx:
return newNginxHandler(r), nil
default:
return nil, ErrUnknownWebserver
}
}
// atomicWrite writes data to a sibling temp file then renames it into
// place so the webserver never sees a half-written snippet. fsync on
// the directory is best-effort; rename + fsync on the file before
// rename gives crash safety on every common Linux filesystem.
func atomicWrite(path string, data []byte, mode os.FileMode) error {
dir := filepath.Dir(path)
tmp, err := os.CreateTemp(dir, ".csm-ws-install-*")
if err != nil {
return err
}
tmpName := tmp.Name()
cleanup := func() { _ = os.Remove(tmpName) }
if _, werr := tmp.Write(data); werr != nil {
_ = tmp.Close()
cleanup()
return werr
}
if serr := tmp.Sync(); serr != nil {
_ = tmp.Close()
cleanup()
return serr
}
if cerr := tmp.Close(); cerr != nil {
cleanup()
return cerr
}
if merr := os.Chmod(tmpName, mode); merr != nil {
cleanup()
return merr
}
return os.Rename(tmpName, path)
}
package webserver
import (
"context"
"fmt"
"time"
"github.com/pidginhost/csm/internal/platform"
)
// lswsHandler manages the LiteSpeed integration. The right snippet path
// depends on how LSWS is wired:
//
// - cPanel + LSWS: LSWS runs with <loadApacheConf>1</loadApacheConf>
// and reads cPanel's Apache config tree. The snippet drops at
// /etc/apache2/conf.d/csm-challenge.conf, same as plain Apache,
// and LSWS picks it up automatically.
//
// - Plain LSWS (no cPanel): the operator runs LSWS in native mode
// with /usr/local/lsws/conf/httpd_config.xml. There is no auto-
// include dir for text-style rewrite rules; the snippet goes in
// /usr/local/lsws/conf/templates/ and the operator must include
// it manually via the LSWS WebAdmin Console -> Server -> General
// -> Rewrite -> External Rewrite Rules. The installer writes the
// file but emits a stderr note pointing at the manual step.
type lswsHandler struct {
cmdRunner cmdRunner
cpanel bool
}
func newLSWSHandler(info platform.Info, r cmdRunner) *lswsHandler {
return &lswsHandler{cmdRunner: r, cpanel: info.IsCPanel()}
}
func (h *lswsHandler) Kind() string { return "lsws" }
func (h *lswsHandler) SnippetPath() string {
if h.cpanel {
return "/etc/apache2/conf.d/csm-challenge.conf"
}
return "/usr/local/lsws/conf/templates/csm-challenge.conf"
}
func (h *lswsHandler) Template() string { return lswsTemplate }
func (h *lswsHandler) Validate() error {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
// LSWS has no `lswsctrl conftest` verb (the documented surface is
// start|stop|restart|reload|condrestart|try-restart|status). The
// equivalent configtest is `lshttpd -t`, which parses the active
// config and exits non-zero on syntax errors without touching the
// running listener.
out, err := h.cmdRunner.Run(ctx, "/usr/local/lsws/bin/lshttpd", "-t")
if err != nil {
return fmt.Errorf("lsws lshttpd -t failed: %v\n%s", err, out)
}
return nil
}
// PostInstallInstructions returns operator-facing follow-up after a
// successful install. v4 of the snippet redirects challenged IPs to
// challenge.public_url directly; no LSWS External App is required, so
// the previous WebAdmin Console steps are obsolete.
func (h *lswsHandler) PostInstallInstructions() string { return "" }
func (h *lswsHandler) Reload() error {
// LSWS does not have a graceful reload equivalent; `restart` is the
// supported way to pick up new config without dropping established
// listener sockets (LSWS's internal supervisor handles the hand-
// off). The full-restart path is what the operator's own toolchain
// invokes on config change too.
ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second)
defer cancel()
out, err := h.cmdRunner.Run(ctx, "/usr/local/lsws/bin/lswsctrl", "restart")
if err != nil {
return fmt.Errorf("lsws restart failed: %v\n%s", err, out)
}
return nil
}
package webserver
import (
"context"
"fmt"
"time"
)
// nginxHandler manages the nginx integration. The snippet lands in
// /etc/nginx/conf.d/ where stock nginx auto-includes everything via
// the default http{} include glob. Validation uses `nginx -t`; reload
// uses `systemctl reload nginx` so existing connections drain
// gracefully.
type nginxHandler struct {
cmdRunner cmdRunner
}
func newNginxHandler(r cmdRunner) *nginxHandler {
return &nginxHandler{cmdRunner: r}
}
func (h *nginxHandler) Kind() string { return "nginx" }
func (h *nginxHandler) SnippetPath() string { return "/etc/nginx/conf.d/csm-challenge.conf" }
func (h *nginxHandler) Template() string { return nginxTemplate }
// PostInstallInstructions reminds the operator that the http{}
// snippet only ships the shared map; each guarded server{} block has
// to opt in with a one-line if-redirect. Nginx cannot apply this
// safely from an http{} drop-in.
func (h *nginxHandler) PostInstallInstructions() string {
return `Nginx detected. Per-server opt-in required so http{} drop-ins do
not blind-redirect already-protected hosts. For each server{} that
should respect the challenge:
server {
...
if ($csm_challenged) {
return 302 <challenge.public_url>?dest=$scheme://$host$request_uri;
}
}
Then run: nginx -t && systemctl reload nginx`
}
func (h *nginxHandler) Validate() error {
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
out, err := h.cmdRunner.Run(ctx, "nginx", "-t")
if err != nil {
return fmt.Errorf("nginx configtest failed: %v\n%s", err, out)
}
return nil
}
func (h *nginxHandler) Reload() error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
out, err := h.cmdRunner.Run(ctx, "systemctl", "reload", "nginx")
if err != nil {
return fmt.Errorf("nginx reload failed: %v\n%s", err, out)
}
return nil
}
package integrity
import (
"fmt"
"os"
"reflect"
"strings"
"sync"
"github.com/pidginhost/csm/internal/config"
)
var configWriteMu sync.Mutex
// ConfigWriteMutex serializes complete read-validate-write transactions for
// the main configuration. Callers must take it before reading csm.yaml, not
// only around the final rename, or an ETag check can race another writer.
func ConfigWriteMutex() *sync.Mutex {
return &configWriteMu
}
// WriteConfigBytesAtomic writes data to path with the same atomic-rename
// semantics SignAndSaveAtomic uses. Intended for paths that ship pre-signed
// bytes (e.g. restoring a snapshot whose hash already matches its content)
// where re-signing would mutate the integrity block we want to preserve.
func WriteConfigBytesAtomic(path string, data []byte) error {
return atomicWriteFile(path, data, 0o600)
}
// SignAndSavePreserving writes editedBytes to path after patching
// integrity.binary_hash and integrity.config_hash inside the byte
// stream itself, not by re-marshaling the cfg. Operator comments and
// untouched formatting outside the integrity block survive
// byte-for-byte.
//
// Verifies the final bytes decode via config.LoadBytes and match
// intendedClone under reflect.DeepEqual (with integrity fields
// normalised). Mismatch aborts the write.
//
// Atomic write semantics match SignAndSaveAtomic: same-directory
// tempfile, fsync, rename. On success, intendedClone.Integrity.BinaryHash
// and .ConfigHash are updated in place to reflect the hashes written to
// disk. intendedClone.ConfigFile must equal path and intendedClone.ConfigDir
// must equal confDir.
func SignAndSavePreserving(path, confDir string, editedBytes []byte, intendedClone *config.Config, binaryHash string) error {
if intendedClone == nil {
return fmt.Errorf("intendedClone is nil")
}
if intendedClone.ConfigFile != path {
return fmt.Errorf("intendedClone.ConfigFile=%q does not match path=%q", intendedClone.ConfigFile, path)
}
if intendedClone.ConfigDir != confDir {
return fmt.Errorf("intendedClone.ConfigDir=%q does not match confDir=%q", intendedClone.ConfigDir, confDir)
}
// Hash the operator-edited bytes before the integrity scalars are
// rewritten. HashConfigStableBytes ignores the integrity block, so
// the stored hash still matches the final file after the integrity
// patch.
newConfigHash := HashConfigStableBytes(editedBytes)
// Cover the conf.d fragments merged on top of this main config. Empty
// when there are none, leaving conf.d-free configs byte-identical to
// their prior baseline.
newConfdHash, err := HashConfDir(confDir, intendedClone.ConfD.IntegrityExempt)
if err != nil {
return fmt.Errorf("hashing conf.d: %w", err)
}
patched, err := config.YAMLEdit(editedBytes, []config.YAMLChange{
{Path: []string{"integrity", "binary_hash"}, Value: binaryHash},
{Path: []string{"integrity", "config_hash"}, Value: newConfigHash},
{Path: []string{"integrity", "confd_hash"}, Value: newConfdHash},
})
if err != nil {
return fmt.Errorf("patch integrity scalars: %w", err)
}
if stripIntegrityBlock(string(patched)) != stripIntegrityBlock(string(editedBytes)) {
return fmt.Errorf("integrity patch drift: bytes outside integrity block changed")
}
decoded, err := config.LoadBytes(patched)
if err != nil {
return fmt.Errorf("verify decode: %w", err)
}
decoded.ConfigFile = path
decoded.ConfigDir = confDir
expected := *intendedClone
expected.Integrity.BinaryHash = binaryHash
expected.Integrity.ConfigHash = newConfigHash
expected.Integrity.ConfdHash = newConfdHash
if !reflect.DeepEqual(decoded, &expected) {
return fmt.Errorf("yaml rewrite drift: decoded config does not match intended clone")
}
intendedClone.Integrity.BinaryHash = binaryHash
intendedClone.Integrity.ConfigHash = newConfigHash
intendedClone.Integrity.ConfdHash = newConfdHash
return atomicWriteFile(path, patched, 0o600)
}
// SignConfigFilePreserving signs path in place without re-marshaling the
// config. It updates only the operator-owned main config file, but folds the
// conf.d fragments under confDir into integrity.confd_hash so a later edit to
// any fragment is detected by Verify.
func SignConfigFilePreserving(path, confDir, binaryHash string) (configHash, confdHash string, err error) {
cfg, err := SignConfigFilePreservingSnapshot(path, confDir, binaryHash)
if err != nil {
return "", "", err
}
return cfg.Integrity.ConfigHash, cfg.Integrity.ConfdHash, nil
}
// SignConfigFilePreservingSnapshot signs path and returns the decoded main
// config whose conf.d exemption policy selected the new ConfdHash. Callers
// that mirror the new integrity metadata into a running config must carry
// ConfD from this same snapshot so later Verify calls use the matching list.
func SignConfigFilePreservingSnapshot(path, confDir, binaryHash string) (*config.Config, error) {
// #nosec G304 -- operator-configured config path.
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read config: %w", err)
}
cfg, err := config.LoadBytes(data)
if err != nil {
return nil, err
}
cfg.ConfigFile = path
cfg.ConfigDir = confDir
if err := SignAndSavePreserving(path, confDir, data, cfg, binaryHash); err != nil {
return nil, err
}
return cfg, nil
}
// stripIntegrityBlock removes the top-level `integrity:` mapping and
// its indented children from s, then returns the remaining text.
// Used both by SignAndSavePreserving's drift guard and by tests, so
// both compare using the same definition of "outside the integrity
// block".
func stripIntegrityBlock(s string) string {
lines := strings.Split(s, "\n")
var out []string
inIntegrity := false
for _, line := range lines {
if strings.HasPrefix(line, "integrity:") {
inIntegrity = true
continue
}
if inIntegrity {
if strings.HasPrefix(line, " ") || strings.HasPrefix(line, "\t") || line == "" {
continue
}
inIntegrity = false
}
out = append(out, line)
}
return strings.Join(out, "\n")
}
package integrity
import (
"bufio"
"bytes"
"crypto/sha256"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"strings"
"gopkg.in/yaml.v3"
"github.com/pidginhost/csm/internal/config"
)
// Sentinel causes for a Verify failure. Callers such as `csm doctor` match on
// them to suggest the right remedy; the wrapped messages carry the hashes.
var (
ErrBinaryHashMismatch = errors.New("binary hash mismatch")
ErrConfigHashMismatch = errors.New("config hash mismatch")
ErrConfdHashMismatch = errors.New("conf.d hash mismatch")
)
// HashFile returns the SHA256 hash of a file.
func HashFile(path string) (string, error) {
// #nosec G304 -- integrity hashing of operator-configured binary/config paths.
f, err := os.Open(path)
if err != nil {
return "", err
}
defer func() { _ = f.Close() }()
h := sha256.New()
if _, err := io.Copy(h, f); err != nil {
return "", err
}
return fmt.Sprintf("sha256:%x", h.Sum(nil)), nil
}
// HashConfigStable hashes the config file excluding the integrity section,
// so that writing hashes back to the config doesn't change the hash.
//
// Line length is capped at 1 MiB; anything longer is treated as a
// corrupted config and the digest reflects whatever was scanned up to
// the truncation point. A corruption-driven hash will trip
// integrity.Verify on the next start, which is the intended response.
func HashConfigStable(path string) (string, error) {
// #nosec G304 -- operator-supplied config file path.
data, err := os.ReadFile(path)
if err != nil {
return "", fmt.Errorf("reading config: %w", err)
}
return HashConfigStableBytes(data), nil
}
// HashConfigStableBytes is the in-memory counterpart to
// HashConfigStable: compute the stable hash from an already-serialised
// config body without touching disk. Used by the SIGHUP reload path
// and `csm rehash` so both can hash a prospective file content before
// committing it, avoiding the two-pass-write dance that could leave
// integrity.config_hash blank on disk if the second save failed.
//
// Returns the same "sha256:..." string shape as HashConfigStable.
// The scanner is sized to cover config lines well beyond any
// realistic CSM yaml; a line longer than 1 MiB is treated as
// corruption and the final digest reflects whatever was scanned
// before the truncation (the resulting mismatch against any stored
// hash will trip integrity.Verify, which is the correct response).
func HashConfigStableBytes(data []byte) string {
h := sha256.New()
scanner := bufio.NewScanner(bytes.NewReader(data))
scanner.Buffer(make([]byte, 0, 64*1024), 1<<20)
inIntegrity := false
for scanner.Scan() {
line := scanner.Text()
// Skip the integrity section
if strings.HasPrefix(line, "integrity:") {
inIntegrity = true
continue
}
if inIntegrity {
// Still inside integrity block (indented lines)
if strings.HasPrefix(line, " ") || strings.HasPrefix(line, "\t") || line == "" {
continue
}
inIntegrity = false
}
_, _ = h.Write([]byte(line + "\n"))
}
return fmt.Sprintf("sha256:%x", h.Sum(nil))
}
// SignAndSaveAtomic re-computes integrity.config_hash for cfg and
// writes the result to cfg.ConfigFile atomically. Atomicity means:
// the on-disk file either carries the prior content (operation
// failed) or the fully-signed new content (operation succeeded).
// There is no window where the file exists on disk with an empty or
// stale config_hash, so a crash between the two passes of the
// previous two-save dance can no longer put the daemon into a
// crash-loop on next startup.
//
// The integrity.binary_hash is set to the supplied binaryHash. The
// CALLER is responsible for picking the right value: `csm rehash`
// hashes /opt/csm/csm afresh; SIGHUP reload preserves the prior
// daemon's binary hash because a reload cannot upgrade the binary.
//
// Implementation: marshal the config with a blank ConfigHash, hash
// the stable form of those bytes, store the hash, marshal again,
// write to a sibling temp file, rename into place. The YAML hashing
// strips the integrity block so both marshals round-trip to the
// same stable hash.
func SignAndSaveAtomic(cfg *config.Config, binaryHash string) error {
cfg.Integrity.BinaryHash = binaryHash
cfg.Integrity.ConfigHash = ""
confdHash, err := HashConfDir(cfg.ConfigDir, cfg.ConfD.IntegrityExempt)
if err != nil {
return fmt.Errorf("hashing conf.d: %w", err)
}
cfg.Integrity.ConfdHash = confdHash
preHash, err := yaml.Marshal(cfg)
if err != nil {
return fmt.Errorf("marshal (pre-hash): %w", err)
}
cfg.Integrity.ConfigHash = HashConfigStableBytes(preHash)
final, err := yaml.Marshal(cfg)
if err != nil {
return fmt.Errorf("marshal (post-hash): %w", err)
}
return atomicWriteFile(cfg.ConfigFile, final, 0o600)
}
// atomicWriteFile writes data to a temp file in the same directory
// as path, fsyncs and closes the temp, then renames it onto path.
// Rename is atomic on POSIX when source and destination are on the
// same filesystem (which is why the temp is created in the target's
// dir, not /tmp). Permission is applied before the rename.
func atomicWriteFile(path string, data []byte, perm os.FileMode) error {
targetPath, err := atomicWriteTarget(path)
if err != nil {
return err
}
dir := filepath.Dir(targetPath)
tmp, err := os.CreateTemp(dir, ".csm-cfg-*.tmp")
if err != nil {
return fmt.Errorf("create temp: %w", err)
}
tmpName := tmp.Name()
// Best-effort cleanup: if any error below leaves the temp
// behind, unlink it so we do not fill the dir with orphans.
cleanup := func() { _ = os.Remove(tmpName) }
if _, err := tmp.Write(data); err != nil {
_ = tmp.Close()
cleanup()
return fmt.Errorf("write temp: %w", err)
}
if err := tmp.Chmod(perm); err != nil {
_ = tmp.Close()
cleanup()
return fmt.Errorf("chmod temp: %w", err)
}
if err := tmp.Sync(); err != nil {
_ = tmp.Close()
cleanup()
return fmt.Errorf("fsync temp: %w", err)
}
if err := tmp.Close(); err != nil {
cleanup()
return fmt.Errorf("close temp: %w", err)
}
if err := os.Rename(tmpName, targetPath); err != nil {
cleanup()
return fmt.Errorf("rename: %w", err)
}
if err := syncDirectory(dir); err != nil {
return err
}
return nil
}
func atomicWriteTarget(path string) (string, error) {
info, err := os.Lstat(path)
if err != nil {
if os.IsNotExist(err) {
return path, nil
}
return "", fmt.Errorf("stat config path: %w", err)
}
if info.Mode()&os.ModeSymlink == 0 {
return path, nil
}
resolved, err := filepath.EvalSymlinks(path)
if err != nil {
return "", fmt.Errorf("resolve config symlink: %w", err)
}
return resolved, nil
}
func syncDirectory(dir string) error {
// #nosec G304 -- dir is derived from the caller-owned config path.
d, err := os.Open(dir)
if err != nil {
return fmt.Errorf("open dir: %w", err)
}
if err := d.Sync(); err != nil {
_ = d.Close()
return fmt.Errorf("fsync dir: %w", err)
}
if err := d.Close(); err != nil {
return fmt.Errorf("close dir: %w", err)
}
return nil
}
// HashConfDir returns a stable digest over every conf.d drop-in fragment that
// would be merged on top of the main config, in merge order. It returns the
// empty string when there are no fragments, so a config without conf.d keeps
// the empty digest and its pre-existing baseline still verifies after upgrade.
//
// Each fragment is domain-separated by name and length so two fragments cannot
// collide by shuffling bytes across the filename boundary.
//
// Fragments named in exempt (confd.integrity_exempt) hash as if absent: their
// owning integration rewrites them on its own schedule, and pinning them would
// turn every one of those rewrites into a failed restart.
func HashConfDir(confDir string, exempt []string) (string, error) {
frags, err := config.ConfDirFragmentDigestInput(confDir)
if err != nil {
return "", err
}
h := sha256.New()
covered := 0
for _, f := range frags {
if isExemptFragment(f.Name, exempt) {
continue
}
covered++
fmt.Fprintf(h, "confd-fragment:%s:%d\n", f.Name, len(f.Data))
_, _ = h.Write(f.Data)
}
if covered == 0 {
return "", nil
}
return fmt.Sprintf("sha256:%x", h.Sum(nil)), nil
}
// hashOrNone keeps "expected , got ..." out of operator-facing errors when one
// side of a conf.d comparison is the empty no-fragments digest.
func hashOrNone(h string) string {
if h == "" {
return "none"
}
return h
}
func isExemptFragment(name string, exempt []string) bool {
for _, e := range exempt {
if e == name {
return true
}
}
return false
}
// Verify checks the binary and config file integrity.
func Verify(binaryPath string, cfg *config.Config) error {
if cfg.Integrity.BinaryHash == "" {
return nil // Not yet baselined
}
currentHash, err := HashFile(binaryPath)
if err != nil {
return fmt.Errorf("hashing binary: %w", err)
}
if currentHash != cfg.Integrity.BinaryHash {
return fmt.Errorf("%w: expected %s, got %s; run `csm rehash` after an intentional binary upgrade",
ErrBinaryHashMismatch, cfg.Integrity.BinaryHash, currentHash)
}
if cfg.Integrity.ConfigHash != "" {
configHash, err := HashConfigStable(cfg.ConfigFile)
if err != nil {
return fmt.Errorf("hashing config: %w", err)
}
if configHash != cfg.Integrity.ConfigHash {
return fmt.Errorf("%w: expected %s, got %s; run `csm rehash` after an intentional csm.yaml change",
ErrConfigHashMismatch, cfg.Integrity.ConfigHash, configHash)
}
// conf.d fragments are merged on top of the main config at load time,
// so they must be covered too. A symmetric comparison closes the gap
// both ways: a tampered or added fragment makes the computed digest
// diverge from the stored one, and a baseline taken without conf.d
// stays empty == empty. Operators who already use conf.d must re-run
// `csm rehash` once after upgrade to populate confd_hash.
confdHash, err := HashConfDir(cfg.ConfigDir, cfg.ConfD.IntegrityExempt)
if err != nil {
return fmt.Errorf("hashing conf.d: %w", err)
}
if confdHash != cfg.Integrity.ConfdHash {
return fmt.Errorf("%w: a drop-in under %s changed since the config was last signed (expected %s, got %s); run `csm rehash` after an intentional conf.d change, or list fragments an integration rewrites under confd.integrity_exempt",
ErrConfdHashMismatch, cfg.ConfigDir, hashOrNone(cfg.Integrity.ConfdHash), hashOrNone(confdHash))
}
}
return nil
}
package jstaint
import "github.com/tdewolff/parse/v2/js"
// canonicalVar follows the parser's Link chain to the declaration identity. The
// parser sets Link when it merges an undeclared use with its later declaration
// or an outer-scope binding, so following it to the end yields one stable
// identity per variable. Shadowed declarations keep distinct identities because
// each is its own declared Var with a nil Link.
func canonicalVar(v *js.Var) *js.Var {
for v.Link != nil {
v = v.Link
}
return v
}
// tdewolff's walker reaches ClassElementName but does not descend into its
// computed expression. Walk it explicitly so limits and semantic visitors see
// code that JavaScript executes while defining a class.
func walkComputedClassName(v js.IVisitor, node js.INode) bool {
name, ok := node.(*js.ClassElementName)
if !ok {
return false
}
js.Walk(v, name.Computed)
return true
}
// staticStringOrIdent returns the literal identifier or string value of expr, if
// expr is a literal. It accepts both concrete literal forms the AST uses: a
// DotExpr member name arrives as a LiteralExpr value, while a bracket index or
// call argument arrives as a *LiteralExpr pointer. Handling only the pointer
// form would silently drop every dotted member name.
//
// String tokens are unquoted but not unescaped: escape spellings such as
// \x6b are out of version 1 scope, so the raw bytes between the quotes are
// returned verbatim.
func staticStringOrIdent(expr js.IExpr) (string, bool) {
data, ok := staticBytesOrIdent(expr)
if !ok {
return "", false
}
return string(data), true
}
// staticBytesOrIdent is the allocation-free form of staticStringOrIdent for
// AST walkers that only need to compare or render the literal bytes.
func staticBytesOrIdent(expr js.IExpr) ([]byte, bool) {
switch lit := expr.(type) {
case *js.LiteralExpr:
return literalBytes(lit.TokenType, lit.Data)
case js.LiteralExpr:
return literalBytes(lit.TokenType, lit.Data)
default:
return nil, false
}
}
func literalText(tt js.TokenType, data []byte) (string, bool) {
literal, ok := literalBytes(tt, data)
if !ok {
return "", false
}
return string(literal), true
}
func literalBytes(tt js.TokenType, data []byte) ([]byte, bool) {
switch tt {
case js.IdentifierToken:
return data, true
case js.StringToken:
if len(data) < 2 {
return nil, false
}
q := data[0]
if (q != '\'' && q != '"') || data[len(data)-1] != q {
return nil, false
}
return data[1 : len(data)-1], true
default:
return nil, false
}
}
package jstaint
import (
"strconv"
"strings"
"unicode"
"unicode/utf8"
)
const (
// maxEvidencePaths bounds the returned evidence flows. TotalResults still
// reports the true count, so a caller can name the additional paths.
maxEvidencePaths = 8
// maxViaSegments bounds a displayed laundering chain: the first headViaSegments
// segments, one omission marker, and the final tailViaSegments segments.
maxViaSegments = 32
headViaSegments = 16
tailViaSegments = 15
// maxSegmentBytes bounds one displayed source, via, or sink segment so a long
// attacker-controlled identifier cannot create unbounded findings.
maxSegmentBytes = 64
)
// finalizeResults returns the sorted evidence flows capped to maxEvidencePaths,
// the true pre-cap flow count, and whether any evidence was shortened. Every
// displayed segment is sanitized and length-bounded, and a chain longer than
// maxViaSegments keeps its head and tail around one omission marker.
func (a *analysis) finalizeResults() ([]Result, int, bool) {
sorted := a.sortedResults()
total := len(sorted)
truncated := false
if total > maxEvidencePaths {
truncated = true
sorted = sorted[:maxEvidencePaths]
}
for i := range sorted {
source, cutSource := boundSegment(sorted[i].Source)
sink, cutSink := boundSegment(sorted[i].Sink)
via, cutVia := truncateVia(sorted[i].Via)
sorted[i].Source = source
sorted[i].Sink = sink
sorted[i].Via = via
truncated = truncated || cutSource || cutSink || cutVia
}
return sorted, total, truncated
}
// truncateVia bounds each segment and, for a chain longer than maxViaSegments,
// retains the head and tail around one marker that names the omitted count.
func truncateVia(via []string) ([]string, bool) {
if len(via) == 0 {
return via, false
}
truncated := false
bounded := make([]string, len(via))
for i, s := range via {
seg, cut := boundSegment(s)
bounded[i] = seg
truncated = truncated || cut
}
if len(bounded) <= maxViaSegments {
return bounded, truncated
}
omitted := len(bounded) - (headViaSegments + tailViaSegments)
out := make([]string, 0, maxViaSegments)
out = append(out, bounded[:headViaSegments]...)
out = append(out, "["+strconv.Itoa(omitted)+" segments omitted]")
out = append(out, bounded[len(bounded)-tailViaSegments:]...)
return out, true
}
// boundSegment sanitizes one display segment and bounds it to maxSegmentBytes at a
// rune boundary. It reports whether the length bound shortened the segment;
// sanitizing invalid bytes alone does not count as truncation.
func boundSegment(s string) (string, bool) {
clean := sanitizeSegment(s)
if len(clean) <= maxSegmentBytes {
return clean, false
}
const marker = "..."
cut := maxSegmentBytes - len(marker)
for cut > 0 && !utf8.RuneStart(clean[cut]) {
cut--
}
return clean[:cut] + marker, true
}
// sanitizeSegment emits valid UTF-8 and replaces control bytes so evidence text
// stays printable and searchable.
func sanitizeSegment(s string) string {
for i := 0; i < len(s); {
r, size := utf8.DecodeRuneInString(s[i:])
if (r == utf8.RuneError && size == 1) || unicode.IsControl(r) {
var clean strings.Builder
clean.Grow(len(s))
clean.WriteString(s[:i])
for i < len(s) {
r, size = utf8.DecodeRuneInString(s[i:])
if (r == utf8.RuneError && size == 1) || unicode.IsControl(r) {
clean.WriteByte('?')
} else {
clean.WriteString(s[i : i+size])
}
i += size
}
return clean.String()
}
i += size
}
return s
}
package jstaint
import (
"bytes"
"encoding/json"
"io"
"strconv"
"strings"
"github.com/tdewolff/parse/v2"
"github.com/tdewolff/parse/v2/css"
"github.com/tdewolff/parse/v2/html"
)
// isNonJSDocument is only used after JavaScript parsing fails. The deep walk
// supplies every file regardless of extension, so token matches in data and
// templates must not be reported as JavaScript coverage failures. Ambiguous
// content stays a parse failure; filenames never decide which files to skip.
func isNonJSDocument(src []byte) bool {
src = bytes.TrimSpace(bytes.TrimPrefix(src, []byte("\xef\xbb\xbf")))
if opensWithPHPTag(src) || json.Valid(src) {
return true
}
lexer := html.NewLexer(parse.NewInputBytes(src[:len(src):len(src)]))
for {
token, data := lexer.Next()
switch token {
case html.CommentToken:
// HTML comments can precede the document's first tag.
continue
case html.TextToken:
if len(bytes.TrimSpace(data)) == 0 {
continue
}
case html.StartTagToken, html.DoctypeToken:
return true
}
break
}
return isStylesheet(src) || isTranslationCatalog(src)
}
func opensWithPHPTag(src []byte) bool {
if bytes.HasPrefix(src, []byte("<?=")) {
return true
}
const tag = "<?php"
if len(src) < len(tag) || !bytes.EqualFold(src[:len(tag)], []byte(tag)) {
return false
}
// The long opening tag requires whitespace, unlike the echo tag.
return len(src) == len(tag) || bytes.ContainsAny(src[len(tag):len(tag)+1], " \t\r\n")
}
func isStylesheet(src []byte) bool {
// CSS parsers recover at EOF, even from unclosed strings and blocks. Only
// classify a complete token stream so malformed JS retains its warning.
if !completeCSSTokens(src) {
return false
}
p := css.NewParser(parse.NewInputBytes(src[:len(src):len(src)]), false)
sawRule := false
for {
grammar, _, _ := p.Next()
switch grammar {
case css.ErrorGrammar:
return p.Err() == io.EOF && sawRule
case css.AtRuleGrammar, css.BeginAtRuleGrammar:
sawRule = true
case css.BeginRulesetGrammar:
if !plausibleCSSSelector(p.Values()) {
return false
}
sawRule = true
}
}
}
// The CSS grammar parser accepts arbitrary selector tokens. Reject JS
// assignments and calls rather than treating their blocks as style rules.
func plausibleCSSSelector(tokens []css.Token) bool {
brackets, functions := 0, 0
previous := css.EmptyToken
for _, token := range tokens {
switch token.TokenType {
case css.WhitespaceToken, css.CommentToken:
continue
case css.LeftBracketToken:
brackets++
case css.RightBracketToken:
brackets--
case css.FunctionToken:
if previous != css.ColonToken && functions == 0 {
return false
}
functions++
case css.LeftParenthesisToken:
if functions == 0 {
return false
}
functions++
case css.RightParenthesisToken:
functions--
case css.DelimToken:
if brackets == 0 && !bytes.ContainsAny(token.Data, ".*>+~|&") {
return false
}
case css.SemicolonToken, css.AtKeywordToken:
return false
}
previous = token.TokenType
}
return true
}
func completeCSSTokens(src []byte) bool {
l := css.NewLexer(parse.NewInputBytes(src[:len(src):len(src)]))
var closes []css.TokenType
for {
token, data := l.Next()
switch token {
case css.ErrorToken:
return l.Err() == io.EOF && len(closes) == 0
case css.BadStringToken, css.BadURLToken:
return false
case css.StringToken:
if len(data) < 2 || data[len(data)-1] != data[0] {
return false
}
case css.CommentToken:
if !bytes.HasSuffix(data, []byte("*/")) {
return false
}
case css.URLToken:
if !bytes.HasSuffix(data, []byte(")")) {
return false
}
case css.LeftBraceToken:
closes = append(closes, css.RightBraceToken)
case css.LeftBracketToken:
closes = append(closes, css.RightBracketToken)
case css.LeftParenthesisToken, css.FunctionToken:
closes = append(closes, css.RightParenthesisToken)
case css.RightBraceToken, css.RightBracketToken, css.RightParenthesisToken:
if len(closes) == 0 || closes[len(closes)-1] != token {
return false
}
closes = closes[:len(closes)-1]
}
}
}
// Catalog strings can contain entire JavaScript programs. Require every
// non-comment line to be a gettext directive or a quoted continuation, so a
// catalog-like prefix cannot discard a trailing program or a parse failure.
func isTranslationCatalog(src []byte) bool {
sawID, sawTranslation, inString := false, false, false
for line := range strings.Lines(string(src)) {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
if line[0] != '"' {
end := strings.IndexAny(line, " \t")
if end < 0 {
return false
}
key, value := line[:end], line[end:]
switch {
case key == "msgid":
sawID = true
case key == "msgstr":
sawTranslation = true
case strings.HasPrefix(key, "msgstr[") && strings.HasSuffix(key, "]"):
index := key[len("msgstr[") : len(key)-1]
if index == "" || strings.Trim(index, "0123456789") != "" {
return false
}
sawTranslation = true
case key == "msgctxt", key == "msgid_plural":
default:
return false
}
line = strings.TrimSpace(value)
inString = true
}
if !inString || len(line) < 2 || line[0] != '"' {
return false
}
if _, err := strconv.Unquote(line); err != nil {
return false
}
}
return sawID && sawTranslation
}
package jstaint
import "github.com/tdewolff/parse/v2/js"
// allocate installs a fresh instance at site and returns a value referencing it.
// Inside a loop the instance is summarized, because the site can produce many
// runtime objects; otherwise it is the current singleton for this control-flow
// path. A prior current instance at the same site is promoted to the summary
// class first, so a re-allocation cannot clear taint on an earlier instance that
// was aliased or published.
func (a *analysis) allocate(st *state, site int, inLoop bool) value {
a.promoteCurrent(st, site)
return a.installLiteral(st, site, inLoop, &object{})
}
// installLiteral publishes a fully evaluated object or array literal at its
// allocation site. Building it separately preserves construction semantics
// before loop instances are merged into a summary.
func (a *analysis) installLiteral(st *state, site int, inLoop bool, fresh *object) value {
id := allocID{site: site, summary: inLoop}
if inLoop {
if old, ok := st.heap[id]; ok {
fresh = mergeObject(old, fresh)
}
}
st.installObject(id, fresh)
a.fact()
return value{allocs: allocSet{{id: id}: true}, allocOnly: true}
}
// promoteCurrent moves any current instance at site into the summary class and
// rewrites every reference to it. This is how two recency classes model an
// unbounded number of runtime objects: the escaped or aliased earlier object
// survives in the summary while a fresh current instance takes its place.
func (a *analysis) promoteCurrent(st *state, site int) {
from := allocID{site: site, summary: false}
cur, ok := st.heap[from]
if !ok {
return
}
to := allocID{site: site, summary: true}
if ex, ok := st.heap[to]; ok {
st.installObject(to, mergeObject(ex, cur))
} else {
// A plain move within one state preserves whatever exclusivity the
// object already had.
st.ownHeap()
st.heap[to] = cur
}
st.dropObject(from)
rewriteAlloc(st, from, to)
}
func rewriteAlloc(st *state, from, to allocID) {
for k, v := range st.env {
if hasAllocID(v.allocs, from) {
v.allocs = replaceAlloc(v.allocs, from, to)
st.setEnv(k, v)
}
}
for node, v := range st.captures {
if hasAllocID(v.allocs, from) {
v.allocs = replaceAlloc(v.allocs, from, to)
st.setCapture(node, v)
}
}
// Read-only scan first: mutObject may replace st.heap with an owned copy,
// which must not happen while ranging over the map being replaced.
var hit []allocID
for id, o := range st.heap {
if objectRefersToAlloc(o, from) {
hit = append(hit, id)
}
}
for _, id := range hit {
o := st.mutObject(id)
for fk, fv := range o.fields {
if hasAllocID(fv.allocs, from) {
fv.allocs = replaceAlloc(fv.allocs, from, to)
o.fields[fk] = fv
}
}
if hasAllocID(o.wild.allocs, from) {
o.wild.allocs = replaceAlloc(o.wild.allocs, from, to)
}
if hasAllocID(o.wildReq.allocs, from) {
o.wildReq.allocs = replaceAlloc(o.wildReq.allocs, from, to)
}
if hasAllocID(o.elem.allocs, from) {
o.elem.allocs = replaceAlloc(o.elem.allocs, from, to)
}
}
}
func objectRefersToAlloc(o *object, id allocID) bool {
for _, fv := range o.fields {
if hasAllocID(fv.allocs, id) {
return true
}
}
return hasAllocID(o.wild.allocs, id) || hasAllocID(o.wildReq.allocs, id) ||
hasAllocID(o.elem.allocs, id)
}
func replaceAlloc(s allocSet, from, to allocID) allocSet {
out := make(allocSet, len(s))
for ref := range s {
if ref.id == from {
ref.id = to
out[ref] = true
} else {
out[ref] = true
}
}
return out
}
func valueCarriesDepthZero(st *state, v value) bool {
if taintCarriesDepthZero(v.scalar) {
return true
}
seen := map[allocRef]bool{}
stack := make([]allocRef, 0, len(v.allocs))
for ref := range v.allocs {
stack = append(stack, ref)
}
push := func(parent allocRef, child value) bool {
child = applyRefDepth(child, parent)
for fact := range child.scalar {
if fact.callDepth == 0 {
return true
}
}
for ref := range child.allocs {
if !seen[ref] {
stack = append(stack, ref)
}
}
return false
}
for len(stack) > 0 {
ref := stack[len(stack)-1]
stack = stack[:len(stack)-1]
if seen[ref] {
continue
}
seen[ref] = true
o := st.heap[ref.id]
if o == nil {
continue
}
for _, field := range o.fields {
if push(ref, field) {
return true
}
}
if push(ref, o.elem) || push(ref, o.wild) || push(ref, o.wildReq) {
return true
}
}
return false
}
func taintCarriesDepth(ts taintSet, depth int) bool {
for fact := range ts {
if int(fact.callDepth) == depth {
return true
}
}
return false
}
func taintCarriesDepthZero(ts taintSet) bool {
return taintCarriesDepth(ts, 0)
}
// resetAllocRefDepth records that an allocation now carries a fact created in
// the current invocation. Existing field facts retain their own depth, but
// aliases must no longer impose an older return-path barrier on the new write.
func resetAllocRefDepth(st *state, ids map[allocID]bool) {
needsReset := func(v value) bool {
for ref := range v.allocs {
if ids[ref.id] && (ref.minDepth != 0 || ref.advance != 0) {
return true
}
}
return false
}
reset := func(v value) value {
refs := make(allocSet, len(v.allocs))
for ref := range v.allocs {
if ids[ref.id] {
ref.minDepth = 0
ref.advance = 0
}
refs[ref] = true
}
v.allocs = refs
return v
}
for cv, v := range st.env {
if needsReset(v) {
st.setEnv(cv, reset(v))
}
}
for node, v := range st.captures {
if needsReset(v) {
st.setCapture(node, reset(v))
}
}
// Read-only scan first: mutObject may replace st.heap with an owned copy,
// which must not happen while ranging over the map being replaced.
var hit []allocID
for id, o := range st.heap {
if objectNeedsDepthReset(o, needsReset) {
hit = append(hit, id)
}
}
for _, id := range hit {
o := st.mutObject(id)
for name, v := range o.fields {
if needsReset(v) {
o.fields[name] = reset(v)
}
}
if needsReset(o.elem) {
o.elem = reset(o.elem)
}
if needsReset(o.wild) {
o.wild = reset(o.wild)
}
if needsReset(o.wildReq) {
o.wildReq = reset(o.wildReq)
}
}
}
func objectNeedsDepthReset(o *object, needsReset func(value) bool) bool {
for _, v := range o.fields {
if needsReset(v) {
return true
}
}
return needsReset(o.elem) || needsReset(o.wild) || needsReset(o.wildReq)
}
func hasAllocID(s allocSet, id allocID) bool {
for ref := range s {
if ref.id == id {
return true
}
}
return false
}
func allocsConstrained(s allocSet) bool {
for ref := range s {
if ref.minDepth != 0 || ref.advance != 0 {
return true
}
}
return false
}
// writeField taints one field on the receiver's allocations. A clean write is a
// strong update (it can clear taint) only when the receiver is exactly one
// current instance; a summary field or a receiver with several possible
// allocations is weak-updated, so it cannot lose taint an aliased path still
// carries. Statically known array indexes are distinct fields; unresolved array
// elements and wildcard writes stay weak because they can represent many keys.
func (a *analysis) writeField(st *state, recv value, key fieldKey, rhs value) {
ids := uniqueAllocIDs(recv.allocs)
strong := recv.allocOnly && len(ids) == 1 && soleCurrent(recv.allocs)
for id := range ids {
o := st.mutObject(id)
if o == nil {
o = &object{}
st.installObject(id, o)
}
switch key.kind {
case fieldNamed:
o.setNamed(key.name, rhs, strong && !id.summary)
case fieldElem:
if key.name != "" {
o.setNamed(key.name, rhs, strong && !id.summary)
} else {
o.weakElem(rhs, strong && !id.summary)
}
default:
o.writeWild(rhs, strong && !id.summary)
}
}
if a.callDepth == 0 && allocsConstrained(recv.allocs) && valueCarriesDepthZero(st, rhs) {
resetAllocRefDepth(st, ids)
}
a.fact()
}
// deleteField removes one statically known field when the receiver is a lone
// current instance. A delete through a summary or ambiguous receiver is a weak
// update, so it cannot prove that every represented runtime field disappeared.
func (a *analysis) deleteField(st *state, recv value, key fieldKey) {
ids := uniqueAllocIDs(recv.allocs)
if len(ids) != 1 || !soleCurrent(recv.allocs) {
return
}
for id := range ids {
o := st.mutObject(id)
if o == nil {
return
}
switch key.kind {
case fieldNamed:
o.deleteNamed(key.name)
case fieldElem:
if key.name == "" {
return
}
o.deleteNamed(key.name)
default:
return
}
a.fact()
}
}
func soleCurrent(s allocSet) bool {
for ref := range s {
if ref.id.summary {
return false
}
}
return len(s) != 0
}
func uniqueAllocIDs(s allocSet) map[allocID]bool {
ids := make(map[allocID]bool, len(s))
for ref := range s {
ids[ref.id] = true
}
return ids
}
// readField reads one field across every allocation the receiver may reference.
func (a *analysis) readField(st *state, recv value, key fieldKey) value {
var out value
have := false
definiteAlloc := recv.allocOnly && len(recv.allocs) != 0
definiteField := definiteAlloc
for ref := range recv.allocs {
o := st.heap[ref.id]
if o == nil {
definiteAlloc = false
definiteField = false
continue
}
fv, definite := getField(o, key)
fv = applyRefDepth(fv, ref)
out, have = mergePresentValue(out, have, fv)
if !definite || !fv.allocOnly {
definiteAlloc = false
}
if !definite {
definiteField = false
}
}
if !have {
return value{}
}
out.allocOnly = definiteAlloc
if !definiteField {
out.scheme = schemeState{}
}
return out
}
func fieldDefinitelyPresent(st *state, recv value, key fieldKey) bool {
if !recv.allocOnly || len(recv.allocs) == 0 {
return false
}
seen := map[allocID]bool{}
for ref := range recv.allocs {
if seen[ref.id] {
continue
}
seen[ref.id] = true
o := st.heap[ref.id]
if o == nil {
return false
}
if _, definite := getField(o, key); !definite {
return false
}
}
return len(seen) != 0
}
// serializeStringify returns the scalar taint JSON.stringify(v) would carry.
// Reachable field taint is folded in by an iterative graph walk. An
// allocation whose reachable graph contains a cycle contributes no value, because
// the runtime throws before producing output; a union that also holds an acyclic
// alternative keeps the acyclic taint. The walk is iterative because a heap graph
// built through assignments can be far deeper than the AST, and a recursive walk
// could overflow the Go stack, a fatal error recover cannot intercept.
func (a *analysis) serializeStringify(st *state, v value) taintSet {
out := v.scalar
for ref := range v.allocs {
if ts, cyclic := a.walkStringify(st, ref); !cyclic {
out = mergeTaint(out, ts)
}
}
return out
}
type stringifySlot struct {
value value
optional bool
collection bool
}
func stringifyValues(o *object, ref allocRef) []stringifySlot {
values := make([]stringifySlot, 0, len(o.fields)+3)
for name, fv := range o.fields {
if !o.array || isArrayIndexName(name) {
values = append(values, stringifySlot{value: applyRefDepth(fv, ref), optional: !o.must[name]})
}
}
if o.elemMay {
values = append(values, stringifySlot{
value: applyRefDepth(o.elem, ref), optional: !o.elemMust, collection: true,
})
}
if o.wildMay {
values = append(values, stringifySlot{value: applyRefDepth(o.wild, ref), optional: true})
}
if o.wildMust && !o.array {
values = append(values, stringifySlot{value: applyRefDepth(o.wildReq, ref)})
}
return values
}
func arrayNeighbors(o *object, ref allocRef, addScalar func(taintSet)) []allocRef {
if !o.array {
return nil
}
var nbrs []allocRef
collect := func(v value) {
v = applyRefDepth(v, ref)
addScalar(v.scalar)
for aid := range v.allocs {
nbrs = append(nbrs, aid)
}
}
for name, fv := range o.fields {
if isArrayIndexName(name) {
collect(fv)
}
}
if o.elemMay {
collect(o.elem)
}
if o.wildMay {
collect(o.wild)
}
return nbrs
}
func (a *analysis) walkStringify(st *state, root allocRef) (taintSet, bool) {
// Preserve each field's allocation set as one group: its members are runtime
// alternatives, while separate fields and collected array elements must all
// serialize. Starting with leaf allocations and satisfying dependent groups
// computes the acyclic choices without mistaking a diamond for a cycle.
nodes := map[allocRef][]stringifySlot{}
stack := []allocRef{root}
for len(stack) > 0 {
if !a.alive() {
return nil, false
}
ref := stack[len(stack)-1]
stack = stack[:len(stack)-1]
if _, ok := nodes[ref]; ok {
continue
}
var values []stringifySlot
if o := st.heap[ref.id]; o != nil {
values = stringifyValues(o, ref)
}
nodes[ref] = values
for _, slot := range values {
for aid := range slot.value.allocs {
if _, seen := nodes[aid]; !seen {
stack = append(stack, aid)
}
}
}
}
type dependencyGroup struct {
owner allocRef
satisfied bool
}
var groups []dependencyGroup
pending := make(map[allocRef]int, len(nodes))
reverse := map[allocRef][]int{}
for ref, values := range nodes {
for _, slot := range values {
fv := slot.value
if slot.optional {
continue
}
if len(fv.allocs) == 0 || (!slot.collection && (!fv.allocOnly || len(fv.scalar) != 0)) {
continue
}
if slot.collection {
for aid := range fv.allocs {
group := len(groups)
groups = append(groups, dependencyGroup{owner: ref})
pending[ref]++
reverse[aid] = append(reverse[aid], group)
}
continue
}
group := len(groups)
groups = append(groups, dependencyGroup{owner: ref})
pending[ref]++
for aid := range fv.allocs {
reverse[aid] = append(reverse[aid], group)
}
}
}
serializable := make(map[allocRef]bool, len(nodes))
queue := make([]allocRef, 0, len(nodes))
for ref := range nodes {
if pending[ref] == 0 {
queue = append(queue, ref)
}
}
for len(queue) > 0 {
if !a.alive() {
return nil, false
}
ref := queue[len(queue)-1]
queue = queue[:len(queue)-1]
if serializable[ref] {
continue
}
serializable[ref] = true
for _, group := range reverse[ref] {
if groups[group].satisfied {
continue
}
groups[group].satisfied = true
owner := groups[group].owner
pending[owner]--
if pending[owner] == 0 {
queue = append(queue, owner)
}
}
}
if !serializable[root] {
return nil, true
}
var out taintSet
done := map[allocRef]bool{}
stack = append(stack, root)
for len(stack) > 0 {
if !a.alive() {
return out, false
}
ref := stack[len(stack)-1]
stack = stack[:len(stack)-1]
if done[ref] || !serializable[ref] {
continue
}
done[ref] = true
for _, slot := range nodes[ref] {
fv := slot.value
out = mergeTaint(out, fv.scalar)
for aid := range fv.allocs {
if serializable[aid] && !done[aid] {
stack = append(stack, aid)
}
}
}
}
return out, false
}
// serializeArray returns the scalar taint a join or array serialization carries.
// A self-referential element contributes no taint when revisited, but other
// tainted elements still propagate and the walk always terminates. It is
// iterative for the same stack-safety reason as serializeStringify.
func (a *analysis) serializeArray(st *state, v value) taintSet {
var out taintSet
done := map[allocRef]bool{}
stack := make([]allocRef, 0, len(v.allocs))
for ref := range v.allocs {
stack = append(stack, ref)
}
for len(stack) > 0 {
if !a.alive() {
return out
}
ref := stack[len(stack)-1]
stack = stack[:len(stack)-1]
if done[ref] {
continue
}
done[ref] = true
o := st.heap[ref.id]
if o == nil || !o.array {
continue
}
for _, aid := range arrayNeighbors(o, ref, func(ts taintSet) { out = mergeTaint(out, ts) }) {
if !done[aid] {
stack = append(stack, aid)
}
}
}
return out
}
// numberAllocSites assigns a deterministic site id to every object, array, new
// expression, and modeled createElement call in source order. The ids drive the
// two-class recency model.
func numberAllocSites(ast *js.AST) map[js.INode]int {
n := &siteNumberer{sites: map[js.INode]int{}}
js.Walk(n, ast)
return n.sites
}
type siteNumberer struct {
sites map[js.INode]int
next int
}
func (n *siteNumberer) Exit(js.INode) {}
func (n *siteNumberer) Enter(node js.INode) js.IVisitor {
if walkComputedClassName(n, node) {
return n
}
switch e := node.(type) {
case *js.ObjectExpr, *js.ArrayExpr, *js.NewExpr:
n.sites[node] = n.next
n.next++
case *js.CallExpr:
// A createElement call allocates a fresh element, so it needs its own site
// for the resource-element receiver model.
if prop, _, ok := memberAccess(ungroupExpr(e.X)); ok && prop == "createElement" {
n.sites[node] = n.next
n.next++
}
}
return n
}
package jstaint
import "github.com/tdewolff/parse/v2/js"
// funcInfo is a normalized view of a function expression, declaration, or arrow
// function. The two AST shapes (FuncDecl and ArrowFunc) share nothing but these
// fields for the analyzer's purposes.
type funcInfo struct {
params js.Params
body *js.BlockStmt
generator bool
async bool
}
// handlerSite is one discovered keyboard-event handler registration. eventVar is
// the canonical identity of the handler's first parameter, or nil when that
// parameter is not a plain identifier (a destructured or absent parameter).
type handlerSite struct {
fn *funcInfo
eventVar *js.Var
}
// discoverHandlers finds every statically resolvable keyboard-handler
// registration in the parse unit and returns one site per resolved function.
// When a handler identifier resolves to more than one function value, each is
// returned so their analyses can be unioned.
func discoverHandlers(ast *js.AST) []handlerSite {
funcs := collectFuncValues(ast)
var sites []handlerSite
seen := map[*funcInfo]bool{}
v := &handlerVisitor{funcs: funcs, add: func(fn *funcInfo) {
if seen[fn] {
return
}
seen[fn] = true
sites = append(sites, handlerSite{fn: fn, eventVar: firstParamVar(fn.params)})
}}
js.Walk(v, ast)
return sites
}
type handlerVisitor struct {
funcs map[*js.Var][]*funcInfo
add func(*funcInfo)
}
func (v *handlerVisitor) Exit(js.INode) {}
func (v *handlerVisitor) Enter(n js.INode) js.IVisitor {
if walkComputedClassName(v, n) {
return v
}
switch e := n.(type) {
case *js.BinaryExpr:
if e.Op == js.EqToken && isKeyHandlerProperty(e.X) {
v.emit(e.Y)
}
case *js.CallExpr:
if name, fn, ok := addEventListenerHandler(e); ok && isDOMEventName(name) {
v.emit(fn)
}
case *js.Property:
if e.Name != nil && !e.Name.IsComputed() && isReactHandlerProp(e.Name.String()) {
v.emit(e.Value)
}
}
return v
}
// emit resolves expr to zero or more concrete functions and records a handler
// site for each. A generator is never a handler.
func (v *handlerVisitor) emit(expr js.IExpr) {
for _, fn := range resolveFuncValues(expr, v.funcs) {
if !fn.generator {
v.add(fn)
}
}
}
// isKeyHandlerProperty reports whether target is a member access naming a DOM
// keyboard on* property, in either dot or static-bracket form.
func isKeyHandlerProperty(target js.IExpr) bool {
name, _, ok := memberAccess(target)
return ok && isDOMHandlerProp(name)
}
// DOM on* properties are lowercase; the DOM does not fire a camelCase
// element.onKeyDown assignment.
func isDOMHandlerProp(name string) bool {
switch name {
case "onkeydown", "onkeypress", "onkeyup":
return true
default:
return false
}
}
// DOM event type strings passed to addEventListener are case-sensitive
// lowercase.
func isDOMEventName(name string) bool {
switch name {
case "keydown", "keypress", "keyup":
return true
default:
return false
}
}
// React object-literal handler props are camelCase.
func isReactHandlerProp(name string) bool {
switch name {
case "onKeyDown", "onKeyPress", "onKeyUp":
return true
default:
return false
}
}
// addEventListenerHandler matches both the receiver form el.addEventListener and
// the bare unshadowed-global form, returning the event name and handler value.
func addEventListenerHandler(call *js.CallExpr) (string, js.IExpr, bool) {
if len(call.Args.List) < 2 {
return "", nil, false
}
if call.Args.List[0].Rest || call.Args.List[1].Rest {
return "", nil, false
}
switch callee := call.X.(type) {
case *js.DotExpr:
if name, ok := staticStringOrIdent(callee.Y); !ok || name != "addEventListener" {
return "", nil, false
}
case *js.Var:
if !isGlobalRef(callee, "addEventListener") {
return "", nil, false
}
default:
return "", nil, false
}
eventName, ok := staticStringOrIdent(ungroupExpr(call.Args.List[0].Value))
if !ok {
return "", nil, false
}
return eventName, call.Args.List[1].Value, true
}
func ungroupExpr(expr js.IExpr) js.IExpr {
for {
group, ok := expr.(*js.GroupExpr)
if !ok {
return expr
}
expr = group.X
}
}
// isGlobalRef reports whether v is an unshadowed reference to the named global
// binding. The parser leaves free identifiers undeclared, so a NoDecl canonical
// identity with the expected name is a global; a local declaration of the same
// name has a declaration type and is therefore not the platform global.
func isGlobalRef(v *js.Var, name string) bool {
n, ok := globalName(v)
return ok && n == name
}
func firstParamVar(params js.Params) *js.Var {
if len(params.List) == 0 {
return nil
}
if v, ok := params.List[0].Binding.(*js.Var); ok {
return canonicalVar(v)
}
return nil
}
// resolveFuncValues returns the concrete functions expr can denote: a direct,
// possibly parenthesized function or arrow literal, or a same-file identifier
// whose collected values are functions.
func resolveFuncValues(expr js.IExpr, funcs map[*js.Var][]*funcInfo) []*funcInfo {
expr = ungroupExpr(expr)
if fn := literalFuncValue(expr); fn != nil {
return []*funcInfo{fn}
}
switch e := expr.(type) {
case *js.Var:
return funcs[canonicalVar(e)]
default:
return nil
}
}
// collectFuncValues maps each canonical variable to the function values assigned
// to it anywhere in the file: named declarations, declaration initializers, and
// plain assignments. Multiple values accumulate so a later union covers every
// possibility.
func collectFuncValues(ast *js.AST) map[*js.Var][]*funcInfo {
c := &funcValueCollector{funcs: map[*js.Var][]*funcInfo{}}
js.Walk(c, ast)
return c.funcs
}
type funcValueCollector struct {
funcs map[*js.Var][]*funcInfo
}
func (c *funcValueCollector) Exit(js.INode) {}
func (c *funcValueCollector) Enter(n js.INode) js.IVisitor {
if walkComputedClassName(c, n) {
return c
}
switch e := n.(type) {
case *js.FuncDecl:
if e.Name != nil {
c.bind(e.Name, newFuncInfo(e.Params, &e.Body, e.Generator, e.Async))
}
case *js.VarDecl:
for i := range e.List {
be := e.List[i]
target, ok := be.Binding.(*js.Var)
if !ok || be.Default == nil {
continue
}
for _, fn := range literalFuncValues(be.Default) {
c.bind(target, fn)
}
}
case *js.BinaryExpr:
if e.Op != js.EqToken {
return c
}
if target, ok := e.X.(*js.Var); ok {
for _, fn := range literalFuncValues(e.Y) {
c.bind(target, fn)
}
}
}
return c
}
func (c *funcValueCollector) bind(v *js.Var, fn *funcInfo) {
key := canonicalVar(v)
c.funcs[key] = append(c.funcs[key], fn)
}
// literalFuncValues returns a function value only for a direct, possibly
// parenthesized function or arrow literal. It does not chase identifier aliases,
// so binding collection cannot recurse without bound.
func literalFuncValues(expr js.IExpr) []*funcInfo {
if fn := literalFuncValue(ungroupExpr(expr)); fn != nil {
return []*funcInfo{fn}
}
return nil
}
func literalFuncValue(expr js.IExpr) *funcInfo {
switch e := expr.(type) {
case *js.FuncDecl:
return newFuncInfo(e.Params, &e.Body, e.Generator, e.Async)
case *js.ArrowFunc:
return newFuncInfo(e.Params, &e.Body, false, e.Async)
default:
}
return nil
}
func newFuncInfo(params js.Params, body *js.BlockStmt, generator, async bool) *funcInfo {
return &funcInfo{params: params, body: body, generator: generator, async: async}
}
// collectFunctionLocals records bindings owned by one invocation. They must not
// escape into the caller state when an inline callee returns.
func collectFunctionLocals(body *js.BlockStmt) map[*js.Var]bool {
c := &functionLocalCollector{locals: map[*js.Var]bool{}}
js.Walk(c, body)
return c.locals
}
type functionLocalCollector struct {
locals map[*js.Var]bool
}
func (c *functionLocalCollector) Exit(js.INode) {}
func (c *functionLocalCollector) Enter(n js.INode) js.IVisitor {
switch x := n.(type) {
case *js.BlockStmt:
for _, v := range x.Declared {
c.locals[canonicalVar(v)] = true
}
case *js.FuncDecl, *js.ArrowFunc, *js.MethodDecl:
return nil
}
return c
}
package jstaint
import (
"math"
"math/big"
"strconv"
"strings"
"sync/atomic"
"github.com/tdewolff/parse/v2/js"
)
// allocID identifies one abstract allocation instance: a deterministic source
// site plus a recency class. The two classes per site are the current instance
// and a summary of older instances, so the analyzer never invents an unbounded
// runtime-object identity.
type allocID struct {
site int
summary bool
}
// allocRef carries the call depth of the path through which an allocation was
// obtained. Heap objects remain keyed by allocation identity.
type allocRef struct {
id allocID
minDepth uint8
advance uint8
}
type allocSet map[allocRef]bool
// value is an abstract value at a program point: scalar taint carried directly
// plus the set of allocation instances the value may reference. Network sinks
// consume only the scalar part, because a URL, body, or header must be a string;
// an object reaches a sink only after a modeled serializer turns its reachable
// field taint into scalar taint.
type value struct {
scalar taintSet
allocs allocSet
// allocOnly is true only when every represented runtime alternative is an
// allocation. A branch that may instead produce an untracked clean primitive
// clears it, preventing an object spread or cycle check from acting definite.
allocOnly bool
// scheme is the URL scheme this value definitely carries when used as a
// destination, independent of taint. It gates whether a tainted URL is a
// network sink.
scheme schemeState
}
// object is the abstract contents of one allocation: statically named fields, an
// array-element field, and a wildcard field for writes whose key is not
// statically known.
type object struct {
fields map[string]value
must map[string]bool
elem value
wild value
wildReq value
elemMay bool
wildMay bool
elemMust bool
wildMust bool
array bool
// Receiver provenance for network-sink method calls. kind names the platform
// object this allocation represents; a generic object is never a sink.
kind objectKind
// XMLHttpRequest path state: whether an open has been seen on this path, and
// the taint remembered from that open's URL and any setRequestHeader values.
xhrOpened bool
xhrURL taintSet
xhrHeader taintSet
xhrScheme schemeState
// WebSocket path state: possibly open once a later callback can observe it,
// definitely closed only when closed on every merged path.
wsMaybeOpen bool
wsClosed bool
// owner is the writer token of the single state allowed to mutate this
// object in place. Zero means frozen: every state must copy before writing.
// clone() leaves it zero; installers set it. It is bookkeeping, not
// semantic state, so objectEqual and mergeObject ignore it.
owner uint64
}
func (o *object) clone() *object {
n := &object{
elem: o.elem,
wild: o.wild,
wildReq: o.wildReq,
elemMay: o.elemMay,
wildMay: o.wildMay,
elemMust: o.elemMust,
wildMust: o.wildMust,
array: o.array,
kind: o.kind,
xhrOpened: o.xhrOpened,
xhrURL: o.xhrURL,
xhrHeader: o.xhrHeader,
xhrScheme: o.xhrScheme,
wsMaybeOpen: o.wsMaybeOpen,
wsClosed: o.wsClosed,
}
if len(o.fields) != 0 {
n.fields = make(map[string]value, len(o.fields))
for k, v := range o.fields {
n.fields[k] = v
}
}
if len(o.must) != 0 {
n.must = make(map[string]bool, len(o.must))
for k := range o.must {
n.must[k] = true
}
}
return n
}
// stateMutSeq issues globally unique writer tokens. Uniqueness across
// goroutines is all that matters; the values never reach output.
var stateMutSeq atomic.Uint64
func nextMut() uint64 { return stateMutSeq.Add(1) }
// state is the flow-sensitive abstract store at a program point: per-variable
// values, transient receiver captures, and the heap of allocation contents.
// Branch forks clone it; merges union it; loop fixed points compare it.
//
// Clones share the three maps copy-on-write, so clone is O(1) instead of a
// deep heap copy at every branch fork. The shared flags mark maps another
// state may still reach, and the per-object owner token marks objects this
// state may mutate in place. Every write must go through setEnv/delEnv,
// setCapture/delCapture, mutObject/installObject/shareObject/dropObject, or
// an own* barrier; a direct map write on a shared state corrupts its siblings.
type state struct {
env map[*js.Var]value
// captures holds transient receivers across computed keys, right-hand sides,
// and call arguments. It is part of state so allocation promotion rewrites
// captured identities on every branch before the operation uses its receiver.
captures map[js.INode]value
heap map[allocID]*object
// mut is this state's writer token. An object is mutable in place only
// when o.owner == s.mut. Cloning hands the maps to a second state, so both
// sides take fresh tokens and thereby abandon in-place rights on every
// previously owned object. Invariant: heapShared implies no object has
// owner == s.mut, because installing an owned object first unshares the
// map and cloning re-tokens s.
mut uint64
// envVer, capsVer, and heapVer identify the underlying map storage. Two
// states with equal versions hold the same map, so merge and equality
// short-circuit; own* stamps a fresh version when it copies.
envVer uint64
capsVer uint64
heapVer uint64
envShared bool
capsShared bool
heapShared bool
continues bool
}
func newState() *state {
return &state{
env: map[*js.Var]value{},
captures: map[js.INode]value{},
heap: map[allocID]*object{},
mut: nextMut(),
envVer: nextMut(),
capsVer: nextMut(),
heapVer: nextMut(),
continues: true,
}
}
func (s *state) clone() *state {
s.mut = nextMut()
s.envShared, s.capsShared, s.heapShared = true, true, true
return &state{
env: s.env, captures: s.captures, heap: s.heap,
mut: nextMut(),
envVer: s.envVer, capsVer: s.capsVer, heapVer: s.heapVer,
envShared: true, capsShared: true, heapShared: true,
continues: s.continues,
}
}
// replaceWith overwrites s with the contents of src, keeping a stable *state
// identity across a may-execute merge. It adopts src's storage and writer
// token, so src must be a temporary that is never used again.
func (s *state) replaceWith(src *state) {
s.env, s.captures, s.heap = src.env, src.captures, src.heap
s.mut = src.mut
s.envVer, s.capsVer, s.heapVer = src.envVer, src.capsVer, src.heapVer
s.envShared, s.capsShared, s.heapShared = src.envShared, src.capsShared, src.heapShared
s.continues = src.continues
}
func (s *state) ownEnv() {
if !s.envShared {
return
}
env := make(map[*js.Var]value, len(s.env))
for k, v := range s.env {
env[k] = v
}
s.env = env
s.envVer = nextMut()
s.envShared = false
}
func (s *state) ownCaptures() {
if !s.capsShared {
return
}
caps := make(map[js.INode]value, len(s.captures))
for k, v := range s.captures {
caps[k] = v
}
s.captures = caps
s.capsVer = nextMut()
s.capsShared = false
}
func (s *state) ownHeap() {
if !s.heapShared {
return
}
heap := make(map[allocID]*object, len(s.heap))
for k, o := range s.heap {
heap[k] = o
}
s.heap = heap
s.heapVer = nextMut()
s.heapShared = false
}
func (s *state) setEnv(cv *js.Var, v value) {
s.ownEnv()
s.env[cv] = v
}
func (s *state) delEnv(cv *js.Var) {
if _, ok := s.env[cv]; !ok {
return
}
s.ownEnv()
delete(s.env, cv)
}
func (s *state) setCapture(node js.INode, v value) {
s.ownCaptures()
s.captures[node] = v
}
func (s *state) delCapture(node js.INode) {
if _, ok := s.captures[node]; !ok {
return
}
s.ownCaptures()
delete(s.captures, node)
}
// mutObject returns the object at id with in-place write rights for s, copying
// it first when any other state may still reach it.
func (s *state) mutObject(id allocID) *object {
o := s.heap[id]
if o == nil {
return nil
}
if o.owner != s.mut {
o = o.clone()
o.owner = s.mut
s.ownHeap()
s.heap[id] = o
}
return o
}
// installObject publishes an object s built or merged exclusively for itself,
// granting in-place write rights.
func (s *state) installObject(id allocID, o *object) {
o.owner = s.mut
s.ownHeap()
s.heap[id] = o
}
// shareObject publishes an object owned elsewhere without copying. Freezing
// the owner makes every later writer, including s, copy first.
func (s *state) shareObject(id allocID, o *object) {
o.owner = 0
s.ownHeap()
s.heap[id] = o
}
func (s *state) dropObject(id allocID) {
if _, ok := s.heap[id]; !ok {
return
}
s.ownHeap()
delete(s.heap, id)
}
func mergeAllocs(a, b allocSet) allocSet {
if len(a) == 0 {
return b
}
if len(b) == 0 {
return a
}
out := make(allocSet, len(a)+len(b))
for k := range a {
out[k] = true
}
for k := range b {
out[k] = true
}
return out
}
func mergeValue(a, b value) value {
return value{
scalar: mergeTaint(a.scalar, b.scalar),
allocs: mergeAllocs(a.allocs, b.allocs),
allocOnly: a.allocOnly && b.allocOnly,
scheme: mergeScheme(a.scheme, b.scheme),
}
}
// widenAbsentValue adds an untracked non-allocation alternative to v. A value
// missing on one path cannot retain a definite allocation or URL scheme.
func widenAbsentValue(v value) value {
v.allocOnly = false
v.scheme = schemeState{}
return v
}
// advanceCallValue moves a value across one user-defined call edge. Facts past
// depth 1 are discarded. Allocation identities remain available at a blocked
// depth so the callee can still apply object mutations without propagating taint.
func advanceCallValue(v value) value {
v.scalar = applyTaintDepth(v.scalar, 0, 1)
if len(v.allocs) != 0 {
refs := make(allocSet, len(v.allocs))
for ref := range v.allocs {
if ref.advance < 2 {
ref.advance++
}
refs[ref] = true
}
v.allocs = refs
}
return v
}
// returnValueAtDepthOne materializes the lazy allocation constraints accumulated
// while a depth-1 callee used an argument, then records the return edge without
// counting the same call twice.
func returnValueAtDepthOne(v value) value {
v.scalar = applyTaintDepth(v.scalar, 1, 0)
if len(v.allocs) == 0 {
return v
}
refs := make(allocSet, len(v.allocs))
for ref := range v.allocs {
refs[composeAllocRef(ref, allocRef{minDepth: 1})] = true
}
v.allocs = refs
return v
}
func applyRefDepth(v value, parent allocRef) value {
v.scalar = applyTaintDepth(v.scalar, parent.minDepth, parent.advance)
if len(v.allocs) != 0 && (parent.minDepth != 0 || parent.advance != 0) {
refs := make(allocSet, len(v.allocs))
for ref := range v.allocs {
refs[composeAllocRef(ref, parent)] = true
}
v.allocs = refs
}
return v
}
// composeAllocRef applies parent after child. Each constraint represents
// max(depth, minDepth) + advance, so a later minimum is discounted by an
// advance the child has already applied.
func composeAllocRef(child, parent allocRef) allocRef {
requiredMin := uint8(0)
if parent.minDepth > child.advance {
requiredMin = parent.minDepth - child.advance
}
if child.minDepth < requiredMin {
child.minDepth = requiredMin
}
child.advance += parent.advance
if child.advance > 2 {
child.advance = 2
}
return child
}
func applyTaintDepth(ts taintSet, minDepth, advance uint8) taintSet {
if len(ts) == 0 || (minDepth == 0 && advance == 0) {
return ts
}
out := make(taintSet, len(ts))
for fact, chain := range ts {
if fact.callDepth < minDepth {
fact.callDepth = minDepth
}
fact.callDepth += advance
if fact.callDepth <= 1 {
if existing, ok := out[fact]; !ok || shorterChain(chain, existing) {
out[fact] = chain
}
}
}
return out
}
func mergePresentValue(current value, present bool, next value) (value, bool) {
if !present {
return next, true
}
return mergeValue(current, next), true
}
func mergeObject(a, b *object) *object {
n := &object{
elemMay: a.elemMay || b.elemMay,
wildMay: a.wildMay || b.wildMay,
elemMust: a.elemMust && b.elemMust,
wildMust: a.wildMust && b.wildMust,
array: a.array,
// The two operands are the same allocation site, so kind agrees; a fresh
// generic instance from one path defers to the provenance of the other.
kind: mergeKind(a.kind, b.kind),
xhrOpened: a.xhrOpened || b.xhrOpened,
xhrURL: mergeTaint(a.xhrURL, b.xhrURL),
xhrHeader: mergeTaint(a.xhrHeader, b.xhrHeader),
xhrScheme: mergeXHRScheme(a, b),
wsMaybeOpen: a.wsMaybeOpen || b.wsMaybeOpen,
wsClosed: a.wsClosed && b.wsClosed,
}
if a.elemMay {
n.elem = a.elem
}
if b.elemMay {
n.elem, _ = mergePresentValue(n.elem, a.elemMay, b.elem)
}
if a.wildMay {
n.wild = a.wild
}
if b.wildMay {
n.wild, _ = mergePresentValue(n.wild, a.wildMay, b.wild)
}
if n.wildMust {
n.wildReq = mergeValue(a.wildReq, b.wildReq)
}
n.fields = make(map[string]value, len(a.fields)+len(b.fields))
for k, v := range a.fields {
if bv, ok := b.fields[k]; ok {
n.fields[k] = mergeValue(v, bv)
} else if b.wildMay {
n.fields[k] = mergeValue(v, b.wild)
} else {
n.fields[k] = v
}
}
for k, v := range b.fields {
if _, ok := n.fields[k]; !ok {
if a.wildMay {
n.fields[k] = mergeValue(a.wild, v)
} else {
n.fields[k] = v
}
}
}
for k := range a.must {
if b.must[k] {
if n.must == nil {
n.must = map[string]bool{}
}
n.must[k] = true
}
}
return n
}
// mergeXHRScheme joins only paths with an active request. A path where open has
// not run cannot contribute a destination scheme because send cannot issue a
// network request there.
func mergeXHRScheme(a, b *object) schemeState {
switch {
case !a.xhrOpened:
return b.xhrScheme
case !b.xhrOpened:
return a.xhrScheme
default:
return mergeScheme(a.xhrScheme, b.xhrScheme)
}
}
func mergeState(a, b *state) *state {
if !a.continues && b.continues {
return b.clone()
}
if a.continues && !b.continues {
return a.clone()
}
n := &state{mut: nextMut(), continues: a.continues || b.continues}
// Equal versions mean both sides still hold the identical shared map, so
// its self-merge is itself and the result becomes one more holder.
if a.envVer == b.envVer {
n.env, n.envVer = a.env, a.envVer
n.envShared, a.envShared, b.envShared = true, true, true
} else {
n.env = make(map[*js.Var]value, len(a.env)+len(b.env))
n.envVer = nextMut()
for k, v := range a.env {
if bv, ok := b.env[k]; ok {
if merged := mergeValue(v, bv); storable(merged) {
n.env[k] = merged
}
} else {
v = widenAbsentValue(v)
if storable(v) {
n.env[k] = v
}
}
}
for k, v := range b.env {
if _, ok := a.env[k]; !ok {
v = widenAbsentValue(v)
if storable(v) {
n.env[k] = v
}
}
}
}
if a.capsVer == b.capsVer {
n.captures, n.capsVer = a.captures, a.capsVer
n.capsShared, a.capsShared, b.capsShared = true, true, true
} else {
n.captures = make(map[js.INode]value, len(a.captures)+len(b.captures))
n.capsVer = nextMut()
for node, v := range a.captures {
if bv, ok := b.captures[node]; ok {
n.captures[node] = mergeValue(v, bv)
} else {
n.captures[node] = widenAbsentValue(v)
}
}
for node, v := range b.captures {
if _, ok := a.captures[node]; !ok {
n.captures[node] = widenAbsentValue(v)
}
}
}
if a.heapVer == b.heapVer {
n.heap, n.heapVer = a.heap, a.heapVer
n.heapShared, a.heapShared, b.heapShared = true, true, true
return n
}
n.heap = make(map[allocID]*object, len(a.heap)+len(b.heap))
n.heapVer = nextMut()
// Objects present on one side, or pointer-identical on both, are shared
// into the merged state frozen; only genuinely diverged objects are merged
// into a fresh object the result owns.
for k, oa := range a.heap {
if ob, ok := b.heap[k]; ok {
if oa == ob {
oa.owner = 0
n.heap[k] = oa
} else {
m := mergeObject(oa, ob)
m.owner = n.mut
n.heap[k] = m
}
} else {
oa.owner = 0
n.heap[k] = oa
}
}
for k, ob := range b.heap {
if _, ok := a.heap[k]; !ok {
ob.owner = 0
n.heap[k] = ob
}
}
return n
}
func allocsEqual(a, b allocSet) bool {
if len(a) != len(b) {
return false
}
for k := range a {
if !b[k] {
return false
}
}
return true
}
func valueEqual(a, b value) bool {
return a.allocOnly == b.allocOnly && a.scheme == b.scheme &&
taintEqual(a.scalar, b.scalar) && allocsEqual(a.allocs, b.allocs)
}
func objectEqual(a, b *object) bool {
if a.array != b.array || a.elemMay != b.elemMay || a.wildMay != b.wildMay ||
a.elemMust != b.elemMust || a.wildMust != b.wildMust ||
a.kind != b.kind || a.xhrOpened != b.xhrOpened ||
a.wsMaybeOpen != b.wsMaybeOpen || a.wsClosed != b.wsClosed ||
a.xhrScheme != b.xhrScheme ||
!taintEqual(a.xhrURL, b.xhrURL) || !taintEqual(a.xhrHeader, b.xhrHeader) ||
len(a.must) != len(b.must) || !valueEqual(a.elem, b.elem) ||
!valueEqual(a.wildReq, b.wildReq) ||
!valueEqual(a.wild, b.wild) || len(a.fields) != len(b.fields) {
return false
}
for k := range a.must {
if !b.must[k] {
return false
}
}
for k, va := range a.fields {
vb, ok := b.fields[k]
if !ok || !valueEqual(va, vb) {
return false
}
}
return true
}
func stateEqual(a, b *state) bool {
if a.continues != b.continues || len(a.env) != len(b.env) ||
len(a.captures) != len(b.captures) || len(a.heap) != len(b.heap) {
return false
}
if a.envVer != b.envVer {
for k, va := range a.env {
vb, ok := b.env[k]
if !ok || !valueEqual(va, vb) {
return false
}
}
}
if a.capsVer != b.capsVer {
for node, va := range a.captures {
vb, ok := b.captures[node]
if !ok || !valueEqual(va, vb) {
return false
}
}
}
if a.heapVer != b.heapVer {
for k, oa := range a.heap {
ob, ok := b.heap[k]
if !ok {
return false
}
if oa == ob {
continue
}
if !objectEqual(oa, ob) {
return false
}
}
}
return true
}
// fieldKind selects which slot of an object a member access reaches.
type fieldKind uint8
const (
fieldNamed fieldKind = iota
fieldElem
fieldWild
)
type fieldKey struct {
kind fieldKind
name string
}
// fieldKeyOf classifies a member-access key expression. A static string or
// identifier names a field, a numeric literal index is an array element, and an
// unresolved computed key is the wildcard.
func fieldKeyOf(key js.IExpr) fieldKey {
key = ungroupExpr(key)
if name, ok := staticStringOrIdent(key); ok {
if isArrayIndexName(name) {
return fieldKey{kind: fieldElem, name: name}
}
return fieldKey{kind: fieldNamed, name: name}
}
if name, ok := numericPropertyNameOf(key); ok {
if isArrayIndexName(name) {
return fieldKey{kind: fieldElem, name: name}
}
return fieldKey{kind: fieldNamed, name: name}
}
return fieldKey{kind: fieldWild}
}
func numericPropertyNameOf(expr js.IExpr) (string, bool) {
negative := false
for {
u, ok := ungroupExpr(expr).(*js.UnaryExpr)
if !ok || (u.Op != js.PosToken && u.Op != js.NegToken) {
break
}
if u.Op == js.NegToken {
negative = !negative
}
expr = u.X
}
lit, ok := ungroupExpr(expr).(*js.LiteralExpr)
if !ok || (lit.TokenType != js.IntegerToken && lit.TokenType != js.DecimalToken) {
return "", false
}
name := numericPropertyName(lit)
if name == "" {
return "", false
}
if negative && name != "0" {
name = "-" + name
}
return name, true
}
func numericPropertyName(lit *js.LiteralExpr) string {
raw := strings.ReplaceAll(string(lit.Data), "_", "")
if strings.HasSuffix(raw, "n") {
integer := new(big.Int)
if _, ok := integer.SetString(strings.TrimSuffix(raw, "n"), 0); ok {
return integer.String()
}
return ""
}
if lit.TokenType == js.IntegerToken {
if n, err := strconv.ParseUint(raw, 0, 64); err == nil {
return jsNumberPropertyName(float64(n))
}
}
n, err := strconv.ParseFloat(raw, 64)
if err != nil || math.IsInf(n, 0) || math.IsNaN(n) {
return ""
}
return jsNumberPropertyName(n)
}
func jsNumberPropertyName(n float64) string {
if n == 0 {
return "0"
}
abs := math.Abs(n)
if abs >= 1e-6 && abs < 1e21 {
return strconv.FormatFloat(n, 'f', -1, 64)
}
s := strconv.FormatFloat(n, 'e', -1, 64)
parts := strings.SplitN(s, "e", 2)
exponent, err := strconv.Atoi(parts[1])
if err != nil {
return ""
}
return parts[0] + "e" + fmtSignedExponent(exponent)
}
func fmtSignedExponent(exponent int) string {
if exponent >= 0 {
return "+" + strconv.Itoa(exponent)
}
return strconv.Itoa(exponent)
}
// isArrayIndexName reports whether a string is a canonical JavaScript array
// index. Bracket access coerces both 0 and "0" to the same property key.
func isArrayIndexName(name string) bool {
if name == "" || (len(name) > 1 && name[0] == '0') {
return false
}
for i := range name {
if name[i] < '0' || name[i] > '9' {
return false
}
}
n, err := strconv.ParseUint(name, 10, 32)
return err == nil && n < 1<<32-1
}
// getField reads one field of an object. A static named read and a wildcard read
// both consume the wildcard, because a wildcard write may have set any property.
func getField(o *object, key fieldKey) (value, bool) {
switch key.kind {
case fieldNamed:
if v, ok := o.fields[key.name]; ok {
return v, o.must[key.name]
}
if o.wildMay {
return o.wild, false
}
return value{}, false
case fieldElem:
if key.name != "" {
if v, ok := o.fields[key.name]; ok {
return v, o.must[key.name]
}
}
var out value
have := false
if o.elemMay {
out, have = mergePresentValue(out, have, o.elem)
}
if o.wildMay {
out, _ = mergePresentValue(out, have, o.wild)
}
return out, false
default:
var out value
have := false
if o.elemMay {
out, have = mergePresentValue(out, have, o.elem)
}
if o.wildMay {
out, have = mergePresentValue(out, have, o.wild)
}
for _, fv := range o.fields {
out, have = mergePresentValue(out, have, fv)
}
return out, false
}
}
func (o *object) setNamed(name string, v value, strong bool) {
if o.fields == nil {
o.fields = map[string]value{}
}
if strong {
o.fields[name] = v
// A preceding wildcard write may have selected this exact property, so it
// is no longer certain that an unresolved field remains after overwrite.
o.wildMust = false
o.wildReq = value{}
if o.must == nil {
o.must = map[string]bool{}
}
o.must[name] = true
} else {
if o.wildMust {
// The weak write may select the required unresolved property. Preserve
// both its old value and the overwrite as runtime alternatives.
o.wildReq = mergeValue(o.wildReq, v)
}
old, ok := o.fields[name]
if !ok {
old = o.wild
}
o.fields[name] = mergeValue(old, v)
}
}
func (o *object) deleteNamed(name string) {
// The required unresolved property may be the deleted name. Other unresolved
// properties remain possible, but none remains certain after this delete.
o.wildMust = false
o.wildReq = value{}
if o.wildMay {
// Keep a clean tombstone so the deleted property does not expose an older
// unresolved write. A later wildcard write folds into the tombstone again.
if o.fields == nil {
o.fields = map[string]value{}
}
o.fields[name] = value{}
} else {
delete(o.fields, name)
}
delete(o.must, name)
}
func (o *object) weakElem(v value, definite bool) {
o.elem, o.elemMay = mergePresentValue(o.elem, o.elemMay, v)
if definite {
o.elemMust = true
}
}
// writeWild applies a write whose key is unresolved. Existing named fields must
// absorb it because the key may select any of them; a later strong named write
// can then replace that field without clearing the wildcard for other names.
func (o *object) writeWild(v value, definite bool) {
o.wild, o.wildMay = mergePresentValue(o.wild, o.wildMay, v)
if definite {
o.wildMust = true
o.wildReq = v
} else if o.wildMust {
o.wildReq = mergeValue(o.wildReq, v)
}
for k, fv := range o.fields {
o.fields[k] = mergeValue(fv, v)
}
}
package jstaint
import (
"strconv"
"github.com/tdewolff/parse/v2/js"
)
// rootSite is a reachable execution root: a callback whose body is analyzed with
// callback-published shared state. eventVar is the canonical keyboard-event
// parameter for a keyboard handler, or nil for a timer, listener, or socket
// callback whose parameters are not keystroke sources.
type rootSite struct {
fn *funcInfo
eventVar *js.Var
argCount int
}
type reachableFunction struct {
body *js.BlockStmt
params *js.Params
scanBody bool
scanBindings bool
defaultFrom int
defaultTo int
}
type reachableCall struct {
fn *funcInfo
provided int
}
// discoverReachableRoots returns the callback roots reachable from the module top
// level. Reachability follows callback registrations and direct same-file calls,
// but never descends into a nested function body until that function is itself
// registered or called. A function that is only declared is therefore not a root,
// so a sink reachable only inside it is never analyzed.
func (a *analysis) discoverReachableRoots(ast *js.AST) []rootSite {
var roots []rootSite
seenRoot := map[*funcInfo]int{}
bodyScanned := map[*funcInfo]bool{}
bindingsScanned := map[*funcInfo]bool{}
defaultFrom := map[*funcInfo]int{}
queue := []reachableFunction{{body: &ast.BlockStmt, scanBody: true}}
enqueue := func(fn *funcInfo, provided int) {
if fn.generator || fn.body == nil {
return
}
job := reachableFunction{body: fn.body, params: &fn.params}
if !bodyScanned[fn] {
bodyScanned[fn] = true
job.scanBody = true
}
if !bindingsScanned[fn] {
bindingsScanned[fn] = true
job.scanBindings = true
}
if provided < 0 {
provided = 0
}
if provided > len(fn.params.List) {
provided = len(fn.params.List)
}
previous, ok := defaultFrom[fn]
if !ok {
previous = len(fn.params.List)
}
if provided < previous {
job.defaultFrom = provided
job.defaultTo = previous
defaultFrom[fn] = provided
}
if job.scanBody || job.scanBindings || job.defaultFrom < job.defaultTo {
queue = append(queue, job)
}
}
for len(queue) > 0 {
current := queue[len(queue)-1]
queue = queue[:len(queue)-1]
regs, calls := a.scanFunctionLevel(current)
for _, r := range regs {
if index, ok := seenRoot[r.fn]; ok {
if roots[index].eventVar == nil && r.eventVar != nil {
roots[index].eventVar = r.eventVar
}
if r.argCount < roots[index].argCount {
roots[index].argCount = r.argCount
}
} else {
seenRoot[r.fn] = len(roots)
roots = append(roots, r)
}
enqueue(r.fn, r.argCount)
}
for _, call := range calls {
enqueue(call.fn, call.provided)
}
}
return roots
}
// scanFunctionLevel walks one function body's own code, collecting callback
// registrations and directly called same-file functions. It does not descend
// into nested function literals, whose bodies belong to their own reachability.
func (a *analysis) scanFunctionLevel(fn reachableFunction) ([]rootSite, []reachableCall) {
s := &funcLevelScanner{a: a}
if fn.params != nil {
if fn.scanBindings {
s.scanParamBindings(fn.params)
}
for i := fn.defaultFrom; i < fn.defaultTo; i++ {
js.Walk(s, fn.params.List[i].Default)
}
}
if fn.scanBody {
for i := range fn.body.List {
js.Walk(s, fn.body.List[i])
}
}
return s.regs, s.calls
}
type funcLevelScanner struct {
a *analysis
regs []rootSite
calls []reachableCall
}
func (s *funcLevelScanner) Exit(js.INode) {}
func (s *funcLevelScanner) Enter(n js.INode) js.IVisitor {
if !s.a.alive() {
return nil
}
if walkComputedClassName(s, n) {
return s
}
switch e := n.(type) {
case *js.FuncDecl:
// A nested function's body is reachable only when it is registered or
// called, which is decided at this level, not by lexical nesting.
return nil
case *js.ArrowFunc:
return nil
case *js.MethodDecl:
js.Walk(s, e.Name.Computed)
return nil
case *js.ClassDecl:
s.scanClassDefinition(e)
return nil
case *js.BinaryExpr:
if e.Op == js.EqToken {
s.registerAssignment(e)
}
case *js.CallExpr:
s.registerCall(e)
case *js.Property:
if e.Name != nil && !e.Name.IsComputed() && isReactHandlerProp(e.Name.String()) {
s.addRoots(e.Value, true, 1)
}
}
return s
}
// scanParamBindings visits computed keys and nested defaults, which can execute
// even when the invocation supplied the containing top-level argument.
func (s *funcLevelScanner) scanParamBindings(params *js.Params) {
for i := range params.List {
s.scanParamBinding(params.List[i].Binding)
}
s.scanParamBinding(params.Rest)
}
func (s *funcLevelScanner) scanParamBinding(binding js.IBinding) {
switch b := binding.(type) {
case *js.BindingArray:
for i := range b.List {
element := &b.List[i]
js.Walk(s, element.Default)
s.scanParamBinding(element.Binding)
}
s.scanParamBinding(b.Rest)
case *js.BindingObject:
for i := range b.List {
item := &b.List[i]
if item.Key != nil {
js.Walk(s, item.Key.Computed)
}
js.Walk(s, item.Value.Default)
s.scanParamBinding(item.Value.Binding)
}
s.scanParamBinding(b.Rest)
}
}
// scanClassDefinition visits only the expressions that run while a class is
// defined. Instance field initializers and method bodies do not execute until a
// construction or method call that version 1 does not model.
func (s *funcLevelScanner) scanClassDefinition(class *js.ClassDecl) {
js.Walk(s, class.Extends)
for i := range class.List {
item := &class.List[i]
if item.Method != nil {
js.Walk(s, item.Method.Name.Computed)
} else if item.StaticBlock == nil {
js.Walk(s, item.Name.Computed)
}
}
for i := range class.List {
item := &class.List[i]
if item.StaticBlock != nil {
js.Walk(s, item.StaticBlock)
} else if item.Method == nil && item.Static {
js.Walk(s, item.Init)
}
}
}
// registerAssignment records a callback assigned to a DOM keyboard on-property or
// to a static onopen/onmessage socket property.
func (s *funcLevelScanner) registerAssignment(e *js.BinaryExpr) {
name, _, ok := memberAccess(e.X)
if !ok {
return
}
switch {
case isDOMHandlerProp(name):
s.addRoots(e.Y, true, 1)
case isSocketCallbackProp(name):
s.addRoots(e.Y, false, 1)
}
}
func (s *funcLevelScanner) registerCall(call *js.CallExpr) {
if name, fn, ok := addEventListenerHandler(call); ok {
s.addRoots(fn, isDOMEventName(name), 1)
return
}
if fn, args, ok := scheduledCallback(call); ok {
s.addRoots(fn, false, args)
return
}
// A direct call to a same-file function makes that function reachable, so its
// own registrations count.
provided := definiteArgCount(call.Args.List)
for _, fn := range resolveFuncValues(ungroupExpr(call.X), s.a.funcs) {
s.calls = append(s.calls, reachableCall{fn: fn, provided: provided})
}
}
func definiteArgCount(args []js.Arg) int {
for i := range args {
if args[i].Rest {
return 0
}
}
return len(args)
}
// addRoots resolves an expression to its concrete callback functions and records
// each as a root. keyHandler marks whether the first parameter is a keystroke
// event.
func (s *funcLevelScanner) addRoots(expr js.IExpr, keyHandler bool, argCount int) {
for _, fn := range resolveFuncValues(expr, s.a.funcs) {
if fn.generator {
continue
}
var ev *js.Var
if keyHandler {
ev = firstParamVar(fn.params)
}
s.regs = append(s.regs, rootSite{fn: fn, eventVar: ev, argCount: argCount})
}
}
// isSocketCallbackProp reports whether a property assignment installs a socket
// event callback. Receiver provenance is refined later; here it only widens the
// reachable set, and findings still require a keystroke to reach a sink.
func isSocketCallbackProp(name string) bool {
switch name {
case "onopen", "onmessage":
return true
default:
return false
}
}
// scheduledCallback returns the function scheduled by a timer or microtask
// registration whose first argument is a callback.
func scheduledCallback(call *js.CallExpr) (js.IExpr, int, bool) {
name, ok := scheduledCallbackName(call.X)
if !ok {
return nil, 0, false
}
if len(call.Args.List) == 0 || call.Args.List[0].Rest {
return nil, 0, false
}
if name == "requestAnimationFrame" {
return call.Args.List[0].Value, 1, true
}
if name == "queueMicrotask" {
return call.Args.List[0].Value, 0, true
}
if len(call.Args.List) <= 2 {
return call.Args.List[0].Value, 0, true
}
return call.Args.List[0].Value, definiteArgCount(call.Args.List[2:]), true
}
func scheduledCallbackName(expr js.IExpr) (string, bool) {
switch callee := ungroupExpr(expr).(type) {
case *js.Var:
name, ok := globalName(callee)
return name, ok && isScheduledCallbackName(name)
case *js.DotExpr, *js.IndexExpr:
name, base, ok := memberAccess(expr)
return name, ok && isGlobalObject(base) && isScheduledCallbackName(name)
default:
return "", false
}
}
func isScheduledCallbackName(name string) bool {
switch name {
case "setTimeout", "setInterval", "queueMicrotask", "requestAnimationFrame":
return true
default:
return false
}
}
// sharedState projects a state onto its file-scope portion: the shared variables
// plus the heap reachable from them. This is the may-state that callbacks publish
// to and consume from one another.
func (a *analysis) sharedState(src *state) *state {
out := newState()
var roots []value
for cv, v := range src.env {
if a.isShared(cv) {
out.setEnv(cv, v)
roots = append(roots, v)
}
}
copyReachableHeap(src, out, roots)
return out
}
// publishShared merges a callback's exit shared writes into the global may-state
// and returns the updated may-state.
func (a *analysis) publishShared(global, src *state) *state {
out := global.clone()
var roots []value
for cv, v := range out.env {
if !a.isShared(cv) {
continue
}
if _, ok := src.env[cv]; ok {
continue
}
v = widenAbsentValue(v)
if storable(v) {
out.setEnv(cv, v)
} else {
out.delEnv(cv)
}
}
for cv, v := range src.env {
if !a.isShared(cv) {
continue
}
if ex, ok := out.env[cv]; ok {
merged := mergeValue(ex, v)
if storable(merged) {
out.setEnv(cv, merged)
} else {
out.delEnv(cv)
}
} else {
v = widenAbsentValue(v)
if storable(v) {
out.setEnv(cv, v)
}
}
roots = append(roots, v)
}
mergeReachableHeap(src, out, roots)
return out
}
func (a *analysis) isShared(cv *js.Var) bool {
return a.sharedVars[cv] || cv.Decl == js.NoDecl
}
// copyReachableHeap copies into dst every heap object reachable from the given
// values in src.
func copyReachableHeap(src, dst *state, roots []value) {
for _, id := range reachableAllocs(src, roots) {
if o := src.heap[id]; o != nil {
dst.shareObject(id, o)
}
}
}
// mergeReachableHeap merges into dst every heap object reachable from the given
// values in src, unioning with any object already present.
func mergeReachableHeap(src, dst *state, roots []value) {
for _, id := range reachableAllocs(src, roots) {
o := src.heap[id]
if o == nil {
continue
}
if ex, ok := dst.heap[id]; ok {
if ex != o {
dst.installObject(id, mergeObject(ex, o))
}
} else {
dst.shareObject(id, o)
}
}
}
// reachableAllocs returns every allocation reachable from roots through the heap.
func reachableAllocs(st *state, roots []value) []allocID {
seen := map[allocID]bool{}
var order []allocID
var stack []allocID
push := func(v value) {
for ref := range v.allocs {
id := ref.id
if !seen[id] {
seen[id] = true
order = append(order, id)
stack = append(stack, id)
}
}
}
for _, v := range roots {
push(v)
}
for len(stack) > 0 {
id := stack[len(stack)-1]
stack = stack[:len(stack)-1]
o := st.heap[id]
if o == nil {
continue
}
for _, fv := range o.fields {
push(fv)
}
push(o.elem)
push(o.wild)
push(o.wildReq)
}
return order
}
// evalUserCall inlines a same-file function call at depth 1. It returns the
// call's tainted return value and true when the callee is a modeled user
// function; otherwise it returns false so the caller falls back to built-in
// handling. A fact already at depth 1 cannot enter another callee.
func (a *analysis) evalUserCall(callee js.IExpr, args []value, st *state) (value, bool) {
if a.callDepth != 0 {
return value{}, false
}
fns := resolveFuncValues(callee, a.funcs)
if len(fns) == 0 {
return value{}, false
}
var ret value
haveRet := false
handled := false
unmodeledAlternative := false
base := st.clone()
effects := base.clone()
for _, fn := range fns {
if fn.generator || fn.body == nil || a.inProgress[fn] {
unmodeledAlternative = true
continue
}
handled = true
branch := base.clone()
branchRet := a.analyzeCallee(fn, args, branch)
if fn.async {
branchRet = value{}
}
ret, haveRet = mergePresentValue(ret, haveRet, branchRet)
effects = mergeState(effects, branch)
if !a.alive() {
break
}
}
if handled {
st.replaceWith(effects)
if unmodeledAlternative {
ret = mergeValue(ret, value{})
}
}
// A user call is not a scheme-preserving operation. Receiver allocations and
// scalar taint return, but destination classification becomes unknown.
ret.scheme = schemeState{}
return ret, handled
}
// analyzeCallee analyzes one callee body at depth 1 with its parameters bound to
// the argument values, then publishes its heap effects back to the caller and
// returns the union of its return values.
func (a *analysis) analyzeCallee(fn *funcInfo, args []value, st *state) value {
caller := st.clone()
sub := st.clone()
suspensionStart := len(a.suspensions)
previousDepth := a.callDepth
a.callDepth = 1
a.inProgress[fn] = true
a.retStack = append(a.retStack, value{})
a.returnExits = append(a.returnExits, nil)
a.bindParams(fn.params, args, sub)
sub = a.analyzeBlock(fn.body, sub)
exits := a.popReturnExits(sub)
a.retStack = a.retStack[:len(a.retStack)-1]
var ret value
haveRet := false
for i := range exits {
var exitRet value
if exits[i].returned {
exitRet = exits[i].ret
}
ret, haveRet = mergePresentValue(ret, haveRet, exitRet)
}
// The callee shares the caller's heap identities, so its object mutations and
// outer-variable writes flow back by unioning its exit state. Invocation-local
// bindings are removed so they cannot seed a later call.
effects := mergeExitStates(exits)
if fn.async && len(a.suspensions) > suspensionStart {
// The complete exit state is callback-published because the continuation
// can run later. A second pass stops each path at its first await and keeps
// no-await alternatives, which is the state visible when the promise returns.
a.suspensions = append(a.suspensions, effects)
effects = a.analyzeAsyncPrefix(fn, args, caller)
}
a.inProgress[fn] = false
a.callDepth = previousDepth
a.removeFunctionLocals(fn, effects)
out := mergeState(caller, effects)
st.replaceWith(out)
return returnValueAtDepthOne(ret)
}
func (a *analysis) analyzeAsyncPrefix(fn *funcInfo, args []value, caller *state) *state {
sub := caller.clone()
previousStop := a.stopAtAwait
a.stopAtAwait = true
a.retStack = append(a.retStack, value{})
a.returnExits = append(a.returnExits, nil)
a.bindParams(fn.params, args, sub)
sub = a.analyzeBlock(fn.body, sub)
exits := a.popReturnExits(sub)
a.retStack = a.retStack[:len(a.retStack)-1]
a.stopAtAwait = previousStop
return mergeExitStates(exits)
}
func mergeExitStates(exits []functionExit) *state {
if len(exits) == 0 {
return newState()
}
out := exits[0].state.clone()
for i := 1; i < len(exits); i++ {
out = mergeState(out, exits[i].state)
}
return out
}
func (a *analysis) removeFunctionLocals(fn *funcInfo, st *state) {
for cv := range a.functionLocals(fn) {
if !a.isShared(cv) {
st.delEnv(cv)
}
}
}
func (a *analysis) functionLocals(fn *funcInfo) map[*js.Var]bool {
if locals, ok := a.localCache[fn.body]; ok {
return locals
}
locals := collectFunctionLocals(fn.body)
a.localCache[fn.body] = locals
return locals
}
// bindParams applies positional, destructured, defaulted, and rest bindings as
// strong invocation-local updates. Rest arguments use a bounded synthetic array
// allocation so only explicit serialization turns their fields into a scalar.
func (a *analysis) bindParams(params js.Params, args []value, st *state) {
for i := range params.List {
p := ¶ms.List[i]
v, ok := p.Binding.(*js.Var)
if !ok {
switch {
case i < len(args):
a.bindParamPattern(p.Binding, advanceCallValue(args[i]), st)
case p.Default != nil:
a.bindParamPattern(p.Binding, a.evalExpr(p.Default, st), st)
default:
a.bindParamPattern(p.Binding, value{}, st)
}
continue
}
cv := canonicalVar(v)
switch {
case i < len(args):
a.bindVar(st, cv, advanceCallValue(args[i]))
case p.Default != nil:
a.bindVar(st, cv, a.evalExpr(p.Default, st))
default:
st.delEnv(cv)
}
}
if params.Rest != nil {
start := len(params.List)
if start > len(args) {
start = len(args)
}
a.bindRestParam(params.Rest, args[start:], st)
}
}
func (a *analysis) bindParamPattern(binding js.IBinding, arg value, st *state) {
switch b := binding.(type) {
case *js.Var:
a.bindVar(st, canonicalVar(b), arg)
case *js.BindingArray:
for i := range b.List {
be := &b.List[i]
if be.Binding == nil {
continue
}
key := fieldKey{kind: fieldElem, name: strconv.Itoa(i)}
a.bindParamElement(be, a.readField(st, arg, key), fieldDefinitelyPresent(st, arg, key), st)
}
if b.Rest != nil {
a.bindArrayRestParam(b.Rest, arg, len(b.List), st)
}
case *js.BindingObject:
excluded := make(map[string]bool, len(b.List))
for i := range b.List {
item := &b.List[i]
key := a.bindingObjectKey(item, st)
if key.kind != fieldWild {
excluded[key.name] = true
}
a.bindParamElement(
&item.Value,
a.readField(st, arg, key),
fieldDefinitelyPresent(st, arg, key),
st,
)
}
if b.Rest != nil {
a.bindObjectRestParam(b.Rest, arg, excluded, st)
}
}
}
func (a *analysis) bindingObjectKey(item *js.BindingObjectItem, st *state) fieldKey {
if item.Key != nil {
if item.Key.Computed != nil {
a.evalExpr(item.Key.Computed, st)
return fieldKeyOf(ungroupExpr(item.Key.Computed))
}
return fieldKeyOf(&item.Key.Literal)
}
if v, ok := item.Value.Binding.(*js.Var); ok {
return fieldKey{kind: fieldNamed, name: string(v.Name())}
}
return fieldKey{kind: fieldWild}
}
func (a *analysis) bindParamElement(be *js.BindingElement, arg value, present bool, st *state) {
if be.Binding == nil {
return
}
if be.Default == nil || present {
a.bindParamPattern(be.Binding, arg, st)
return
}
defaultSt := st.clone()
a.bindParamPattern(be.Binding, a.evalExpr(be.Default, defaultSt), defaultSt)
argSt := st.clone()
a.bindParamPattern(be.Binding, arg, argSt)
st.replaceWith(mergeState(defaultSt, argSt))
}
func (a *analysis) bindRestParam(binding js.IBinding, args []value, st *state) {
fresh := &object{array: true}
for i := range args {
fresh.setNamed(strconv.Itoa(i), advanceCallValue(args[i]), true)
a.fact()
}
a.bindRestObject(binding, fresh, st)
}
func (a *analysis) bindArrayRestParam(binding js.IBinding, arg value, start int, st *state) {
fresh := &object{array: true}
if len(arg.scalar) != 0 {
fresh.weakElem(value{scalar: arg.scalar}, false)
}
ustart := uint64(start) // #nosec G115 -- start is a parameter index, always non-negative
for ref := range arg.allocs {
o := st.heap[ref.id]
if o == nil || !o.array {
continue
}
for name, field := range o.fields {
if !isArrayIndexName(name) {
continue
}
index, err := strconv.ParseUint(name, 10, 32)
if err != nil || index < ustart {
continue
}
fresh.setNamed(strconv.FormatUint(index-ustart, 10), applyRefDepth(field, ref), false)
}
if o.elemMay {
fresh.weakElem(applyRefDepth(o.elem, ref), false)
}
if o.wildMay {
fresh.weakElem(applyRefDepth(o.wild, ref), false)
}
}
a.bindRestObject(binding, fresh, st)
}
func (a *analysis) bindObjectRestParam(binding js.IBinding, arg value, excluded map[string]bool, st *state) {
fresh := &object{}
if len(arg.scalar) != 0 {
fresh.writeWild(value{scalar: arg.scalar}, false)
}
for ref := range arg.allocs {
o := st.heap[ref.id]
if o == nil {
continue
}
for name, field := range o.fields {
if excluded[name] {
continue
}
fresh.setNamed(name, applyRefDepth(field, ref), false)
}
if o.elemMay {
fresh.weakElem(applyRefDepth(o.elem, ref), false)
}
if o.wildMay {
fresh.writeWild(applyRefDepth(o.wild, ref), false)
}
}
a.bindRestObject(binding, fresh, st)
}
func (a *analysis) bindRestObject(binding js.IBinding, fresh *object, st *state) {
site, ok := a.restSites[binding]
if !ok {
a.evalBindingPattern(binding, st)
return
}
a.promoteCurrent(st, site)
rest := a.installLiteral(st, site, a.loopDepth > 0, fresh)
a.bindParamPattern(binding, rest, st)
}
func numberRestSites(ast *js.AST) map[js.IBinding]int {
n := &restSiteNumberer{sites: map[js.IBinding]int{}, next: -1}
js.Walk(n, ast)
return n.sites
}
type restSiteNumberer struct {
sites map[js.IBinding]int
next int
}
func (n *restSiteNumberer) Exit(js.INode) {}
func (n *restSiteNumberer) Enter(node js.INode) js.IVisitor {
switch binding := node.(type) {
case *js.Params:
n.number(binding.Rest)
case *js.BindingArray:
n.number(binding.Rest)
case *js.BindingObject:
if binding.Rest != nil {
n.number(binding.Rest)
}
}
return n
}
func (n *restSiteNumberer) number(binding js.IBinding) {
if binding == nil {
return
}
if _, exists := n.sites[binding]; exists {
return
}
n.sites[binding] = n.next
n.next--
}
// Package jstaint reports keystroke values that reach a network sink in a
// JavaScript source file.
//
// It exists because the regex keylogger rule can only match a syntactically
// direct capture. A keylogger that moves the keystroke through variables needs
// variable identity to detect, which requires backreferences that neither
// YARA-X nor RE2 provides. See
// docs/superpowers/specs/2026-08-07-js-keystroke-taint-analyzer-design.md.
//
// The package is pure: bytes in, report out. It touches no filesystem, config,
// store, or process global, so callers own every I/O and persistence decision.
package jstaint
import (
"bytes"
"context"
"errors"
"fmt"
"strings"
"unicode"
"unicode/utf8"
"github.com/tdewolff/parse/v2"
"github.com/tdewolff/parse/v2/js"
)
// Status is the outcome of an analysis attempt. Callers must not infer a clean
// file from an empty result slice: only StatusAnalyzed means the content was
// examined end to end. StatusNotCandidate excludes unsupported documents and
// content without flow tokens; all remaining statuses are coverage gaps.
type Status uint8
const (
// StatusNotCandidate means the content lacks the required flow tokens or
// is a recognized non-JavaScript document. Embedded scripts are not examined.
StatusNotCandidate Status = iota
// StatusAnalyzed means the content was parsed and examined to completion.
StatusAnalyzed
StatusOversize
StatusParseError
StatusResourceLimit
StatusCanceled
StatusPanic
)
// String names the status for metrics labels and operator-facing text.
func (s Status) String() string {
switch s {
case StatusNotCandidate:
return "not_candidate"
case StatusAnalyzed:
return "analyzed"
case StatusOversize:
return "oversize"
case StatusParseError:
return "parse_error"
case StatusResourceLimit:
return "resource_limit"
case StatusCanceled:
return "canceled"
case StatusPanic:
return "panic"
}
return "unknown"
}
// MaxSourceBytes bounds the complete source an analysis accepts. Parsing
// allocates on the order of 18x the input, so this also bounds transient
// memory. Callers should read one byte past it to tell an exact-limit file
// from a truncated prefix.
const MaxSourceBytes = 2 << 20
// MaxReasonBytes bounds Report.Reason. Parser diagnostics can contain
// attacker-controlled text, so reports retain only sanitized, bounded context.
const MaxReasonBytes = 256
const (
// maxAnalysisDepth must cover real minified bundles, whose parse trees
// reach depth 905, while staying below the parser's pinned 1000 nesting
// limit so a parseable fixture can still distinguish this limit from a
// parser error.
maxAnalysisDepth = 950
maxPropagatedFacts = 200_000
)
type analysisLimitError uint8
const (
errAnalysisDepthLimit analysisLimitError = iota
errFactLimit
errNodeLimit
)
func (e analysisLimitError) Error() string {
switch e {
case errAnalysisDepthLimit:
return "maximum AST recursion depth exceeded"
case errNodeLimit:
return "maximum AST node count exceeded"
default:
return "maximum propagated fact count exceeded"
}
}
// Result is one keystroke-to-sink flow.
type Result struct {
// Source is the keyboard property read, such as "e.which".
Source string
// Via lists the canonical names the value passed through, in order.
Via []string
// Sink names the network operation that received the value.
Sink string
}
// Report is the outcome of analysing one source file.
type Report struct {
Status Status
// Results is non-empty only when Status is StatusAnalyzed.
Results []Result
// TotalResults counts every distinct flow before evidence truncation. It is
// non-zero only when Status is StatusAnalyzed.
TotalResults int
// Reason carries bounded diagnostic context for a non-analyzed status.
Reason string
// EvidenceTruncated reports that a flow or display segment exceeded an
// evidence limit.
EvidenceTruncated bool
}
// Analyze examines src for keystroke data reaching a network sink.
//
// It never panics: a panic anywhere inside is converted to StatusPanic so one
// malformed input cannot take down a scan.
func Analyze(ctx context.Context, src []byte) Report {
return analyzeWithPass(ctx, src, taintPass)
}
type analysisPass func(context.Context, *js.AST, *resourceBudget) (results []Result, total int, truncated bool, err error)
func analyzeWithPass(ctx context.Context, src []byte, pass analysisPass) (report Report) {
defer func() {
if r := recover(); r != nil {
report = Report{Status: StatusPanic, Reason: "recovered panic during analysis"}
}
report = finalizeReport(report)
}()
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: cancellationReason(err)}
}
// The size gate runs before the pre-filter so an oversize file is reported
// as a coverage gap even when its first bytes look uninteresting. Deciding
// "not a candidate" from a prefix would let padding hide a payload.
if len(src) > MaxSourceBytes {
return Report{Status: StatusOversize}
}
if !isCandidate(src) {
return Report{Status: StatusNotCandidate}
}
ast, err := js.Parse(parse.NewInputBytes(src), js.Options{})
if err != nil {
// A document marker inside a JavaScript literal or comment must never
// suppress analysis. Classify other formats only after JS parsing fails.
if isNonJSDocument(src) {
return Report{Status: StatusNotCandidate}
}
return Report{Status: StatusParseError, Reason: parseFailureContext(err)}
}
results, total, evidenceTruncated, err := pass(ctx, ast, &resourceBudget{})
if err != nil {
status := analysisErrorStatus(err)
return Report{
Status: status,
Results: results,
TotalResults: total,
Reason: analysisErrorReason(status, err),
EvidenceTruncated: evidenceTruncated,
}
}
return Report{
Status: StatusAnalyzed,
Results: results,
TotalResults: total,
EvidenceTruncated: evidenceTruncated,
}
}
// MayBeJSSource reports whether a prefix could be JavaScript source at all.
// It exists for files too large to analyze: the deep walk hands every
// readable file to this analyzer, and without a content check each oversize
// one became "JavaScript we failed to examine" -- 118,688 claimed skips over
// eight weeks on a live host, whose examples were .jpg, .png, .zip and
// .mmdb.
//
// The test is deliberately "is this source at all" rather than "is this a
// candidate". Analyze runs its size gate ahead of isCandidate on purpose, so
// padding cannot hide a payload behind an uninteresting prefix; deciding
// candidacy from a prefix here would reintroduce exactly that. NUL suggests
// binary content but is legal inside JS literals and comments. Only reject
// binary bytes the lexer encounters outside those tokens; incomplete tokens
// and ambiguous syntax must keep their coverage gap.
func MayBeJSSource(prefix []byte) bool {
if bytes.IndexByte(prefix, 0) < 0 {
return true
}
// The parser input appends a sentinel. Cap the slice so it cannot write
// into a caller's source beyond the peek, including concurrent readers.
input := parse.NewInputBytes(prefix[:len(prefix):len(prefix)])
lexer := js.NewLexer(input)
for {
token, data := lexer.Next()
switch token {
case js.ErrorToken:
// EOF can split a literal, comment or UTF-8 sequence. Other
// syntax errors are not proof of binary content either.
if input.Err() != nil {
return true
}
if !utf8.Valid(data) {
return false
}
return len(data) != 1 || (data[0] >= 0x20 && data[0] != 0x7f)
case js.DivToken, js.DivEqToken:
// Distinguishing division from a regexp needs parser context.
// Either may lead to a literal containing binary bytes.
return true
}
}
}
// isCandidate reports whether src carries both a key-handler token and a sink
// token. Matching is ASCII-case-insensitive so React's onKeyDown and a quoted
// "KeyDown" event name are admitted alongside the lowercase spellings.
func isCandidate(src []byte) bool {
return containsAnyASCIIFold(src, "keydown", "keypress", "keyup") &&
containsAnyASCIIFold(src, "fetch", "sendbeacon", "send", "open", "src", "websocket")
}
func containsAnyASCIIFold(src []byte, tokens ...string) bool {
for i, c := range src {
c = asciiLower(c)
for _, token := range tokens {
if c != token[0] || len(src)-i < len(token) {
continue
}
matched := true
for j := 1; j < len(token); j++ {
if asciiLower(src[i+j]) != token[j] {
matched = false
break
}
}
if matched {
return true
}
}
}
return false
}
func asciiLower(c byte) byte {
if c >= 'A' && c <= 'Z' {
return c + 'a' - 'A'
}
return c
}
type resourceBudget struct {
depth int
facts int
nodes int
}
func (b *resourceBudget) enterAST() error {
if b.depth >= maxAnalysisDepth {
return errAnalysisDepthLimit
}
b.depth++
return nil
}
func (b *resourceBudget) leaveAST() {
b.depth--
}
func (b *resourceBudget) addFact() error {
if b.facts >= maxPropagatedFacts {
return errFactLimit
}
b.facts++
return nil
}
type limitVisitor struct {
ctx context.Context
budget *resourceBudget
err error
}
func (v *limitVisitor) Enter(node js.INode) js.IVisitor {
if v.err != nil {
return nil
}
if err := v.ctx.Err(); err != nil {
v.err = err
return nil
}
if err := v.budget.enterAST(); err != nil {
v.err = err
return nil
}
// Count each parser node once, here, so the taint pass need not recount as it
// revisits basic blocks during fixed-point iteration.
v.budget.nodes++
if v.budget.nodes > maxASTNodes {
v.err = errNodeLimit
v.budget.leaveAST()
return nil
}
walkComputedClassName(v, node)
return v
}
func (v *limitVisitor) Exit(js.INode) {
v.budget.leaveAST()
}
func analysisErrorStatus(err error) Status {
switch {
case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded):
return StatusCanceled
case errors.Is(err, errAnalysisDepthLimit), errors.Is(err, errFactLimit), errors.Is(err, errNodeLimit):
return StatusResourceLimit
default:
panic(fmt.Sprintf("unexpected analyzer error: %v", err))
}
}
func analysisErrorReason(status Status, err error) string {
if status == StatusCanceled {
return cancellationReason(err)
}
return err.Error()
}
func cancellationReason(err error) string {
if errors.Is(err, context.DeadlineExceeded) {
return context.DeadlineExceeded.Error()
}
return context.Canceled.Error()
}
func parseFailureContext(err error) string {
var parseErr *parse.Error
if errors.As(err, &parseErr) {
return fmt.Sprintf("invalid JavaScript at line %d, column %d", parseErr.Line, parseErr.Column)
}
return "invalid JavaScript"
}
func finalizeReport(report Report) Report {
if report.Status == StatusAnalyzed {
report.Reason = ""
return report
}
report.Results = nil
report.TotalResults = 0
report.EvidenceTruncated = false
detail := sanitizeReason(report.Reason)
report.Reason = report.Status.String()
if detail != "" {
report.Reason += ": " + detail
}
report.Reason = boundReason(report.Reason)
return report
}
func sanitizeReason(reason string) string {
reason = strings.ToValidUTF8(reason, "?")
reason = strings.Map(func(r rune) rune {
if unicode.IsControl(r) {
return ' '
}
return r
}, reason)
return strings.Join(strings.Fields(reason), " ")
}
func boundReason(reason string) string {
reason = sanitizeReason(reason)
if len(reason) <= MaxReasonBytes {
return reason
}
cut := MaxReasonBytes - 3
for cut > 0 && !utf8.RuneStart(reason[cut]) {
cut--
}
return reason[:cut] + "..."
}
package jstaint
import "github.com/tdewolff/parse/v2/js"
// objectKind names the platform object an allocation represents. Receiver
// provenance is what separates a real network sink from a generic application
// object that merely has a method named send, open, or src.
type objectKind uint8
const (
kindGeneric objectKind = iota
kindXHR
kindWebSocket
kindResource // Image or a typed resource element from createElement
)
func mergeKind(a, b objectKind) objectKind {
if a != kindGeneric {
return a
}
return b
}
// Provenance sink display strings.
const (
sinkXHRBody = "XMLHttpRequest.send body argument"
sinkXHRURL = "XMLHttpRequest.open url argument"
sinkXHRHeader = "XMLHttpRequest.setRequestHeader value argument"
sinkWSURL = "WebSocket url argument"
sinkWSProtocol = "WebSocket protocols argument"
sinkWSSend = "WebSocket.send argument"
sinkResourceSrc = "resource element src assignment"
)
// schemeState is the URL-scheme lattice element. The zero value is the unknown,
// possibly-networked scheme (a relative or dynamic URL). A set scheme names a
// single definite scheme; a merge of two different schemes returns to unknown.
type schemeState struct {
set bool
name string
}
func mergeScheme(a, b schemeState) schemeState {
if a == b {
return a
}
return schemeState{}
}
// isNonNetworkScheme reports whether a definite scheme never performs a network
// fetch, so a captured key placed after it is not exfiltrated.
func (s schemeState) isNonNetwork() bool {
if !s.set {
return false
}
switch s.name {
case "about", "blob", "data", "file", "javascript":
return true
default:
return false
}
}
// schemeOfLiteral extracts the definite scheme of a string literal's value. A URL
// scheme is an ASCII letter followed by letters, digits, +, -, or ., ending at
// the first colon. Anything else (a relative path, a dynamic prefix) is unknown.
func schemeOfLiteral(expr js.IExpr) schemeState {
lit, ok := expr.(*js.LiteralExpr)
if !ok || lit.TokenType != js.StringToken || len(lit.Data) < 2 {
return schemeState{}
}
q := lit.Data[0]
if (q != '\'' && q != '"') || lit.Data[len(lit.Data)-1] != q {
return schemeState{}
}
return schemeOfBytes(lit.Data[1 : len(lit.Data)-1])
}
func schemeOfBytes(b []byte) schemeState {
if len(b) == 0 || !isASCIILetter(b[0]) {
return schemeState{}
}
for i := 0; i < len(b); i++ {
c := b[i]
if c == ':' {
if i == 0 {
return schemeState{}
}
return schemeState{set: true, name: asciiFoldString(b[:i])}
}
if !isSchemeChar(c) {
return schemeState{}
}
}
return schemeState{}
}
func isASCIILetter(c byte) bool {
return (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z')
}
func isSchemeChar(c byte) bool {
return isASCIILetter(c) || (c >= '0' && c <= '9') || c == '+' || c == '-' || c == '.'
}
func asciiFoldString(b []byte) string {
out := make([]byte, len(b))
for i := range b {
out[i] = asciiLower(b[i])
}
return string(out)
}
// templateScheme reads the definite scheme from a template literal's leading
// text, or from its first substitution when no literal prefix precedes it.
func templateScheme(x *js.TemplateExpr, first schemeState) schemeState {
if x.Tag != nil {
return schemeState{}
}
if len(x.List) == 0 {
return schemeOfBytes(templateCooked(x.Tail))
}
leading := templateCooked(x.List[0].Value)
if scheme := schemeOfBytes(leading); scheme.set {
return scheme
}
if len(leading) == 0 {
return first
}
return schemeState{}
}
// templateCooked strips the surrounding backtick or `}`/`${` delimiters from a
// template part's raw bytes so only the literal text remains.
func templateCooked(raw []byte) []byte {
if len(raw) > 0 && (raw[0] == '`' || raw[0] == '}') {
raw = raw[1:]
}
if n := len(raw); n >= 2 && raw[n-2] == '$' && raw[n-1] == '{' {
raw = raw[:n-2]
} else if n >= 1 && raw[n-1] == '`' {
raw = raw[:n-1]
}
return raw
}
// resourceElementTag reports whether a createElement tag names an element that
// fetches its src attribute, compared ASCII case-insensitively.
func resourceElementTag(name string) bool {
switch asciiFoldString([]byte(name)) {
case "audio", "embed", "iframe", "img", "script", "source", "track", "video":
return true
default:
return false
}
}
// allocateKinded installs a fresh current platform-object allocation. Recency is
// retained inside loops because constructor-local protocol state must not inherit
// from older instances; promoteCurrent still bounds each site to two identities.
func (a *analysis) allocateKinded(st *state, site int, kind objectKind) value {
a.promoteCurrent(st, site)
fresh := &object{kind: kind}
return a.installLiteral(st, site, false, fresh)
}
// evalNewSink handles new-expression receivers with network provenance. It
// returns the constructed value and true when it recognized one; the caller
// evaluates arguments before calling this.
func (a *analysis) evalNewSink(x *js.NewExpr, args []value, site int, st *state) (value, bool) {
switch {
case a.isGlobalCallee(x.X, "XMLHttpRequest"):
return a.allocateKinded(st, site, kindXHR), true
case a.isGlobalCallee(x.X, "Image"):
return a.allocateKinded(st, site, kindResource), true
case a.isGlobalCallee(x.X, "WebSocket"):
if len(args) == 0 {
return value{}, true
}
a.recordURL(args[0], x, 0, sinkWSURL)
if args[0].scheme.isNonNetwork() {
return value{}, true
}
if len(args) >= 2 {
a.record(argScalar(args, 1), x, 1, sinkWSProtocol)
a.record(a.serializeArray(st, args[1]), x, 1, sinkWSProtocol)
}
return a.allocateKinded(st, site, kindWebSocket), true
}
return value{}, false
}
// createElementValue returns a typed resource-element allocation when the call is
// document.createElement with a static resource tag, and true when it applied.
func (a *analysis) createElementValue(call *js.CallExpr, st *state) (value, bool) {
prop, base, ok := memberAccess(ungroupExpr(call.X))
if !ok || prop != "createElement" || !a.isDocument(base) {
return value{}, false
}
if len(call.Args.List) == 0 {
return value{}, false
}
tag, ok := staticStringOrIdent(ungroupExpr(call.Args.List[0].Value))
if !ok || !resourceElementTag(tag) {
return value{}, false
}
site, ok := a.sites[call]
if !ok {
return value{}, false
}
return a.allocateKinded(st, site, kindResource), true
}
// isDocument reports whether expr is the unshadowed global document object.
func (a *analysis) isDocument(expr js.IExpr) bool {
switch v := ungroupExpr(expr).(type) {
case *js.Var:
return isGlobalRef(v, "document")
case *js.DotExpr, *js.IndexExpr:
prop, base, ok := memberAccess(expr)
return ok && prop == "document" && isGlobalObject(base)
default:
return false
}
}
// xhrOrWSMethod applies a method call on an XHR or WebSocket receiver and records
// any sink. It returns true when the receiver's provenance claimed the call.
func (a *analysis) xhrOrWSMethod(call *js.CallExpr, recv value, args []value, st *state) bool {
prop, _, ok := memberAccess(ungroupExpr(call.X))
if !ok {
return false
}
ids := xhrWSTargets(st, recv)
if len(ids) == 0 {
return false
}
strong := recv.allocOnly && len(uniqueAllocIDs(recv.allocs)) == 1 && soleCurrent(recv.allocs)
handled := false
for id := range ids {
o := st.mutObject(id)
switch o.kind {
case kindXHR:
handled = a.applyXHRMethod(st, o, prop, call, args, recv, id, strong && !id.summary) || handled
case kindWebSocket:
handled = a.applyWSMethod(o, prop, call, args, strong && !id.summary) || handled
}
}
return handled
}
// xhrWSTargets returns the receiver allocations that carry XHR or WebSocket
// provenance.
func xhrWSTargets(st *state, recv value) map[allocID]bool {
var ids map[allocID]bool
for ref := range recv.allocs {
o := st.heap[ref.id]
if o == nil || (o.kind != kindXHR && o.kind != kindWebSocket) {
continue
}
if ids == nil {
ids = map[allocID]bool{}
}
ids[ref.id] = true
}
return ids
}
func (a *analysis) applyXHRMethod(
st *state,
o *object,
prop string,
call *js.CallExpr,
args []value,
recv value,
id allocID,
strong bool,
) bool {
switch prop {
case "open":
if len(args) < 2 {
return true
}
wasOpened := o.xhrOpened
// A later open reinitializes the request. Only a lone current receiver can
// prove the reset clears an earlier path's remembered URL and headers.
o.xhrOpened = true
switch {
case strong:
o.xhrURL = argScalar(args, 1)
o.xhrHeader = nil
o.xhrScheme = args[1].scheme
case !wasOpened:
o.xhrURL = argScalar(args, 1)
o.xhrScheme = args[1].scheme
default:
o.xhrURL = mergeTaint(o.xhrURL, argScalar(args, 1))
o.xhrScheme = mergeScheme(o.xhrScheme, args[1].scheme)
}
a.resetReceiverDepth(st, id, argScalar(args, 1))
return true
case "setRequestHeader":
if o.xhrOpened {
o.xhrHeader = mergeTaint(o.xhrHeader, argScalar(args, 1))
a.resetReceiverDepth(st, id, argScalar(args, 1))
}
return true
case "send":
if o.xhrOpened && !o.xhrScheme.isNonNetwork() {
a.record(argScalar(args, 0), call, 0, sinkXHRBody)
a.record(receiverTaint(recv, id, o.xhrURL), call, 1, sinkXHRURL)
a.record(receiverTaint(recv, id, o.xhrHeader), call, 2, sinkXHRHeader)
}
return true
case "abort":
if strong {
o.xhrOpened = false
o.xhrURL = nil
o.xhrHeader = nil
o.xhrScheme = schemeState{}
}
return true
}
return false
}
func (a *analysis) resetReceiverDepth(st *state, id allocID, ts taintSet) {
if taintCarriesDepth(ts, a.callDepth) {
resetAllocRefDepth(st, map[allocID]bool{id: true})
}
}
// receiverTaint applies every call-depth constraint carried by receiver aliases
// to state stored on one allocation.
func receiverTaint(recv value, id allocID, ts taintSet) taintSet {
var out taintSet
for ref := range recv.allocs {
if ref.id == id {
out = mergeTaint(out, applyTaintDepth(ts, ref.minDepth, ref.advance))
}
}
return out
}
func (a *analysis) applyWSMethod(o *object, prop string, call *js.CallExpr, args []value, strong bool) bool {
switch prop {
case "send":
if o.wsMaybeOpen && !o.wsClosed {
a.record(argScalar(args, 0), call, 0, sinkWSSend)
}
return true
case "close":
if strong {
o.wsClosed = true
}
return true
}
return false
}
// markSocketsObservable marks every WebSocket allocation in a shared state as
// possibly open. A socket published to file-scope state can be observed open by a
// later callback, which is exactly when a send on it becomes a network sink.
func markSocketsObservable(st *state) {
var hit []allocID
for id, o := range st.heap {
if o.kind == kindWebSocket && !o.wsMaybeOpen {
hit = append(hit, id)
}
}
for _, id := range hit {
st.mutObject(id).wsMaybeOpen = true
}
}
// resourceSrcSink records a src-assignment sink when the receiver is a typed
// resource element, gated by the assigned value's URL scheme.
func (a *analysis) resourceSrcSink(recv value, key fieldKey, rhs value, node js.INode, st *state) {
if key.kind != fieldNamed || key.name != "src" {
return
}
for ref := range recv.allocs {
if o := st.heap[ref.id]; o != nil && o.kind == kindResource {
a.recordURL(rhs, node, 0, sinkResourceSrc)
return
}
}
}
// recordURL records a destination-URL sink unless the destination is proven to
// use a non-network scheme.
func (a *analysis) recordURL(dest value, sink js.INode, arg int, text string) {
if dest.scheme.isNonNetwork() {
return
}
a.record(dest.scalar, sink, arg, text)
}
// recordBody records a body, data, or header sink carried alongside a
// destination. A proven non-network destination controls the whole operation, so
// its body is not exfiltrated either.
func (a *analysis) recordBody(dest, payload value, sink js.INode, arg int, text string) {
if dest.scheme.isNonNetwork() {
return
}
a.record(payload.scalar, sink, arg, text)
}
package jstaint
import "github.com/tdewolff/parse/v2/js"
// Sink display strings. They are stable because they appear in deterministic
// evidence output.
const (
sinkFetchURL = "fetch url argument"
sinkFetchBody = "fetch body option"
sinkFetchReferrer = "fetch referrer option"
sinkBeaconURL = "navigator.sendBeacon url argument"
sinkBeaconData = "navigator.sendBeacon data argument"
)
// evalCall evaluates a call's receiver and arguments once, records any network
// sink the call represents, and returns the value the call propagates.
func (a *analysis) evalCall(call *js.CallExpr, st *state) value {
callee := ungroupExpr(call.X)
recv := a.evalCallCallee(callee, st)
st.setCapture(call, recv)
isFetch := a.isGlobalCallee(callee, "fetch")
isBeacon := a.isBeaconCallee(callee)
argsSt := st
argsMayBeSkipped := call.Optional || optionalChainMaySkip(call.X)
if argsMayBeSkipped {
argsSt = st.clone()
}
args := make([]value, len(call.Args.List))
for i := range call.Args.List {
if isFetch && i == 1 {
a.evalFetchInit(call, args[0], call.Args.List[i].Value, argsSt)
continue
}
args[i] = a.evalExpr(call.Args.List[i].Value, argsSt)
}
if argsMayBeSkipped {
st.replaceWith(mergeState(st, argsSt))
}
recv = st.captures[call]
st.delCapture(call)
a.checkCallSink(call, args, isFetch, isBeacon)
if a.xhrOrWSMethod(call, recv, args, st) {
return value{}
}
if v, ok := a.createElementValue(call, st); ok {
return v
}
if ret, ok := a.evalUserCall(callee, args, st); ok {
return ret
}
return a.evalCallReturn(st, callee, args, recv)
}
func optionalChainMaySkip(expr js.IExpr) bool {
switch x := expr.(type) {
case *js.DotExpr:
return x.Optional || optionalChainMaySkip(x.X)
case *js.IndexExpr:
return x.Optional || optionalChainMaySkip(x.X)
case *js.CallExpr:
return x.Optional || optionalChainMaySkip(x.X)
case *js.TemplateExpr:
return x.Optional || optionalChainMaySkip(x.Tag)
default:
return false
}
}
// evalCallCallee evaluates a call's receiver and returns its value, which the
// array-method and string-method handling consume.
func (a *analysis) evalCallCallee(callee js.IExpr, st *state) value {
switch c := callee.(type) {
case *js.DotExpr:
return a.evalExpr(c.X, st)
case *js.IndexExpr:
recv := a.evalExpr(c.X, st)
st.setCapture(c, recv)
if c.Optional || optionalChainMaySkip(c.X) {
skipped := st.clone()
taken := st.clone()
a.evalExpr(c.Y, taken)
st.replaceWith(mergeState(skipped, taken))
} else {
a.evalExpr(c.Y, st)
}
recv = st.captures[c]
st.delCapture(c)
return recv
default:
a.evalExpr(callee, st)
return value{}
}
}
// checkCallSink records a finding when a network-call sink receives a tainted
// scalar. Only the scalar part of an argument counts: a URL or data argument must
// be a string, so an unserialized object never taints a sink.
func (a *analysis) checkCallSink(call *js.CallExpr, args []value, isFetch, isBeacon bool) {
if isFetch {
if len(args) >= 1 {
a.recordURL(args[0], call, 0, sinkFetchURL)
}
return
}
if isBeacon {
if len(args) >= 1 {
a.recordURL(args[0], call, 0, sinkBeaconURL)
}
if len(args) >= 2 {
a.recordBody(argValue(args, 0), args[1], call, 1, sinkBeaconData)
}
}
}
// evalFetchInit records a tainted scalar body or referrer in fetch's second
// argument, which must be a plain object literal in version 1. Evaluating the
// literal through the heap preserves computed, spread, and duplicate-property
// semantics without evaluating any property twice.
func (a *analysis) evalFetchInit(call *js.CallExpr, dest value, initExpr js.IExpr, st *state) {
_, ok := ungroupExpr(initExpr).(*js.ObjectExpr)
if !ok {
a.evalExpr(initExpr, st)
return
}
init := a.evalExpr(initExpr, st)
body := a.readField(st, init, fieldKey{kind: fieldNamed, name: "body"})
referrer := a.readField(st, init, fieldKey{kind: fieldNamed, name: "referrer"})
a.recordBody(dest, body, call, 1, sinkFetchBody)
a.recordBody(dest, referrer, call, 1, sinkFetchReferrer)
}
// evalCallReturn returns the value a call propagates for the value-preserving
// built-ins and serializers version 1 models, and applies the array-mutating
// effect of push. Any other call returns clean.
func (a *analysis) evalCallReturn(st *state, callee js.IExpr, args []value, recv value) value {
for _, name := range []string{"String", "encodeURIComponent", "encodeURI", "escape", "btoa"} {
if a.isGlobalCallee(callee, name) {
return value{scalar: argScalar(args, 0)}
}
}
prop, base, ok := memberAccess(callee)
if !ok {
return value{}
}
if prop == "fromCharCode" && a.isGlobalCallee(base, "String") {
return value{scalar: unionArgScalars(args)}
}
if prop == "stringify" && a.isGlobalCallee(base, "JSON") {
return value{scalar: a.serializeStringify(st, argValue(args, 0))}
}
switch prop {
case "toString", "trim", "charAt", "slice", "substr", "substring":
return value{scalar: recv.scalar}
case "concat":
if isDefiniteNonString(base) || (recv.allocOnly && len(recv.allocs) != 0) {
return value{}
}
return value{scalar: mergeTaint(recv.scalar, unionArgScalars(args))}
case "push":
a.arrayPush(st, recv, args)
return value{}
case "join":
return value{scalar: a.serializeArray(st, recv)}
default:
return value{}
}
}
// arrayPush taints the array-element field of the receiver's allocations with the
// pushed values.
func (a *analysis) arrayPush(st *state, recv value, args []value) {
if len(args) == 0 {
return
}
v := args[0]
for i := 1; i < len(args); i++ {
v = mergeValue(v, args[i])
}
a.writeArrayElements(st, recv, v, true)
}
func (a *analysis) writeArrayElements(st *state, recv value, v value, definite bool) {
ids := uniqueAllocIDs(recv.allocs)
strong := len(ids) == 1 && soleCurrent(recv.allocs)
for id := range ids {
if o := st.heap[id]; o == nil || !o.array {
continue
}
st.mutObject(id).weakElem(v, definite && strong && !id.summary)
}
if a.callDepth == 0 && allocsConstrained(recv.allocs) && valueCarriesDepthZero(st, v) {
resetAllocRefDepth(st, ids)
}
a.fact()
}
func argValue(args []value, i int) value {
if i >= len(args) {
return value{}
}
return args[i]
}
func argScalar(args []value, i int) taintSet {
return argValue(args, i).scalar
}
func unionArgScalars(args []value) taintSet {
var out taintSet
for i := range args {
out = mergeTaint(out, args[i].scalar)
}
return out
}
// isDefiniteNonString reports whether concat's receiver is provably not a string,
// so array concat is not treated as string laundering.
func isDefiniteNonString(expr js.IExpr) bool {
switch x := ungroupExpr(expr).(type) {
case *js.LiteralExpr:
return x.TokenType != js.StringToken
case *js.ArrayExpr, *js.ObjectExpr, *js.NewExpr, *js.ClassDecl, *js.FuncDecl, *js.ArrowFunc:
return true
default:
return false
}
}
// isGlobalCallee reports whether callee is the named unshadowed global function,
// either bare or as a property of window/self/globalThis.
func (a *analysis) isGlobalCallee(callee js.IExpr, name string) bool {
switch c := ungroupExpr(callee).(type) {
case *js.Var:
return isGlobalRef(c, name)
case *js.DotExpr, *js.IndexExpr:
prop, base, ok := memberAccess(callee)
return ok && prop == name && isGlobalObject(base)
default:
return false
}
}
func (a *analysis) isBeaconCallee(callee js.IExpr) bool {
prop, base, ok := memberAccess(callee)
return ok && prop == "sendBeacon" && a.isNavigator(base)
}
// isNavigator reports whether expr is the unshadowed global navigator object.
func (a *analysis) isNavigator(expr js.IExpr) bool {
switch v := ungroupExpr(expr).(type) {
case *js.Var:
return isGlobalRef(v, "navigator")
case *js.DotExpr, *js.IndexExpr:
prop, base, ok := memberAccess(expr)
return ok && prop == "navigator" && isGlobalObject(base)
default:
return false
}
}
// isGlobalObject reports whether expr is one of the unshadowed global object
// aliases.
func isGlobalObject(expr js.IExpr) bool {
v, ok := ungroupExpr(expr).(*js.Var)
if !ok {
return false
}
name, ok := globalName(v)
if !ok {
return false
}
switch name {
case "window", "self", "globalThis":
return true
default:
return false
}
}
// globalName returns the name of an unshadowed global reference. The parser
// leaves free identifiers undeclared, so a NoDecl canonical identity is a global
// binding rather than a local of the same spelling.
func globalName(v *js.Var) (string, bool) {
c := canonicalVar(v)
if c.Decl != js.NoDecl {
return "", false
}
return string(c.Name()), true
}
package jstaint
import (
"bytes"
"strings"
"github.com/tdewolff/parse/v2/js"
)
// isKeyboardProp reports whether name is an event property whose value is a
// keystroke. e.target is deliberately absent: its .value is the input's current
// text, which legitimate search-as-you-type widgets read and post.
func isKeyboardProp(name []byte) bool {
return bytes.Equal(name, []byte("key")) ||
bytes.Equal(name, []byte("keyCode")) ||
bytes.Equal(name, []byte("charCode")) ||
bytes.Equal(name, []byte("which")) ||
bytes.Equal(name, []byte("code"))
}
// The only member accesses allowed between the event variable and a keyboard
// property are these framework wrappers, for example e.originalEvent.key.
func isEventWrapperProp(name []byte) bool {
return bytes.Equal(name, []byte("originalEvent")) ||
bytes.Equal(name, []byte("nativeEvent"))
}
// keyboardSource reports whether expr statically reads a keyboard property off
// the handler's event variable, possibly through originalEvent/nativeEvent
// wrappers, and returns a dotted display string for evidence. The boolean result
// of a comparison is not handled here; barriers are enforced during propagation.
func keyboardSource(expr js.IExpr, eventVar *js.Var) (string, bool) {
name, base, ok := memberAccessBytes(expr)
if !ok || !isKeyboardProp(name) {
return "", false
}
if !resolvesToEventBase(base, eventVar) {
return "", false
}
return memberDisplay(expr), true
}
// resolvesToEventBase reports whether base is the event variable itself or a
// wrapper chain rooted at it.
func resolvesToEventBase(base js.IExpr, eventVar *js.Var) bool {
if eventVar == nil {
return false
}
return eventBaseVar(base) == canonicalVar(eventVar)
}
// eventBaseVar returns the canonical event-variable candidate at the root of a
// supported wrapper chain.
func eventBaseVar(base js.IExpr) *js.Var {
for {
base = ungroupExpr(base)
if v, ok := base.(*js.Var); ok {
return canonicalVar(v)
}
name, inner, ok := memberAccessBytes(base)
if !ok || !isEventWrapperProp(name) {
return nil
}
base = inner
}
}
// memberAccess splits a dot or static-bracket member access into its property
// name and base expression.
func memberAccess(expr js.IExpr) (name string, base js.IExpr, ok bool) {
data, base, ok := memberAccessBytes(expr)
if !ok {
return "", nil, false
}
return string(data), base, true
}
func memberAccessBytes(expr js.IExpr) (name []byte, base js.IExpr, ok bool) {
switch e := ungroupExpr(expr).(type) {
case *js.DotExpr:
if n, ok := staticBytesOrIdent(ungroupExpr(e.Y)); ok {
return n, e.X, true
}
case *js.IndexExpr:
if n, ok := staticBytesOrIdent(ungroupExpr(e.Y)); ok {
return n, e.X, true
}
}
return nil, nil, false
}
// memberDisplay renders a member-access chain as dotted names for evidence,
// normalizing bracket access to dot form. An unresolvable base is shown as "?".
func memberDisplay(expr js.IExpr) string {
var names [][]byte
for {
expr = ungroupExpr(expr)
if v, ok := expr.(*js.Var); ok {
return joinMemberDisplay(v.Name(), names)
}
name, base, ok := memberAccessBytes(expr)
if !ok {
return joinMemberDisplay([]byte("?"), names)
}
names = append(names, name)
expr = base
}
}
func joinMemberDisplay(base []byte, reversedNames [][]byte) string {
size := len(base)
for _, name := range reversedNames {
size += 1 + len(name)
}
var display strings.Builder
display.Grow(size)
display.Write(base)
for i := len(reversedNames) - 1; i >= 0; i-- {
display.WriteByte('.')
display.Write(reversedNames[i])
}
return display.String()
}
// binaryOpPropagates reports whether a binary operator's result carries taint
// from its operands. Comparisons are barriers because their boolean result does
// not contain the captured key; logical operators propagate because JavaScript
// returns one of the operand values.
func binaryOpPropagates(op js.TokenType) bool {
switch op {
case js.AddToken, js.SubToken, js.MulToken, js.DivToken, js.ModToken, js.ExpToken,
js.LtLtToken, js.GtGtToken, js.GtGtGtToken,
js.BitAndToken, js.BitOrToken, js.BitXorToken,
js.AndToken, js.OrToken, js.NullishToken:
return true
default:
return false
}
}
// unaryOpPropagates reports whether a unary operator's result carries taint.
// !, typeof, void, and delete are barriers; +, -, ~, increment/decrement, and
// await keep a value derived from the operand.
func unaryOpPropagates(op js.TokenType) bool {
switch op {
case js.PosToken, js.NegToken, js.BitNotToken, js.AwaitToken,
js.PreIncrToken, js.PreDecrToken, js.PostIncrToken, js.PostDecrToken:
return true
default:
return false
}
}
package jstaint
import (
"context"
"sort"
"github.com/tdewolff/parse/v2/js"
)
// maxASTNodes bounds the parser nodes the structural pass visits. It is a safety
// ceiling: exceeding it fails the whole analysis rather than returning a partial
// decision. Real minified bundles average 4-6 bytes per node, so the cap must
// admit about 2.6 bytes per node at the 2 MiB input limit or real files become
// permanent coverage gaps; it still rejects pathological synthetic nesting.
const maxASTNodes = 800_000
// taintChain is the shortest known list of via display names a captured value
// passed through to reach a program point.
type taintChain = []string
// taintFact identifies one source at one permitted user-call depth. Depth is
// part of the fact so a returned value cannot enter a second user-defined
// callee after control returns to a root.
type taintFact struct {
source int
callDepth uint8
}
// taintSet maps a source/depth fact to its shortest via chain.
type taintSet map[taintFact]taintChain
// flowKey is the identity of one source-to-sink flow: which source occurrence,
// sink occurrence, sink kind, and argument. Multiple propagation routes to the
// same endpoints collapse to one result with the shortest chain.
type flowKey struct {
source int
sink js.INode
arg int
kind string
}
type sourceOccurrence struct {
id int
display string
}
type functionExit struct {
state *state
ret value
returned bool
}
// analysis holds the mutable state for one Analyze call.
type analysis struct {
ctx context.Context
budget *resourceBudget
sources map[js.IExpr]sourceOccurrence
sites map[js.INode]int
funcs map[*js.Var][]*funcInfo
sharedVars map[*js.Var]bool
inProgress map[*funcInfo]bool
localCache map[*js.BlockStmt]map[*js.Var]bool
restSites map[js.IBinding]int
retStack []value
returnExits [][]functionExit
suspensions []*state
display map[int]string
results map[flowKey]Result
loopDepth int
callDepth int
stopAtAwait bool
err error
}
// taintPass is the production analysis pass wired into Analyze. It first enforces
// the structural depth, node, and cancellation limits with a bounded walk, then
// runs the taint analysis on the within-limit tree.
func taintPass(ctx context.Context, ast *js.AST, budget *resourceBudget) ([]Result, int, bool, error) {
lv := &limitVisitor{ctx: ctx, budget: budget}
js.Walk(lv, ast)
if lv.err != nil {
return nil, 0, false, lv.err
}
sharedVars := map[*js.Var]bool{}
for _, v := range ast.Declared {
sharedVars[canonicalVar(v)] = true
}
a := &analysis{
ctx: ctx,
budget: budget,
sites: numberAllocSites(ast),
funcs: collectFuncValues(ast),
sharedVars: sharedVars,
inProgress: map[*funcInfo]bool{},
localCache: map[*js.BlockStmt]map[*js.Var]bool{},
restSites: numberRestSites(ast),
results: map[flowKey]Result{},
}
roots := a.discoverReachableRoots(ast)
eventVars := map[*js.Var]bool{}
for _, r := range roots {
if r.eventVar != nil {
eventVars[r.eventVar] = true
}
}
a.sources, a.display = numberSources(ast, eventVars)
a.analyzeReachable(ast, roots)
if a.err != nil {
return nil, 0, false, a.err
}
results, total, truncated := a.finalizeResults()
return results, total, truncated, nil
}
// analyzeReachable runs the top level once, then every reachable callback root to
// a fixed point over the file-scope may-state that callbacks publish to and read
// from one another.
func (a *analysis) analyzeReachable(ast *js.AST, roots []rootSite) {
// The top level runs once with a clean published state, so a request that runs
// before any event cannot observe a taint a later handler produces.
top := newState()
top = a.analyzeBlock(&ast.BlockStmt, top)
if !a.alive() {
return
}
global := a.sharedState(top)
// A socket that survives to file-scope state can be observed open by a later
// callback, which is when a send on it becomes a network sink.
markSocketsObservable(global)
for a.alive() {
changed := false
for i := range roots {
st := global.clone()
a.suspensions = make([]*state, 0)
a.bindRootParams(roots[i], st)
a.returnExits = append(a.returnExits, nil)
st = a.analyzeBlock(roots[i].fn.body, st)
exits := a.popReturnExits(st)
if !a.alive() {
return
}
next := global
for _, exit := range exits {
next = a.publishShared(next, exit.state)
}
for _, suspended := range a.suspensions {
next = a.publishShared(next, suspended)
}
a.suspensions = nil
markSocketsObservable(next)
if !stateEqual(next, global) {
global = next
changed = true
}
}
if !changed || !a.fact() {
return
}
}
}
// popReturnExits closes the current function-flow frame. Explicit return states
// and the normal fallthrough state are separate execution alternatives; code
// after a return must not inherit writes from the returned path.
func (a *analysis) popReturnExits(normalExit *state) []functionExit {
index := len(a.returnExits) - 1
exits := a.returnExits[index]
a.returnExits = a.returnExits[:index]
if normalExit.continues {
exits = append(exits, functionExit{state: normalExit})
}
return exits
}
// bindRootParams binds a callback's parameters before its body. A keyboard
// handler's first parameter is the event object whose keystroke reads are the
// numbered sources, so only its later parameters are bound here.
func (a *analysis) bindRootParams(r rootSite, st *state) {
start := 0
if r.eventVar != nil {
start = 1
}
for i := start; i < len(r.fn.params.List); i++ {
if i < r.argCount {
a.bindParamPattern(r.fn.params.List[i].Binding, value{}, st)
continue
}
a.analyzeBinding(&r.fn.params.List[i], st, true)
}
}
// sortedResults returns the recorded flows in deterministic content order.
func (a *analysis) sortedResults() []Result {
out := make([]Result, 0, len(a.results))
for _, r := range a.results {
out = append(out, r)
}
sort.Slice(out, func(i, j int) bool { return lessResult(out[i], out[j]) })
return out
}
func lessResult(x, y Result) bool {
if x.Source != y.Source {
return x.Source < y.Source
}
xv, yv := joinChain(x.Via), joinChain(y.Via)
if xv != yv {
return xv < yv
}
return x.Sink < y.Sink
}
func joinChain(c []string) string {
out := ""
for i, s := range c {
if i != 0 {
out += "\x00"
}
out += s
}
return out
}
// record inserts a source-to-sink flow, keeping the shortest via chain when a
// route to the same endpoints already exists. Ties break lexicographically by
// the joined via names so output does not depend on map iteration order.
func (a *analysis) record(ts taintSet, sink js.INode, arg int, sinkText string) {
for fact, chain := range ts {
if !a.alive() {
return
}
key := flowKey{source: fact.source, sink: sink, arg: arg, kind: sinkText}
cand := Result{Source: a.display[fact.source], Via: append([]string(nil), chain...), Sink: sinkText}
if ex, ok := a.results[key]; ok {
if shorterChain(cand.Via, ex.Via) {
a.results[key] = cand
}
continue
}
if !a.fact() {
return
}
a.results[key] = cand
}
}
func shorterChain(cand, existing []string) bool {
if len(cand) != len(existing) {
return len(cand) < len(existing)
}
return joinChain(cand) < joinChain(existing)
}
func (a *analysis) fact() bool {
if a.err != nil {
return false
}
if err := a.budget.addFact(); err != nil {
a.err = err
return false
}
return true
}
func (a *analysis) alive() bool {
if a.err != nil {
return false
}
if err := a.ctx.Err(); err != nil {
a.err = err
return false
}
return true
}
// numberSources assigns a deterministic occurrence id to every keyboard-source
// read whose base resolves to a discovered event variable. The traversal order
// does not affect output because results are content-sorted; it only needs to be
// stable across identical inputs.
func numberSources(ast *js.AST, eventVars map[*js.Var]bool) (map[js.IExpr]sourceOccurrence, map[int]string) {
n := &sourceNumberer{eventVars: eventVars, out: map[js.IExpr]sourceOccurrence{}, display: map[int]string{}}
js.Walk(n, ast)
return n.out, n.display
}
type sourceNumberer struct {
eventVars map[*js.Var]bool
out map[js.IExpr]sourceOccurrence
display map[int]string
next int
}
func (n *sourceNumberer) Exit(js.INode) {}
func (n *sourceNumberer) Enter(node js.INode) js.IVisitor {
if walkComputedClassName(n, node) {
return n
}
expr, ok := node.(js.IExpr)
if !ok {
return n
}
name, base, ok := memberAccessBytes(expr)
if !ok || !isKeyboardProp(name) || !n.resolves(base) {
return n
}
disp := memberDisplay(expr)
n.out[expr] = sourceOccurrence{id: n.next, display: disp}
n.display[n.next] = disp
n.next++
return n
}
func (n *sourceNumberer) resolves(base js.IExpr) bool {
v := eventBaseVar(base)
return v != nil && n.eventVars[v]
}
// mergeTaint unions two taint sets, keeping the shortest chain per source.
func mergeTaint(a, b taintSet) taintSet {
if len(a) == 0 {
return b
}
if len(b) == 0 {
return a
}
out := make(taintSet, len(a)+len(b))
for fact, ch := range a {
out[fact] = ch
}
for fact, ch := range b {
if ex, ok := out[fact]; !ok || shorterChain(ch, ex) {
out[fact] = ch
}
}
return out
}
// appendVia extends every chain in ts with name, skipping a redundant repeat so
// loop fixed points do not grow chains without bound.
func appendVia(ts taintSet, name string) taintSet {
if len(ts) == 0 {
return nil
}
out := make(taintSet, len(ts))
for fact, ch := range ts {
if len(ch) > 0 && ch[len(ch)-1] == name {
out[fact] = ch
continue
}
nc := make(taintChain, len(ch)+1)
copy(nc, ch)
nc[len(ch)] = name
out[fact] = nc
}
return out
}
func taintEqual(a, b taintSet) bool {
if len(a) != len(b) {
return false
}
for fact, ca := range a {
cb, ok := b[fact]
if !ok || !chainEqual(ca, cb) {
return false
}
}
return true
}
func chainEqual(a, b taintChain) bool {
if len(a) != len(b) {
return false
}
for i := range a {
if a[i] != b[i] {
return false
}
}
return true
}
package jstaint
import (
"strconv"
"github.com/tdewolff/parse/v2/js"
)
// srcValue is the abstract value of a keyboard-source read: scalar taint tagged
// with the source occurrence and an empty laundering chain.
func (a *analysis) srcValue(id int) value {
depth := uint8(a.callDepth) // #nosec G115 -- callDepth is 0 or 1 under the depth-1 interprocedural cap
return value{scalar: taintSet{{source: id, callDepth: depth}: taintChain{}}}
}
// analyzeBlock threads the state through a statement list in source order.
func (a *analysis) analyzeBlock(block *js.BlockStmt, st *state) *state {
if block == nil {
return st
}
for i := range block.List {
if !a.alive() || !st.continues {
return st
}
st = a.analyzeStmt(block.List[i], st)
}
return st
}
func (a *analysis) analyzeStmt(stmt js.IStmt, st *state) *state {
if !a.alive() || !st.continues {
return st
}
switch s := stmt.(type) {
case *js.VarDecl:
for i := range s.List {
a.analyzeBinding(&s.List[i], st, s.TokenType != js.VarToken)
}
case *js.ClassDecl:
a.evalClass(s, st)
case *js.ExprStmt:
a.analyzeExprStmt(s.Value, st)
case *js.BlockStmt:
st = a.analyzeBlock(s, st)
case *js.IfStmt:
a.evalExpr(s.Cond, st)
thenSt := a.analyzeStmt(s.Body, st.clone())
var elseSt *state
if s.Else != nil {
elseSt = a.analyzeStmt(s.Else, st.clone())
} else {
elseSt = st.clone()
}
st = mergeState(thenSt, elseSt)
case *js.ForStmt:
a.analyzeInit(s.Init, st)
st = a.analyzeLoop(s.Body, s.Cond, s.Post, st)
case *js.WhileStmt:
st = a.analyzeLoop(s.Body, s.Cond, nil, st)
case *js.DoWhileStmt:
st = a.analyzeDoLoop(s.Body, s.Cond, st)
case *js.ForInStmt:
a.evalExpr(s.Value, st)
st = a.analyzeIterationLoop(s.Body, s.Init, value{}, false, st)
case *js.ForOfStmt:
iter := a.evalExpr(s.Value, st)
st = a.analyzeIterationLoop(s.Body, s.Init, iter, true, st)
case *js.SwitchStmt:
st = a.analyzeSwitch(s, st)
case *js.TryStmt:
st = a.analyzeTry(s, st)
case *js.ReturnStmt:
var rv value
if s.Value != nil {
rv = a.evalExpr(s.Value, st)
}
if n := len(a.returnExits); n > 0 {
a.returnExits[n-1] = append(a.returnExits[n-1], functionExit{
state: st.clone(), ret: rv, returned: true,
})
st.continues = false
}
case *js.ThrowStmt:
a.evalExpr(s.Value, st)
case *js.WithStmt:
a.evalExpr(s.Cond, st)
bodySt := a.analyzeStmt(s.Body, st.clone())
st = mergeState(st, bodySt)
case *js.LabelledStmt:
st = a.analyzeStmt(s.Value, st)
case *js.BranchStmt:
// break/continue: no data effect for this model.
}
return st
}
// analyzeInit handles a for-loop initializer, which may be a var declaration or
// an expression.
func (a *analysis) analyzeInit(init js.IExpr, st *state) {
if init == nil {
return
}
if vd, ok := init.(*js.VarDecl); ok {
for i := range vd.List {
a.analyzeBinding(&vd.List[i], st, vd.TokenType != js.VarToken)
}
return
}
a.evalExpr(init, st)
}
// analyzeBinding applies a declaration binding. A tainted initializer taints the
// bound variable, while a clean initializer is a strong update. An absent lexical
// initializer writes undefined; an absent var initializer is only a declaration.
func (a *analysis) analyzeBinding(be *js.BindingElement, st *state, clearAbsent bool) {
v, ok := be.Binding.(*js.Var)
if !ok {
if be.Default != nil {
a.evalExpr(be.Default, st)
}
a.evalBindingPattern(be.Binding, st)
return
}
cv := canonicalVar(v)
if be.Default == nil {
if clearAbsent {
st.delEnv(cv)
}
return
}
rhs := a.evalExpr(be.Default, st)
a.bindVar(st, cv, rhs)
}
// bindVar performs a strong update binding a variable to a value, appending the
// variable name to the laundering chain and preserving aliased allocations.
func (a *analysis) bindVar(st *state, cv *js.Var, rhs value) value {
nt := value{
scalar: appendVia(rhs.scalar, string(cv.Name())),
allocs: rhs.allocs, allocOnly: rhs.allocOnly, scheme: rhs.scheme,
}
if !storable(nt) {
st.delEnv(cv)
return value{}
}
st.setEnv(cv, nt)
a.fact()
return nt
}
// storable reports whether a value carries information worth keeping in the
// environment: taint, an allocation identity, or a definite URL scheme.
func storable(v value) bool {
return len(v.scalar) != 0 || len(v.allocs) != 0 || v.scheme.set
}
func (a *analysis) analyzeExprStmt(expr js.IExpr, st *state) {
expr = ungroupExpr(expr)
if be, ok := expr.(*js.BinaryExpr); ok && isAssignOp(be.Op) {
a.handleAssign(be, st)
return
}
a.evalExpr(expr, st)
}
// handleAssign updates state for target = rhs. A plain assignment is a strong
// update that can clear taint; a compound assignment reads the old value and
// combines it. A logical assignment writes only when its short-circuit condition
// allows, so the write and its right side are may-state.
func (a *analysis) handleAssign(be *js.BinaryExpr, st *state) value {
target := ungroupExpr(be.X)
if isLogicalAssignOp(be.Op) {
return a.handleLogicalAssign(be, target, st)
}
switch t := target.(type) {
case *js.Var:
cv := canonicalVar(t)
old := st.env[cv]
rhs := a.evalExpr(be.Y, st)
return a.assignVar(st, cv, old, rhs, be.Op)
case *js.DotExpr, *js.IndexExpr:
return a.assignMember(be, target, st)
default:
a.evalExpr(target, st)
return a.evalExpr(be.Y, st)
}
}
func (a *analysis) assignVar(st *state, cv *js.Var, old, rhs value, op js.TokenType) value {
name := string(cv.Name())
var nt value
if op == js.EqToken {
nt = value{scalar: appendVia(rhs.scalar, name), allocs: rhs.allocs, allocOnly: rhs.allocOnly, scheme: rhs.scheme}
} else {
// A compound assignment coerces to a string or number, so the result is a
// scalar and carries no allocation identity.
nt = value{scalar: appendVia(mergeTaint(old.scalar, rhs.scalar), name)}
if op == js.AddEqToken {
nt.scheme = old.scheme
}
}
if !storable(nt) {
st.delEnv(cv)
return value{}
}
st.setEnv(cv, nt)
a.fact()
return nt
}
func (a *analysis) assignMember(be *js.BinaryExpr, target js.IExpr, st *state) value {
recv, key := a.evalMemberTarget(target, st)
st.setCapture(be, recv)
var old value
if be.Op != js.EqToken {
// Compound assignment captures the current property value before the RHS
// runs. The RHS may overwrite the same field.
old = a.readField(st, recv, key)
}
rhs := a.evalExpr(be.Y, st)
recv = st.captures[be]
st.delCapture(be)
writeVal := rhs
if be.Op != js.EqToken {
writeVal = value{scalar: mergeTaint(old.scalar, rhs.scalar)}
if be.Op == js.AddEqToken {
writeVal.scheme = old.scheme
}
}
a.resourceSrcSink(recv, key, writeVal, be, st)
a.writeField(st, recv, key, writeVal)
return writeVal
}
// handleLogicalAssign models &&=, ||=, and ??=, whose right side and write are
// only reached on the short-circuit path, so both are merged as may-state.
func (a *analysis) handleLogicalAssign(be *js.BinaryExpr, target js.IExpr, st *state) value {
switch t := target.(type) {
case *js.Var:
cv := canonicalVar(t)
skipped := st.clone()
taken := st.clone()
rhs := a.evalExpr(be.Y, taken)
a.assignVar(taken, cv, taken.env[cv], rhs, js.EqToken)
st.replaceWith(mergeState(skipped, taken))
return st.env[cv]
default:
recv, key := a.evalMemberTarget(target, st)
st.setCapture(be, recv)
skipped := st.clone()
taken := st.clone()
rhs := a.evalExpr(be.Y, taken)
takenRecv := taken.captures[be]
resultRecv := mergeValue(skipped.captures[be], takenRecv)
taken.delCapture(be)
skipped.delCapture(be)
a.resourceSrcSink(takenRecv, key, rhs, be, taken)
a.writeField(taken, takenRecv, key, rhs)
st.replaceWith(mergeState(skipped, taken))
return a.readField(st, resultRecv, key)
}
}
// evalMemberTarget evaluates the receiver and key of a member assignment target
// before the right-hand side, and returns the receiver value and field key.
func (a *analysis) evalMemberTarget(target js.IExpr, st *state) (value, fieldKey) {
switch t := ungroupExpr(target).(type) {
case *js.DotExpr:
recv := a.evalExpr(t.X, st)
return recv, fieldKeyOf(ungroupExpr(t.Y))
case *js.IndexExpr:
recv := a.evalExpr(t.X, st)
st.setCapture(t, recv)
a.evalExpr(t.Y, st)
recv = st.captures[t]
st.delCapture(t)
return recv, fieldKeyOf(ungroupExpr(t.Y))
}
return value{}, fieldKey{}
}
// analyzeLoop iterates the body to a fixed point over the loop state.
func (a *analysis) analyzeLoop(body js.IStmt, cond, post js.IExpr, st *state) *state {
a.loopDepth++
defer func() { a.loopDepth-- }()
cur := st.clone()
for {
if !a.alive() {
return cur
}
if cond != nil {
a.evalExpr(cond, cur)
}
bodyOut := a.analyzeStmt(body, cur.clone())
if post != nil {
a.evalExpr(post, bodyOut)
}
merged := mergeState(cur, bodyOut)
if stateEqual(merged, cur) {
return merged
}
if !a.fact() {
return merged
}
cur = merged
}
}
// analyzeDoLoop applies the first body execution before merging later iterations,
// because a do-while body always runs at least once.
func (a *analysis) analyzeDoLoop(body js.IStmt, cond js.IExpr, st *state) *state {
a.loopDepth++
defer func() { a.loopDepth-- }()
cur := a.analyzeStmt(body, st.clone())
if !cur.continues {
return cur
}
if cond != nil {
a.evalExpr(cond, cur)
}
for {
if !a.alive() {
return cur
}
bodyOut := a.analyzeStmt(body, cur.clone())
if cond != nil {
a.evalExpr(cond, bodyOut)
}
merged := mergeState(cur, bodyOut)
if stateEqual(merged, cur) {
return merged
}
if !a.fact() {
return merged
}
cur = merged
}
}
// analyzeIterationLoop applies the implicit iteration assignment before each body
// execution. The incoming state remains an exit alternative because an iterable
// can be empty.
func (a *analysis) analyzeIterationLoop(
body js.IStmt,
init js.IExpr,
iter value,
forOf bool,
st *state,
) *state {
a.loopDepth++
defer func() { a.loopDepth-- }()
cur := st.clone()
for {
if !a.alive() {
return cur
}
bodyIn := cur.clone()
elem := value{}
if forOf {
elem = a.iterationElement(bodyIn, iter)
}
a.applyIterationBinding(init, elem, bodyIn)
bodyOut := a.analyzeStmt(body, bodyIn)
merged := mergeState(cur, bodyOut)
if stateEqual(merged, cur) {
return merged
}
if !a.fact() {
return merged
}
cur = merged
}
}
// iterationElement is the value each element of a for-of iterable can carry:
// scalar taint from iterating a tainted string, plus the iterable's element
// field for an array of values.
func (a *analysis) iterationElement(st *state, iter value) value {
return a.collectArrayElements(st, iter)
}
func (a *analysis) applyIterationBinding(init js.IExpr, elem value, st *state) {
if decl, ok := init.(*js.VarDecl); ok {
if len(decl.List) != 0 {
a.assignIterationVar(decl.List[0].Binding, elem, st)
}
return
}
if v, ok := ungroupExpr(init).(*js.Var); ok {
a.bindVar(st, canonicalVar(v), elem)
return
}
target := ungroupExpr(init)
switch target.(type) {
case *js.DotExpr, *js.IndexExpr:
recv, key := a.evalMemberTarget(target, st)
a.resourceSrcSink(recv, key, elem, target, st)
a.writeField(st, recv, key, elem)
default:
a.evalExpr(target, st)
}
}
func (a *analysis) assignIterationVar(binding js.IBinding, elem value, st *state) {
if v, ok := binding.(*js.Var); ok {
a.bindVar(st, canonicalVar(v), elem)
return
}
a.evalBindingPattern(binding, st)
}
// evalBindingPattern evaluates computed property names and nested defaults. This
// phase does not bind values extracted from a destructuring pattern, but the
// pattern's expressions still execute and can contain sinks or scalar writes.
func (a *analysis) evalBindingPattern(binding js.IBinding, st *state) {
switch b := binding.(type) {
case *js.BindingArray:
for i := range b.List {
a.evalNestedBinding(&b.List[i], st)
}
if b.Rest != nil {
a.evalBindingPattern(b.Rest, st)
}
case *js.BindingObject:
for i := range b.List {
item := &b.List[i]
if item.Key != nil {
a.evalExpr(item.Key.Computed, st)
}
a.evalNestedBinding(&item.Value, st)
}
}
}
func (a *analysis) evalNestedBinding(be *js.BindingElement, st *state) {
if be.Default != nil {
taken := st.clone()
a.evalExpr(be.Default, taken)
st.replaceWith(mergeState(st, taken))
}
a.evalBindingPattern(be.Binding, st)
}
func (a *analysis) analyzeSwitch(s *js.SwitchStmt, st *state) *state {
a.evalExpr(s.Init, st)
out := st.clone()
var fall *state
for i := range s.List {
cl := &s.List[i]
if cl.Cond != nil {
a.evalExpr(cl.Cond, st)
}
in := st.clone()
if fall != nil {
in = mergeState(in, fall)
}
for j := range cl.List {
in = a.analyzeStmt(cl.List[j], in)
}
fall = in
out = mergeState(out, in)
}
return out
}
// analyzeTry models that an exception can occur anywhere in the body, so the
// catch clause sees the pre-body state merged with taint the body may have set,
// and the finally clause applies to every outgoing edge.
func (a *analysis) analyzeTry(s *js.TryStmt, st *state) *state {
frame := len(a.returnExits) - 1
returnStart := 0
if frame >= 0 {
returnStart = len(a.returnExits[frame])
}
bodySt := a.analyzeBlock(s.Body, st.clone())
merged := bodySt
if s.Catch != nil {
catchIn := mergeState(st, bodySt)
catchSt := a.analyzeBlock(s.Catch, catchIn)
merged = mergeState(bodySt, catchSt)
}
if s.Finally != nil {
var pending []functionExit
if frame >= 0 {
pending = append(pending, a.returnExits[frame][returnStart:]...)
a.returnExits[frame] = a.returnExits[frame][:returnStart]
}
if merged.continues {
merged = a.analyzeBlock(s.Finally, merged.clone())
}
if frame >= 0 {
for i := range pending {
out := a.analyzeBlock(s.Finally, pending[i].state.clone())
if out.continues {
pending[i].state = out
a.returnExits[frame] = append(a.returnExits[frame], pending[i])
}
}
}
}
return merged
}
func (a *analysis) evalClass(class *js.ClassDecl, st *state) {
a.evalExpr(class.Extends, st)
// Every computed key is evaluated while the class elements are defined. Static
// fields and blocks initialize only after all keys have been computed.
for i := range class.List {
item := &class.List[i]
if item.Method != nil {
a.evalExpr(item.Method.Name.Computed, st)
} else if item.StaticBlock == nil {
a.evalExpr(item.Name.Computed, st)
}
}
for i := range class.List {
item := &class.List[i]
if item.StaticBlock != nil {
blockSt := a.analyzeBlock(item.StaticBlock, st.clone())
st.replaceWith(blockSt)
} else if item.Method == nil && item.Static {
a.evalExpr(item.Init, st)
}
}
}
// evalExpr returns the abstract value of expr and records any sink it reaches. It
// mutates state only through an assignment used as a value or a may-state merge.
func (a *analysis) evalExpr(expr js.IExpr, st *state) value {
if !a.alive() || !st.continues {
return value{}
}
expr = ungroupExpr(expr)
switch x := expr.(type) {
case nil:
return value{}
case *js.Var:
return st.env[canonicalVar(x)]
case *js.LiteralExpr:
return value{scheme: schemeOfLiteral(x)}
case *js.DotExpr:
if occ, ok := a.sources[expr]; ok {
return a.srcValue(occ.id)
}
base := a.evalExpr(x.X, st)
name, ok := staticStringOrIdent(ungroupExpr(x.Y))
if !ok {
return value{}
}
return a.readField(st, base, fieldKey{kind: fieldNamed, name: name})
case *js.IndexExpr:
if occ, ok := a.sources[expr]; ok {
return a.srcValue(occ.id)
}
base := a.evalExpr(x.X, st)
st.setCapture(x, base)
key := a.evalReadIndexKey(x, st)
base = st.captures[x]
st.delCapture(x)
return a.readField(st, base, key)
case *js.BinaryExpr:
if isAssignOp(x.Op) {
return a.handleAssign(x, st)
}
lx := a.evalExpr(x.X, st)
var rx value
if isLogicalOp(x.Op) {
skipped := st.clone()
taken := st.clone()
rx = a.evalExpr(x.Y, taken)
st.replaceWith(mergeState(skipped, taken))
} else {
rx = a.evalExpr(x.Y, st)
}
if binaryOpPropagates(x.Op) {
if isLogicalOp(x.Op) {
// ||, &&, and ?? return one operand value, so both allocations and
// scalars can flow through. The operator itself is not one of the
// scheme-preserving forms, so destination classification is reset.
out := mergeValue(lx, rx)
out.scheme = schemeState{}
return out
}
out := value{scalar: mergeTaint(lx.scalar, rx.scalar)}
if x.Op == js.AddToken {
// A concatenation's URL scheme is fixed by its leftmost prefix.
out.scheme = lx.scheme
}
return out
}
return value{}
case *js.UnaryExpr:
return a.evalUnary(x, st)
case *js.CondExpr:
a.evalExpr(x.Cond, st)
thenSt := st.clone()
elseSt := st.clone()
thenVal := a.evalExpr(x.X, thenSt)
elseVal := a.evalExpr(x.Y, elseSt)
st.replaceWith(mergeState(thenSt, elseSt))
return mergeValue(thenVal, elseVal)
case *js.TemplateExpr:
return a.evalTemplate(x, st)
case *js.CommaExpr:
var last value
for i := range x.List {
last = a.evalExpr(x.List[i], st)
}
last.scheme = schemeState{}
return last
case *js.CallExpr:
return a.evalCall(x, st)
case *js.NewExpr:
return a.evalNew(x, st)
case *js.ArrayExpr:
return a.evalArray(x, st)
case *js.ObjectExpr:
return a.evalObject(x, st)
case *js.ClassDecl:
a.evalClass(x, st)
return value{}
default:
return value{}
}
}
func (a *analysis) evalUnary(x *js.UnaryExpr, st *state) value {
if x.Op == js.DeleteToken {
a.evalDelete(x.X, st)
return value{}
}
if isUpdateOp(x.Op) {
return a.evalUpdate(x.X, x.Op, st)
}
xt := a.evalExpr(x.X, st)
if x.Op == js.AwaitToken {
if a.stopAtAwait {
st.continues = false
return xt
}
if a.suspensions != nil {
a.suspensions = append(a.suspensions, st.clone())
}
xt.scheme = schemeState{}
return xt
}
if unaryOpPropagates(x.Op) {
return value{scalar: xt.scalar}
}
return value{}
}
func (a *analysis) evalDelete(target js.IExpr, st *state) {
target = ungroupExpr(target)
switch target.(type) {
case *js.DotExpr, *js.IndexExpr:
if optionalChainMaySkip(target) {
skipped := st.clone()
taken := st.clone()
recv, key := a.evalMemberTarget(target, taken)
a.deleteField(taken, recv, key)
st.replaceWith(mergeState(skipped, taken))
return
}
recv, key := a.evalMemberTarget(target, st)
a.deleteField(st, recv, key)
default:
a.evalExpr(target, st)
}
}
func (a *analysis) evalUpdate(target js.IExpr, op js.TokenType, st *state) value {
target = ungroupExpr(target)
switch t := target.(type) {
case *js.Var:
cv := canonicalVar(t)
old := st.env[cv]
return a.assignVar(st, cv, old, value{}, op)
case *js.DotExpr, *js.IndexExpr:
recv, key := a.evalMemberTarget(target, st)
old := a.readField(st, recv, key)
updated := value{scalar: old.scalar}
a.writeField(st, recv, key, updated)
return updated
default:
old := a.evalExpr(target, st)
return value{scalar: old.scalar}
}
}
func isUpdateOp(op js.TokenType) bool {
switch op {
case js.PreIncrToken, js.PreDecrToken, js.PostIncrToken, js.PostDecrToken:
return true
default:
return false
}
}
func (a *analysis) evalTemplate(x *js.TemplateExpr, st *state) value {
if x.Tag != nil {
a.evalExpr(x.Tag, st)
valuesSt := st
if x.Optional {
valuesSt = st.clone()
}
for i := range x.List {
a.evalExpr(x.List[i].Expr, valuesSt)
}
if x.Optional {
st.replaceWith(mergeState(st, valuesSt))
}
// A tag function's return value is not modeled, so the result is clean.
return value{}
}
var ts taintSet
var firstScheme schemeState
for i := range x.List {
v := a.evalExpr(x.List[i].Expr, st)
if i == 0 {
firstScheme = v.scheme
}
ts = mergeTaint(ts, v.scalar)
}
return value{scalar: ts, scheme: templateScheme(x, firstScheme)}
}
// evalReadIndexKey evaluates a read index expression, honoring optional-chain
// short-circuit for its side effects, and returns the field key it selects.
func (a *analysis) evalReadIndexKey(x *js.IndexExpr, st *state) fieldKey {
if x.Optional || optionalChainMaySkip(x.X) {
skipped := st.clone()
taken := st.clone()
a.evalExpr(x.Y, taken)
st.replaceWith(mergeState(skipped, taken))
} else {
a.evalExpr(x.Y, st)
}
return fieldKeyOf(ungroupExpr(x.Y))
}
func (a *analysis) evalNew(x *js.NewExpr, st *state) value {
isArray := a.isGlobalCallee(x.X, "Array")
a.evalExpr(x.X, st)
var args []value
if x.Args != nil {
args = make([]value, len(x.Args.List))
for i := range x.Args.List {
args[i] = a.evalExpr(x.Args.List[i].Value, st)
}
}
site, ok := a.sites[x]
if !ok {
return value{}
}
if v, handled := a.evalNewSink(x, args, site, st); handled {
return v
}
if !isArray {
return a.allocate(st, site, a.loopDepth > 0)
}
a.promoteCurrent(st, site)
fresh := &object{array: true}
if len(args) != 1 || !isNumericLiteralExpr(x.Args.List[0].Value) {
for i := range args {
fresh.setNamed(strconv.Itoa(i), args[i], true)
a.fact()
}
}
return a.installLiteral(st, site, a.loopDepth > 0, fresh)
}
func isNumericLiteralExpr(expr js.IExpr) bool {
_, ok := numericPropertyNameOf(expr)
return ok
}
func (a *analysis) evalArray(x *js.ArrayExpr, st *state) value {
site, ok := a.sites[x]
if !ok {
for i := range x.List {
if x.List[i].Value != nil {
a.evalExpr(x.List[i].Value, st)
}
}
return value{}
}
a.promoteCurrent(st, site)
fresh := &object{array: true}
unknownIndex := false
for i := range x.List {
el := &x.List[i]
if el.Value == nil {
continue
}
ev := a.evalExpr(el.Value, st)
if el.Spread {
fresh.weakElem(a.spreadElements(st, ev), false)
unknownIndex = true
a.fact()
continue
}
if unknownIndex {
fresh.weakElem(ev, true)
} else {
fresh.setNamed(strconv.Itoa(i), ev, true)
}
a.fact()
}
return a.installLiteral(st, site, a.loopDepth > 0, fresh)
}
func (a *analysis) evalObject(x *js.ObjectExpr, st *state) value {
site, ok := a.sites[x]
if !ok {
a.evalObjectSideEffects(x, st)
return value{}
}
// A literal's properties are evaluated into one fresh runtime object. Only
// after the final property is known is that object merged into a loop summary.
a.promoteCurrent(st, site)
fresh := &object{}
for i := range x.List {
p := &x.List[i]
if p.Spread {
sv := a.evalExpr(p.Value, st)
a.spreadInto(st, fresh, sv)
continue
}
if p.Name != nil && p.Name.Computed != nil {
a.evalExpr(p.Name.Computed, st)
}
pv := a.evalExpr(p.Value, st)
if p.Init != nil {
a.evalExpr(p.Init, st)
}
key := fieldKey{kind: fieldWild}
if p.Name != nil {
if p.Name.IsComputed() {
key = fieldKeyOf(ungroupExpr(p.Name.Computed))
} else {
key = fieldKeyOf(&p.Name.Literal)
}
}
switch key.kind {
case fieldNamed:
fresh.setNamed(key.name, pv, true)
case fieldElem:
if key.name != "" {
fresh.setNamed(key.name, pv, true)
} else {
fresh.weakElem(pv, true)
}
default:
fresh.writeWild(pv, true)
}
a.fact()
}
return a.installLiteral(st, site, a.loopDepth > 0, fresh)
}
// evalObjectSideEffects evaluates an object literal's expressions without
// recording fields, used only when the literal had no allocation site assigned.
func (a *analysis) evalObjectSideEffects(x *js.ObjectExpr, st *state) {
for i := range x.List {
p := &x.List[i]
if p.Name != nil && p.Name.Computed != nil {
a.evalExpr(p.Name.Computed, st)
}
a.evalExpr(p.Value, st)
if p.Init != nil {
a.evalExpr(p.Init, st)
}
}
}
// spreadElements collapses a spread source's element taint into one value for
// insertion into the target array element field.
func (a *analysis) spreadElements(st *state, src value) value {
return a.collectArrayElements(st, src)
}
func (a *analysis) collectArrayElements(st *state, src value) value {
var out value
have := false
if len(src.scalar) != 0 {
out, have = mergePresentValue(out, have, value{scalar: src.scalar})
}
for ref := range src.allocs {
o := st.heap[ref.id]
if o == nil || !o.array {
continue
}
for name, fv := range o.fields {
if isArrayIndexName(name) {
out, have = mergePresentValue(out, have, applyRefDepth(fv, ref))
}
}
if o.elemMay {
out, have = mergePresentValue(out, have, applyRefDepth(o.elem, ref))
}
if o.wildMay {
out, have = mergePresentValue(out, have, applyRefDepth(o.wild, ref))
}
}
if !have {
return value{}
}
if !src.allocOnly {
out.allocOnly = false
}
return out
}
// spreadInto copies a spread source onto a fresh object literal. A named field
// present on every possible source definitely overwrites the earlier property;
// partial and unresolved fields remain weak updates.
func (a *analysis) spreadInto(st *state, dst *object, src value) {
fields := map[string]value{}
definite := map[string]int{}
elemDefinite := 0
wildDefinite := 0
var elem, wild, wildReq value
elemMay := false
wildMay := false
wildReqSeen := false
for ref := range src.allocs {
so := st.heap[ref.id]
if so == nil {
continue
}
for k, fv := range so.fields {
fv = applyRefDepth(fv, ref)
if old, ok := fields[k]; ok {
fields[k] = mergeValue(old, fv)
} else {
fields[k] = fv
}
if so.must[k] {
definite[k]++
}
}
if so.elemMay {
elem, elemMay = mergePresentValue(elem, elemMay, applyRefDepth(so.elem, ref))
}
if so.wildMay {
wild, wildMay = mergePresentValue(wild, wildMay, applyRefDepth(so.wild, ref))
}
if so.elemMust {
elemDefinite++
}
if so.wildMust {
wildDefinite++
if wildReqSeen {
wildReq = mergeValue(wildReq, applyRefDepth(so.wildReq, ref))
} else {
wildReq = applyRefDepth(so.wildReq, ref)
wildReqSeen = true
}
}
}
// A scalar spread contributes character properties at numeric keys.
if len(src.scalar) != 0 {
elem, elemMay = mergePresentValue(elem, elemMay, value{scalar: src.scalar})
}
allSources := src.allocOnly && len(src.allocs) != 0
if wildMay {
dst.writeWild(wild, false)
}
if allSources && wildDefinite == len(src.allocs) {
dst.wildMust = true
dst.wildReq = wildReq
}
for k, fv := range fields {
dst.setNamed(k, fv, allSources && definite[k] == len(src.allocs))
}
if elemMay {
dst.weakElem(elem, allSources && elemDefinite == len(src.allocs))
}
a.fact()
}
func isLogicalOp(op js.TokenType) bool {
return op == js.AndToken || op == js.OrToken || op == js.NullishToken
}
func isLogicalAssignOp(op js.TokenType) bool {
return op == js.AndEqToken || op == js.OrEqToken || op == js.NullishEqToken
}
func isAssignOp(op js.TokenType) bool {
switch op {
case js.EqToken, js.AddEqToken, js.SubEqToken, js.MulEqToken, js.DivEqToken,
js.ModEqToken, js.ExpEqToken, js.LtLtEqToken, js.GtGtEqToken, js.GtGtGtEqToken,
js.BitAndEqToken, js.BitOrEqToken, js.BitXorEqToken,
js.AndEqToken, js.OrEqToken, js.NullishEqToken:
return true
default:
return false
}
}
// Package log provides a structured-logging wrapper around log/slog.
//
// Rationale: CSM's daemon currently emits timestamped log lines via direct
// fmt.Fprintf(os.Stderr, ...) calls. That works for journalctl but loses
// structure when operators ship logs to Loki, ELK, or Datadog. This
// package provides a drop-in replacement that:
//
// - Emits the legacy "[YYYY-MM-DD HH:MM:SS] msg" format in text mode so
// mixing csmlog calls with legacy fmt.Fprintf calls produces a
// uniform log stream (important during incremental migration)
// - Switches to JSON on CSM_LOG_FORMAT=json for log-shipping pipelines,
// emitting the slog-native level/msg/time/fields structure
// - Honors CSM_LOG_LEVEL={debug|info|warn|error} (default: info)
//
// Usage (preferred for new code):
//
// log.Info("daemon starting", "version", v, "pid", os.Getpid())
// log.Warn("log not found, will retry", "path", path)
// log.Error("alert dispatch failed", "err", err)
//
// Legacy call sites that still use fmt.Fprintf will keep working —
// migration is incremental. See docs/src/development.md for guidance.
package log
import (
"context"
"io"
"log/slog"
"os"
"strings"
"sync"
"sync/atomic"
"time"
)
// global is the package-level logger, loaded lazily on first use.
// atomic.Pointer so Init can swap it without data races.
var global atomic.Pointer[slog.Logger]
// Init configures the global logger from environment variables:
//
// CSM_LOG_FORMAT = "text" (default) | "json"
// CSM_LOG_LEVEL = "debug" | "info" (default) | "warn" | "error"
//
// Safe to call multiple times. Returns the installed logger so callers can
// also pass it into subsystems that take a *slog.Logger.
func Init() *slog.Logger {
level := parseLevel(os.Getenv("CSM_LOG_LEVEL"))
handler := buildHandler(os.Getenv("CSM_LOG_FORMAT"), level)
logger := slog.New(handler)
global.Store(logger)
slog.SetDefault(logger)
return logger
}
// L returns the current global logger, initializing it on first call.
// Cheap on the hot path (single atomic load).
func L() *slog.Logger {
if l := global.Load(); l != nil {
return l
}
return Init()
}
// Helpers that mirror slog's method set on the global logger. Provided so
// call sites don't need to write log.L().Info(...) on every line.
func Debug(msg string, args ...any) { L().Debug(msg, args...) }
func Info(msg string, args ...any) { L().Info(msg, args...) }
func Warn(msg string, args ...any) { L().Warn(msg, args...) }
func Error(msg string, args ...any) { L().Error(msg, args...) }
func parseLevel(s string) slog.Level {
switch strings.ToLower(strings.TrimSpace(s)) {
case "debug":
return slog.LevelDebug
case "warn", "warning":
return slog.LevelWarn
case "error", "err":
return slog.LevelError
default:
return slog.LevelInfo
}
}
func buildHandler(format string, level slog.Level) slog.Handler {
opts := &slog.HandlerOptions{Level: level}
switch strings.ToLower(strings.TrimSpace(format)) {
case "json":
return slog.NewJSONHandler(os.Stderr, opts)
default:
return newLegacyTextHandler(os.Stderr, level)
}
}
// legacyTextHandler emits log records in CSM's historical "[timestamp] msg"
// format so callers migrating from fmt.Fprintf produce the same output. Key
// differences from slog.NewTextHandler:
//
// - No "time=... level=... msg=..." prefix; just "[YYYY-MM-DD HH:MM:SS] msg"
// - Structured fields are appended as " key=value" when present
// - Level is prepended as "WARN:" / "ERROR:" only for non-info records
//
// This lets operators mix csmlog calls with the ~180 remaining fmt.Fprintf
// call sites in the daemon without introducing a mixed-format log stream.
type legacyTextHandler struct {
w io.Writer
mu *sync.Mutex
level slog.Level
attrs []slog.Attr
group string
}
func newLegacyTextHandler(w io.Writer, level slog.Level) *legacyTextHandler {
return &legacyTextHandler{
w: w,
mu: &sync.Mutex{},
level: level,
}
}
func (h *legacyTextHandler) Enabled(_ context.Context, level slog.Level) bool {
return level >= h.level
}
func (h *legacyTextHandler) Handle(_ context.Context, r slog.Record) error {
var sb strings.Builder
sb.Grow(128)
ts := r.Time
if ts.IsZero() {
ts = time.Now()
}
sb.WriteByte('[')
sb.WriteString(ts.Format("2006-01-02 15:04:05"))
sb.WriteString("] ")
// Prepend a level marker only for non-info records so the info path
// exactly matches the legacy "[ts] msg" format.
switch r.Level {
case slog.LevelWarn:
sb.WriteString("WARN: ")
case slog.LevelError:
sb.WriteString("ERROR: ")
case slog.LevelDebug:
sb.WriteString("DEBUG: ")
}
sb.WriteString(r.Message)
// Append pre-bound attrs then record attrs as " key=value" pairs.
for _, a := range h.attrs {
writeAttr(&sb, a)
}
r.Attrs(func(a slog.Attr) bool {
writeAttr(&sb, a)
return true
})
sb.WriteByte('\n')
h.mu.Lock()
defer h.mu.Unlock()
_, err := io.WriteString(h.w, sb.String())
return err
}
func writeAttr(sb *strings.Builder, a slog.Attr) {
if a.Key == "" {
return
}
sb.WriteString(" ")
sb.WriteString(a.Key)
sb.WriteByte('=')
v := a.Value.Resolve()
s := v.String()
// Quote values that contain whitespace so the key=value pairs stay
// parseable when operators grep the log.
if strings.ContainsAny(s, " \t") {
sb.WriteByte('"')
sb.WriteString(s)
sb.WriteByte('"')
} else {
sb.WriteString(s)
}
}
func (h *legacyTextHandler) WithAttrs(attrs []slog.Attr) slog.Handler {
merged := make([]slog.Attr, 0, len(h.attrs)+len(attrs))
merged = append(merged, h.attrs...)
merged = append(merged, attrs...)
return &legacyTextHandler{
w: h.w,
mu: h.mu,
level: h.level,
attrs: merged,
group: h.group,
}
}
func (h *legacyTextHandler) WithGroup(name string) slog.Handler {
// Groups flatten into the attr key via a prefix; simple implementation
// that's sufficient for CSM's usage (we don't use groups today).
return &legacyTextHandler{
w: h.w,
mu: h.mu,
level: h.level,
attrs: h.attrs,
group: name,
}
}
// Ensure legacyTextHandler satisfies the slog.Handler contract at compile time.
var _ slog.Handler = (*legacyTextHandler)(nil)
// Package adapter renders and applies the MTA-native forward-guard rule. On
// cPanel/exim it writes a router + transport into the cPanel-preserved
// /etc/exim.conf.local include sections and regenerates exim.conf via
// buildeximconf, so the rule survives cPanel exim rebuilds. CSM is never in the
// live mail path: exim evaluates the rule and writes held copies to the
// CSM-owned quarantine Maildir itself.
//
// The exact router/transport here was validated on a real cPanel exim 4.99
// host (null-sender forward held while the local copy delivers; normal mail
// forwarded unchanged; Remove restores normal forwarding).
package adapter
import (
"bytes"
"context"
"errors"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/mailfwd/policy"
"github.com/pidginhost/csm/internal/systemdrun"
)
// Status reports whether the guard rule is currently installed.
type Status struct {
Installed bool `json:"installed"`
}
// ForwardGuard renders and (un)installs the MTA forward-guard rule.
type ForwardGuard interface {
Apply(cfg policy.Config, badIPs []string) error
Remove() error
Status() (Status, error)
RefreshBadIPs(badIPs []string) error
}
// Default on-disk locations (overridable in tests).
const (
defaultLocalConf = "/etc/exim.conf.local"
defaultBadIPsPath = "/var/lib/csm/forward_guard/bad_ips"
defaultQuarantineDir = "/var/lib/csm/forward_quarantine/held"
// transportUser delivers held copies. It must NOT be root: cPanel lists
// root on exim's never_users, so an appendfile as root fails. mailnull is
// exim's own non-root identity and exists on every cPanel host.
transportUser = "mailnull"
)
// Managed-block sentinels. Apply replaces whatever is between them, so a
// re-apply is idempotent and Remove can strip the block cleanly.
const (
routerBegin = "# CSM-FORWARD-GUARD ROUTER BEGIN (managed by csm; do not edit)"
routerEnd = "# CSM-FORWARD-GUARD ROUTER END"
transportBegin = "# CSM-FORWARD-GUARD TRANSPORT BEGIN (managed by csm; do not edit)"
transportEnd = "# CSM-FORWARD-GUARD TRANSPORT END"
)
// eximLocalSkeleton is the full set of cPanel exim.conf.local section markers,
// created when the host has no exim.conf.local yet. cPanel's buildeximconf
// injects whatever follows each marker into the generated exim.conf.
const eximLocalSkeleton = `@AUTH@
@BEGINACL@
@CONFIG@
@DIRECTOREND@
@DIRECTORMIDDLE@
@DIRECTORSTART@
@ENDACL@
@RETRYEND@
@RETRYSTART@
@REWRITE@
@ROUTEREND@
@ROUTERSTART@
@TRANSPORTEND@
@TRANSPORTMIDDLE@
@TRANSPORTSTART@
`
// EximAdapter is the cPanel/exim ForwardGuard.
type EximAdapter struct {
localConf string
badIPsPath string
quarantineDir string
// Injected side effects (real implementations on a live host; fakes in tests).
rebuild func() error // runs buildeximconf
chown func(path, user string) error // chowns the quarantine dir
mkdirAll func(path string, perm os.FileMode) error
}
// NewEximAdapter returns an adapter targeting the standard cPanel locations.
func NewEximAdapter() *EximAdapter {
return &EximAdapter{
localConf: defaultLocalConf,
badIPsPath: defaultBadIPsPath,
quarantineDir: defaultQuarantineDir,
rebuild: runBuildEximConf,
chown: chownToUser,
mkdirAll: os.MkdirAll,
}
}
// Apply installs (or refreshes) the forward-guard rule for an enabled,
// non-dry-run policy. It is transactional: on any failure the previous
// exim.conf.local is restored and exim is rebuilt back to its prior state, so a
// failed apply never leaves a half-installed rule.
func (a *EximAdapter) Apply(cfg policy.Config, badIPs []string) error {
if !cfg.Enabled {
return fmt.Errorf("forward-guard adapter: cannot apply disabled policy")
}
if cfg.DryRun {
return fmt.Errorf("forward-guard adapter: cannot apply dry-run policy")
}
router, err := a.renderRouter(cfg.HoldSignals)
if err != nil {
return err
}
prev, hadPrev, err := a.readLocalConf()
if err != nil {
return err
}
base := prev
if !hadPrev || strings.TrimSpace(base) == "" {
base = eximLocalSkeleton
}
next, err := injectBlock(base, "@ROUTERSTART@", router)
if err != nil {
return err
}
next, err = injectBlock(next, "@TRANSPORTSTART@", a.renderTransport())
if err != nil {
return err
}
// Quarantine dir must exist and be writable by the transport user before
// exim can deliver into it.
if err := a.mkdirAll(a.quarantineDir, 0700); err != nil {
return fmt.Errorf("creating quarantine dir: %w", err)
}
if err := a.chown(a.quarantineDir, transportUser); err != nil {
return fmt.Errorf("chowning quarantine dir to %s: %w", transportUser, err)
}
if err := a.writeBadIPs(badIPs); err != nil {
return err
}
if err := writeFileAtomic(a.localConf, []byte(next)); err != nil {
return err
}
if err := a.rebuild(); err != nil {
// Roll back to the prior config so mail keeps flowing as before.
if restoreErr := a.restore(prev, hadPrev); restoreErr != nil {
return fmt.Errorf("buildeximconf failed: %w; rollback failed: %v", err, restoreErr)
}
return fmt.Errorf("buildeximconf failed, rolled back: %w", err)
}
return nil
}
// Remove strips the managed blocks and rebuilds, restoring normal forwarding.
func (a *EximAdapter) Remove() error {
cur, had, err := a.readLocalConf()
if err != nil {
return err
}
if !had {
return nil // nothing installed
}
stripped := stripBlock(cur, routerBegin, routerEnd)
stripped = stripBlock(stripped, transportBegin, transportEnd)
if stripped == cur {
return nil // not installed; no rebuild needed
}
if err := writeFileAtomic(a.localConf, []byte(stripped)); err != nil {
return err
}
if err := a.rebuild(); err != nil {
if restoreErr := a.restore(cur, true); restoreErr != nil {
return fmt.Errorf("buildeximconf failed during remove: %w; rollback failed: %v", err, restoreErr)
}
return fmt.Errorf("buildeximconf failed during remove, rolled back: %w", err)
}
return nil
}
// Status reports whether both managed blocks are present.
func (a *EximAdapter) Status() (Status, error) {
cur, had, err := a.readLocalConf()
if err != nil {
return Status{}, err
}
if !had {
return Status{}, nil
}
routerInstalled := strings.Contains(cur, routerBegin)
transportInstalled := strings.Contains(cur, transportBegin)
if routerInstalled != transportInstalled {
return Status{}, fmt.Errorf("forward-guard adapter: partial install in exim.conf.local (router=%t transport=%t)", routerInstalled, transportInstalled)
}
return Status{Installed: routerInstalled}, nil
}
func (a *EximAdapter) renderRouter(sig policy.HoldSignals) (string, error) {
var clauses []string
if sig.BounceBackscatter {
clauses = append(clauses, "{eq{$sender_address}{}}")
}
if sig.BadSenderIP {
clauses = append(clauses, fmt.Sprintf("{eq{${lookup{$sender_host_address}lsearch{%s}{1}{0}}}{1}}", a.badIPsPath))
}
if len(clauses) == 0 {
// Config validation forbids enforce mode with neither signal; guard here
// so the adapter never installs a router that holds everything or nothing.
return "", fmt.Errorf("forward-guard adapter: no routing-time-enforceable signal enabled (need bounce_backscatter or bad_sender_ip)")
}
cond := fmt.Sprintf("${if and{ {def:parent_local_part} {or{ %s } } }{yes}{no}}", strings.Join(clauses, " "))
return strings.Join([]string{
routerBegin,
"csm_forward_guard:",
" driver = accept",
" domains = ! +local_domains",
" condition = " + cond,
" transport = csm_forward_hold",
routerEnd,
}, "\n"), nil
}
func (a *EximAdapter) renderTransport() string {
headers := strings.Join([]string{
"X-CSM-Forwarder: $parent_local_part@$parent_domain",
"X-CSM-Recipient: $local_part@$domain",
"X-CSM-Sender: $sender_address",
"X-CSM-Reasons: ${if eq{$sender_address}{}{bounce_backscatter}{bad_sender_ip}}",
}, "\\n")
return strings.Join([]string{
transportBegin,
"csm_forward_hold:",
" driver = appendfile",
" directory = " + a.quarantineDir,
" maildir_format",
" create_directory",
" directory_mode = 0700",
" mode = 0600",
" user = " + transportUser,
` headers_add = "` + headers + `"`,
transportEnd,
}, "\n")
}
// RefreshBadIPs rewrites only the bad-IP lookup file. exim reads the lsearch
// file at lookup time, so the change takes effect immediately with no rebuild
// or reload -- cheap to call on a schedule as the attack DB changes.
func (a *EximAdapter) RefreshBadIPs(ips []string) error {
return a.writeBadIPs(ips)
}
func (a *EximAdapter) writeBadIPs(ips []string) error {
if err := a.mkdirAll(filepath.Dir(a.badIPsPath), 0755); err != nil {
return fmt.Errorf("creating bad IP lookup dir: %w", err)
}
var buf bytes.Buffer
for _, ip := range ips {
ip = strings.TrimSpace(ip)
if ip == "" || strings.ContainsAny(ip, " \t\r\n:") {
continue // lsearch keys are one bare token per line
}
fmt.Fprintf(&buf, "%s: 1\n", ip)
}
return writeFileAtomic(a.badIPsPath, buf.Bytes())
}
func (a *EximAdapter) readLocalConf() (string, bool, error) {
data, err := os.ReadFile(a.localConf) // #nosec G304 -- operator-fixed exim.conf.local path
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return "", false, nil
}
return "", false, fmt.Errorf("reading %s: %w", a.localConf, err)
}
return string(data), true, nil
}
func (a *EximAdapter) restore(prev string, had bool) error {
if had {
if err := writeFileAtomic(a.localConf, []byte(prev)); err != nil {
return fmt.Errorf("restoring exim.conf.local: %w", err)
}
} else {
if err := os.Remove(a.localConf); err != nil && !errors.Is(err, os.ErrNotExist) {
return fmt.Errorf("removing new exim.conf.local: %w", err)
}
}
if err := a.rebuild(); err != nil {
return fmt.Errorf("rebuilding restored exim config: %w", err)
}
return nil
}
// injectBlock removes any existing managed block of the same kind, then inserts
// block immediately after the marker line. Idempotent: re-injecting yields the
// same file.
func injectBlock(conf, marker, block string) (string, error) {
// Strip a prior copy of this block so re-apply doesn't duplicate it.
begin, end := blockSentinels(block)
conf = stripBlock(conf, begin, end)
markerLine := marker + "\n"
idx := strings.Index(conf, markerLine)
if idx < 0 {
return "", fmt.Errorf("exim.conf.local missing %s marker", marker)
}
at := idx + len(markerLine)
return conf[:at] + block + "\n" + conf[at:], nil
}
var blockSentinelRe = regexp.MustCompile(`^(# CSM-FORWARD-GUARD \w+ BEGIN)`)
func blockSentinels(block string) (begin, end string) {
lines := strings.SplitN(block, "\n", 2)
begin = lines[0]
// Derive the END sentinel from the BEGIN kind.
if m := blockSentinelRe.FindStringSubmatch(begin); m != nil {
kind := strings.Fields(m[1])[2] // ROUTER or TRANSPORT
return begin, "# CSM-FORWARD-GUARD " + kind + " END"
}
return begin, ""
}
// stripBlock removes the inclusive begin..end region (and a trailing newline).
func stripBlock(conf, begin, end string) string {
for {
bi := strings.Index(conf, begin)
if bi < 0 {
return conf
}
ei := strings.Index(conf[bi:], end)
if ei < 0 {
return conf
}
stop := bi + ei + len(end)
if stop < len(conf) && conf[stop] == '\n' {
stop++
}
conf = conf[:bi] + conf[stop:]
}
}
func writeFileAtomic(path string, data []byte) error {
dir := filepath.Dir(path)
f, err := os.CreateTemp(dir, "."+filepath.Base(path)+".*.csmtmp") // #nosec G304 -- caller owns the destination path
if err != nil {
return fmt.Errorf("opening temp file for %s: %w", path, err)
}
tmp := f.Name()
if err := f.Chmod(0644); err != nil {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("chmod temp file for %s: %w", path, err)
}
if n, err := f.Write(data); err != nil {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("writing %s: %w", path, err)
} else if n != len(data) {
_ = f.Close()
_ = os.Remove(tmp)
return fmt.Errorf("writing %s: %w", path, io.ErrShortWrite)
}
if err := f.Close(); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("closing temp file for %s: %w", path, err)
}
if err := os.Rename(tmp, path); err != nil {
_ = os.Remove(tmp)
return fmt.Errorf("committing %s: %w", path, err)
}
return nil
}
// buildEximConfScript is cPanel's exim config builder (/scripts is the standard
// symlink to /usr/local/cpanel/scripts on cPanel hosts).
const (
buildEximConfScript = "/scripts/buildeximconf"
buildEximConfTimeout = 120 * time.Second
)
// buildEximConfArgv returns the argv that rebuilds exim. buildeximconf is a
// cPanel script that writes an unbounded set of system paths (exim.conf,
// exim.pl.local, cPanel state, etc.), which csm.service's ProtectSystem=strict
// allow-list cannot reasonably enumerate, so the rebuild runs as a transient
// unit forked by PID 1.
func buildEximConfArgv(systemdRunPath string) (string, []string) {
return systemdrun.Argv(systemdRunPath, systemdrun.Options{
RuntimeMax: buildEximConfTimeout,
}, buildEximConfScript)
}
type commandRunner func(context.Context, string, ...string) ([]byte, error)
func runBuildEximConf() error {
ctx, cancel := context.WithTimeout(context.Background(), buildEximConfTimeout)
defer cancel()
systemdRun, _ := exec.LookPath("systemd-run")
return runBuildEximConfCommand(ctx, systemdRun, runCommand)
}
func runBuildEximConfCommand(ctx context.Context, systemdRunPath string, run commandRunner) error {
lookPath := func(string) (string, error) {
if systemdRunPath == "" {
return "", exec.ErrNotFound
}
return systemdRunPath, nil
}
output, err := systemdrun.Run(ctx, lookPath, systemdrun.RunnerFunc(run), systemdrun.Options{
RuntimeMax: buildEximConfTimeout,
}, buildEximConfScript)
if err != nil {
return commandFailure(buildEximConfScript, output, err)
}
return nil
}
func runCommand(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- argv is built from fixed constants and the resolved
// systemd-run path; no attacker-controlled input.
return exec.CommandContext(ctx, name, args...).CombinedOutput()
}
func commandFailure(name string, output []byte, err error) error {
trimmed := strings.TrimSpace(string(output))
if trimmed == "" {
return fmt.Errorf("%s failed: %w", name, err)
}
return fmt.Errorf("%s failed: %w: %s", name, err, trimmed)
}
func chownToUser(path, user string) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// #nosec G204 -- user is the constant transportUser ("mailnull") and path is
// the operator-fixed quarantine dir; neither is attacker-controlled.
return exec.CommandContext(ctx, "chown", "-R", user+":"+user, path).Run()
}
package adapter
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/netip"
"os"
"os/exec"
"time"
"github.com/pidginhost/csm/internal/mailfwd/policy"
"github.com/pidginhost/csm/internal/systemdrun"
)
const eximMutationLimit = 8 << 20
const eximMutationTimeout = 5 * time.Minute
// The helper accepts policy data, never a destination path or command.
type eximMutation struct {
Operation string `json:"operation"`
Config policy.Config `json:"config"`
BadIPs []string `json:"bad_ips"`
}
type eximServiceAdapter struct {
*EximAdapter
mutate func(eximMutation) error
}
// NewEximServiceAdapter keeps the entire config transaction, including rollback,
// outside the daemon's mount namespace. Read-only status and lookup refreshes
// still use the daemon's existing access to the CSM state directory.
func NewEximServiceAdapter() ForwardGuard {
return &eximServiceAdapter{EximAdapter: NewEximAdapter(), mutate: runEximMutation}
}
func (a *eximServiceAdapter) Apply(cfg policy.Config, ips []string) error {
return a.mutate(eximMutation{Operation: "apply", Config: cfg, BadIPs: ips})
}
func (a *eximServiceAdapter) Remove() error {
return a.mutate(eximMutation{Operation: "remove"})
}
func runEximMutation(request eximMutation) error {
binary, err := os.Executable()
if err != nil {
return err
}
body, err := json.Marshal(request)
if err != nil {
return err
}
if len(body) > eximMutationLimit {
return fmt.Errorf("forward-guard request exceeds limit")
}
ctx, cancel := context.WithTimeout(context.Background(), eximMutationTimeout)
defer cancel()
run := func(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- command is this executable with a fixed helper subcommand, or systemd-run with fixed flags.
cmd := exec.CommandContext(ctx, name, args...)
cmd.Stdin = bytes.NewReader(body)
return cmd.CombinedOutput()
}
return executeEximMutation(ctx, binary, exec.LookPath, run)
}
func executeEximMutation(ctx context.Context, binary string, lookup systemdrun.LookPathFunc, run systemdrun.RunnerFunc) error {
output, err := systemdrun.Run(ctx, lookup, run, systemdrun.Options{Pipe: true, RuntimeMax: eximMutationTimeout}, binary, "forward-guard-worker")
if err != nil {
return commandFailure("forward-guard worker", output, err)
}
return nil
}
// HandleEximMutation handles one bounded request at the privileged CLI boundary.
// Only the fixed cPanel locations in NewEximAdapter are available to callers.
func HandleEximMutation(input io.Reader) error {
return handleEximMutation(input, lockedEximAdapter{NewEximAdapter()})
}
func handleEximMutation(input io.Reader, target ForwardGuard) error {
body, err := io.ReadAll(io.LimitReader(input, eximMutationLimit+1))
if err != nil {
return err
}
if len(body) > eximMutationLimit {
return fmt.Errorf("forward-guard request exceeds limit")
}
var request eximMutation
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&request); err != nil {
return fmt.Errorf("invalid forward-guard request: %w", err)
}
if err := decoder.Decode(new(any)); !errors.Is(err, io.EOF) {
return fmt.Errorf("forward-guard request has trailing data")
}
switch request.Operation {
case "apply":
if !request.Config.Enabled || request.Config.DryRun || (!request.Config.HoldSignals.BounceBackscatter && !request.Config.HoldSignals.BadSenderIP) {
return fmt.Errorf("forward-guard request must enable an enforceable policy")
}
for _, ip := range request.BadIPs {
addr, err := netip.ParseAddr(ip)
if err != nil || addr.Zone() != "" {
return fmt.Errorf("forward-guard request has an invalid address")
}
}
return target.Apply(request.Config, request.BadIPs)
case "remove":
if request.Config != (policy.Config{}) || len(request.BadIPs) != 0 {
return fmt.Errorf("remove request must not contain policy data")
}
return target.Remove()
default:
return fmt.Errorf("unsupported forward-guard operation")
}
}
package adapter
import (
"fmt"
"os"
"path/filepath"
"syscall"
"github.com/pidginhost/csm/internal/mailfwd/policy"
)
type lockedEximAdapter struct{ *EximAdapter }
func (a lockedEximAdapter) Apply(cfg policy.Config, ips []string) error {
return withEximMutationLock(filepath.Dir(a.badIPsPath), func() error { return a.EximAdapter.Apply(cfg, ips) })
}
func (a lockedEximAdapter) Remove() error {
return withEximMutationLock(filepath.Dir(a.badIPsPath), a.EximAdapter.Remove)
}
func withEximMutationLock(dir string, mutate func() error) error {
// #nosec G301 -- Exim reads the bad-IP lookup in this directory as the mail transport user.
if err := os.MkdirAll(dir, 0755); err != nil {
return err
}
// #nosec G304 -- fixed CSM state directory; callers cannot supply a path through the helper protocol.
file, err := os.OpenFile(filepath.Join(dir, "mutation.lock"), os.O_CREATE|os.O_RDWR, 0600)
if err != nil {
return err
}
// Keep the inode after close: unlinking a lock allows two callers to lock
// different inodes. Close releases the lock even after a failed transaction.
defer func() { _ = file.Close() }()
// #nosec G115 -- POSIX file descriptors fit in int.
if err := syscall.Flock(int(file.Fd()), syscall.LOCK_EX|syscall.LOCK_NB); err != nil {
return fmt.Errorf("forward-guard transaction lock unavailable: %w", err)
}
return mutate()
}
// Package guard glues the operator config to the pure forward-guard policy.
// It exists so internal/mailfwd/policy stays dependency-free of internal/config
// (policy is a low-level leaf; only this glue knows about both).
package guard
import (
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/mailfwd/policy"
)
// PolicyFromConfig projects the operator's forward-guard config onto the policy
// input. Only the fields the verdict needs are mapped; skip-list and retention
// are consumed by the adapter and quarantine, not by Verdict.
func PolicyFromConfig(fg config.ForwardGuardConfig) policy.Config {
return policy.Config{
Enabled: fg.Enabled,
DryRun: fg.DryRun,
HoldSignals: policy.HoldSignals{
BounceBackscatter: fg.HoldSignals.BounceBackscatter,
SpamFlagged: fg.HoldSignals.SpamFlagged,
Malware: fg.HoldSignals.Malware,
BadSenderIP: fg.HoldSignals.BadSenderIP,
AuthFail: fg.HoldSignals.AuthFail,
},
}
}
package guard
import (
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/mailfwd/adapter"
)
// Reconciler drives the exim forward-guard from operator config. The daemon
// calls Reconcile on startup and on every config reload, and RefreshBadIPs on a
// schedule. It is the only thing that decides apply-vs-remove, so the live mail
// path and the dry-run path can never both be active.
type Reconciler struct {
// Guard is the MTA adapter (nil on platforms without one).
Guard adapter.ForwardGuard
// Active gates the whole reconciler; the daemon sets it true only on a
// cPanel/exim host. When false, Reconcile/RefreshBadIPs are no-ops.
Active bool
// BadIPs supplies the current bad-sender-IP set (from the reputation DB).
BadIPs func() []string
}
// Reconcile installs the guard when it is enabled and enforcing, and removes it
// otherwise (disabled or dry-run). Dry-run never installs an MTA rule -- its
// accounting is CSM-side only.
func (r Reconciler) Reconcile(fg config.ForwardGuardConfig) error {
if !r.Active || r.Guard == nil {
return nil
}
if fg.Enabled && !fg.DryRun {
return r.Guard.Apply(PolicyFromConfig(fg), r.badIPs())
}
return r.Guard.Remove()
}
// RefreshBadIPs rewrites the bad-IP lookup file while the guard is enforcing.
// It is a no-op when the guard is not installed, so it is safe to call on a
// timer regardless of config.
func (r Reconciler) RefreshBadIPs(fg config.ForwardGuardConfig) error {
if !r.Active || r.Guard == nil || !fg.Enabled || fg.DryRun {
return nil
}
return r.Guard.RefreshBadIPs(r.badIPs())
}
func (r Reconciler) badIPs() []string {
if r.BadIPs == nil {
return nil
}
return r.BadIPs()
}
package intel
import (
"context"
"errors"
"fmt"
"strings"
"time"
)
// FlushResult reports the observable outcome of a queue flush.
type FlushResult struct {
Removed int `json:"removed"`
// Targeted is the number of message IDs submitted for removal. It is kept
// off the wire because the public response historically exposes only the
// post-action count, but callers need it to audit an unconfirmed outcome.
Targeted int `json:"-"`
}
// QueueFlusher removes safe-to-delete backscatter from the mail queue.
type QueueFlusher interface {
FlushBackscatter() (FlushResult, error)
}
// FrozenBackscatterIDs returns the message IDs of messages that are BOTH frozen
// AND null-sender (<>) in `exim -bp` output. This is the only set the flush
// touches: a frozen null-sender message is undeliverable bounce backscatter,
// so removing it cannot lose a real sender's mail or interrupt a live retry.
func FrozenBackscatterIDs(out string) []string {
var ids []string
for _, line := range strings.Split(out, "\n") {
if id, _, _, bounce, frozen, ok := parseQueueHeader(line); ok && bounce && frozen {
ids = append(ids, id)
}
}
return ids
}
// eximRemoveBatch bounds how many message IDs are passed to one `exim -Mrm`
// invocation so a huge queue cannot overflow the command line.
const eximRemoveBatch = 100
// EximQueueFlusher lists the queue, selects frozen null-sender messages, and
// removes them with `exim -Mrm`.
type EximQueueFlusher struct {
list func() ([]byte, error)
remove func(ids []string) error
}
// NewEximQueueFlusher returns a flusher backed by the live exim binary.
func NewEximQueueFlusher() *EximQueueFlusher {
return &EximQueueFlusher{list: runEximBp, remove: runEximRemove}
}
// survivorsListed bounds how many surviving message IDs are named in the
// error text so a large stuck queue cannot produce an unreadable message.
const survivorsListed = 5
// FlushBackscatter removes every frozen null-sender message currently queued.
//
// The queue, not the exit status, decides the outcome. exim -Mrm has been
// observed removing every message it was handed and still exiting non-zero,
// which reported a failure for work that had already completed and skipped
// the caller's audit record. Re-reading the queue afterwards establishes
// what actually left, so a lying exit code can neither invent a failure nor
// hide a removal that silently did nothing.
//
// Removed is the number of targeted IDs absent from the post-action queue and
// stays meaningful alongside a non-nil error. A concurrent delivery can also
// make an ID disappear, so audit callers must describe it as no longer queued,
// not necessarily deleted by this command.
func (f *EximQueueFlusher) FlushBackscatter() (FlushResult, error) {
out, err := f.list()
if err != nil {
return FlushResult{}, err
}
ids := FrozenBackscatterIDs(string(out))
if len(ids) == 0 {
return FlushResult{}, nil
}
result := FlushResult{Targeted: len(ids)}
removeErr := f.remove(ids)
after, listErr := f.list()
if listErr != nil {
confirmErr := fmt.Errorf("removal could not be confirmed: %w", listErr)
if removeErr != nil {
return result, errors.Join(fmt.Errorf("removal command failed: %w", removeErr), confirmErr)
}
return result, confirmErr
}
targeted := make(map[string]struct{}, len(ids))
for _, id := range ids {
targeted[id] = struct{}{}
}
stillQueued := make(map[string]struct{}, len(ids))
for _, line := range strings.Split(string(after), "\n") {
id, _, _, _, _, ok := parseQueueHeader(line)
if _, isTargeted := targeted[id]; ok && isTargeted {
stillQueued[id] = struct{}{}
}
}
var survived []string
for _, id := range ids {
if _, ok := stillQueued[id]; ok {
survived = append(survived, id)
}
}
result.Removed = len(ids) - len(survived)
if len(survived) == 0 {
return result, nil
}
named := survived
if len(named) > survivorsListed {
named = named[:survivorsListed]
}
msg := fmt.Sprintf("%d of %d frozen backscatter messages still queued after removal (%s)",
len(survived), len(ids), strings.Join(named, " "))
if removeErr != nil {
return result, fmt.Errorf("%s: %w", msg, removeErr)
}
return result, errors.New(msg)
}
func runEximRemove(ids []string) error {
var firstErr error
failedBatches := 0
for start := 0; start < len(ids); start += eximRemoveBatch {
end := start + eximRemoveBatch
if end > len(ids) {
end = len(ids)
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
args := append([]string{"-Mrm"}, ids[start:end]...)
_, err := runEximCommand(ctx, 30*time.Second, args...)
cancel()
if err != nil {
// Exim can remove every ID and still return non-zero. Continue so an
// unreliable status from one batch cannot prevent later attempts;
// the queue re-read decides which messages remain. Retaining only the
// first error keeps the eventual operator message bounded.
failedBatches++
if firstErr == nil {
firstErr = err
}
}
}
if firstErr != nil {
return fmt.Errorf("%d removal batch(es) reported an error; first: %w", failedBatches, firstErr)
}
return nil
}
// Package intel turns exim_mainlog deferral lines into operator-facing
// reputation signals: which outbound IPs are being throttled, by which mail
// providers, and for what stated reason. It answers "why is the queue backing
// up" without CSM sitting in the mail path -- it only reads the log the MTA
// already writes.
//
// The parser is deliberate about attacker-controlled content: a deferral line
// echoes a remote server's free-text error, so every parsed field is bounded
// and the line is rejected unless it has the exact exim deferral shape.
package intel
import (
"net"
"regexp"
"sort"
"strings"
"time"
"unicode/utf8"
"github.com/pidginhost/csm/internal/mailfwd/inventory"
)
// eximTimeLayout is exim_mainlog's default timestamp ("2026-06-07 10:15:23").
const eximTimeLayout = "2006-01-02 15:04:05"
// maxTextLen bounds the stored remote error text so a hostile MTA cannot bloat
// the report with a multi-kilobyte error string.
const (
maxAddressLen = 254
maxDomainLen = 253
maxHostLen = 253
maxIPLen = 45
maxTextLen = 240
)
// Deferral is one parsed exim "==" deferral event.
type Deferral struct {
Time time.Time
Recipient string
Domain string
Provider inventory.ProviderClass
RemoteHost string // the deferring MX (from H=)
RemoteIP string // the deferring MX address
OutboundIP string // this server's sending IP, as echoed in the error, "" if absent
SMTPCode string // "421", "" if none
ReasonCode string // bracketed code (TSS04) or a keyword label (spamhaus, rate_limit)
Text string // bounded remote error text
}
var (
hostBoundaryRe = regexp.MustCompile(`\bH=(\S{1,253})\s+\[([0-9a-fA-F:.]{1,45})\]:\s*`)
smtpCodeRe = regexp.MustCompile(`\b([45]\d{2})\b`)
reasonRe = regexp.MustCompile(`\[([A-Za-z][A-Za-z0-9]{1,7})\]`)
ipv4Re = regexp.MustCompile(`\b\d{1,3}(?:\.\d{1,3}){3}\b`)
whitespaceR = regexp.MustCompile(`\s+`)
)
// parseDeferralLine parses one exim_mainlog line. ok is false for anything that
// is not a deferral (delivery "=>", arrival "<=", failure "**", blanks, junk).
func parseDeferralLine(line string) (Deferral, bool) {
fields := strings.Fields(line)
// <date> <time> <msgid> == <recipient> ...
if len(fields) < 5 || fields[3] != "==" {
return Deferral{}, false
}
recipient := strings.Trim(fields[4], "<>")
if recipient == "" || len(recipient) > maxAddressLen || !strings.Contains(recipient, "@") {
return Deferral{}, false
}
d := Deferral{
Recipient: recipient,
Provider: inventory.ClassifyAddress(recipient),
}
if at := strings.LastIndexByte(recipient, '@'); at >= 0 && at < len(recipient)-1 {
d.Domain = strings.ToLower(recipient[at+1:])
}
if d.Domain == "" || len(d.Domain) > maxDomainLen {
return Deferral{}, false
}
if t, err := time.ParseInLocation(eximTimeLayout, fields[0]+" "+fields[1], time.Local); err == nil {
d.Time = t
}
errText := fallbackDeferralText(line)
canExtractOutboundIP := false
if m := hostBoundaryRe.FindStringSubmatchIndex(line); m != nil {
d.RemoteHost = line[m[2]:m[3]]
d.RemoteIP = parseIPLiteral(line[m[4]:m[5]])
errText = line[m[1]:]
canExtractOutboundIP = true
}
d.SMTPCode = firstSMTPCode(errText)
if canExtractOutboundIP {
d.OutboundIP = firstIPv4(errText)
}
d.ReasonCode = classifyReason(errText)
d.Text = boundText(errText)
return d, true
}
func fallbackDeferralText(line string) string {
if i := strings.Index(line, " defer ("); i >= 0 {
if j := strings.Index(line[i:], "):"); j >= 0 {
return strings.TrimSpace(line[i+j+2:])
}
}
return line
}
func parseIPLiteral(s string) string {
if len(s) > maxIPLen {
return ""
}
ip := net.ParseIP(s)
if ip == nil {
return ""
}
return s
}
func firstSMTPCode(s string) string {
for _, m := range smtpCodeRe.FindAllStringSubmatchIndex(s, -1) {
start, end := m[2], m[3]
if smtpCodeIsAddressFragment(s, start, end) {
continue
}
return s[start:end]
}
return ""
}
func smtpCodeIsAddressFragment(s string, start, end int) bool {
if start > 0 {
switch s[start-1] {
case '.', ':':
return true
}
}
if end < len(s) {
switch s[end] {
case '.', ':':
return true
}
}
return false
}
// firstIPv4 returns the first syntactically valid IPv4 address in s, or "".
func firstIPv4(s string) string {
for _, cand := range ipv4Re.FindAllString(s, -1) {
if ip := net.ParseIP(cand); ip != nil && ip.To4() != nil {
return cand
}
}
return ""
}
// classifyReason resolves a stable reason token: a bracketed provider code
// (e.g. TSS04) when present, otherwise a keyword label derived from the error
// text. Returns "" when nothing recognizable is found.
func classifyReason(errText string) string {
for _, m := range reasonRe.FindAllStringSubmatch(errText, -1) {
if validReasonCode(m[1]) {
return m[1]
}
}
low := strings.ToLower(errText)
switch {
case strings.Contains(low, "spamhaus"):
return "spamhaus"
case strings.Contains(low, "unusual rate"), strings.Contains(low, "rate limit"),
strings.Contains(low, "too many"), strings.Contains(low, "unexpected volume"):
return "rate_limit"
case strings.Contains(low, "complaint"):
return "complaint"
case strings.Contains(low, "greylist"), strings.Contains(low, "grey-list"),
strings.Contains(low, "try again later"):
return "greylist"
case strings.Contains(low, "blocked"), strings.Contains(low, "blacklist"),
strings.Contains(low, "listed"):
return "blocked"
}
return ""
}
func validReasonCode(code string) bool {
upper := strings.ToUpper(code)
if strings.HasPrefix(upper, "TLS") {
return false
}
trailingDigits := 0
for i := len(code) - 1; i >= 0; i-- {
if code[i] < '0' || code[i] > '9' {
break
}
trailingDigits++
}
return trailingDigits >= 2
}
func boundText(s string) string {
s = strings.ToValidUTF8(s, "?")
s = strings.TrimSpace(whitespaceR.ReplaceAllString(s, " "))
if len(s) <= maxTextLen {
return s
}
// Truncate on a rune boundary: the deferral text is attacker-influenced,
// so a byte-position cut could split a multi-byte rune and store invalid
// UTF-8 that JSON would then mangle.
cut := maxTextLen
for cut > 0 && !utf8.RuneStart(s[cut]) {
cut--
}
return s[:cut]
}
// ReasonCount is a reason token and how often it occurred.
type ReasonCount struct {
Code string `json:"code"`
Count int `json:"count"`
}
// ProviderCount is a provider class and how often it appeared.
type ProviderCount struct {
Provider string `json:"provider"`
Count int `json:"count"`
}
// ProviderRollup aggregates deferrals to one provider class.
type ProviderRollup struct {
Provider string `json:"provider"`
Deferrals int `json:"deferrals"`
Reasons []ReasonCount `json:"reasons"`
LastSeen time.Time `json:"last_seen,omitzero"`
Sample string `json:"sample"`
}
// OutboundIPRollup aggregates deferrals affecting one of this server's sending
// IPs -- the reputation picture for that address.
type OutboundIPRollup struct {
IP string `json:"ip"`
Deferrals int `json:"deferrals"`
Providers []ProviderCount `json:"providers"`
Reasons []ReasonCount `json:"reasons"`
LastSeen time.Time `json:"last_seen,omitzero"`
}
// Report is the aggregated deferral picture over a window of log lines.
type Report struct {
Deferrals int `json:"deferrals"`
Providers []ProviderRollup `json:"providers"`
OutboundIPs []OutboundIPRollup `json:"outbound_ips"`
}
// emptyReport returns a zero report with non-nil slices so it serializes as []
// rather than null.
func emptyReport() Report {
return Report{Providers: []ProviderRollup{}, OutboundIPs: []OutboundIPRollup{}}
}
// BuildReport parses every line and aggregates deferrals by provider and by
// outbound IP. Non-deferral lines are ignored.
func BuildReport(lines []string) Report {
provAgg := map[string]*provAccum{}
ipAgg := map[string]*ipAccum{}
deferrals := 0
for _, line := range lines {
d, ok := parseDeferralLine(line)
if !ok {
continue
}
deferrals++
prov := string(d.Provider)
pa := provAgg[prov]
if pa == nil {
pa = &provAccum{reasons: map[string]int{}}
provAgg[prov] = pa
}
pa.add(d)
if d.OutboundIP != "" {
ia := ipAgg[d.OutboundIP]
if ia == nil {
ia = &ipAccum{providers: map[string]int{}, reasons: map[string]int{}}
ipAgg[d.OutboundIP] = ia
}
ia.add(d)
}
}
rep := emptyReport()
rep.Deferrals = deferrals
for prov, pa := range provAgg {
rep.Providers = append(rep.Providers, ProviderRollup{
Provider: prov,
Deferrals: pa.count,
Reasons: sortedReasons(pa.reasons),
LastSeen: pa.lastSeen,
Sample: pa.sample,
})
}
for ip, ia := range ipAgg {
rep.OutboundIPs = append(rep.OutboundIPs, OutboundIPRollup{
IP: ip,
Deferrals: ia.count,
Providers: sortedProviders(ia.providers),
Reasons: sortedReasons(ia.reasons),
LastSeen: ia.lastSeen,
})
}
sort.Slice(rep.Providers, func(i, j int) bool {
if rep.Providers[i].Deferrals != rep.Providers[j].Deferrals {
return rep.Providers[i].Deferrals > rep.Providers[j].Deferrals
}
return rep.Providers[i].Provider < rep.Providers[j].Provider
})
sort.Slice(rep.OutboundIPs, func(i, j int) bool {
if rep.OutboundIPs[i].Deferrals != rep.OutboundIPs[j].Deferrals {
return rep.OutboundIPs[i].Deferrals > rep.OutboundIPs[j].Deferrals
}
return rep.OutboundIPs[i].IP < rep.OutboundIPs[j].IP
})
return rep
}
type provAccum struct {
count int
reasons map[string]int
lastSeen time.Time
sample string
}
func (p *provAccum) add(d Deferral) {
p.count++
if d.ReasonCode != "" {
p.reasons[d.ReasonCode]++
}
if d.Time.After(p.lastSeen) {
p.lastSeen = d.Time
}
if p.sample == "" {
p.sample = d.Text
}
}
type ipAccum struct {
count int
providers map[string]int
reasons map[string]int
lastSeen time.Time
}
func (a *ipAccum) add(d Deferral) {
a.count++
a.providers[string(d.Provider)]++
if d.ReasonCode != "" {
a.reasons[d.ReasonCode]++
}
if d.Time.After(a.lastSeen) {
a.lastSeen = d.Time
}
}
func sortedReasons(m map[string]int) []ReasonCount {
out := make([]ReasonCount, 0, len(m))
for code, n := range m {
out = append(out, ReasonCount{Code: code, Count: n})
}
sort.Slice(out, func(i, j int) bool {
if out[i].Count != out[j].Count {
return out[i].Count > out[j].Count
}
return out[i].Code < out[j].Code
})
return out
}
func sortedProviders(m map[string]int) []ProviderCount {
out := make([]ProviderCount, 0, len(m))
for prov, n := range m {
out = append(out, ProviderCount{Provider: prov, Count: n})
}
sort.Slice(out, func(i, j int) bool {
if out[i].Count != out[j].Count {
return out[i].Count > out[j].Count
}
return out[i].Provider < out[j].Provider
})
return out
}
package intel
import (
"context"
"os/exec"
"regexp"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/systemdrun"
)
// QueueComposition is the makeup of the exim queue: how much is real mail still
// trying to deliver versus null-sender bounce backscatter, how much is frozen,
// and which recipients are stuck the most.
type QueueComposition struct {
Total int `json:"total"`
Bounce int `json:"bounce"` // null-sender <> messages (backscatter)
Real int `json:"real"`
Frozen int `json:"frozen"`
FlushableBackscatter int `json:"flushable_backscatter"` // frozen AND null-sender: safe to flush
// OldestAgeSeconds is the age of the oldest queued message; left out
// when the queue is empty.
OldestAgeSeconds *int `json:"oldest_age_seconds,omitempty"`
TopRecipients []RecipientCount `json:"top_recipients"`
}
// RecipientCount is a recipient address and how many queued messages target it.
type RecipientCount struct {
Address string `json:"address"`
Count int `json:"count"`
}
const (
topRecipientLimit = 10
// Queue headers are padded for age alignment. Recipient and continuation
// lines are indented much deeper and must not be parsed as new messages.
maxQueueHeaderIndent = 4
)
var (
// A queue header line: "<age> <size> <msgid> [(user)] <sender> [*** frozen ***]".
// Accept both the legacy 6-6-2 message id and the longer base62 form exim
// 4.97+ emits (6-11-4, e.g. "1wVR8E-0000000C9po-1DDg").
queueMsgIDRe = regexp.MustCompile(`^[0-9A-Za-z]{6}-(?:[0-9A-Za-z]{6}-[0-9A-Za-z]{2}|[0-9A-Za-z]{11}-[0-9A-Za-z]{4})$`)
queueAgeRe = regexp.MustCompile(`^\d+[smhdw]$`)
queueSizeRe = regexp.MustCompile(`(?i)^\d+(?:\.\d+)?[kmgt]?$`)
)
var (
eximLookPath = exec.LookPath
eximRun = func(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- name is either the resolved systemd-run path or Exim,
// and args contain fixed flags plus validated Exim message IDs.
return exec.CommandContext(ctx, name, args...).Output()
}
)
// ParseQueue parses `exim -bp` output into a composition summary.
func ParseQueue(out string) QueueComposition {
comp := QueueComposition{TopRecipients: []RecipientCount{}}
recipients := map[string]int{}
oldestSeconds := -1
inMessage := false
for _, line := range strings.Split(out, "\n") {
if _, _, ageSec, bounce, frozen, ok := parseQueueHeader(line); ok {
comp.Total++
if bounce {
comp.Bounce++
} else {
comp.Real++
}
if frozen {
comp.Frozen++
}
if bounce && frozen {
comp.FlushableBackscatter++
}
if ageSec > oldestSeconds {
oldestSeconds = ageSec
oldest := ageSec
comp.OldestAgeSeconds = &oldest
}
inMessage = true
continue
}
if !inMessage {
continue
}
trimmed := strings.TrimSpace(line)
if trimmed == "" {
continue
}
if queueHeaderCandidate(line) {
inMessage = false
continue
}
if queueHeaderIndent(line) <= maxQueueHeaderIndent {
inMessage = false
continue
}
if !queueRecipientLine(line) {
inMessage = false
continue
}
if addr, ok := queueRecipientAddress(trimmed); ok {
recipients[addr]++
}
}
comp.TopRecipients = topRecipients(recipients)
return comp
}
func parseQueueHeader(line string) (msgID, age string, ageSeconds int, bounce, frozen, ok bool) {
if queueHeaderIndent(line) > maxQueueHeaderIndent {
return "", "", 0, false, false, false
}
fields := strings.Fields(line)
if len(fields) != 4 && len(fields) != 5 && len(fields) != 7 && len(fields) != 8 {
return "", "", 0, false, false, false
}
if !queueAgeRe.MatchString(fields[0]) || !queueSizeRe.MatchString(fields[1]) || !queueMsgIDRe.MatchString(fields[2]) {
return "", "", 0, false, false, false
}
senderIndex := 3
frozenIndex := 4
if len(fields) == 5 || len(fields) == 8 {
if !queueLocalUserField(fields[3]) {
return "", "", 0, false, false, false
}
senderIndex = 4
frozenIndex = 5
}
if len(fields) == 7 || len(fields) == 8 {
if fields[frozenIndex] != "***" || fields[frozenIndex+1] != "frozen" || fields[frozenIndex+2] != "***" {
return "", "", 0, false, false, false
}
frozen = true
}
bounce = fields[senderIndex] == "<>"
return fields[2], fields[0], AgeToSeconds(fields[0]), bounce, frozen, true
}
func queueHeaderCandidate(line string) bool {
if queueHeaderIndent(line) > maxQueueHeaderIndent {
return false
}
fields := strings.Fields(line)
return (len(fields) == 4 || len(fields) == 7) && queueMsgIDRe.MatchString(fields[2])
}
func queueRecipientLine(line string) bool {
return len(line) > 0 && (line[0] == ' ' || line[0] == '\t')
}
func queueRecipientAddress(trimmed string) (string, bool) {
fields := strings.Fields(trimmed)
// Exim prefixes an already-delivered recipient with "D"; it is not stuck.
if len(fields) == 2 && fields[0] == "D" {
return "", false
}
if len(fields) != 1 {
return "", false
}
addr := strings.Trim(fields[0], "<>")
if len(addr) > maxAddressLen || !strings.Contains(addr, "@") {
return "", false
}
return addr, true
}
func queueLocalUserField(field string) bool {
if len(field) < 3 || field[0] != '(' || field[len(field)-1] != ')' {
return false
}
return !strings.ContainsAny(field[1:len(field)-1], "() \t\r\n")
}
func queueHeaderIndent(line string) int {
for i := 0; i < len(line); i++ {
switch line[i] {
case ' ':
continue
case '\t':
return maxQueueHeaderIndent + 1
default:
return i
}
}
return len(line)
}
// AgeToSeconds converts an exim age token (e.g. "25m", "4d") to seconds. An
// unrecognized token returns 0.
func AgeToSeconds(age string) int {
if len(age) < 2 {
return 0
}
n, err := strconv.Atoi(age[:len(age)-1])
if err != nil {
return 0
}
switch age[len(age)-1] {
case 's':
return n
case 'm':
return n * 60
case 'h':
return n * 3600
case 'd':
return n * 86400
case 'w':
return n * 604800
}
return 0
}
func topRecipients(m map[string]int) []RecipientCount {
out := make([]RecipientCount, 0, len(m))
for addr, n := range m {
out = append(out, RecipientCount{Address: addr, Count: n})
}
sort.Slice(out, func(i, j int) bool {
if out[i].Count != out[j].Count {
return out[i].Count > out[j].Count
}
return out[i].Address < out[j].Address
})
if len(out) > topRecipientLimit {
out = out[:topRecipientLimit]
}
return out
}
// QueueReporter produces the queue composition for the host.
type QueueReporter interface {
Composition() (QueueComposition, error)
}
// EmptyQueueReporter yields an empty composition. It stands in on platforms
// with no exim queue (non-cPanel).
type EmptyQueueReporter struct{}
func (EmptyQueueReporter) Composition() (QueueComposition, error) {
return QueueComposition{TopRecipients: []RecipientCount{}}, nil
}
// EximQueueSource runs `exim -bp` and parses the result.
type EximQueueSource struct {
run func() ([]byte, error)
}
// NewEximQueueSource returns a source that reads the live exim queue.
func NewEximQueueSource() *EximQueueSource {
return &EximQueueSource{run: runEximBp}
}
// Composition lists the queue and summarizes it. An exim error yields an empty
// composition, not a hard error: this is a read-only visibility surface.
func (s *EximQueueSource) Composition() (QueueComposition, error) {
out, err := s.run()
if err != nil {
// exim absent or failing means no observable queue, not an API error;
// surface an empty composition rather than a 500.
return QueueComposition{TopRecipients: []RecipientCount{}}, nil //nolint:nilerr
}
return ParseQueue(string(out)), nil
}
func runEximBp() ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
return runEximCommand(ctx, 10*time.Second, "-bp")
}
func runEximCommand(ctx context.Context, runtimeMax time.Duration, args ...string) ([]byte, error) {
eximPath, err := eximLookPath("exim")
if err != nil {
return nil, err
}
return systemdrun.Run(ctx, eximLookPath, eximRun, systemdrun.Options{
Pipe: true,
RuntimeMax: runtimeMax,
}, eximPath, args...)
}
package intel
import (
"io"
"os"
"strings"
)
// eximMainLog is the cPanel/exim delivery log. Deferrals to remote providers
// are recorded here; CSM only reads it.
const eximMainLog = "/var/log/exim_mainlog"
// Default read bounds: enough recent history to see a throttle pattern without
// reading an unbounded multi-gigabyte log.
const (
defaultTailBytes = 8 << 20 // 8 MiB
defaultTailLines = 20000
)
// Reporter produces a deferral Report for the host.
type Reporter interface {
Report() (Report, error)
}
// EmptyReporter yields an empty report. It stands in on platforms with no exim
// log (non-cPanel) until their adapters land.
type EmptyReporter struct{}
func (EmptyReporter) Report() (Report, error) { return emptyReport(), nil }
// EximSource reads the tail of exim_mainlog and builds a deferral report.
type EximSource struct {
path string
tailBytes int64
tailLines int
readTail func(path string, maxBytes int64, maxLines int) []string
}
// NewEximSource returns a source reading the standard exim_mainlog location.
func NewEximSource() *EximSource {
return &EximSource{
path: eximMainLog,
tailBytes: defaultTailBytes,
tailLines: defaultTailLines,
readTail: tailLines,
}
}
// Report reads the recent tail of the log and aggregates deferrals. A missing
// or unreadable log yields an empty report, not an error: no log just means no
// observed deferrals, and this is a read-only visibility surface.
func (s *EximSource) Report() (Report, error) {
return BuildReport(s.readTail(s.path, s.tailBytes, s.tailLines)), nil
}
// tailLines returns up to maxLines trailing lines of the file, reading at most
// the last maxBytes so a huge log never loads whole into memory. A partial
// first line (from the byte-window cut) is dropped.
func tailLines(path string, maxBytes int64, maxLines int) []string {
if maxBytes <= 0 || maxLines <= 0 {
return nil
}
f, err := os.Open(path) // #nosec G304 -- fixed exim_mainlog path, operator-scoped.
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil
}
var data []byte
if info.Size() > maxBytes {
if _, err := f.Seek(-maxBytes, io.SeekEnd); err != nil {
return nil
}
data, _ = io.ReadAll(f)
// Drop the first (likely partial) line after a mid-file seek.
if nl := strings.IndexByte(string(data), '\n'); nl >= 0 {
data = data[nl+1:]
}
} else {
data, _ = io.ReadAll(f)
}
text := strings.TrimSuffix(string(data), "\n")
if text == "" {
return nil
}
lines := strings.Split(text, "\n")
if len(lines) > maxLines {
lines = lines[len(lines)-maxLines:]
}
return lines
}
package inventory
import (
"path/filepath"
"sort"
"strings"
)
// FS is the minimal filesystem surface the enumerator needs. Injected so the
// cPanel source can be tested against fixture directories without root.
type FS interface {
Glob(pattern string) ([]string, error)
ReadFile(name string) ([]byte, error)
}
// Source enumerates the forwarders configured on a host.
type Source interface {
Forwarders() ([]Forwarder, error)
}
// EmptySource reports no forwarders. It stands in on platforms whose
// enumeration is not wired yet (non-cPanel), so callers always hold a usable
// Source instead of a nil.
type EmptySource struct{}
func (EmptySource) Forwarders() ([]Forwarder, error) { return []Forwarder{}, nil }
// CPanelSource reads forwarders from cPanel's /etc/valiases directory, with
// local domains from /etc/localdomains and /etc/virtualdomains and owners from
// /etc/userdomains.
type CPanelSource struct {
fs FS
valiasGlob string
localDomainsPath string
virtualDomainsPath string
userDomainsPath string
}
// NewCPanelSource returns a source reading the standard cPanel locations.
func NewCPanelSource() *CPanelSource {
return &CPanelSource{
fs: osFS{},
valiasGlob: "/etc/valiases/*",
localDomainsPath: "/etc/localdomains",
virtualDomainsPath: "/etc/virtualdomains",
userDomainsPath: "/etc/userdomains",
}
}
// Forwarders enumerates every forwarder across all hosted domains. A missing
// or unreadable valias file is skipped, not fatal: partial inventory beats no
// inventory on a server with thousands of domains.
func (s *CPanelSource) Forwarders() ([]Forwarder, error) {
localDomains := s.loadLocalDomains()
owners := s.loadOwners()
files, err := s.fs.Glob(s.valiasGlob)
if err != nil {
return nil, err
}
var out []Forwarder
for _, path := range files {
domain := normalizeDomain(filepath.Base(path))
content, err := s.fs.ReadFile(path)
if err != nil {
continue
}
owner := owners[domain]
for _, line := range strings.Split(string(content), "\n") {
fwd, ok := parseForwarderLine(domain, line, localDomains)
if !ok {
continue
}
fwd.Owner = owner
out = append(out, fwd)
}
}
sort.Slice(out, func(i, j int) bool { return out[i].Source < out[j].Source })
return out, nil
}
// loadLocalDomains reads cPanel's local-domain files into a normalized set.
// Returns an empty set when the files are unavailable, which makes every
// destination classify as external -- the safe direction for a reputation tool
// (over-report external, never hide it).
func (s *CPanelSource) loadLocalDomains() map[string]bool {
domains := make(map[string]bool)
for _, path := range []string{s.localDomainsPath, s.virtualDomainsPath} {
if path == "" {
continue
}
content, err := s.fs.ReadFile(path)
if err != nil {
continue
}
for _, line := range strings.Split(string(content), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
domain := configDomain(line)
if domain != "" {
domains[domain] = true
}
}
}
return domains
}
// loadOwners reads /etc/userdomains ("domain: user") into a domain->owner map.
func (s *CPanelSource) loadOwners() map[string]string {
owners := make(map[string]string)
content, err := s.fs.ReadFile(s.userDomainsPath)
if err != nil {
return owners
}
for _, line := range strings.Split(string(content), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
idx := strings.IndexByte(line, ':')
if idx <= 0 {
continue
}
domain := configDomain(line[:idx])
owner := strings.TrimSpace(line[idx+1:])
if domain != "" && owner != "" {
owners[domain] = owner
}
}
return owners
}
func configDomain(line string) string {
line = strings.TrimSpace(line)
if idx := strings.IndexByte(line, ':'); idx >= 0 {
line = strings.TrimSpace(line[:idx])
}
domain := normalizeDomain(line)
if domain == "" || strings.ContainsAny(domain, " \t\r\n/:\\") {
return ""
}
return domain
}
package inventory
import (
"os"
"path/filepath"
)
// osFS is the production FS backed by the real filesystem.
type osFS struct{}
func (osFS) Glob(pattern string) ([]string, error) { return filepath.Glob(pattern) }
// ReadFile reads a file. Paths come from a fixed glob of operator-owned mail
// config directories, not from untrusted input.
func (osFS) ReadFile(name string) ([]byte, error) {
return os.ReadFile(name) // #nosec G304 -- name is a valias path from a fixed glob, operator-scoped.
}
// Package inventory enumerates mail forwarders on a host and classifies their
// destinations, so operators can see which accounts relay mail off-server and
// to which providers. It is the canonical home for forwarder-domain logic in
// CSM. Enumeration is platform-abstracted: cPanel reads /etc/valiases, other
// hosts read /etc/aliases, ~/.forward, or postfix virtual maps (wired
// incrementally). Reading the filesystem is injected so the parsing logic is
// testable without root or a live mail server.
package inventory
import (
"strings"
"golang.org/x/net/publicsuffix"
)
// ProviderClass labels a forwarder destination by where it delivers. Free
// providers (Yahoo/Gmail/Outlook) are split out because forwarding spam to
// them is what degrades a server's outbound reputation; local stays on-server.
type ProviderClass string
const (
ProviderLocal ProviderClass = "local"
ProviderYahoo ProviderClass = "yahoo"
ProviderGmail ProviderClass = "gmail"
ProviderOutlook ProviderClass = "outlook"
ProviderExternal ProviderClass = "external"
)
// freeProviderExact maps known free-provider mail domains to their class.
// Lowercase keys; lookups lowercase the input.
var freeProviderExact = map[string]ProviderClass{
"gmail.com": ProviderGmail,
"googlemail.com": ProviderGmail,
"rocketmail.com": ProviderYahoo,
"yahoo.ca": ProviderYahoo,
"yahoo.co.in": ProviderYahoo,
"yahoo.co.jp": ProviderYahoo,
"yahoo.co.uk": ProviderYahoo,
"yahoo.com": ProviderYahoo,
"yahoo.com.au": ProviderYahoo,
"yahoo.com.br": ProviderYahoo,
"yahoo.com.mx": ProviderYahoo,
"yahoo.de": ProviderYahoo,
"yahoo.es": ProviderYahoo,
"yahoo.fr": ProviderYahoo,
"yahoo.it": ProviderYahoo,
"yahoo.ro": ProviderYahoo,
"ymail.com": ProviderYahoo,
"hotmail.co.uk": ProviderOutlook,
"hotmail.com": ProviderOutlook,
"hotmail.de": ProviderOutlook,
"hotmail.fr": ProviderOutlook,
"live.co.uk": ProviderOutlook,
"live.com": ProviderOutlook,
"live.com.au": ProviderOutlook,
"live.de": ProviderOutlook,
"live.fr": ProviderOutlook,
"live.it": ProviderOutlook,
"live.ro": ProviderOutlook,
"msn.com": ProviderOutlook,
"outlook.com": ProviderOutlook,
"outlook.de": ProviderOutlook,
}
// Destination is one resolved target of a forwarder.
type Destination struct {
Address string `json:"address"`
Domain string `json:"domain"`
Provider ProviderClass `json:"provider"`
}
// Forwarder is a single source address and everything it relays to.
type Forwarder struct {
Source string `json:"source"` // local_part@domain
Domain string `json:"domain"` // hosting domain
Owner string `json:"owner"` // panel account, "" if unknown
Destinations []Destination `json:"destinations"` // address targets only
KeepLocal bool `json:"keep_local"` // also delivers to a local mailbox
ForwardOnly bool `json:"forward_only"` // only remote targets, no local copy
}
// HasExternal reports whether any destination leaves the server.
func (f Forwarder) HasExternal() bool {
for _, d := range f.Destinations {
if d.Provider != ProviderLocal {
return true
}
}
return false
}
// HasFreeProvider reports whether any destination is a free-provider mailbox
// (the reputation-risk case).
func (f Forwarder) HasFreeProvider() bool {
for _, d := range f.Destinations {
switch d.Provider {
case ProviderYahoo, ProviderGmail, ProviderOutlook:
return true
}
}
return false
}
// ClassifyAddress returns the provider class of a mail address with no
// local-domain context. Use it where the address is known to be a remote
// recipient (e.g. a deferral target parsed from exim_mainlog), so a bare
// free-provider domain classifies as that provider rather than local.
func ClassifyAddress(addr string) ProviderClass {
return classifyProvider(addr, nil)
}
// classifyProvider returns the provider class of a destination address.
// localDomains are the domains hosted on this server (lowercased keys).
func classifyProvider(addr string, localDomains map[string]bool) ProviderClass {
addr = strings.TrimSpace(addr)
at := strings.LastIndexByte(addr, '@')
if at < 0 || at >= len(addr)-1 {
// No domain: a bare local part is delivered to the local mailbox.
return ProviderLocal
}
domain := normalizeDomain(addr[at+1:])
if localDomains[domain] {
return ProviderLocal
}
if c, ok := freeProviderExact[domain]; ok {
return c
}
if registered, ok := registeredDomain(domain); ok {
if c, ok := freeProviderExact[registered]; ok {
return c
}
}
return ProviderExternal
}
func registeredDomain(domain string) (string, bool) {
registered, err := publicsuffix.EffectiveTLDPlusOne(domain)
if err != nil {
return "", false
}
return normalizeDomain(registered), true
}
func normalizeDomain(domain string) string {
return strings.TrimSuffix(strings.ToLower(strings.TrimSpace(domain)), ".")
}
// parseForwarderLine parses one alias/valias line ("local_part: dest[, dest]")
// for the given hosting domain. Returns ok=false for blanks, comments,
// malformed lines, and non-address forwarders (pipes, :fail:, :blackhole:,
// /dev/null) -- those are not mail relayed to a mailbox and carry no
// reputation risk. A line is still returned for purely local aliases so the
// inventory is complete; callers filter on HasExternal as needed.
func parseForwarderLine(domain, line string, localDomains map[string]bool) (Forwarder, bool) {
domain = normalizeDomain(domain)
if domain == "" {
return Forwarder{}, false
}
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
return Forwarder{}, false
}
colon := strings.IndexByte(line, ':')
if colon < 0 {
return Forwarder{}, false
}
localPart := strings.TrimSpace(line[:colon])
rest := strings.TrimSpace(line[colon+1:])
if localPart == "" || rest == "" {
return Forwarder{}, false
}
fwd := Forwarder{
Source: localPart + "@" + domain,
Domain: domain,
}
for _, raw := range strings.Split(rest, ",") {
dest := strings.TrimSpace(raw)
if dest == "" || !isAddressDestination(dest) {
// Pipe / :fail: / :blackhole: / /dev/null / autoresponder:
// not an address relay. Skip the target but keep the record.
continue
}
fwd.Destinations = append(fwd.Destinations, newDestination(dest, localDomains))
}
if len(fwd.Destinations) == 0 {
return Forwarder{}, false
}
localCount := 0
for _, d := range fwd.Destinations {
if d.Provider == ProviderLocal {
localCount++
}
}
fwd.KeepLocal = localCount > 0
fwd.ForwardOnly = localCount == 0
return fwd, true
}
// isAddressDestination reports whether a valias destination is a mailbox
// address (as opposed to a pipe, discard, fail, or file delivery).
func isAddressDestination(dest string) bool {
dest = strings.TrimSpace(dest)
if quotedLocalPartAddress(dest) {
return true
}
if len(dest) >= 2 && dest[0] == '"' && dest[len(dest)-1] == '"' {
dest = strings.TrimSpace(dest[1 : len(dest)-1])
}
switch {
case dest == "":
return false
case strings.HasPrefix(dest, "|"): // pipe to a program
return false
case strings.HasPrefix(dest, ":"): // :fail:, :blackhole:, :defer:
return false
case strings.HasPrefix(dest, "/"): // file / /dev/null
return false
case strings.HasPrefix(dest, "\""): // malformed quote or quoted directive
return false
}
return true
}
func quotedLocalPartAddress(dest string) bool {
if !strings.HasPrefix(dest, "\"") {
return false
}
closing := strings.IndexByte(dest[1:], '"')
if closing < 0 {
return false
}
closing++ // convert offset in dest[1:] to index in dest
return closing+1 < len(dest) && dest[closing+1] == '@'
}
func newDestination(addr string, localDomains map[string]bool) Destination {
addr = normalizeAddressDestination(addr)
domain := ""
if at := strings.LastIndexByte(addr, '@'); at >= 0 && at < len(addr)-1 {
domain = normalizeDomain(addr[at+1:])
}
return Destination{
Address: addr,
Domain: domain,
Provider: classifyProvider(addr, localDomains),
}
}
func normalizeAddressDestination(addr string) string {
addr = strings.TrimSpace(addr)
if len(addr) >= 2 && addr[0] == '"' && addr[len(addr)-1] == '"' {
inner := strings.TrimSpace(addr[1 : len(addr)-1])
if strings.Contains(inner, "@") {
return inner
}
}
return addr
}
// Package policy is the single source of truth for the forward-guard hold
// decision: given the signals observed for a message, should the external
// forward copy be held? The same Verdict function feeds both the dry-run
// "would-hold" accounting and (in Phase 2) the generated MTA rule, so the two
// can never drift apart.
//
// The verdict is layered: any enabled signal that matches holds the message.
// Holding is conservative-by-omission -- a signal whose toggle is off, or a
// message with no matching signal, is never held.
package policy
// MessageMeta is the set of signals known about a single forwarded message.
// All fields default false (unknown == not flagged), so a partially-populated
// meta can only ever reduce the chance of a hold, never invent one.
type MessageMeta struct {
NullSender bool // envelope sender is <> (bounce/backscatter)
SpamFlagged bool // SpamAssassin marked it spam
MalwareHit bool // ClamAV / YARA-X matched
SenderIPBad bool // sender IP is in the CSM attack DB / reputation
SPFFail bool
DKIMFail bool
DMARCFail bool
}
// HoldSignals toggles which layered signals are allowed to hold a message.
// Each is individually switchable so operators can roll out one signal at a time.
type HoldSignals struct {
BounceBackscatter bool
SpamFlagged bool
Malware bool
BadSenderIP bool
AuthFail bool
}
// Config is the forward-guard policy input. Enabled is the master switch; when
// off, Verdict never holds regardless of signals. DryRun does not affect the
// verdict itself -- it tells callers whether to enforce the hold or only
// account for it -- so it lives here for callers but is not read by Verdict.
type Config struct {
Enabled bool
DryRun bool
HoldSignals HoldSignals
}
// Verdict reports whether a message's external-forward copy should be held and
// the matching reason codes (in a fixed, deterministic order). Reasons are
// stable identifiers safe to surface in the UI, log, and generated MTA rule.
func Verdict(meta MessageMeta, cfg Config) (hold bool, reasons []string) {
if !cfg.Enabled {
return false, nil
}
sig := cfg.HoldSignals
// Fixed evaluation order -> deterministic reason slice (no map iteration).
if sig.AuthFail && meta.SPFFail && meta.DKIMFail && meta.DMARCFail {
reasons = append(reasons, "auth_fail")
}
if sig.BadSenderIP && meta.SenderIPBad {
reasons = append(reasons, "bad_sender_ip")
}
if sig.BounceBackscatter && meta.NullSender {
reasons = append(reasons, "bounce_backscatter")
}
if sig.Malware && meta.MalwareHit {
reasons = append(reasons, "malware")
}
if sig.SpamFlagged && meta.SpamFlagged {
reasons = append(reasons, "spam_flagged")
}
return len(reasons) > 0, reasons
}
// Package quarantine is the CSM-owned Maildir that holds external forward
// copies the forward-guard decided to withhold. The exim transport (Phase 2
// Slice C) appends held copies here with X-CSM-* control headers; CSM lists
// them, releases (re-injects to the original external recipient) or deletes,
// and prunes by age. CSM is never in the live delivery path -- exim writes the
// file, CSM only acts on it afterwards.
package quarantine
import (
"bytes"
"context"
"fmt"
"io"
"os"
"os/exec"
"path/filepath"
"sort"
"strings"
"sync/atomic"
"syscall"
"time"
)
// Control headers the exim transport adds and CSM parses. They are stripped
// before a message is re-injected so they never leak to the recipient.
const (
hdrForwarder = "X-CSM-Forwarder"
hdrRecipient = "X-CSM-Recipient"
hdrSender = "X-CSM-Sender"
hdrReasons = "X-CSM-Reasons"
hdrPrefix = "X-CSM-"
)
// HeldMessage is the operator-facing view of one held forward copy.
type HeldMessage struct {
ID string `json:"id"`
Forwarder string `json:"forwarder"` // local source address that forwarded
Recipient string `json:"recipient"` // external destination that was held
Sender string `json:"sender"` // envelope sender ("" = null-sender bounce)
Reasons []string `json:"reasons"`
HeldAt time.Time `json:"held_at"`
Size int64 `json:"size"`
}
// Quarantine manages the held-forward Maildir.
type Quarantine struct {
base string
counter atomic.Uint64
sendmail func(sender, recipient string, body []byte) error
}
// New returns a quarantine rooted at dir (a Maildir; new/cur/tmp are created on
// demand). The default re-injector shells out to the platform sendmail.
func New(dir string) *Quarantine {
return &Quarantine{base: dir, sendmail: runSendmail}
}
func (q *Quarantine) sub(name string) string { return filepath.Join(q.base, name) }
// Hold writes a held copy into the Maildir with the X-CSM-* control headers and
// returns its id (the Maildir filename). In production exim's appendfile writes
// these files; Hold produces the identical format for CSM-side paths and tests.
func (q *Quarantine) Hold(m HeldMessage, body []byte) (string, error) {
for _, d := range []string{"tmp", "new", "cur"} {
if err := os.MkdirAll(q.sub(d), 0700); err != nil {
return "", fmt.Errorf("creating maildir %s: %w", d, err)
}
}
var buf bytes.Buffer
fmt.Fprintf(&buf, "%s: %s\r\n", hdrForwarder, headerValue(m.Forwarder))
fmt.Fprintf(&buf, "%s: %s\r\n", hdrRecipient, headerValue(m.Recipient))
fmt.Fprintf(&buf, "%s: %s\r\n", hdrSender, headerValue(m.Sender))
fmt.Fprintf(&buf, "%s: %s\r\n", hdrReasons, headerValue(strings.Join(m.Reasons, ",")))
buf.Write(body)
id := fmt.Sprintf("%d.%d.csm", time.Now().UnixNano(), q.counter.Add(1))
tmp := filepath.Join(q.sub("tmp"), id)
if err := os.WriteFile(tmp, buf.Bytes(), 0600); err != nil { // #nosec G306 -- 0600 is intended
return "", fmt.Errorf("writing held message: %w", err)
}
dst := filepath.Join(q.sub("new"), id)
if err := os.Rename(tmp, dst); err != nil {
_ = os.Remove(tmp)
return "", fmt.Errorf("committing held message: %w", err)
}
return id, nil
}
// List returns every held message with parsed metadata. A missing Maildir is
// not an error: it just means nothing has been held.
func (q *Quarantine) List() ([]HeldMessage, error) {
out := []HeldMessage{}
for _, dir := range []string{"new", "cur"} {
entries, err := os.ReadDir(q.sub(dir))
if err != nil {
if os.IsNotExist(err) {
continue
}
return nil, err
}
for _, e := range entries {
if e.IsDir() {
continue
}
m, err := q.read(dir, e.Name())
if err != nil {
continue // skip unreadable/partial entries rather than fail the whole list
}
out = append(out, m)
}
}
sort.Slice(out, func(i, j int) bool { return out[i].HeldAt.Before(out[j].HeldAt) })
return out, nil
}
func (q *Quarantine) read(dir, id string) (HeldMessage, error) {
path := filepath.Join(q.sub(dir), id)
data, info, err := readMessageFile(path)
if err != nil {
return HeldMessage{}, err
}
hdr := parseControlHeaders(data)
return HeldMessage{
ID: id,
Forwarder: hdr[hdrForwarder],
Recipient: hdr[hdrRecipient],
Sender: hdr[hdrSender],
Reasons: splitReasons(hdr[hdrReasons]),
HeldAt: info.ModTime(),
Size: info.Size(),
}, nil
}
// Release re-injects the held copy to its original external recipient (operator
// decided it was a false positive), then removes it. The message is removed
// only after a successful re-injection, so a sendmail failure never loses mail.
func (q *Quarantine) Release(id string) error {
_, path := q.locate(id)
if path == "" {
return fmt.Errorf("held message %q not found", id)
}
data, _, err := readMessageFile(path)
if err != nil {
return err
}
hdr := parseControlHeaders(data)
clean := stripControlHeaders(data)
if err := q.sendmail(hdr[hdrSender], hdr[hdrRecipient], clean); err != nil {
return fmt.Errorf("re-injecting held message: %w", err)
}
return os.Remove(path)
}
// Delete discards a held copy without delivering it.
func (q *Quarantine) Delete(id string) error {
_, path := q.locate(id)
if path == "" {
return fmt.Errorf("held message %q not found", id)
}
return os.Remove(path)
}
// PruneOlderThan removes held copies older than maxAge and returns how many were
// removed.
func (q *Quarantine) PruneOlderThan(maxAge time.Duration) (int, error) {
cutoff := time.Now().Add(-maxAge)
removed := 0
for _, dir := range []string{"new", "cur"} {
entries, err := os.ReadDir(q.sub(dir))
if err != nil {
if os.IsNotExist(err) {
continue
}
return removed, err
}
for _, e := range entries {
if e.IsDir() {
continue
}
info, err := e.Info()
if err != nil {
continue
}
if info.Mode().IsRegular() && info.ModTime().Before(cutoff) {
if err := os.Remove(filepath.Join(q.sub(dir), e.Name())); err == nil {
removed++
}
}
}
}
return removed, nil
}
// CountsByForwarder returns how many held copies each forwarder produced.
func (q *Quarantine) CountsByForwarder() (map[string]int, error) {
msgs, err := q.List()
if err != nil {
return nil, err
}
counts := make(map[string]int)
for _, m := range msgs {
counts[m.Forwarder]++
}
return counts, nil
}
// pathOf returns the on-disk path of a held message id, or "" if absent.
func (q *Quarantine) pathOf(id string) string {
_, path := q.locate(id)
return path
}
func (q *Quarantine) locate(id string) (dir, path string) {
id = filepath.Base(id) // defend against traversal in a caller-supplied id
if id == "" || id == "." || id == ".." || id == string(filepath.Separator) {
return "", ""
}
for _, d := range []string{"new", "cur"} {
p := filepath.Join(q.sub(d), id)
info, err := os.Lstat(p)
if err != nil {
continue
}
if info.Mode().IsRegular() {
return d, p
}
}
return "", ""
}
func readMessageFile(path string) ([]byte, os.FileInfo, error) {
// #nosec G304 -- path is a located Maildir entry under the CSM-owned base;
// O_NOFOLLOW rejects symlink entries and symlink swaps before reading.
f, err := os.OpenFile(path, os.O_RDONLY|syscall.O_NOFOLLOW, 0)
if err != nil {
return nil, nil, err
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return nil, nil, err
}
if !info.Mode().IsRegular() {
return nil, nil, fmt.Errorf("held message %q is not a regular file", filepath.Base(path))
}
data, err := io.ReadAll(f)
if err != nil {
return nil, nil, err
}
return data, info, nil
}
// headerValue collapses a control-header value to a single safe line (no CR/LF
// so a hostile address cannot inject extra headers).
func headerValue(v string) string {
v = strings.ReplaceAll(v, "\r", "")
v = strings.ReplaceAll(v, "\n", "")
return strings.TrimSpace(v)
}
func splitReasons(v string) []string {
if strings.TrimSpace(v) == "" {
return nil
}
parts := strings.Split(v, ",")
out := make([]string, 0, len(parts))
for _, p := range parts {
if p = strings.TrimSpace(p); p != "" {
out = append(out, p)
}
}
return out
}
// headerBlockEnd returns the index just past the first blank line (the
// header/body separator), or len(data) if there is no blank line. A leading
// blank line means the header block is empty -- the rest is body. Defining the
// boundary by the first empty line (rather than searching for "\n\n") keeps
// parsing and stripping consistent on degenerate inputs.
func headerBlockEnd(data []byte) int {
off := 0
for off < len(data) {
nl := bytes.IndexByte(data[off:], '\n')
if nl < 0 {
return len(data) // no terminating newline: all headers, no body
}
line := data[off : off+nl+1]
if strings.TrimRight(string(line), "\r\n") == "" {
return off + nl + 1
}
off += nl + 1
}
return len(data)
}
func parseControlHeaders(data []byte) map[string]string {
out := map[string]string{}
header := data[:headerBlockEnd(data)]
for _, line := range strings.Split(string(header), "\n") {
line = strings.TrimRight(line, "\r")
colon := strings.IndexByte(line, ':')
if colon < 0 {
continue
}
key, ok := canonicalControlHeader(line[:colon])
if !ok {
continue
}
if _, exists := out[key]; exists {
continue
}
out[key] = headerValue(line[colon+1:])
}
return out
}
func canonicalControlHeader(name string) (string, bool) {
name = strings.TrimSpace(name)
switch {
case strings.EqualFold(name, hdrForwarder):
return hdrForwarder, true
case strings.EqualFold(name, hdrRecipient):
return hdrRecipient, true
case strings.EqualFold(name, hdrSender):
return hdrSender, true
case strings.EqualFold(name, hdrReasons):
return hdrReasons, true
case hasControlHeaderPrefix(name):
return name, true
default:
return "", false
}
}
func hasControlHeaderPrefix(line string) bool {
return len(line) >= len(hdrPrefix) && strings.EqualFold(line[:len(hdrPrefix)], hdrPrefix)
}
// stripControlHeaders removes every X-CSM-* header line, leaving the original
// message intact for re-injection. Only the header lines before the
// header/body separator are filtered; the separator and body are preserved
// byte-for-byte so the recipient sees exactly the original message.
func stripControlHeaders(data []byte) []byte {
var kept bytes.Buffer
rest := data
droppingControl := false
for len(rest) > 0 {
nl := bytes.IndexByte(rest, '\n')
var line []byte
if nl < 0 {
line, rest = rest, nil
} else {
line, rest = rest[:nl+1], rest[nl+1:]
}
trimmed := strings.TrimRight(string(line), "\r\n")
if trimmed == "" {
// Blank line ends the header block; keep it and the body verbatim.
kept.Write(line)
kept.Write(rest)
return kept.Bytes()
}
if line[0] == ' ' || line[0] == '\t' {
if droppingControl {
continue
}
kept.Write(line)
continue
}
droppingControl = hasControlHeaderPrefix(trimmed)
if droppingControl {
continue
}
kept.Write(line)
}
return kept.Bytes()
}
func runSendmail(sender, recipient string, body []byte) error {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
sender = headerValue(sender)
recipient = headerValue(recipient)
// -i: do not treat lone "." as end; -f: envelope sender ("" yields null
// sender); -- terminates options so a hostile recipient cannot be a flag.
cmd := exec.CommandContext(ctx, sendmailPath, "-i", "-f", sender, "--", recipient) // #nosec G204 -- recipient guarded by --, args are envelope addresses
cmd.Stdin = bytes.NewReader(body)
return cmd.Run()
}
const sendmailPath = "/usr/sbin/sendmail"
package maillog
import (
"errors"
"fmt"
"os"
"github.com/pidginhost/csm/internal/config"
)
// New returns the mail-log Reader appropriate for cfg.Source. A "platform"
// default file path is supplied by the caller (computed from
// internal/platform.Detect()). Pass an empty string to skip the platform
// default - useful in tests.
//
// auto - try file (must exist); fall back to journal if file missing
// but units are configured.
// file - error if the file doesn't exist.
// journal - error if the journal reader is unavailable (default builds).
func New(cfg config.MailLogsConfig, platformDefaultFile string, queue *Queue) (Reader, error) {
path := cfg.File
if path == "" {
path = platformDefaultFile
}
switch cfg.Source {
case "file":
if _, err := os.Stat(path); err != nil {
return nil, fmt.Errorf("mail_logs.source=file but %s: %w", path, err)
}
return NewFileReader(path, queue), nil
case "journal":
if len(cfg.Units) == 0 {
return nil, fmt.Errorf("mail_logs.source=journal requires units")
}
return NewJournalReader(cfg.Units, queue), nil
case "auto":
if path != "" {
if _, err := os.Stat(path); err == nil {
return NewFileReader(path, queue), nil
}
}
if len(cfg.Units) == 0 {
return nil, errors.New("mail_logs.source=auto: log file not found and no units configured for journal fallback")
}
return NewJournalReader(cfg.Units, queue), nil
default:
return nil, fmt.Errorf("mail_logs.source=%q: unknown", cfg.Source)
}
}
package maillog
import (
"bufio"
"context"
"errors"
"fmt"
"io"
"os"
"strings"
"time"
)
// maxLogLineBytes caps a single mail-log line. Real syslog lines top
// out around 8 KB; 64 KB is generous yet bounded. Without this cap a
// malformed source could ship a multi-gigabyte "line" and turn the
// reader into an OOM vector.
const maxLogLineBytes = 64 * 1024
// defaultGoneGrace is how long the source path must stay missing before
// the reader declares the source gone. Long enough to ride out a
// logrotate create-delay (rename old -> create new), short enough that an
// operator notices a real syslog->journald migration quickly.
const defaultGoneGrace = 90 * time.Second
// FileReader tails a single log file. It uses a 2-second polling loop
// because rsyslog/syslog-ng don't reliably trigger inotify events on
// every line written, and periodic path re-stat checks for log rotation.
//
// On context cancel the reader closes the output channel and returns.
type FileReader struct {
path string
queue *Queue
// onGone, when set, fires once when the source path has been missing
// continuously for goneGrace. A FileReader whose path vanishes mid-run
// (e.g. a syslog->journald migration) otherwise tails a dead fd
// silently; the callback lets the daemon surface a finding and mark the
// watcher unhealthy. onRestored fires after the path returns and the
// reader can use it again.
onGone func(error)
onRestored func()
goneGrace time.Duration
nowFn func() time.Time
// gone-tracking state, touched only by the single loop goroutine.
firstMissing time.Time
goneFired bool
restoreReady bool
}
// NewFileReader constructs a FileReader for the given path.
func NewFileReader(path string, queue *Queue) *FileReader {
return &FileReader{path: path, queue: queue, goneGrace: defaultGoneGrace, nowFn: time.Now}
}
// SetOnGone installs a callback invoked once when the source path has been
// missing for longer than the grace period. Must be called before Run.
func (r *FileReader) SetOnGone(fn func(error)) { r.onGone = fn }
// SetOnRestored installs a callback invoked once after a previously-gone
// source path returns and the reader is using it again. Must be called
// before Run.
func (r *FileReader) SetOnRestored(fn func()) { r.onRestored = fn }
// recordStat advances the missing-source state machine from one stat result.
// It fires onGone once after the path is missing past the grace period and
// arms the restore callback when the path returns.
func (r *FileReader) recordStat(missing bool, missErr error) {
if !missing {
if !r.firstMissing.IsZero() && r.goneFired {
r.restoreReady = true
}
r.firstMissing = time.Time{}
return
}
now := r.nowFn()
if r.firstMissing.IsZero() {
r.firstMissing = now
}
if !r.goneFired && !r.restoreReady && now.Sub(r.firstMissing) >= r.goneGrace {
r.goneFired = true
if r.onGone != nil {
r.onGone(missErr)
}
}
}
func (r *FileReader) recordRestored() {
if !r.restoreReady {
return
}
r.restoreReady = false
r.goneFired = false
if r.onRestored != nil {
r.onRestored()
}
}
// Run starts the polling loop and returns the line channel. Returns an
// error only when the path can't be opened at all; runtime errors during
// polling are best-effort logged via stderr but do not stop the reader.
func (r *FileReader) Run(ctx context.Context) (<-chan Line, error) {
input, err := r.open()
if err != nil {
return nil, fmt.Errorf("open %s: %w", r.path, err)
}
// A usable replacement retires the old source's current failure while
// retaining its historical uncertainty and delivery loss evidence.
r.queue.journal.outcome(false, false)
out := r.queue.channel()
go r.loop(ctx, out, input)
return out, nil
}
// Temporary EOF does not finish a log record. Both its bounded prefix and
// discard state must survive until the newline or a change of file generation.
type pendingLogLine struct {
data strings.Builder
truncated bool
onRead func(int, int, error)
}
func (p *pendingLogLine) reset() {
p.data.Reset()
p.truncated = false
}
func (p *pendingLogLine) read(ctx context.Context, r *bufio.Reader, maxBytes int) (string, bool, error) {
for {
if err := ctx.Err(); err != nil {
return "", false, err
}
chunk, err := r.ReadSlice('\n')
if p.onRead != nil {
p.onRead(len(chunk), r.Buffered(), err)
}
if len(chunk) > 0 {
switch {
case p.truncated:
// drain remainder to align on next newline
case p.data.Len()+len(chunk) <= maxBytes:
p.data.Write(chunk)
default:
if room := maxBytes - p.data.Len(); room > 0 {
p.data.Write(chunk[:room])
}
p.truncated = true
}
}
if errors.Is(err, bufio.ErrBufferFull) {
continue
}
if err != nil {
return "", false, err
}
line, truncated := p.data.String(), p.truncated
p.reset()
return line, truncated, nil
}
}
func (r *FileReader) open() (*mailFileInput, error) {
return r.openAt(0, io.SeekEnd)
}
func (r *FileReader) openRotated() (*mailFileInput, error) {
return r.openAt(0, io.SeekStart)
}
func (r *FileReader) openAt(offset int64, whence int) (*mailFileInput, error) {
f, err := os.Open(r.path) // #nosec G304 -- operator-supplied log path
if err != nil {
return nil, err
}
position, seekErr := f.Seek(offset, whence)
if seekErr != nil {
_ = f.Close()
return nil, seekErr
}
st, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, err
}
return &mailFileInput{file: f, reader: bufio.NewReader(f), ino: inode(st), offset: position, size: st.Size()}, nil
}
func (r *FileReader) loop(ctx context.Context, out chan<- Line, input *mailFileInput) {
defer close(out)
f, reader, lastIno := input.file, input.reader, input.ino
source := r.queue.file.attach(input, true)
sampleCtx, stopSample := context.WithCancel(ctx)
sampleDone := make(chan struct{})
go r.queue.file.sampleLoop(sampleCtx, sampleDone)
normal := false
defer func() {
defer r.queue.file.finish()
defer func() { r.queue.discardFile(source, bufferedMailRecords(reader)) }()
stopSample()
if !normal {
r.queue.file.outcome(fileExitFault, true)
}
r.queue.file.operation()
// The sampler still owns its descriptor until the last Stat returns.
// Closing first can turn normal shutdown into a false source failure.
<-sampleDone
closed := false
defer func() {
if !closed {
r.queue.file.outcome(fileExitFault, true)
}
}()
if err := f.Close(); err != nil {
r.queue.file.outcome(fileCloseFault, true)
}
closed = true
}()
poll := time.NewTicker(2 * time.Second)
defer poll.Stop()
// Rotation safety-net: even if every poll tick finds zero EOFs (a
// continuously-active log), still re-stat once per minute so a
// rotation that happens during a sustained write burst is caught
// without waiting for the next idle period.
rotate := time.NewTicker(time.Minute)
defer rotate.Stop()
pending := pendingLogLine{onRead: func(n, buffered int, err error) { r.queue.file.consume(source, n, buffered, err) }}
rewindOnTruncate := func() {
buffered := bufferedMailRecords(reader)
reset, err := rewindTruncatedFile(f, reader)
r.queue.file.outcome(fileCursorFault, err != nil)
if err != nil {
fmt.Fprintf(os.Stderr, "maillog file_reader %s rewind: %v\n", r.path, err)
} else if reset {
r.queue.discardFile(source, buffered)
source = r.queue.file.attach(&mailFileInput{file: f}, false)
pending.reset()
}
}
reopenOnRotate := func() {
st, err := os.Stat(r.path)
if err != nil {
// Track persistent disappearance so a source that vanishes
// mid-run (syslog->journald migration) surfaces instead of
// tailing a dead fd in silence.
r.recordStat(os.IsNotExist(err), err)
return
}
r.recordStat(false, nil)
if inode(st) == lastIno {
r.queue.file.outcome(fileOpenFault, false)
rewindOnTruncate()
r.recordRestored()
return
}
next, err := r.openRotated()
if err != nil {
r.queue.file.outcome(fileOpenFault, true)
fmt.Fprintf(os.Stderr, "maillog file_reader %s reopen: %v\n", r.path, err)
return
}
_ = f.Close()
r.queue.discardFile(source, bufferedMailRecords(reader))
f, reader, lastIno = next.file, next.reader, next.ino
source = r.queue.file.attach(next, true)
pending.reset()
r.recordRestored()
}
for {
select {
case <-ctx.Done():
normal = true
return
case <-poll.C:
r.queue.file.operation()
rewindOnTruncate()
for {
line, truncated, err := pending.read(ctx, reader, maxLogLineBytes)
if err != nil {
if ctx.Err() != nil {
normal = true
return
}
r.queue.file.outcome(fileReadFault, !errors.Is(err, io.EOF))
// Tight rotation detection: every time the reader
// hits EOF or any I/O error we re-stat the path so a
// post-rotate log is picked up by the next poll tick
// rather than waiting for the safety-net ticker.
reopenOnRotate()
break
}
r.queue.file.outcome(fileReadFault, false)
if truncated {
r.queue.lose()
r.queue.file.complete(source)
fmt.Fprintf(os.Stderr, "maillog file_reader %s: oversized line skipped at %d bytes\n", r.path, maxLogLineBytes)
continue
}
if !r.queue.sendFile(ctx, out, Line{Source: "file", Message: line}, source) {
normal = true
return
}
}
r.queue.file.idle()
case <-rotate.C:
r.queue.file.operation()
reopenOnRotate()
r.queue.file.idle()
}
}
}
func rewindTruncatedFile(f mailLogFile, reader *bufio.Reader) (bool, error) {
st, err := f.Stat()
if err != nil {
return false, err
}
offset, err := f.Seek(0, io.SeekCurrent)
if err != nil || st.Size() >= offset {
return false, err
}
if _, err := f.Seek(0, io.SeekStart); err != nil {
return false, err
}
// Read-ahead bytes belong to the generation that was truncated. Compare
// against the descriptor position so those bytes cannot conceal shrinkage.
reader.Reset(f)
return true, nil
}
package maillog
import (
"bufio"
"bytes"
"context"
"errors"
"io"
"os"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type mailLogFile interface {
io.Reader
io.Seeker
Stat() (os.FileInfo, error)
Close() error
}
type mailFileInput struct {
file mailLogFile
reader *bufio.Reader
ino uint64
offset, size int64
}
// One identity per open or rewind prevents a late Stat from joining bytes from
// different file generations. All mutable generation fields use the owner lock.
type fileSourceGeneration struct {
file mailLogFile
read, settled, end int64
physical int64
revision uint64
eof, selected bool
known, invalid bool
lagAt time.Time
}
type fileSourceQueue struct {
mu sync.Mutex
seen bool
failures fileSourceFault
sampleFailed, uncertain bool
current *fileSourceGeneration
readerAt, sampleAt time.Time
}
type fileSourceFault uint8
const (
fileReadFault fileSourceFault = 1 << iota
fileCursorFault
fileOpenFault
fileCloseFault
fileExitFault
)
func (s *fileSourceQueue) attach(input *mailFileInput, known bool) *fileSourceGeneration {
s.mu.Lock()
defer s.mu.Unlock()
g := &fileSourceGeneration{file: input.file, read: input.offset, settled: input.offset, physical: input.offset, end: input.size, known: known}
g.invalid = input.size < input.offset
s.current, s.seen = g, true
s.failures, s.sampleFailed = 0, false
s.readerAt = time.Now()
g.refreshLag(s.readerAt, false)
return g
}
func (g *fileSourceGeneration) refreshLag(now time.Time, progress bool) {
if g.end <= g.settled || g.eof && g.end <= g.read {
g.lagAt = time.Time{}
} else if progress || g.lagAt.IsZero() {
g.lagAt = now
}
}
func (s *fileSourceQueue) operation() {
s.mu.Lock()
s.readerAt = time.Now()
s.mu.Unlock()
}
func (s *fileSourceQueue) idle() {
s.mu.Lock()
s.readerAt = time.Time{}
s.mu.Unlock()
}
func (s *fileSourceQueue) outcome(fault fileSourceFault, failed bool) {
s.mu.Lock()
if failed {
s.failures |= fault
s.uncertain = true
} else {
s.failures &^= fault
}
s.mu.Unlock()
}
func (s *fileSourceQueue) consume(g *fileSourceGeneration, n, buffered int, err error) {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
g.read += int64(n)
g.physical = g.read + int64(buffered)
g.end = max(g.end, g.physical)
g.revision++
g.eof, g.selected = errors.Is(err, io.EOF), err == nil
if g.eof && !g.invalid {
g.end, g.known = g.read, true
}
if n > 0 {
s.readerAt = now
}
g.refreshLag(now, n > 0)
}
func (s *fileSourceQueue) complete(g *fileSourceGeneration) {
s.mu.Lock()
defer s.mu.Unlock()
g.settled, g.selected = g.read, false
g.revision++
s.readerAt = time.Now()
g.refreshLag(s.readerAt, true)
}
func (q *Queue) discardFile(g *fileSourceGeneration, bufferedRecords int) {
s := &q.file
s.mu.Lock()
defer s.mu.Unlock()
if g.selected {
bufferedRecords++
}
if bufferedRecords > 0 {
q.health.Lose(time.Now(), uint64(bufferedRecords))
}
s.current = nil
// Disk bytes do not establish record boundaries, and writers can append
// until close. Retain uncertainty without manufacturing a record count.
s.uncertain = true
}
func (s *fileSourceQueue) finish() {
s.mu.Lock()
s.readerAt = time.Time{}
s.mu.Unlock()
}
func bufferedMailRecords(reader *bufio.Reader) int {
// Peek only bytes already in memory; this must not read a closing source.
data, _ := reader.Peek(reader.Buffered())
return bytes.Count(data, []byte{'\n'})
}
func (s *fileSourceQueue) sample() {
s.mu.Lock()
g := s.current
if g == nil {
s.mu.Unlock()
return
}
revision := g.revision
s.sampleAt = time.Now()
s.mu.Unlock()
info, err := g.file.Stat()
s.mu.Lock()
defer s.mu.Unlock()
s.sampleAt = time.Time{}
if s.current != g {
return
}
s.sampleFailed = err != nil
s.uncertain = s.uncertain || err != nil
if err != nil || g.revision != revision {
return
}
if info.Size() < g.physical {
g.invalid, s.uncertain = true, true
} else if !g.invalid {
g.end, g.known = info.Size(), true
}
g.refreshLag(time.Now(), false)
}
func (s *fileSourceQueue) sampleLoop(ctx context.Context, done chan<- struct{}) {
defer close(done)
poll := time.NewTicker(2 * time.Second)
defer poll.Stop()
for {
select {
case <-ctx.Done():
return
case <-poll.C:
s.sample()
}
}
}
func (s *fileSourceQueue) snapshot(now time.Time) (queuehealth.Status, bool) {
s.mu.Lock()
defer s.mu.Unlock()
row := queuehealth.Status{Status: "ok", DepthUnit: "bytes", CapacityUnavailable: true, DroppedLowerBound: s.uncertain, LagBasis: "consumer_progress"}
if g := s.current; g != nil {
row.DepthUnavailable = !g.known || g.invalid || s.sampleFailed
if !row.DepthUnavailable {
row.Depth = int(max(0, g.end-g.settled))
}
if !g.lagAt.IsZero() {
row.LagSeconds = max(0, now.Sub(g.lagAt).Seconds())
}
} else {
// No descriptor left to measure. Zero bytes would be an invented
// reading of a source that is no longer open.
row.DepthUnavailable = true
}
for _, at := range []time.Time{s.readerAt, s.sampleAt} {
if !at.IsZero() {
row.ProcessingSeconds = max(row.ProcessingSeconds, now.Sub(at).Seconds())
}
}
switch {
case s.failures != 0 || s.sampleFailed:
row.Status, row.Reason = "degraded", "source_io"
case row.ProcessingSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "processing_lag"
case row.LagSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "consumer_stalled"
}
return row, s.seen
}
//go:build unix
package maillog
import (
"os"
"syscall"
)
func inode(fi os.FileInfo) uint64 {
if st, ok := fi.Sys().(*syscall.Stat_t); ok {
return st.Ino
}
return 0
}
//go:build !linux || !journal
package maillog
import (
"context"
"errors"
)
// JournalReader is a no-op stub on builds without the `journal` tag.
// The factory (T4) uses this to produce a clear error rather than
// silently downgrading to file mode when the operator explicitly asked
// for journald.
type JournalReader struct{}
// NewJournalReader satisfies the same constructor signature as the
// linux+journal build, so the factory and tests compile identically
// on default builds.
func NewJournalReader(_ []string, _ *Queue) *JournalReader { return &JournalReader{} }
func JournalSupported() bool { return false }
// ErrJournalUnsupported is returned when the build was produced without
// the `journal` tag (default builds).
var ErrJournalUnsupported = errors.New("journal reader not compiled in (build with JOURNAL=1)")
// Run returns ErrJournalUnsupported immediately on stub builds.
func (*JournalReader) Run(_ context.Context) (<-chan Line, error) {
return nil, ErrJournalUnsupported
}
package maillog
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
// The journal cursor cannot measure its unread backlog. Track actual reader
// progress and the one selected entry without taking the journal's I/O locks.
type journalSourceQueue struct {
mu sync.Mutex
seen, active, selected bool
failed, uncertain bool
at time.Time
}
func (s *journalSourceQueue) outcome(failed, uncertain bool) {
s.mu.Lock()
s.failed = failed
s.uncertain = s.uncertain || uncertain
s.mu.Unlock()
}
func (s *journalSourceQueue) snapshot(now time.Time) (queuehealth.Status, bool) {
s.mu.Lock()
defer s.mu.Unlock()
// The cursor exposes no unread count and no waiting age, so the row
// measures the current operation and says so instead of reporting a
// backlog of zero.
row := queuehealth.Status{Status: "ok", DepthUnavailable: true, CapacityUnavailable: true, DroppedLowerBound: s.uncertain, LagBasis: "unavailable"}
if s.selected {
row.InFlight = 1
}
if s.active {
row.ProcessingSeconds = max(0, now.Sub(s.at).Seconds())
}
switch {
case s.failed:
row.Status, row.Reason = "degraded", "source_io"
case row.ProcessingSeconds >= time.Minute.Seconds():
row.Status, row.Reason = "degraded", "processing_lag"
}
return row, s.seen
}
package maillog
import (
"context"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
const queueCapacity = 64
// Queue retains delivery health across reader replacements and source changes.
// The supervisor finishes the old reader before starting its replacement.
type Queue struct {
health *queuehealth.Tracker
journal journalSourceQueue
file fileSourceQueue
}
func NewQueue() *Queue {
return &Queue{health: queuehealth.New(queueCapacity, time.Minute)}
}
func (q *Queue) QueueStatuses(now time.Time) map[string]queuehealth.Status {
rows := map[string]queuehealth.Status{"delivery": q.health.Snapshot(now)}
if journal, seen := q.journal.snapshot(now); seen {
rows["journal_source"] = journal
}
if file, seen := q.file.snapshot(now); seen {
rows["file_source"] = file
}
return rows
}
func (q *Queue) channel() chan Line { return make(chan Line, queueCapacity) }
func (q *Queue) send(ctx context.Context, out chan<- Line, line Line) bool {
line.ticket = q.health.Begin(time.Now())
return q.sendTracked(ctx, out, line)
}
func (q *Queue) sendFile(ctx context.Context, out chan<- Line, line Line, source *fileSourceGeneration) bool {
line.ticket = q.health.Begin(time.Now())
q.file.complete(source)
return q.sendTracked(ctx, out, line)
}
func (q *Queue) sendTracked(ctx context.Context, out chan<- Line, line Line) bool {
select {
case out <- line:
return true
case <-ctx.Done():
line.reject()
return false
}
}
func (q *Queue) lose() { q.health.Lose(time.Now(), 1) }
func (line Line) reject() { line.ticket.Reject(time.Now()) }
// Process accounts for delivery through the entire consumer callback. False
// means delivery was abandoned. A panic is counted and continues to the owner.
func (line Line) Process(consume func(Line) bool) (completed bool) {
ticket := line.ticket
line.ticket = queuehealth.Ticket{}
ticket.Start(time.Now())
defer func() {
if completed {
ticket.Finish(time.Now())
} else {
ticket.Reject(time.Now())
}
}()
return consume(line)
}
package maillog
import (
"context"
"errors"
"sync"
"time"
)
// Supervise retries source selection and attachment until ctx is canceled or
// consume returns false. status receives nil only after successful attachment,
// and an error while unavailable. Each old reader finishes before its replacement
// starts, so source migration cannot count the same event through two readers.
func Supervise(ctx context.Context, factory func() (Reader, error), status func(error), consume func(Line) bool) {
delay := time.Second
var statusMu sync.Mutex
reported, failed := false, false
lastFailure := ""
report := func(err error) {
// File disappearance must change health even while consume is blocked.
// Serialize that callback with attachment and retry status updates.
statusMu.Lock()
defer statusMu.Unlock()
failure := ""
if err != nil {
failure = err.Error()
}
if reported && failed == (err != nil) && lastFailure == failure {
return
}
reported, failed, lastFailure = true, err != nil, failure
status(err)
}
ready := func() {
delay = time.Second
report(nil)
}
for ctx.Err() == nil {
reader, err := factory()
if err == nil {
err = consumeReader(ctx, reader, ready, report, consume)
}
if ctx.Err() != nil || errors.Is(err, context.Canceled) {
return
}
report(err)
timer := time.NewTimer(delay)
select {
case <-ctx.Done():
timer.Stop()
return
case <-timer.C:
}
delay = min(delay*2, 30*time.Second)
}
}
func consumeReader(ctx context.Context, reader Reader, ready func(), unavailable func(error), consume func(Line) bool) error {
readerCtx, cancel := context.WithCancel(ctx)
defer cancel()
gone := make(chan error, 1)
if file, ok := reader.(*FileReader); ok {
file.SetOnGone(func(err error) {
unavailable(err)
select {
case gone <- err:
default:
}
})
}
lines, err := reader.Run(readerCtx)
if err != nil {
return err
}
defer func() {
cancel()
for line := range lines {
line.reject()
}
}()
if err := ctx.Err(); err != nil {
return err
}
ready()
for {
select {
case <-ctx.Done():
return ctx.Err()
case err := <-gone:
cancel()
for line := range lines {
if ctxErr := ctx.Err(); ctxErr != nil {
line.reject()
return ctxErr
}
if !line.Process(consume) {
return context.Canceled
}
}
return err
case line, ok := <-lines:
if !ok {
return errors.New("mail log reader stopped")
}
if !line.Process(consume) {
return context.Canceled
}
}
}
}
// Package mailranges maintains an atomic in-memory map of mail-provider IP
// ranges used to exempt shared-source ranges (carrier CGNAT, mail providers)
// from firewall DoS heuristics. The package is self-contained: it embeds a
// seed snapshot so the binary ships with a usable fallback and holds no
// references to any identity-verification or bot-allowlist logic.
package mailranges
import (
_ "embed"
"encoding/json"
"errors"
"net"
"os"
"strings"
"sync/atomic"
"time"
)
// embeddedSnapshot is the seed snapshot compiled into the binary. It contains
// a handful of real Google and Microsoft outbound-mail CIDRs as placeholders.
// Task 6, Step 5 replaces this file with resolver-generated content so the
// shipped binary always carries a current snapshot.
//
//go:embed snapshot.json
var embeddedSnapshot []byte
// cacheFile is the shared JSON schema for both the on-disk cache and the
// embedded snapshot. The same decoder handles both paths.
type cacheFile struct {
RefreshedAt int64 `json:"refreshed_at"`
Providers map[string][]string `json:"providers"`
}
// providerMap is the concrete type stored in providerAtom. atomic.Value
// requires the same concrete type on every Store call; a named type satisfies
// that without an extra allocation.
type providerMap map[string][]*net.IPNet
var (
providerAtom atomic.Value // stores providerMap; nil load = never published
lastRefreshAt atomic.Int64 // Unix timestamp; 0 = never published
)
// PublishProviderSnapshot installs a new provider snapshot atomically. The
// input is deep-copied before storing so callers may modify m after returning.
// A nil or empty m publishes an empty map (clears the effective set).
func PublishProviderSnapshot(m map[string][]*net.IPNet) {
cp := make(providerMap, len(m))
for k, v := range m {
nets := make([]*net.IPNet, len(v))
for i, n := range v {
nets[i] = cloneIPNet(n)
}
cp[k] = nets
}
providerAtom.Store(cp)
}
// ProviderNets returns a flat slice of all current provider nets across every
// provider. Each call returns a freshly allocated deep copy so callers may
// inspect or modify the slice and its elements without affecting the store.
func ProviderNets() []*net.IPNet {
m := loadProviderMap()
var out []*net.IPNet
for _, nets := range m {
for _, n := range nets {
out = append(out, cloneIPNet(n))
}
}
return out
}
// ProviderSnapshot returns the current provider map as a deep copy keyed by
// provider name. The returned map and its *net.IPNet values are independent of
// the atomic store; mutations do not propagate back.
func ProviderSnapshot() map[string][]*net.IPNet {
m := loadProviderMap()
out := make(map[string][]*net.IPNet, len(m))
for k, v := range m {
nets := make([]*net.IPNet, len(v))
for i, n := range v {
nets[i] = cloneIPNet(n)
}
out[k] = nets
}
return out
}
// LastRefresh returns when PublishProviderSnapshot was last called (via
// LoadCache or an external refresh), or the zero time if it never has been.
func LastRefresh() time.Time {
ts := lastRefreshAt.Load()
if ts == 0 {
return time.Time{}
}
return time.Unix(ts, 0)
}
// LoadCache reads the on-disk provider range cache at path and publishes it.
// On any read or parse failure it falls back to embeddedSnapshot. If both
// the on-disk cache and the embedded snapshot fail to parse, LoadCache
// publishes an empty provider map and returns the error so the caller can log
// it. A missing file is a normal first-run condition and is not an error when
// the embedded snapshot parses successfully.
func LoadCache(path string) error {
data, readErr := os.ReadFile(path) // #nosec G304 -- daemon-owned state path
if readErr == nil {
m, ts, err := parseCacheData(data)
if err == nil {
PublishProviderSnapshot(m)
lastRefreshAt.Store(ts)
return nil
}
// on-disk cache unreadable; fall through to embedded snapshot
}
m, ts, embErr := parseCacheData(embeddedSnapshot)
if embErr != nil {
// Both sources failed; publish an empty map so readers get a safe zero value.
PublishProviderSnapshot(nil)
return embErr
}
PublishProviderSnapshot(m)
lastRefreshAt.Store(ts)
return nil
}
// parseCacheData decodes the shared cacheFile JSON format into a provider map.
// Malformed entries make the whole cache unusable so LoadCache can fall back to
// the embedded snapshot instead of publishing a narrowed partial set.
func parseCacheData(data []byte) (map[string][]*net.IPNet, int64, error) {
var c cacheFile
if err := json.Unmarshal(data, &c); err != nil {
return nil, 0, err
}
m := make(map[string][]*net.IPNet, len(c.Providers))
total := 0
for provider, strs := range c.Providers {
for _, s := range strs {
_, n, err := net.ParseCIDR(strings.TrimSpace(s))
if err != nil {
return nil, 0, err
}
m[provider] = append(m[provider], n)
total++
}
}
if total == 0 {
return nil, 0, errors.New("mailranges: cache has no usable provider prefixes")
}
return m, c.RefreshedAt, nil
}
func loadProviderMap() providerMap {
v := providerAtom.Load()
if v == nil {
return nil
}
return v.(providerMap) //nolint:forcetypeassert -- only providerMap is ever stored
}
// cloneIPNet returns a deep copy of n. Both the IP and Mask byte slices are
// freshly allocated so mutations to the original do not affect the clone.
func cloneIPNet(n *net.IPNet) *net.IPNet {
if n == nil {
return nil
}
ip := make(net.IP, len(n.IP))
copy(ip, n.IP)
mask := make(net.IPMask, len(n.Mask))
copy(mask, n.Mask)
return &net.IPNet{IP: ip, Mask: mask}
}
package mailranges
import (
"context"
"encoding/json"
"errors"
"fmt"
"net"
"sort"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/atomicio"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/metrics"
)
// Providers maps provider name to the SPF root domain that covers all
// outbound mail ranges for that provider. Only the two major mail providers
// whose shared IPs are relevant for DoS-exempt range building are included.
// No JSON feed URLs, ASN feeds, or CIDR lists; only SPF root domains.
var Providers = map[string]string{
"google": "_spf.google.com",
"microsoft": "spf.protection.outlook.com",
}
// staleCacheThreshold is how old the last-good cache must be before a failed
// refresh emits a warning so operators know the data may be outdated.
const staleCacheThreshold = 7 * 24 * time.Hour
// Package-level atomics for metrics. Separate from providerAtom and
// lastRefreshAt (defined in mailranges.go) to keep each concern isolated.
var (
mailrangesRefreshTotal atomic.Uint64 // number of refreshes with >= 1 success
mailrangesPrefixes atomic.Int64 // total prefix count after last successful refresh
staleCacheWarnings atomic.Uint64 // stale-cache warnings emitted (test-observable)
)
// RegisterMailrangesMetrics binds the mailranges counters and gauges to reg.
// Production callers pass metrics.Default(); tests pass metrics.NewRegistry()
// to keep registration isolated. Idempotent per-registry (the underlying
// RegisterCounterFunc/RegisterGaugeFunc panic on duplicate names; callers
// must not register the same registry twice).
func RegisterMailrangesMetrics(reg *metrics.Registry) {
reg.RegisterCounterFunc(
"csm_mailranges_refresh_total",
"Total provider range refreshes with at least one successful SPF resolve.",
func() float64 { return float64(mailrangesRefreshTotal.Load()) },
)
reg.RegisterGaugeFunc(
"csm_mailranges_prefixes",
"Current number of mail-provider prefixes across all providers.",
func() float64 { return float64(mailrangesPrefixes.Load()) },
)
reg.RegisterGaugeFunc(
"csm_mailranges_last_success_timestamp_seconds",
"Unix timestamp of the last successful provider range refresh.",
func() float64 { return float64(lastRefreshAt.Load()) },
)
}
// Refresh resolves every provider in Providers using r. For each provider that
// fails, the provider's previous last-good ranges are kept in the merged map.
// When at least one provider resolves successfully, the merged map is serialized
// and written atomically to cachePath, then published as the active snapshot and
// metrics are updated. The active snapshot is only updated after the write
// succeeds so a torn write cannot narrow the effective set.
//
// Returns (total prefixes across the merged map, nil) on full success.
// Returns (total, joinedErr) when some providers fail but at least one
// succeeds, where joinedErr wraps every failed provider's error.
// Returns (0, joinedErr) when all providers fail or the cache cannot be
// written; in the all-failure case, if the existing cache is older than 7 days,
// a stale-cache warning is emitted.
func Refresh(ctx context.Context, r Resolver, cachePath string) (int, error) {
// Take a snapshot of the current active provider map to use as the
// last-good baseline. Providers that fail this cycle keep their entry.
prev := ProviderSnapshot()
// Prime the merged map with last-good values. Successful resolves overwrite.
merged := make(map[string][]*net.IPNet, len(Providers))
for name := range Providers {
if nets, ok := prev[name]; ok {
merged[name] = nets
}
}
var errs []error
successCount := 0
for name, root := range Providers {
nets, err := ResolveSPF(ctx, r, root)
if err != nil {
csmlog.Warn("mailranges: provider SPF resolve failed",
"provider", name, "root", root, "err", err)
errs = append(errs, fmt.Errorf("provider %q: %w", name, err))
continue
}
if len(nets) == 0 {
err := fmt.Errorf("spf: no usable public prefixes")
csmlog.Warn("mailranges: provider SPF resolve failed",
"provider", name, "root", root, "err", err)
errs = append(errs, fmt.Errorf("provider %q: %w", name, err))
continue
}
merged[name] = nets
successCount++
}
// All providers failed: do not write or publish; warn if the cache is stale.
if successCount == 0 {
ts := lastRefreshAt.Load()
if ts != 0 && time.Since(time.Unix(ts, 0)) > staleCacheThreshold {
staleCacheWarnings.Add(1)
csmlog.Warn("mailranges: provider refresh failed; cache is stale",
"age_hours", int(time.Since(time.Unix(ts, 0)).Hours()))
}
return 0, errors.Join(errs...)
}
// Count total prefixes across all merged providers (successful + last-good).
total := 0
for _, nets := range merged {
total += len(nets)
}
// Build the on-disk cache payload. Sort each provider's slice for stable output.
cf := cacheFile{
RefreshedAt: time.Now().Unix(),
Providers: make(map[string][]string, len(merged)),
}
for name, nets := range merged {
strs := make([]string, 0, len(nets))
for _, n := range nets {
strs = append(strs, n.String())
}
sort.Strings(strs)
cf.Providers[name] = strs
}
data, err := json.Marshal(cf)
if err != nil {
return 0, err
}
// Atomic write first. Do not touch the active snapshot if the write fails.
if err := atomicio.AtomicWrite(cachePath, 0o600, data); err != nil {
return 0, err
}
// Publish only after the write commits.
PublishProviderSnapshot(merged)
lastRefreshAt.Store(cf.RefreshedAt)
// Update metrics only when at least one provider resolved successfully.
mailrangesRefreshTotal.Add(1)
mailrangesPrefixes.Store(int64(total))
// errs is non-nil only on partial failure; errors.Join(nil...) returns nil
// for the full-success path.
return total, errors.Join(errs...)
}
package mailranges
import (
"context"
"fmt"
"net"
"strings"
)
// Resolver is the DNS lookup interface used by ResolveSPF. The standard
// library's net.Resolver satisfies this interface; tests supply a fake.
type Resolver interface {
LookupTXT(ctx context.Context, name string) ([]string, error)
}
// maxSPFDepth is the maximum recursion depth for include/redirect chains.
// Chains deeper than this are rejected; a well-formed SPF record never needs
// more than a handful of levels, and deeper chains are a sign of misconfiguration
// or an attempt to exhaust resolver resources.
const maxSPFDepth = 10
// maxSPFLookups bounds total DNS TXT lookups during one ResolveSPF call.
// Depth alone does not bound a record that fans out to many unique includes.
const maxSPFLookups = 64
// nonPublicCIDRs lists finite (non-/0) reserved ranges that must never appear
// in a mail-provider SPF record. Default routes (0.0.0.0/0, ::/0) are handled
// separately because they contain every address and would incorrectly reject all
// public prefixes if tested with the Contains-based overlap check.
var nonPublicCIDRs = func() []*net.IPNet {
ranges := []string{
"10.0.0.0/8", // RFC 1918 private
"172.16.0.0/12", // RFC 1918 private
"192.168.0.0/16", // RFC 1918 private
"100.64.0.0/10", // RFC 6598 Shared Address Space
"127.0.0.0/8", // loopback (IPv4)
"::1/128", // loopback (IPv6)
"169.254.0.0/16", // link-local (IPv4)
"fe80::/10", // link-local (IPv6)
"fc00::/7", // ULA (unique-local)
"192.0.2.0/24", // RFC 5737 TEST-NET-1 (documentation)
"198.51.100.0/24", // RFC 5737 TEST-NET-2 (documentation)
"203.0.113.0/24", // RFC 5737 TEST-NET-3 (documentation)
"2001:db8::/32", // RFC 3849 (documentation)
}
cidrs := make([]*net.IPNet, 0, len(ranges))
for _, r := range ranges {
_, n, err := net.ParseCIDR(r)
if err != nil {
// All entries are compile-time constants; a parse error is a bug.
panic(fmt.Sprintf("mailranges: bad nonPublicCIDR %q: %v", r, err))
}
cidrs = append(cidrs, n)
}
return cidrs
}()
// isPublicPrefix reports whether n is routable public address space. It
// returns false for:
// - Any default route (prefix length 0), either IPv4 or IPv6.
// - Any prefix whose network address falls within a reserved range.
// - Any prefix that is a supernet containing a reserved range's network address,
// so attackers cannot smuggle private space inside a wide covering prefix.
//
// IPv4-mapped IPv6 prefixes (::ffff:<private>/N) are caught by the existing
// IPv4 range checks because Go's net.Contains normalises IPv4-mapped addresses
// to their 4-byte form before comparing against 4-byte reserved ranges.
func isPublicPrefix(n *net.IPNet) bool {
// Reject default routes: mask with all zero bits in either address family.
ones, bits := n.Mask.Size()
if bits > 0 && ones == 0 {
return false
}
// Bidirectional overlap check against every finite reserved range.
// reserved.Contains(n.IP): n's network address is within the reserved range.
// n.Contains(reserved.IP): n is a supernet that contains a reserved range.
for _, reserved := range nonPublicCIDRs {
if reserved.Contains(n.IP) || n.Contains(reserved.IP) {
return false
}
}
return true
}
// spfRecord holds the result of parsing one SPF TXT string. It carries only
// the token types ResolveSPF cares about; all other mechanisms (a, mx, ptr,
// exists, qualifiers) are silently ignored because they reference the domain
// itself, not static CIDR blocks useful for DoS-exempt range building.
type spfRecord struct {
nets []*net.IPNet
includes []string
redirect string // at most one per record
}
// parseSPFRecord parses a single SPF TXT record string and returns its tokens.
// It makes no network calls; all recursion lives in ResolveSPF. Records that
// do not begin with "v=spf1" are rejected. Malformed ip4:/ip6: CIDRs cause an
// error so ResolveSPF falls back to the last-good provider set rather than
// silently omitting them.
func parseSPFRecord(txt string) (spfRecord, error) {
fields := strings.Fields(txt)
if len(fields) == 0 || !strings.EqualFold(fields[0], "v=spf1") {
return spfRecord{}, fmt.Errorf("spf: not a v=spf1 record")
}
var rec spfRecord
var hasRedirect bool
for _, tok := range fields[1:] {
rawLower := strings.ToLower(tok)
if strings.HasPrefix(rawLower, "redirect=") {
// Track presence with a separate flag: an empty first redirect=
// value must not let a second redirect= slip past the guard.
if hasRedirect {
return spfRecord{}, fmt.Errorf("spf: multiple redirect= directives in one record")
}
hasRedirect = true
rec.redirect = tok[9:]
if rec.redirect == "" {
return spfRecord{}, fmt.Errorf("spf: empty redirect= directive")
}
continue
}
mech, pass, hadQualifier := spfMechanismToken(tok)
lower := strings.ToLower(mech)
if hadQualifier && strings.HasPrefix(lower, "redirect=") {
return spfRecord{}, fmt.Errorf("spf: redirect= directive cannot carry a qualifier")
}
if !pass {
continue
}
switch {
case strings.HasPrefix(lower, "ip4:"):
cidr := mech[4:]
n, err := parseSPFIPNet("ip4", cidr)
if err != nil {
return spfRecord{}, fmt.Errorf("spf: bad ip4 prefix %q: %w", cidr, err)
}
rec.nets = append(rec.nets, n)
case strings.HasPrefix(lower, "ip6:"):
cidr := mech[4:]
n, err := parseSPFIPNet("ip6", cidr)
if err != nil {
return spfRecord{}, fmt.Errorf("spf: bad ip6 prefix %q: %w", cidr, err)
}
rec.nets = append(rec.nets, n)
case strings.HasPrefix(lower, "include:"):
domain := mech[8:]
if domain == "" {
return spfRecord{}, fmt.Errorf("spf: empty include directive")
}
rec.includes = append(rec.includes, domain)
}
// All other tokens (all, a, mx, ptr, exists, qualifiers) are ignored.
}
return rec, nil
}
// spfMechanismToken returns the mechanism token after SPF qualifier handling.
// Only pass mechanisms (no qualifier or explicit '+') can contribute ranges.
func spfMechanismToken(tok string) (mech string, pass bool, hadQualifier bool) {
if tok == "" {
return "", false, false
}
switch tok[0] {
case '+':
return tok[1:], true, true
case '-', '~', '?':
return tok[1:], false, true
default:
return tok, true, false
}
}
func parseSPFIPNet(family, value string) (*net.IPNet, error) {
value = strings.TrimSpace(value)
if value == "" {
return nil, fmt.Errorf("empty prefix")
}
if strings.Contains(value, "/") {
_, n, err := net.ParseCIDR(value)
if err != nil {
return nil, err
}
return normalizeSPFIPNetFamily(family, n)
}
ip := net.ParseIP(value)
if ip == nil {
return nil, fmt.Errorf("invalid IP address")
}
switch family {
case "ip4":
ip4 := ip.To4()
if ip4 == nil {
return nil, fmt.Errorf("not an IPv4 prefix")
}
return &net.IPNet{IP: ip4, Mask: net.CIDRMask(32, 32)}, nil
case "ip6":
if ip.To4() != nil {
return nil, fmt.Errorf("not an IPv6 prefix")
}
ip16 := ip.To16()
if ip16 == nil {
return nil, fmt.Errorf("not an IPv6 prefix")
}
return &net.IPNet{IP: ip16, Mask: net.CIDRMask(128, 128)}, nil
default:
return nil, fmt.Errorf("unknown SPF IP family %q", family)
}
}
func normalizeSPFIPNetFamily(family string, n *net.IPNet) (*net.IPNet, error) {
_, bits := n.Mask.Size()
switch family {
case "ip4":
ip4 := n.IP.To4()
if ip4 == nil || bits != 32 {
return nil, fmt.Errorf("not an IPv4 prefix")
}
n.IP = ip4
return n, nil
case "ip6":
if n.IP.To4() != nil || bits != 128 {
return nil, fmt.Errorf("not an IPv6 prefix")
}
ip16 := n.IP.To16()
if ip16 == nil {
return nil, fmt.Errorf("not an IPv6 prefix")
}
n.IP = ip16
return n, nil
default:
return nil, fmt.Errorf("unknown SPF IP family %q", family)
}
}
// ResolveSPF resolves the SPF record for root, following include: and redirect=
// directives recursively up to maxSPFDepth levels. It collects all ip4: and
// ip6: prefixes found across the chain, rejects non-public prefixes, detects
// loops, and de-duplicates the result.
//
// The function returns an error rather than partial results on any anomaly
// (loop, depth exceeded, malformed TXT, non-public prefix, malformed CIDR) so
// callers can keep the last-good set instead of publishing a poisoned one.
func ResolveSPF(ctx context.Context, r Resolver, root string) ([]*net.IPNet, error) {
st := &spfResolveState{
onPath: make(map[string]bool),
memo: make(map[string][]*net.IPNet),
}
nets, err := resolveSPFRec(ctx, r, root, st, 0)
if err != nil {
return nil, err
}
return dedupNets(nets), nil
}
// spfResolveState carries the per-resolution bookkeeping shared across the
// recursive walk. onPath holds the domains on the current DFS branch so a true
// ancestor cycle errors, while memo caches fully-resolved domains so a diamond
// (the same sub-domain reached via two different include branches) resolves once
// and is reused instead of being mistaken for a loop.
type spfResolveState struct {
onPath map[string]bool
memo map[string][]*net.IPNet
lookups int
}
// resolveSPFRec is the internal recursive worker for ResolveSPF. It returns the
// prefixes collected from domain's record and everything it transitively
// references.
func resolveSPFRec(ctx context.Context, r Resolver, domain string, st *spfResolveState, depth int) ([]*net.IPNet, error) {
// A domain present on the current branch is a genuine ancestor cycle.
if st.onPath[domain] {
return nil, fmt.Errorf("spf: include/redirect loop detected at %q", domain)
}
// A fully-resolved domain reached again off-path is a diamond, not a loop;
// reuse its cached result without re-querying or re-counting depth.
if cached, ok := st.memo[domain]; ok {
return cached, nil
}
// Depth bounds only NEW resolutions; cached/cyclic cases are handled above.
if depth >= maxSPFDepth {
return nil, fmt.Errorf("spf: depth limit (%d) exceeded at %q", maxSPFDepth, domain)
}
if st.lookups >= maxSPFLookups {
return nil, fmt.Errorf("spf: DNS lookup limit (%d) exceeded at %q", maxSPFLookups, domain)
}
st.lookups++
// Check context before every network call so callers can abort the chain.
select {
case <-ctx.Done():
return nil, ctx.Err()
default:
}
txts, err := r.LookupTXT(ctx, domain)
if err != nil {
return nil, fmt.Errorf("spf: lookup %q: %w", domain, err)
}
// Select exactly one SPF record. Match only an exact "v=spf1" or a
// "v=spf1 " prefix so a malformed "v=spf1foo..." string is not selected
// over a real later record.
var spfTxt string
for _, t := range txts {
if isSPFTXT(t) {
if spfTxt != "" {
return nil, fmt.Errorf("spf: multiple v=spf1 records for %q", domain)
}
spfTxt = t
}
}
if spfTxt == "" {
return nil, fmt.Errorf("spf: no v=spf1 record for %q", domain)
}
rec, err := parseSPFRecord(spfTxt)
if err != nil {
return nil, err
}
// Enter the branch; leave it on return so siblings can revisit shared nodes.
st.onPath[domain] = true
defer delete(st.onPath, domain)
// Validate and collect ip4:/ip6: prefixes before recursing so a bad record
// causes an immediate error rather than a partial result.
var collected []*net.IPNet
for _, n := range rec.nets {
if !isPublicPrefix(n) {
return nil, fmt.Errorf("spf: non-public prefix %s in record for %q", n, domain)
}
collected = append(collected, n)
}
// Recurse into includes.
for _, inc := range rec.includes {
sub, err := resolveSPFRec(ctx, r, inc, st, depth+1)
if err != nil {
return nil, err
}
collected = append(collected, sub...)
}
// Follow redirect= (at most one per record, enforced by parseSPFRecord).
if rec.redirect != "" {
sub, err := resolveSPFRec(ctx, r, rec.redirect, st, depth+1)
if err != nil {
return nil, err
}
collected = append(collected, sub...)
}
// Cache only on full success so a partially-resolved domain is never reused.
st.memo[domain] = collected
return collected, nil
}
func isSPFTXT(txt string) bool {
lower := strings.ToLower(txt)
return lower == "v=spf1" || strings.HasPrefix(lower, "v=spf1 ")
}
// dedupNets returns nets with duplicate prefixes (by CIDR string) removed.
// Order of first occurrence is preserved.
func dedupNets(nets []*net.IPNet) []*net.IPNet {
if len(nets) == 0 {
return nil
}
seen := make(map[string]bool, len(nets))
out := make([]*net.IPNet, 0, len(nets))
for _, n := range nets {
key := n.String()
if !seen[key] {
seen[key] = true
out = append(out, n)
}
}
return out
}
// Package metrics is CSM's local OpenMetrics implementation. It exists
// so the daemon can expose a `/metrics` endpoint (ROADMAP item 4)
// without pulling in `github.com/prometheus/client_golang`, which
// would add ~20 transitive dependencies for the handful of counters,
// gauges, and histograms this project actually needs.
//
// The surface is intentionally narrow: Counter, Gauge, Histogram, and
// their labelled vector siblings. No summaries, no collectors, no
// custom exposition formats. Metric objects are safe for concurrent
// use; registration is idempotent.
package metrics
import (
"errors"
"fmt"
"io"
"math"
"sort"
"strings"
"sync"
"sync/atomic"
)
// metricType discriminates the OpenMetrics TYPE line.
type metricType string
const (
typeCounter metricType = "counter"
typeGauge metricType = "gauge"
typeHistogram metricType = "histogram"
)
// collectable is any metric that can write its exposition form to w.
// Kept unexported; registry drives it.
type collectable interface {
writeTo(w *bufferedWriter)
}
// Registry holds a set of metrics and a lazily-refreshed snapshot of
// callback-driven gauges. Scraping takes a read lock; registration
// takes a write lock. Both are short-lived.
type Registry struct {
mu sync.RWMutex
entries []registered
names map[string]struct{}
gaugeHooks []gaugeHook
counterHooks []counterHook
}
type registered struct {
name string
c collectable
}
type gaugeHook struct {
name string
help string
fn func() float64
}
type counterHook struct {
name string
help string
fn func() float64
}
// NewRegistry returns an empty registry.
func NewRegistry() *Registry {
return &Registry{names: map[string]struct{}{}}
}
// MustRegister panics if a metric of the same name is already
// registered. Daemons call this once at startup; a duplicate is a
// programming error.
func (r *Registry) MustRegister(name string, c Collector) {
r.mu.Lock()
defer r.mu.Unlock()
if _, dup := r.names[name]; dup {
panic(fmt.Sprintf("metrics: duplicate registration %q", name))
}
r.names[name] = struct{}{}
r.entries = append(r.entries, registered{name: name, c: c})
}
// RegisterGaugeFunc exposes a value produced by calling fn at scrape
// time. Useful for "ask the OS for the bbolt file size" metrics where
// caching the value would be wrong.
func (r *Registry) RegisterGaugeFunc(name, help string, fn func() float64) {
r.mu.Lock()
defer r.mu.Unlock()
if _, dup := r.names[name]; dup {
panic(fmt.Sprintf("metrics: duplicate registration %q", name))
}
r.names[name] = struct{}{}
r.gaugeHooks = append(r.gaugeHooks, gaugeHook{name: name, help: help, fn: fn})
}
// RegisterCounterFunc is the counter equivalent of RegisterGaugeFunc.
// Exposition must be monotonically non-decreasing across calls;
// callers are on the hook for that invariant.
func (r *Registry) RegisterCounterFunc(name, help string, fn func() float64) {
r.mu.Lock()
defer r.mu.Unlock()
if _, dup := r.names[name]; dup {
panic(fmt.Sprintf("metrics: duplicate registration %q", name))
}
r.names[name] = struct{}{}
r.counterHooks = append(r.counterHooks, counterHook{name: name, help: help, fn: fn})
}
// WriteOpenMetrics renders a scrape in the OpenMetrics text format.
// The output ends with the `# EOF` marker that Prometheus requires
// when served as `Content-Type: application/openmetrics-text`.
func (r *Registry) WriteOpenMetrics(w io.Writer) error {
r.mu.RLock()
defer r.mu.RUnlock()
bw := newBufferedWriter(w)
// Stable ordering matters for diffs and for human reading. Sort
// by name at scrape time; registration order is not stable across
// restarts because goroutines may register concurrently.
entries := make([]registered, len(r.entries))
copy(entries, r.entries)
sort.Slice(entries, func(i, j int) bool { return entries[i].name < entries[j].name })
for _, e := range entries {
e.c.writeTo(bw)
}
gauges := append([]gaugeHook(nil), r.gaugeHooks...)
sort.Slice(gauges, func(i, j int) bool { return gauges[i].name < gauges[j].name })
for _, h := range gauges {
bw.writeMeta(h.name, h.help, typeGauge)
bw.writeSample(h.name, nil, h.fn())
}
counters := append([]counterHook(nil), r.counterHooks...)
sort.Slice(counters, func(i, j int) bool { return counters[i].name < counters[j].name })
for _, h := range counters {
bw.writeMeta(h.name, h.help, typeCounter)
bw.writeSample(h.name, nil, h.fn())
}
bw.writeEOF()
return bw.err
}
// -----------------------------------------------------------------------
// Counter
// -----------------------------------------------------------------------
// Counter is a monotonically non-decreasing float value.
type Counter struct {
name string
help string
// Stored as bits of a float64 so Add can safely work on values
// that never need to fit into int64 (e.g., byte counts).
bits uint64
}
// NewCounter constructs an unregistered Counter. Register with
// Registry.MustRegister(name, c).
func NewCounter(name, help string) *Counter {
return &Counter{name: name, help: help}
}
// Add increments the counter by v. Panics on negative v; counters are
// monotonic by contract.
func (c *Counter) Add(v float64) {
if v < 0 {
panic(fmt.Sprintf("metrics: counter %q Add(%g): negative delta", c.name, v))
}
for {
old := atomic.LoadUint64(&c.bits)
newVal := math.Float64frombits(old) + v
if atomic.CompareAndSwapUint64(&c.bits, old, math.Float64bits(newVal)) {
return
}
}
}
// Inc adds 1.
func (c *Counter) Inc() { c.Add(1) }
// Value returns the current counter value. Useful in tests.
func (c *Counter) Value() float64 {
return math.Float64frombits(atomic.LoadUint64(&c.bits))
}
func (c *Counter) writeTo(w *bufferedWriter) {
w.writeMeta(c.name, c.help, typeCounter)
w.writeSample(c.name, nil, c.Value())
}
// -----------------------------------------------------------------------
// Gauge
// -----------------------------------------------------------------------
// Gauge is a point-in-time numeric value that can go up or down.
type Gauge struct {
name string
help string
bits uint64
}
// NewGauge constructs an unregistered Gauge.
func NewGauge(name, help string) *Gauge {
return &Gauge{name: name, help: help}
}
// Set replaces the gauge value.
func (g *Gauge) Set(v float64) {
atomic.StoreUint64(&g.bits, math.Float64bits(v))
}
// Add updates the gauge by v (may be negative).
func (g *Gauge) Add(v float64) {
for {
old := atomic.LoadUint64(&g.bits)
newVal := math.Float64frombits(old) + v
if atomic.CompareAndSwapUint64(&g.bits, old, math.Float64bits(newVal)) {
return
}
}
}
// Inc adds 1. Dec subtracts 1.
func (g *Gauge) Inc() { g.Add(1) }
func (g *Gauge) Dec() { g.Add(-1) }
// Value returns the current gauge value.
func (g *Gauge) Value() float64 {
return math.Float64frombits(atomic.LoadUint64(&g.bits))
}
func (g *Gauge) writeTo(w *bufferedWriter) {
w.writeMeta(g.name, g.help, typeGauge)
w.writeSample(g.name, nil, g.Value())
}
// -----------------------------------------------------------------------
// Histogram
// -----------------------------------------------------------------------
// Histogram is a cumulative histogram with fixed upper-bound buckets.
// Buckets must be strictly increasing; the implicit +Inf bucket is
// appended automatically.
type Histogram struct {
name string
help string
upper []float64
bucketCnt []uint64 // atomic counters per bucket (last entry is +Inf)
sum uint64 // atomic float64 bits
count uint64 // atomic total count
}
// NewHistogram constructs an unregistered Histogram. upperBounds must
// be strictly increasing; the +Inf bucket is implicit.
func NewHistogram(name, help string, upperBounds []float64) *Histogram {
for i := 1; i < len(upperBounds); i++ {
if upperBounds[i] <= upperBounds[i-1] {
panic(fmt.Sprintf("metrics: histogram %q bounds must be strictly increasing", name))
}
}
return &Histogram{
name: name,
help: help,
upper: append([]float64{}, upperBounds...),
bucketCnt: make([]uint64, len(upperBounds)+1), // one extra for +Inf
}
}
// Observe records a single sample.
func (h *Histogram) Observe(v float64) {
for i, up := range h.upper {
if v <= up {
atomic.AddUint64(&h.bucketCnt[i], 1)
}
}
// Always increment the +Inf bucket (cumulative semantics).
atomic.AddUint64(&h.bucketCnt[len(h.upper)], 1)
atomic.AddUint64(&h.count, 1)
for {
old := atomic.LoadUint64(&h.sum)
newSum := math.Float64frombits(old) + v
if atomic.CompareAndSwapUint64(&h.sum, old, math.Float64bits(newSum)) {
return
}
}
}
func (h *Histogram) writeTo(w *bufferedWriter) {
w.writeMeta(h.name, h.help, typeHistogram)
for i, up := range h.upper {
labels := []labelPair{{"le", formatFloat(up)}}
w.writeSample(h.name+"_bucket", labels, float64(atomic.LoadUint64(&h.bucketCnt[i])))
}
w.writeSample(h.name+"_bucket", []labelPair{{"le", "+Inf"}}, float64(atomic.LoadUint64(&h.bucketCnt[len(h.upper)])))
w.writeSample(h.name+"_sum", nil, math.Float64frombits(atomic.LoadUint64(&h.sum)))
w.writeSample(h.name+"_count", nil, float64(atomic.LoadUint64(&h.count)))
}
// -----------------------------------------------------------------------
// Labelled variants (vectors)
// -----------------------------------------------------------------------
// CounterVec is a family of counters indexed by a fixed set of label
// keys. Label values are provided per sample.
type CounterVec struct {
name string
help string
labelKeys []string
mu sync.Mutex
children map[string]*Counter
keys []string // insertion-ordered; stable for scrape ordering within a vec
maxChildren int // 0 = unlimited
}
// defaultVecCardinalityCap bounds the per-vec child count. Operators
// can raise the limit per metric via SetMaxChildren. The cap exists
// because user-controlled labels (per-IP, per-domain) can otherwise
// grow the children map without bound and exhaust memory.
const defaultVecCardinalityCap = 1000
// overflowLabelValue collapses all label values past the cap into a
// single sentinel bucket so cardinality stays bounded.
const overflowLabelValue = "_overflow_"
// NewCounterVec constructs a vector counter. labelKeys must be non-
// empty; use NewCounter for an unlabelled counter.
func NewCounterVec(name, help string, labelKeys []string) *CounterVec {
if len(labelKeys) == 0 {
panic(fmt.Sprintf("metrics: counter vec %q needs at least one label key", name))
}
return &CounterVec{
name: name,
help: help,
labelKeys: append([]string{}, labelKeys...),
children: map[string]*Counter{},
maxChildren: defaultVecCardinalityCap,
}
}
// SetMaxChildren overrides the per-vec cardinality cap. Pass 0 to
// disable the cap entirely (only safe when label values come from a
// fixed enum the operator controls).
func (cv *CounterVec) SetMaxChildren(n int) {
cv.mu.Lock()
defer cv.mu.Unlock()
cv.maxChildren = n
}
// ChildCount returns the current number of distinct label-value
// children, including the overflow sentinel if used. Exposed for
// tests and operator health checks.
func (cv *CounterVec) ChildCount() int {
cv.mu.Lock()
defer cv.mu.Unlock()
return len(cv.children)
}
// With returns the child counter for the given label values. Values
// are identified by the concatenation of label values; caller supplies
// them in the same order as labelKeys from NewCounterVec. Once the cap
// is reached, additional label-value combinations collapse to a
// single "_overflow_" bucket so the map stays bounded.
func (cv *CounterVec) With(values ...string) *Counter {
if len(values) != len(cv.labelKeys) {
panic(fmt.Sprintf("metrics: counter vec %q: got %d label values, want %d", cv.name, len(values), len(cv.labelKeys)))
}
key := joinLabelValues(values)
cv.mu.Lock()
defer cv.mu.Unlock()
if c, ok := cv.children[key]; ok {
return c
}
if cv.maxChildren > 0 && len(cv.children) >= explicitChildLimit(cv.maxChildren) {
key = joinLabelValues(overflowLabelValuesForArity(len(cv.labelKeys)))
if c, ok := cv.children[key]; ok {
return c
}
}
c := &Counter{name: cv.name, help: cv.help}
cv.children[key] = c
cv.keys = append(cv.keys, key)
return c
}
// overflowLabelValuesForArity returns a values slice of length n with
// every entry set to the overflow sentinel. Used so the cap path
// produces a single canonical key regardless of arity.
func overflowLabelValuesForArity(n int) []string {
out := make([]string, n)
for i := range out {
out[i] = overflowLabelValue
}
return out
}
func explicitChildLimit(maxChildren int) int {
if maxChildren <= 1 {
return 0
}
return maxChildren - 1
}
func (cv *CounterVec) writeTo(w *bufferedWriter) {
w.writeMeta(cv.name, cv.help, typeCounter)
cv.mu.Lock()
keys := append([]string(nil), cv.keys...)
childMap := make(map[string]*Counter, len(cv.children))
for k, v := range cv.children {
childMap[k] = v
}
cv.mu.Unlock()
sort.Strings(keys)
for _, k := range keys {
values := splitLabelValues(k)
pairs := make([]labelPair, len(cv.labelKeys))
for i, lk := range cv.labelKeys {
pairs[i] = labelPair{key: lk, value: values[i]}
}
w.writeSample(cv.name, pairs, childMap[k].Value())
}
}
// HistogramVec is the labelled variant of Histogram. All children
// share the same upper-bound set.
type HistogramVec struct {
name string
help string
labelKeys []string
upper []float64
mu sync.Mutex
children map[string]*Histogram
keys []string
maxChildren int
}
// NewHistogramVec constructs a vector histogram.
func NewHistogramVec(name, help string, labelKeys []string, upperBounds []float64) *HistogramVec {
if len(labelKeys) == 0 {
panic(fmt.Sprintf("metrics: histogram vec %q needs at least one label key", name))
}
for i := 1; i < len(upperBounds); i++ {
if upperBounds[i] <= upperBounds[i-1] {
panic(fmt.Sprintf("metrics: histogram vec %q bounds must be strictly increasing", name))
}
}
return &HistogramVec{
name: name,
help: help,
labelKeys: append([]string{}, labelKeys...),
upper: append([]float64{}, upperBounds...),
children: map[string]*Histogram{},
maxChildren: defaultVecCardinalityCap,
}
}
// SetMaxChildren overrides the per-vec cardinality cap.
func (hv *HistogramVec) SetMaxChildren(n int) {
hv.mu.Lock()
defer hv.mu.Unlock()
hv.maxChildren = n
}
// ChildCount returns the current number of distinct label-value children.
func (hv *HistogramVec) ChildCount() int {
hv.mu.Lock()
defer hv.mu.Unlock()
return len(hv.children)
}
// With returns the child histogram for the given label values. See
// CounterVec.With for cap semantics.
func (hv *HistogramVec) With(values ...string) *Histogram {
if len(values) != len(hv.labelKeys) {
panic(fmt.Sprintf("metrics: histogram vec %q: got %d label values, want %d", hv.name, len(values), len(hv.labelKeys)))
}
key := joinLabelValues(values)
hv.mu.Lock()
defer hv.mu.Unlock()
if h, ok := hv.children[key]; ok {
return h
}
if hv.maxChildren > 0 && len(hv.children) >= explicitChildLimit(hv.maxChildren) {
key = joinLabelValues(overflowLabelValuesForArity(len(hv.labelKeys)))
if h, ok := hv.children[key]; ok {
return h
}
}
h := &Histogram{
name: hv.name,
help: hv.help,
upper: hv.upper,
bucketCnt: make([]uint64, len(hv.upper)+1),
}
hv.children[key] = h
hv.keys = append(hv.keys, key)
return h
}
func (hv *HistogramVec) writeTo(w *bufferedWriter) {
w.writeMeta(hv.name, hv.help, typeHistogram)
hv.mu.Lock()
keys := append([]string(nil), hv.keys...)
childMap := make(map[string]*Histogram, len(hv.children))
for k, v := range hv.children {
childMap[k] = v
}
hv.mu.Unlock()
sort.Strings(keys)
for _, k := range keys {
values := splitLabelValues(k)
labelPairs := make([]labelPair, len(hv.labelKeys))
for i, lk := range hv.labelKeys {
labelPairs[i] = labelPair{key: lk, value: values[i]}
}
h := childMap[k]
for i, up := range h.upper {
pairs := append([]labelPair(nil), labelPairs...)
pairs = append(pairs, labelPair{"le", formatFloat(up)})
w.writeSample(hv.name+"_bucket", pairs, float64(atomicLoad(&h.bucketCnt[i])))
}
pairsInf := append([]labelPair(nil), labelPairs...)
pairsInf = append(pairsInf, labelPair{"le", "+Inf"})
w.writeSample(hv.name+"_bucket", pairsInf, float64(atomicLoad(&h.bucketCnt[len(h.upper)])))
w.writeSample(hv.name+"_sum", labelPairs, math.Float64frombits(atomicLoad(&h.sum)))
w.writeSample(hv.name+"_count", labelPairs, float64(atomicLoad(&h.count)))
}
}
// atomicLoad is a small helper so the HistogramVec writeTo reads match
// the base Histogram's atomic semantics without pulling sync/atomic
// across every line.
func atomicLoad(p *uint64) uint64 { return atomic.LoadUint64(p) }
// GaugeVec is the labelled variant of Gauge.
type GaugeVec struct {
name string
help string
labelKeys []string
mu sync.Mutex
children map[string]*Gauge
keys []string
maxChildren int
}
// NewGaugeVec constructs a vector gauge.
func NewGaugeVec(name, help string, labelKeys []string) *GaugeVec {
if len(labelKeys) == 0 {
panic(fmt.Sprintf("metrics: gauge vec %q needs at least one label key", name))
}
return &GaugeVec{
name: name,
help: help,
labelKeys: append([]string{}, labelKeys...),
children: map[string]*Gauge{},
maxChildren: defaultVecCardinalityCap,
}
}
// SetMaxChildren overrides the per-vec cardinality cap.
func (gv *GaugeVec) SetMaxChildren(n int) {
gv.mu.Lock()
defer gv.mu.Unlock()
gv.maxChildren = n
}
// ChildCount returns the current number of distinct label-value children.
func (gv *GaugeVec) ChildCount() int {
gv.mu.Lock()
defer gv.mu.Unlock()
return len(gv.children)
}
// With returns the child gauge for the given label values. See
// CounterVec.With for cap semantics.
func (gv *GaugeVec) With(values ...string) *Gauge {
if len(values) != len(gv.labelKeys) {
panic(fmt.Sprintf("metrics: gauge vec %q: got %d label values, want %d", gv.name, len(values), len(gv.labelKeys)))
}
key := joinLabelValues(values)
gv.mu.Lock()
defer gv.mu.Unlock()
if g, ok := gv.children[key]; ok {
return g
}
if gv.maxChildren > 0 && len(gv.children) >= explicitChildLimit(gv.maxChildren) {
key = joinLabelValues(overflowLabelValuesForArity(len(gv.labelKeys)))
if g, ok := gv.children[key]; ok {
return g
}
}
g := &Gauge{name: gv.name, help: gv.help}
gv.children[key] = g
gv.keys = append(gv.keys, key)
return g
}
func (gv *GaugeVec) writeTo(w *bufferedWriter) {
w.writeMeta(gv.name, gv.help, typeGauge)
gv.mu.Lock()
keys := append([]string(nil), gv.keys...)
childMap := make(map[string]*Gauge, len(gv.children))
for k, v := range gv.children {
childMap[k] = v
}
gv.mu.Unlock()
sort.Strings(keys)
for _, k := range keys {
values := splitLabelValues(k)
pairs := make([]labelPair, len(gv.labelKeys))
for i, lk := range gv.labelKeys {
pairs[i] = labelPair{key: lk, value: values[i]}
}
w.writeSample(gv.name, pairs, childMap[k].Value())
}
}
// -----------------------------------------------------------------------
// Internal exposition
// -----------------------------------------------------------------------
type labelPair struct {
key, value string
}
type bufferedWriter struct {
w io.Writer
err error
}
func newBufferedWriter(w io.Writer) *bufferedWriter {
return &bufferedWriter{w: w}
}
func (bw *bufferedWriter) writef(format string, args ...any) {
if bw.err != nil {
return
}
if _, err := fmt.Fprintf(bw.w, format, args...); err != nil {
bw.err = err
}
}
func (bw *bufferedWriter) writeMeta(name, help string, typ metricType) {
bw.writef("# HELP %s %s\n", name, escapeHelp(help))
bw.writef("# TYPE %s %s\n", name, typ)
}
func (bw *bufferedWriter) writeSample(name string, labels []labelPair, value float64) {
var sb strings.Builder
sb.WriteString(name)
if len(labels) > 0 {
sb.WriteByte('{')
for i, p := range labels {
if i > 0 {
sb.WriteByte(',')
}
sb.WriteString(p.key)
sb.WriteString(`="`)
sb.WriteString(escapeLabel(p.value))
sb.WriteByte('"')
}
sb.WriteByte('}')
}
bw.writef("%s %s\n", sb.String(), formatFloat(value))
}
func (bw *bufferedWriter) writeEOF() {
bw.writef("# EOF\n")
}
// Label values are joined with an ASCII unit-separator so `a|b` and
// `ab|` do not collide. Label values themselves cannot contain 0x1F
// (we would reject it in validation); joinLabelValues panics if
// someone smuggles one in.
const labelSep = "\x1f"
func joinLabelValues(vs []string) string {
for _, v := range vs {
if strings.Contains(v, labelSep) {
panic("metrics: label value contains unit-separator")
}
}
return strings.Join(vs, labelSep)
}
func splitLabelValues(k string) []string {
return strings.Split(k, labelSep)
}
// formatFloat renders a float64 in OpenMetrics-friendly form: integer
// samples look integer, floats keep precision, NaN and infinities use
// the OpenMetrics tokens.
func formatFloat(v float64) string {
switch {
case math.IsNaN(v):
return "NaN"
case math.IsInf(v, 1):
return "+Inf"
case math.IsInf(v, -1):
return "-Inf"
case v == math.Trunc(v) && math.Abs(v) < 1e15:
return fmt.Sprintf("%d", int64(v))
default:
return fmt.Sprintf("%g", v)
}
}
// escapeHelp replaces characters that break the HELP line format
// (newline, backslash). Strings the caller supplies are not attacker-
// controlled (they are developer literals), but defensive escaping
// keeps the scrape well-formed even if a future contributor gets
// creative.
func escapeHelp(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, "\n", `\n`)
return s
}
// escapeLabel is stricter: OpenMetrics requires \\, \n, and \" inside
// double-quoted label values.
func escapeLabel(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, `"`, `\"`)
s = strings.ReplaceAll(s, "\n", `\n`)
return s
}
// ErrNotRegistered can be returned by callers that look up a metric by
// name without finding it. Currently unused internally; kept for the
// external helper surface so test doubles can standardise on it.
var ErrNotRegistered = errors.New("metrics: not registered")
// -----------------------------------------------------------------------
// Process-wide default registry
// -----------------------------------------------------------------------
// defaultRegistry is the Registry the daemon shares across packages.
// Tests that need isolation should construct their own NewRegistry().
var defaultRegistry = NewRegistry()
// Default returns the process-wide Registry.
func Default() *Registry { return defaultRegistry }
// MustRegister is shorthand for Default().MustRegister.
func MustRegister(name string, c Collector) { defaultRegistry.MustRegister(name, c) }
// RegisterGaugeFunc is shorthand for Default().RegisterGaugeFunc.
func RegisterGaugeFunc(name, help string, fn func() float64) {
defaultRegistry.RegisterGaugeFunc(name, help, fn)
}
// RegisterCounterFunc is shorthand for Default().RegisterCounterFunc.
func RegisterCounterFunc(name, help string, fn func() float64) {
defaultRegistry.RegisterCounterFunc(name, help, fn)
}
// WriteOpenMetrics is shorthand for Default().WriteOpenMetrics.
func WriteOpenMetrics(w io.Writer) error {
return defaultRegistry.WriteOpenMetrics(w)
}
// Collector is the type accepted by Registry.MustRegister. Exported
// so external packages have a name for the interface, even though the
// useful method on it is unexported (only metrics-package types can
// implement it, which is the intent).
type Collector = collectable
package metrics
import (
"runtime"
"sync"
)
var runtimeMetricsOnce sync.Once
// goRuntimeCollector emits Go runtime memory and scheduler stats. It reads
// runtime.ReadMemStats once per scrape (ReadMemStats briefly stops the world,
// so a single read for the whole family is deliberate -- do not split it into
// per-metric gauge hooks).
type goRuntimeCollector struct{}
func (goRuntimeCollector) writeTo(bw *bufferedWriter) {
var m runtime.MemStats
runtime.ReadMemStats(&m)
g := func(name, help string, v float64) {
bw.writeMeta(name, help, typeGauge)
bw.writeSample(name, nil, v)
}
g("go_memstats_heap_alloc_bytes", "Heap bytes allocated and still in use.", float64(m.HeapAlloc))
g("go_memstats_heap_inuse_bytes", "Heap bytes in in-use spans.", float64(m.HeapInuse))
g("go_memstats_heap_idle_bytes", "Heap bytes idle, waiting to be used.", float64(m.HeapIdle))
g("go_memstats_heap_released_bytes", "Heap bytes released to the OS.", float64(m.HeapReleased))
g("go_memstats_heap_sys_bytes", "Heap bytes obtained from the OS.", float64(m.HeapSys))
g("go_memstats_heap_objects", "Number of currently allocated heap objects.", float64(m.HeapObjects))
g("go_memstats_stack_inuse_bytes", "Bytes in use by the stack allocator.", float64(m.StackInuse))
g("go_memstats_sys_bytes", "Total bytes obtained from the OS.", float64(m.Sys))
g("go_memstats_next_gc_bytes", "Heap size target for the next GC cycle.", float64(m.NextGC))
g("go_memstats_gc_cpu_fraction", "Fraction of CPU time used by GC since program start.", m.GCCPUFraction)
g("go_goroutines", "Number of goroutines that currently exist.", float64(runtime.NumGoroutine()))
}
// RegisterRuntimeMetrics adds Go runtime memory/scheduler stats to the default
// registry so they appear on the /metrics endpoint. Idempotent: safe to call on
// every daemon start (repeated starts in a test binary would otherwise panic on
// the duplicate registration).
func RegisterRuntimeMetrics() {
runtimeMetricsOnce.Do(func() {
defaultRegistry.MustRegister("go_runtime_stats", goRuntimeCollector{})
})
}
package mime
import (
"archive/tar"
"archive/zip"
"bufio"
"bytes"
"compress/gzip"
"fmt"
"io"
"mime"
"mime/multipart"
"net/textproto"
"os"
"path/filepath"
"strings"
"unicode"
)
// ExtractedPart represents a single extracted attachment.
type ExtractedPart struct {
Filename string
ContentType string
Size int64
TempPath string
Nested bool
ArchiveName string
}
// ExtractionResult holds all extracted parts and envelope metadata.
type ExtractionResult struct {
Parts []ExtractedPart
Partial bool
PartialReason string
EncryptedEntries []EncryptedArchiveEntry
// Report names are bounded separately from extraction: encrypted members
// beyond the reporting limit still must not cause delivery retries.
EncryptedEntriesOmitted int
Direction string
From string
To []string
Subject string
// Shared across multipart recursion so ambiguous wrappers cannot branch
// into exponentially many decoding passes.
transferVariants int
}
// EncryptedArchiveEntry names an archive member CSM cannot read because the
// entry is encrypted. This is deliberately not a partial extraction: a partial
// extraction is a limit or a fault that a later attempt may get past, while an
// encrypted member stays unreadable however many times delivery is retried.
type EncryptedArchiveEntry struct {
ArchiveName string
Filename string
}
// Limits controls resource bounds during extraction.
type Limits struct {
MaxAttachmentSize int64
MaxArchiveDepth int
MaxArchiveFiles int
MaxExtractionSize int64
// TempDir is the directory CreateTemp uses for extracted parts.
// Empty falls back to os.TempDir() (/tmp on Linux). Operators
// should set this to a daemon-owned 0700 path so extracted email
// attachments are not staged in a world-writable directory where
// another local uid can race the scanner via symlink swaps.
TempDir string
}
// DefaultLimits returns the default extraction limits.
func DefaultLimits() Limits {
return Limits{
MaxAttachmentSize: 25 * 1024 * 1024,
MaxArchiveDepth: 1,
MaxArchiveFiles: 50,
MaxExtractionSize: 100 * 1024 * 1024,
}
}
// ParseSpoolMessage parses an Exim spool message (-H and -D files) and
// extracts attachments to a temp directory. Caller must remove the temp
// files in result.Parts[*].TempPath when done.
func ParseSpoolMessage(headerPath, bodyPath string, limits Limits) (*ExtractionResult, error) {
envelope, hdrs, err := parseEximHeader(headerPath)
if err != nil {
return nil, fmt.Errorf("parsing header file: %w", err)
}
result := &ExtractionResult{
From: envelope.from,
To: envelope.to,
Subject: envelope.subject,
}
// Determine direction from Received headers
result.Direction = detectDirection(hdrs)
maxBodyBytes := bodyReadLimit(limits)
bodyData, partial, err := readBodyFileLimited(bodyPath, maxBodyBytes)
if err != nil {
return nil, fmt.Errorf("reading body file: %w", err)
}
if partial {
result.Partial = true
result.PartialReason = "message body exceeds parser memory budget"
return result, nil
}
// Exim -D files open with a "<message-id>-D" marker line; strip it before
// any body decode so single-part base64/QP payloads are not corrupted by
// the marker bytes.
bodyData = stripSpoolBodyMarker(bodyData, bodyPath)
ct := hdrs.Get("Content-Type")
if ct == "" {
ct = "text/plain"
}
mediaType, params, parseErr := mime.ParseMediaType(ct)
if parseErr != nil {
// Unparseable content type - treat as plain text, no attachments
return result, nil //nolint:nilerr // fail-open by design
}
if strings.HasPrefix(mediaType, "multipart/") {
boundary := params["boundary"]
if boundary == "" {
return result, nil
}
var totalSize int64
cte := strings.ToLower(hdrs.Get("Content-Transfer-Encoding"))
extractEncodedMultipart(bytes.NewReader(bodyData), cte, boundary, limits, result, &totalSize, 0, 0)
} else if !strings.HasPrefix(mediaType, "text/") {
// Single-part non-text message (e.g. application/octet-stream,
// application/pdf, image/*). These are attachment-like payloads
// that must be scanned even without a multipart wrapper.
cte := strings.ToLower(hdrs.Get("Content-Transfer-Encoding"))
readers, _ := transferReaders(cte, bytes.NewReader(bodyData), result)
var totalSize int64
for _, reader := range readers {
decoded, truncated, decodeErr := decodeSinglePart(reader, limits.MaxAttachmentSize+1)
if decodeErr != nil {
// The decoded prefix is still scanned; see extractMultipartNested.
markPartial(result, fmt.Sprintf("could not decode single-part attachment: %v", decodeErr))
}
switch {
case decodeErr != nil && len(decoded) == 0:
// Nothing decoded, nothing to scan.
case !truncated && int64(len(decoded)) <= limits.MaxAttachmentSize:
if int64(len(decoded)) > limits.MaxExtractionSize-totalSize {
markPartial(result, "total extraction size exceeds limit")
break
}
tmpFile, tmpErr := os.CreateTemp(limits.TempDir, "csm-emailav-single-*")
if tmpErr == nil {
n, writeErr := tmpFile.Write(decoded)
closeErr := tmpFile.Close()
if writeErr != nil || closeErr != nil || n != len(decoded) {
os.Remove(tmpFile.Name())
markPartial(result, "could not stage single-part attachment for scanning")
} else {
totalSize += int64(n)
filename := params["name"]
if filename == "" {
filename = "attachment"
}
filename = sanitizeAttachmentName(filename)
result.Parts = append(result.Parts, ExtractedPart{
Filename: filename,
ContentType: mediaType,
Size: int64(len(decoded)),
TempPath: tmpFile.Name(),
})
}
} else {
markPartial(result, "could not stage single-part attachment for scanning")
}
default:
result.Partial = true
result.PartialReason = "single-part attachment exceeds max size"
}
}
}
// text/* bodies are not attachments - skip
return result, nil
}
func bodyReadLimit(limits Limits) int64 {
limit := limits.MaxExtractionSize
if limit < limits.MaxAttachmentSize*2 {
limit = limits.MaxAttachmentSize * 2
}
if limit <= 0 {
limit = DefaultLimits().MaxExtractionSize
}
return limit
}
func readFileLimited(path string, limit int64) ([]byte, error) {
// #nosec G304 -- path is mail queue file path from scanner walk.
f, err := os.Open(path)
if err != nil {
return nil, err
}
defer f.Close()
data, err := io.ReadAll(io.LimitReader(f, limit+1))
if err != nil {
return nil, err
}
if int64(len(data)) > limit {
return nil, fmt.Errorf("message body exceeds parser memory budget")
}
return data, nil
}
// readBodyFileLimited opens the spool body file once and reads up to
// limit+1 bytes. Returns (data, partial=true, nil) when the file
// exceeds the limit so the caller can mark the result as partial.
// Folding the size check into the same file descriptor closes the
// TOCTOU window: previously the Stat-then-Open sequence let an
// attacker swap the file for a larger one between the two syscalls
// and bypass the limit.
func readBodyFileLimited(path string, limit int64) ([]byte, bool, error) {
// #nosec G304 -- path is mail queue file path from scanner walk.
f, err := os.Open(path)
if err != nil {
return nil, false, err
}
defer f.Close()
data, err := io.ReadAll(io.LimitReader(f, limit+1))
if err != nil {
return nil, false, err
}
if int64(len(data)) > limit {
return nil, true, nil
}
return data, false, nil
}
func markPartial(result *ExtractionResult, reason string) {
result.Partial = true
if result.PartialReason == "" {
result.PartialReason = reason
}
}
// decodeSinglePart returns the decoded body, whether it was cut at limit, and
// any decode error. On error the bytes decoded before it are still returned.
func decodeSinglePart(r io.Reader, limit int64) ([]byte, bool, error) {
decoded, err := io.ReadAll(io.LimitReader(r, limit+1))
if int64(len(decoded)) > limit {
return decoded[:limit], true, err
}
return decoded, false, err
}
type envelope struct {
from string
to []string
subject string
}
// parseEximHeader reads an Exim -H file and extracts envelope info and headers.
func parseEximHeader(path string) (*envelope, textproto.MIMEHeader, error) {
// #nosec G304 -- path is Exim -H file path from mail queue scanner walk.
data, err := os.ReadFile(path)
if err != nil {
return nil, nil, err
}
env, hdrs := parseEximHeaderData(data)
return env, hdrs, nil
}
// parseEximHeaderData parses the bytes of an Exim -H spool file into the
// envelope fields and the RFC 5322 header set. The -H format is:
//
// line 1: <message-id>-H
// line 2: <envelope-user> <uid> <gid>
// <envelope metadata / options / recipient block>
// <blank line>
// <RFC headers, each prefixed "<byte-count><flag> ">
//
// Each header line carries a decimal byte count and a single flag character
// ('F' From, 'T' To, 'P' Received, 'R' Reply-To, '*' deleted, space for the
// rest); folded continuations start with whitespace. Real cPanel-Exim writes
// this format, so parsing bare "From:" lines (as an earlier version did) never
// matched and left every field empty. Parsing is fail-open: unparseable input
// yields empty headers, never an error.
func parseEximHeaderData(data []byte) (*envelope, textproto.MIMEHeader) {
env := &envelope{}
hdrs := make(textproto.MIMEHeader)
reconstructed := reconstructEximHeaderBytes(data)
tp := textproto.NewReader(bufio.NewReader(bytes.NewReader(reconstructed)))
// A malformed trailing line makes ReadMIMEHeader return an error alongside
// the headers it parsed before the fault; keep those and fail open.
parsed, _ := tp.ReadMIMEHeader()
if len(parsed) == 0 {
return env, hdrs // fail-open: no recognizable headers
}
hdrs = parsed
env.from = hdrs.Get("From")
env.subject = hdrs.Get("Subject")
if to := hdrs.Get("To"); to != "" {
for _, addr := range strings.Split(to, ",") {
env.to = append(env.to, strings.TrimSpace(addr))
}
}
return env, hdrs
}
// ParseSpoolMIMEHeaders returns the live RFC 5322 headers from Exim -H data.
// Folded values are unfolded and deleted headers are omitted, just as they
// are for attachment extraction.
func ParseSpoolMIMEHeaders(data []byte) textproto.MIMEHeader {
_, headers := parseEximHeaderData(data)
return headers
}
// reconstructEximHeaderBytes rebuilds a plain RFC 5322 header block from an
// Exim -H file by dropping the two leading metadata lines and the envelope
// preamble, then stripping the "<byte-count><flag> " prefix from each header
// line. Deleted headers (flag '*') and their folds are omitted; folded
// continuation lines are preserved verbatim so textproto can rejoin them.
func reconstructEximHeaderBytes(data []byte) []byte {
sc := bufio.NewScanner(bytes.NewReader(data))
// A single header line (DKIM/ARC signatures, long base64) can be large but
// is bounded by the -H file already in memory.
sc.Buffer(make([]byte, 0, 8192), len(data)+64)
var out bytes.Buffer
lineNum := 0
inHeaders := false
skippingDeleted := false
haveLiveHeader := false
for sc.Scan() {
line := sc.Text()
lineNum++
if lineNum <= 2 {
continue // message-id marker, then envelope-user line
}
if !inHeaders {
if line == "" {
inHeaders = true
}
continue // envelope metadata / recipient block
}
// Folded continuation: RFC 5322 lines starting with WSP belong to the
// preceding header.
if len(line) > 0 && (line[0] == ' ' || line[0] == '\t') {
if !skippingDeleted && haveLiveHeader {
out.WriteString(line)
out.WriteString("\r\n")
}
continue
}
rest, flag, ok := stripEximPrefix(line)
if !ok {
skippingDeleted = false
haveLiveHeader = false
continue // not a recognizable header start
}
if flag == '*' {
skippingDeleted = true // deleted header: skip it and its folds
haveLiveHeader = false
continue
}
if !isMIMEHeaderStart(rest) {
skippingDeleted = false
haveLiveHeader = false
continue
}
skippingDeleted = false
haveLiveHeader = true
out.WriteString(rest)
out.WriteString("\r\n")
}
out.WriteString("\r\n") // terminate the header block for ReadMIMEHeader
return out.Bytes()
}
// stripEximPrefix removes the leading "<byte-count><flag> " from an Exim -H
// header line and returns the remaining "Name: value" text plus the flag byte.
// ok is false when the line does not carry a valid prefix.
func stripEximPrefix(line string) (rest string, flag byte, ok bool) {
i := 0
for i < len(line) && line[i] >= '0' && line[i] <= '9' {
i++
}
if i == 0 || i+1 >= len(line) {
return "", 0, false
}
flag = line[i]
if !isEximFlagByte(flag) || line[i+1] != ' ' {
return "", 0, false
}
return line[i+2:], flag, true
}
func isMIMEHeaderStart(line string) bool {
colon := strings.IndexByte(line, ':')
if colon <= 0 {
return false
}
for i := 0; i < colon; i++ {
if !isMIMEHeaderFieldNameByte(line[i]) {
return false
}
}
return true
}
func isMIMEHeaderFieldNameByte(b byte) bool {
return (b >= '0' && b <= '9') ||
(b >= 'a' && b <= 'z') ||
(b >= 'A' && b <= 'Z') ||
b == '!' || b == '#' || b == '$' || b == '%' ||
b == '&' || b == '\'' || b == '*' || b == '+' ||
b == '-' || b == '.' || b == '^' || b == '_' ||
b == '`' || b == '|' || b == '~'
}
func isEximFlagByte(b byte) bool {
return b == ' ' || b == '*' ||
(b >= 'A' && b <= 'Z') || (b >= 'a' && b <= 'z')
}
// stripSpoolBodyMarker removes the first line of an Exim -D body file when it
// is the "<message-id>-D" marker Exim always writes there. The marker equals
// the -D file's base name, so matching on that is exact: bodies that lack a
// marker (or callers that pass a raw body) are left untouched. Without this,
// single-part base64/QP payloads decode the marker bytes as content and fail.
func stripSpoolBodyMarker(bodyData []byte, bodyPath string) []byte {
marker := filepath.Base(bodyPath)
nl := bytes.IndexByte(bodyData, '\n')
var first []byte
if nl < 0 {
first = bodyData
} else {
first = bodyData[:nl]
}
if string(bytes.TrimRight(first, "\r")) != marker {
return bodyData
}
if nl < 0 {
return nil
}
return bodyData[nl+1:]
}
// detectDirection guesses inbound vs outbound from Received headers.
func detectDirection(hdrs textproto.MIMEHeader) string {
received := hdrs.Values("Received")
if len(received) == 0 {
return "outbound" // locally generated, no Received headers
}
// If the first (topmost) Received header contains "authenticated" it's outbound
first := strings.ToLower(received[0])
if strings.Contains(first, "(authenticated") || strings.Contains(first, "auth=") {
return "outbound"
}
return "inbound"
}
// maxMIMENestingDepth caps multipart-in-multipart recursion. Legitimate
// mail rarely nests beyond mixed > alternative > related; a crafted message
// can nest arbitrarily and would otherwise consume one stack frame per
// wrapper while hiding attachments below any scanner's patience.
const maxMIMENestingDepth = 16
// extractEncodedMultipart recursively walks MIME parts, extracting attachments.
// depth counts archive nesting (zip-in-zip), not MIME nesting: an archive
// attached five multipart levels down is still archive depth 0.
func extractEncodedMultipart(r io.Reader, cte, boundary string, limits Limits, result *ExtractionResult, totalSize *int64, depth, mimeDepth int) {
readers, readErr := transferReaders(cte, r, result)
if readErr != nil {
markPartial(result, fmt.Sprintf("could not decode multipart body: %v", readErr))
}
for _, reader := range readers {
body := &readErrRecorder{r: reader}
parseErr := extractMultipartNested(body, boundary, limits, result, totalSize, depth, mimeDepth)
// A closing MIME delimiter ends parsing before the transfer decoder has
// necessarily reported its error. Drain the bounded spool body (or this
// outer part only), including any epilogue, to finish the decoder.
_, _ = io.Copy(io.Discard, body)
if body.err != nil {
markPartial(result, fmt.Sprintf("could not decode multipart body: %v", body.err))
}
if parseErr != nil {
markPartial(result, fmt.Sprintf("could not parse multipart body: %v", parseErr))
}
}
}
func extractMultipartNested(r io.Reader, boundary string, limits Limits, result *ExtractionResult, totalSize *int64, depth, mimeDepth int) error {
if mimeDepth >= maxMIMENestingDepth {
result.Partial = true
if result.PartialReason == "" {
result.PartialReason = "MIME nesting exceeds depth limit"
}
return nil
}
mr := multipart.NewReader(r, boundary)
for {
part, err := mr.NextRawPart()
if err == io.EOF {
return nil
}
if err != nil {
return err
}
ct := part.Header.Get("Content-Type")
if ct == "" {
ct = "text/plain"
}
mediaType, params, _ := mime.ParseMediaType(ct)
cte := strings.ToLower(part.Header.Get("Content-Transfer-Encoding"))
// Recurse into nested multipart. RFC 2045 forbids encoding a
// multipart body, but a sender can still do it, so decode first.
if strings.HasPrefix(mediaType, "multipart/") {
if b := params["boundary"]; b != "" {
// Failure within one wrapper must not hide its outer siblings.
extractEncodedMultipart(part, cte, b, limits, result, totalSize, depth, mimeDepth+1)
}
continue
}
// Skip inline text bodies - only extract attachments
disp := part.Header.Get("Content-Disposition")
filename := part.FileName()
if filename == "" {
// Try Content-Disposition filename param
if disp != "" {
_, dparams, _ := mime.ParseMediaType(disp)
filename = dparams["filename"]
}
}
if filename == "" {
// No filename and text/* content - this is a body part, skip
if strings.HasPrefix(mediaType, "text/") {
continue
}
// Non-text without filename - use generic name
filename = "unnamed_attachment"
}
rawFilename := filename
filename = sanitizeAttachmentName(filename)
readers, readErr := transferReaders(cte, part, result)
if readErr != nil {
markPartial(result, fmt.Sprintf("could not decode attachment %q: %v", filename, readErr))
}
for _, reader := range readers {
// Decode the part body based on Content-Transfer-Encoding
bodyReader := &readErrRecorder{r: reader}
// Write to temp file with size limit
tmpFile, err := os.CreateTemp(limits.TempDir, "csm-emailav-*")
if err != nil {
markPartial(result, "could not stage attachment for scanning")
return fmt.Errorf("creating temp file: %w", err)
}
limited := io.LimitReader(bodyReader, limits.MaxAttachmentSize+1)
n, err := io.Copy(tmpFile, limited)
closeErr := tmpFile.Close()
if closeErr != nil || (err != nil && err != bodyReader.err) {
os.Remove(tmpFile.Name())
markPartial(result, "could not stage attachment for scanning")
continue // fail-open: skip this part
}
if bodyReader.err != nil {
// Keep what decoded: a mail client shows those bytes, so they
// must be scanned even though the rest of the part is lost.
markPartial(result, fmt.Sprintf("could not decode attachment %q: %v", filename, bodyReader.err))
if n == 0 {
os.Remove(tmpFile.Name())
continue
}
}
if n > limits.MaxAttachmentSize {
os.Remove(tmpFile.Name())
result.Partial = true
result.PartialReason = fmt.Sprintf("attachment %q exceeds max size %d", filename, limits.MaxAttachmentSize)
continue
}
*totalSize += n
if *totalSize > limits.MaxExtractionSize {
os.Remove(tmpFile.Name())
result.Partial = true
result.PartialReason = fmt.Sprintf("total extraction size exceeds %d bytes", limits.MaxExtractionSize)
return nil // stop extracting
}
result.Parts = append(result.Parts, ExtractedPart{
Filename: filename,
ContentType: mediaType,
Size: n,
TempPath: tmpFile.Name(),
})
// Attempt archive extraction
if depth < limits.MaxArchiveDepth {
switch archiveKindForAttachmentName(rawFilename, filename) {
case "zip":
extractZIP(tmpFile.Name(), filename, limits, result, totalSize, depth+1)
case "tar.gz":
extractTarGz(tmpFile.Name(), filename, limits, result, totalSize, depth+1)
}
}
}
}
}
// zipEncryptedFlag is bit 0 of the ZIP general purpose bit flag: the entry is
// encrypted. APPNOTE.TXT section 4.4.4.
const zipEncryptedFlag = 0x1
func extractZIP(zipPath, archiveName string, limits Limits, result *ExtractionResult, totalSize *int64, depth int) {
// #nosec G304 -- zipPath is CreateTemp-produced path from the caller.
f, err := os.Open(zipPath)
if err != nil {
return // fail-open: skip corrupt archives
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return
}
zr, err := zip.NewReader(f, info.Size())
if err != nil {
return
}
extracted := 0
for _, zf := range zr.File {
if zf.FileInfo().IsDir() {
continue
}
safeName := sanitizeAttachmentName(zf.Name)
// Bit 0 of the general purpose flag marks an encrypted entry. Both
// encryption schemes in the wild set it, which matters because they
// fail in different places otherwise: legacy ZipCrypto keeps the
// deflate method and fails during the read, while WinZip AES uses
// method 99 and fails when the entry is opened. Reading the flag
// catches both before either produces a misleading error.
if zf.Flags&zipEncryptedFlag != 0 {
if len(result.EncryptedEntries) < limits.MaxArchiveFiles {
result.EncryptedEntries = append(result.EncryptedEntries, EncryptedArchiveEntry{
ArchiveName: archiveName,
Filename: safeName,
})
} else {
result.EncryptedEntriesOmitted++
}
continue
}
if extracted >= limits.MaxArchiveFiles {
markPartial(result, fmt.Sprintf("archive %q exceeds max files %d", archiveName, limits.MaxArchiveFiles))
return
}
rc, err := zf.Open()
if err != nil {
markPartial(result, fmt.Sprintf("could not decompress file %q in archive %q: %v", safeName, archiveName, err))
continue
}
tmpFile, err := os.CreateTemp(limits.TempDir, "csm-emailav-zip-*")
if err != nil {
rc.Close()
markPartial(result, fmt.Sprintf("could not stage file %q in archive %q for scanning: %v", safeName, archiveName, err))
continue
}
limited := io.LimitReader(rc, limits.MaxAttachmentSize+1)
n, err := io.Copy(tmpFile, limited)
closeErr := tmpFile.Close()
rc.Close()
if err != nil || closeErr != nil || n > limits.MaxAttachmentSize {
os.Remove(tmpFile.Name())
switch {
case n > limits.MaxAttachmentSize:
result.Partial = true
result.PartialReason = fmt.Sprintf("file %q in archive exceeds max size", safeName)
case err != nil:
markPartial(result, fmt.Sprintf("could not decompress file %q in archive %q: %v", safeName, archiveName, err))
default:
markPartial(result, fmt.Sprintf("could not stage file %q in archive %q for scanning: %v", safeName, archiveName, closeErr))
}
continue
}
*totalSize += n
if *totalSize > limits.MaxExtractionSize {
os.Remove(tmpFile.Name())
result.Partial = true
result.PartialReason = fmt.Sprintf("total extraction size exceeds %d bytes", limits.MaxExtractionSize)
return
}
result.Parts = append(result.Parts, ExtractedPart{
Filename: safeName,
ContentType: "application/octet-stream",
Size: n,
TempPath: tmpFile.Name(),
Nested: true,
ArchiveName: archiveName,
})
extracted++
}
}
func extractTarGz(tgzPath, archiveName string, limits Limits, result *ExtractionResult, totalSize *int64, depth int) {
// #nosec G304 -- tgzPath is CreateTemp-produced path from the caller.
f, err := os.Open(tgzPath)
if err != nil {
return
}
defer f.Close()
gr, err := gzip.NewReader(f)
if err != nil {
return
}
defer func() { _ = gr.Close() }()
tr := tar.NewReader(gr)
extracted := 0
for {
hdr, err := tr.Next()
if err != nil {
return // EOF or error - done
}
if hdr.Typeflag != tar.TypeReg {
continue
}
if extracted >= limits.MaxArchiveFiles {
result.Partial = true
result.PartialReason = fmt.Sprintf("archive %q exceeds max files %d", archiveName, limits.MaxArchiveFiles)
return
}
safeName := sanitizeAttachmentName(hdr.Name)
tmpFile, err := os.CreateTemp(limits.TempDir, "csm-emailav-tgz-*")
if err != nil {
markPartial(result, fmt.Sprintf("could not stage file %q in archive %q for scanning: %v", safeName, archiveName, err))
continue
}
limited := io.LimitReader(tr, limits.MaxAttachmentSize+1)
n, err := io.Copy(tmpFile, limited)
closeErr := tmpFile.Close()
if err != nil || closeErr != nil || n > limits.MaxAttachmentSize {
os.Remove(tmpFile.Name())
switch {
case n > limits.MaxAttachmentSize:
result.Partial = true
result.PartialReason = fmt.Sprintf("file %q in archive exceeds max size", safeName)
case err != nil:
markPartial(result, fmt.Sprintf("could not read file %q from archive %q: %v", safeName, archiveName, err))
default:
markPartial(result, fmt.Sprintf("could not stage file %q in archive %q for scanning: %v", safeName, archiveName, closeErr))
}
continue
}
*totalSize += n
if *totalSize > limits.MaxExtractionSize {
os.Remove(tmpFile.Name())
result.Partial = true
result.PartialReason = fmt.Sprintf("total extraction size exceeds %d bytes", limits.MaxExtractionSize)
return
}
result.Parts = append(result.Parts, ExtractedPart{
Filename: safeName,
ContentType: "application/octet-stream",
Size: n,
TempPath: tmpFile.Name(),
Nested: true,
ArchiveName: archiveName,
})
extracted++
}
}
// sanitizeAttachmentName trims an attachment or archive-entry name to
// its base name and truncates at control characters before the name
// reaches logs, alerts, or JSON responses.
func sanitizeAttachmentName(name string) string {
name = strings.ReplaceAll(name, "\\", "/")
name = filepath.Base(name)
// Truncate at the first control character so a crafted entry
// like "good.txt\nFAKE-LOG-LINE" cannot smuggle a forged log
// record past the visible filename.
if i := strings.IndexFunc(name, unicode.IsControl); i >= 0 {
name = name[:i]
}
name = strings.TrimSpace(name)
switch name {
case "", ".", "..", "/":
return "_unnamed_"
}
return name
}
func archiveKindForAttachmentName(rawName, safeName string) string {
for _, name := range []string{safeName, archiveDetectionName(rawName)} {
lower := strings.ToLower(name)
switch {
case strings.HasSuffix(lower, ".zip"):
return "zip"
case strings.HasSuffix(lower, ".tar.gz") || strings.HasSuffix(lower, ".tgz"):
return "tar.gz"
}
}
return ""
}
// archiveDetectionName removes controls instead of truncating so a
// filename like "payload\u0085.zip" still gets unpacked while the
// public filename remains log-safe.
func archiveDetectionName(name string) string {
name = strings.ReplaceAll(name, "\\", "/")
name = filepath.Base(name)
name = strings.Map(func(r rune) rune {
if unicode.IsControl(r) {
return -1
}
return r
}, name)
return strings.TrimSpace(name)
}
package mime
import (
"bytes"
"errors"
"io"
"strings"
)
// transferDecoder returns a reader that undoes a part's
// Content-Transfer-Encoding. Decoding is as lenient as mail clients are:
// anything a client would render as an attachment has to reach the scanners
// too, or a malformed encoding becomes a way to deliver unscanned content.
// Malformed input still decodes as far as clients decode it; the reader then
// reports an error so the part is marked incompletely scanned. Unknown
// encodings (7bit, 8bit, binary) pass through unchanged.
func transferDecoder(cte string, r io.Reader) io.Reader {
switch cte {
case "base64":
return &decodeReader{src: r, dec: &base64Decoder{}}
case "quoted-printable":
return &decodeReader{src: r, dec: &qpDecoder{}}
default:
return r
}
}
// transferReaders preserves client interpretations of malformed padding:
// consuming full quartets, ignoring misplaced padding, or stopping at valid
// padding. A long suffix can hide an otherwise readable archive, so scanning
// a continued decode cannot replace scanning the padded prefix. The source
// is a part of the already memory-bounded spool body.
func transferReaders(cte string, r io.Reader, result *ExtractionResult) ([]io.Reader, error) {
if cte != "base64" {
return []io.Reader{transferDecoder(cte, r)}, nil
}
body, err := io.ReadAll(r)
readers := []io.Reader{transferDecoder(cte, bytes.NewReader(body))}
var alternatives []io.Reader
if prefix := base64PaddedPrefix(body); prefix != nil {
alternatives = append(alternatives, transferDecoder(cte, bytes.NewReader(prefix)))
}
if ambiguousBase64Padding(body) {
alternatives = append(alternatives, &decodeReader{src: bytes.NewReader(body), dec: &base64AlphabetDecoder{}})
}
for _, alternative := range alternatives {
const maxTransferVariants = 16
if result.transferVariants >= maxTransferVariants {
markPartial(result, "transfer decoding exceeds alternate interpretation limit")
} else {
result.transferVariants++
readers = append(readers, alternative)
}
}
return readers, err
}
// Return a normally padded prefix only when meaningful data follows it.
func base64PaddedPrefix(body []byte) []byte {
slot := 0
padded := false
for i, b := range body {
_, alphabet := base64Value(b)
switch {
case b == '=':
if slot < 2 {
return nil
}
padded = true
case !alphabet:
continue
case padded:
return nil // the quartet itself has misplaced padding
}
slot++
if slot == 4 {
if padded {
for _, tail := range body[i+1:] {
if _, ok := base64Value(tail); ok || tail == '=' {
return body[:i+1]
}
}
return nil
}
slot = 0
}
}
return nil
}
func ambiguousBase64Padding(body []byte) bool {
slot, pads := 0, 0
for _, b := range body {
if b == '=' {
if slot < 2 {
return true
}
pads++
} else if _, ok := base64Value(b); !ok {
continue
} else if pads > 0 {
return true
}
slot++
if slot == 4 {
slot, pads = 0, 0
}
}
return false
}
type base64AlphabetDecoder struct{ base64Decoder }
func (d *base64AlphabetDecoder) feed(b byte, out []byte) []byte {
if b == '=' {
return out
}
return d.base64Decoder.feed(b, out)
}
var (
errBase64AfterPadding = errors.New("base64 data after padding")
errBase64Truncated = errors.New("base64 data ends with an incomplete byte")
errBase64Padding = errors.New("misplaced base64 padding")
)
// byteDecoder turns encoded bytes into decoded bytes one input byte at a
// time. finish flushes held state at end of input and reports malformed input
// that was nonetheless decoded.
type byteDecoder interface {
feed(b byte, out []byte) []byte
finish(out []byte) ([]byte, error)
}
// decodeReader drives a byteDecoder. Decoded output is drained before any
// error is returned, so a caller always sees every byte a client would.
type decodeReader struct {
src io.Reader
dec byteDecoder
buf [4096]byte
pending []byte
err error
}
func (d *decodeReader) Read(p []byte) (int, error) {
for len(d.pending) == 0 && d.err == nil {
n, err := d.src.Read(d.buf[:])
out := d.pending[:0]
for _, b := range d.buf[:n] {
out = d.dec.feed(b, out)
}
if err != nil {
var decodeErr error
out, decodeErr = d.dec.finish(out)
// A truncated MIME part returns UnexpectedEOF. Its buffered
// tail still carries bytes; preserve the source error too.
if err == io.EOF && decodeErr != nil {
err = decodeErr
}
}
d.pending = out
d.err = err
}
if len(d.pending) == 0 {
return 0, d.err
}
n := copy(p, d.pending)
d.pending = d.pending[n:]
return n, nil
}
// base64Decoder decodes quartet by quartet and ignores every byte outside
// the alphabet, as RFC 2045 section 6.8 requires. Padding occupies a slot
// in the quartet, even when misplaced. Resetting at '=' instead would shift
// all subsequent bytes compared with clients that consume full quartets.
type base64Decoder struct {
quad [4]byte
n int
padding int
badPadding bool
sawPad bool
afterPadData bool
}
func (d *base64Decoder) feed(b byte, out []byte) []byte {
v, ok := base64Value(b)
switch {
case b == '=':
if d.n < 2 {
d.badPadding = true
}
d.padding++
d.sawPad = true
case !ok:
return out
case d.sawPad:
d.afterPadData = true
}
d.quad[d.n] = v
d.n++
if d.n == 4 {
out = d.flush(out)
d.n = 0
d.padding = 0
}
return out
}
// Clients such as Thunderbird emit at least one byte even for a quartet
// containing three or four padding characters.
func (d *base64Decoder) flush(out []byte) []byte {
q := d.quad
out = append(out, q[0]<<2|q[1]>>4)
if d.padding < 2 {
out = append(out, q[1]<<4|q[2]>>2)
}
if d.padding == 0 {
out = append(out, q[2]<<6|q[3])
}
return out
}
func (d *base64Decoder) finish(out []byte) ([]byte, error) {
var err error
switch {
case d.n >= 2:
for d.n < 4 {
d.quad[d.n] = 0
d.n++
d.padding++
}
out = d.flush(out) // missing final padding
case d.n == 1:
err = errBase64Truncated
}
d.n = 0
if d.badPadding {
err = errBase64Padding
}
if d.afterPadData {
err = errBase64AfterPadding
}
return out, err
}
func base64Value(b byte) (byte, bool) {
switch {
case b >= 'A' && b <= 'Z':
return b - 'A', true
case b >= 'a' && b <= 'z':
return b - 'a' + 26, true
case b >= '0' && b <= '9':
return b - '0' + 52, true
case b == '+':
return 62, true
case b == '/':
return 63, true
}
return 0, false
}
// qpDecoder decodes quoted-printable without ever failing: an escape that is
// not two hex digits is kept literally, a soft line break may end in CR, LF or
// CRLF, raw control bytes pass through, and lines have no length limit.
// Whitespace before a line break is transport padding and is dropped
// (RFC 2045 section 6.7, rule 3).
type qpDecoder struct {
state qpState
hex1 byte
ws []byte // whitespace held until the next byte shows whether it ends a line
}
type qpState int
const (
qpText qpState = iota
qpEquals // saw '='
qpEqualsH1 // saw '=' and one hex digit
qpEqualsWS // saw '=' followed by whitespace: soft break if a line break follows
qpSoftCR // soft break ended in CR; swallow a following LF
)
func (d *qpDecoder) feed(b byte, out []byte) []byte {
switch d.state {
case qpEquals:
switch {
case isHexDigit(b):
d.hex1 = b
d.state = qpEqualsH1
return out
case b == ' ' || b == '\t':
d.ws = append(d.ws[:0], b)
d.state = qpEqualsWS
return out
case b == '\r':
d.state = qpSoftCR
return out
case b == '\n':
d.state = qpText
return out
}
out = append(out, '=')
d.state = qpText
case qpEqualsH1:
d.state = qpText
if isHexDigit(b) {
return append(out, hexValue(d.hex1)<<4|hexValue(b))
}
out = append(out, '=', d.hex1)
case qpEqualsWS:
switch b {
case ' ', '\t':
d.ws = append(d.ws, b)
return out
case '\r':
d.ws = d.ws[:0]
d.state = qpSoftCR
return out
case '\n':
d.ws = d.ws[:0]
d.state = qpText
return out
}
out = append(out, '=')
out = append(out, d.ws...)
d.ws = d.ws[:0]
d.state = qpText
case qpSoftCR:
d.state = qpText
if b == '\n' {
return out
}
}
switch b {
case '=':
out = append(out, d.ws...)
d.ws = d.ws[:0]
d.state = qpEquals
case ' ', '\t':
d.ws = append(d.ws, b)
case '\r', '\n':
d.ws = d.ws[:0]
out = append(out, b)
default:
out = append(out, d.ws...)
d.ws = d.ws[:0]
out = append(out, b)
}
return out
}
func (d *qpDecoder) finish(out []byte) ([]byte, error) {
switch d.state {
case qpEquals:
out = append(out, '=')
case qpEqualsH1:
out = append(out, '=', d.hex1)
}
// Trailing whitespace on the last line and a trailing "= " soft break
// carry no data.
d.ws = d.ws[:0]
d.state = qpText
return out, nil
}
func isHexDigit(b byte) bool {
return b >= '0' && b <= '9' || b >= 'A' && b <= 'F' || b >= 'a' && b <= 'f'
}
func hexValue(b byte) byte {
switch {
case b >= '0' && b <= '9':
return b - '0'
case b >= 'a' && b <= 'f':
return b - 'a' + 10
default:
return b - 'A' + 10
}
}
// readErrRecorder remembers the error its source returned, so a failed copy
// can be attributed to decoding rather than to writing the staged file.
type readErrRecorder struct {
r io.Reader
err error
}
func (rr *readErrRecorder) Read(p []byte) (int, error) {
n, err := rr.r.Read(p)
if err != nil && err != io.EOF {
rr.err = err
}
return n, err
}
// DecodeTransferVariants decodes data under Content-Transfer-Encoding cte the
// way attachment extraction does: leniently, and under every reading of
// ambiguous base64 a mail client could apply. Decode errors are ignored; each
// variant holds every byte that decoded.
func DecodeTransferVariants(cte string, data []byte) [][]byte {
readers, _ := transferReaders(strings.ToLower(strings.TrimSpace(cte)), bytes.NewReader(data), &ExtractionResult{})
variants := make([][]byte, 0, len(readers))
for _, r := range readers {
decoded, _ := io.ReadAll(r)
variants = append(variants, decoded)
}
return variants
}
package modsec
import "github.com/pidginhost/csm/internal/platform"
// RuleDirs returns the candidate directories where vendor ModSecurity rules
// live for the detected web server / panel combination. The list is ordered
// from most-specific to least-specific so callers walking the dirs encounter
// the operator's installed pack before any system fallback.
func RuleDirs(info platform.Info) []string {
var dirs []string
switch info.WebServer {
case platform.WSApache:
if info.IsDebianFamily() {
dirs = append(dirs,
"/etc/apache2/conf.d/modsec_vendor_configs/",
"/etc/modsecurity/",
"/usr/share/modsecurity-crs/rules/",
)
}
if info.IsRHELFamily() {
dirs = append(dirs,
"/etc/httpd/modsecurity.d/",
"/etc/httpd/modsecurity.d/activated_rules/",
"/usr/share/modsecurity-crs/rules/",
)
}
dirs = append(dirs, "/usr/local/apache/conf/modsec_vendor_configs/")
case platform.WSNginx:
dirs = append(dirs,
"/etc/nginx/modsec/",
"/etc/modsecurity/",
"/usr/share/modsecurity-crs/rules/",
)
}
// cPanel + LiteSpeed: cPanel's modsec_assemble job writes vendor rules
// into the apache2 tree even when the front-end is LiteSpeed. Without
// this branch the rule probe has no filesystem evidence during the
// window in which modsec_assemble itself is rewriting the tree.
if info.IsCPanel() && info.WebServer == platform.WSLiteSpeed {
dirs = append(dirs,
"/etc/apache2/conf.d/modsec_vendor_configs/",
"/usr/local/apache/conf/modsec_vendor_configs/",
)
}
return dirs
}
package modsec
import (
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"os"
)
// RuleTreeFingerprint summarises the rule files BuildRegistry would parse:
// their paths, contents and directory precedence. Content hashing catches
// vendor updates that preserve timestamps and changes behind .conf symlinks,
// while avoiding the rule parser's allocations and action extraction.
// An empty fingerprint means a read failed and must never skip a rebuild.
//
// present is false when none of the dirs exist. That is the signal that the
// resolved directories no longer describe this host (a web server swapped out,
// a vendor pack removed) and detection has to run again; it is not the same as
// a rule tree that exists and happens to be empty.
func RuleTreeFingerprint(dirs []string) (string, bool) {
digest := sha256.New()
present := false
complete := true
for _, dir := range dirs {
if dir == "" {
continue
}
if info, err := os.Stat(dir); err == nil && info.IsDir() {
present = true
}
fmt.Fprintf(digest, "dir:%q\n", dir)
err := walkRuleFiles(dir, func(path string) error {
fileDigest, err := readRuleFile(path, func(reader io.Reader) error {
_, err := io.Copy(io.Discard, reader)
return err
})
if err != nil {
return err
}
fmt.Fprintf(digest, "file:%q:%x\n", path, fileDigest)
return nil
})
if err != nil {
complete = false
}
}
if !complete {
return "", present
}
return hex.EncodeToString(digest.Sum(nil)), present
}
// Hash the same stream the consumer reads so a concurrent rewrite cannot
// associate parsed actions with a fingerprint of different file contents.
func readRuleFile(path string, consume func(io.Reader) error) ([]byte, error) {
f, err := openRuleFile(path)
if err != nil {
return nil, err
}
digest := sha256.New()
readErr := consume(&ruleFileReader{reader: io.TeeReader(f, digest)})
if err := errors.Join(readErr, f.Close()); err != nil {
return nil, err
}
return digest.Sum(nil), nil
}
// Tests inject read failures and appends at EOF through the file boundary.
var openRuleFile = func(path string) (io.ReadCloser, error) {
// #nosec G304 -- path comes from the operator's ModSec rule directories.
return os.Open(path)
}
// The parser and drain share one terminal result. Retrying an I/O error
// would misclassify an incomplete parse as cacheable. Reading past EOF
// could hash a concurrent append that the parser never saw.
type ruleFileReader struct {
reader io.Reader
err error
}
func (r *ruleFileReader) Read(p []byte) (int, error) {
if r.err != nil {
return 0, r.err
}
var n int
n, r.err = r.reader.Read(p)
return n, r.err
}
package modsec
import (
"bufio"
"fmt"
"io"
"log"
"os"
"sort"
"strconv"
"strings"
)
const overridesHeader = "# CSM ModSecurity Rule Overrides\n# Managed by CSM - do not edit manually.\n"
// WriteOverrides writes the overrides file with SecRuleRemoveById directives.
// Atomic write (tmp + rename). Sorts IDs for deterministic output.
func WriteOverrides(path string, disabledIDs []int) error {
sort.Ints(disabledIDs)
var sb strings.Builder
sb.WriteString(overridesHeader)
for _, id := range disabledIDs {
fmt.Fprintf(&sb, "SecRuleRemoveById %d\n", id)
}
tmpPath := path + ".tmp"
// #nosec G306 -- ModSec config files are read by the webserver process
// (Apache/nginx) which runs as a different user. 0640 lets root write
// and the webserver group read; world-read stays off.
if err := os.WriteFile(tmpPath, []byte(sb.String()), 0640); err != nil {
return fmt.Errorf("writing overrides tmp: %w", err)
}
if err := os.Rename(tmpPath, path); err != nil {
os.Remove(tmpPath)
return fmt.Errorf("renaming overrides: %w", err)
}
return nil
}
// ReadOverrides reads the overrides file and returns disabled rule IDs.
// Returns empty list (not error) if the file does not exist.
func ReadOverrides(path string) ([]int, error) {
// #nosec G304 -- path is operator-configured ModSec overrides file.
f, err := os.Open(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
return nil, err
}
defer f.Close()
var ids []int
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 0, 64*1024), maxModsecLineBytes)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if strings.HasPrefix(line, "SecRuleRemoveById ") {
idStr := strings.TrimPrefix(line, "SecRuleRemoveById ")
if id, err := strconv.Atoi(strings.TrimSpace(idStr)); err == nil && id >= 900000 && id <= 900999 {
ids = append(ids, id)
}
}
}
return ids, scanner.Err()
}
// ReadOverridesRaw reads the overrides file content for rollback purposes.
// Returns nil (not error) if the file does not exist.
func ReadOverridesRaw(path string) []byte {
// #nosec G304 -- path is operator-configured ModSec overrides file.
data, err := os.ReadFile(path)
if err != nil {
return nil
}
return data
}
// RestoreOverrides writes raw content back to the overrides file (for rollback).
// Uses atomic tmp+rename to prevent partial writes on crash.
func RestoreOverrides(path string, content []byte) error {
if content == nil {
// File didn't exist before - remove it
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("removing overrides during rollback: %w", err)
}
return nil
}
tmpPath := path + ".tmp"
// #nosec G306 -- see note in WriteOverrides: webserver-readable ModSec config.
if err := os.WriteFile(tmpPath, content, 0640); err != nil {
return fmt.Errorf("writing rollback tmp: %w", err)
}
if err := os.Rename(tmpPath, path); err != nil {
os.Remove(tmpPath)
return fmt.Errorf("renaming rollback: %w", err)
}
return nil
}
// EnsureOverridesInclude appends an Include directive for the overrides file
// to a ModSecurity config file if not already present, and creates an empty
// overrides file if it doesn't exist. Idempotent: re-reads content under the
// write-open to avoid appending duplicate Include directives.
func EnsureOverridesInclude(rulesFile, overridesFile string) {
// Open for read+write to check-then-append atomically (same fd).
// #nosec G302 G304 -- webserver-readable ModSec rules file; 0640 means root
// can write and the webserver group can read. No world read.
f, err := os.OpenFile(rulesFile, os.O_RDWR|os.O_APPEND, 0640)
if err != nil {
return
}
data, err := io.ReadAll(f)
if err != nil {
_ = f.Close()
return
}
if !strings.Contains(string(data), overridesFile) {
if _, writeErr := fmt.Fprintf(f, "\n# CSM overrides - managed by CSM rule management\nInclude %s\n", overridesFile); writeErr != nil {
_ = f.Close()
return
}
}
// Close error here drops the appended Include directive, leaving the
// override file unreferenced -- worth surfacing instead of swallowing.
if closeErr := f.Close(); closeErr != nil {
log.Printf("modsec: overrides include close failed for %s: %v", rulesFile, closeErr)
}
// Create empty overrides file if it doesn't exist
if _, err := os.Stat(overridesFile); os.IsNotExist(err) {
// #nosec G306 -- see note in WriteOverrides: webserver-readable.
_ = os.WriteFile(overridesFile, []byte(overridesHeader), 0640)
}
}
package modsec
import (
"bufio"
"fmt"
"io"
"os"
"regexp"
"strconv"
"strings"
)
// maxModsecLineBytes bounds a single logical line for the rule scanners. Far
// above any legitimate directive, but finite so a pathological file cannot
// drive an unbounded allocation.
const maxModsecLineBytes = 8 << 20 // 8 MiB
// Rule represents a parsed ModSecurity rule from the CSM custom config.
type Rule struct {
ID int // e.g. 900112
Description string // from msg:'...' field
Action string // disposition keyword: deny|drop|block|redirect|proxy|pause|allow|pass; "" if rule has only metadata (log, msg, ...) and inherits SecDefaultAction
StatusCode int // 403, 429, 0 (for pass)
Phase int // 1 or 2
Raw string // full rule text including chains
IsCounter bool // true if pass,nolog (bookkeeping rule, hidden in UI)
}
var (
reID = regexp.MustCompile(`[,"]id:(\d+)`)
reMsg = regexp.MustCompile(`msg:'([^']*)'`)
rePhase = regexp.MustCompile(`phase:(\d)`)
reStatus = regexp.MustCompile(`status:(\d+)`)
)
// dispositionPriority lists the ModSecurity action keywords that decide
// what happens to the request, ordered most-disruptive first. The first
// keyword present as a standalone token in the action string wins; this
// matches how ModSecurity itself resolves multiple disruptive directives
// in a single rule. Metadata keywords (log, msg, severity, tag, ...) are
// intentionally excluded - a rule that carries only metadata inherits
// SecDefaultAction, which CSM does not parse, so the registry leaves
// Action empty. A populated registry treats that rule as unknown and the
// LiteSpeed classifier defaults it to block.
var dispositionPriority = []string{
"deny", "drop", "block", "redirect", "proxy", "pause", "allow", "pass",
}
// dispositionSet is dispositionPriority as a lookup table. Pre-built for
// O(1) membership checks during action-string tokenisation.
var dispositionSet = func() map[string]struct{} {
m := make(map[string]struct{}, len(dispositionPriority))
for _, k := range dispositionPriority {
m[k] = struct{}{}
}
return m
}()
// ParseRulesFile reads a ModSecurity config file and extracts CSM-owned rules
// (IDs in 900000-900999). Use ParseRulesFileAll for the rule-action registry,
// which needs every rule including vendor packs.
func ParseRulesFile(path string) ([]Rule, error) {
all, err := ParseRulesFileAll(path)
if err != nil {
return nil, err
}
var csm []Rule
for _, r := range all {
if r.ID >= 900000 && r.ID <= 900999 {
csm = append(csm, r)
}
}
return csm, nil
}
// ParseRulesFileAll reads a ModSecurity config file and extracts every rule,
// regardless of ID range. Handles line continuations (\) and chained rules
// (chain action keyword). Vendor packs (Comodo, OWASP CRS, Imunify360) all
// use this entrypoint via the rule-action registry so the daemon can tell
// pass-action rules apart from deny rules.
func ParseRulesFileAll(path string) ([]Rule, error) {
// #nosec G304 -- path is operator-configured ModSec rules file.
f, err := os.Open(path)
if err != nil {
return nil, fmt.Errorf("opening rules file: %w", err)
}
defer f.Close()
return parseRules(f)
}
func parseRules(reader io.Reader) ([]Rule, error) {
// Phase 1: Read lines, joining backslash continuations into logical lines.
var logicalLines []string
var current strings.Builder
appendCurrent := func(part string) error {
if current.Len()+len(part) > maxModsecLineBytes {
return fmt.Errorf("logical modsec line exceeds %d bytes", maxModsecLineBytes)
}
current.WriteString(part)
return nil
}
scanner := bufio.NewScanner(reader)
// Vendor packs (OWASP CRS, Comodo, Imunify360, cPanel modsec_assemble)
// ship assembled/minified directives that can exceed the default 64 KB
// token. Without a larger buffer Scan stops at ErrTooLong and the file's
// rules drop out of the action registry, where unknown IDs then default
// to "deny" and skew the modsec signal. Raise the ceiling so real files
// parse in full.
scanner.Buffer(make([]byte, 0, 64*1024), maxModsecLineBytes)
for scanner.Scan() {
raw := scanner.Text()
trimmed := strings.TrimSpace(raw)
if strings.HasSuffix(trimmed, "\\") {
// Continuation: strip trailing \ and keep accumulating.
if err := appendCurrent(strings.TrimSuffix(trimmed, "\\")); err != nil {
return nil, err
}
if err := appendCurrent(" "); err != nil {
return nil, err
}
continue
}
if err := appendCurrent(trimmed); err != nil {
return nil, err
}
logicalLines = append(logicalLines, current.String())
current.Reset()
}
if current.Len() > 0 {
logicalLines = append(logicalLines, current.String())
}
if err := scanner.Err(); err != nil {
return nil, err
}
// Phase 2: Group logical lines into blocks.
// Each block starts with a SecRule and may include chained SecRules.
// When a directive's action string contains "chain", the next SecRule
// is part of the same block.
var blocks []string
var block strings.Builder
chainPending := false
flushBlock := func() {
if block.Len() > 0 {
blocks = append(blocks, block.String())
block.Reset()
}
chainPending = false
}
for _, line := range logicalLines {
if strings.HasPrefix(line, "SecRule ") {
if block.Len() > 0 && !chainPending {
flushBlock()
}
} else if block.Len() == 0 {
continue // skip comments and blank lines outside blocks
}
if block.Len() > 0 || strings.HasPrefix(line, "SecRule ") {
block.WriteString(line)
block.WriteString("\n")
// Only update chainPending for SecRule lines - non-SecRule
// directives between chained rules must not reset the flag.
if strings.HasPrefix(line, "SecRule ") {
chainPending = hasChainAction(line)
}
}
}
flushBlock()
// Phase 3: Parse each block into a Rule.
var rules []Rule
for _, b := range blocks {
if r, ok := parseBlock(b); ok {
rules = append(rules, r)
}
}
return rules, nil
}
// hasChainAction checks whether a logical line (continuations already joined)
// contains "chain" as a ModSecurity action keyword. Strips whitespace to handle
// action strings split across continuation lines like: "id:900004,..., chain"
func hasChainAction(line string) bool {
// Remove all whitespace so ", chain\"" becomes ",chain\""
stripped := strings.Map(func(r rune) rune {
if r == ' ' || r == '\t' {
return -1
}
return r
}, line)
return strings.Contains(stripped, ",chain\"") ||
strings.Contains(stripped, ",chain'") ||
strings.Contains(stripped, ",chain,") ||
strings.Contains(stripped, "\"chain\"") ||
strings.Contains(stripped, "\"chain,")
}
func parseBlock(block string) (Rule, bool) {
// Extract ID
m := reID.FindStringSubmatch(block)
if m == nil {
return Rule{}, false
}
id, _ := strconv.Atoi(m[1])
r := Rule{
ID: id,
Raw: strings.TrimSpace(block),
}
// Extract description from msg
if mm := reMsg.FindStringSubmatch(block); mm != nil {
r.Description = mm[1]
}
// Extract phase
if pm := rePhase.FindStringSubmatch(block); pm != nil {
r.Phase, _ = strconv.Atoi(pm[1])
}
// Extract the action string (the quoted segment that carries id:N) and
// pull the disposition keyword from its tokens. Substring matching
// against the whole block would false-match on action-like text inside
// regex operators, msg:'...' literals, and the like (e.g. "passive"
// looks like "pass"). Token parsing inside the bounded action string
// avoids those collisions.
if actionStr := extractActionString(block, id); actionStr != "" {
r.Action = pickDisposition(actionStr)
// nolog/log are honoured only as flags, never as the action.
actionLower := strings.ToLower(actionStr)
if r.Action == "pass" && strings.Contains(actionLower, "nolog") {
r.IsCounter = true
}
}
// Extract status code
if sm := reStatus.FindStringSubmatch(block); sm != nil {
r.StatusCode, _ = strconv.Atoi(sm[1])
}
return r, true
}
// extractActionString returns the body of the quoted segment that contains
// "id:<ruleID>". ModSecurity rule blocks have one such segment per rule
// (chained sub-rules carry "chain,capture"-style action lists with no id),
// so finding the id-bearing quotes uniquely identifies the primary action
// list. Returns "" if the segment cannot be located, in which case the
// caller leaves Action empty and the registry treats the rule as unknown.
func extractActionString(block string, ruleID int) string {
needle := "id:" + strconv.Itoa(ruleID)
for i := 0; i < len(block); i++ {
if block[i] != '"' {
continue
}
start := i + 1
escaped := false
for j := start; j < len(block); j++ {
switch {
case escaped:
escaped = false
case block[j] == '\\':
escaped = true
case block[j] == '"':
segment := block[start:j]
if actionStringHasRuleID(segment, needle) {
return segment
}
i = j
j = len(block)
}
}
}
return ""
}
func actionStringHasRuleID(actionStr, needle string) bool {
for _, tok := range tokenizeActionString(actionStr) {
name, value, ok := strings.Cut(tok, ":")
if !ok || strings.ToLower(strings.TrimSpace(name)) != "id" {
continue
}
if strings.Trim(strings.TrimSpace(value), `'"`) == strings.TrimPrefix(needle, "id:") {
return true
}
}
return false
}
// pickDisposition returns the ModSecurity disposition keyword present in
// the action string, preferring more-disruptive keywords when several are
// present (defensive: a rule labelled "deny" wins over a stray "allow").
// Returns "" if the action string carries only metadata and the rule
// therefore inherits SecDefaultAction.
func pickDisposition(actionStr string) string {
tokens := tokenizeActionString(actionStr)
seen := make(map[string]struct{}, len(tokens))
for _, t := range tokens {
name := t
if i := strings.IndexByte(name, ':'); i >= 0 {
name = name[:i]
}
name = strings.ToLower(strings.TrimSpace(name))
if _, ok := dispositionSet[name]; ok {
seen[name] = struct{}{}
}
}
for _, kw := range dispositionPriority {
if _, ok := seen[kw]; ok {
return kw
}
}
return ""
}
// tokenizeActionString splits a ModSecurity action list on top-level commas
// while respecting single-quoted string values (msg:'foo, bar', logdata:'...'),
// where commas are part of the literal and must not be treated as token
// separators. Backslash-escaped quotes inside the literals are preserved.
func tokenizeActionString(s string) []string {
var out []string
var cur strings.Builder
inSingle := false
escape := false
for _, r := range s {
switch {
case escape:
cur.WriteRune(r)
escape = false
case r == '\\':
cur.WriteRune(r)
escape = true
case r == '\'':
inSingle = !inSingle
cur.WriteRune(r)
case r == ',' && !inSingle:
tok := strings.TrimSpace(cur.String())
if tok != "" {
out = append(out, tok)
}
cur.Reset()
default:
cur.WriteRune(r)
}
}
if tok := strings.TrimSpace(cur.String()); tok != "" {
out = append(out, tok)
}
return out
}
// IsBlockingAction reports whether an action causes the request to be denied
// or otherwise diverted away from normal processing. Used by the LiteSpeed
// log-line classifier - error_log records every match as "triggered!"
// regardless of action, so the action lookup is the only way to tell a real
// deny apart from a pass-action informational rule. redirect, proxy and
// pause are disruptive: the original request never reaches the upstream
// application as intended, so they are classified the same as deny.
func IsBlockingAction(action string) bool {
switch action {
case "deny", "drop", "block", "redirect", "proxy", "pause":
return true
}
return false
}
package modsec
import (
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"io/fs"
"path/filepath"
"strings"
"sync/atomic"
)
// Registry maps parsed ModSecurity rule IDs with a decisive disposition to
// that action (deny, drop, block, redirect, proxy, pause, pass, allow).
// Rules with only metadata actions are intentionally absent so callers use
// the unknown-rule default. It is consulted by the LiteSpeed
// log-line classifier - error_log records every match as "triggered!"
// regardless of whether the rule's action denied the request, so the action
// lookup is the only signal that distinguishes a real deny from a noisy
// pass-action informational rule.
type Registry struct {
actions map[int]string
fingerprint string
}
// Action returns the declared action for ruleID and whether it is known.
// An unknown ID is the safe default: callers should treat it as a potential
// block when they have a populated registry but not this specific rule.
func (r *Registry) Action(ruleID int) (action string, known bool) {
if r == nil {
return "", false
}
a, ok := r.actions[ruleID]
return a, ok
}
// Len returns the number of rules in the registry. Useful for startup
// telemetry: zero typically means rule directories are missing.
func (r *Registry) Len() int {
if r == nil {
return 0
}
return len(r.actions)
}
// Fingerprint identifies the exact file contents used to build the registry.
// Empty means the build was incomplete and must not be cached.
func (r *Registry) Fingerprint() string {
if r == nil {
return ""
}
return r.fingerprint
}
// BuildRegistry walks every directory in dirs (recursively), parses each
// .conf file, and returns a Registry mapping rule IDs
// to actions. Read and parse errors are returned alongside the usable rules;
// a vendor pack with one malformed file should not blank the whole registry,
// but only builds that read every file in full can be cached as unchanged.
//
// Precedence: dirs is treated as most-specific-first. Within a single
// directory, files are walked in lexical order and a duplicate rule ID
// uses last-write-wins, mirroring how ModSecurity itself resolves two
// SecRule directives that share an ID. Across directories, the first
// directory to define a rule keeps it - that way an operator override in
// /etc/apache2/conf.d/modsec_vendor_configs/ is not silently replaced by
// a stale system fallback in /usr/share/modsecurity-crs/rules/.
func BuildRegistry(dirs []string) (*Registry, error) {
actions := make(map[int]string)
claimed := make(map[int]struct{})
digest := sha256.New()
// Read failures and parse failures are not the same for caching. Bytes
// that could not be read leave the tree unknown, so the build must not be
// cached. A file that was read in full but cannot be parsed -- a vendor
// file past the line ceiling, a truncated rule -- produces the same
// result on every pass, so it must not force a reparse every refresh.
var readErr, parseErr error
for _, dir := range dirs {
if dir == "" {
continue
}
perDirActions := make(map[int]string)
perDirClaimed := make(map[int]struct{})
fmt.Fprintf(digest, "dir:%q\n", dir)
walkErr := walkRuleFiles(dir, func(path string) error {
var rules []Rule
var fileParseErr error
fileDigest, err := readRuleFile(path, func(reader io.Reader) error {
rules, fileParseErr = parseRules(reader)
// Read the rest even when the parser stopped early, so the
// digest always describes the whole file and matches what
// RuleTreeFingerprint computes for the same tree.
_, drainErr := io.Copy(io.Discard, reader)
return drainErr
})
if err != nil {
return err
}
fmt.Fprintf(digest, "file:%q:%x\n", path, fileDigest)
if fileParseErr != nil {
parseErr = errors.Join(parseErr, fmt.Errorf("%s: %w", path, fileParseErr))
}
for _, r := range rules {
perDirClaimed[r.ID] = struct{}{}
if r.Action != "" {
perDirActions[r.ID] = r.Action
} else {
delete(perDirActions, r.ID)
}
}
return nil
})
// Promote the per-directory map into the global map only for IDs
// that no earlier (more-specific) directory has already claimed.
for id := range perDirClaimed {
if _, exists := claimed[id]; !exists {
claimed[id] = struct{}{}
action, hasAction := perDirActions[id]
if !hasAction {
continue
}
actions[id] = action
}
}
readErr = errors.Join(readErr, walkErr)
}
reg := &Registry{actions: actions}
if readErr == nil {
reg.fingerprint = hex.EncodeToString(digest.Sum(nil))
}
return reg, errors.Join(readErr, parseErr)
}
// Both the parser and fingerprint must see the same paths and precedence.
// WalkDir visits files lexically and does not descend into symlinked dirs;
// symlinked .conf files are opened by the visitor, following their targets.
func walkRuleFiles(dir string, visit func(string) error) error {
var readErr error
walkErr := filepath.WalkDir(dir, func(path string, d fs.DirEntry, err error) error {
if err != nil {
// Candidate directories commonly do not exist on this platform.
// A failure inside an existing tree means the walk is incomplete.
if path != dir || !errors.Is(err, fs.ErrNotExist) {
readErr = errors.Join(readErr, err)
}
return nil
}
if d.IsDir() || !strings.HasSuffix(strings.ToLower(d.Name()), ".conf") {
return nil
}
if err := visit(path); err != nil {
readErr = errors.Join(readErr, fmt.Errorf("%s: %w", path, err))
}
return nil
})
return errors.Join(readErr, walkErr)
}
var globalRegistry atomic.Pointer[Registry]
// SetGlobal installs r as the daemon-wide registry. Callers are expected to
// rebuild and re-set on a refresh interval. Safe for concurrent use.
func SetGlobal(r *Registry) {
globalRegistry.Store(r)
}
// ReplaceGlobal installs r as the daemon-wide registry, EXCEPT when r is empty
// (or nil) while the currently-installed registry is non-empty: in that case
// the previous registry is kept and false is returned.
//
// The vendor rule tree is transiently empty or unreadable during cPanel's
// nightly modsec_assemble rewrite and during a boot-time web-server
// mis-detection window (a LiteSpeed host probed before lsws has finished
// starting resolves to the wrong rule directories). Replacing a populated
// registry with an empty one would discard known pass and deny actions until
// the next successful refresh.
//
// Refresh callers should use this instead of SetGlobal. Returns true if r was
// installed, false if the previous registry was kept.
func ReplaceGlobal(r *Registry) bool {
for {
prev := globalRegistry.Load()
if r == nil || r.Len() == 0 {
if prev != nil && prev.Len() > 0 {
return false
}
}
if globalRegistry.CompareAndSwap(prev, r) {
return true
}
}
}
// Global returns the currently installed registry, or nil if none has been
// set yet (e.g. during very early daemon startup, or in unit tests that
// did not seed one). Callers must nil-check.
func Global() *Registry {
return globalRegistry.Load()
}
// ResetGlobalForTest clears the global registry. Test-only helper.
func ResetGlobalForTest() {
globalRegistry.Store(nil)
}
package modsec
import (
"context"
"errors"
"fmt"
"os/exec"
"time"
)
const reloadTimeout = 30 * time.Second
// Reload executes the configured web server reload command.
// Returns combined stdout+stderr output and any error.
func Reload(command string) (string, error) {
if command == "" {
return "", errors.New("reload command is empty")
}
ctx, cancel := context.WithTimeout(context.Background(), reloadTimeout)
defer cancel()
// Run through shell to support compound commands, quoted paths, etc.
// #nosec G204 -- `command` is the operator-configured reload command
// from csm.yaml (e.g. "apachectl graceful"), loaded at daemon startup
// from a root-owned config. Not webui-settable.
cmd := exec.CommandContext(ctx, "sh", "-c", command)
out, err := cmd.CombinedOutput()
output := string(out)
if ctx.Err() == context.DeadlineExceeded {
return output, fmt.Errorf("reload timed out after %v", reloadTimeout)
}
if err != nil {
return output, fmt.Errorf("reload failed: %w (output: %s)", err, output)
}
return output, nil
}
// Package mysqlclient wraps the database/sql + go-sql-driver/mysql
// pair for the read-only queries CSM issues against host-local
// MySQL/MariaDB. It replaces the per-call `mysql -e <query>`
// shell-outs across CMS DB scans, performance metrics, and forensic
// dumps so the daemon no longer forks and tears down a child process
// (plus its libc/libmariadbclient relocations) every query.
//
// Two cred modes are supported:
//
// - Root: implicit, mirrors the historical `mysql` CLI behaviour
// that reads /root/.my.cnf [client] section when present. RootDB
// parses that file and returns a *sql.DB pointed at unix-socket
// auth, falling back to the mysql CLI's root socket-auth default
// when absent.
//
// - Per-account: the caller passes explicit user / password / host
// / dbname / port / socket via DSN. PerAccountQuery opens,
// runs, closes.
//
// Output shape mirrors the previous `mysql -N -B -e` shell-out: rows
// are returned as []string where each entry is the tab-joined column
// values of one result row, with MySQL batch-mode escaping applied.
// Existing scanner code paths consume this format unchanged.
package mysqlclient
import (
"bufio"
"context"
"database/sql"
"fmt"
"net"
"os"
"strconv"
"strings"
"sync"
"time"
"github.com/go-sql-driver/mysql"
)
// queryTimeout caps a single SELECT to a defensive 30 s -- well under
// the legacy 2 min CLI cmdTimeout but enough for any realistic CMS
// scan / SHOW STATUS / SHOW PROCESSLIST result.
const queryTimeout = 30 * time.Second
// Creds carries per-account database credentials. Mirrors wpDBCreds in
// internal/checks so callers can pass the same shape without an extra
// adapter type.
type Creds struct {
User string
Password string
Host string
Port int
Socket string
DBName string
}
// dsn returns a go-sql-driver/mysql connection string. Empty Host
// falls back to the Unix socket path the mysql CLI uses on cPanel
// hosts; this matches the historical default when -h was omitted.
func (c Creds) dsn() string {
cfg := mysql.NewConfig()
cfg.User = c.User
cfg.Passwd = c.Password
cfg.DBName = c.DBName
// Read timeouts mirror queryTimeout's budget.
cfg.Timeout = 5 * time.Second
cfg.ReadTimeout = queryTimeout
cfg.WriteTimeout = 5 * time.Second
cfg.Loc = time.Local
cfg.Net, cfg.Addr = c.networkAddr()
return cfg.FormatDSN()
}
func (c Creds) networkAddr() (string, string) {
host := strings.TrimSpace(c.Host)
socket := strings.TrimSpace(c.Socket)
if socket != "" && (host == "" || host == "localhost") {
return "unix", socket
}
if strings.HasPrefix(host, "/") {
return "unix", host
}
if _, socket, ok := splitHostSocket(host); ok {
return "unix", socket
}
if host == "" || host == "localhost" {
return "unix", defaultUnixSocket()
}
host, port := splitHostPort(host, c.Port)
return "tcp", net.JoinHostPort(host, strconv.Itoa(port))
}
func splitHostSocket(host string) (string, string, bool) {
idx := strings.LastIndex(host, ":/")
if idx < 0 || idx == len(host)-1 {
return "", "", false
}
return host[:idx], host[idx+1:], true
}
func splitHostPort(host string, fallbackPort int) (string, int) {
if h, p, err := net.SplitHostPort(host); err == nil {
if port, perr := strconv.Atoi(p); perr == nil {
return trimIPv6Brackets(h), port
}
}
if strings.Count(host, ":") == 1 {
h, p, _ := strings.Cut(host, ":")
if port, err := strconv.Atoi(p); err == nil {
return h, port
}
}
if fallbackPort == 0 {
fallbackPort = 3306
}
return trimIPv6Brackets(host), fallbackPort
}
func trimIPv6Brackets(host string) string {
return strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")
}
// defaultUnixSocket returns the first known mysql socket path that
// exists on disk, falling back to the cPanel canonical location.
func defaultUnixSocket() string {
candidates := []string{
"/var/lib/mysql/mysql.sock",
"/tmp/mysql.sock",
"/var/run/mysqld/mysqld.sock",
}
for _, p := range candidates {
if _, err := os.Stat(p); err == nil {
return p
}
}
return "/var/lib/mysql/mysql.sock"
}
// PerAccountQuery opens a short-lived connection with the supplied
// credentials, runs the query, and returns each row as a tab-joined
// string. Empty result set returns (nil, nil). Errors include open,
// query, and scan failures.
//
// Tests can intercept via SetPerAccountQueryForTest.
func PerAccountQuery(ctx context.Context, creds Creds, query string, args ...any) ([]string, error) {
if fn := getPerAccountQueryMock(); fn != nil {
return fn(ctx, creds, query, args...)
}
db, err := sql.Open("mysql", creds.dsn())
if err != nil {
return nil, fmt.Errorf("mysqlclient: open: %w", err)
}
defer func() { _ = db.Close() }()
// Single short-lived call: cap idle pool to avoid leaking
// connections on per-account scans across many tenants.
db.SetMaxOpenConns(1)
db.SetMaxIdleConns(0)
db.SetConnMaxLifetime(queryTimeout)
return runQuery(ctx, db, query, args...)
}
// PerAccountQueryFunc is the signature SetPerAccountQueryForTest accepts.
type PerAccountQueryFunc func(ctx context.Context, creds Creds, query string, args ...any) ([]string, error)
var (
perAccountMockMu sync.RWMutex
perAccountMock PerAccountQueryFunc
)
// SetPerAccountQueryForTest installs an interceptor for
// PerAccountQuery. Pass nil to clear and restore the real database/sql
// path. Production code paths must NOT call this.
func SetPerAccountQueryForTest(fn PerAccountQueryFunc) {
perAccountMockMu.Lock()
defer perAccountMockMu.Unlock()
perAccountMock = fn
}
func getPerAccountQueryMock() PerAccountQueryFunc {
perAccountMockMu.RLock()
defer perAccountMockMu.RUnlock()
return perAccountMock
}
// runQuery is the shared execution path for any *sql.DB. Internal so
// the package can grow a RootDB-backed singleton later without
// duplicating the row-iteration code.
func runQuery(ctx context.Context, db *sql.DB, query string, args ...any) ([]string, error) {
cctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
rows, err := db.QueryContext(cctx, query, args...)
if err != nil {
return nil, fmt.Errorf("mysqlclient: query: %w", err)
}
defer func() { _ = rows.Close() }()
cols, err := rows.Columns()
if err != nil {
return nil, fmt.Errorf("mysqlclient: columns: %w", err)
}
out := make([]string, 0)
raw := make([]sql.NullString, len(cols))
scanArgs := make([]any, len(cols))
for i := range raw {
scanArgs[i] = &raw[i]
}
for rows.Next() {
if err := rows.Scan(scanArgs...); err != nil {
return nil, fmt.Errorf("mysqlclient: scan: %w", err)
}
parts := make([]string, len(cols))
for i, v := range raw {
if v.Valid {
parts[i] = mysqlBatchEscape(v.String)
} else {
// `mysql -N -B` prints NULL for SQL NULL; preserve
// that so legacy parsers that key off the literal
// see the same bytes.
parts[i] = "NULL"
}
}
out = append(out, strings.Join(parts, "\t"))
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("mysqlclient: iterate: %w", err)
}
return out, nil
}
func mysqlBatchEscape(s string) string {
var b strings.Builder
for i := 0; i < len(s); i++ {
switch s[i] {
case 0:
b.WriteString(`\0`)
case '\n':
b.WriteString(`\n`)
case '\r':
b.WriteString(`\r`)
case '\t':
b.WriteString(`\t`)
case '\\':
b.WriteString(`\\`)
default:
b.WriteByte(s[i])
}
}
return b.String()
}
// BatchUnescape is the exact inverse of mysqlBatchEscape (and of `mysql -N -B`
// batch-mode output): it turns the escaped sequences \0 \n \r \t \\ back into
// their raw bytes. Because the escaper backslash-escapes every literal
// backslash, a backslash in batch output is always the start of one of these
// sequences, so the inverse is unambiguous and round-trips byte-for-byte.
//
// Callers that write a value read from batch output back to the database must
// unescape first; otherwise a real newline persists as the two-byte literal
// "\n", corrupting the value and breaking PHP-serialized length prefixes.
func BatchUnescape(s string) string {
if !strings.ContainsRune(s, '\\') {
return s
}
var b strings.Builder
b.Grow(len(s))
for i := 0; i < len(s); i++ {
if s[i] != '\\' || i+1 >= len(s) {
b.WriteByte(s[i])
continue
}
i++
switch s[i] {
case '0':
b.WriteByte(0)
case 'n':
b.WriteByte('\n')
case 'r':
b.WriteByte('\r')
case 't':
b.WriteByte('\t')
case '\\':
b.WriteByte('\\')
default:
// Not a batch-escape sequence: keep both bytes verbatim.
b.WriteByte('\\')
b.WriteByte(s[i])
}
}
return b.String()
}
// --- Root creds via /root/.my.cnf ---------------------------------------
var (
rootMu sync.Mutex
rootDB *sql.DB
rootPath = "/root/.my.cnf"
)
// SetRootCnfPath overrides the [client] config path before the first
// RootDB() call. Tests set this to a tempdir-rooted .my.cnf.
func SetRootCnfPath(p string) {
rootMu.Lock()
defer rootMu.Unlock()
rootPath = p
// Drop any cached DB so the next RootDB() rebuilds from the new path.
if rootDB != nil {
_ = rootDB.Close()
rootDB = nil
}
}
// RootQuery runs a query as root using the credentials in
// /root/.my.cnf (or the path set via SetRootCnfPath). It mirrors the
// historical `mysql -N -B -e <query>` invocation that picked up
// /root/.my.cnf implicitly. The connection is pooled across calls.
//
// Tests can intercept via SetRootQueryForTest.
func RootQuery(ctx context.Context, query string, args ...any) ([]string, error) {
if fn := getRootQueryMock(); fn != nil {
return fn(ctx, "", query, args...)
}
db, err := RootDB()
if err != nil {
return nil, err
}
return runQuery(ctx, db, query, args...)
}
// RootExec runs a non-SELECT statement (DDL/DML) as root and returns
// the affected-row count plus any error. Mirrors `mysql -e <stmt>` for
// callers that previously relied on exit code only.
//
// Tests can intercept via SetRootExecForTest.
func RootExec(ctx context.Context, stmt string, args ...any) (int64, error) {
if fn := getRootExecMock(); fn != nil {
return fn(ctx, "", stmt, args...)
}
db, err := RootDB()
if err != nil {
return 0, err
}
cctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
res, err := db.ExecContext(cctx, stmt, args...)
if err != nil {
return 0, fmt.Errorf("mysqlclient: exec: %w", err)
}
rows, _ := res.RowsAffected()
return rows, nil
}
// RootExecSchema runs a non-SELECT statement against an explicit schema
// using root creds. Mirrors `mysql <schema> -e <stmt>`.
func RootExecSchema(ctx context.Context, schema, stmt string, args ...any) (int64, error) {
if fn := getRootExecMock(); fn != nil {
return fn(ctx, schema, stmt, args...)
}
creds, err := loadRootCreds()
if err != nil {
return 0, err
}
creds.DBName = schema
db, err := sql.Open("mysql", creds.dsn())
if err != nil {
return 0, fmt.Errorf("mysqlclient: open %s: %w", schema, err)
}
defer func() { _ = db.Close() }()
db.SetMaxOpenConns(1)
db.SetMaxIdleConns(0)
db.SetConnMaxLifetime(queryTimeout)
cctx, cancel := context.WithTimeout(ctx, queryTimeout)
defer cancel()
res, err := db.ExecContext(cctx, stmt, args...)
if err != nil {
return 0, fmt.Errorf("mysqlclient: exec %s: %w", schema, err)
}
rows, _ := res.RowsAffected()
return rows, nil
}
// RootQueryFunc is the signature SetRootQueryForTest accepts. schema is
// empty for RootQuery, non-empty for RootQuerySchema.
type RootQueryFunc func(ctx context.Context, schema, query string, args ...any) ([]string, error)
// RootExecFunc is the signature SetRootExecForTest accepts. schema is
// empty for RootExec, non-empty for RootExecSchema.
type RootExecFunc func(ctx context.Context, schema, stmt string, args ...any) (int64, error)
var (
rootQueryMockMu sync.RWMutex
rootQueryMock RootQueryFunc
rootExecMockMu sync.RWMutex
rootExecMock RootExecFunc
)
// SetRootQueryForTest installs an interceptor for RootQuery /
// RootQuerySchema. Pass nil to clear and restore the real database/sql
// path. Production code paths must NOT call this.
func SetRootQueryForTest(fn RootQueryFunc) {
rootQueryMockMu.Lock()
defer rootQueryMockMu.Unlock()
rootQueryMock = fn
}
// SetRootExecForTest installs an interceptor for RootExec /
// RootExecSchema. Pass nil to clear and restore the real database/sql
// path. Production code paths must NOT call this.
func SetRootExecForTest(fn RootExecFunc) {
rootExecMockMu.Lock()
defer rootExecMockMu.Unlock()
rootExecMock = fn
}
func getRootQueryMock() RootQueryFunc {
rootQueryMockMu.RLock()
defer rootQueryMockMu.RUnlock()
return rootQueryMock
}
func getRootExecMock() RootExecFunc {
rootExecMockMu.RLock()
defer rootExecMockMu.RUnlock()
return rootExecMock
}
// RootDB returns the pooled root MySQL handle backed by /root/.my.cnf
// when present (or the path set via SetRootCnfPath), falling back to
// the mysql CLI's root socket-auth default. sql.Open does not contact
// the server until the first query.
func RootDB() (*sql.DB, error) {
return rootSingleton()
}
// RootQuerySchema runs a query against an explicit schema using root
// creds. Mirrors `mysql <schema> -e <query>`.
func RootQuerySchema(ctx context.Context, schema, query string, args ...any) ([]string, error) {
if fn := getRootQueryMock(); fn != nil {
return fn(ctx, schema, query, args...)
}
creds, err := loadRootCreds()
if err != nil {
return nil, err
}
creds.DBName = schema
return PerAccountQuery(ctx, creds, query, args...)
}
func rootSingleton() (*sql.DB, error) {
rootMu.Lock()
if rootDB != nil {
defer rootMu.Unlock()
return rootDB, nil
}
rootMu.Unlock()
creds, err := loadRootCreds()
if err != nil {
return nil, err
}
db, err := sql.Open("mysql", creds.dsn())
if err != nil {
return nil, fmt.Errorf("mysqlclient: root open: %w", err)
}
// Conservative pool: root scans run at most a few queries per
// minute, sharing one idle connection is plenty.
db.SetMaxOpenConns(2)
db.SetMaxIdleConns(1)
db.SetConnMaxLifetime(5 * time.Minute)
rootMu.Lock()
defer rootMu.Unlock()
if rootDB != nil {
_ = db.Close()
return rootDB, nil
}
rootDB = db
return rootDB, nil
}
// loadRootCreds parses /root/.my.cnf's [client] section (or the
// path set via SetRootCnfPath) when the file exists. The format is the
// same INI subset mysql_config_editor / the official client honour:
// section header in `[client]`, key=value pairs (values may be
// unquoted, single quoted, or double quoted; matching quote is
// stripped).
//
// Missing files and omitted user keys fall back to user=root with an
// empty password so Unix socket-auth setups keep matching `mysql -e`.
func loadRootCreds() (Creds, error) {
rootMu.Lock()
path := rootPath
rootMu.Unlock()
creds := Creds{User: "root"}
// #nosec G304 -- path is operator-configured at init time, not
// attacker-controlled.
f, err := os.Open(path)
if err != nil {
if os.IsNotExist(err) {
return creds, nil
}
return Creds{}, fmt.Errorf("mysqlclient: open %s: %w", path, err)
}
defer f.Close()
inClient := false
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 0, 4096), 1<<20)
for scanner.Scan() {
line := strings.TrimSpace(scanner.Text())
if line == "" || strings.HasPrefix(line, "#") || strings.HasPrefix(line, ";") {
continue
}
if strings.HasPrefix(line, "[") && strings.HasSuffix(line, "]") {
inClient = strings.EqualFold(line, "[client]") || strings.EqualFold(line, "[mysql]")
continue
}
if !inClient {
continue
}
key, val, ok := strings.Cut(line, "=")
if !ok {
continue
}
key = strings.TrimSpace(key)
val = unquote(strings.TrimSpace(val))
switch strings.ToLower(key) {
case "user":
if val != "" {
creds.User = val
}
case "password", "pass":
creds.Password = val
case "host":
creds.Host = val
case "port":
if p, perr := strconv.Atoi(val); perr == nil {
creds.Port = p
}
case "socket":
creds.Socket = val
}
}
if err := scanner.Err(); err != nil {
return Creds{}, fmt.Errorf("mysqlclient: read %s: %w", path, err)
}
return creds, nil
}
func unquote(s string) string {
if len(s) >= 2 {
first, last := s[0], s[len(s)-1]
if (first == '"' && last == '"') || (first == '\'' && last == '\'') {
return s[1 : len(s)-1]
}
}
return s
}
package netutil
import (
"net"
"sync"
"time"
)
// hostAddrCacheTTL bounds how long a cached interface enumeration is reused.
// Alias IPs come and go on panel hosts, so the set cannot be read once at
// startup, but enumerating on every finding would be wasteful.
const hostAddrCacheTTL = 5 * time.Minute
// hostAddrLookup returns every IP bound to a local interface. Package-level so
// tests can inject deterministic addresses.
var hostAddrLookup = enumerateHostAddresses
var (
hostAddrMu sync.Mutex
hostAddrCache map[string]struct{}
hostAddrCachedAt time.Time
hostAddrCacheGood bool
hostAddrGeneration uint64
)
func enumerateHostAddresses() ([]net.IP, error) {
addrs, err := net.InterfaceAddrs()
if err != nil {
return nil, err
}
out := make([]net.IP, 0, len(addrs))
for _, a := range addrs {
switch v := a.(type) {
case *net.IPNet:
out = append(out, v.IP)
case *net.IPAddr:
out = append(out, v.IP)
}
}
return out, nil
}
// IsHostAddress reports whether ip is an address bound to one of this host's
// own interfaces.
//
// Traffic a machine sends to itself is not an attack on itself. cPanel hosts
// proxy nginx to Apache over the machine's public address rather than
// loopback, so without this guard every proxied request counts as inbound
// traffic from the server, and a site that answers its own cron with an error
// drives the host's own address up the local threat score.
//
// Loopback and link-local addresses deliberately do not count: callers that
// want to accept local traffic outright test for that separately, and folding
// it in here would let a check discard findings that really did originate on
// the box. A lookup failure fails open (reports false) so a transient syscall
// error cannot suppress every finding.
func IsHostAddress(ip string) bool {
parsed := net.ParseIP(ip)
if parsed == nil {
return false
}
if parsed.IsLoopback() || parsed.IsLinkLocalUnicast() || parsed.IsLinkLocalMulticast() {
return false
}
set, ok := hostAddresses()
if !ok {
return false
}
_, found := set[hostAddrKey(parsed)]
return found
}
// hostAddrKey normalizes an address so the v4-mapped-v6 and non-canonical
// textual forms of the same address share one key.
func hostAddrKey(ip net.IP) string {
if v4 := ip.To4(); v4 != nil {
return v4.String()
}
return ip.String()
}
func hostAddresses() (map[string]struct{}, bool) {
hostAddrMu.Lock()
if hostAddrCacheGood && time.Since(hostAddrCachedAt) < hostAddrCacheTTL {
set := hostAddrCache
hostAddrMu.Unlock()
return set, true
}
lookup, generation := hostAddrLookup, hostAddrGeneration
hostAddrMu.Unlock()
ips, err := lookup()
if err != nil {
return nil, false
}
set := make(map[string]struct{}, len(ips))
for _, ip := range ips {
if ip == nil || ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() {
continue
}
set[hostAddrKey(ip)] = struct{}{}
}
hostAddrMu.Lock()
defer hostAddrMu.Unlock()
// Replacing the source invalidates any lookup already in flight. Never
// let an old callback repopulate the replacement's cache.
if generation != hostAddrGeneration {
return nil, false
}
if hostAddrCacheGood && time.Since(hostAddrCachedAt) < hostAddrCacheTTL {
return hostAddrCache, true
}
hostAddrCache = set
hostAddrCachedAt = time.Now()
hostAddrCacheGood = true
return set, true
}
// SetHostAddressLookup swaps the interface enumeration and drops the cache,
// returning a function that restores the previous lookup. Tests in other
// packages use it to pin a deterministic address set.
func SetHostAddressLookup(fn func() ([]net.IP, error)) func() {
hostAddrMu.Lock()
prev := hostAddrLookup
hostAddrLookup = fn
hostAddrGeneration++
hostAddrCache = nil
hostAddrCacheGood = false
hostAddrMu.Unlock()
return func() {
hostAddrMu.Lock()
hostAddrLookup = prev
hostAddrGeneration++
hostAddrCache = nil
hostAddrCacheGood = false
hostAddrMu.Unlock()
}
}
package netutil
import (
"net"
"strings"
)
// ParseIPToken extracts an IP address from one whitespace-delimited log or
// message token: surrounding brackets and trailing punctuation are removed,
// and a trailing ':' separator is dropped only when the token does not
// already parse as an address, so the "::" that ends many IPv6 addresses
// survives. The result is the canonical text form; ok is false when no
// address is left. It replaces per-caller TrimRight cutsets that contained
// ':' and cut IPv6 addresses short without re-validating the remnant.
func ParseIPToken(token string) (string, bool) {
t := strings.Trim(strings.TrimSpace(token), ",;.!?()<>\"'")
if ip := parseIPLiteral(t); ip != nil {
return ip.String(), true
}
if host, _, err := net.SplitHostPort(t); err == nil {
if ip := net.ParseIP(strings.Trim(host, "[]")); ip != nil {
return ip.String(), true
}
}
for strings.HasSuffix(t, ":") {
t = strings.TrimSuffix(t, ":")
if ip := parseIPLiteral(t); ip != nil {
return ip.String(), true
}
}
return "", false
}
func parseIPLiteral(token string) net.IP {
if len(token) >= 2 && token[0] == '[' && token[len(token)-1] == ']' {
token = token[1 : len(token)-1]
}
return net.ParseIP(token)
}
// Package netutil holds the shared public-range guard used to validate
// operator- and vendor-supplied IP ranges. It lives in its own package so the
// special-use CIDR list and the public-IP predicate have a single definition
// instead of being copy-pasted across internal/config and internal/threatintel.
package netutil
import (
"net"
"strings"
)
// nonPublicSpecialUseNets are address blocks that are routable-looking but must
// never be treated as public crawler space: an allowlist entry inside any of
// them would let an attacker claim addresses CSM should always scan.
var nonPublicSpecialUseNets = mustParseCIDRs(
"0.0.0.0/8", // "this network"
"100.64.0.0/10", // carrier-grade NAT
"192.0.0.0/24", // IETF protocol assignments
"192.0.2.0/24", // documentation
"192.88.99.0/24", // deprecated 6to4 relay anycast
"198.18.0.0/15", // benchmarking
"198.51.100.0/24", // documentation
"203.0.113.0/24", // documentation
"240.0.0.0/4", // reserved
"100::/64", // discard-only
"2001:2::/48", // benchmarking
"2001:db8::/32", // documentation
"2002::/16", // 6to4
"64:ff9b::/96", // IPv4/IPv6 translation
"64:ff9b:1::/48", // IPv4/IPv6 translation
)
func mustParseCIDRs(cidrs ...string) []*net.IPNet {
out := make([]*net.IPNet, 0, len(cidrs))
for _, cidr := range cidrs {
_, n, err := net.ParseCIDR(cidr)
if err != nil {
panic(err)
}
out = append(out, n)
}
return out
}
// IPInAnyNet reports whether ip is contained in any of nets.
func IPInAnyNet(ip net.IP, nets []*net.IPNet) bool {
for _, n := range nets {
if n.Contains(ip) {
return true
}
}
return false
}
// IsPublicIP reports whether ip is a public, globally-routable unicast address:
// not private, link-local, multicast, or in any non-public special-use block
// (CGNAT, documentation, benchmarking, reserved, IPv6 special-use). A nil IP is
// never public.
func IsPublicIP(ip net.IP) bool {
return ip.IsGlobalUnicast() && !ip.IsPrivate() &&
!ip.IsLinkLocalUnicast() && !ip.IsLinkLocalMulticast() && !ip.IsMulticast() &&
!IPInAnyNet(ip, nonPublicSpecialUseNets)
}
// ParseCIDROrIP parses a CIDR or a bare IP, returning the normalized network. A
// bare IPv4 becomes a /32 and a bare IPv6 a /128. Returns nil on garbage.
func ParseCIDROrIP(s string) *net.IPNet {
s = strings.TrimSpace(s)
if _, n, err := net.ParseCIDR(s); err == nil {
return NormalizeIPNet(n)
}
if ip := net.ParseIP(s); ip != nil {
if v4 := ip.To4(); v4 != nil {
return &net.IPNet{IP: v4, Mask: net.CIDRMask(32, 32)}
}
return &net.IPNet{IP: ip, Mask: net.CIDRMask(128, 128)}
}
return nil
}
// NormalizeIPNet returns n with its IP masked to the network address and IPv4
// kept in 4-byte form. An IPv4-mapped IPv6 CIDR is reduced to its effective
// IPv4 prefix so Contains matches on the IPv4 form. Returns nil for a malformed
// mask/IP combination.
func NormalizeIPNet(n *net.IPNet) *net.IPNet {
if n == nil {
return nil
}
if v4 := n.IP.To4(); v4 != nil {
mask := n.Mask
if len(mask) == net.IPv6len {
if _, bits := mask.Size(); bits != net.IPv6len*8 {
return nil
}
// Go treats IPv4-mapped IPv6 CIDRs as IPv4 ranges for Contains.
// Keep validation and matching on that same effective prefix.
mask = net.IPMask(mask[12:])
}
if _, bits := mask.Size(); bits != net.IPv4len*8 {
return nil
}
return &net.IPNet{IP: v4.Mask(mask), Mask: mask}
}
ip := n.IP.To16()
if ip == nil {
return nil
}
if _, bits := n.Mask.Size(); bits != net.IPv6len*8 {
return nil
}
return &net.IPNet{IP: ip.Mask(n.Mask), Mask: n.Mask}
}
// Package obs centralises crash reporting and selective error capture via
// Sentry. Init is a one-shot called from the daemon entry point; after
// that, callers use Go/SafeGo to launch goroutines with panic recovery
// and Capture/CaptureMsg to forward selected errors.
//
// When Sentry is disabled or the DSN is empty, every function becomes a
// no-op wrapper with the same semantics as a plain `go func()` call,
// so guarded call sites work unchanged in tests and in operator
// builds that opt out of telemetry.
package obs
import (
"fmt"
"os"
"sync/atomic"
"time"
"github.com/getsentry/sentry-go"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/platform"
)
var enabled atomic.Bool
// flushTimeout bounds how long shutdown flushes block before giving up.
// Sentry enforces its own deadline on Flush; we pass this value so
// stuck HTTP calls don't hang systemd past TimeoutStopSec.
const flushTimeout = 2 * time.Second
// Init configures the Sentry SDK once at daemon startup. Returns nil if
// Sentry is disabled or the DSN is empty so callers can treat init as
// optional. Safe to call when cfg is nil.
func Init(cfg *config.Config, version, buildHash string) error {
if cfg == nil || !cfg.Sentry.Enabled || cfg.Sentry.DSN == "" {
return nil
}
env := cfg.Sentry.Environment
if env == "" {
env = "production"
}
rate := cfg.Sentry.SampleRate
if rate <= 0 {
rate = 1.0
}
hostname, _ := os.Hostname()
release := "csm@" + version
if buildHash != "" && buildHash != "unknown" {
release = release + "+" + buildHash
}
err := sentry.Init(sentry.ClientOptions{
Dsn: cfg.Sentry.DSN,
Environment: env,
Release: release,
ServerName: hostname,
SampleRate: rate,
TracesSampleRate: 0.0,
Debug: cfg.Sentry.Debug,
AttachStacktrace: true,
})
if err != nil {
return fmt.Errorf("sentry init: %w", err)
}
info := platform.Detect()
sentry.ConfigureScope(func(scope *sentry.Scope) {
scope.SetTag("os", string(info.OS))
if info.OSVersion != "" {
scope.SetTag("os_version", info.OSVersion)
}
if info.Panel != "" {
scope.SetTag("panel", string(info.Panel))
}
if info.WebServer != "" {
scope.SetTag("webserver", string(info.WebServer))
}
})
enabled.Store(true)
return nil
}
// Enabled reports whether Init succeeded and Sentry is live.
func Enabled() bool { return enabled.Load() }
// Flush waits for queued events to be sent before returning. Call from
// shutdown paths before exit. Safe to call when Sentry is disabled.
func Flush() {
if !enabled.Load() {
return
}
sentry.Flush(flushTimeout)
}
// Go launches fn in a new goroutine. A panic in fn is captured with the
// given component tag, the event is flushed, and the panic is
// re-raised so the existing crash-and-systemd-restart behavior is
// preserved. Use this for long-lived supervisor goroutines where a
// silent death would leave the daemon in a degraded state.
func Go(component string, fn func()) {
go func() {
defer func() {
if r := recover(); r != nil {
report(component, r)
panic(r)
}
}()
fn()
}()
}
// SafeGo is Go but swallows the panic after capture. Use this for
// per-request handlers (socket accept, HTTP request) where one bad
// input should not crash the whole daemon.
func SafeGo(component string, fn func()) {
go func() {
defer func() {
if r := recover(); r != nil {
report(component, r)
}
}()
fn()
}()
}
// Capture sends an error to Sentry with a component tag. No-op when
// disabled or err is nil. Reserve for unexpected states and invariant
// violations; expected-failure errors (permission denied, transient
// network) should stay out of Sentry to avoid noise.
func Capture(component string, err error) {
if !enabled.Load() || err == nil {
return
}
sentry.WithScope(func(scope *sentry.Scope) {
scope.SetTag("component", component)
sentry.CaptureException(err)
})
}
// CaptureMsg is Capture for string-only events (e.g. invariant
// violations without a wrapped error).
func CaptureMsg(component, msg string) {
if !enabled.Load() {
return
}
sentry.WithScope(func(scope *sentry.Scope) {
scope.SetTag("component", component)
sentry.CaptureMessage(msg)
})
}
func report(component string, r any) {
if !enabled.Load() {
return
}
sentry.WithScope(func(scope *sentry.Scope) {
scope.SetTag("component", component)
sentry.CurrentHub().Recover(r)
})
sentry.Flush(flushTimeout)
}
package phptaint
import (
"fmt"
"strings"
"unicode"
"unicode/utf8"
)
// Display bounds. Attacker-controlled identifiers must not be able to
// produce unbounded findings or logs.
const (
maxSegmentBytes = 64
maxChainSegments = 32
chainHeadKeep = 16
chainTailKeep = 15
)
// sanitizeSegment renders one untrusted display segment: valid UTF-8,
// non-printing and invalid bytes become '?', truncation happens at a rune
// boundary and the marker counts into the cap.
func sanitizeSegment(s string) string {
clean, _ := sanitize(s, maxSegmentBytes)
return clean
}
// sanitizeReason bounds Report.Reason the same way.
func sanitizeReason(s string) string {
clean, _ := sanitize(s, MaxReasonBytes)
return clean
}
func sanitize(s string, maxBytes int) (string, bool) {
s = strings.ToValidUTF8(s, "?")
s = strings.Map(func(r rune) rune {
// IsPrint excludes format controls (including bidi overrides) and the
// Unicode line/paragraph separators as well as ordinary control bytes.
if !unicode.IsPrint(r) {
return '?'
}
return r
}, s)
if len(s) <= maxBytes {
return s, false
}
cut := maxBytes - 3
for cut > 0 && !utf8.RuneStart(s[cut]) {
cut--
}
return s[:cut] + "...", true
}
// truncateChain sanitizes and bounds every segment, then keeps a bounded head
// and tail with one marker naming any omitted segments.
func truncateChain(via []string) ([]string, bool) {
bounded := make([]string, len(via))
truncated := false
for i, segment := range via {
var cut bool
bounded[i], cut = sanitize(segment, maxSegmentBytes)
truncated = truncated || cut
}
if len(bounded) <= maxChainSegments {
return bounded, truncated
}
omitted := len(bounded) - chainHeadKeep - chainTailKeep
out := make([]string, 0, maxChainSegments)
out = append(out, bounded[:chainHeadKeep]...)
out = append(out, fmt.Sprintf("... %d segment(s) omitted ...", omitted))
out = append(out, bounded[len(bounded)-chainTailKeep:]...)
return out, true
}
package phptaint
import (
"sort"
"strings"
"sync/atomic"
"github.com/VKCOM/php-parser/pkg/ast"
"github.com/VKCOM/php-parser/pkg/visitor"
"github.com/VKCOM/php-parser/pkg/visitor/traverser"
)
// callSinks are code-execution sinks that appear as ordinary function calls,
// mapped to the index of the argument that gets executed. assert() evaluates
// its first argument on PHP 7, which hosts still run; create_function()
// evaluates the code in its second argument.
var callSinks = map[string]int{"assert": 0, "create_function": 1}
// alwaysBoolPredicates are PHP builtins whose return type is bool for every
// input. A call to one cannot hand assert() a string to execute.
var alwaysBoolPredicates = map[string]bool{
"is_array": true, "is_bool": true, "is_callable": true, "is_countable": true,
"is_double": true, "is_float": true, "is_int": true, "is_integer": true,
"is_iterable": true, "is_long": true, "is_null": true, "is_numeric": true,
"is_object": true, "is_resource": true, "is_scalar": true, "is_string": true,
"is_subclass_of": true, "is_a": true,
"array_key_exists": true, "in_array": true, "property_exists": true,
"method_exists": true, "function_exists": true, "class_exists": true,
"interface_exists": true, "defined": true, "file_exists": true,
"is_dir": true, "is_file": true, "is_readable": true, "is_writable": true,
"str_contains": true, "str_starts_with": true, "str_ends_with": true,
"ctype_digit": true, "ctype_alpha": true, "ctype_alnum": true,
}
// assertArgumentCouldBeString reports whether e's top-level shape permits a
// string value. This is exactly what decides whether assert(e) can execute
// code on any PHP version this analyzer targets: PHP 7 only ever evaluated a
// *string* argument as code, and PHP 8 removed the eval form entirely. A
// logical, comparison, identity, or instanceof expression can only ever
// produce a bool (or, for <=>, an int) -- never a string -- so it is not a
// code-execution sink regardless of what it compares.
//
// The check is shallow by design: only e's own node type is inspected, not
// its operands. assert($a && fetchCode()) must still be excluded, because
// the value actually passed to assert() is the bool && produces, not the
// string one of its operands would have produced on its own.
func assertArgumentCouldBeString(e ast.Vertex) bool {
// A call to a predicate returns a bool whatever its argument was, so the
// value assert() receives can never be a string. This stays shallow for
// the same reason as the operator cases: only the outermost node decides
// what assert() is actually handed.
if call, ok := e.(*ast.ExprFunctionCall); ok {
return !alwaysBoolPredicates[calleeName(call.Function)]
}
if _, ok := e.(*ast.ExprEmpty); ok {
return false
}
if _, ok := e.(*ast.ExprIsset); ok {
return false
}
switch e.(type) {
case *ast.ExprBinaryBooleanAnd, *ast.ExprBinaryBooleanOr,
*ast.ExprBinaryLogicalAnd, *ast.ExprBinaryLogicalOr, *ast.ExprBinaryLogicalXor,
*ast.ExprBooleanNot,
*ast.ExprBinaryEqual, *ast.ExprBinaryNotEqual,
*ast.ExprBinaryIdentical, *ast.ExprBinaryNotIdentical,
*ast.ExprBinaryGreater, *ast.ExprBinaryGreaterOrEqual,
*ast.ExprBinarySmaller, *ast.ExprBinarySmallerOrEqual,
*ast.ExprBinarySpaceship,
*ast.ExprInstanceOf:
return false
}
return true
}
// sinkSite is one code-execution construct and the expression it executes.
type sinkSite struct {
kind string
expr ast.Vertex
}
type callSite struct {
name string
node ast.Vertex
}
// scopeFacts is everything the taint pass needs from one lexical scope.
// Collection is one pass; the analysis reads it repeatedly.
type scopeFacts struct {
assigns []*ast.ExprAssign
references []*ast.ExprAssignReference
concats []*ast.ExprAssignConcat
returns []*ast.StmtReturn
funcs []*ast.StmtFunction
methods []*ast.StmtClassMethod
// closures and arrowFuncs hold every closure and arrow function found in
// this scope. Both are declarations in the same sense funcs and methods
// are: each gets its own entry in declarationTree (so its span is
// excluded from the enclosing scope) and its own per-scope analysis in
// analyze, exactly like a named function's body. Without both halves, a
// closure parameter or local reassignment that merely shares a name with
// an outer variable would either borrow that outer variable's taint (if
// only excluded, never analysed) or stop being examined for its own
// sinks (if only analysed, never excluded) -- see analyze's closure and
// arrow-function loops for the second half.
closures []*ast.ExprClosure
arrowFuncs []*ast.ExprArrowFunction
// classLikes holds every class, interface, trait, and enum declaration
// (named or anonymous) found in this scope. Their own contents are
// never read from here directly -- methods are already tracked in
// methods above -- this list exists only to supply declarationTree
// with the position span of the class body itself.
classLikes []ast.Vertex
sinks []sinkSite
callNodes []*ast.ExprFunctionCall
callSites []callSite
calls map[string]bool
vars map[string]bool
varNodes []*ast.ExprVariable
// propNodes holds every property-fetch node (both "->" and "?->") found
// in this scope, whether it appears as a read or as an assignment
// target. readVarNodes keys each one to the specific property it fetches
// (see assignedTargetKey) so a write to one property never taints a read
// of a different one or of the bare base object.
propNodes []ast.Vertex
writes []ast.Vertex
// fileWrites holds every file_put_contents(path, data) call in this
// scope. It is the bridge the fetch-write-include dropper relies on:
// remote content lands in a local file that an include then executes
// by path, with no variable ever carrying the taint into the sink.
fileWrites []fileWriteSite
precisionLoss map[string]bool
visited int
budgetExceeded bool
// declTreeCache memoizes declarationTree's result for this scopeFacts
// instance. Safe to cache: a scopeFacts is fully populated by the
// traversal in collectScope/collectAll/collectOwnStmts before any
// caller can reach it and is never mutated afterward, and each
// instance is freshly allocated per collection call, so nothing here
// carries state across separate Analyze invocations. This is what lets
// functionSummaries call f.declarationTree() again on the same f
// analyze already indexed without repeating the O(D log D) build.
declTreeCache *declTree
}
func newScopeFacts() *scopeFacts {
return &scopeFacts{
calls: map[string]bool{},
vars: map[string]bool{},
precisionLoss: map[string]bool{},
}
}
// count enforces the collection budget. Once exceeded, collection stops
// recording so a hostile file cannot grow analysis state without bound.
func (f *scopeFacts) count() bool {
if f.budgetExceeded {
return false
}
f.visited++
if f.visited > maxCollectedNodes {
f.budgetExceeded = true
return false
}
return true
}
// factVisitor implements only the node types the analysis needs; every other
// node type is a no-op inherited from visitor.Null.
type factVisitor struct {
visitor.Null
f *scopeFacts
functionAliases map[string]string
// exclude, when set, marks the position spans of nested declarations
// this scope must not absorb facts from. The library traverser recurses
// unconditionally through control-flow wrappers (if/while/switch/try/
// foreach/...), so a declaration nested inside one of those -- the
// WordPress `if (!function_exists(...))` guard is the common case -- is
// still reached by this same traversal even though it belongs to a
// separately analysed scope. Position-based exclusion catches it
// regardless of which wrapper (or how many, nested how deep) sits
// between this scope and the declaration; a filter keyed on the direct
// statement type of the top of the list cannot, because it only ever
// sees the wrapper, never what the wrapper contains.
exclude *spanIndex
}
// excluded reports whether n's position falls inside a nested declaration
// this scope must not record facts from. A node without a determinable
// position is not excluded: parsed nodes normally always have positions,
// and failing open (keep, don't drop) matches the conservative choice
// already made in readVarNodes for the same edge case -- an unrecordable
// position must never cost the enclosing scope one of its own facts.
func (v *factVisitor) excluded(n ast.Vertex) bool {
if v.exclude == nil {
return false
}
span, ok := spanOf(n)
if !ok {
return false
}
return v.exclude.contains(span)
}
func (v *factVisitor) StmtNamespace(*ast.StmtNamespace) {
v.functionAliases = nil
}
func (v *factVisitor) StmtUse(n *ast.StmtUseList) {
for _, useNode := range n.Uses {
use, ok := useNode.(*ast.StmtUse)
if ok {
v.addFunctionAlias(n.Type, nil, use)
}
}
}
func (v *factVisitor) StmtGroupUse(n *ast.StmtGroupUseList) {
prefix, _ := n.Prefix.(*ast.Name)
for _, useNode := range n.Uses {
use, ok := useNode.(*ast.StmtUse)
if ok {
v.addFunctionAlias(n.Type, prefix, use)
}
}
}
func (v *factVisitor) addFunctionAlias(listType ast.Vertex, prefix *ast.Name, use *ast.StmtUse) {
useType := listType
if use.Type != nil {
useType = use.Type
}
id, ok := useType.(*ast.Identifier)
if !ok || !strings.EqualFold(string(id.Value), "function") {
return
}
name, ok := use.Use.(*ast.Name)
if !ok || len(name.Parts) == 0 {
return
}
parts := make([]ast.Vertex, 0, len(name.Parts)+4)
if prefix != nil {
parts = append(parts, prefix.Parts...)
}
parts = append(parts, name.Parts...)
target := joinNameParts(parts)
alias := ""
if use.Alias != nil {
alias = calleeName(use.Alias)
} else if last, ok := name.Parts[len(name.Parts)-1].(*ast.NamePart); ok {
alias = strings.ToLower(string(last.Value))
}
if alias == "" || target == "" {
return
}
if v.functionAliases == nil {
v.functionAliases = map[string]string{}
}
v.functionAliases[alias] = target
}
func (v *factVisitor) ExprVariable(n *ast.ExprVariable) {
if !v.f.count() || v.excluded(n) {
return
}
name := varName(n.Name)
v.f.vars[name] = true
v.f.varNodes = append(v.f.varNodes, n)
if name == "" {
v.f.precisionLoss["variable-variable"] = true
}
}
// ExprPropertyFetch and ExprNullsafePropertyFetch record every "->"/"?->"
// property fetch this scope contains, whether it is a read or (for the
// non-nullsafe form; nullsafe cannot appear as a write target in valid PHP)
// an assignment target. The traverser still visits the embedded base
// variable (e.g. $this inside $this->body) as an ordinary ExprVariable
// regardless of this hook, which is what lets a genuinely whole-object
// taint keep reaching a property read -- see readVarNodes.
func (v *factVisitor) ExprPropertyFetch(n *ast.ExprPropertyFetch) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.propNodes = append(v.f.propNodes, n)
}
func (v *factVisitor) ExprNullsafePropertyFetch(n *ast.ExprNullsafePropertyFetch) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.propNodes = append(v.f.propNodes, n)
}
func (v *factVisitor) ExprFunctionCall(n *ast.ExprFunctionCall) {
if !v.f.count() || v.excluded(n) {
return
}
name, resolved := v.resolveFunctionAlias(n)
v.f.calls[name] = true
v.f.callNodes = append(v.f.callNodes, resolved)
v.f.callSites = append(v.f.callSites, callSite{name: name, node: n})
if name == "" {
v.f.precisionLoss["dynamic-call"] = true
}
// assert() and create_function() are ordinary calls in the grammar, not
// dedicated nodes, so they are recognised here rather than by node type.
// The executed argument differs per sink (assert's is first,
// create_function's is second), so the index is looked up rather than
// assumed to be the last argument. assert() additionally requires its
// argument to be capable of holding a string: see
// assertArgumentCouldBeString for why a boolean-shaped argument (a
// comparison, a logical operator, instanceof, ...) is not a
// code-execution sink on any PHP version this analyzer targets.
if sink, ok := callSinkSite(name, n); ok {
v.f.sinks = append(v.f.sinks, sink)
}
if w, ok := callFileWriteSite(name, n); ok {
v.f.fileWrites = append(v.f.fileWrites, w)
}
}
// fileWriteSite is one file_put_contents(path, data) call.
type fileWriteSite struct {
path ast.Vertex
data ast.Vertex
}
// fileWriteCalls write their second argument to the path in their first.
var fileWriteCalls = map[string]bool{"file_put_contents": true}
func callFileWriteSite(name string, node ast.Vertex) (fileWriteSite, bool) {
call, ok := node.(*ast.ExprFunctionCall)
if !ok || !fileWriteCalls[name] || len(call.Args) < 2 {
return fileWriteSite{}, false
}
pathArg, okPath := call.Args[0].(*ast.Argument)
dataArg, okData := call.Args[1].(*ast.Argument)
if !okPath || !okData {
return fileWriteSite{}, false
}
return fileWriteSite{path: pathArg.Expr, data: dataArg.Expr}, true
}
func callSinkSite(name string, node ast.Vertex) (sinkSite, bool) {
call, ok := node.(*ast.ExprFunctionCall)
if !ok {
return sinkSite{}, false
}
idx, ok := callSinks[name]
if !ok || len(call.Args) <= idx {
return sinkSite{}, false
}
arg, ok := call.Args[idx].(*ast.Argument)
if !ok || (name == "assert" && !assertArgumentCouldBeString(arg.Expr)) {
return sinkSite{}, false
}
return sinkSite{kind: name, expr: arg.Expr}, true
}
// resolveFunctionAlias resolves a call target through this scope's
// function-alias table (populated from `use function` imports) and reports
// the canonical name. When the call is aliased, it returns a detached copy
// of the call node whose Function names the canonical target, so downstream
// name lookups (sourceGrade, decoder checks) see the resolved name
// without the parsed tree itself ever being rewritten mid-traversal: the
// traverser reads call.Function right after visiting call, so mutating it in
// place would orphan the original alias-name subtree from the walk in
// progress. The copy shares Args and Position with the original node, so
// argument inspection and source-span correlation are unaffected.
func (v *factVisitor) resolveFunctionAlias(call *ast.ExprFunctionCall) (string, *ast.ExprFunctionCall) {
name := calleeName(call.Function)
plain, ok := call.Function.(*ast.Name)
if !ok || len(plain.Parts) != 1 {
return name, call
}
target, ok := v.functionAliases[name]
if !ok {
return name, call
}
parts := strings.Split(target, "\\")
canonical := make([]ast.Vertex, 0, len(parts))
for _, part := range parts {
canonical = append(canonical, &ast.NamePart{Value: []byte(part)})
}
resolved := *call
resolved.Function = &ast.NameFullyQualified{Parts: canonical}
return target, &resolved
}
func (v *factVisitor) ExprMethodCall(n *ast.ExprMethodCall) {
v.call(calleeName(n.Method), n)
}
func (v *factVisitor) ExprNullsafeMethodCall(n *ast.ExprNullsafeMethodCall) {
v.call(calleeName(n.Method), n)
}
func (v *factVisitor) ExprStaticCall(n *ast.ExprStaticCall) {
v.call(calleeName(n.Call), n)
}
func (v *factVisitor) call(name string, node ast.Vertex) {
if !v.f.count() || v.excluded(node) {
return
}
v.f.calls[name] = true
v.f.callSites = append(v.f.callSites, callSite{name: name, node: node})
if name == "" {
v.f.precisionLoss["dynamic-call"] = true
}
}
func (v *factVisitor) ExprAssign(n *ast.ExprAssign) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.assigns = append(v.f.assigns, n)
v.f.writes = append(v.f.writes, n.Var)
}
func (v *factVisitor) ExprAssignReference(n *ast.ExprAssignReference) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.references = append(v.f.references, n)
v.f.writes = append(v.f.writes, n.Var)
}
func (v *factVisitor) ExprAssignConcat(n *ast.ExprAssignConcat) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.concats = append(v.f.concats, n)
}
func (v *factVisitor) StmtReturn(n *ast.StmtReturn) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.returns = append(v.f.returns, n)
}
func (v *factVisitor) StmtFunction(n *ast.StmtFunction) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.funcs = append(v.f.funcs, n)
}
func (v *factVisitor) StmtClassMethod(n *ast.StmtClassMethod) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.methods = append(v.f.methods, n)
}
func (v *factVisitor) ExprClosure(n *ast.ExprClosure) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.closures = append(v.f.closures, n)
}
func (v *factVisitor) ExprArrowFunction(n *ast.ExprArrowFunction) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.arrowFuncs = append(v.f.arrowFuncs, n)
}
func (v *factVisitor) StmtClass(n *ast.StmtClass) { v.classLike(n) }
func (v *factVisitor) StmtInterface(n *ast.StmtInterface) { v.classLike(n) }
func (v *factVisitor) StmtTrait(n *ast.StmtTrait) { v.classLike(n) }
func (v *factVisitor) StmtEnum(n *ast.StmtEnum) { v.classLike(n) }
// classLike records a class, interface, trait, or enum declaration's span.
// n's Name is nil for an anonymous class (`new class { ... }`); the node
// itself still carries a real position, so anonymous classes are covered by
// the same declaration-span exclusion as named ones, with no separate case.
func (v *factVisitor) classLike(n ast.Vertex) {
if !v.f.count() || v.excluded(n) {
return
}
v.f.classLikes = append(v.f.classLikes, n)
}
func (v *factVisitor) ExprEval(n *ast.ExprEval) { v.sink("eval", n.Expr) }
func (v *factVisitor) ExprInclude(n *ast.ExprInclude) { v.sink("include", n.Expr) }
func (v *factVisitor) ExprIncludeOnce(n *ast.ExprIncludeOnce) { v.sink("include_once", n.Expr) }
func (v *factVisitor) ExprRequire(n *ast.ExprRequire) { v.sink("require", n.Expr) }
func (v *factVisitor) ExprRequireOnce(n *ast.ExprRequireOnce) { v.sink("require_once", n.Expr) }
func (v *factVisitor) sink(kind string, expr ast.Vertex) {
if !v.f.count() || v.excluded(expr) {
return
}
v.f.sinks = append(v.f.sinks, sinkSite{kind: kind, expr: expr})
}
// collectScope gathers facts from one subtree using the parser library's own
// traverser. The library traverser cannot be stopped mid-walk, so it visits
// every node in the subtree regardless of budget; scopeFacts.count() bounds
// the actual cost by making every visitor method a no-op once the budget is
// exceeded, so an over-budget file does no further per-node work beyond the
// traversal dispatch itself. That traverser was measured surviving 2,000,000
// levels of nesting without a Go stack overflow, so a hand-rolled iterative
// walker is not needed to protect against attacker-controlled AST depth.
func collectScope(n ast.Vertex) *scopeFacts {
f := newScopeFacts()
if n == nil {
return f
}
traverser.NewTraverser(&factVisitor{f: f}).Traverse(n)
return f
}
// collectTopLevel gathers facts from the whole file, excluding anything
// positioned inside a declaration named in exclude. Folding a declaration's
// body into the flat top-level map lets a variable local to it taint an
// unrelated top-level variable that merely shares its name, which reports
// clean code as malicious: names like $data, $content and $tmp recur
// constantly in real PHP. exclude is normally built by
// declarationTree(...).exclusionFor(nil) over the whole file, so every
// declaration in the file is excluded (there is no declaration whose own
// statements this call needs to keep).
func collectTopLevel(root ast.Vertex, exclude *spanIndex) *scopeFacts {
r, ok := root.(*ast.Root)
if !ok {
return collectScope(root)
}
return collectOwnStmts(r.Stmts, exclude)
}
// collectOwnStmts gathers facts from a statement list -- a function body, a
// method body, or the file's top-level statements -- excluding anything
// positioned inside a nested declaration named in exclude. This is the same
// cross-scope leak collectTopLevel guards against, one level deeper: a
// function can declare another function (or a class, including an
// anonymous one) in its own body, and that nested declaration gets its own
// entry in the whole-file funcs/methods inventory, analysed separately with
// its own taint state. Folding it into this scope's flat map too would let
// its local variables taint an identically-named local in the enclosing
// body.
//
// Filtering happens by position, not by skipping statements of a
// declaration type at the top of stmts, because the library traverser
// recurses unconditionally through control-flow wrappers: a function
// declared inside `if (!function_exists('f')) { function f() {...} }` --
// the standard WordPress conditional-declaration guard -- is not a direct
// member of stmts, so a type-based filter on stmts itself would miss it
// (and miss it again for every layer of if/while/switch/try/foreach it is
// wrapped in). Every statement is traversed unconditionally here; exclude,
// consulted per node inside factVisitor, is what actually drops facts whose
// position falls inside one of those nested declarations, regardless of
// how they are reached.
func collectOwnStmts(stmts []ast.Vertex, exclude *spanIndex) *scopeFacts {
f := newScopeFacts()
t := traverser.NewTraverser(&factVisitor{f: f, exclude: exclude})
for _, stmt := range stmts {
if stmt != nil {
t.Traverse(stmt)
}
}
return f
}
// methodStmts unwraps a class method's single body vertex into a statement
// list, the shape collectOwnStmts and collectAll both expect. A braced
// method body is *ast.StmtStmtList; an abstract or interface method has a
// nil Stmt.
func methodStmts(stmt ast.Vertex) []ast.Vertex {
switch v := stmt.(type) {
case nil:
return nil
case *ast.StmtStmtList:
return v.Stmts
default:
return []ast.Vertex{stmt}
}
}
// arrowFunctionBody wraps an arrow function's single implicit-return
// expression into the one-element statement list collectOwnStmts expects,
// the same shape methodStmts gives a method body. Unlike a closure (whose
// Stmts is already a real statement list), `fn($x) => expr` has no braces or
// statements to unwrap -- expr itself is both the body and the return value.
func arrowFunctionBody(n *ast.ExprArrowFunction) []ast.Vertex {
if n.Expr == nil {
return nil
}
return []ast.Vertex{n.Expr}
}
// closureCaptureNames names the outer bindings a closure's use() clause
// receives. Only Uses is consulted, never Params: a parameter that happens to
// share an outer variable's name is the closure's own binding and receives
// nothing from the enclosing scope, which is the shadowing case that scoping a
// closure's body exists to keep clean.
func closureCaptureNames(cl *ast.ExprClosure) map[string]bool {
names := make(map[string]bool, len(cl.Uses))
for _, u := range cl.Uses {
use, ok := u.(*ast.ExprClosureUse)
if !ok {
continue
}
v, ok := use.Var.(*ast.ExprVariable)
if !ok {
continue
}
if name := varName(v.Name); name != "" {
names[name] = true
}
}
return names
}
// paramNames names the plain-variable parameters of a parameter list. A
// parameter with any other shape is skipped rather than guessed at.
func paramNames(params []ast.Vertex) map[string]bool {
names := make(map[string]bool, len(params))
for _, p := range params {
param, ok := p.(*ast.Parameter)
if !ok {
continue
}
v, ok := param.Var.(*ast.ExprVariable)
if !ok {
continue
}
if name := varName(v.Name); name != "" {
names[name] = true
}
}
return names
}
// arrowCaptureNames names the outer bindings an arrow function receives.
// Unlike a closure, which lists them in a use() clause, fn(...) => expr
// captures by value every enclosing variable its body mentions. body.vars
// already lists every variable name the body reads, so subtracting the arrow
// function's own parameters leaves exactly the names that must have come from
// outside. Static arrows are the one exception: they still capture ordinary
// variables, but PHP deliberately does not bind $this to them. Over-inclusive
// by design: a name that is not in fact tainted in the enclosing scope simply
// produces no marker.
func arrowCaptureNames(af *ast.ExprArrowFunction, body *scopeFacts) map[string]bool {
params := paramNames(af.Params)
names := make(map[string]bool, len(body.vars))
for name := range body.vars {
if name != "" && !params[name] && (name != "this" || af.StaticTkn == nil) {
names[name] = true
}
}
return names
}
type nodeSpan struct {
start int
end int
}
type declarationSpan struct {
nodeSpan
node ast.Vertex
}
type namedNodeSpan struct {
nodeSpan
name string
node ast.Vertex
}
// resolvedCallIndex restores the canonical call names from a whole-file
// collection to independently collected lexical scopes. Namespace-level
// `use function` imports are outside a function or method body, so collecting
// that body alone cannot resolve its aliases. Exact source spans join the two
// views without importing variables, assignments, or calls from another
// scope.
type resolvedCallIndex struct {
functions map[nodeSpan]*ast.ExprFunctionCall
sites map[nodeSpan]callSite
}
func newResolvedCallIndex(f *scopeFacts) resolvedCallIndex {
index := resolvedCallIndex{
functions: make(map[nodeSpan]*ast.ExprFunctionCall, len(f.callNodes)),
sites: make(map[nodeSpan]callSite, len(f.callSites)),
}
for _, call := range f.callNodes {
if span, ok := spanOf(call); ok {
index.functions[span] = call
}
}
for _, call := range f.callSites {
if span, ok := spanOf(call.node); ok {
index.sites[span] = call
}
}
return index
}
func (index resolvedCallIndex) apply(f *scopeFacts) *scopeFacts {
out := *f
out.callNodes = append([]*ast.ExprFunctionCall(nil), f.callNodes...)
for i, call := range out.callNodes {
if span, ok := spanOf(call); ok {
if resolved, found := index.functions[span]; found {
out.callNodes[i] = resolved
}
}
}
out.callSites = append([]callSite(nil), f.callSites...)
for i, call := range out.callSites {
if span, ok := spanOf(call.node); ok {
if resolved, found := index.sites[span]; found {
out.callSites[i] = resolved
}
}
}
out.calls = make(map[string]bool, len(out.callSites))
for _, call := range out.callSites {
out.calls[call.name] = true
}
out.sinks = make([]sinkSite, 0, len(f.sinks))
for _, sink := range f.sinks {
if _, callSink := callSinks[sink.kind]; !callSink {
out.sinks = append(out.sinks, sink)
}
}
for _, call := range out.callSites {
if sink, ok := callSinkSite(call.name, call.node); ok {
out.sinks = append(out.sinks, sink)
}
}
return &out
}
// declTree links each declaration in a whole-file collection to the
// position spans of its IMMEDIATE child declarations only, plus the total
// declaration count. Built once per file by declarationTree, it lets every
// scope's exclusion index be produced by a map lookup instead of scanning
// and re-sorting the whole file's declaration list once per declaration.
type declTree struct {
// children maps a declaration node -- or nil, for the top-level scope --
// to the spans of its immediate child declarations.
children map[ast.Vertex][]nodeSpan
// parent is the reverse: each indexed declaration to the declaration
// immediately enclosing it, or the nil interface when the top-level
// scope encloses it. Captured during the same sweep that builds
// children, since the sweep already knows both ends of the edge. It
// answers the question children cannot: given a closure, which scope did
// it capture its outer bindings FROM.
parent map[ast.Vertex]ast.Vertex
// ordered holds the same declarations in source nesting order. Retaining
// the already-sorted sweep input lets capture analysis walk lexical scopes
// without rebuilding or re-sorting the attacker-controlled declaration
// list.
ordered []declarationSpan
// count is the total number of declarations indexed, checked against
// maxDeclarations before any per-scope work begins.
count int
}
// declTreeBuilds counts how many times declarationTree has actually
// performed the O(D log D) sort-and-sweep build, as opposed to returning an
// already-cached result. It exists solely so a same-package white-box test
// can observe that this count stays a small constant as declaration count
// grows -- the structural invariant this whole file's design establishes --
// rather than asserting elapsed wall-clock time, which is flaky under
// shared CI load and, worse, would not fail reliably if a per-declaration
// rebuild were reintroduced at a scale too small to visibly stall a test
// run. It is read-only from every caller's perspective (nothing in this
// package branches on its value, and it never influences a Report), so it
// is test-only observation of internal behaviour, not the kind of mutable
// process-global configuration or cross-call state this package's purity
// contract forbids.
var declTreeBuilds atomic.Int64
// declarationTree indexes every declaration recorded in f: functions,
// methods, classes/interfaces/traits/enums (anonymous classes included, via
// classLikes), closures, and arrow functions. f must come from an
// unfiltered, whole-file collection (collectScope(root)) so every
// declaration in the file is present, regardless of how deeply any of them
// is nested inside another or wrapped in control flow. The result is cached
// on f (see scopeFacts.declTreeCache), so calling this more than once on the
// same f -- analyze and functionSummaries both do, independently -- costs
// one build, not two.
//
// Sorting happens ONCE, over all D declarations together, rather than once
// per declaration: a per-declaration rebuild-and-sort of the whole span
// list costs O(D) work times D declarations, O(D^2 log D) overall, which is
// a real CPU-exhaustion surface against attacker-controlled PHP source (a
// file of trivial one-line function declarations reaches tens of thousands
// of them well within MaxSourceBytes and maxCollectedNodes). Declarations in
// a legitimately parsed file nest properly -- one is either fully disjoint
// from another or fully contained inside it, never a partial overlap -- so
// a single stack sweep over the sorted spans (the same interval-stack
// technique distributeOrigins already uses in taint.go) assigns each
// declaration to its immediate parent in one pass: pop any open declaration
// that has already closed before this one starts, and whatever remains on
// top of the stack (if anything) is the immediate parent.
//
// Indexing only immediate children, not every descendant, is what keeps the
// total cost linearithmic in D: a fact positioned inside a deeper
// descendant is already covered by its immediate parent's span, so nothing
// beyond direct children is ever needed to exclude an arbitrarily deep
// nested declaration (see exclusionFor). Each declaration contributes to
// exactly one parent's child list, so those lists sum to D across the whole
// file, and sorting each of them once, at exclusionFor time, costs at most
// D log D in total across every scope in the file.
func (f *scopeFacts) declarationTree() declTree {
if f.declTreeCache != nil {
return *f.declTreeCache
}
declTreeBuilds.Add(1)
nodes := make([]declarationSpan, 0, len(f.funcs)+len(f.methods)+len(f.classLikes)+len(f.closures)+len(f.arrowFuncs))
add := func(n ast.Vertex) {
if span, ok := spanOf(n); ok {
nodes = append(nodes, declarationSpan{nodeSpan: span, node: n})
}
}
for _, fn := range f.funcs {
add(fn)
}
for _, m := range f.methods {
add(m)
}
for _, c := range f.classLikes {
add(c)
}
for _, cl := range f.closures {
add(cl)
}
for _, af := range f.arrowFuncs {
add(af)
}
sort.Slice(nodes, func(i, j int) bool {
if nodes[i].start != nodes[j].start {
return nodes[i].start < nodes[j].start
}
return nodes[i].end > nodes[j].end
})
tree := declTree{
children: make(map[ast.Vertex][]nodeSpan, len(nodes)+1),
parent: make(map[ast.Vertex]ast.Vertex, len(nodes)),
ordered: nodes,
count: len(nodes),
}
stack := make([]int, 0, len(nodes))
for i, n := range nodes {
// A declaration's end position is EXCLUSIVE, so a declaration
// starting exactly where the previous one ends is its sibling, not
// its child. PHP allows `}function` with nothing between, and
// minifiers emit it, so `<` here would silently reparent every such
// pair and leak the second one's locals into the enclosing scope.
for len(stack) > 0 && nodes[stack[len(stack)-1]].end <= n.start {
stack = stack[:len(stack)-1]
}
var parent ast.Vertex // nil selects the top-level scope's own children
if len(stack) > 0 {
parent = nodes[stack[len(stack)-1]].node
}
tree.children[parent] = append(tree.children[parent], n.nodeSpan)
tree.parent[n.node] = parent
stack = append(stack, i)
}
f.declTreeCache = &tree
return tree
}
// exclusionFor builds the spanIndex scope self must exclude from its own
// collection: the spans of self's immediate child declarations only (self
// nil selects the top-level scope's own children). See declarationTree for
// why immediate children alone are sufficient regardless of nesting depth.
func (t declTree) exclusionFor(self ast.Vertex) spanIndex {
return newSpanIndex(t.children[self])
}
// withoutNestedDeclarationVars returns a view of f whose VARIABLE facts drop
// anything positioned inside one of the excluded declaration spans, leaving
// its CALL facts whole.
//
// A return expression is collected without excluding nested declarations, and
// deliberately so: that is what lets a callee invoked inside a closure within
// the return contribute a dependency edge and a summary. But the taint state
// that expression is graded against belongs to the ENCLOSING body, and a
// variable belonging to a nested declaration -- a closure's own local, an
// arrow function's parameter -- was never that state's variable. Names like
// $data and $content recur constantly in real PHP, so a bare name collision
// is otherwise enough to report a flow on a file where the two never meet.
//
// Nodes without position information are retained, matching the rest of this
// package: an unplaceable fact fails open, toward keeping a dependency rather
// than silently dropping one.
//
// This is a trade, not a pure win, and the losing side is worth naming. A
// closure that is immediately invoked really does yield what it captured, so
// `return (function() use ($tainted) { return $tainted; })();` stops being
// reported. Taint the closure PRODUCES by calling a source inside itself is
// still followed, because calls are not filtered; only taint carried IN by a
// capture is lost. Every such case still records closure-capture precision
// loss, so the reduction is reported rather than silent, and it buys the
// removal of a false positive that fires whenever a nested declaration's
// parameter merely shares a name with a tainted variable outside it.
func (f *scopeFacts) withoutNestedDeclarationVars(exclude *spanIndex) *scopeFacts {
// Most functions have no nested declarations. Avoid copying three fact
// slices on every return evaluation when the exclusion index cannot
// possibly remove anything.
if exclude == nil || len(exclude.spans) == 0 {
return f
}
keep := func(n ast.Vertex) bool {
span, ok := spanOf(n)
return !ok || !exclude.contains(span)
}
out := *f
out.varNodes = make([]*ast.ExprVariable, 0, len(f.varNodes))
for _, n := range f.varNodes {
if keep(n) {
out.varNodes = append(out.varNodes, n)
}
}
out.propNodes = make([]ast.Vertex, 0, len(f.propNodes))
for _, n := range f.propNodes {
if keep(n) {
out.propNodes = append(out.propNodes, n)
}
}
out.writes = make([]ast.Vertex, 0, len(f.writes))
for _, n := range f.writes {
if keep(n) {
out.writes = append(out.writes, n)
}
}
// vars feeds the evidence strings only, never a taint decision, but it is
// rebuilt from the surviving nodes anyway so a report never names a
// variable that belongs to a nested declaration and was excluded from the
// grading that produced the report.
out.vars = make(map[string]bool, len(out.varNodes))
for _, n := range out.varNodes {
if name := varName(n.Name); name != "" {
out.vars[name] = true
}
}
return &out
}
// readVarNodes returns variables whose value is read in this subtree. The
// parser visitor also visits assignment targets; those are writes, not
// inputs to the assignment expression, and must not borrow taint from an
// earlier assignment.
func (f *scopeFacts) readVarNodes() []namedNodeSpan {
writes := make([]nodeSpan, 0, len(f.writes))
for _, n := range f.writes {
if span, ok := spanOf(n); ok {
writes = append(writes, span)
}
}
sort.Slice(writes, func(i, j int) bool { return writes[i].start < writes[j].start })
vars := make([]namedNodeSpan, 0, len(f.varNodes)+len(f.propNodes))
reads := make([]namedNodeSpan, 0, len(f.varNodes)+len(f.propNodes))
for _, n := range f.varNodes {
name := varName(n.Name)
if name == "" {
continue
}
if span, ok := spanOf(n); ok {
vars = append(vars, namedNodeSpan{nodeSpan: span, name: name, node: n})
} else {
// Parsed nodes normally always have positions. If a future parser
// omits one, retain the conservative taint dependency.
reads = append(reads, namedNodeSpan{name: name, node: n})
}
}
// Property-fetch reads are keyed to the specific property chain (see
// assignedTargetKey), not the bare base variable, so a write to one
// property cannot leak into a read of a different one. A property whose
// name is not statically known is skipped here entirely rather than
// keyed to "": the embedded base variable visited above already carries
// the conservative base-variable dependency for that case, via
// assignedTargetKey's own fallback at write time.
for _, n := range f.propNodes {
key := assignedTargetKey(n)
if key == "" {
continue
}
if span, ok := spanOf(n); ok {
vars = append(vars, namedNodeSpan{nodeSpan: span, name: key, node: n})
} else {
reads = append(reads, namedNodeSpan{name: key, node: n})
}
}
sort.Slice(vars, func(i, j int) bool { return vars[i].start < vars[j].start })
writeIndex := 0
maxWriteEnd := -1
for _, variable := range vars {
for writeIndex < len(writes) && writes[writeIndex].start <= variable.start {
if writes[writeIndex].end > maxWriteEnd {
maxWriteEnd = writes[writeIndex].end
}
writeIndex++
}
if maxWriteEnd >= variable.end {
continue
}
reads = append(reads, variable)
}
return reads
}
func spanOf(n ast.Vertex) (nodeSpan, bool) {
if n == nil || n.GetPosition() == nil {
return nodeSpan{}, false
}
pos := n.GetPosition()
if pos.StartPos < 0 || pos.EndPos < pos.StartPos {
return nodeSpan{}, false
}
return nodeSpan{start: pos.StartPos, end: pos.EndPos}, true
}
// calleeName renders a call target as a lowercase name. PHP function names
// are case-insensitive; variable and dynamic targets yield "".
//
// Multi-part names, including *ast.NameRelative (a "namespace\name" call),
// are never collapsed onto their last segment: PHP resolves an unqualified
// call in the current namespace before falling back to the global one, so
// treating "Foo\curl_exec" as global curl_exec would invent false
// positives. *ast.NameRelative is therefore left unhandled and yields "";
// missing an explicitly namespaced wrapper is the accepted trade.
func calleeName(n ast.Vertex) string {
switch v := n.(type) {
case *ast.Name:
return joinNameParts(v.Parts)
case *ast.NameFullyQualified:
return joinNameParts(v.Parts)
case *ast.Identifier:
return strings.ToLower(string(v.Value))
}
return ""
}
func joinNameParts(parts []ast.Vertex) string {
out := make([]string, 0, len(parts))
for _, p := range parts {
if np, ok := p.(*ast.NamePart); ok {
out = append(out, string(np.Value))
}
}
return strings.ToLower(strings.Join(out, "\\"))
}
// varName renders a variable's name. The parser's T_VARIABLE token includes
// the leading '$' in Identifier.Value, so it is trimmed here to yield the
// bare name. Variable variables yield "" because their identity is not
// statically known.
func varName(n ast.Vertex) string {
if id, ok := n.(*ast.Identifier); ok {
return strings.TrimPrefix(string(id.Value), "$")
}
return ""
}
package phptaint
// grade is the analyzer's internal lattice element. The public Confidence is
// its first component and keeps its meaning; basis and offset only explain
// which proof produced that confidence, so a stronger flow must carry its own
// explanation rather than an older, weaker one.
type grade struct {
conf Confidence
basis Basis
offset int
}
func directGrade(c Confidence, b Basis) grade {
return grade{conf: c, basis: b, offset: -1}
}
// rankedBases is the spec's fixed tie priority, strongest first. It is also
// the index order of gradeSet, so the two can never disagree.
var rankedBases = [...]Basis{
BasisAlwaysRemote,
BasisLiteral,
BasisDecoded,
BasisRequest,
BasisCallArgument,
BasisUnresolved,
}
// basisRank is the position of b in rankedBases; -1 means undefined.
func basisRank(b Basis) int {
for i, r := range rankedBases {
if r == b {
return i
}
}
return -1
}
// stronger orders by confidence, then basis priority, then the lowest
// resolution offset, with -1 (direct) lowest. The order is total over valid
// grades so every join is deterministic regardless of worklist order.
func (g grade) stronger(o grade) bool {
if g.conf != o.conf {
return g.conf > o.conf
}
if rg, ro := basisRank(g.basis), basisRank(o.basis); rg != ro {
return rg < ro
}
return g.offset < o.offset
}
// directGradeOf rebuilds the ordering key from a finished Result.
func directGradeOf(r Result) grade {
return grade{conf: r.Confidence, basis: r.Basis, offset: r.ResolutionOffset}
}
// confidenceLevels is the number of Confidence values, the second index of
// a gradeSet.
const confidenceLevels = int(ConfidenceCertain) + 1
// gradeSet is the value the fixpoints store for a variable, an assignment,
// an origin and a summary: for each (basis, confidence) pair, the lowest
// resolution offset any proof with that key reached. A single grade cannot
// be stored there, because the decoder upgrade is not monotone under
// stronger: High always-remote is weaker than Certain literal, yet after
// both are decoded Certain always-remote wins. Keying by basis alone is not
// enough either: (High, 3) is weaker than (Certain, 9) on one basis, yet
// after the upgrade offset 3 wins. With confidence in the key, the upgrade
// just moves each entry to its basis's Certain key and joins it there, which
// distributes over the join. The join keeps the lowest offset per key, there
// are at most len(rankedBases) x confidenceLevels keys, so every fixpoint
// over gradeSet still terminates and has one answer. The single strongest
// grade is chosen only when a Result is emitted.
type gradeSet struct {
entries [len(rankedBases)][confidenceLevels]gradeEntry
}
// gradeEntry is one key's lowest resolution offset, stored biased by
// entryBias so the zero value means absent and the zero gradeSet is the
// empty set. A resolution offset is -1 or a position inside the source,
// which MaxSourceBytes bounds far below int32. The solver keeps one gradeSet
// per variable, assignment output and summary, so the entry stays four
// bytes rather than a padded (bool, int) pair.
type gradeEntry int32
const entryBias = 2
// entryAt encodes offset. An offset outside [-1, MaxSourceBytes] is an
// analyzer defect: it panics, which the package boundary recovers into a
// visible coverage gap, rather than storing a value that decodes to a
// different proof.
func entryAt(offset int) gradeEntry {
if offset < -1 || offset > MaxSourceBytes {
panic("phptaint: resolution offset out of range")
}
return gradeEntry(offset + entryBias)
}
func (e gradeEntry) present() bool { return e != 0 }
func (e gradeEntry) offset() int { return int(e) - entryBias }
// joinEntry keeps the lower offset for one key and reports whether it grew.
// The bias preserves order, so present entries compare directly.
func joinEntry(cur *gradeEntry, e gradeEntry) bool {
if !e.present() || (cur.present() && *cur <= e) {
return false
}
*cur = e
return true
}
// setOf is the set holding one proof. An undefined basis or confidence is
// an analyzer defect: dropping it would silently turn a flow into no flow,
// so it panics instead, and the package boundary (recovered) reports the
// file as a StatusPanic coverage gap.
func setOf(g grade) gradeSet {
r := basisRank(g.basis)
if r < 0 {
panic("phptaint: undefined basis")
}
if int(g.conf) >= confidenceLevels {
panic("phptaint: undefined confidence")
}
var s gradeSet
s.entries[r][g.conf] = entryAt(g.offset)
return s
}
func (s gradeSet) isEmpty() bool {
return s == gradeSet{}
}
// add joins o into s per key and reports whether s grew: a key appeared or
// an existing key's offset dropped.
func (s *gradeSet) add(o gradeSet) bool {
grew := false
for b := range o.entries {
for c := range o.entries[b] {
if joinEntry(&s.entries[b][c], o.entries[b][c]) {
grew = true
}
}
}
return grew
}
// decoded applies the decoder upgrade to every proof: each entry moves to
// its basis's Certain key, and each basis still explains its own
// acquisition.
func (s gradeSet) decoded() gradeSet {
var out gradeSet
for b := range s.entries {
for c := range s.entries[b] {
joinEntry(&out.entries[b][ConfidenceCertain], s.entries[b][c])
}
}
return out
}
// strongest selects the proof a Result reports, by the stronger order. The
// empty set yields the zero grade; callers emit only non-empty sets.
func (s gradeSet) strongest() grade {
var best grade
found := false
for b := range s.entries {
for c, e := range s.entries[b] {
if !e.present() {
continue
}
g := grade{conf: Confidence(c), basis: rankedBases[b], offset: e.offset()}
if !found || g.stronger(best) {
best = g
}
found = true
}
}
return best
}
package phptaint
import (
"fmt"
"github.com/VKCOM/php-parser/pkg/ast"
"github.com/VKCOM/php-parser/pkg/conf"
"github.com/VKCOM/php-parser/pkg/errors"
"github.com/VKCOM/php-parser/pkg/parser"
"github.com/VKCOM/php-parser/pkg/version"
)
// parserVersion is the highest grammar this parser implements. Hosts run
// newer PHP; constructs beyond this ceiling recover into a partial tree and
// are accounted as coverage gaps rather than analysed.
const parserVersion = "8.1"
// parseSource returns the tree plus a status. StatusAnalyzed here means only
// that parsing completed cleanly; the data-flow pass decides the final status.
func parseSource(src []byte) (ast.Vertex, Status, string) {
ver, err := version.New(parserVersion)
if err != nil {
return nil, StatusParseError, "parser version unavailable"
}
var syntaxErrs int
var firstLine int
root, err := parser.Parse(src, conf.Config{
Version: ver,
ErrorHandlerFunc: func(e *errors.Error) {
syntaxErrs++
if firstLine == 0 && e != nil && e.Pos != nil {
firstLine = e.Pos.StartLine
}
},
})
switch {
case err != nil:
// parser.Parse currently returns only configuration errors here. Keep
// this diagnostic generic so a future parser error cannot expose input.
return nil, StatusParseError, "parser failed"
case root == nil:
return nil, StatusParseError, "parser produced no tree"
case syntaxErrs > 0:
reason := fmt.Sprintf("recovered from %d syntax error(s)", syntaxErrs)
if firstLine > 0 {
reason += fmt.Sprintf(" starting at line %d", firstLine)
}
return root, StatusPartialParse, sanitizeReason(reason)
}
return root, StatusAnalyzed, ""
}
// Package phptaint reports remotely-fetched content that reaches a PHP
// code-execution sink.
//
// It exists because binding an executed variable to a fetched one requires a
// backreference, which neither YARA-X nor Go's RE2 provides. See
// docs/superpowers/specs/2026-08-18-php-remote-source-taint-analyzer-design.md.
//
// The package is pure: bytes in, report out. It touches no filesystem,
// config, store, or process global that any caller can observe or that
// carries state across calls. The exceptions are declTreeBuilds and
// summaryBodyEvals, package-level counters incremented during analysis purely
// so same-package tests can assert structural invariants. Nothing in this
// package branches on either, so neither is observable to a caller; see their
// own doc comments.
package phptaint
import (
"context"
"sort"
"github.com/VKCOM/php-parser/pkg/ast"
)
// Status is the outcome of an analysis attempt. Callers must not infer a
// clean file from an empty result slice. StatusAnalyzed and StatusNotCandidate
// are the two completed outcomes; every other status is a coverage gap and must
// be accounted for as such rather than counted as a clean file.
type Status uint8
const (
// StatusNotCandidate means the pre-filter alone proved the content
// cannot contain a flow this analyzer reports.
StatusNotCandidate Status = iota
// StatusAnalyzed means the content parsed cleanly and the data-flow pass
// ran to completion.
StatusAnalyzed
StatusOversize
// StatusParseError means no usable tree was produced.
StatusParseError
// StatusPartialParse means the parser recovered from syntax errors and
// returned an incomplete tree. Analysing it would under-report silently,
// so it is a coverage gap rather than a result.
StatusPartialParse
StatusResourceLimit
StatusCanceled
StatusPanic
// StatusTimeout means analysis was still running when its deadline
// expired and the process running it was killed. It is distinct from
// StatusCanceled, which is the caller withdrawing: a timeout means the
// analyzer did not stop on its own, which for this package is the
// expected outcome of a parser loop rather than a surprise.
StatusTimeout
// StatusWorkerFailure means analysis could not be carried out because the
// isolated process failed -- it exited unexpectedly, could not be
// started, or its reply was unusable. The content was never examined.
StatusWorkerFailure
)
// String names the status for metrics labels and operator-facing text.
func (s Status) String() string {
switch s {
case StatusNotCandidate:
return "not_candidate"
case StatusAnalyzed:
return "analyzed"
case StatusOversize:
return "oversize"
case StatusParseError:
return "parse_error"
case StatusPartialParse:
return "partial_parse"
case StatusResourceLimit:
return "resource_limit"
case StatusCanceled:
return "canceled"
case StatusPanic:
return "panic"
case StatusTimeout:
return "timeout"
case StatusWorkerFailure:
return "worker_failure"
}
return "unknown"
}
// Confidence grades how firmly the source was shown to be remote.
type Confidence uint8
const (
// ConfidenceLow means the source function also reads local files and its
// argument could not be proven either way.
ConfidenceLow Confidence = iota
// ConfidenceHigh means the argument carries a remote scheme.
ConfidenceHigh
// ConfidenceCertain means a remote source additionally passed through a
// decoder before reaching the sink.
ConfidenceCertain
)
// String names the confidence for operator-facing text.
func (c Confidence) String() string {
switch c {
case ConfidenceLow:
return "low"
case ConfidenceHigh:
return "high"
case ConfidenceCertain:
return "certain"
}
return "unknown"
}
// Basis names how a result's source was shown to be remote, or that it was
// not. It is explanation for reviewers and for the reporting decision in
// internal/checks; nothing inside this package detects on it.
type Basis string
const (
// BasisAlwaysRemote: the acquiring call can only read over the network.
BasisAlwaysRemote Basis = "always-remote"
// BasisLiteral: the path argument's exact text carries a remote scheme.
BasisLiteral Basis = "literal"
// BasisDecoded: the scheme appears only after PHP escape or builtin decoding.
BasisDecoded Basis = "decoded"
// BasisRequest: a requester can supply the beginning of the path.
BasisRequest Basis = "request"
// BasisCallArgument: a same-file call site supplied the remote argument.
BasisCallArgument Basis = "call-argument"
// BasisUnresolved: the analyzer followed every modeled construct and
// could not decide the argument's locality.
BasisUnresolved Basis = "unresolved"
)
// Valid reports whether b is one of the defined bases.
func (b Basis) Valid() bool {
return basisRank(b) >= 0
}
// MaxSourceBytes bounds the complete source an analysis accepts. Callers
// should read one byte past it to tell an exact-limit file from a truncated
// prefix.
const MaxSourceBytes = 2 << 20
// MaxReasonBytes bounds Report.Reason. Parser diagnostics can contain
// attacker-controlled text, so reports retain only sanitized, bounded context.
const MaxReasonBytes = 256
// Analyzer budgets. These bound work regardless of the parser's own
// resilience: a Go stack overflow is fatal and cannot be recovered, so
// analyzer recursion is capped explicitly.
const (
// maxCollectedNodes bounds nodes of interest recorded per scope. Total
// traversal work is separately bounded by MaxSourceBytes.
maxCollectedNodes = 800_000
maxSummarizedFuncs = 2_000
maxAnalysisDepth = 256
// MaxEvidenceResults bounds returned evidence paths; TotalResults still
// reports the full count. It is exported so the worker protocol can reject
// a clean reply that silently omits evidence.
MaxEvidenceResults = 8
// maxDeclarations bounds how many functions/methods/classes one file's
// declaration tree is built from. declarationTree is linearithmic in
// this count (not quadratic -- see its doc comment), so this exists as
// defence in depth against a future regression rather than as the
// primary complexity fix, and is set far above any real single PHP
// file's declaration count.
maxDeclarations = 20_000
)
// Result is one remote-source-to-sink flow.
type Result struct {
// Source is the acquiring call, such as "curl_exec".
Source string
// Identifiers lists every distinct variable and call name that appears
// in the sink's own expression, sorted alphabetically. This is context
// for a reviewer reading the report, not a laundering path: it includes
// names that never carried the tainted value, and carries no ordering
// information about how the value actually moved.
Identifiers []string
// Sink names the executing construct, such as "eval".
Sink string
// Confidence grades how firmly the source was shown to be remote.
Confidence Confidence
// Basis explains how the source's locality was established. It never
// changes Confidence; see the Basis type.
Basis Basis
// ResolutionOffset is the zero-based byte offset of the outermost call
// site where a parameter relation was resolved, or -1 for a direct result.
ResolutionOffset int
}
// Report is the outcome of analysing one source file.
type Report struct {
Status Status
// Results is non-empty only when Status is StatusAnalyzed.
Results []Result
// TotalResults counts every distinct source-sink endpoint pair before
// evidence truncation. It is non-zero only when Status is StatusAnalyzed.
TotalResults int
// Reason carries a stable status label plus bounded sanitized context.
// It never contains source excerpts.
Reason string
// PrecisionLoss names constructs that defeat variable identity, such as
// variable variables or extract(). Recorded, never silently ignored.
PrecisionLoss []string
// EvidenceTruncated reports that displayed evidence was shortened.
EvidenceTruncated bool
}
// CoverageGap constructs an incomplete report through the same boundary used
// by Analyze. Supervisors use it for failures that happen outside the analyzer
// so error text cannot bypass the report's sanitizing and size bounds.
func CoverageGap(status Status, reason string) Report {
return finalizeReport(Report{Status: status, Reason: reason})
}
// recovered recovers a panic at the package boundary so a parser or analyzer
// defect degrades to a coverage gap instead of taking down the caller.
func recovered(fn func() Report) (report Report) {
defer func() {
if r := recover(); r != nil {
// A panic value may contain parser tokens or source text. The status
// is actionable without reflecting that attacker-controlled value.
report = Report{Status: StatusPanic, Reason: "recovered panic during analysis"}
}
report = finalizeReport(report)
}()
return fn()
}
// finalizeReport enforces the package boundary invariants in one place. Only a
// completed analysis may carry findings; every evidence segment leaving the
// package is printable and bounded even if an internal producer misses a cap.
func finalizeReport(report Report) Report {
if report.Status == StatusAnalyzed {
report.Reason = ""
for i := range report.Results {
var cutSource, cutSink, cutIdentifiers bool
report.Results[i].Source, cutSource = sanitize(report.Results[i].Source, maxSegmentBytes)
report.Results[i].Sink, cutSink = sanitize(report.Results[i].Sink, maxSegmentBytes)
report.Results[i].Identifiers, cutIdentifiers = truncateChain(report.Results[i].Identifiers)
report.EvidenceTruncated = report.EvidenceTruncated || cutSource || cutSink || cutIdentifiers
}
for i := range report.PrecisionLoss {
report.PrecisionLoss[i] = sanitizeSegment(report.PrecisionLoss[i])
}
return report
}
report.Results = nil
report.TotalResults = 0
report.PrecisionLoss = nil
report.EvidenceTruncated = false
if report.Status == StatusNotCandidate {
report.Reason = ""
return report
}
detail := sanitizeReason(report.Reason)
report.Reason = report.Status.String()
if detail != "" {
report.Reason += ": " + detail
}
report.Reason = sanitizeReason(report.Reason)
return report
}
// Analyze owns the size check, pre-filter, parse, and data-flow pass. It
// recovers a panic at the package boundary via the recovered helper.
func Analyze(ctx context.Context, src []byte) Report {
return recovered(func() Report { return analyze(ctx, src) })
}
func analyze(ctx context.Context, src []byte) Report {
if len(src) > MaxSourceBytes {
return Report{Status: StatusOversize, Reason: "source exceeds maximum analyzed size"}
}
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
if !IsCandidate(src) {
return Report{Status: StatusNotCandidate}
}
root, status, reason := parseSource(src)
if status != StatusAnalyzed {
return Report{Status: status, Reason: reason}
}
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
// all: whole-file inventory, used for the declaration list and the
// budget. top: the file's top-level statements only, collected below via
// collectTopLevel. These MUST stay separate. A single flat taint map over
// the whole file lets a function-local variable taint an unrelated
// top-level variable that merely shares its name, which fires on clean
// code -- names like $data and $content are ubiquitous in real PHP.
// Function and method bodies get their own state below.
all := collectScope(root)
if all.budgetExceeded {
return Report{Status: StatusResourceLimit, Reason: "collection budget exceeded"}
}
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
// tree indexes every declaration in the file once (see declarationTree
// for why this must happen exactly once, not once per declaration), so
// each scope below can derive its own exclusion index by lookup.
// functionSummaries calls f.declarationTree() again on this same all;
// that hits scopeFacts' own cache rather than rebuilding. Check the
// defence-in-depth cap before any further work, including
// functionSummaries, so a file that trips it does no summarization work
// either. This is a coverage gap, not a silent skip: StatusResourceLimit
// means the file was not examined, matching the same contract already
// used for a collection-budget overrun above.
tree := all.declarationTree()
if tree.count > maxDeclarations {
return Report{Status: StatusResourceLimit, Reason: "too many declarations to analyze"}
}
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
callIndex := newResolvedCallIndex(all)
summaries, loss, err := functionSummaries(ctx, all)
if err != nil {
if err == errSummaryLimit {
return Report{Status: StatusResourceLimit, Reason: err.Error()}
}
return Report{Status: StatusCanceled, Reason: err.Error()}
}
// Each scope below excludes every declaration nested inside it except,
// when the scope IS a declaration's own body, that declaration's own
// span -- otherwise a function would exclude its own statements from
// itself. exclusionFor derives this from tree by lookup rather than by
// rescanning the file's declarations for every scope.
topExclude := tree.exclusionFor(nil)
top := callIndex.apply(collectTopLevel(root, &topExclude))
// factsByScope keeps the facts each loop below already builds, keyed by the
// declaration they belong to (nil for the top level). Nothing here
// re-collects: this only retains what would otherwise be discarded, so a
// closure can later ask what its enclosing scope held.
factsByScope := map[ast.Vertex]*scopeFacts{nil: top}
flows, err := findFlows(ctx, top, summaries, callIndex, &topExclude)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
// droppedTaint tracks whether ANY scope assigns tainted content to a
// target this package cannot key (see hasUnresolvableTaintedTarget):
// method-call-chain, static-property, list()-destructuring, and
// variable-variable targets all silently drop the flow rather than
// risk a false positive by guessing. Checked per scope (top level, each
// function, each method) using the SAME scope-isolated facts already
// collected here for findFlows, so "tainted" means tainted within that
// specific scope's own taint state -- never a same-named variable
// leaking taint in from an unrelated scope. Direct call origins come from
// all so namespace-level `use function` aliases stay resolved inside
// separately collected declaration bodies; source positions still limit
// them to the exact RHS being checked.
droppedTaint, err := hasUnresolvableTaintedTarget(ctx, top, summaries)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
// Sinks inside function bodies count too, each against its own state.
// collectOwnStmts (not collectAll) skips any function/class/method/
// closure/arrow-function nested inside this body, for the same reason
// collectTopLevel skips declarations at the file level: a nested
// declaration's locals must not taint an identically-named local in the
// enclosing body.
for _, fn := range all.funcs {
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
fnExclude := tree.exclusionFor(fn)
body := callIndex.apply(collectOwnStmts(fn.Stmts, &fnExclude))
factsByScope[fn] = body
bodyFlows, err := findFlows(ctx, body, summaries, callIndex, &fnExclude)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
flows = append(flows, bodyFlows...)
if !droppedTaint {
droppedTaint, err = hasUnresolvableTaintedTarget(ctx, body, summaries)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
}
}
for _, m := range all.methods {
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
// StmtClassMethod carries ONE Stmt vertex, not a Stmts slice.
mExclude := tree.exclusionFor(m)
body := callIndex.apply(collectOwnStmts(methodStmts(m.Stmt), &mExclude))
factsByScope[m] = body
bodyFlows, err := findFlows(ctx, body, summaries, callIndex, &mExclude)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
flows = append(flows, bodyFlows...)
if !droppedTaint {
droppedTaint, err = hasUnresolvableTaintedTarget(ctx, body, summaries)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
}
}
// Closures are the most common nested scope in real PHP (WordPress hooks,
// shutdown/callback registration, ...), so they get the exact same
// treatment as a named function's body: excluded from the enclosing
// scope above via declarationTree, analysed here in their own. Recording
// the span without this loop (or the reverse) would either let a closure
// local borrow an unrelated outer variable's taint or stop examining the
// closure's own sinks entirely -- see the closures field doc in facts.go.
for _, cl := range all.closures {
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
clExclude := tree.exclusionFor(cl)
body := callIndex.apply(collectOwnStmts(cl.Stmts, &clExclude))
factsByScope[cl] = body
bodyFlows, err := findFlows(ctx, body, summaries, callIndex, &clExclude)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
flows = append(flows, bodyFlows...)
if !droppedTaint {
droppedTaint, err = hasUnresolvableTaintedTarget(ctx, body, summaries)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
}
}
// An arrow function's body is a single implicit-return expression rather
// than a statement list (see arrowFunctionBody), but it is the same kind
// of nested scope as a closure and needs the same two-sided treatment.
for _, af := range all.arrowFuncs {
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
afExclude := tree.exclusionFor(af)
body := callIndex.apply(collectOwnStmts(arrowFunctionBody(af), &afExclude))
factsByScope[af] = body
bodyFlows, err := findFlows(ctx, body, summaries, callIndex, &afExclude)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
flows = append(flows, bodyFlows...)
if !droppedTaint {
droppedTaint, err = hasUnresolvableTaintedTarget(ctx, body, summaries)
if err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
}
}
if err := ctx.Err(); err != nil {
return Report{Status: StatusCanceled, Reason: err.Error()}
}
capturedTaint, captureErr := hasDroppedCapture(ctx, all, tree, factsByScope, summaries)
if captureErr != nil {
return Report{Status: StatusCanceled, Reason: captureErr.Error()}
}
if droppedTaint {
loss = append(loss, "unresolvable-assign-target")
}
if capturedTaint {
loss = append(loss, "closure-capture")
}
if droppedTaint || capturedTaint {
// loss is already sorted (functionSummaries sorts it); re-sort once
// after appending rather than inserting in place, so this stays a
// single obviously-correct call regardless of where the appended
// markers fall alphabetically among the other markers.
sort.Strings(loss)
}
flows = dedupeAndSort(flows)
total := len(flows)
truncated := total > MaxEvidenceResults
flows = retainStrongestEvidence(flows)
results := make([]Result, len(flows))
for i, flow := range flows {
results[i] = flow.Result
truncated = truncated || flow.evidenceTruncated
}
return Report{
Status: StatusAnalyzed,
Results: results,
TotalResults: total,
PrecisionLoss: loss,
EvidenceTruncated: truncated,
}
}
package phptaint
import "bytes"
// sinkKeywords and sourceKeywords drive admission only. They are language
// constructs and library function names, never variable or file names.
var (
sinkKeywords = [][]byte{
[]byte("eval"), []byte("include"), []byte("require"),
[]byte("assert"), []byte("create_function"),
}
sourceKeywords = [][]byte{
[]byte("curl_exec"), []byte("curl_multi_getcontent"),
[]byte("file_get_contents"), []byte("fopen"), []byte("fread"),
[]byte("fgets"), []byte("stream_get_contents"), []byte("readfile"),
[]byte("wp_remote_retrieve_body"), []byte("wp_remote_get"),
[]byte("fsockopen"),
}
phpOpenTags = [][]byte{[]byte("<?php"), []byte("<?="), []byte("<?")}
)
// foldChunkBytes bounds the lowercased working copy containsAnyFold folds
// src into. Matching case-insensitively requires comparing against a
// lowercased view of the bytes somewhere; the previous implementation
// lowercased the entire file up front, which meant a full-size allocation
// and copy for every file scanned, paid even on the reject path -- the hot
// path for a daemon walking millions of files, most of which are not PHP at
// all. Folding a small fixed-size window at a time bounds that cost to a
// constant instead of MaxSourceBytes. 64KiB is also small enough that the
// compiler keeps the working buffer on the stack rather than the heap, so
// this costs zero allocations, not just a smaller one.
const foldChunkBytes = 64 << 10
// maxNeedleBytes is the longest literal in sinkKeywords, sourceKeywords, or
// phpOpenTags, computed rather than hand-maintained so a future longer
// keyword cannot silently undersize the window overlap below and reopen a
// false negative. Consecutive folded windows overlap by this many bytes
// minus one so a match straddling a window boundary is never missed.
var maxNeedleBytes = longestNeedle(sinkKeywords, sourceKeywords, phpOpenTags)
func longestNeedle(groups ...[][]byte) int {
max := 0
for _, group := range groups {
for _, n := range group {
if len(n) > max {
max = len(n)
}
}
}
return max
}
// MayBePHPSource reports whether content could be PHP at all, judged only by
// the presence of an open tag. It exists for callers that must decide
// something about a file they cannot analyze in full -- an oversize file, say
// -- and would otherwise have to guess.
//
// It is deliberately weaker than IsCandidate: no sink or source keyword is
// required, because a caller holding only a prefix cannot conclude anything
// from their absence. Judging by content rather than by name or extension is
// the point; a scanner that decided what to examine from a path would be
// telling an attacker where to hide.
func MayBePHPSource(prefix []byte) bool {
// PHP emits arbitrary bytes before its opening tag as inline HTML.
// A NUL or binary header cannot rule out executable PHP later on.
return containsAnyFold(prefix, phpOpenTags)
}
// IsCandidate is the cheap byte scan Analyze runs before parsing; content
// it rejects is StatusNotCandidate. No parser runs, so a caller that isolates
// Analyze in another process can apply it in its own process first and skip
// the round trip. It is intentionally over-inclusive; the AST pass decides
// whether a real flow exists. PHP
// function names and language constructs are case-insensitive (EVAL, Eval
// and eval all execute the same construct), so admission matches
// case-insensitively too, on pain of a false negative admitting less than
// the AST rules would. The open-tag check runs first so a file that never
// looks like PHP never pays for the sink/source keyword scans either.
func IsCandidate(src []byte) bool {
if !MayBePHPSource(src) {
return false
}
return containsAnyFold(src, sinkKeywords) && containsAnyFold(src, sourceKeywords)
}
// containsAnyFold reports whether any needle occurs in hay under ASCII case
// folding. It folds hay one bounded window at a time into a stack-local
// buffer -- never allocating a lowercased copy of the whole input -- and
// runs the standard library's optimized bytes.Contains against each folded
// window, so search speed stays close to a plain bytes.Contains scan.
// Windows overlap by maxNeedleBytes-1 bytes so a needle split across a
// window boundary still lands whole inside the next window.
func containsAnyFold(hay []byte, needles [][]byte) bool {
if len(hay) == 0 {
return false
}
overlap := maxNeedleBytes - 1
step := foldChunkBytes - overlap
var buf [foldChunkBytes]byte
for start := 0; start < len(hay); start += step {
end := start + foldChunkBytes
if end > len(hay) {
end = len(hay)
}
window := hay[start:end]
folded := buf[:len(window)]
for i, b := range window {
folded[i] = lowerASCII(b)
}
for _, n := range needles {
if bytes.Contains(folded, n) {
return true
}
}
if end == len(hay) {
break
}
}
return false
}
func lowerASCII(b byte) byte {
if 'A' <= b && b <= 'Z' {
return b + ('a' - 'A')
}
return b
}
package phptaint
import (
"strings"
"github.com/VKCOM/php-parser/pkg/ast"
"github.com/pidginhost/csm/internal/cms"
)
// buildLocalPathConstants collects every bootstrap path constant declared by
// the given CMS descriptors into a lookup set.
func buildLocalPathConstants(descriptors []cms.Descriptor) map[string]bool {
out := make(map[string]bool)
for _, d := range descriptors {
for _, c := range d.PathConstants {
out[c] = true
}
}
return out
}
// alwaysRemote functions can only acquire content over the network.
var alwaysRemote = map[string]bool{
"curl_exec": true,
"curl_multi_getcontent": true,
"wp_remote_get": true,
"wp_remote_retrieve_body": true,
"fsockopen": true,
}
// dualUse functions return either local or remote content, so their argument
// decides. fread, fgets and stream_get_contents are deliberately excluded:
// they take a stream resource, never a path, so their argument carries no
// locality signal. readfile is excluded because it returns a byte count, not
// the content it writes, so no fetched bytes can flow from its return value.
// For stream readers, the acquiring fopen or fsockopen call is already the
// source and its tainted handle carries through the expression.
var dualUse = map[string]bool{
"file_get_contents": true,
"fopen": true,
}
// remoteSchemes mark an argument as remote. Only php://input is intrinsically
// request-controlled; php://memory, php://temp and local php://filter resources
// are not remote sources. A filter around a remote resource still contains its
// nested remote scheme and is classified accordingly.
var remoteSchemes = []string{"http://", "https://", "ftp://", "ftps://", "php://input", "data://"}
// These are the specific PHP and CMS constructs whose result is known to be a
// local path. Arbitrary constants and calls remain undecidable: their runtime
// value can be a remote URL.
//
// Every supported CMS defines filesystem path constants during bootstrap, and
// all of them compile templates or caches by writing generated PHP under one
// of those paths and including it afterwards. Without the constant, that read
// is undecidable, the generated file looks remotely acquired, and the include
// reports as remote execution on a stock installation. The constants come
// from the supported CMS table so a CMS declared there cannot be missing here.
var (
localPathConstants = buildLocalPathConstants(cms.All())
localPathResults = map[string]bool{
"get_template_directory": true,
"realpath": true,
"sys_get_temp_dir": true,
}
pathTransforms = map[string]bool{
"dirname": true,
"plugin_dir_path": true,
}
)
type locality uint8
const (
localityUnknown locality = iota
localityLocal
localityRemote
)
// sourceGrade reports whether a call acquires remote content, how firmly
// that was shown, and on what basis.
func sourceGrade(call *ast.ExprFunctionCall) (grade, bool) {
name := calleeName(call.Function)
if alwaysRemote[name] {
return directGrade(ConfidenceHigh, BasisAlwaysRemote), true
}
if !dualUse[name] {
return grade{}, false
}
unresolved := directGrade(ConfidenceLow, BasisUnresolved)
if len(call.Args) == 0 {
return unresolved, true
}
arg, ok := call.Args[0].(*ast.Argument)
if !ok {
return unresolved, true
}
switch argLocality(arg.Expr) {
case localityLocal:
return grade{}, false
case localityRemote:
return directGrade(ConfidenceHigh, BasisLiteral), true
}
return unresolved, true
}
// argLocality classifies an argument by shape. Literal text and
// concatenations of literals are decidable; anything else is unknown.
func argLocality(arg ast.Vertex) locality {
text, decidable := staticText(arg)
lower := strings.ToLower(text)
for _, scheme := range remoteSchemes {
if strings.Contains(lower, scheme) {
return localityRemote
}
}
if !decidable {
return localityUnknown
}
if strings.HasPrefix(text, "/") || strings.HasPrefix(text, "./") ||
strings.HasPrefix(text, "../") {
return localityLocal
}
if text != "" {
return localityLocal
}
return localityUnknown
}
// staticText folds an expression to text when every part is statically known.
// Known fragments survive an undecidable concatenation so a literal scheme can
// still prove remoteness; arbitrary constants and calls stay undecidable.
//
// Recursion is depth-bounded because this is the one place the package
// recurses over attacker-controlled structure, and a Go stack overflow is
// fatal: recover() cannot catch it. Exceeding the bound yields "undecidable",
// which degrades to a reduced-confidence source rather than a wrong answer.
func staticText(n ast.Vertex) (string, bool) { return staticTextAt(n, 0) }
func staticTextAt(n ast.Vertex, depth int) (string, bool) {
if depth >= maxAnalysisDepth {
return "", false
}
switch v := n.(type) {
case *ast.ScalarString:
raw := string(v.Value)
if len(raw) < 2 || raw[0] != raw[len(raw)-1] || (raw[0] != '\'' && raw[0] != '"') {
return raw, false
}
text := raw[1 : len(raw)-1]
// Double-quoted PHP strings interpret hex, octal and other escapes.
// Keeping visible fragments lets a literal scheme still prove remote,
// but an escape makes the complete runtime text undecidable.
if raw[0] == '"' && strings.ContainsRune(text, '\\') {
return text, false
}
return text, true
case *ast.ExprBinaryConcat:
left, okL := staticTextAt(v.Left, depth+1)
right, okR := staticTextAt(v.Right, depth+1)
// Preserve known literal fragments even when another fragment is
// dynamic. Separate undecidable pieces so two fragments cannot invent
// a scheme across an unknown runtime value.
if !okL || !okR {
return left + "\x00" + right, false
}
return left + right, true
case *ast.ExprBrackets:
return staticTextAt(v.Expr, depth+1)
case *ast.ExprConstFetch:
name := calleeName(v.Const)
if localPathConstants[name] {
return name, true
}
return "", false
case *ast.ExprFunctionCall:
name := calleeName(v.Function)
if localPathResults[name] {
return name, true
}
if !pathTransforms[name] || len(v.Args) == 0 {
return "", false
}
arg, ok := v.Args[0].(*ast.Argument)
if !ok {
return "", false
}
return staticTextAt(arg.Expr, depth+1)
case *ast.ScalarMagicConstant:
return strings.ToLower(string(v.Value)), true
}
return "", false
}
package phptaint
import (
"context"
"errors"
"sort"
"sync/atomic"
"github.com/VKCOM/php-parser/pkg/ast"
)
// precisionLossMarkers name calls that defeat static variable identity.
// extract() and compact() move values between named variables and an array
// at runtime; call_user_func(_array) dispatches through a value rather than
// a lexical name. Their presence is recorded so a caller knows coverage was
// reduced, rather than the loss passing silently.
var precisionLossMarkers = map[string]string{
"extract": "extract",
"compact": "compact",
"call_user_func": "dynamic-call",
"call_user_func_array": "dynamic-call",
}
var errSummaryLimit = errors.New("too many function summaries to analyze")
type bodyKind uint8
const (
bodyFunction bodyKind = iota
bodyMethod
)
// summaryKey names one entry in a summaryTables namespace: which of the two
// tables it lives in, plus the bare name within that table.
type summaryKey struct {
kind bodyKind
name string
}
// callSiteSummaryKey reports the summaryKey a call site would resolve
// against, mirroring the node-type switch summaryTables.lookup uses. A call
// shape outside the three kinds facts.go records resolves to nothing, same
// as lookup itself.
func callSiteSummaryKey(call callSite) (summaryKey, bool) {
switch call.node.(type) {
case *ast.ExprFunctionCall:
return summaryKey{kind: bodyFunction, name: call.name}, true
case *ast.ExprMethodCall, *ast.ExprNullsafeMethodCall, *ast.ExprStaticCall:
return summaryKey{kind: bodyMethod, name: call.name}, true
}
return summaryKey{}, false
}
// summaryBodyEvals counts how many times the interprocedural worklist below
// actually re-evaluated a body (taintedLocals plus its return expressions),
// as opposed to a body sitting untouched in the queue. It exists so a
// same-package white-box test can observe that this count grows with the
// call graph's edge count, not with (declaration count) x (declaration
// count), which is the structural invariant a dependency worklist is
// supposed to buy over a round-robin sweep. Nothing in this package branches
// on its value, so this is test-only observation, not shared process state.
var summaryBodyEvals atomic.Int64
// funcBody pairs a summarizable function or method with its collected facts
// and which summary namespace its name belongs to. Facts are collected once
// up front: they do not change across fixpoint rounds, only the summary
// tables consulted while interpreting them do.
type funcBody struct {
name string
kind bodyKind
facts *scopeFacts
// calls is the whole-file resolved-call view, kept so each evaluation can
// re-collect the return expressions with namespace-level aliases intact.
calls resolvedCallIndex
// exclude holds the spans of the declarations nested inside this body,
// the same index its own facts were collected with. A return expression
// is collected WITHOUT it, so its calls still reach nested declarations,
// and then its variables are filtered THROUGH it, so a nested
// declaration's locals are never graded against this body's taint state.
exclude spanIndex
// returnKeys names the summaries reachable from this body's return
// expressions. Only the KEYS are kept, never the collected facts: a
// return expression is collected without excluding nested declarations,
// so for a body whose return lexically contains further declarations
// those facts cover the whole remaining subtree, and retaining one per
// body across the fixpoint is quadratic in nesting depth. Attacker-
// written input reaches a gigabyte of live heap that way for a few
// hundred KB of source. The keys are a few bytes each and are all the
// dependency edges need.
returnKeys []summaryKey
}
// dependencyKeys names every summary a single evaluation of this body can
// read: those reachable from its own statements, consulted while computing
// its local taint state, plus those reachable from each return expression,
// consulted while grading what it returns. The worklist's correctness rests
// on this being a SUPERSET of what an evaluation actually reads. If a body
// can read a summary it has no edge on, a later rise in that summary never
// wakes it, and the fixpoint's answer starts depending on the order bodies
// were queued in -- which is the order they were declared in, which an
// attacker writing the file chooses.
func (b funcBody) dependencyKeys() []summaryKey {
keys := make([]summaryKey, 0, len(b.facts.callSites)+len(b.returnKeys))
for _, call := range b.facts.callSites {
if key, ok := callSiteSummaryKey(call); ok {
keys = append(keys, key)
}
}
return append(keys, b.returnKeys...)
}
// functionSummaries reports which user-defined functions and methods return
// remotely-sourced data, computed to a fixpoint over the call graph, plus any
// precision-loss markers observed while collecting their bodies.
//
// Functions and methods are kept in separate namespaces (summaryTables): PHP
// allows a function and a method, or methods on unrelated classes, to share
// a bare name, and a single shared table would let a tainted method
// anywhere in the file poison every same-named plain function call site. A
// method name declared by more than one class is dropped from the method
// table entirely rather than resolved to any one class's definition - a
// deliberate false-negative trade recorded as "ambiguous-method" precision
// loss, because this analyzer's zero-false-positive bar makes guessing
// wrong more expensive than a visible gap.
//
// This is what lets the analyzer see the motivating shape: a fetch in one
// function and a sink in another, joined only by a call. Without a summary
// for the fetching function, that flow is invisible to a single-scope pass.
func functionSummaries(ctx context.Context, f *scopeFacts) (summaryTables, []string, error) {
bodies, loss, err := summaryBodies(ctx, f)
if err != nil {
return summaryTables{}, nil, err
}
tables, err := solveSummaries(ctx, bodies)
if err != nil {
return summaryTables{}, nil, err
}
names := make([]string, 0, len(loss))
for name := range loss {
names = append(names, name)
}
sort.Strings(names)
if err := ctx.Err(); err != nil {
return summaryTables{}, nil, err
}
return tables, names, nil
}
// summaryBodies collects every summarizable function and method body, with
// the precision loss observed while collecting them. Separated from the
// fixpoint below so the two can be exercised apart: a test can build the
// bodies for a file and then run its own reference fixpoint over them, which
// is what pins the worklist to the answer an exhaustive sweep would give.
func summaryBodies(ctx context.Context, f *scopeFacts) ([]funcBody, map[string]bool, error) {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
loss := map[string]bool{}
// The enclosing facts already cover every declaration body. Record loss
// from them before omitting ambiguous methods, because that filter can
// exclude the only body containing a precision-loss construct.
recordPrecisionLoss(f, loss)
// A single class cannot legally declare the same method name twice, so
// any name appearing more than once in this flat, whole-file list was
// declared by more than one class - ambiguous, without needing to track
// which class owns which method. Ambiguous methods do not consume the
// summary-body budget because they are deliberately omitted from the
// summary table regardless of that budget.
methodCounts := make(map[string]int, len(f.methods))
for _, m := range f.methods {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
methodCounts[calleeName(m.Name)]++
}
summarizable := len(f.funcs)
for _, count := range methodCounts {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
if count > 1 {
loss["ambiguous-method"] = true
continue
}
summarizable++
}
if summarizable > maxSummarizedFuncs {
return nil, nil, errSummaryLimit
}
// tree indexes every declaration in f (see declarationTree), so a
// function or method nested inside another (however deeply, however
// many if/while/switch/try/foreach wrappers it sits behind) is excluded
// from its enclosing body's own facts by lookup rather than by
// rescanning the file's declarations for every body, and so cannot
// pollute that enclosing declaration's own interprocedural summary. In
// production f is the same *scopeFacts analyze already indexed, so this
// hits declarationTree's own cache rather than rebuilding, and is
// already known to be within maxDeclarations here (analyze's cap check
// runs first).
tree := f.declarationTree()
callIndex := newResolvedCallIndex(f)
bodies := make([]funcBody, 0, summarizable)
for _, fn := range f.funcs {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
exclude := tree.exclusionFor(fn)
facts := callIndex.apply(collectOwnStmts(fn.Stmts, &exclude))
returnKeys, err := returnSummaryKeys(ctx, facts, callIndex)
if err != nil {
return nil, nil, err
}
bodies = append(bodies, funcBody{
name: calleeName(fn.Name), kind: bodyFunction, facts: facts,
calls: callIndex, exclude: exclude, returnKeys: returnKeys,
})
}
for _, m := range f.methods {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
name := calleeName(m.Name)
if methodCounts[name] > 1 {
continue
}
// StmtClassMethod carries ONE Stmt vertex (normally a StmtStmtList),
// unlike StmtFunction which carries a Stmts slice.
exclude := tree.exclusionFor(m)
facts := callIndex.apply(collectOwnStmts(methodStmts(m.Stmt), &exclude))
returnKeys, err := returnSummaryKeys(ctx, facts, callIndex)
if err != nil {
return nil, nil, err
}
bodies = append(bodies, funcBody{
name: name, kind: bodyMethod, facts: facts,
calls: callIndex, exclude: exclude, returnKeys: returnKeys,
})
}
// Recheck every included body using its independently collected facts.
// The enclosing collection has one aggregate node budget, so it can stop
// recording inside a declaration that was already discovered; the fresh
// per-body budget can still preserve that declaration's loss markers.
for _, b := range bodies {
if err := ctx.Err(); err != nil {
return nil, nil, err
}
recordPrecisionLoss(b.facts, loss)
}
return bodies, loss, nil
}
// solveSummaries runs the interprocedural fixpoint over prebuilt bodies.
func solveSummaries(ctx context.Context, bodies []funcBody) (summaryTables, error) {
// Summaries only ever move from absent to present, or grow pointwise as
// a gradeSet, so this is a monotone fixpoint over a finite lattice (one
// entry per (basis, confidence) key, each holding the lowest offset drawn
// from this file's call sites) whose result does not depend on the order
// bodies are (re)evaluated in - only on eventually evaluating every body
// whose inputs changed since it was last evaluated. A body's inputs are
// exactly the summaries of the functions and methods its own facts call,
// so a worklist keyed on that call graph re-evaluates a body only when a
// callee it actually references just changed, instead of re-sweeping
// every body on every pass regardless of whether anything it depends on
// moved. Which summaries a body depends on is exactly what
// dependencyKeys reports, and that must stay a superset of what an
// evaluation reads, or the order-independence claimed above is simply
// false: edges drawn from a body's own statements alone miss a callee
// invoked from inside a closure within a return, because a body's facts
// exclude nested declarations while the collection used to grade its
// return expression does not.
//
// This is the same dependency-worklist shape solveAssignments
// already runs intraprocedurally in taint.go, applied one level up: a
// body here plays the role an assignment plays there, and a produced
// summary name plays the role a variable plays there. Termination
// follows from the lattice being finite - at most one entry per body
// name, each raised a bounded number of times - rather than from any
// iteration count, so it needs no cap sized to the input the way a
// round-robin sweep would.
produced := make(map[summaryKey]bool, len(bodies))
for _, b := range bodies {
if b.name == "" {
continue
}
produced[summaryKey{kind: b.kind, name: b.name}] = true
}
// dependents maps a produced key to the bodies whose own facts call it,
// i.e. the bodies to wake when that key's grade rises. Built once
// from each body's already-collected call sites, so this costs one pass
// over the call sites this file already gathered rather than a rescan
// per round.
dependents := make(map[summaryKey][]int)
queue := make([]int, 0, len(bodies))
queued := make([]bool, len(bodies))
for i, b := range bodies {
if err := ctx.Err(); err != nil {
return summaryTables{}, err
}
if b.name == "" {
continue
}
for _, key := range b.dependencyKeys() {
if !produced[key] {
continue
}
dependents[key] = append(dependents[key], i)
}
queue = append(queue, i)
queued[i] = true
}
enqueue := func(indices []int) {
for _, i := range indices {
if !queued[i] {
queued[i] = true
queue = append(queue, i)
}
}
}
tables := summaryTables{funcs: map[string]gradeSet{}, methods: map[string]gradeSet{}}
for head := 0; head < len(queue); head++ {
if err := ctx.Err(); err != nil {
return summaryTables{}, err
}
i := queue[head]
queued[i] = false
b := bodies[i]
best, found, err := evalBodySummary(ctx, b, tables)
if err != nil {
return summaryTables{}, err
}
if !found {
continue
}
target := tables.funcs
if b.kind == bodyMethod {
target = tables.methods
}
if cur, ok := target[b.name]; ok {
if !cur.add(best) {
continue
}
best = cur
}
target[b.name] = best
enqueue(dependents[summaryKey{kind: b.kind, name: b.name}])
}
return tables, nil
}
// evalBodySummary grades what one body returns against the summaries known so
// far: the joined proofs of every tainted return, plus whether anything
// tainted is returned at all. It reads
// only the summaries dependencyKeys reports, which is what lets the worklist
// wake exactly the bodies an update can affect.
func evalBodySummary(ctx context.Context, b funcBody, tables summaryTables) (gradeSet, bool, error) {
summaryBodyEvals.Add(1)
st := taintedLocals(b.facts, tables)
var best gradeSet
found := false
for _, ret := range b.facts.returns {
if err := ctx.Err(); err != nil {
return gradeSet{}, false, err
}
if ret.Expr == nil {
continue
}
retFacts := b.calls.apply(collectScope(ret.Expr)).withoutNestedDeclarationVars(&b.exclude)
c, tainted := exprTaintFacts(retFacts, st, tables)
if !tainted {
continue
}
best.add(c)
found = true
}
return best, found, nil
}
// returnSummaryKeys names every summary reachable from a body's return
// expressions. calls is the whole-file resolved-call view, not an index
// rebuilt from f: f excludes nested declarations, while a return expression
// can read through an invoked closure or arrow function inside one. The
// whole-file view is what preserves namespace-level aliases for those nested
// calls.
//
// Each collection is transient. Only the keys survive, because the collected
// facts of a return expression cover every declaration nested inside it, and
// holding one per body for the life of the fixpoint costs memory quadratic in
// nesting depth on input an attacker writes.
func returnSummaryKeys(
ctx context.Context, f *scopeFacts, calls resolvedCallIndex,
) ([]summaryKey, error) {
var keys []summaryKey
for _, ret := range f.returns {
if err := ctx.Err(); err != nil {
return nil, err
}
if ret.Expr == nil {
continue
}
for _, call := range calls.apply(collectScope(ret.Expr)).callSites {
if key, ok := callSiteSummaryKey(call); ok {
keys = append(keys, key)
}
}
}
return keys, nil
}
// recordPrecisionLoss folds one body's precision-loss facts into the running
// set. Variable-variables and dynamic calls are already flagged during
// collection (facts.go); this adds the calls whose names are known but whose
// effect on variable identity is not visible to the collector, such as
// extract() writing into caller-invisible names.
func recordPrecisionLoss(f *scopeFacts, loss map[string]bool) {
for marker := range f.precisionLoss {
loss[marker] = true
}
for name := range f.calls {
if marker, ok := precisionLossMarkers[name]; ok {
loss[marker] = true
}
}
}
package phptaint
import (
"context"
"sort"
"strings"
"github.com/VKCOM/php-parser/pkg/ast"
)
// decoders preserve taint and raise confidence: fetching plaintext and
// executing it has a small benign population, but fetching, decoding, then
// executing has effectively none.
var decoders = map[string]bool{
"base64_decode": true, "gzinflate": true, "gzuncompress": true,
"gzdecode": true, "str_rot13": true, "hex2bin": true,
"convert_uudecode": true, "unserialize": true, "pack": true,
}
// There is no passthrough allowlist: taint propagation is structural, not
// name-based. exprTaint marks any expression tainted if it references a
// tainted variable anywhere in its subtree, so trim($a), sprintf($a), and an
// unlisted some_helper($a) are all already covered without naming a single
// one of them. A name list here could only ever be narrower than that rule.
// taintState maps a variable name to every proof, one per (basis,
// confidence) key, with which it carries remote content. See gradeSet for
// why it is not a single grade.
type taintState map[string]gradeSet
// summaryTables holds interprocedural summaries in two namespaces so a
// function and a method that happen to share a name can never collide: PHP
// resolves f(), $obj->f(), and Class::f() through distinct call syntax, and
// callSite.node retains which syntax was used, so a lookup can always pick
// the namespace the call site itself selects. A single shared map would let
// a tainted method anywhere in the file poison every same-named plain
// function call site, which is false-positive-only but unacceptable given
// this analyzer's zero-false-positive bar against real WordPress/plugin code.
type summaryTables struct {
funcs map[string]gradeSet
methods map[string]gradeSet
}
// lookup resolves a call site's summary in the namespace its call syntax
// selects. A node shape outside the three call kinds facts.go records
// (should not occur) resolves to nothing rather than guessing a namespace.
func (s summaryTables) lookup(call callSite) (gradeSet, bool) {
switch call.node.(type) {
case *ast.ExprFunctionCall:
c, ok := s.funcs[call.name]
return c, ok
case *ast.ExprMethodCall, *ast.ExprNullsafeMethodCall, *ast.ExprStaticCall:
c, ok := s.methods[call.name]
return c, ok
}
return gradeSet{}, false
}
// raise joins g into name's proofs and reports whether they grew.
func (s taintState) raise(name string, g gradeSet) bool {
if name == "" || g.isEmpty() {
return false
}
cur, ok := s[name]
if !cur.add(g) && ok {
return false
}
s[name] = cur
return true
}
type taintAssignment struct {
target string
node ast.Vertex
rhs nodeSpan
origins []taintOrigin
}
// taintOrigin is one input of an assignment: a variable read, another
// assignment's output, or a fixed value from a source or summarized call.
// Only fixed origins carry proofs of their own, so those live out of line in
// compiledAssignments.fixed and the origin holds an index: an attacker picks
// how many origins a file has, and a gradeSet in every one of them tripled
// what analysis allocated.
type taintOrigin struct {
variable string
assignment int
// fixed is the 1-based index of this origin's value in
// compiledAssignments.fixed; 0 means the origin is not fixed.
fixed int32
decoded bool
}
// compiledAssignments is one scope's solver input. fixed holds each distinct
// fixed origin value once.
type compiledAssignments struct {
assignments []taintAssignment
fixed []gradeSet
}
// fixedValues interns fixed origin values for one compileAssignments call.
type fixedValues struct {
values []gradeSet
index map[gradeSet]int32
}
// ref returns the 1-based reference to set, adding it on first use.
func (v *fixedValues) ref(set gradeSet) int32 {
if i, ok := v.index[set]; ok {
return i
}
if v.index == nil {
v.index = make(map[gradeSet]int32)
}
v.values = append(v.values, set)
i := int32(len(v.values)) // #nosec G115 -- one value per collected call, and maxCollectedNodes bounds those far below int32
v.index[set] = i
return i
}
type positionedOrigin struct {
nodeSpan
origin taintOrigin
self int
}
type assignmentInterval struct {
nodeSpan
assignment int
}
// taintedLocals computes the tainted variables of one scope to a fixpoint. It
// is flow-insensitive: assignment order within the scope is not modelled. A
// dependency worklist reaches arbitrarily long local chains without borrowing
// the interprocedural summary-round limit or silently returning a partial state.
func taintedLocals(f *scopeFacts, summaries summaryTables) taintState {
compiled, ok := compileAssignments(f, summaries)
if !ok {
return taintedLocalsFallback(f, summaries)
}
return solveAssignments(compiled)
}
func compileAssignments(f *scopeFacts, summaries summaryTables) (compiledAssignments, bool) {
assignments := make([]taintAssignment, 0, len(f.assigns)+len(f.references)*2+len(f.concats))
for _, a := range f.assigns {
var ok bool
assignments, _, ok = appendAssignment(assignments, a, a.Var, a.Expr)
if !ok {
return compiledAssignments{}, false
}
}
for _, a := range f.references {
var ok bool
assignments, _, ok = appendAssignment(assignments, a, a.Var, a.Expr)
if !ok {
return compiledAssignments{}, false
}
}
for _, a := range f.concats {
var index int
var ok bool
assignments, index, ok = appendAssignment(assignments, a, a.Var, a.Expr)
if !ok {
return compiledAssignments{}, false
}
if index >= 0 {
assignments[index].origins = append(assignments[index].origins, taintOrigin{
variable: assignedTargetKey(a.Var), assignment: -1,
})
}
}
forwardCount := len(assignments)
var fixed fixedValues
positioned := make([]positionedOrigin, 0, len(f.varNodes)+len(f.callSites)+forwardCount)
for _, variable := range f.readVarNodes() {
span, ok := spanOf(variable.node)
if !ok {
return compiledAssignments{}, false
}
positioned = append(positioned, positionedOrigin{
nodeSpan: span,
origin: taintOrigin{variable: variable.name, assignment: -1},
self: -1,
})
}
for _, call := range f.callNodes {
if g, source := sourceGrade(call); source {
span, ok := spanOf(call)
if !ok {
return compiledAssignments{}, false
}
positioned = append(positioned, positionedOrigin{
nodeSpan: span,
origin: taintOrigin{
assignment: -1, fixed: fixed.ref(setOf(g)),
},
self: -1,
})
}
}
for _, call := range f.callSites {
set, summarized := summaries.lookup(call)
if !summarized {
continue
}
span, ok := spanOf(call.node)
if !ok {
return compiledAssignments{}, false
}
positioned = append(positioned, positionedOrigin{
nodeSpan: span,
origin: taintOrigin{
assignment: -1, fixed: fixed.ref(set),
},
self: -1,
})
}
for i := 0; i < forwardCount; i++ {
span, ok := spanOf(assignments[i].node)
if !ok {
return compiledAssignments{}, false
}
positioned = append(positioned, positionedOrigin{
nodeSpan: span,
origin: taintOrigin{assignment: i},
self: i,
})
}
decoderSpans := make([]nodeSpan, 0)
for _, call := range f.callNodes {
if !decoders[calleeName(call.Function)] {
continue
}
for _, input := range decoderInputs(call) {
if span, ok := spanOf(input); ok {
decoderSpans = append(decoderSpans, span)
} else {
return compiledAssignments{}, false
}
}
}
decoders := newSpanIndex(decoderSpans)
for i := range positioned {
positioned[i].origin.decoded = decoders.contains(positioned[i].nodeSpan)
}
distributeOrigins(assignments, positioned)
// PHP references alias both names. A later write through either side
// changes the other, so add the persistent reverse dependency after the
// concrete assignment expression has received its own RHS origins.
for _, reference := range f.references {
target := assignedTargetKey(reference.Expr)
source := assignedTargetKey(reference.Var)
if target == "" || source == "" {
continue
}
assignments = append(assignments, taintAssignment{
target: target,
origins: []taintOrigin{{
variable: source, assignment: -1,
}},
})
}
return compiledAssignments{assignments: assignments, fixed: fixed.values}, true
}
func appendAssignment(assignments []taintAssignment, node, target, expr ast.Vertex) ([]taintAssignment, int, bool) {
name := assignedTargetKey(target)
if name == "" {
return assignments, -1, true
}
rhs, ok := spanOf(expr)
if !ok {
return nil, -1, false
}
assignments = append(assignments, taintAssignment{target: name, node: node, rhs: rhs})
return assignments, len(assignments) - 1, true
}
func distributeOrigins(assignments []taintAssignment, origins []positionedOrigin) {
intervals := make([]assignmentInterval, 0, len(assignments))
for i := range assignments {
if assignments[i].node != nil {
intervals = append(intervals, assignmentInterval{nodeSpan: assignments[i].rhs, assignment: i})
}
}
sort.Slice(intervals, func(i, j int) bool {
if intervals[i].start != intervals[j].start {
return intervals[i].start < intervals[j].start
}
return intervals[i].end > intervals[j].end
})
sort.Slice(origins, func(i, j int) bool {
if origins[i].start != origins[j].start {
return origins[i].start < origins[j].start
}
return origins[i].end < origins[j].end
})
stack := make([]assignmentInterval, 0)
next := 0
for _, origin := range origins {
for next < len(intervals) && intervals[next].start <= origin.start {
interval := intervals[next]
for len(stack) > 0 && stack[len(stack)-1].end < interval.start {
stack = stack[:len(stack)-1]
}
stack = append(stack, interval)
next++
}
for len(stack) > 0 && stack[len(stack)-1].end < origin.end {
stack = stack[:len(stack)-1]
}
for i := len(stack) - 1; i >= 0; i-- {
interval := stack[i]
if interval.assignment == origin.self || interval.start > origin.start || interval.end < origin.end {
continue
}
assignments[interval.assignment].origins = append(assignments[interval.assignment].origins, origin.origin)
break
}
}
}
type spanIndex struct {
spans []nodeSpan
maxEnds []int
}
func newSpanIndex(spans []nodeSpan) spanIndex {
sort.Slice(spans, func(i, j int) bool { return spans[i].start < spans[j].start })
index := spanIndex{spans: spans, maxEnds: make([]int, len(spans))}
maxEnd := -1
for i, span := range spans {
if span.end > maxEnd {
maxEnd = span.end
}
index.maxEnds[i] = maxEnd
}
return index
}
func (index spanIndex) contains(span nodeSpan) bool {
i := sort.Search(len(index.spans), func(i int) bool { return index.spans[i].start > span.start }) - 1
return i >= 0 && index.maxEnds[i] >= span.end
}
func solveAssignments(compiled compiledAssignments) taintState {
assignments := compiled.assignments
st := taintState{}
variableDependents := make(map[string][]int)
assignmentDependents := make([][]int, len(assignments))
for i, assignment := range assignments {
for _, origin := range assignment.origins {
switch {
case origin.fixed > 0:
case origin.assignment >= 0:
assignmentDependents[origin.assignment] = append(assignmentDependents[origin.assignment], i)
case origin.variable != "":
variableDependents[origin.variable] = append(variableDependents[origin.variable], i)
}
}
}
queue := make([]int, len(assignments))
queued := make([]bool, len(assignments))
outputs := make([]gradeSet, len(assignments))
outputSet := make([]bool, len(assignments))
for i := range assignments {
queue[i] = i
queued[i] = true
}
enqueue := func(indices []int) {
for _, index := range indices {
if !queued[index] {
queued[index] = true
queue = append(queue, index)
}
}
}
for head := 0; head < len(queue); head++ {
i := queue[head]
queued[i] = false
var best gradeSet
found := false
for _, origin := range assignments[i].origins {
var value gradeSet
var active bool
switch {
case origin.fixed > 0:
value, active = compiled.fixed[origin.fixed-1], true
case origin.assignment >= 0:
value, active = outputs[origin.assignment], outputSet[origin.assignment]
default:
value, active = st[origin.variable]
}
if !active {
continue
}
if origin.decoded {
// Decoding fetched content raises the confidence of this
// origin's proofs only; each keeps its own basis, and an
// origin outside every decoder keeps its own confidence.
value = value.decoded()
}
best.add(value)
found = true
}
if !found {
continue
}
if outputSet[i] {
merged := outputs[i]
if !merged.add(best) {
continue
}
best = merged
}
outputs[i], outputSet[i] = best, true
enqueue(assignmentDependents[i])
if st.raise(assignments[i].target, best) {
enqueue(variableDependents[assignments[i].target])
}
}
return st
}
// taintedLocalsFallback is the oracle path: no compiled solver, just direct
// fixpoint iteration over the raw facts. Its round cap must terminate AND
// stay complete. A fixed cap (the original defect) is neither: a reverse-
// ordered assignment chain propagates taint exactly one hop per round, so a
// chain longer than the cap evades detection entirely. Capping at the number
// of assignment facts plus one is always enough, because no dependency chain
// in this scope can have more hops than this scope has assignments, and it
// still terminates because it is a fixed bound.
func taintedLocalsFallback(f *scopeFacts, summaries summaryTables) taintState {
st := taintState{}
maxRounds := len(f.assigns) + len(f.references) + len(f.concats) + 1
for round := 0; round < maxRounds; round++ {
changed := false
for _, assignment := range f.assigns {
if g, tainted := exprTaint(assignment.Expr, st, summaries); tainted {
changed = st.raise(assignedTargetKey(assignment.Var), g) || changed
}
}
for _, assignment := range f.references {
if g, tainted := exprTaint(assignment.Expr, st, summaries); tainted {
changed = st.raise(assignedTargetKey(assignment.Var), g) || changed
}
if g, tainted := exprTaint(assignment.Var, st, summaries); tainted {
changed = st.raise(assignedTargetKey(assignment.Expr), g) || changed
}
}
for _, assignment := range f.concats {
if g, tainted := exprTaint(assignment.Expr, st, summaries); tainted {
changed = st.raise(assignedTargetKey(assignment.Var), g) || changed
}
}
if !changed {
return st
}
}
return st
}
// assignedVarName returns the root local variable written by an assignment,
// flattening any array-element or property-fetch chain onto its base
// variable. This remains the array-element behavior (the spec permits
// over-approximating containers, and array keys are not tracked at all), and
// it is also assignedTargetKey's fallback for a property chain it cannot key
// more precisely -- see there for when that applies.
func assignedVarName(target ast.Vertex) string {
for target != nil {
switch n := target.(type) {
case *ast.ExprVariable:
return varName(n.Name)
case *ast.ExprArrayDimFetch:
target = n.Var
case *ast.ExprPropertyFetch:
target = n.Var
default:
return ""
}
}
return ""
}
// propertyFetchParts returns the base expression and the (possibly dynamic)
// property vertex of a property fetch, regardless of whether it used the
// regular ("->") or nullsafe ("?->") operator. Nullsafe fetches cannot
// appear as an assignment target in valid PHP, but a property that was
// WRITTEN through the regular operator must still be found by a later READ
// spelled with "?->", so both node types resolve through the one path
// assignedTargetKey uses for both directions.
func propertyFetchParts(n ast.Vertex) (base, prop ast.Vertex, ok bool) {
switch v := n.(type) {
case *ast.ExprPropertyFetch:
return v.Var, v.Prop, true
case *ast.ExprNullsafePropertyFetch:
return v.Var, v.Prop, true
}
return nil, nil, false
}
// assignedTargetKey returns the taint-state key an assignment target -- or,
// via readVarNodes, a later read of the identical expression shape --
// resolves to. A bare variable resolves to its own name.
//
// The WHOLE access chain is walked, in whatever order property fetches
// ("->"/"?->") and array-dim fetches ("[...]") appear and however deeply
// nested: each property fetch contributes its name to the compound key, and
// each array-dim fetch is transparent, unwrapped without contributing a
// segment of its own. That is what keeps array indices over-approximated
// (the spec permits over-approximating containers: $a->log[] = X and
// $a->log[3] = Y both key to "a->log" as a whole -- array keys are not
// tracked) while still scoping precisely to the property chain around them.
//
// This generality is load-bearing, not incidental: an earlier version
// only unwrapped property fetches, so an array-dim fetch sitting
// OUTERMOST over a property fetch ($this->log[] = <tainted>, one of the
// most common idioms in OO PHP -- logs, queues, error collections, caches)
// broke the walk immediately and fell back to the bare base variable,
// reintroducing the exact wholesale-$this false-positive class this key
// scoping exists to close. Handling only that one shape and stopping would
// leave the same gap for every other ordering ($a[0]->b, $a->b[0]->c[1],
// ...), so the loop below handles both access kinds generically, in any
// order, rather than adding a case for the reported shape and calling it
// done.
//
// A property fetch is different from a bare variable or array element:
// $this->body = <tainted> must NOT taint every later use of bare $this,
// because $this appears in nearly every method of a class and one tainted
// property would otherwise poison every sink in it. This is deliberately
// not special-cased to the name "$this": $obj->prop = <tainted> gets
// exactly the same treatment. The key scopes to
// the full static property path (so $this->a->b is distinct from both
// $this->a and $this->a->c), joining base and property names with "->", a
// sequence no PHP variable name can contain, so a compound key can never
// collide with a bare one.
//
// The walk is iterative, not recursive: an access chain's depth is
// attacker-controlled PHP source ($a->b->c->... or $a[0][0][0]...), and Go
// recursion over unbounded attacker-controlled depth is an unrecoverable
// stack overflow, the same hazard TestTaintHandlesDeeplyNestedAssignments
// guards elsewhere in this package.
//
// A property whose name is not statically known ($obj->$name = X) has no
// specific key to scope to, so the ENTIRE chain falls back to
// assignedVarName's base-variable over-approximation, rather than a
// partially-keyed name nothing else would ever match.
func assignedTargetKey(target ast.Vertex) string {
var props []string
node := target
for {
if base, prop, ok := propertyFetchParts(node); ok {
name := calleeName(prop)
if name == "" {
return assignedVarName(target)
}
props = append(props, name)
node = base
continue
}
if dim, ok := node.(*ast.ExprArrayDimFetch); ok {
node = dim.Var
continue
}
break
}
base := assignedVarName(node)
if base == "" || len(props) == 0 {
return base
}
for i, j := 0, len(props)-1; i < j; i, j = i+1, j-1 {
props[i], props[j] = props[j], props[i]
}
return base + "->" + strings.Join(props, "->")
}
// exprTaint reports whether an expression carries remote content, and with
// which proofs. It collects the subtree once and correlates decoder inputs by
// source positions, without recursing over the parsed structure.
func exprTaint(e ast.Vertex, st taintState, summaries summaryTables) (gradeSet, bool) {
if e == nil {
return gradeSet{}, false
}
return exprTaintFacts(collectScope(e), st, summaries)
}
func exprTaintFacts(sub *scopeFacts, st taintState, summaries summaryTables) (gradeSet, bool) {
// A decoder raises confidence to Certain only for the origins that
// passed through THAT decoder's own argument, not for anything else in
// the same expression: f(base64_decode($clean), $tainted) must stay at
// the source's own grade, because the decode never touched $tainted, and
// in base64_decode(a()) . b() only a()'s proofs become Certain. The
// upgrade changes confidence only: each source's basis still explains
// how its content was acquired.
decoderSpans := make([]nodeSpan, 0)
for _, call := range sub.callNodes {
if !decoders[calleeName(call.Function)] {
continue
}
for _, input := range decoderInputs(call) {
if span, ok := spanOf(input); ok {
decoderSpans = append(decoderSpans, span)
}
}
}
return activeTaint(sub, st, summaries, newSpanIndex(decoderSpans))
}
// decoderInputs returns only arguments that carry data through the decoder.
// Most decoder APIs consume their payload in argument zero. pack is the
// exception: argument zero is a format string and the values begin at one.
func decoderInputs(call *ast.ExprFunctionCall) []ast.Vertex {
start := 0
if calleeName(call.Function) == "pack" {
start = 1
}
if start >= len(call.Args) {
return nil
}
inputs := make([]ast.Vertex, 0, len(call.Args)-start)
end := start + 1
if start == 1 {
end = len(call.Args)
}
for _, argNode := range call.Args[start:end] {
if arg, ok := argNode.(*ast.Argument); ok {
inputs = append(inputs, arg.Expr)
}
}
return inputs
}
// flowResult retains whether its already-bounded display evidence was
// shortened. Keeping this beside the Result until deduplication is complete
// lets Analyze report truncation only for evidence that actually survives to
// the returned result slice.
type flowResult struct {
Result
evidenceTruncated bool
}
// capturableNames expands a scope's taint keys into the binding names a
// capture can actually spell. Taint on a property is keyed by its whole access
// path ("o->body"), but a closure captures the base variable ($o) and receives
// the property along with it, so the base name has to be markable too or a
// genuinely dropped capture is never recorded.
func capturableNames(st taintState) map[string]bool {
names := make(map[string]bool, len(st))
for key := range st {
names[key] = true
if base, _, found := strings.Cut(key, "->"); found {
names[base] = true
}
}
return names
}
// hasDroppedCapture reports whether any closure or arrow function receives an
// outer binding that was tainted in the scope it captured from.
//
// A closure's body is analysed as its own scope, which is what stops an
// unrelated same-named outer variable from firing on clean code. The cost is
// that a value the closure genuinely does receive through use(), or that an
// arrow function picks up implicitly, stops being tracked at the boundary.
// This package's contract is that a reduction in coverage is recorded rather
// than passed over, so the drop is reported as precision loss.
//
// The check is deliberately gated on the captured value ACTUALLY being
// tainted. Capturing is ordinary PHP and appears in roughly one in seven of
// the files this analyzer examines, so a marker raised on the shape alone
// would be noise an operator cannot act on. Gated this way it means exactly
// one thing: taint was dropped here.
//
// Enclosing states are computed lazily and memoised, so a file with no
// capturing closure pays nothing and a scope with several capturing children
// is solved once. Nothing is seeded INTO a closure: taint deliberately does
// not cross the boundary, it is only reported as lost.
func hasDroppedCapture(
ctx context.Context, all *scopeFacts, tree declTree,
factsByScope map[ast.Vertex]*scopeFacts, summaries summaryTables,
) (bool, error) {
if len(all.closures) == 0 && len(all.arrowFuncs) == 0 {
return false, nil
}
// Record only declarations that actually capture a name. Building this
// cheap structural index first preserves the important lazy path: a file
// containing only non-capturing declarations never solves another taint
// state merely to decide that there was no capture.
captures := make(map[ast.Vertex]map[string]bool)
for _, cl := range all.closures {
if err := ctx.Err(); err != nil {
return false, err
}
names := closureCaptureNames(cl)
// A non-static closure created in object context binds $this
// implicitly, so it appears in no use() clause while still carrying
// whatever the enclosing object holds.
if body := factsByScope[cl]; cl.StaticTkn == nil && body != nil && body.vars["this"] {
names["this"] = true
}
if len(names) > 0 {
captures[cl] = names
}
}
for _, af := range all.arrowFuncs {
if err := ctx.Err(); err != nil {
return false, err
}
body := factsByScope[af]
if body == nil {
continue
}
if names := arrowCaptureNames(af, body); len(names) > 0 {
captures[af] = names
}
}
// $this is the exception to an ordinary closure being a hard capture
// boundary. Every non-static closure created in object context is bound to
// that object, even when only a declaration nested inside the closure reads
// $this. The independently collected closure facts exclude that nested
// declaration, so surface the implicit capture on every non-static
// closure/arrow along the path. Static declarations and named lexical
// scopes remain hard boundaries. Stopping at an already surfaced node keeps
// the total walk linear when many captures share a deeply nested path.
thisForwarded := make(map[ast.Vertex]bool)
var thisSeeds []ast.Vertex
for node, names := range captures {
if names["this"] {
thisSeeds = append(thisSeeds, node)
}
}
for _, node := range thisSeeds {
for scope := tree.parent[node]; scope != nil && !thisForwarded[scope]; scope = tree.parent[scope] {
if err := ctx.Err(); err != nil {
return false, err
}
forwards := false
switch declaration := scope.(type) {
case *ast.ExprClosure:
forwards = declaration.StaticTkn == nil
case *ast.ExprArrowFunction:
forwards = declaration.StaticTkn == nil
}
if !forwards {
break
}
thisForwarded[scope] = true
names := captures[scope]
if names == nil {
names = map[string]bool{}
captures[scope] = names
}
names["this"] = true
}
}
if len(captures) == 0 {
return false, nil
}
// A capture nested directly in an arrow function also makes that arrow
// capture the same binding implicitly. Mark every enclosing arrow on a
// capture path, plus the first ordinary lexical scope that supplies the
// value. Stopping when an already-marked scope is reached makes the total
// walk linear in declaration count even for deeply nested arrow trees.
needed := make(map[ast.Vertex]bool, len(captures)+1)
for node := range captures {
scope := tree.parent[node]
for !needed[scope] {
needed[scope] = true
if _, arrow := scope.(*ast.ExprArrowFunction); !arrow {
break
}
scope = tree.parent[scope]
}
}
// Taint states remain lazy and are memoised by their lexical scope. A
// class body can enclose a declaration but has no local variable facts of
// its own, so its state is intentionally empty.
states := make(map[ast.Vertex]taintState, len(needed))
stateFor := func(scope ast.Vertex) (taintState, error) {
if st, solved := states[scope]; solved {
return st, nil
}
if err := ctx.Err(); err != nil {
return nil, err
}
var st taintState
if f := factsByScope[scope]; f != nil {
st = taintedLocals(f, summaries)
}
if err := ctx.Err(); err != nil {
return nil, err
}
states[scope] = st
return st, nil
}
// active is a marker-only view of which tainted bindings are available at
// the current declaration boundary. It is never passed to taintedLocals or
// findFlows and therefore never seeds a closure or arrow body. Arrow scopes
// inherit the view because PHP propagates captures through nested arrows;
// every other declaration starts a new lexical variable scope. Boundary
// IDs make those resets O(1), while per-name stacks make enter/leave O(1).
type binding struct {
boundary int
tainted bool
}
type frame struct {
end int
previousBoundary int
pushed []string
}
active := make(map[string][]binding)
boundary, nextBoundary := 0, 0
push := func(name string, tainted bool, pushed *[]string) {
active[name] = append(active[name], binding{boundary: boundary, tainted: tainted})
*pushed = append(*pushed, name)
}
pop := func(f frame) {
for i := len(f.pushed) - 1; i >= 0; i-- {
name := f.pushed[i]
stack := active[name]
if len(stack) == 1 {
delete(active, name)
} else {
active[name] = stack[:len(stack)-1]
}
}
boundary = f.previousBoundary
}
capturesActive := func(names map[string]bool) bool {
for name := range names {
stack := active[name]
if len(stack) == 0 {
continue
}
value := stack[len(stack)-1]
if value.boundary == boundary && value.tainted {
return true
}
}
return false
}
if needed[nil] {
st, err := stateFor(nil)
if err != nil {
return false, err
}
for name := range capturableNames(st) {
active[name] = []binding{{boundary: boundary, tainted: true}}
}
}
frames := make([]frame, 0, 8)
seen := make(map[ast.Vertex]bool, len(captures))
for _, declaration := range tree.ordered {
if err := ctx.Err(); err != nil {
return false, err
}
// Exclusive end position, matching declarationTree's sweep: a
// declaration starting exactly where a frame ends is outside it.
for len(frames) > 0 && frames[len(frames)-1].end <= declaration.start {
pop(frames[len(frames)-1])
frames = frames[:len(frames)-1]
}
if names := captures[declaration.node]; len(names) > 0 {
seen[declaration.node] = true
if capturesActive(names) {
if err := ctx.Err(); err != nil {
return false, err
}
return true, nil
}
}
f := frame{end: declaration.end, previousBoundary: boundary}
if needed[declaration.node] {
if af, arrow := declaration.node.(*ast.ExprArrowFunction); arrow {
// A static arrow still captures ordinary outer variables, but
// it neither receives $this nor forwards it to a declaration
// nested inside the arrow.
if af.StaticTkn != nil {
push("this", false, &f.pushed)
}
// Parameters are the arrow's own bindings and shadow an
// identically named capture from any enclosing scope.
for name := range paramNames(af.Params) {
push(name, false, &f.pushed)
}
} else {
nextBoundary++
boundary = nextBoundary
}
st, err := stateFor(declaration.node)
if err != nil {
return false, err
}
for name := range capturableNames(st) {
push(name, true, &f.pushed)
}
}
frames = append(frames, f)
}
// Parser-produced declaration nodes always carry positions and therefore
// appear in tree.ordered. Retain the conservative direct-scope behavior if
// a future parser version emits an unpositioned capture node.
for node, names := range captures {
if seen[node] {
continue
}
st, err := stateFor(tree.parent[node])
if err != nil {
return false, err
}
available := capturableNames(st)
for name := range names {
if available[name] {
return true, nil
}
}
}
if err := ctx.Err(); err != nil {
return false, err
}
return false, nil
}
// unresolvableAssignRHS returns the right-hand-side expression of every
// assignment (=, =&, .=) in f whose target assignedTargetKey cannot resolve
// to a taint-state key: a method-call result mid-chain (`$a->b()->c = X`),
// a static property (`Foo::$cache = X`), a list()/[] destructuring target,
// or a variable-variable base. This is a purely structural scan (no taint
// fixpoint), used to cheaply decide whether hasUnresolvableTaintedTarget
// needs to do any further work at all.
func unresolvableAssignRHS(ctx context.Context, f *scopeFacts) ([]ast.Vertex, error) {
var out []ast.Vertex
for _, a := range f.assigns {
if err := ctx.Err(); err != nil {
return nil, err
}
if assignedTargetKey(a.Var) == "" {
out = append(out, a.Expr)
}
}
for _, a := range f.references {
if err := ctx.Err(); err != nil {
return nil, err
}
if assignedTargetKey(a.Var) == "" {
out = append(out, a.Expr)
}
}
for _, a := range f.concats {
if err := ctx.Err(); err != nil {
return nil, err
}
if assignedTargetKey(a.Var) == "" {
out = append(out, a.Expr)
}
}
return out, nil
}
// hasUnresolvableTaintedTarget reports whether f contains an assignment
// whose target cannot be keyed (see unresolvableAssignRHS) AND whose
// right-hand side is actually tainted -- the only case that represents a
// real, silent loss of tracking. `list($a, $b) = ['x', 'y']` and
// `Foo::$cache = 'literal'` drop nothing at all and must not be flagged;
// `list($a, $b) = curl_exec($u)` and `Foo::$cache = curl_exec($u)`
// genuinely drop taint this package cannot track further and must be.
//
// An earlier version of this check fired on the unkeyable SHAPE alone,
// regardless of taint. Measured against the reference corpus that reported
// the marker on 38.75% of analyzed files -- overwhelmingly ordinary,
// completely benign idioms (list() destructuring, static properties used
// as singletons/caches) that were not dropping anything -- noise no
// operator could act on. Gating on the RHS's own taint is what makes the
// marker mean "we dropped taint here" instead of "this file uses PHP".
//
// f.assigns/f.references/f.concats are scanned for an unresolvable target
// FIRST (see unresolvableAssignRHS), before the per-scope taint fixpoint
// below runs at all, so the (relatively expensive) taintedLocals call is
// skipped entirely for the overwhelming majority of scopes that have no
// unresolvable target -- the same early-return findFlows already uses for
// f.sinks being empty. Active origins are then correlated against the RHS
// spans from f's existing facts. Besides avoiding a traversal per RHS, this
// preserves canonical names already resolved from `use function` aliases;
// recollecting an isolated RHS has no access to the enclosing alias imports.
func hasUnresolvableTaintedTarget(
ctx context.Context, f *scopeFacts, summaries summaryTables,
) (bool, error) {
if err := ctx.Err(); err != nil {
return false, err
}
rhs, err := unresolvableAssignRHS(ctx, f)
if err != nil {
return false, err
}
if len(rhs) == 0 {
return false, nil
}
rhsSpans := make([]nodeSpan, 0, len(rhs))
var unpositioned []ast.Vertex
for _, expr := range rhs {
if err := ctx.Err(); err != nil {
return false, err
}
if span, ok := spanOf(expr); ok {
rhsSpans = append(rhsSpans, span)
} else {
unpositioned = append(unpositioned, expr)
}
}
var rhsReads []namedNodeSpan
if len(rhsSpans) > 0 {
index := newSpanIndex(rhsSpans)
for _, call := range f.callNodes {
if err := ctx.Err(); err != nil {
return false, err
}
span, positioned := spanOf(call)
if positioned && index.contains(span) {
if _, source := sourceGrade(call); source {
return true, nil
}
}
}
for _, call := range f.callSites {
if err := ctx.Err(); err != nil {
return false, err
}
span, positioned := spanOf(call.node)
if positioned && index.contains(span) {
if _, summarized := summaries.lookup(call); summarized {
return true, nil
}
}
}
for _, variable := range f.readVarNodes() {
if err := ctx.Err(); err != nil {
return false, err
}
span, positioned := spanOf(variable.node)
if positioned && index.contains(span) {
rhsReads = append(rhsReads, variable)
}
}
}
// Literal RHS values and direct source/summary calls are already decided
// above. Only variable-carried taint (or a synthetic positionless AST)
// needs the full per-scope assignment fixpoint.
if len(rhsReads) == 0 && len(unpositioned) == 0 {
return false, nil
}
st := taintedLocals(f, summaries)
if err := ctx.Err(); err != nil {
return false, err
}
for _, variable := range rhsReads {
if _, tainted := st[variable.name]; tainted {
return true, nil
}
}
// Parser-produced nodes have positions. Keep the helper total for
// synthetic or future positionless ASTs without turning shape alone into
// a precision-loss marker.
for _, expr := range unpositioned {
if err := ctx.Err(); err != nil {
return false, err
}
if _, tainted := exprTaint(expr, st, summaries); tainted {
return true, nil
}
}
return false, nil
}
// findFlows reports each sink in a scope that receives remote content.
// exclude names the declarations nested inside this scope. A sink's argument
// expression is collected WITHOUT it, so a call reached only through a closure
// there is still followed -- `include array_map(function(){ return fetch(); },
// $r)[0]` really does receive what that closure returns. Its variables are then
// filtered THROUGH it, because they are graded against this scope's taint
// state and a nested declaration's own parameter or local was never part of
// it. Same split, and same reason, as the return expressions in summaries.go.
// wholeCalls is the resolved-call index built from the unfiltered whole-file
// facts: the enclosing f correctly excludes nested declarations, so an index
// rebuilt from f cannot resolve namespace-level aliases on the calls
// deliberately kept from those declarations.
func findFlows(
ctx context.Context, f *scopeFacts, summaries summaryTables,
wholeCalls resolvedCallIndex, exclude *spanIndex,
) ([]flowResult, error) {
if err := ctx.Err(); err != nil {
return nil, err
}
if len(f.sinks) == 0 {
return nil, nil
}
st := taintedLocals(f, summaries)
if err := ctx.Err(); err != nil {
return nil, err
}
written := taintedWritePaths(f, st, summaries, wholeCalls, exclude)
out := make([]flowResult, 0, len(f.sinks))
for _, s := range f.sinks {
if err := ctx.Err(); err != nil {
return nil, err
}
sub := wholeCalls.apply(collectScope(s.expr)).withoutNestedDeclarationVars(exclude)
set, tainted := exprTaintFacts(sub, st, summaries)
if !tainted {
// A file this scope wrote from remote content and now includes
// by the same path expression executes that content just as a
// tainted variable would: the file is the carrier.
if ev, ok := written[pathExprKey(s.expr)]; ok && includeSink(s.kind) {
identifiers, identifiersTruncated := identifiersFor(sub)
out = append(out, flowResult{
Result: Result{
Source: ev.source,
Identifiers: identifiers,
Sink: s.kind,
Confidence: ev.value.conf,
Basis: ev.value.basis,
ResolutionOffset: ev.value.offset,
},
evidenceTruncated: ev.truncated || identifiersTruncated,
})
}
continue
}
c := set.strongest()
source, sourceTruncated := sourceLabel(sub, st, summaries)
identifiers, identifiersTruncated := identifiersFor(sub)
out = append(out, flowResult{
Result: Result{
Source: source,
Identifiers: identifiers,
Sink: s.kind,
Confidence: c.conf,
Basis: c.basis,
ResolutionOffset: c.offset,
},
evidenceTruncated: sourceTruncated || identifiersTruncated,
})
}
return out, nil
}
// sourceLabel names the acquiring construct for evidence. It prefers a
// directly visible source call, then a taint-returning callee, then the
// tainted variable that carried the value in. Resolution goes through
// callSites (rather than the bare f.calls name set) so a summarized method
// is matched via the same call-syntax-scoped lookup taintedLocals uses,
// instead of guessing across the function/method namespaces.
func sourceLabel(sub *scopeFacts, st taintState, summaries summaryTables) (string, bool) {
for _, call := range sub.callNodes {
if _, ok := sourceGrade(call); ok {
return sanitize(calleeName(call.Function), maxSegmentBytes)
}
}
names := make([]string, 0, len(sub.callSites))
for _, call := range sub.callSites {
if _, ok := summaries.lookup(call); ok {
names = append(names, call.name)
}
}
sort.Strings(names)
if len(names) > 0 {
return sanitize(names[0], maxSegmentBytes)
}
vars := make([]string, 0, len(sub.vars)+len(sub.propNodes))
for name := range sub.vars {
if _, ok := st[name]; ok {
vars = append(vars, name)
}
}
// A property key (e.g. "this->body") is checked separately from the
// bare-variable loop above: the base variable's own name is not tainted
// by a property write, so a flow carried entirely by a specific property
// would otherwise fall through every branch above and report "unknown"
// instead of naming the property that actually carried it.
for _, n := range sub.propNodes {
key := assignedTargetKey(n)
if key == "" {
continue
}
if _, ok := st[key]; ok {
vars = append(vars, key)
}
}
sort.Strings(vars)
if len(vars) > 0 {
return sanitize("$"+vars[0], maxSegmentBytes)
}
return "unknown", false
}
// identifiersFor renders the sanitized, bounded Identifiers list for
// evidence: every distinct variable and call name appearing anywhere in the
// sink's own expression, tainted or not. This is not a laundering path --
// there is no ordering or filtering by whether a name actually carried the
// tainted value, only alphabetical sort for deterministic output.
func identifiersFor(sub *scopeFacts) ([]string, bool) {
segs := make([]string, 0, len(sub.vars)+len(sub.calls))
for name := range sub.vars {
if name != "" {
segs = append(segs, "$"+name)
}
}
for name := range sub.calls {
segs = append(segs, name)
}
sort.Strings(segs)
return truncateChain(segs)
}
// dedupeAndSort collapses flows that render the same display endpoint pair
// and orders results so repeated runs are byte-identical.
func dedupeAndSort(in []flowResult) []flowResult {
seen := map[string]int{}
out := make([]flowResult, 0, len(in))
for _, r := range in {
key := r.Source + "\x00" + r.Sink
if idx, ok := seen[key]; ok {
// The stronger flow replaces the whole grade, so a raised
// confidence always carries its own basis and resolution point.
cur := directGradeOf(out[idx].Result)
if next := directGradeOf(r.Result); next.stronger(cur) {
out[idx].Confidence = r.Confidence
out[idx].Basis = r.Basis
out[idx].ResolutionOffset = r.ResolutionOffset
}
continue
}
seen[key] = len(out)
out = append(out, r)
}
sort.Slice(out, func(i, j int) bool {
if out[i].Source != out[j].Source {
return out[i].Source < out[j].Source
}
if out[i].Sink != out[j].Sink {
return out[i].Sink < out[j].Sink
}
return strings.Join(out[i].Identifiers, ",") < strings.Join(out[j].Identifiers, ",")
})
return out
}
// retainStrongestEvidence applies the display cap without hiding the grade
// that downstream alerting must use. The input is endpoint-sorted, so when the
// first strongest flow falls beyond the cap, replacing the last retained flow
// with it keeps the returned evidence deterministic and sorted.
func retainStrongestEvidence(flows []flowResult) []flowResult {
if len(flows) <= MaxEvidenceResults {
return flows
}
strongest := 0
for i := 1; i < len(flows); i++ {
// Basis only explains a flow; it must not change the retained
// endpoints (and therefore the finding's identity) on a confidence tie.
if flows[i].Confidence > flows[strongest].Confidence {
strongest = i
}
}
retained := append([]flowResult(nil), flows[:MaxEvidenceResults]...)
if strongest >= MaxEvidenceResults {
retained[len(retained)-1] = flows[strongest]
}
return retained
}
// activeTaint joins the proofs of an already-collected subtree's source
// calls, summarized calls, and tainted variable reads. Each origin inside a
// decoder input span is upgraded on its own before the join, so correlation
// stays linearithmic rather than recursively recollecting every nested
// decoder argument, and a proof outside every decoder keeps its confidence.
// A variable name is looked up in a taintState in exactly five places in this
// package, and every one of them must be fed by facts collected WITH
// declaration exclusion, or a nested declaration's own parameter or local
// borrows an identically named outer variable's taint and reports a flow on
// clean code. Names like $data, $content and $url recur constantly in real
// PHP, so a bare collision is enough. The five, and what keeps each scoped:
//
// solveAssignments origins built from the scope's own facts
// hasUnresolvableTaintedTarget the scope's own facts
// sourceLabel a filtered sub (vars rebuilt to match)
// stateFor taintedLocals over a per-scope facts value
// activeTaint (below) see the three feeders named next
//
// activeTaint is reached only through exprTaintFacts, which has three
// feeders: evalBodySummary and findFlows, both of which collect their
// expression WITHOUT exclusion so that a call reached only through a nested
// closure is still followed, and then filter its variable side back through
// the scope's exclusion index; and exprTaint, whose two callers are both
// reachable only for an AST node carrying no position, which this parser does
// not produce.
//
// Anything new that grades an expression against a scope's taint state joins
// this list and needs the same treatment. Fixing one instance at a time does
// not work here: this defect was found and fixed three separate times -- in
// the summaries path, in the sink path, and in the capture walk -- before the
// enumeration above made it possible to say the class was closed rather than
// merely that no more instances had turned up.
func activeTaint(sub *scopeFacts, st taintState, summaries summaryTables, decoderArgs spanIndex) (gradeSet, bool) {
var best gradeSet
found := false
join := func(set gradeSet, node ast.Vertex) {
// An origin without a position cannot be placed inside a decoder
// input, so it keeps its own grade.
if span, ok := spanOf(node); ok && decoderArgs.contains(span) {
set = set.decoded()
}
best.add(set)
found = true
}
for _, call := range sub.callNodes {
if c, ok := sourceGrade(call); ok {
join(setOf(c), call)
}
}
for _, call := range sub.callSites {
if c, ok := summaries.lookup(call); ok {
join(c, call.node)
}
}
for _, variable := range sub.readVarNodes() {
if c, ok := st[variable.name]; ok {
join(c, variable.node)
}
}
return best, found
}
// writtenEvidence is what a tainted file write contributes to an include of
// the same path.
type writtenEvidence struct {
source string
value grade
truncated bool
}
// taintedWritePaths maps the path key of every file_put_contents whose data
// argument is tainted in this scope to the evidence of that taint.
func taintedWritePaths(
f *scopeFacts, st taintState, summaries summaryTables,
wholeCalls resolvedCallIndex, exclude *spanIndex,
) map[string]writtenEvidence {
if len(f.fileWrites) == 0 {
return nil
}
out := make(map[string]writtenEvidence, len(f.fileWrites))
for _, w := range f.fileWrites {
key := pathExprKey(w.path)
if key == "" {
continue
}
sub := wholeCalls.apply(collectScope(w.data)).withoutNestedDeclarationVars(exclude)
set, tainted := exprTaintFacts(sub, st, summaries)
if !tainted {
continue
}
c := set.strongest()
if prev, ok := out[key]; ok {
if !c.stronger(prev.value) {
continue
}
if c.conf == prev.value.conf {
// Keep the first strongest-confidence endpoint as before;
// only its explanation changes when another basis wins.
prev.value = c
out[key] = prev
continue
}
}
source, truncated := sourceLabel(sub, st, summaries)
out[key] = writtenEvidence{source: source, value: c, truncated: truncated}
}
return out
}
func includeSink(kind string) bool {
switch kind {
case "include", "include_once", "require", "require_once":
return true
}
return false
}
// pathExprKey renders a path expression into a canonical string so a write
// and an include of the same path match syntactically: literals, magic
// constants, constants, variables, concatenations and argument-less or
// keyable calls. Anything else (a property, an array element, an unresolved
// call) yields "" and never matches, so no guess can produce a flow.
func pathExprKey(n ast.Vertex) string { return pathExprKeyAt(n, 0) }
func pathExprKeyAt(n ast.Vertex, depth int) string {
if depth >= maxAnalysisDepth {
return ""
}
switch v := n.(type) {
case *ast.ScalarString:
return "s:" + string(v.Value)
case *ast.ScalarMagicConstant:
return "m:" + strings.ToLower(string(v.Value))
case *ast.ExprConstFetch:
return "c:" + strings.ToLower(calleeName(v.Const))
case *ast.ExprVariable:
name := varName(v.Name)
if name == "" {
return ""
}
return "v:" + name
case *ast.ExprBrackets:
return pathExprKeyAt(v.Expr, depth+1)
case *ast.ExprBinaryConcat:
left := pathExprKeyAt(v.Left, depth+1)
right := pathExprKeyAt(v.Right, depth+1)
if left == "" || right == "" {
return ""
}
return "(" + left + "." + right + ")"
case *ast.ExprFunctionCall:
name := calleeName(v.Function)
if name == "" {
return ""
}
parts := make([]string, 0, len(v.Args))
for _, a := range v.Args {
arg, ok := a.(*ast.Argument)
if !ok {
return ""
}
key := pathExprKeyAt(arg.Expr, depth+1)
if key == "" {
return ""
}
parts = append(parts, key)
}
return "f:" + strings.ToLower(name) + "(" + strings.Join(parts, ",") + ")"
}
return ""
}
// Package phptaintipc defines the wire protocol spoken between the CSM daemon
// and the supervised `csm phptaint-worker` child process.
//
// The worker exists because the PHP parser this analyzer depends on can enter
// an infinite loop on attacker-controlled input. Two such inputs are known and
// there is no reason to believe they are the only ones: it is a hand-written
// lexer in an unmaintained package, and the second class was found outside a
// search that had concluded the first was the only one. A loop is not a panic,
// so recover() cannot catch it, and the parser never checks context, so a
// deadline around the call cannot stop it either. The only thing that reliably
// stops it is killing the process it runs in -- which is what the supervisor on
// the other side of this protocol does, and why analysis runs out of process at
// all.
//
// The protocol is length-prefixed JSON frames over private pipes. Pipes rather
// than a Unix socket: the daemon runs under ProtectSystem=strict, and a pipe
// inherited across fork needs no socket path and so no new writable directory.
//
// Frames carry phptaint's own Report type rather than a restatement of it.
// Worker and daemon are the same binary, so the two definitions can never
// disagree, and a translation layer could only introduce drift. That matters
// more here than it looks: an unmapped status would decode as the zero value,
// which is StatusNotCandidate -- "this file is clean". A coverage gap silently
// becoming a clean result is the one failure this package must not have.
package phptaintipc
import (
"bytes"
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
"strings"
"github.com/pidginhost/csm/internal/phptaint"
)
// MaxFrameBytes caps a single request or response body. The largest legitimate
// frame is an analyze request carrying phptaint.MaxSourceBytes of source, which
// JSON base64-encodes at 4/3 expansion, plus envelope. Responses are far
// smaller: a Report's evidence is already bounded by phptaint itself.
const MaxFrameBytes = 8 << 20
// Op selects the handler on the worker side. Strings rather than iota ints so
// an unrecognised op is reported as such instead of silently matching whatever
// constant happens to share its number.
const (
OpAnalyze = "analyze"
OpPing = "ping"
)
// ErrSourceTooLarge is returned when a source buffer exceeds what phptaint
// would analyze anyway. Rejecting it here keeps the ceiling in one place and
// avoids marshalling a multi-megabyte frame only to have the worker decline it.
var ErrSourceTooLarge = errors.New("phptaintipc: source exceeds maximum analyzed size")
// Frame is the envelope. A request carries an Op and a typed Payload; a
// response leaves Op empty and carries either a Payload or an Error.
type Frame struct {
Op string `json:"op,omitempty"`
Payload json.RawMessage `json:"payload,omitempty"`
Error string `json:"error,omitempty"`
}
// AnalyzeArgs carries the bytes to analyze. The daemon has already read the
// file and run the pre-filter, so only admitted content crosses the boundary,
// and the worker is never given a path to open on its own.
type AnalyzeArgs struct {
Source []byte `json:"source"`
}
// AnalyzeResult carries the completed report.
type AnalyzeResult struct {
Report phptaint.Report `json:"report"`
}
// PingResult answers a readiness probe.
type PingResult struct {
OK bool `json:"ok"`
}
// EncodePayload validates and marshals v into a frame under op. AnalyzeArgs is
// size-checked here so an oversize source fails before it is ever encoded.
func EncodePayload(op string, v any) (Frame, error) {
if v == nil {
return Frame{Op: op}, nil
}
if err := validatePayload(v); err != nil {
return Frame{}, err
}
raw, err := json.Marshal(v)
if err != nil {
return Frame{}, fmt.Errorf("phptaintipc: marshal %s payload: %w", op, err)
}
return Frame{Op: op, Payload: raw}, nil
}
// DecodePayload unmarshals and validates a frame's payload into v.
func DecodePayload(f Frame, v any) error {
if len(f.Payload) == 0 {
return errors.New("phptaintipc: frame has no payload")
}
// Decode the security-sensitive payloads into fresh values. A missing JSON
// field must not retain a prior request's source or a prior response's clean
// status when a caller reuses storage across frames.
switch out := v.(type) {
case *AnalyzeArgs:
if out == nil {
return errors.New("phptaintipc: nil AnalyzeArgs decode target")
}
if _, err := requiredJSONField(f.Payload, "source"); err != nil {
return err
}
var decoded AnalyzeArgs
if err := json.Unmarshal(f.Payload, &decoded); err != nil {
return fmt.Errorf("phptaintipc: unmarshal payload: %w", err)
}
if err := validatePayload(decoded); err != nil {
return err
}
*out = decoded
return nil
case *AnalyzeResult:
if out == nil {
return errors.New("phptaintipc: nil AnalyzeResult decode target")
}
reportJSON, err := requiredJSONField(f.Payload, "report")
if err != nil {
return err
}
statusJSON, err := requiredJSONField(reportJSON, "status")
if err != nil {
return err
}
var status *phptaint.Status
if err := json.Unmarshal(statusJSON, &status); err != nil {
return fmt.Errorf("phptaintipc: unmarshal report status: %w", err)
}
if status == nil {
return errors.New("phptaintipc: report status is null")
}
if err := requireResultEvidenceKeys(reportJSON); err != nil {
return err
}
var decoded AnalyzeResult
if err := json.Unmarshal(f.Payload, &decoded); err != nil {
return fmt.Errorf("phptaintipc: unmarshal payload: %w", err)
}
if decoded.Report.Status != *status {
return errors.New("phptaintipc: ambiguous report status")
}
if err := validatePayload(decoded); err != nil {
return err
}
*out = decoded
return nil
}
if err := json.Unmarshal(f.Payload, v); err != nil {
return fmt.Errorf("phptaintipc: unmarshal payload: %w", err)
}
return nil
}
func validatePayload(v any) error {
switch payload := v.(type) {
case AnalyzeArgs:
return validateSourceSize(payload.Source)
case *AnalyzeArgs:
if payload == nil {
return errors.New("phptaintipc: nil AnalyzeArgs payload")
}
return validateSourceSize(payload.Source)
case AnalyzeResult:
return validateReport(payload.Report)
case *AnalyzeResult:
if payload == nil {
return errors.New("phptaintipc: nil AnalyzeResult payload")
}
return validateReport(payload.Report)
}
return nil
}
func validateSourceSize(source []byte) error {
if len(source) > phptaint.MaxSourceBytes {
return fmt.Errorf("%w (%d > %d bytes)", ErrSourceTooLarge, len(source), phptaint.MaxSourceBytes)
}
return nil
}
func validateReport(report phptaint.Report) error {
if report.Status.String() == "unknown" {
return fmt.Errorf("phptaintipc: unknown report status %d", report.Status)
}
hasEvidence := len(report.Results) != 0 || report.TotalResults != 0 ||
len(report.PrecisionLoss) != 0 || report.EvidenceTruncated
switch report.Status {
case phptaint.StatusTimeout, phptaint.StatusWorkerFailure:
// These describe what the parent did to, or could not do with, the
// worker. A process that returned a reply cannot report either one.
return errors.New("phptaintipc: worker reply carries a supervisor-only status")
case phptaint.StatusNotCandidate:
if hasEvidence || report.Reason != "" {
return errors.New("phptaintipc: not-candidate report carries analysis data")
}
case phptaint.StatusAnalyzed:
if report.Reason != "" {
return errors.New("phptaintipc: analyzed report carries an error reason")
}
if report.TotalResults < 0 {
return errors.New("phptaintipc: analyzed report has inconsistent result counts")
}
expectedResults := report.TotalResults
if expectedResults > phptaint.MaxEvidenceResults {
expectedResults = phptaint.MaxEvidenceResults
}
if len(report.Results) != expectedResults {
return errors.New("phptaintipc: analyzed report has inconsistent result counts")
}
if report.TotalResults > len(report.Results) && !report.EvidenceTruncated {
return errors.New("phptaintipc: analyzed report omitted evidence without marking truncation")
}
for _, result := range report.Results {
if result.Confidence.String() == "unknown" {
return errors.New("phptaintipc: analyzed report has unknown confidence")
}
if !result.Basis.Valid() {
return errors.New("phptaintipc: analyzed report has missing or unknown basis")
}
if result.ResolutionOffset < -1 {
return errors.New("phptaintipc: analyzed report has an invalid resolution offset")
}
}
default:
if hasEvidence {
return errors.New("phptaintipc: incomplete report carries analysis data")
}
if report.Reason == "" {
return errors.New("phptaintipc: incomplete report has no reason")
}
}
if len(report.Reason) > phptaint.MaxReasonBytes {
return fmt.Errorf("phptaintipc: report reason is %d bytes, exceeds cap %d", len(report.Reason), phptaint.MaxReasonBytes)
}
return nil
}
// ValidateReportForSource applies the checks that need the submitted source:
// a resolution offset must point inside it. The parent calls it after
// DecodePayload, which has already applied every source-independent check.
func ValidateReportForSource(report phptaint.Report, sourceLen int) error {
if err := validateReport(report); err != nil {
return err
}
for _, result := range report.Results {
if result.ResolutionOffset >= sourceLen {
return fmt.Errorf("phptaintipc: resolution offset %d outside %d-byte source", result.ResolutionOffset, sourceLen)
}
}
return nil
}
// requireResultEvidenceKeys checks every result for its basis and resolution
// offset keys. A missing offset would decode to 0, a real position, so a reply
// from a worker that predates either field must fail rather than pose as
// evidence resolved at the first byte.
func requireResultEvidenceKeys(reportJSON json.RawMessage) error {
resultsJSON, ok, err := lookupJSONField(reportJSON, "results")
if err != nil {
return err
}
if !ok {
return nil
}
var results []json.RawMessage
if err := json.Unmarshal(resultsJSON, &results); err != nil {
return fmt.Errorf("phptaintipc: unmarshal report results: %w", err)
}
for _, result := range results {
if _, err := requiredJSONField(result, "basis"); err != nil {
return err
}
offsetJSON, err := requiredJSONField(result, "resolutionoffset")
if err != nil {
return err
}
var offset *int
if err := json.Unmarshal(offsetJSON, &offset); err != nil {
return fmt.Errorf("phptaintipc: unmarshal resolution offset: %w", err)
}
if offset == nil {
return errors.New("phptaintipc: result resolution offset is null")
}
}
return nil
}
func requiredJSONField(raw []byte, name string) (json.RawMessage, error) {
found, ok, err := lookupJSONField(raw, name)
if err != nil {
return nil, err
}
if !ok {
return nil, fmt.Errorf("phptaintipc: payload has no %s field", name)
}
return found, nil
}
// lookupJSONField finds name in a JSON object the way encoding/json matches
// struct fields, case-insensitively, and rejects more than one spelling so a
// reply cannot carry two values for one field.
func lookupJSONField(raw []byte, name string) (json.RawMessage, bool, error) {
// A map discards repeated keys, but struct decoding can merge their
// values. Inspect every occurrence so validation sees what decoding sees.
if !json.Valid(raw) {
return nil, false, fmt.Errorf("phptaintipc: invalid JSON inspecting %s field", name)
}
decoder := json.NewDecoder(bytes.NewReader(raw))
start, err := decoder.Token()
if err != nil || start != json.Delim('{') {
return nil, false, fmt.Errorf("phptaintipc: expected object inspecting %s field", name)
}
var found json.RawMessage
ok := false
for decoder.More() {
field, err := decoder.Token()
if err != nil {
return nil, false, fmt.Errorf("phptaintipc: inspect %s field: %w", name, err)
}
var value json.RawMessage
if err := decoder.Decode(&value); err != nil {
return nil, false, fmt.Errorf("phptaintipc: inspect %s value: %w", name, err)
}
if !strings.EqualFold(field.(string), name) {
continue
}
if ok {
return nil, false, fmt.Errorf("phptaintipc: payload has ambiguous %s fields", name)
}
found, ok = value, true
}
return found, ok, nil
}
// WriteFrame writes one length-prefixed frame.
func WriteFrame(w io.Writer, f Frame) error {
body, err := json.Marshal(f)
if err != nil {
return fmt.Errorf("phptaintipc: marshal frame: %w", err)
}
if len(body) > MaxFrameBytes {
return fmt.Errorf("phptaintipc: frame body %d bytes exceeds cap %d", len(body), MaxFrameBytes)
}
var hdr [4]byte
// #nosec G115 -- len(body) is bounded above by MaxFrameBytes, which fits in uint32.
binary.BigEndian.PutUint32(hdr[:], uint32(len(body)))
if err := writeFull(w, hdr[:]); err != nil {
return fmt.Errorf("phptaintipc: write header: %w", err)
}
if err := writeFull(w, body); err != nil {
return fmt.Errorf("phptaintipc: write body: %w", err)
}
return nil
}
func writeFull(w io.Writer, p []byte) error {
for len(p) != 0 {
n, err := w.Write(p)
if err != nil {
return err
}
if n <= 0 || n > len(p) {
return io.ErrShortWrite
}
p = p[n:]
}
return nil
}
// ReadFrame reads one length-prefixed frame.
//
// The length prefix comes from the peer, so it is checked against the cap
// BEFORE any buffer is allocated. Reading the declared size first and
// validating afterwards would let a four-byte write ask for gigabytes.
func ReadFrame(r io.Reader) (Frame, error) {
var hdr [4]byte
if _, err := io.ReadFull(r, hdr[:]); err != nil {
return Frame{}, fmt.Errorf("phptaintipc: read header: %w", err)
}
size := binary.BigEndian.Uint32(hdr[:])
if size > MaxFrameBytes {
return Frame{}, fmt.Errorf("phptaintipc: declared frame size %d exceeds cap %d", size, MaxFrameBytes)
}
body := make([]byte, size)
if _, err := io.ReadFull(r, body); err != nil {
return Frame{}, fmt.Errorf("phptaintipc: read body: %w", err)
}
var f Frame
if err := json.Unmarshal(body, &f); err != nil {
return Frame{}, fmt.Errorf("phptaintipc: unmarshal frame: %w", err)
}
return f, nil
}
package phptaintworker
import (
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type requestQueue struct {
mu sync.Mutex
pending map[*requestWork]struct{}
losses *queuehealth.Tracker
}
type requestWork struct {
queue *requestQueue
queued, started, phaseDeadline, rpcDeadline time.Time
callerDone, rpcOutstanding, awaitingRPC, failed bool
}
func newRequestQueue() *requestQueue {
return &requestQueue{pending: make(map[*requestWork]struct{}), losses: queuehealth.New(0, time.Minute)}
}
func (q *requestQueue) begin() *requestWork {
w := &requestWork{queue: q, queued: time.Now()}
q.mu.Lock()
q.pending[w] = struct{}{}
q.mu.Unlock()
return w
}
func (w *requestWork) start() {
w.queue.mu.Lock()
w.started = time.Now()
w.phaseDeadline = w.started.Add(time.Minute)
w.queue.mu.Unlock()
}
func (w *requestWork) beginRPC(deadline time.Time) {
w.queue.mu.Lock()
w.rpcOutstanding = true
w.awaitingRPC = true
w.rpcDeadline, w.phaseDeadline = deadline, deadline
w.queue.mu.Unlock()
}
func (w *requestWork) progress() {
w.queue.mu.Lock()
w.awaitingRPC = false
w.phaseDeadline = time.Now().Add(time.Minute)
w.queue.mu.Unlock()
}
func (w *requestWork) finishRPC(completed bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
w.rpcOutstanding = false
if w.awaitingRPC {
w.awaitingRPC = false
w.phaseDeadline = time.Now().Add(time.Minute)
}
if !completed {
w.failLocked()
}
w.releaseLocked()
}
func (w *requestWork) finishCaller(completed bool) {
q := w.queue
q.mu.Lock()
defer q.mu.Unlock()
w.callerDone = true
if !completed {
w.failLocked()
}
w.releaseLocked()
}
func (w *requestWork) releaseLocked() {
// Closing the pipes normally releases the RPC goroutine. A caller that
// already returned must not hide an operation still blocked in I/O.
if w.callerDone && !w.rpcOutstanding {
delete(w.queue.pending, w)
}
}
func (w *requestWork) fail() {
w.queue.mu.Lock()
w.failLocked()
w.queue.mu.Unlock()
}
func (w *requestWork) failLocked() {
if !w.failed {
w.failed = true
w.queue.losses.Lose(time.Now(), 1)
}
}
// QueueStatuses does not take the supervisor lock or inspect child processes.
func (s *Supervisor) QueueStatuses(now time.Time) map[string]queuehealth.Status {
q := s.health
q.mu.Lock()
defer q.mu.Unlock()
status := q.losses.Snapshot(now)
status.CapacityUnavailable = true
var waitingLate, runningLate bool
for w := range q.pending {
if w.started.IsZero() {
status.Depth++
status.LagSeconds = max(status.LagSeconds, now.Sub(w.queued).Seconds())
waitingLate = waitingLate || now.Sub(w.queued) >= time.Minute
} else {
status.InFlight++
status.ProcessingSeconds = max(status.ProcessingSeconds, now.Sub(w.started).Seconds())
runningLate = runningLate || (!w.callerDone && !now.Before(w.phaseDeadline)) || (w.rpcOutstanding && !now.Before(w.rpcDeadline))
}
}
switch {
case waitingLate:
status.Reason = "backlog_lag"
case runningLate:
status.Reason = "processing_lag"
}
if status.Reason != "" {
status.Status = "degraded"
}
return map[string]queuehealth.Status{"requests": status}
}
package phptaintworker
import (
"context"
"errors"
"fmt"
"io"
"os/exec"
"sync"
"time"
"github.com/pidginhost/csm/internal/phptaint"
"github.com/pidginhost/csm/internal/phptaintipc"
)
// ConsecutiveFailureLimit is how many failed workers in a row are tolerated
// before the supervisor pauses spawning replacements for one cooldown.
//
// It is a constant rather than a setting because nothing an operator can
// observe would tell them a better value: the right number depends on how the
// parser fails, not on the host. It is also harmful in both directions -- too
// low and one transient failure blinds the rest of a scan, too high and a
// crafted account gets the exec storm the limit exists to prevent. When it
// trips, the coverage gap says so, so the condition stays visible even though
// the threshold is not tunable.
const ConsecutiveFailureLimit = 3
// reapTimeout bounds how long killLocked waits for a SIGKILLed child while
// holding the supervisor lock.
const reapTimeout = 10 * time.Second
// breakerCooldown is how long the supervisor refuses to spawn replacements
// after ConsecutiveFailureLimit failures in a row, before it tries once more.
//
// The breaker bounds an exec storm; it must not switch the detector off. With
// no way back it would be a kill switch any tenant could throw with a handful
// of crafted files, disabling PHP analysis across a shared host for the life of
// the daemon -- worse than not deploying the detector, because an operator
// would believe it was running. One trial per cooldown keeps the storm bounded
// (a hostile account costs one timeout per cooldown, not one per file) while
// letting a host recover without an operator having to notice and restart.
//
// The cost of each retry lands on the shared deep-scan budget, which the YARA
// and JavaScript consumers walk too, so a longer cooldown buys their coverage
// back on a host that stays broken.
const breakerCooldown = 60 * time.Second
// SupervisorConfig describes how to start the worker process.
type SupervisorConfig struct {
// Command and Args start the worker. In production this is the CSM binary
// re-executing itself as a subcommand.
Command string
Args []string
Env []string
// Timeout bounds a single analysis. It must exceed the slowest legitimate
// file by a wide margin: every expiry costs a killed process, so a value
// tuned too tightly converts slow-but-fine files into coverage gaps.
Timeout time.Duration
// Log is optional.
Log func(string, ...any)
}
// Supervisor runs one worker process at a time and guarantees that a request
// which does not return kills the process it was running in.
//
// The guarantee is the point. A deadline on its own is not containment: the
// call returns to the caller while the child keeps spinning on a core, and the
// next request starts another one. That is the failure mode of a timeout
// without a kill, and it is why this type does not reuse the YARA supervisor,
// which restarts a child on exit and never kills one that simply stops
// answering.
//
// A worker that misses its deadline is poisoned: killed, reaped, and never
// used again. Replacement is lazy, so a scan with no further candidates pays
// nothing for the last file's failure.
type Supervisor struct {
cfg SupervisorConfig
health *requestQueue
mu sync.Mutex
child *child
lastPID int
spawns int
consecutive int
// openedAt is when the breaker last opened; the zero value means closed.
openedAt time.Time
mode string
stopped bool
// now is a seam so the cooldown can be exercised without sleeping.
now func() time.Time
}
type child struct {
cmd *exec.Cmd
stdin io.WriteCloser
stdout io.ReadCloser
done chan struct{}
}
// NewSupervisor validates config. It does not start a worker: the first
// analysis does, so a scan that never admits a candidate never forks.
func NewSupervisor(cfg SupervisorConfig) (*Supervisor, error) {
if cfg.Command == "" {
return nil, errors.New("phptaintworker: empty command")
}
if cfg.Timeout <= 0 {
return nil, errors.New("phptaintworker: timeout must be positive")
}
return &Supervisor{cfg: cfg, now: time.Now, health: newRequestQueue()}, nil
}
// SetMode is a test seam for the helper child, which selects its behaviour
// from the environment. Production callers never use it.
func (s *Supervisor) SetMode(mode string) {
s.mu.Lock()
defer s.mu.Unlock()
s.mode = mode
}
// LastChildPID reports the pid of the most recently started worker.
func (s *Supervisor) LastChildPID() int {
s.mu.Lock()
defer s.mu.Unlock()
return s.lastPID
}
// SpawnCount reports how many workers have been started.
func (s *Supervisor) SpawnCount() int {
s.mu.Lock()
defer s.mu.Unlock()
return s.spawns
}
// Stop shuts down the current worker.
func (s *Supervisor) Stop() error {
s.mu.Lock()
defer s.mu.Unlock()
s.stopped = true
s.killLocked()
return nil
}
// Analyze runs one analysis in the worker. It never returns a status a caller
// could read as clean unless the worker actually produced one, or unless the
// pre-filter the worker applies before parsing already rejects the content.
func (s *Supervisor) Analyze(ctx context.Context, src []byte) phptaint.Report {
// The deep scan hands every file it reads to this analyzer, and the
// worker serves one request at a time. Content the pre-filter rejects
// needs no parser, so answering it here keeps images and plain text from
// queuing behind real analyses or forking a worker. Size and
// cancellation are the answers the worker path gives first, so they
// still take precedence.
if len(src) <= phptaint.MaxSourceBytes && ctx.Err() == nil && !phptaint.IsCandidate(src) {
// Cancellation may arrive during the byte scan. The IPC path
// observes it while waiting for a reply; a local answer must too.
if err := ctx.Err(); err != nil {
return gap(phptaint.StatusCanceled, err.Error())
}
return phptaint.Report{Status: phptaint.StatusNotCandidate}
}
work := s.health.begin()
completed := false
defer func() { work.finishCaller(completed) }()
s.mu.Lock()
defer s.mu.Unlock()
work.start()
report := s.analyzeLocked(ctx, src, work)
completed = true
return report
}
func (s *Supervisor) analyzeLocked(ctx context.Context, src []byte, work *requestWork) phptaint.Report {
if s.stopped {
return gap(phptaint.StatusWorkerFailure, "supervisor stopped")
}
if s.breakerOpenLocked() {
work.fail()
return gap(phptaint.StatusWorkerFailure, fmt.Sprintf(
"worker failed %d times in a row; not analyzed, next attempt after %s",
s.consecutive, breakerCooldown))
}
req, err := phptaintipc.EncodePayload(phptaintipc.OpAnalyze, phptaintipc.AnalyzeArgs{Source: src})
if err != nil {
// An oversize source is the caller's own ceiling, not a worker fault,
// so it must not count toward the breaker.
return gap(phptaint.StatusOversize, err.Error())
}
if ctxErr := ctx.Err(); ctxErr != nil {
return gap(phptaint.StatusCanceled, ctxErr.Error())
}
if startErr := s.ensureChildLocked(); startErr != nil {
work.fail()
s.recordFailureLocked()
return gap(phptaint.StatusWorkerFailure, startErr.Error())
}
report, err := s.roundTripLocked(ctx, req, len(src), work)
if err != nil {
if !errors.Is(err, context.Canceled) && !errors.Is(err, context.DeadlineExceeded) {
work.fail()
}
// The child is not trusted after any failure: it may be mid-parse and
// spinning, or it may have left a partial frame in the pipe that would
// desynchronise every later request.
s.killLocked()
if errors.Is(err, errDeadline) {
s.recordFailureLocked()
return gap(phptaint.StatusTimeout, fmt.Sprintf(
"analysis exceeded %s; worker killed", s.cfg.Timeout))
}
if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
// The caller withdrew, so this is not a worker fault and must not
// count toward the breaker. It did consume the trial, though: the
// cooldown has to be re-armed or a caller that cancels can spend
// one spawn per file while the breaker is nominally open.
s.rearmBreakerLocked()
return gap(phptaint.StatusCanceled, err.Error())
}
s.recordFailureLocked()
return gap(phptaint.StatusWorkerFailure, err.Error())
}
s.consecutive = 0
if report.Status == phptaint.StatusPanic {
work.fail()
}
return report
}
var errDeadline = errors.New("phptaintworker: analysis deadline exceeded")
// recordFailureLocked counts a worker failure and, once the run of failures
// reaches the limit, restarts the cooldown. Stamping on every failure at or
// past the limit -- a failed trial request included -- is what makes a
// persistently broken host retry once per cooldown instead of once per file.
func (s *Supervisor) recordFailureLocked() {
s.consecutive++
s.rearmBreakerLocked()
}
// rearmBreakerLocked restarts the cooldown when the breaker is at or past its
// limit. openedAt is never cleared on success: breakerOpenLocked reads it only
// once consecutive has reached the limit, and every path that reaches the limit
// stamps it first, so a stale value cannot be observed.
func (s *Supervisor) rearmBreakerLocked() {
if s.consecutive >= ConsecutiveFailureLimit {
s.openedAt = s.now()
}
}
// breakerOpenLocked reports whether this request must be refused without
// spawning. Once the cooldown elapses the breaker goes half-open: the request
// is let through as a trial, which either resets the counter on success or
// restarts the cooldown on failure.
func (s *Supervisor) breakerOpenLocked() bool {
if s.consecutive < ConsecutiveFailureLimit {
return false
}
return s.now().Sub(s.openedAt) < breakerCooldown
}
// roundTripLocked writes one request and waits for its reply, bounded by the
// configured timeout. sourceLen is the size of the submitted source, which
// bounds every offset the reply may carry. The pipe round trip runs on its own
// goroutine because the child may stop before reading the complete request or
// never answer. The goroutine ends when killLocked closes the pipes.
func (s *Supervisor) roundTripLocked(ctx context.Context, req phptaintipc.Frame, sourceLen int, work *requestWork) (phptaint.Report, error) {
c := s.child
type result struct {
frame phptaintipc.Frame
err error
}
deadline := time.Now().Add(s.cfg.Timeout)
timer := time.NewTimer(s.cfg.Timeout)
defer timer.Stop()
if parent, ok := ctx.Deadline(); ok && parent.Before(deadline) {
deadline = parent
}
work.beginRPC(deadline)
defer work.progress()
done := make(chan result, 1)
go func() {
completed := false
defer func() { work.finishRPC(completed) }()
if err := phptaintipc.WriteFrame(c.stdin, req); err != nil {
done <- result{err: fmt.Errorf("phptaintworker: write request: %w", err)}
completed = true
return
}
f, err := phptaintipc.ReadFrame(c.stdout)
done <- result{frame: f, err: err}
completed = true
}()
select {
case <-timer.C:
return phptaint.Report{}, errDeadline
case <-ctx.Done():
return phptaint.Report{}, ctx.Err()
case res := <-done:
work.progress()
if res.err != nil {
return phptaint.Report{}, fmt.Errorf("phptaintworker: read reply: %w", res.err)
}
if res.frame.Error != "" {
return phptaint.Report{}, fmt.Errorf("phptaintworker: worker: %s", res.frame.Error)
}
if res.frame.Op != "" {
return phptaint.Report{}, fmt.Errorf("phptaintworker: response carries op %q", res.frame.Op)
}
var out phptaintipc.AnalyzeResult
if err := phptaintipc.DecodePayload(res.frame, &out); err != nil {
return phptaint.Report{}, err
}
if err := phptaintipc.ValidateReportForSource(out.Report, sourceLen); err != nil {
return phptaint.Report{}, err
}
return out.Report, nil
}
}
func (s *Supervisor) ensureChildLocked() error {
if s.child != nil {
return nil
}
env := s.cfg.Env
if s.mode != "" {
env = append(append([]string(nil), env...), "CSM_HELPER_WORKER="+s.mode)
}
cmd := exec.Command(s.cfg.Command, s.cfg.Args...) // #nosec G204 -- command comes from the daemon's own config, not from scanned content.
cmd.Env = env
stdin, err := cmd.StdinPipe()
if err != nil {
return fmt.Errorf("phptaintworker: stdin pipe: %w", err)
}
stdout, err := cmd.StdoutPipe()
if err != nil {
return fmt.Errorf("phptaintworker: stdout pipe: %w", err)
}
if err := cmd.Start(); err != nil {
return fmt.Errorf("phptaintworker: start worker: %w", err)
}
c := &child{cmd: cmd, stdin: stdin, stdout: stdout, done: make(chan struct{})}
go func() {
_ = cmd.Wait()
close(c.done)
}()
s.child = c
s.spawns++
s.lastPID = cmd.Process.Pid
return nil
}
// killLocked terminates the current worker and waits for it to be reaped.
//
// SIGKILL, not SIGTERM: the process this exists to stop is spinning inside a
// parser loop and never returns to a point where a catchable signal would be
// handled. Closing the pipes around the kill unblocks the round-trip goroutine
// so it cannot outlive the child.
func (s *Supervisor) killLocked() {
c := s.child
if c == nil {
return
}
s.child = nil
_ = c.stdin.Close()
if c.cmd.Process != nil {
_ = c.cmd.Process.Kill()
}
_ = c.stdout.Close()
// Bounded, because this runs while the supervisor lock is held: an
// unbounded wait would turn one unreapable child into a frozen analyzer
// for every later file. SIGKILL cannot be caught or ignored, so the only
// way to reach the timeout is a process wedged in uninterruptible sleep.
// Giving up on the reap leaks one process; blocking here would stop the
// scan, and the scan is what protects the host.
select {
case <-c.done:
case <-time.After(reapTimeout):
if s.cfg.Log != nil {
s.cfg.Log("phptaint worker %d did not exit after SIGKILL", c.cmd.Process.Pid)
}
return
}
if s.cfg.Log != nil {
s.cfg.Log("phptaint worker %d killed", c.cmd.Process.Pid)
}
}
func gap(status phptaint.Status, reason string) phptaint.Report {
return phptaint.CoverageGap(status, reason)
}
// Package phptaintworker is the child side of PHP taint analysis.
//
// The analyzer's parser can enter an infinite loop on attacker-controlled
// input. Two such inputs are known, both found by accident rather than by
// audit, in a hand-written lexer that has been unmaintained since 2022. A loop
// is not a panic, so recover() cannot catch it, and the parser never checks
// context, so no deadline inside this process can stop it. Running the analysis
// here means the parent can stop it the only way that works: by killing this
// process.
//
// Consequences for anything written in this package: a request that never
// returns is an expected outcome, not a bug to defend against locally, and this
// side must stay simple enough that the parent's timeout is the sole liveness
// mechanism. Do not add an internal watchdog -- it would only ever fire on
// inputs the parent already handles, while giving a false impression that this
// process can rescue itself.
package phptaintworker
import (
"context"
"errors"
"fmt"
"io"
"os"
"os/signal"
"github.com/pidginhost/csm/internal/phptaint"
"github.com/pidginhost/csm/internal/phptaintipc"
)
// Serve reads request frames from r and writes response frames to w. It checks
// ctx between frames, not from inside the parser; the parent process remains
// the sole liveness boundary. It is the whole worker: one request at a time, no
// concurrency, no state carried between requests.
//
// A valid frame with a bad operation or payload is answered with an error and
// the loop continues. Dropping the process for that peer mistake would make the
// parent replace its worker repeatedly. Broken framing ends the loop because a
// trustworthy boundary for the next request no longer exists.
func Serve(ctx context.Context, r io.Reader, w io.Writer) error {
for {
if err := ctx.Err(); err != nil {
return err
}
req, err := phptaintipc.ReadFrame(r)
if err != nil {
if errors.Is(err, io.EOF) {
return nil
}
return err
}
resp := handle(ctx, req)
if err := phptaintipc.WriteFrame(w, resp); err != nil {
return err
}
}
}
func handle(ctx context.Context, req phptaintipc.Frame) phptaintipc.Frame {
switch req.Op {
case phptaintipc.OpPing:
return reply(phptaintipc.OpPing, phptaintipc.PingResult{OK: true})
case phptaintipc.OpAnalyze:
var args phptaintipc.AnalyzeArgs
if err := phptaintipc.DecodePayload(req, &args); err != nil {
return phptaintipc.Frame{Error: err.Error()}
}
// Analyze already contains its own panic boundary and reports one as a
// coverage gap, so there is nothing to recover here.
report := phptaint.Analyze(ctx, args.Source)
return reply(phptaintipc.OpAnalyze, phptaintipc.AnalyzeResult{Report: report})
default:
return phptaintipc.Frame{Error: fmt.Sprintf("phptaintworker: unknown op %q", req.Op)}
}
}
func reply(op string, v any) phptaintipc.Frame {
frame, err := phptaintipc.EncodePayload(op, v)
if err != nil {
return phptaintipc.Frame{Error: err.Error()}
}
// A response carries no op: the parent matches replies by order on a
// single-flight connection, and echoing the op would invite a reader to
// match on it instead.
frame.Op = ""
return frame
}
// signalIgnore makes a signal a no-op for this process. The supervisor tests
// use it to build a child that a catchable signal cannot stop, which is what a
// parser loop behaves like.
func signalIgnore(sig os.Signal) {
signal.Ignore(sig)
}
package platform
import (
"errors"
"fmt"
"os"
"path/filepath"
"slices"
"strings"
"github.com/pidginhost/csm/internal/safepath"
)
// ValidateAccountRootPattern rejects paths that cannot define a confined
// content tree. Missing directories are allowed for later provisioning.
func ValidateAccountRootPattern(pattern string) error {
if !filepath.IsAbs(pattern) || filepath.Clean(pattern) != pattern {
return fmt.Errorf("account root must be an absolute, clean path: %q", pattern)
}
for _, char := range pattern {
if char < 0x20 || char == 0x7f || char == '\\' {
return fmt.Errorf("account root contains an unsupported character: %q", pattern)
}
}
parts := strings.Split(strings.TrimPrefix(pattern, "/"), "/")
if len(parts) < 2 || strings.IndexAny(pattern, "*?[") == 1 {
return fmt.Errorf("account root must select a content directory below a fixed top-level directory: %q", pattern)
}
if _, err := filepath.Match(pattern, ""); err != nil {
return fmt.Errorf("invalid account root pattern %q: %w", pattern, err)
}
return nil
}
// ResolveAccountRoots expands configured content directories without following
// symlinks. Unsafe directories are excluded and reported alongside valid roots,
// so a broken tenant tree does not block remediation of other accounts.
func ResolveAccountRoots(patterns []string) ([]string, error) {
var roots []string
var problems []error
for _, pattern := range patterns {
if err := ValidateAccountRootPattern(pattern); err != nil {
return nil, err
}
matches, err := filepath.Glob(pattern)
if err != nil {
return nil, err
}
for _, path := range matches {
info, err := os.Lstat(path)
if os.IsNotExist(err) {
continue
}
if err != nil {
problems = append(problems, err)
continue
}
if !info.IsDir() && info.Mode()&os.ModeSymlink == 0 {
continue
}
if err := validateAccountRootDirectory(path); err != nil {
problems = append(problems, err)
continue
}
roots = append(roots, path)
}
}
slices.Sort(roots)
return slices.Compact(roots), errors.Join(problems...)
}
func validateAccountRootDirectory(path string) error {
dir, err := safepath.OpenDirNoFollow(path)
if err != nil {
return fmt.Errorf("open account root %s without symlinks: %w", path, err)
}
return dir.Close()
}
package platform
import (
"context"
"os"
"os/exec"
"time"
)
// MTAIdents lists local users and process basenames belonging to the
// host's Mail Transfer Agent stack. Direct SMTP egress detection uses
// this allowlist to skip legitimate local MTA traffic instead of
// path-allowlisting a directory.
type MTAIdents struct {
Users []string
Processes []string
}
// IsMTAUser reports whether name is one of the known MTA usernames.
// Match is exact and case-sensitive (Linux usernames are).
func (m MTAIdents) IsMTAUser(name string) bool {
for _, u := range m.Users {
if u == name {
return true
}
}
return false
}
// IsMTAProcess reports whether basename is one of the known MTA process
// basenames. Exact match; the caller passes comm or basename(exe), not
// a full path.
func (m MTAIdents) IsMTAProcess(basename string) bool {
for _, p := range m.Processes {
if p == basename {
return true
}
}
return false
}
// LocalMTAIdentities returns the MTA users and process basenames that
// should be considered legitimate on the detected platform. cPanel
// hosts get exim variants; non-cPanel hosts get the postfix/dovecot
// baseline.
func LocalMTAIdentities(info Info) MTAIdents {
users := []string{
"mail",
"mailnull",
"postfix",
"dovecot",
"dovenull",
"mailman",
}
processes := []string{
"postfix",
"smtpd",
"smtp",
"qmgr",
"pickup",
"cleanup",
"local",
"dovecot",
"imap-login",
"pop3-login",
"lmtp",
}
if info.IsCPanel() {
users = append(users, "exim")
processes = append(processes, "exim", "exim4")
}
return MTAIdents{Users: users, Processes: processes}
}
// MTAKind identifies the host's Mail Transfer Agent.
type MTAKind string
const (
MTAUnknown MTAKind = ""
MTAExim MTAKind = "exim"
MTAPostfix MTAKind = "postfix"
)
var (
mtaLookPath = exec.LookPath
mtaStat = os.Stat
mtaServiceActive = systemdMTAServiceActive
)
// DetectMTA reports the delivery agent rather than whichever package happens
// to leave a binary on disk. cPanel is authoritative because it runs Exim
// while retaining Postfix binaries; elsewhere an active service wins.
func DetectMTA() MTAKind {
if mtaPathExists("/usr/local/cpanel/version") {
return MTAExim
}
eximActive := mtaServiceActive("exim") || mtaServiceActive("exim4")
postfixActive := mtaServiceActive("postfix")
// When both are active, keep Exim-specific weaknesses visible instead of
// dropping the checks behind an ambiguous result.
if eximActive {
return MTAExim
}
if postfixActive {
return MTAPostfix
}
eximInstalled := mtaInstalled([]string{"exim", "exim4"}, []string{
"/usr/sbin/exim",
"/usr/sbin/exim4",
})
postfixInstalled := mtaInstalled([]string{"postfix"}, []string{
"/usr/sbin/postfix",
"/usr/libexec/postfix/master",
"/usr/lib/postfix/sbin/master",
})
if eximInstalled == postfixInstalled {
return MTAUnknown
}
if eximInstalled {
return MTAExim
}
return MTAPostfix
}
func mtaPathExists(path string) bool {
_, err := mtaStat(path)
return err == nil
}
func systemdMTAServiceActive(unit string) bool {
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
// #nosec G204 -- callers pass only the literal MTA unit names above.
return exec.CommandContext(ctx, "systemctl", "is-active", "--quiet", unit).Run() == nil
}
// mtaInstalled reports whether any of the named binaries is on PATH or at
// one of the absolute locations. The daemon runs with a minimal PATH, so
// the absolute probes carry most installs.
func mtaInstalled(binaries, paths []string) bool {
for _, b := range binaries {
if _, err := mtaLookPath(b); err == nil {
return true
}
}
for _, p := range paths {
info, err := mtaStat(p)
if err == nil && info.Mode().IsRegular() && info.Mode().Perm()&0o111 != 0 {
return true
}
}
return false
}
// Package platform detects the host OS, control panel, and web server so
// CSM checks can pick the right config/log paths instead of hardcoding
// cPanel+Apache layouts.
package platform
import (
"bufio"
"os"
"os/exec"
"slices"
"strings"
"sync"
"sync/atomic"
)
type OSFamily string
const (
OSUnknown OSFamily = ""
OSUbuntu OSFamily = "ubuntu"
OSDebian OSFamily = "debian"
OSAlma OSFamily = "almalinux"
OSRocky OSFamily = "rocky"
OSCentOS OSFamily = "centos"
OSRHEL OSFamily = "rhel"
OSCloudLinux OSFamily = "cloudlinux"
)
type Panel string
const (
PanelNone Panel = ""
PanelCPanel Panel = "cpanel"
PanelPlesk Panel = "plesk"
PanelDA Panel = "directadmin"
)
type WebServer string
const (
WSNone WebServer = ""
WSApache WebServer = "apache"
WSNginx WebServer = "nginx"
WSLiteSpeed WebServer = "litespeed"
)
// Info holds everything a check needs to locate web server resources.
type Info struct {
OS OSFamily
OSVersion string
Panel Panel
WebServer WebServer
// Config locations for the detected web server.
ApacheConfigDir string // e.g. /etc/apache2 or /etc/httpd
NginxConfigDir string // e.g. /etc/nginx
// Candidate log files. Populated based on detected web server + OS.
AccessLogPaths []string
ErrorLogPaths []string
ModSecAuditLogPaths []string
// DomlogGlobs is the list of glob patterns used to enumerate
// per-vhost access logs. Populated per panel + OS in populatePaths.
DomlogGlobs []string
// Binary paths useful for reload/control.
ApacheBinary string
NginxBinary string
LiteSpeedBinary string
}
// IsCPanel is a convenience for checks that still need to gate cPanel-only
// behavior (WHM API calls, /home/*/public_html enumeration, exim log
// tailing, etc.) without re-detecting each time.
func (i Info) IsCPanel() bool { return i.Panel == PanelCPanel }
// ApacheCompatibleConfigDir returns the Apache-style config root consumed by
// the active web server. cPanel LiteSpeed reads the EA4 Apache tree even when
// Apache itself is not the detected server, and older detection snapshots may
// therefore have no ApacheConfigDir populated.
func (i Info) ApacheCompatibleConfigDir() string {
if i.ApacheConfigDir != "" {
return i.ApacheConfigDir
}
if i.IsCPanel() {
return "/etc/apache2"
}
return ""
}
// IsRHELFamily reports whether the OS uses rpm/dnf and /etc/httpd style paths.
func (i Info) IsRHELFamily() bool {
switch i.OS {
case OSAlma, OSRocky, OSCentOS, OSRHEL, OSCloudLinux:
return true
}
return false
}
// IsDebianFamily reports whether the OS uses dpkg/apt and /etc/apache2 style paths.
func (i Info) IsDebianFamily() bool {
return i.OS == OSUbuntu || i.OS == OSDebian
}
// CronSpoolDir returns the directory holding per-user crontabs: cronie
// (RHEL family, cPanel) writes them directly under /var/spool/cron, while
// Debian's cron keeps them in /var/spool/cron/crontabs.
func (i Info) CronSpoolDir() string {
if i.IsDebianFamily() {
return "/var/spool/cron/crontabs"
}
return "/var/spool/cron"
}
// AccountHomeRoots returns the directories whose children are hosting
// accounts: /home on cPanel, DirectAdmin and plain hosts, /var/www/vhosts
// on Plesk. Every account-scoped check, remediation and re-check resolves
// "<root>/<account>" through this instead of assuming /home.
func (i Info) AccountHomeRoots() []string {
if i.Panel == PanelPlesk {
return []string{"/var/www/vhosts"}
}
return []string{"/home"}
}
// PanelToolRoots returns the directories holding the control panel's own
// root-run tooling. cPanel keeps its maintenance scripts and binaries under
// /usr/local/cpanel; a panel-less host has no such tree. Provenance checks
// use this to tell the panel's own nightly work from a process that merely
// named itself after it, so the path has to come from detection rather than
// being written into the detector.
func (i Info) PanelToolRoots() []string {
if i.Panel == PanelCPanel {
return []string{"/usr/local/cpanel"}
}
return nil
}
// WebServerUsers returns the account(s) the web server serves requests as
// on this platform: cPanel's nobody, DirectAdmin's apache (plus nobody for
// its suEXEC fallback), Plesk's distribution web user, and for a panel-less
// host the distribution user of the detected server. Checks that audit
// "the web server's" crontab or group ownership must use this instead of
// assuming cPanel's nobody.
func (i Info) WebServerUsers() []string {
switch i.Panel {
case PanelCPanel:
return []string{"nobody"}
case PanelDA:
return []string{"apache", "nobody"}
}
if i.WebServer == WSLiteSpeed {
return []string{"nobody"}
}
if i.IsDebianFamily() {
return []string{"www-data"}
}
if i.WebServer == WSNginx {
return []string{"nginx"}
}
return []string{"apache"}
}
// MailLogPath returns the platform-default mail log file. Empty string
// means "no file source available" (operator must use journal).
func (i Info) MailLogPath() string {
if i.IsDebianFamily() {
return "/var/log/mail.log"
}
return "/var/log/maillog"
}
// AuthLogPath returns the platform-default SSH/PAM auth log file.
// Debian-family uses /var/log/auth.log; RHEL-family uses /var/log/secure.
func (i Info) AuthLogPath() string {
if i.IsDebianFamily() {
return "/var/log/auth.log"
}
return "/var/log/secure"
}
// Overrides lets the operator override auto-detected values from csm.yaml.
// Any field left blank or nil falls back to the auto-detected value.
//
// Panel and WebServer use pointer types so callers can distinguish "leave
// auto-detected" (nil) from "explicitly override to none" (pointer to
// PanelNone / WSNone). The non-pointer string/slice fields use the
// zero-value-means-unset convention since they have no legitimate "none"
// value to override to.
type Overrides struct {
Panel *Panel
WebServer *WebServer
AccessLogPaths []string
ErrorLogPaths []string
ModSecAuditLogPaths []string
DomlogGlobs []string
ApacheConfigDir string
NginxConfigDir string
}
var (
// detectedPtr caches the merged Detect() result. Stored behind an atomic
// pointer (not a sync.Once + plain field) so Refresh can replace it after
// the daemon has already probed once -- a host that mis-detected its web
// server at boot (lsws not yet running) self-heals on a later refresh.
detectedPtr atomic.Pointer[Info]
detectMu sync.Mutex // serialises first Detect/Refresh with SetOverrides
overrideMu sync.Mutex
pendingOverride *Overrides
)
// SetOverrides installs config-supplied overrides to be merged into the next
// (and all subsequent) Detect() result. Call this once from daemon startup,
// BEFORE the first Detect() call, so the merged info is what every check
// sees. Subsequent SetOverrides calls before Detect() replace the previous
// override; calls after Detect() are no-ops and log a warning via the
// returned bool.
//
// Returns true if the override was installed, false if Detect() had already
// cached an un-overridden result.
func SetOverrides(o Overrides) bool {
detectMu.Lock()
defer detectMu.Unlock()
overrideMu.Lock()
defer overrideMu.Unlock()
// sync.Once has no public "was called" query, so detectedFlag is the
// guard that makes late override attempts a true no-op for both Detect
// and fresh re-probes. Re-installing the overrides that are already in
// force is not late: several startup paths install and then detect, and
// the daemon's own call must not read "already detected" as "lost".
if isDetected() {
return pendingOverride != nil && overridesEqual(*pendingOverride, o)
}
pendingOverride = &o
return true
}
func overridesEqual(a, b Overrides) bool {
return optionalEqual(a.Panel, b.Panel) &&
optionalEqual(a.WebServer, b.WebServer) &&
slices.Equal(a.AccessLogPaths, b.AccessLogPaths) &&
slices.Equal(a.ErrorLogPaths, b.ErrorLogPaths) &&
slices.Equal(a.ModSecAuditLogPaths, b.ModSecAuditLogPaths) &&
slices.Equal(a.DomlogGlobs, b.DomlogGlobs) &&
a.ApacheConfigDir == b.ApacheConfigDir &&
a.NginxConfigDir == b.NginxConfigDir
}
func optionalEqual[T comparable](a, b *T) bool {
if a == nil || b == nil {
return a == nil && b == nil
}
return *a == *b
}
// isDetected returns true if Detect() has already cached a result.
// Internal helper — uses a separate flag because sync.Once has no query API.
var detectedFlag atomic.Bool
func isDetected() bool { return detectedFlag.Load() }
// Detect inspects the host and returns platform info. The result is cached
// until the next Refresh — callers that need a fresh probe without touching
// the cache should use DetectFresh instead.
func Detect() Info {
if p := detectedPtr.Load(); p != nil {
return *p
}
detectMu.Lock()
defer detectMu.Unlock()
// Double-check: another goroutine may have populated the cache while we
// waited for the lock.
if p := detectedPtr.Load(); p != nil {
return *p
}
i := DetectFreshWithOverrides()
detectedPtr.Store(&i)
detectedFlag.Store(true)
return i
}
// Refresh re-runs detection (applying any operator overrides) and replaces the
// cached Detect() result. Intended for a periodic caller so a detection that
// was wrong at boot -- for example a cPanel+LiteSpeed host probed before lsws
// finished starting, which pins Apache for everything -- self-heals instead of
// staying wrong for the daemon's lifetime. An operator web_server.type pin
// still wins, exactly as it does through Detect(), because
// DetectFreshWithOverrides re-applies the pending override.
func Refresh() Info {
detectMu.Lock()
defer detectMu.Unlock()
i := DetectFreshWithOverrides()
detectedPtr.Store(&i)
detectedFlag.Store(true)
return i
}
// DetectFresh always re-runs detection, ignoring any cached result.
// Intended for tests and for operator-triggered rescan. Does not apply
// config overrides — use Detect() for the operator-visible view.
func DetectFresh() Info {
i := Info{}
detectOS(&i)
detectPanel(&i)
detectWebServer(&i)
populatePaths(&i)
return i
}
// DetectFreshWithOverrides re-runs detection (ignoring the one-shot Detect()
// cache) and merges in any operator-supplied overrides, returning the same
// view Detect() would produce but without the process-lifetime cache.
//
// Periodic re-probes (the ModSec rule-action registry refresh) use this so a
// detection that was wrong at boot -- for example a LiteSpeed host probed
// before lsws finished starting, which resolves to the wrong rule
// directories and yields an empty registry -- self-heals on the next refresh
// instead of staying wrong for the daemon's lifetime. A configured
// web_server.type override still wins, exactly as it does through Detect().
func DetectFreshWithOverrides() Info {
i := DetectFresh()
var o Overrides
hasOverride := false
overrideMu.Lock()
if pendingOverride != nil {
o = *pendingOverride
hasOverride = true
}
overrideMu.Unlock()
if hasOverride {
i = applyOverrides(i, o)
}
return i
}
// applyOverrides merges non-empty override fields into info. Always returns
// a new Info — never mutates the input. Paths are replaced, not appended:
// if the operator configured an explicit access-log list, the auto-detected
// list is discarded so operators have full control.
func applyOverrides(info Info, o Overrides) Info {
// Panel override must happen before path rebuild so populatePaths
// picks up the cPanel overlay (or drops it) correctly. Nil means
// "leave auto-detected"; a non-nil pointer always wins, even when it
// points at PanelNone, so operators can explicitly force a host to
// look panel-less.
if o.Panel != nil {
info.Panel = *o.Panel
}
if o.WebServer != nil {
// Web server type changed -- rebuild paths from scratch unless the
// operator also supplied path overrides below. Same nil-vs-pointer
// semantics as Panel: a pointer at WSNone forces "no web server"
// instead of being silently ignored.
info.WebServer = *o.WebServer
info.AccessLogPaths = nil
info.ErrorLogPaths = nil
info.ModSecAuditLogPaths = nil
populatePaths(&info)
}
// When Panel or WebServer changed, recompute DomlogGlobs so the globs
// stay consistent with the new panel/web-server combination. Explicit
// DomlogGlobs override below wins over this recompute.
if o.Panel != nil || o.WebServer != nil {
info.DomlogGlobs = nil
populateDomlogGlobs(&info)
}
if len(o.AccessLogPaths) > 0 {
info.AccessLogPaths = append([]string(nil), o.AccessLogPaths...)
}
if len(o.ErrorLogPaths) > 0 {
info.ErrorLogPaths = append([]string(nil), o.ErrorLogPaths...)
}
if len(o.ModSecAuditLogPaths) > 0 {
info.ModSecAuditLogPaths = append([]string(nil), o.ModSecAuditLogPaths...)
}
if len(o.DomlogGlobs) > 0 {
info.DomlogGlobs = append([]string(nil), o.DomlogGlobs...)
}
if o.ApacheConfigDir != "" {
info.ApacheConfigDir = o.ApacheConfigDir
}
if o.NginxConfigDir != "" {
info.NginxConfigDir = o.NginxConfigDir
}
return info
}
// ResetForTest clears the cached Detect() result so tests can re-run with
// different fixtures. Never call from production code.
func ResetForTest() {
detectMu.Lock()
detectedPtr.Store(nil)
detectedFlag.Store(false)
detectMu.Unlock()
overrideMu.Lock()
pendingOverride = nil
overrideMu.Unlock()
}
func detectOS(i *Info) {
f, err := os.Open("/etc/os-release")
if err != nil {
return
}
defer func() { _ = f.Close() }()
scanner := bufio.NewScanner(f)
var id, versionID string
for scanner.Scan() {
line := scanner.Text()
key, val, ok := strings.Cut(line, "=")
if !ok {
continue
}
val = strings.Trim(val, `"'`)
switch key {
case "ID":
id = strings.ToLower(val)
case "VERSION_ID":
versionID = val
}
}
i.OSVersion = versionID
switch id {
case "ubuntu":
i.OS = OSUbuntu
case "debian":
i.OS = OSDebian
case "almalinux":
i.OS = OSAlma
case "rocky":
i.OS = OSRocky
case "centos":
i.OS = OSCentOS
case "rhel":
i.OS = OSRHEL
case "cloudlinux":
i.OS = OSCloudLinux
}
}
func detectPanel(i *Info) {
if _, err := os.Stat("/usr/local/cpanel/version"); err == nil {
i.Panel = PanelCPanel
return
}
if _, err := os.Stat("/usr/local/psa/version"); err == nil {
i.Panel = PanelPlesk
return
}
if _, err := os.Stat("/usr/local/directadmin/directadmin"); err == nil {
i.Panel = PanelDA
return
}
}
func detectWebServer(i *Info) {
// Prefer the process that's actually running. Fall back to installed
// binaries if nothing is running yet (first boot, non-systemd env).
running := runningServices()
// Always record binary paths for reload/control, even if not primary.
if bin, err := exec.LookPath("nginx"); err == nil {
i.NginxBinary = bin
}
if bin, err := exec.LookPath("apache2"); err == nil {
i.ApacheBinary = bin
} else if bin, err := exec.LookPath("httpd"); err == nil {
i.ApacheBinary = bin
}
// cPanel compiles its own httpd under /usr/local/apache/bin/httpd,
// which isn't always in PATH for root under the CSM service unit.
if i.ApacheBinary == "" {
const cpHttpd = "/usr/local/apache/bin/httpd"
if _, err := os.Stat(cpHttpd); err == nil {
i.ApacheBinary = cpHttpd
}
}
// LiteSpeed ships lshttpd/litespeed under /usr/local/lsws/bin and is not
// in PATH under the service unit. A host that serves via LiteSpeed has
// this binary on disk even before the lsws service is up, so probing for
// it lets detection correct the boot-order misdetect below.
for _, p := range []string{"/usr/local/lsws/bin/lshttpd", "/usr/local/lsws/bin/litespeed"} {
if _, err := os.Stat(p); err == nil {
i.LiteSpeedBinary = p
break
}
}
if i.LiteSpeedBinary == "" {
if bin, err := exec.LookPath("lshttpd"); err == nil {
i.LiteSpeedBinary = bin
} else if bin, err := exec.LookPath("litespeed"); err == nil {
i.LiteSpeedBinary = bin
}
}
i.WebServer = selectWebServer(i.Panel, running, i.ApacheBinary != "", i.NginxBinary != "")
i.WebServer = liteSpeedBinaryFallback(i.WebServer, running, i.LiteSpeedBinary != "")
}
// liteSpeedBinaryFallback corrects the historical boot-order misdetect: when
// no web server process is running yet, selectWebServer falls back to a
// binary-only guess and (on cPanel, which always keeps an httpd binary) picks
// Apache. A LiteSpeed binary present on disk is the strongest signal that this
// host serves through LiteSpeed, so it outranks that binary-only fallback.
//
// It only overrides when nothing is actually running: a live Apache/Nginx
// process is trusted over an installed-but-idle lsws binary.
func liteSpeedBinaryFallback(ws WebServer, running map[string]bool, hasLiteSpeedBinary bool) WebServer {
if !hasLiteSpeedBinary || anyWebServerRunning(running) {
return ws
}
switch ws {
case WSApache, WSNginx, WSNone:
return WSLiteSpeed
default:
return ws
}
}
// anyWebServerRunning reports whether any known web server unit is active, so
// the binary-only fallbacks only fire when nothing is genuinely serving.
func anyWebServerRunning(running map[string]bool) bool {
for _, up := range running {
if up {
return true
}
}
return false
}
// runningServices returns which web server process units are currently
// active. Uses systemctl when available; falls back to checking /proc.
func runningServices() map[string]bool {
active := map[string]bool{}
for _, unit := range []string{"nginx", "apache2", "httpd", "litespeed", "lshttpd", "lsws"} {
// #nosec G204 -- systemctl hardcoded; unit iterates a literal slice.
cmd := exec.Command("systemctl", "is-active", "--quiet", unit)
if err := cmd.Run(); err == nil {
active[unit] = true
}
}
return active
}
func selectWebServer(panel Panel, running map[string]bool, hasApacheBinary, hasNginxBinary bool) WebServer {
apacheRunning := running["apache2"] || running["httpd"]
litespeedRunning := running["litespeed"] || running["lshttpd"] || running["lsws"]
nginxRunning := running["nginx"]
// cPanel commonly runs Nginx as a reverse proxy in front of Apache.
// Prefer the origin server logs when Apache is active so real-time
// access and ModSecurity watchers tail the paths cPanel actually writes.
if panel == PanelCPanel {
switch {
case litespeedRunning:
return WSLiteSpeed
case apacheRunning:
return WSApache
case nginxRunning:
return WSNginx
case hasApacheBinary:
return WSApache
case hasNginxBinary:
return WSNginx
default:
return WSNone
}
}
switch {
case nginxRunning:
return WSNginx
case apacheRunning:
return WSApache
case litespeedRunning:
return WSLiteSpeed
case hasApacheBinary:
return WSApache
case hasNginxBinary:
return WSNginx
default:
return WSNone
}
}
func populatePaths(i *Info) {
// Apache config dir. cPanel compiles Apache from source and installs
// under /usr/local/apache, separate from the OS package tree; that
// override wins over the distro default when cPanel is present.
switch {
case i.Panel == PanelCPanel && dirExists("/usr/local/apache/conf"):
i.ApacheConfigDir = "/usr/local/apache/conf"
case i.IsDebianFamily():
if dirExists("/etc/apache2") {
i.ApacheConfigDir = "/etc/apache2"
}
case i.IsRHELFamily():
if dirExists("/etc/httpd") {
i.ApacheConfigDir = "/etc/httpd"
}
}
if dirExists("/etc/nginx") {
i.NginxConfigDir = "/etc/nginx"
}
// Log paths: pick candidates based on detected web server and OS layout.
// We include ALL plausible locations so log watchers can try each;
// missing paths are handled upstream by the retry logic.
switch i.WebServer {
case WSApache:
if i.IsDebianFamily() {
i.AccessLogPaths = []string{"/var/log/apache2/access.log", "/var/log/apache2/other_vhosts_access.log"}
i.ErrorLogPaths = []string{"/var/log/apache2/error.log"}
i.ModSecAuditLogPaths = []string{"/var/log/apache2/modsec_audit.log"}
} else {
i.AccessLogPaths = []string{"/var/log/httpd/access_log"}
i.ErrorLogPaths = []string{"/var/log/httpd/error_log"}
i.ModSecAuditLogPaths = []string{"/var/log/httpd/modsec_audit.log"}
}
case WSNginx:
i.AccessLogPaths = []string{"/var/log/nginx/access.log"}
i.ErrorLogPaths = []string{"/var/log/nginx/error.log"}
i.ModSecAuditLogPaths = []string{"/var/log/nginx/modsec_audit.log"}
case WSLiteSpeed:
i.AccessLogPaths = []string{"/usr/local/lsws/logs/access.log"}
i.ErrorLogPaths = []string{"/usr/local/lsws/logs/error.log"}
i.ModSecAuditLogPaths = []string{"/usr/local/lsws/logs/auditmodsec.log"}
}
// cPanel overlays its own access/error logs on top of the OS defaults.
if i.Panel == PanelCPanel {
i.AccessLogPaths = append([]string{
"/usr/local/apache/logs/access_log",
"/usr/local/cpanel/logs/access_log",
}, i.AccessLogPaths...)
i.ErrorLogPaths = append([]string{
"/usr/local/apache/logs/error_log",
}, i.ErrorLogPaths...)
i.ModSecAuditLogPaths = append([]string{
"/usr/local/apache/logs/modsec_audit.log",
"/var/log/modsec_audit.log",
}, i.ModSecAuditLogPaths...)
}
populateDomlogGlobs(i)
}
// populateDomlogGlobs sets DomlogGlobs based on panel type and, for
// bare-metal installs, the web server + OS family. Panel takes
// precedence over web server because panel-specific layouts write
// per-vhost logs to panel-owned directories regardless of what web
// server is running underneath.
func populateDomlogGlobs(i *Info) {
switch i.Panel {
case PanelCPanel:
i.DomlogGlobs = []string{
"/home/*/access-logs/*-ssl_log",
"/home/*/access-logs/*_log",
}
case PanelPlesk:
i.DomlogGlobs = []string{
"/var/www/vhosts/*/logs/access_ssl_log",
"/var/www/vhosts/*/logs/access_log",
"/var/www/vhosts/*/logs/proxy_access_ssl_log",
}
case PanelDA:
i.DomlogGlobs = []string{"/var/log/httpd/domains/*.log"}
default:
switch i.WebServer {
case WSApache:
if i.IsDebianFamily() {
i.DomlogGlobs = []string{
"/var/log/apache2/*-access.log",
"/var/log/apache2/*_access.log",
}
} else if i.IsRHELFamily() {
i.DomlogGlobs = []string{
"/var/log/httpd/*-access_log",
"/var/log/httpd/*_access_log",
}
}
case WSNginx:
i.DomlogGlobs = []string{
"/var/log/nginx/*.access.log",
"/var/log/nginx/*-access.log",
}
}
}
}
func dirExists(p string) bool {
fi, err := os.Stat(p)
return err == nil && fi.IsDir()
}
// Package privops is the inventory of every CSM operation that needs privilege
// beyond reading its own files, or that writes outside CSM's own directories.
//
// It exists so an operator can answer three questions before granting a
// security daemon host-level access: what does it touch, what does each of
// those need root (or a capability) for, and how do I turn any of it off.
// `csm privileges` renders it, docs/src/capability-matrix.md ships the
// rendered form, and gates in internal/ci keep both in step with the systemd
// sandbox that actually constrains the daemon.
package privops
import (
"encoding/json"
"fmt"
"path"
"slices"
"sort"
"strings"
)
// Privilege is a capability or credential an operation needs.
type Privilege string
const (
// Root means the operation needs uid 0 rather than one capability: it
// reads or writes files owned by many different accounts, or it drives a
// panel tool that assumes root.
Root Privilege = "root"
// CapDACReadSearch is the read side of scanning every account's files.
CapDACReadSearch Privilege = "CAP_DAC_READ_SEARCH"
// CapSysAdmin covers fanotify and mount inspection.
CapSysAdmin Privilege = "CAP_SYS_ADMIN"
// CapNetAdmin covers nftables mutation.
CapNetAdmin Privilege = "CAP_NET_ADMIN"
// CapKill covers signalling processes owned by other accounts.
CapKill Privilege = "CAP_KILL"
// CapBPF covers loading and attaching BPF programs.
CapBPF Privilege = "CAP_BPF"
// CapPerfmon is required alongside CAP_BPF for tracing and LSM programs.
CapPerfmon Privilege = "CAP_PERFMON"
// CapSyslog permits reading a restricted kernel message buffer.
CapSyslog Privilege = "CAP_SYSLOG"
// CapAuditControl permits querying and loading kernel audit rules.
CapAuditControl Privilege = "CAP_AUDIT_CONTROL"
// CapSysModule permits loading or unloading kernel modules.
CapSysModule Privilege = "CAP_SYS_MODULE"
// CapLinuxImmutable permits changing the executable's immutable flag.
CapLinuxImmutable Privilege = "CAP_LINUX_IMMUTABLE"
// Unprivileged marks work CSM does inside its own directories.
Unprivileged Privilege = "none"
)
// KnownPrivileges lists every privilege an operation may declare.
func KnownPrivileges() []Privilege {
return []Privilege{Root, CapDACReadSearch, CapSysAdmin, CapNetAdmin, CapKill, CapBPF, CapPerfmon, CapSyslog, CapAuditControl, CapSysModule, CapLinuxImmutable, Unprivileged}
}
// Trigger says who starts an operation.
type Trigger string
const (
// Automatic operations run on the daemon's own schedule or in response to
// a detection, with no operator in the loop.
Automatic Trigger = "automatic"
// Operator operations run only when someone issues a command or clicks a
// button. Not running the command is how they are turned off.
Operator Trigger = "operator"
)
// RiskTier is an operation's action-risk tier: what can go wrong if it runs
// on a wrong target. The zero value is unclassified and fails
// TestEveryOperationHasARiskTier.
type RiskTier uint8
const (
RiskUnclassified RiskTier = iota
// RiskObserve (tier 0) changes nothing outside CSM's own trees.
RiskObserve
// RiskPreview (tier 1) records a recommendation or dry-run decision only.
// No inventory row currently represents previews separately; rows carry
// their maximum live effect, even when an execution can be a dry run.
RiskPreview
// RiskReversible (tier 2) makes a low-risk host change such as attaching
// a probe, opening a challenge gate or holding mail. The tier alone
// does not promise automatic rollback or reversal of incidental writes.
RiskReversible
// RiskContain (tier 3) quarantines, blocks or denies one target.
RiskContain
// RiskDestructive (tier 4) signals processes, restarts or reloads services,
// or rewrites content or configuration, including an existing archive.
RiskDestructive
)
// Number is the tier as the safety model numbers it, 0 to 4, or -1 when the
// operation is unclassified or invalid.
func (r RiskTier) Number() int {
if r < RiskObserve || r > RiskDestructive {
return -1
}
return int(r) - 1
}
// MarshalJSON uses the same public tier number as the text and Markdown
// views. The internal zero value is an unclassified sentinel, not tier 0.
func (r RiskTier) MarshalJSON() ([]byte, error) {
return []byte(fmt.Sprint(r.Number())), nil
}
// UnmarshalJSON translates public tier numbers back to their internal values.
// JSON null leaves the destination unchanged, as it does for other scalars.
func (r *RiskTier) UnmarshalJSON(data []byte) error {
var number *int
if err := json.Unmarshal(data, &number); err != nil {
return err
}
if number == nil {
return nil
}
if *number < -1 || *number > RiskDestructive.Number() {
return fmt.Errorf("invalid risk tier %d", *number)
}
*r = RiskTier(*number + 1) // #nosec G115 -- public tiers -1 through 4 map to 0 through 5.
return nil
}
// SafetyContract describes current authority, identity, recovery and limits.
// It is inventory metadata, not an enforcement mechanism or a claim that
// the full action lifecycle is implemented. Remaining gaps stay explicit.
type SafetyContract struct {
// Authority is the evidence and opt-ins required before it may run.
Authority string
// Identity is how the target is revalidated immediately before the change.
Identity string
// Recovery is how the change is reversed, or what it cannot undo.
Recovery string
// Limit names current bounds and where they do not apply.
Limit string
}
// csmOwnedPrefixes are the trees CSM creates and manages for itself. Writing
// inside them is not a host change: an operator who removes CSM removes them.
var csmOwnedPrefixes = []string{
"/var/lib/csm",
"/opt/csm",
"/var/log/csm",
"/var/log/csm-php-shield",
"/etc/csm",
"/var/cache/csm",
"/var/run/csm",
}
// Op is one privileged or state-changing operation.
type Op struct {
// ID is stable and namespaced as <subsystem>.<action>. Operators and
// panel integrations may key on it.
ID string
// Subsystem groups related operations in the rendered matrix.
Subsystem string
// Summary is one line: what the operation does.
Summary string
// Privileges is what the operation needs from the kernel or from uid 0.
Privileges []Privilege
// Trigger says whether the daemon starts this on its own.
Trigger Trigger
// Writes lists what the operation writes while it runs, including writes
// made by a tool it invokes. Filesystem paths start with "/"; anything
// else is a resource written as <kind>:<name>. Empty means read-only.
Writes []string
// Unsandboxed marks operations outside the daemon's systemd sandbox:
// transient services or standalone CLI commands. Mixed operations must
// have separate rows for their in-daemon writes.
Unsandboxed bool
// DisableKey is the config key that stops the operation, and DisableValue
// the YAML value to give it. An empty key makes no claim of a config switch.
DisableKey string
DisableValue string
// DisableReason explains why an automatic operation cannot be stopped
// through config. It must not invent a switch that only stops some callers.
DisableReason string
// Audited reports whether the operation writes a record to the action
// log. False is not a claim that the operation is silent, only that it is
// not yet on that stream; the daemon log still carries it.
Audited bool
// Risk is the operation's action-risk tier.
Risk RiskTier
// Contract is the operation's safety contract; nil until that contract has
// been specified for the operation.
Contract *SafetyContract
// RecoveryGap names recovery work not covered by this inventory's
// contracts. Required for host-changing operations without a contract.
RecoveryGap string
// WithoutPrivilege says what an operator loses by withholding the
// privilege, so the matrix reads as a decision, not a demand.
WithoutPrivilege string
}
// ChangesHost reports whether the operation writes outside CSM's own trees.
func (o Op) ChangesHost() bool {
for _, w := range o.Writes {
if !strings.HasPrefix(w, "/") {
return true
}
if !csmOwned(w) {
return true
}
}
return false
}
func csmOwned(name string) bool {
name = path.Clean(name)
for _, prefix := range csmOwnedPrefixes {
if name == prefix || strings.HasPrefix(name, prefix+"/") {
return true
}
}
return false
}
// Operations returns the inventory, ordered by subsystem then ID.
func Operations() []Op {
ops := append([]Op(nil), operations...)
for i := range ops {
ops[i].Privileges = slices.Clone(ops[i].Privileges)
ops[i].Writes = slices.Clone(ops[i].Writes)
if ops[i].Contract != nil {
c := *ops[i].Contract
ops[i].Contract = &c
}
}
sort.Slice(ops, func(i, j int) bool {
if ops[i].Subsystem != ops[j].Subsystem {
return ops[i].Subsystem < ops[j].Subsystem
}
return ops[i].ID < ops[j].ID
})
return ops
}
// DisableInstruction is shared by terminal and documentation output.
func (o Op) DisableInstruction() string {
if o.DisableKey != "" {
return o.DisableKey + ": " + o.DisableValue
}
if o.Trigger == Operator {
return "do not run the command"
}
if o.DisableReason != "" {
return "not configurable: " + o.DisableReason
}
return "not configurable"
}
// Markdown renders the inventory as the table shipped in the docs.
func Markdown() string {
var b strings.Builder
b.WriteString("| Operation | Needs | Trigger | Risk tier | Writes | Turn it off | Action record | Without the privilege |\n")
b.WriteString("| --- | --- | --- | --- | --- | --- | --- | --- |\n")
for _, op := range Operations() {
privs := make([]string, 0, len(op.Privileges))
for _, p := range op.Privileges {
privs = append(privs, string(p))
}
writes := "nothing (read-only)"
if len(op.Writes) > 0 {
writes = strings.Join(op.Writes, ", ")
}
if op.Unsandboxed {
writes += " (outside the systemd sandbox)"
}
off := op.DisableInstruction()
if op.DisableKey != "" {
off = "`" + off + "`"
}
audited := "no"
if op.Audited {
audited = "yes"
}
fmt.Fprintf(&b, "| `%s`<br>%s | %s | %s | %d | %s | %s | %s | %s |\n",
op.ID, op.Summary, strings.Join(privs, ", "), op.Trigger, op.Risk.Number(), writes, off, audited, op.WithoutPrivilege)
}
return b.String()
}
package processctx
import (
"container/list"
"sync"
"sync/atomic"
"time"
)
// Cache is a bounded LRU cache of processEntry keyed by PID with a per-entry
// TTL. Reads and writes are safe for concurrent use. now is a seam for
// deterministic testing.
type Cache struct {
mu sync.Mutex
cap int
ttl time.Duration
ll *list.List // front = most-recently-used
index map[int]*list.Element // pid -> element holding *processEntry
now func() time.Time
evictions atomic.Uint64
ttlPurges atomic.Uint64
misses atomic.Uint64
}
// Stats is a snapshot of cache counters. Safe to call concurrently.
type Stats struct {
Entries int
Evictions uint64 // LRU evictions (cap exceeded)
TTLPurges uint64 // entries dropped because ttl expired on Get
Misses uint64 // Get returned no entry (includes ttl purges)
}
// NewCache returns a cache with the given hard cap and TTL.
func NewCache(cap int, ttl time.Duration) *Cache {
if cap <= 0 {
cap = 1
}
return &Cache{
cap: cap,
ttl: ttl,
ll: list.New(),
index: make(map[int]*list.Element, cap),
now: time.Now,
}
}
// Put inserts or updates an entry. lastTouch is set to now.
func (c *Cache) Put(e processEntry) {
c.mu.Lock()
defer c.mu.Unlock()
e.lastTouch = c.now()
if el, ok := c.index[e.PID]; ok {
el.Value = &e
c.ll.MoveToFront(el)
return
}
el := c.ll.PushFront(&e)
c.index[e.PID] = el
for c.ll.Len() > c.cap {
c.evictOldestLocked()
}
}
// Get returns the entry for pid if present and not TTL-expired. Touching
// an entry promotes it in LRU order.
func (c *Cache) Get(pid int) (processEntry, bool) {
c.mu.Lock()
defer c.mu.Unlock()
el, ok := c.index[pid]
if !ok {
c.misses.Add(1)
return processEntry{}, false
}
entry := el.Value.(*processEntry)
if c.ttl > 0 && c.now().Sub(entry.lastTouch) > c.ttl {
c.removeLocked(el)
c.ttlPurges.Add(1)
c.misses.Add(1)
return processEntry{}, false
}
entry.lastTouch = c.now()
c.ll.MoveToFront(el)
return *entry, true
}
// Len returns the number of live entries (without forcing TTL purge).
func (c *Cache) Len() int {
c.mu.Lock()
defer c.mu.Unlock()
return c.ll.Len()
}
// Stats returns a counter snapshot.
func (c *Cache) Stats() Stats {
c.mu.Lock()
n := c.ll.Len()
c.mu.Unlock()
return Stats{
Entries: n,
Evictions: c.evictions.Load(),
TTLPurges: c.ttlPurges.Load(),
Misses: c.misses.Load(),
}
}
func (c *Cache) evictOldestLocked() {
el := c.ll.Back()
if el == nil {
return
}
c.removeLocked(el)
c.evictions.Add(1)
}
func (c *Cache) removeLocked(el *list.Element) {
entry := el.Value.(*processEntry)
delete(c.index, entry.PID)
c.ll.Remove(el)
}
// PutFromExec is a minimal constructor for callers that have only
// PID/UID/comm/exe from an exec event. UIDKnown is true even for UID 0.
func (c *Cache) PutFromExec(pid, ppid, uid int, comm, exe string) {
c.PutFromExecStartedAt(pid, ppid, uid, comm, exe, time.Time{})
}
// PutFromExecStartedAt is PutFromExec with a detector-supplied process start
// time for PID-reuse validation.
func (c *Cache) PutFromExecStartedAt(pid, ppid, uid int, comm, exe string, startedAt time.Time) {
c.Put(processEntry{PID: pid, PPID: ppid, UID: uid, UIDKnown: true, Comm: comm, Exe: exe, StartedAt: startedAt})
}
// PutFromProc inserts a fully populated entry from /proc-style data. The
// current enricher validates and writes inside processctx; this helper keeps a
// public constructor for tests and future non-daemon callers without exposing
// processEntry.
func (c *Cache) PutFromProc(pid, ppid, uid int, user, account, comm, exe string, cmdline []string) {
c.PutFromProcStartedAt(pid, ppid, uid, user, account, comm, exe, cmdline, time.Time{})
}
// PutFromProcStartedAt is PutFromProc with a known process start time.
func (c *Cache) PutFromProcStartedAt(pid, ppid, uid int, user, account, comm, exe string, cmdline []string, startedAt time.Time) {
c.Put(processEntry{
PID: pid, PPID: ppid, UID: uid, UIDKnown: true,
User: user, Account: account,
Comm: comm, Exe: exe, Cmdline: cmdline, StartedAt: startedAt, ProcRead: true,
})
}
package processctx
import (
"errors"
"maps"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
// procReader is the slice of ProcReader the Enricher needs. Allows fakes.
type procReader interface {
Read(pid int) (processEntry, error)
}
// EnrichRequest is the immutable event snapshot queued off the ring-buffer
// path. UID/Comm/StartedAt are used to reject stale PID reuse before caching
// /proc data.
type EnrichRequest struct {
PID int
UID int
UIDKnown bool
Comm string
StartedAt time.Time
}
type enrichWork struct {
req EnrichRequest
ticket queuehealth.Ticket
}
// IdentityResolver maps a process UID to username/account metadata. It must be
// cache-only in the common path. The daemon implementation uses
// checks.LookupUser's cached /etc/passwd reader and simple local account
// inference; it must not call NSS, LDAP, whmapi1, network services, or any
// blocking account enumerator from the enricher worker.
//
// Implementations SHOULD return within ~1ms in the common case. If a future
// implementation needs a backing data source that can block, refresh it in a
// separate cache outside the worker and have Resolve return ("", "") on cache
// miss rather than stalling the enrichment queue.
type IdentityResolver interface {
Resolve(uid int) (user, account string)
}
type noopResolver struct{}
func (noopResolver) Resolve(int) (string, string) { return "", "" }
// EnricherConfig sizes the worker pool and queue.
type EnricherConfig struct {
Workers int
QueueCap int
Resolver IdentityResolver
}
// EnricherStats is a snapshot of enricher counters.
type EnricherStats struct {
Enqueued uint64
Drops uint64
Reads uint64
Errors uint64
Stale uint64
}
// Enricher consumes PIDs and populates Cache from ProcReader.Read off the
// hot path. Enqueue is nonblocking. On overflow it drops the oldest queued
// request and records that drop, so the producer keeps moving and the queue
// favors fresher process snapshots.
type Enricher struct {
cache *Cache
reader procReader
resolver IdentityResolver
cfg EnricherConfig
queue chan enrichWork
queueStats *queuehealth.Tracker
admission sync.Mutex
wg sync.WaitGroup
stopCh chan struct{}
started bool
stopped bool
stopOnce sync.Once
enqueued atomic.Uint64
drops atomic.Uint64
reads atomic.Uint64
errors atomic.Uint64
stale atomic.Uint64
latencyMu sync.RWMutex
observeLatency func(float64)
}
// NewEnricher prepares a pool; Start launches its workers.
func NewEnricher(cache *Cache, reader procReader, cfg EnricherConfig) *Enricher {
if cfg.Workers <= 0 {
cfg.Workers = 2
}
if cfg.QueueCap <= 0 {
cfg.QueueCap = 1024
}
resolver := cfg.Resolver
if resolver == nil {
resolver = noopResolver{}
}
return &Enricher{
cache: cache,
reader: reader,
resolver: resolver,
cfg: cfg,
queue: make(chan enrichWork, cfg.QueueCap),
queueStats: queuehealth.New(cfg.QueueCap, time.Minute),
stopCh: make(chan struct{}),
}
}
// Start launches the worker goroutines. Idempotent.
func (e *Enricher) Start() {
e.admission.Lock()
defer e.admission.Unlock()
if e.started || e.stopped {
return
}
e.started = true
e.wg.Add(e.cfg.Workers)
for i := 0; i < e.cfg.Workers; i++ {
go e.worker()
}
}
// Stop signals workers and waits for them to exit. Safe to call multiple times.
//
// The pool cannot restart. Running reads finish; buffered requests are discarded
// and counted after workers relinquish ownership.
func (e *Enricher) Stop() {
e.stopOnce.Do(func() {
e.admission.Lock()
e.stopped = true
close(e.stopCh)
close(e.queue)
e.admission.Unlock()
e.wg.Wait()
for work := range e.queue {
work.ticket.Reject(time.Now())
e.drops.Add(1)
}
})
}
// Enqueue adds a request to the work queue. Returns false for an invalid PID
// or a stopped enricher. If the queue is full, the oldest pending request is
// dropped and the new one is queued.
func (e *Enricher) Enqueue(req EnrichRequest) bool {
e.admission.Lock()
defer e.admission.Unlock()
if req.PID <= 0 || e.stopped {
e.queueStats.Lose(time.Now(), 1)
e.drops.Add(1)
return false
}
work := enrichWork{req: req, ticket: e.queueStats.Begin(time.Now())}
select {
case e.queue <- work:
default:
select {
case oldest := <-e.queue:
oldest.ticket.Reject(time.Now())
e.drops.Add(1)
default:
}
// Other producers and close are excluded; consumers can only free slots.
e.queue <- work
}
e.enqueued.Add(1)
return true
}
func (e *Enricher) QueueStatuses(now time.Time) map[string]queuehealth.Status {
out := map[string]queuehealth.Status{"enrichment": e.queueStats.Snapshot(now)}
if reader, ok := e.reader.(interface {
QueueStatuses(time.Time) map[string]queuehealth.Status
}); ok {
maps.Copy(out, reader.QueueStatuses(now))
}
return out
}
// SetLatencyObserver installs an optional callback used by metrics.
func (e *Enricher) SetLatencyObserver(fn func(float64)) {
e.latencyMu.Lock()
defer e.latencyMu.Unlock()
e.observeLatency = fn
}
func (e *Enricher) observe(seconds float64) {
e.latencyMu.RLock()
fn := e.observeLatency
e.latencyMu.RUnlock()
if fn != nil {
fn(seconds)
}
}
func (e *Enricher) shouldCache(req EnrichRequest, entry processEntry) bool {
reqUIDKnown := req.UIDKnown || req.UID != 0
if reqUIDKnown {
if !entry.UIDKnown || req.UID != entry.UID {
return false
}
} else if !entry.UIDKnown {
return false
}
if req.Comm != "" && entry.Comm != req.Comm {
return false
}
// PID-reuse guard: when the detector supplied a process-start
// snapshot, the /proc-derived start time must be present and match
// within a small tolerance. A missing or mismatched /proc value means
// the enricher cannot prove it is looking at the same process.
if !processStartMatches(req.StartedAt, entry.StartedAt) {
return false
}
return true
}
// processStartTimeTolerance bounds the allowed clock skew between a
// detector's process-start snapshot and /proc/<pid>/stat's starttime. Five
// seconds covers clock granularity and slow pickup without letting a PID-reuse
// race slip past.
const processStartTimeTolerance = 5 * time.Second
func processStartMatches(want, got time.Time) bool {
if want.IsZero() {
return true
}
if got.IsZero() {
return false
}
diff := want.Sub(got)
if diff < 0 {
diff = -diff
}
return diff <= processStartTimeTolerance
}
func (e *Enricher) enrichIdentity(entry *processEntry) {
if !entry.UIDKnown {
return
}
user, account := e.resolver.Resolve(entry.UID)
entry.User = user
entry.Account = account
}
// Stats returns a counter snapshot.
func (e *Enricher) Stats() EnricherStats {
return EnricherStats{
Enqueued: e.enqueued.Load(),
Drops: e.drops.Load(),
Reads: e.reads.Load(),
Errors: e.errors.Load(),
Stale: e.stale.Load(),
}
}
func (e *Enricher) worker() {
defer e.wg.Done()
for {
select {
case <-e.stopCh:
return
default:
}
select {
case <-e.stopCh:
return
case work, ok := <-e.queue:
if !ok {
return
}
e.process(work)
}
}
}
func (e *Enricher) process(work enrichWork) {
start := time.Now()
work.ticket.Start(start)
completed := false
defer func() {
if completed {
work.ticket.Finish(time.Now())
} else {
work.ticket.Reject(time.Now())
}
}()
e.reads.Add(1)
entry, err := e.reader.Read(work.req.PID)
e.observe(time.Since(start).Seconds())
if err != nil {
if errors.Is(err, ErrProcessGone) {
completed = true
} else {
e.errors.Add(1)
}
return
}
if !e.shouldCache(work.req, entry) {
e.stale.Add(1)
completed = true
return
}
e.enrichIdentity(&entry)
e.cache.Put(entry)
completed = true
}
package processctx
import "time"
// Materialize walks the PPID chain starting at pid up to MaxParentDepth and
// returns a ProcessContext tree. Returns nil if pid is not in the cache.
// Cycle-safe: tracks visited PIDs.
func (c *Cache) Materialize(pid int) *ProcessContext {
root, ok := c.Get(pid)
if !ok {
return nil
}
return c.materializeFromRoot(root)
}
// MaterializeVerified returns a materialized context only when the cached root
// entry still matches the event snapshot. The bool return reports whether the
// root entry still needs an off-path /proc read.
func (c *Cache) MaterializeVerified(pid, uid int, uidKnown bool, comm string) (*ProcessContext, bool) {
root, ok := c.Get(pid)
if !ok {
return nil, false
}
if !matchesSnapshot(root, uid, uidKnown, comm) {
return nil, false
}
return c.materializeFromRoot(root), !root.ProcRead
}
// MaterializeVerifiedSnapshot returns a materialized context only when the
// cached root entry still matches the full detector snapshot, including the
// process start time when one was captured.
func (c *Cache) MaterializeVerifiedSnapshot(req EnrichRequest) (*ProcessContext, bool) {
root, ok := c.Get(req.PID)
if !ok {
return nil, false
}
if !matchesSnapshot(root, req.UID, req.UIDKnown, req.Comm) {
return nil, false
}
if !processStartMatches(req.StartedAt, root.StartedAt) {
return nil, false
}
return c.materializeFromRoot(root), !root.ProcRead
}
func (c *Cache) materializeFromRoot(root processEntry) *ProcessContext {
visited := map[int]bool{root.PID: true}
head := toContext(root)
cur := head
parentPID := root.PPID
for depth := 1; depth < MaxParentDepth; depth++ {
if parentPID <= 0 || visited[parentPID] {
break
}
entry, ok := c.Get(parentPID)
if !ok {
break
}
visited[entry.PID] = true
cur.Parent = toContext(entry)
cur = cur.Parent
parentPID = entry.PPID
}
return head
}
func matchesSnapshot(e processEntry, uid int, uidKnown bool, comm string) bool {
if uidKnown {
if !e.UIDKnown || e.UID != uid {
return false
}
}
if comm != "" && e.Comm != comm {
return false
}
return true
}
func toContext(e processEntry) *ProcessContext {
var startedAt *time.Time
if !e.StartedAt.IsZero() {
t := e.StartedAt
startedAt = &t
}
return &ProcessContext{
PID: e.PID,
PPID: e.PPID,
UID: e.UID,
User: e.User,
Account: e.Account,
Comm: e.Comm,
Exe: e.Exe,
ExeResolved: e.ProcRead && e.Exe != "",
Cmdline: append([]string(nil), e.Cmdline...),
StartedAt: startedAt,
}
}
package processctx
import "github.com/pidginhost/csm/internal/metrics"
// RegisterMetrics binds cache and enricher counters/gauges to reg. The
// registry argument allows tests to use a private registry. Production
// callers should pass metrics.Default().
func RegisterMetrics(reg *metrics.Registry, cache *Cache, enr *Enricher) {
reg.RegisterGaugeFunc(
"csm_process_context_cache_entries",
"Live process-context cache entries.",
func() float64 { return float64(cache.Stats().Entries) },
)
reg.RegisterCounterFunc(
"csm_process_context_cache_evictions_total",
"Process-context cache LRU evictions (cap exceeded).",
func() float64 { return float64(cache.Stats().Evictions) },
)
reg.RegisterCounterFunc(
"csm_process_context_cache_ttl_purges_total",
"Process-context cache entries dropped because TTL expired on lookup.",
func() float64 { return float64(cache.Stats().TTLPurges) },
)
reg.RegisterCounterFunc(
"csm_process_context_cache_misses_total",
"Process-context cache lookup misses (includes TTL purges).",
func() float64 { return float64(cache.Stats().Misses) },
)
reg.RegisterCounterFunc(
"csm_process_context_enrich_queue_drops_total",
"Process-context enrichment requests refused, evicted or abandoned at shutdown.",
func() float64 { return float64(enr.Stats().Drops) },
)
reg.RegisterCounterFunc(
"csm_process_context_enrich_reads_total",
"Process-context /proc reads attempted by enricher workers.",
func() float64 { return float64(enr.Stats().Reads) },
)
reg.RegisterCounterFunc(
"csm_process_context_enrich_errors_total",
"Process-context /proc read errors (excluding ProcessGone).",
func() float64 { return float64(enr.Stats().Errors) },
)
reg.RegisterCounterFunc(
"csm_process_context_enrich_stale_total",
"Process-context enrichment results rejected as stale PID reuse.",
func() float64 { return float64(enr.Stats().Stale) },
)
latency := metrics.NewHistogram(
"csm_process_context_enrich_latency_seconds",
"Process-context enrichment worker latency in seconds.",
[]float64{0.001, 0.005, 0.01, 0.05, 0.1, 1},
)
reg.MustRegister("csm_process_context_enrich_latency_seconds", latency)
enr.SetLatencyObserver(latency.Observe)
}
package processctx
import (
"errors"
"io/fs"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
var procReads = newProcReadPool(procReadConcurrency)
type procReadPool struct {
slots chan struct{}
stats *queuehealth.Tracker
}
func newProcReadPool(capacity int) *procReadPool {
return &procReadPool{
slots: make(chan struct{}, capacity),
stats: queuehealth.NewSharedCapacity(capacity, time.Minute),
}
}
type procReadWork struct {
pool *procReadPool
ticket queuehealth.Ticket
remaining atomic.Int32
failOnce sync.Once
}
func (p *procReadPool) acquire() *procReadWork {
ticket := p.stats.Begin(time.Now())
select {
case p.slots <- struct{}{}:
work := &procReadWork{pool: p, ticket: ticket}
work.remaining.Store(2)
return work
default:
ticket.Reject(time.Now())
return nil
}
}
func (w *procReadWork) fail() {
w.failOnce.Do(func() { w.pool.stats.Lose(time.Now(), 1) })
}
func (w *procReadWork) release() {
// A timeout finishes the caller, not the syscall. Conversely, a returned
// syscall still owns its result until the caller takes it or times out.
if w.remaining.Add(-1) == 0 {
w.ticket.Finish(time.Now())
<-w.pool.slots
}
}
type procReadResult[T any] struct {
value T
err error
}
func executeProcRead[T any](work *procReadWork, fn func() (T, error), out chan<- procReadResult[T]) {
work.ticket.Start(time.Now())
completed := false
defer func() {
if !completed {
work.fail()
}
work.release()
}()
value, err := fn()
if err != nil && !errors.Is(err, fs.ErrNotExist) {
work.fail()
}
out <- procReadResult[T]{value: value, err: err}
completed = true
}
func runProcReadWithDeadline[T any](pool *procReadPool, d time.Duration, fn func() (T, error)) (T, bool) {
if d <= 0 {
value, err := fn()
return value, err == nil
}
var zero T
work := pool.acquire()
if work == nil {
return zero, false
}
defer work.release()
out := make(chan procReadResult[T], 1)
go executeProcRead(work, fn, out)
timer := time.NewTimer(d)
defer timer.Stop()
select {
case result := <-out:
if result.err != nil {
return zero, false
}
return result.value, true
case <-timer.C:
work.fail()
return zero, false
}
}
// QueueStatuses reports deadline-bounded process reads. The row is advisory:
// an expired read costs one detail on a finding that is still raised, and a
// loaded host expires reads without losing protection work.
func (*ProcReader) QueueStatuses(now time.Time) map[string]queuehealth.Status {
status := procReads.stats.Snapshot(now)
status.Advisory = true
return map[string]queuehealth.Status{"proc_reads": status}
}
package processctx
import (
"bytes"
"errors"
"io/fs"
"os"
"path/filepath"
"strconv"
"strings"
"time"
)
// ErrProcessGone is returned by ProcReader.Read when the /proc/<pid> tree
// no longer exists. Callers must treat this as a soft miss, not an error
// finding - short-lived processes are expected.
var ErrProcessGone = errors.New("process gone")
// ProcReader reads /proc/<pid>/{status,cmdline,exe} with a per-file deadline.
type ProcReader struct {
root string
perFileDeadline time.Duration
}
// NewProcReader constructs a reader rooted at procRoot ("/proc" in production,
// a temp dir in tests). perFileDeadline bounds each individual file read.
func NewProcReader(procRoot string, perFileDeadline time.Duration) *ProcReader {
return &ProcReader{root: procRoot, perFileDeadline: perFileDeadline}
}
// Read returns a processEntry populated from /proc/<pid>. Fields that cannot
// be read within the deadline are left at their zero value; the caller still
// gets an entry with whatever was retrievable. Returns ErrProcessGone when
// the /proc/<pid> directory does not exist.
func (r *ProcReader) Read(pid int) (processEntry, error) {
dir := filepath.Join(r.root, strconv.Itoa(pid))
if _, err := os.Stat(dir); err != nil {
if errors.Is(err, fs.ErrNotExist) {
return processEntry{}, ErrProcessGone
}
return processEntry{}, err
}
e := processEntry{PID: pid, ProcRead: true}
if data, ok := readFileWithDeadline(filepath.Join(dir, "status"), r.perFileDeadline); ok {
e.PPID = parseStatusPPID(string(data))
e.UID, e.UIDKnown = parseStatusUIDKnown(string(data))
e.Comm = parseStatusName(string(data))
}
if data, ok := readFileWithDeadline(filepath.Join(dir, "cmdline"), r.perFileDeadline); ok {
e.Cmdline = parseCmdline(data)
}
if target, ok := readlinkWithDeadline(filepath.Join(dir, "exe"), r.perFileDeadline); ok {
e.Exe = target
}
if data, ok := readFileWithDeadline(filepath.Join(dir, "stat"), r.perFileDeadline); ok {
if t, ok := r.parseStartedAt(data); ok {
e.StartedAt = t
}
}
return e, nil
}
// ReadStartedAt returns only /proc/<pid>/stat's process start time. It lets
// detector hot paths capture a lightweight PID-reuse token without reading
// status, cmdline, exe, or identity data.
func (r *ProcReader) ReadStartedAt(pid int) (time.Time, bool) {
if pid <= 0 {
return time.Time{}, false
}
path := filepath.Join(r.root, strconv.Itoa(pid), "stat")
data, ok := readFileWithDeadline(path, r.perFileDeadline)
if !ok {
return time.Time{}, false
}
return r.parseStartedAt(data)
}
// procStatStartTime extracts field 22 of /proc/<pid>/stat (starttime in
// clock ticks since boot). Field positions are deterministic except
// that the second field (comm) is parenthesized and may contain
// arbitrary bytes including spaces -- so we anchor on the final ")"
// before splitting the rest.
func procStatStartTime(data []byte) (int64, bool) {
end := bytes.LastIndexByte(data, ')')
if end < 0 || end+1 >= len(data) {
return 0, false
}
rest := strings.TrimSpace(string(data[end+1:]))
fields := strings.Fields(rest)
// rest starts at field 3 (state); starttime is field 22, i.e. index 19 in rest.
const starttimeIdx = 19
if len(fields) <= starttimeIdx {
return 0, false
}
v, err := strconv.ParseInt(fields[starttimeIdx], 10, 64)
if err != nil || v < 0 {
return 0, false
}
return v, true
}
// parseStartedAt converts /proc/<pid>/stat's starttime field into an
// absolute time using the host's boot time. Returns (zero, false) on
// any parse or btime resolution failure so callers leave the field
// unset rather than emitting bogus timestamps.
func (r *ProcReader) parseStartedAt(stat []byte) (time.Time, bool) {
ticks, ok := procStatStartTime(stat)
if !ok {
return time.Time{}, false
}
boot, ok := r.bootTime()
if !ok {
return time.Time{}, false
}
hz := clockTicksPerSecond()
if hz <= 0 {
return time.Time{}, false
}
sec := ticks / hz
rem := ticks % hz
ns := rem * int64(time.Second) / hz
return boot.Add(time.Duration(sec)*time.Second + time.Duration(ns)), true
}
// readFileWithDeadline reads up to 4 KiB from path; returns (data, true) on
// success or (nil, false) on any error or deadline expiry. /proc files are
// small; using ReadFile keeps the normal path simple. The generic deadline
// helper is tested with an injected slow function instead of a FIFO because
// general filesystem opens cannot be cancelled safely on every platform.
func readFileWithDeadline(path string, d time.Duration) ([]byte, bool) {
return runBytesWithDeadline(d, func() ([]byte, error) {
// #nosec G304 -- path is constructed from ProcReader.root + numeric PID;
// callers only pass procfs entries under r.root.
data, err := os.ReadFile(path)
if len(data) > 4096 {
data = data[:4096]
}
return data, err
})
}
// procReadConcurrency bounds how many deadline-bound /proc reads run at once.
// A blocking syscall goroutine cannot be cancelled in Go, so a wedged /proc
// entry (NFS-backed, D-state) leaks its goroutine until the kernel returns --
// which may be never. The cap turns what was an unbounded leak under PID-reuse
// churn into a fixed ceiling: once it is reached, further reads fail fast
// instead of spawning more abandonable goroutines. A goroutine releases its
// slot only when its syscall finally returns, so genuinely-stuck reads keep
// their slot (correctly counting against the ceiling).
const procReadConcurrency = 64
func runBytesWithDeadline(d time.Duration, fn func() ([]byte, error)) ([]byte, bool) {
return runProcReadWithDeadline(procReads, d, fn)
}
// readlinkWithDeadline runs Readlink in a goroutine and gives up after d.
func readlinkWithDeadline(path string, d time.Duration) (string, bool) {
return runProcReadWithDeadline(procReads, d, func() (string, error) { return os.Readlink(path) })
}
func parseStatusName(s string) string {
for _, line := range strings.Split(s, "\n") {
if rest, ok := strings.CutPrefix(line, "Name:\t"); ok {
return strings.TrimSpace(rest)
}
}
return ""
}
func parseStatusPPID(s string) int {
for _, line := range strings.Split(s, "\n") {
if rest, ok := strings.CutPrefix(line, "PPid:\t"); ok {
v, _ := strconv.Atoi(strings.TrimSpace(rest))
return v
}
}
return 0
}
func parseStatusUID(s string) int {
uid, _ := parseStatusUIDKnown(s)
return uid
}
func parseStatusUIDKnown(s string) (int, bool) {
for _, line := range strings.Split(s, "\n") {
if rest, ok := strings.CutPrefix(line, "Uid:\t"); ok {
fields := strings.Fields(rest)
if len(fields) == 0 {
return 0, false
}
v, err := strconv.Atoi(fields[0])
return v, err == nil
}
}
return 0, false
}
func parseCmdline(b []byte) []string {
if len(b) == 0 {
return nil
}
parts := bytes.Split(b, []byte{0})
out := make([]string, 0, len(parts))
for _, p := range parts {
if len(p) == 0 {
continue
}
out = append(out, string(p))
}
if len(out) == 0 {
return nil
}
return sanitizeCmdline(out)
}
const maxCmdlineArgLen = 256
var sensitiveCmdlineKeys = []string{"password", "passwd", "secret", "token", "api_key", "apikey"}
func sanitizeCmdline(args []string) []string {
out := make([]string, 0, len(args))
redactNext := false
for _, arg := range args {
if redactNext {
out = append(out, "<redacted>")
redactNext = false
continue
}
lower := strings.ToLower(arg)
redacted := false
for _, key := range sensitiveCmdlineKeys {
if strings.Contains(lower, key+"=") {
prefix, _, _ := strings.Cut(arg, "=")
out = append(out, truncateCmdlineArg(prefix+"=<redacted>"))
redacted = true
break
}
if lower == "--"+key || lower == "-"+key {
out = append(out, truncateCmdlineArg(arg))
redactNext = true
redacted = true
break
}
}
if redacted {
continue
}
out = append(out, truncateCmdlineArg(arg))
}
return out
}
func truncateCmdlineArg(arg string) string {
if len(arg) <= maxCmdlineArgLen {
return arg
}
return arg[:maxCmdlineArgLen]
}
package processctx
import (
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"time"
)
// bootTimeCache caches the host boot time per proc root.
// Boot time is invariant for the life of a kernel; rereading /proc/stat
// on every Read() would only spend syscalls for an unchanging value.
type bootTimeCache struct {
once sync.Once
t time.Time
ok bool
}
var procReaderBoot sync.Map // root -> *bootTimeCache
func (r *ProcReader) bootTime() (time.Time, bool) {
v, _ := procReaderBoot.LoadOrStore(r.root, &bootTimeCache{})
c := v.(*bootTimeCache)
c.once.Do(func() {
c.t, c.ok = readBootTime(filepath.Join(r.root, "stat"))
})
return c.t, c.ok
}
func readBootTime(statPath string) (time.Time, bool) {
// #nosec G304 -- statPath is procReader.root + "stat"; root is operator-pinned.
data, err := os.ReadFile(statPath)
if err != nil {
return time.Time{}, false
}
for _, line := range strings.Split(string(data), "\n") {
if rest, ok := strings.CutPrefix(line, "btime "); ok {
sec, err := strconv.ParseInt(strings.TrimSpace(rest), 10, 64)
if err != nil {
return time.Time{}, false
}
return time.Unix(sec, 0), true
}
}
return time.Time{}, false
}
// clockTicksPerSecondOverride lets tests inject a known _SC_CLK_TCK
// without calling sysconf on the host.
var clockTicksPerSecondOverride int64
func clockTicksPerSecond() int64 {
if clockTicksPerSecondOverride > 0 {
return clockTicksPerSecondOverride
}
return defaultClockTicksPerSecond
}
// defaultClockTicksPerSecond is 100 on every mainstream Linux kernel
// CSM runs on (CONFIG_HZ_100=y). Reading sysconf would require cgo or
// a build-tagged platform shim; the constant is correct for every
// distribution kernel we target. Tests that want a different value
// set clockTicksPerSecondOverride.
const defaultClockTicksPerSecond = 100
// Package processhandle signals verified Linux processes through kernel handles.
package processhandle
import (
"context"
"errors"
"fmt"
"math"
"syscall"
)
var ErrUnsupported = errors.New("safe process signaling requires kernel pidfd support")
// Signal acquires a process handle before running verify. Verification may read
// numeric procfs paths: exit checks surrounding it reject a recycled PID. The
// final signal uses the captured handle even if the process exits afterward.
func Signal(ctx context.Context, pid int, sig syscall.Signal, verify func() error) error {
if err := ctx.Err(); err != nil {
return err
}
if pid <= 1 || pid > math.MaxInt32 {
return fmt.Errorf("refuse to signal PID %d", pid)
}
handle, err := openHandle(pid)
if err != nil {
return err
}
defer handle.close()
if err := handle.alive(); err != nil {
return err
}
if err := verify(); err != nil {
return err
}
if err := handle.alive(); err != nil {
return err
}
if err := ctx.Err(); err != nil {
return err
}
return handle.signal(sig)
}
// Available reports whether this kernel can pin and signal a process handle.
// Probe each snapshot: descriptor pressure or service restrictions can change
// while the daemon runs, even though the kernel itself remains the same.
func Available() error {
return probe()
}
//go:build linux
package processhandle
import (
"bytes"
"errors"
"fmt"
"io"
"os"
"strconv"
"syscall"
"golang.org/x/sys/unix"
)
var (
pidfdOpen = unix.PidfdOpen
pidfdPoll = unix.Poll
pidfdSend = unix.PidfdSendSignal
pidfdClose = unix.Close
procRoot = "/proc"
)
// maxProcStatBytes bounds the tenant-visible /proc/<pid>/stat read. The comm
// field is capped at 16 bytes by the kernel, so a real record is far smaller.
const maxProcStatBytes = 4096
// procfs reports whether fd is a /proc/<pid> directory descriptor rather than
// a descriptor returned by pidfd_open. Both pin one struct pid, so
// pidfd_send_signal cannot be redirected to a recycled PID through either.
type handle struct {
fd int
procfs bool
}
// openHandle pins the target process. pidfd_open needs Linux 5.3; the 4.18
// kernels on EL8 and CloudLinux 8 do not have it, but their pidfd_send_signal
// (Linux 5.1) accepts a /proc/<pid> directory descriptor, which pins the same
// struct pid. Falling back preserves the identity guarantee instead of
// disabling termination on every supported EL8 host.
func openHandle(pid int) (*handle, error) {
fd, err := pidfdOpen(pid, 0)
if err == nil {
return &handle{fd: fd}, nil
}
if !errors.Is(err, unix.ENOSYS) {
return nil, signalError("pidfd_open", err)
}
dirfd, dirErr := unix.Open(procRoot+"/"+strconv.Itoa(pid), unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC, 0)
if dirErr != nil {
if errors.Is(dirErr, unix.ENOENT) || errors.Is(dirErr, unix.ESRCH) {
return nil, fmt.Errorf("open %s/%d: %w: %w", procRoot, pid, os.ErrProcessDone, dirErr)
}
return nil, fmt.Errorf("open %s/%d: %w (pidfd_open: %w)", procRoot, pid, dirErr, err)
}
return &handle{fd: dirfd, procfs: true}, nil
}
func (h *handle) close() { _ = pidfdClose(h.fd) }
func (h *handle) alive() error {
if h.procfs {
return h.procfsAlive()
}
// #nosec G115 -- Linux file descriptors are non-negative signed C ints.
fds := []unix.PollFd{{Fd: int32(h.fd), Events: unix.POLLIN}}
if _, err := pidfdPoll(fds, 0); err != nil {
return fmt.Errorf("poll process handle: %w", err)
}
if fds[0].Revents&(unix.POLLIN|unix.POLLHUP) != 0 {
return os.ErrProcessDone
}
if fds[0].Revents != 0 {
return fmt.Errorf("invalid process handle poll events: %#x", fds[0].Revents)
}
return nil
}
// procfsAlive reads the pinned directory's own stat file. A reaped process
// removes those entries, and an unreaped one reports a terminal state, so a
// replacement holding the same numeric PID is never mistaken for the target.
func (h *handle) procfsAlive() error {
fd, err := unix.Openat(h.fd, "stat", unix.O_RDONLY|unix.O_CLOEXEC, 0)
if err != nil {
if errors.Is(err, unix.ENOENT) || errors.Is(err, unix.ESRCH) {
return os.ErrProcessDone
}
return fmt.Errorf("open process state: %w", err)
}
// #nosec G115 -- Successful openat returns a non-negative file descriptor.
file := os.NewFile(uintptr(fd), "stat")
defer func() { _ = file.Close() }()
data, err := io.ReadAll(io.LimitReader(file, maxProcStatBytes))
if err != nil {
if errors.Is(err, unix.ESRCH) {
return os.ErrProcessDone
}
return fmt.Errorf("read process state: %w", err)
}
state, err := procStatState(data)
if err != nil {
return err
}
// Z: exited, awaiting reap. X/x: released. Neither can receive a signal.
if state == 'Z' || state == 'X' || state == 'x' {
return os.ErrProcessDone
}
return nil
}
// procStatState returns the third field of /proc/<pid>/stat. The second field
// is a process-controlled name in parentheses that may itself contain spaces
// and parentheses, so the scan starts at its final closing parenthesis.
func procStatState(data []byte) (byte, error) {
end := bytes.LastIndexByte(data, ')')
if end < 0 || end+2 >= len(data) || data[end+1] != ' ' {
return 0, errors.New("unreadable process state")
}
return data[end+2], nil
}
func (h *handle) signal(sig syscall.Signal) error {
if err := pidfdSend(h.fd, sig, nil, 0); err != nil {
return signalError("pidfd_send_signal", err)
}
return nil
}
func signalError(operation string, err error) error {
if errors.Is(err, unix.ENOSYS) {
return fmt.Errorf("%s: %w: %w", operation, ErrUnsupported, err)
}
if errors.Is(err, unix.ESRCH) {
return fmt.Errorf("%s: %w: %w", operation, os.ErrProcessDone, err)
}
return fmt.Errorf("%s: %w", operation, err)
}
// probe pins this process and asks the kernel to deliver signal 0, which
// validates the whole path (handle acquisition, liveness, delivery) without
// disturbing any process.
func probe() error {
handle, err := openHandle(os.Getpid())
if err != nil {
return err
}
defer handle.close()
if err := handle.alive(); err != nil {
return err
}
return handle.signal(0)
}
// Package quarantinefs makes recovery copies durable before callers change
// live files. Its paths must be under an operator-owned quarantine directory.
package quarantinefs
import (
"bytes"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"golang.org/x/sys/unix"
)
var (
copyContent = io.Copy
syncFile = (*os.File).Sync
closeFile = (*os.File).Close
)
// EnsureDir persists newly created directory entries, including missing parents.
func EnsureDir(path string, mode os.FileMode) error {
info, err := os.Stat(path)
if err == nil {
if !info.IsDir() {
return fmt.Errorf("quarantine path is not a directory: %s", path)
}
// Another quarantine transaction may have just created this entry.
return SyncDir(filepath.Dir(path))
}
if !os.IsNotExist(err) {
return err
}
parent := filepath.Dir(path)
if err := EnsureDir(parent, mode); err != nil {
return err
}
if err := os.Mkdir(path, mode); err != nil && !os.IsExist(err) {
return err
}
return SyncDir(parent)
}
// WriteExclusive never overwrites older evidence. Success includes the file's
// contents, its close result, and the containing directory entry reaching disk.
// The caller must first persist the containing directory with EnsureDir.
func WriteExclusive(path string, content io.Reader, mode os.FileMode) (err error) {
// #nosec G304 -- caller supplies a path under operator-owned quarantine; exclusive creation refuses existing files and symlinks.
f, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_EXCL, mode)
if err != nil {
return err
}
closed := false
defer func() {
if !closed {
err = errors.Join(err, closeFile(f))
}
if err != nil {
if removeErr := os.Remove(path); removeErr != nil && !os.IsNotExist(removeErr) {
err = errors.Join(err, fmt.Errorf("partial copy retained at %s: %w", path, removeErr))
}
}
}()
if _, err := copyContent(f, content); err != nil {
return fmt.Errorf("writing %s: %w", path, err)
}
if err := syncFile(f); err != nil {
return fmt.Errorf("syncing %s: %w", path, err)
}
closed = true
if err := closeFile(f); err != nil {
return fmt.Errorf("closing %s: %w", path, err)
}
return SyncDir(filepath.Dir(path))
}
// Store writes a private content copy and its sidecar before the caller may
// remove or replace the original. On failure the caller must retain the original.
func Store(path string, content io.Reader, metadata []byte, mode os.FileMode) error {
if err := EnsureDir(filepath.Dir(path), 0700); err != nil {
return err
}
if err := WriteExclusive(path, content, mode); err != nil {
return err
}
if err := WriteExclusive(path+".meta", bytes.NewReader(metadata), 0600); err != nil {
if removeErr := os.Remove(path); removeErr != nil && !os.IsNotExist(removeErr) {
return errors.Join(err, fmt.Errorf("copy retained at %s: %w", path, removeErr))
}
return fmt.Errorf("writing quarantine metadata: %w", err)
}
return nil
}
func SyncDir(path string) error {
// #nosec G304 -- caller supplies the containing directory of a quarantine or restored file, opened read-only to persist its entries.
dir, err := os.Open(path)
if err != nil {
return err
}
return errors.Join(syncFile(dir), closeFile(dir))
}
func SyncFilePath(path string) error {
// #nosec G304 -- caller supplies a quarantine or remediation path; no-follow and regular-file checks reject substituted links and special files.
f, err := os.OpenFile(path, os.O_RDONLY|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if err != nil {
return err
}
info, err := f.Stat()
if err != nil {
return errors.Join(err, closeFile(f))
}
if !info.Mode().IsRegular() {
return errors.Join(fmt.Errorf("cannot sync non-regular file %s", path), closeFile(f))
}
return errors.Join(syncFile(f), closeFile(f))
}
// RemoveEvidence runs only after a restored copy and its directory are durable.
// A failed content removal leaves its metadata intact for another recovery attempt.
func RemoveEvidence(path, metaPath string) error {
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("restored, but quarantine content remains at %s: %w", path, err)
}
if err := SyncDir(filepath.Dir(path)); err != nil {
return fmt.Errorf("restored, but quarantine removal is not durable; metadata retained at %s: %w", metaPath, err)
}
if err := os.Remove(metaPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("restored, but quarantine metadata remains at %s: %w", metaPath, err)
}
return SyncDir(filepath.Dir(metaPath))
}
package quarantinefs
import (
"errors"
"fmt"
"io/fs"
"os"
"golang.org/x/sys/unix"
)
// SyncTree persists a directory being moved without following symlinks inside
// it. expected binds the traversal to the directory admitted by the caller.
func SyncTree(path string, expected os.FileInfo) error {
root, err := os.OpenRoot(path)
if err != nil {
return err
}
defer func() { _ = root.Close() }()
info, err := root.Stat(".")
if err != nil {
return err
}
if !info.IsDir() || !os.SameFile(info, expected) {
return fmt.Errorf("quarantine directory changed before syncing")
}
var directories []string
err = fs.WalkDir(root.FS(), ".", func(name string, entry fs.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if entry.IsDir() {
directories = append(directories, name)
return nil
}
if entry.Type()&os.ModeSymlink != 0 {
return nil
}
file, openErr := root.OpenFile(name, os.O_RDONLY|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0)
if openErr != nil {
return openErr
}
info, statErr := file.Stat()
if statErr != nil || !info.Mode().IsRegular() {
_ = file.Close()
return fmt.Errorf("cannot sync non-regular quarantine entry %s: %v", name, statErr)
}
return errors.Join(syncFile(file), closeFile(file))
})
if err != nil {
return err
}
// Persist child contents before the directory entries that make them reachable.
for i := len(directories) - 1; i >= 0; i-- {
dir, openErr := root.OpenFile(directories[i], os.O_RDONLY|unix.O_DIRECTORY|unix.O_NOFOLLOW, 0)
if openErr != nil {
return openErr
}
if err := errors.Join(syncFile(dir), closeFile(dir)); err != nil {
return err
}
}
return nil
}
package queuehealth
import "time"
// Work retains accounting after a channel receive. The consumer starts its
// ticket, then finishes or rejects it after processing, including error paths.
type Work[T any] struct {
Value T
Ticket Ticket
}
// Process keeps filtered work accounted through the consumer's whole callback.
// A panicking consumer loses this item; the panic continues to its owner.
func (w Work[T]) Process(fn func(T)) {
w.Ticket.Start(time.Now())
completed := false
defer func() {
if completed {
w.Ticket.Finish(time.Now())
} else {
w.Ticket.Reject(time.Now())
}
}()
fn(w.Value)
completed = true
}
// Channel accounts for bounded waiting work and its consumers. Its owner
// closes it after all producers have stopped; stopped consumers must leave
// queued work to DiscardPending instead of silently abandoning it.
type Channel[T any] struct {
items chan Work[T]
tracker *Tracker
now func() time.Time
}
func NewChannel[T any](capacity int, maxLag time.Duration) *Channel[T] {
return &Channel[T]{
items: make(chan Work[T], capacity), tracker: New(capacity, maxLag), now: time.Now,
}
}
func (q *Channel[T]) Items() <-chan Work[T] { return q.items }
func (q *Channel[T]) TrySend(value T) bool {
work := Work[T]{Value: value, Ticket: q.tracker.Begin(q.now())}
select {
case q.items <- work:
return true
default:
work.Ticket.Reject(q.now())
return false
}
}
func (q *Channel[T]) Send(value T, stop <-chan struct{}) bool {
work := Work[T]{Value: value, Ticket: q.tracker.Begin(q.now())}
select {
case q.items <- work:
return true
case <-stop:
work.Ticket.Reject(q.now())
return false
}
}
func (q *Channel[T]) Close() { close(q.items) }
// DiscardPending runs after Close and after consumers have stopped. A work
// item already received by a consumer still belongs to that consumer.
func (q *Channel[T]) DiscardPending() {
for work := range q.items {
work.Ticket.Reject(q.now())
}
}
func (q *Channel[T]) Snapshot(now time.Time) Status { return q.tracker.Snapshot(now) }
// Lose counts work rejected before it could be represented by a queued item.
func (q *Channel[T]) Lose(count uint64) { q.tracker.Lose(q.now(), count) }
package queuehealth
import "time"
// MeasurementWindow is how long a queue must stay unmeasurable before that
// counts as a degradation. A kernel counter read can fail once, and positions
// loaded separately can produce one incoherent sample, without the queue
// having a problem.
const MeasurementWindow = 30 * time.Second
// Dwell reports whether a condition has held continuously for a window. The
// owner serializes its calls with the rest of its snapshot.
type Dwell struct {
since time.Time
}
func (d *Dwell) Held(now time.Time, condition bool, window time.Duration) bool {
if !condition {
d.since = time.Time{}
return false
}
if d.since.IsZero() {
d.since = now
}
return now.Sub(d.since) >= window
}
package queuehealth
import "fmt"
// Evidence labels sampled and unavailable measurements so notifications and
// doctor cannot present bytes as events or consumer stalls as event ages.
func (s Status) Evidence() string {
unit := s.DepthUnit
if unit == "" {
unit = "items"
}
pending, capacity := fmt.Sprint(s.Depth), fmt.Sprint(s.Capacity)
if s.DepthUnavailable {
pending = "unknown"
}
if s.CapacityUnavailable {
capacity = "unknown"
}
depth := pending + "/" + capacity
lag := fmt.Sprintf("lag=%.0fs", s.LagSeconds)
switch s.LagBasis {
case "observed_age":
lag = fmt.Sprintf("observed_lag=%.0fs", s.LagSeconds)
case "consumer_progress":
lag = fmt.Sprintf("consumer_stall=%.0fs", s.LagSeconds)
case "operation_progress":
lag = fmt.Sprintf("operation_stall=%.0fs", s.LagSeconds)
case "deferred_checkpoint":
lag = fmt.Sprintf("deferred_age=%.0fs", s.LagSeconds)
case "unavailable":
lag = "lag=unavailable"
}
dropped := fmt.Sprintf("dropped=%d", s.DroppedTotal)
if s.DroppedLowerBound {
dropped = fmt.Sprintf("dropped>=%d", s.DroppedTotal)
}
return fmt.Sprintf("depth=%s %s running=%d %s recent_drops=%d %s processing=%.0fs", depth, unit, s.InFlight, dropped, s.RecentDrops, lag, s.ProcessingSeconds)
}
package queuehealth
import (
"sort"
"time"
)
// reminderInterval bounds how often one queue can announce. It also bounds a
// flapping queue: a degradation that returns within the interval of the last
// announcement is carried in status but not notified again.
const reminderInterval = 5 * time.Minute
// Event records a degradation, a bounded reminder or a recovery. Healthy
// startup produces no event, and one queue cannot suppress another's change.
type Event struct {
Name string
Current Status
Recovered bool
}
// incident is one queue's notification state. It outlives the degradation so a
// queue that recovers and degrades again stays inside the same bound.
type incident struct {
announcedAt time.Time
announced bool
degraded bool
}
// Reporter belongs to one health loop. Polling status does not consume its
// state or the trackers' counters, so API traffic cannot swallow an alert.
type Reporter struct {
last map[string]incident
}
func (r *Reporter) Events(now time.Time, states map[string]Status) []Event {
if r.last == nil {
r.last = make(map[string]incident)
}
names := make([]string, 0, len(states))
for name := range states {
names = append(names, name)
}
sort.Strings(names)
var events []Event
for _, name := range names {
s := states[name]
state := r.last[name]
switch {
case s.Advisory:
// Best-effort work is visible in status and doctor only.
continue
case s.Status == "degraded":
if state.announcedAt.IsZero() || now.Sub(state.announcedAt) >= reminderInterval {
events = append(events, Event{Name: name, Current: s})
state.announcedAt, state.announced = now, true
} else if !state.degraded {
// A degradation the bound suppressed must not announce a
// recovery either, or the pair count is unchanged.
state.announced = false
}
state.degraded = true
r.last[name] = state
case state.degraded:
if state.announced {
events = append(events, Event{Name: name, Current: s, Recovered: true})
state.announced = false
}
state.degraded = false
r.last[name] = state
case !state.announcedAt.IsZero() && now.Sub(state.announcedAt) >= reminderInterval:
delete(r.last, name)
}
}
return events
}
package queuehealth
import (
"sync"
"time"
)
// Sampled measures opaque queues whose entries cannot carry tickets. Its lag
// is time since observed consumer progress while work remains, not the age of
// an unseen event. Owners sample depth and a monotonic consumption counter
// together; observing new arrivals alone must not reset a stalled consumer.
type Sampled struct {
mu sync.Mutex
capacity int
unit string
maxLag time.Duration
depth int
progress uint64
stalledSince time.Time
fullSince time.Time
losses *Tracker
}
func NewSampled(capacity int, unit string, maxLag time.Duration) *Sampled {
return &Sampled{capacity: capacity, unit: unit, maxLag: maxLag, losses: New(0, maxLag)}
}
func (q *Sampled) Observe(now time.Time, depth int, progress uint64) {
q.mu.Lock()
defer q.mu.Unlock()
if depth == 0 {
q.stalledSince = time.Time{}
} else if q.stalledSince.IsZero() || progress != q.progress {
q.stalledSince = now
}
if q.capacity > 0 && depth >= q.capacity {
if q.fullSince.IsZero() {
q.fullSince = now
}
} else {
q.fullSince = time.Time{}
}
q.depth, q.progress = depth, progress
}
func (q *Sampled) Lose(now time.Time, count uint64) { q.losses.Lose(now, count) }
// SnapshotUnavailable retains independently counted losses without deriving
// pressure from an occupancy reading the owner could not validate.
func (q *Sampled) SnapshotUnavailable(now time.Time) Status {
q.mu.Lock()
defer q.mu.Unlock()
s := q.losses.Snapshot(now)
s.Capacity, s.DepthUnit = q.capacity, q.unit
s.DepthUnavailable, s.LagBasis = true, "unavailable"
return s
}
func (q *Sampled) Snapshot(now time.Time) Status {
q.mu.Lock()
defer q.mu.Unlock()
s := q.losses.Snapshot(now)
s.Depth, s.Capacity, s.DepthUnit = q.depth, q.capacity, q.unit
s.LagBasis = "consumer_progress"
if !q.stalledSince.IsZero() {
s.LagSeconds = max(0, now.Sub(q.stalledSince).Seconds())
}
switch {
case s.LagSeconds >= q.maxLag.Seconds():
s.Reason = "consumer_stalled"
case !q.fullSince.IsZero() && now.Sub(q.fullSince) >= fullWindow:
s.Reason = "queue_full"
}
if s.Reason != "" {
s.Status = "degraded"
}
return s
}
// Package queuehealth measures waiting, running and lost work independently
// of the channel that carries findings about a protection failure.
package queuehealth
import (
"sync"
"time"
)
const (
dropWindow = time.Minute
fullWindow = 30 * time.Second
dropThreshold = 3
)
// Status is the queue evidence carried by status, the API and doctor.
// LagSeconds measures the oldest waiting item unless LagBasis names another
// measurement. ProcessingSeconds measures
// the oldest running item, so an empty queue cannot conceal a stuck worker.
// Advisory marks a queue whose work is best effort, so its degradation is
// reported but does not make the host degraded and raises no notification.
type Status struct {
Status string `json:"status"`
Reason string `json:"reason,omitempty"`
Advisory bool `json:"advisory,omitempty"`
Depth int `json:"depth"`
DepthUnit string `json:"depth_unit,omitempty"`
DepthUnavailable bool `json:"depth_unavailable,omitempty"`
LagBasis string `json:"lag_basis,omitempty"`
Capacity int `json:"capacity"`
CapacityUnavailable bool `json:"capacity_unavailable,omitempty"`
InFlight int `json:"in_flight"`
DroppedTotal uint64 `json:"dropped_total"`
DroppedLowerBound bool `json:"dropped_lower_bound,omitempty"`
RecentDrops uint64 `json:"recent_drops"`
LagSeconds float64 `json:"lag_seconds"`
ProcessingSeconds float64 `json:"processing_seconds"`
}
type work struct {
queued time.Time
started time.Time
}
type dropBucket struct {
second int64
count uint64
}
// Tracker accounts for a bounded queue and its workers. Every Begin must
// end in Finish or Reject, including cancellation and panic paths. Tracking
// starts before a send so a fast consumer cannot finish an unrecorded item.
// Its retained work is bounded by the queue and the producer/worker counts.
type Tracker struct {
mu sync.Mutex
capacity int
maxLag time.Duration
sharedCapacity bool
next uint64
pending map[uint64]work
waiting int
fullSince time.Time
heldSince time.Time
dropped uint64
drops [60]dropBucket
}
func New(capacity int, maxLag time.Duration) *Tracker {
return &Tracker{capacity: capacity, maxLag: maxLag, pending: make(map[uint64]work)}
}
// NewSharedCapacity measures queues whose running work still occupies
// admission slots. Ordinary channels free those slots when a consumer reads.
func NewSharedCapacity(capacity int, maxLag time.Duration) *Tracker {
q := New(capacity, maxLag)
q.sharedCapacity = true
return q
}
// Ticket follows one item from enqueue through processing. A zero ticket
// represents synchronous work that never entered a queue.
type Ticket struct {
tracker *Tracker
id uint64
}
func (q *Tracker) Begin(now time.Time) Ticket {
return q.BeginAt(now, now)
}
// BeginAt sets the waiting-age origin independently of admission. Earlier
// delays must not backdate saturation; a future origin defers lag until the
// work is eligible to run.
func (q *Tracker) BeginAt(queuedAt, now time.Time) Ticket {
q.mu.Lock()
defer q.mu.Unlock()
q.next++
q.pending[q.next] = work{queued: queuedAt}
q.waiting++
q.updateFull(now)
return Ticket{tracker: q, id: q.next}
}
func (t Ticket) Start(now time.Time) {
if t.tracker == nil {
return
}
q := t.tracker
q.mu.Lock()
defer q.mu.Unlock()
w := q.pending[t.id]
w.started = now
q.pending[t.id] = w
q.waiting--
q.updateFull(now)
}
// Requeue returns a running item to waiting without resetting its admission
// time. Repeated attempts must not conceal a backlog that never completes.
func (t Ticket) Requeue(now time.Time) {
if t.tracker == nil {
return
}
q := t.tracker
q.mu.Lock()
defer q.mu.Unlock()
w := q.pending[t.id]
w.started = time.Time{}
q.pending[t.id] = w
q.waiting++
q.updateFull(now)
}
// RetainQueuedAt preserves earlier eligibility when observations coalesce.
// A later observation must never postpone work that was already overdue.
func (t Ticket) RetainQueuedAt(queuedAt time.Time) {
if t.tracker == nil {
return
}
q := t.tracker
q.mu.Lock()
defer q.mu.Unlock()
w := q.pending[t.id]
if queuedAt.Before(w.queued) {
w.queued = queuedAt
q.pending[t.id] = w
}
}
// MergeRunning absorbs a distinct running ticket on the same tracker into
// this waiting ticket. The caller retains only this ticket. Keeping the
// older age and completing the running ticket atomically avoids inventing
// an available waiting slot while coalescing a retry with new work.
func (t Ticket) MergeRunning(running Ticket, now time.Time) {
if t.tracker == nil {
return
}
q := t.tracker
q.mu.Lock()
defer q.mu.Unlock()
w := q.pending[t.id]
if earlier := q.pending[running.id].queued; earlier.Before(w.queued) {
w.queued = earlier
}
q.pending[t.id] = w
q.finish(running.id, now)
}
func (t Ticket) Finish(now time.Time) {
if t.tracker == nil {
return
}
q := t.tracker
q.mu.Lock()
defer q.mu.Unlock()
q.finish(t.id, now)
}
func (t Ticket) Reject(now time.Time) {
if t.tracker == nil {
return
}
q := t.tracker
q.mu.Lock()
defer q.mu.Unlock()
q.finish(t.id, now)
q.lose(now, 1)
}
func (q *Tracker) finish(id uint64, now time.Time) {
if q.pending[id].started.IsZero() {
q.waiting--
}
delete(q.pending, id)
q.updateFull(now)
}
func (q *Tracker) updateFull(now time.Time) {
occupied := q.waiting
if q.sharedCapacity {
occupied = len(q.pending)
}
if q.capacity > 0 && occupied >= q.capacity {
if q.fullSince.IsZero() {
q.fullSince = now
}
} else {
q.fullSince = time.Time{}
}
}
// Hold parks the queue while its owner deliberately withholds the consumer,
// such as the startup baseline. Ages stop at the hold so intentional waiting
// cannot report a stall; losses still count.
func (q *Tracker) Hold(now time.Time) {
q.mu.Lock()
defer q.mu.Unlock()
if q.heldSince.IsZero() {
q.heldSince = now
}
}
// Release resumes measurement. The held time is removed from every age, so
// work parked by the hold is measured from the release.
func (q *Tracker) Release(now time.Time) {
q.mu.Lock()
defer q.mu.Unlock()
if q.heldSince.IsZero() {
return
}
shift := now.Sub(q.heldSince)
rebase := func(t time.Time) time.Time {
switch {
case t.IsZero():
return t
case t.After(now):
// Deferred work has not begun aging and keeps its eligibility.
return t
case t.Before(q.heldSince):
return t.Add(shift)
default:
return now
}
}
for id, w := range q.pending {
w.queued, w.started = rebase(w.queued), rebase(w.started)
q.pending[id] = w
}
q.fullSince = rebase(q.fullSince)
q.heldSince = time.Time{}
}
// Lose records upstream loss for which no userspace work item exists, such
// as a kernel overflow. The cumulative count is never drained by a reporter.
func (q *Tracker) Lose(now time.Time, count uint64) {
q.mu.Lock()
defer q.mu.Unlock()
q.lose(now, count)
}
func (q *Tracker) lose(now time.Time, count uint64) {
q.dropped += count
second := now.Unix()
b := &q.drops[uint64(second)%uint64(len(q.drops))] // #nosec G115 -- modulo index is bounded even for a pre-epoch clock
if b.second != second {
*b = dropBucket{second: second}
}
b.count += count
}
func (q *Tracker) Snapshot(now time.Time) Status {
q.mu.Lock()
defer q.mu.Unlock()
s := Status{Status: "ok", Capacity: q.capacity, Depth: q.waiting, InFlight: len(q.pending) - q.waiting, DroppedTotal: q.dropped}
// A held queue measures ages up to the hold; drops keep the real clock.
clock := now
if !q.heldSince.IsZero() && q.heldSince.Before(now) {
clock = q.heldSince
}
for _, w := range q.pending {
if w.started.IsZero() {
s.LagSeconds = max(s.LagSeconds, clock.Sub(w.queued).Seconds())
} else {
s.ProcessingSeconds = max(s.ProcessingSeconds, clock.Sub(w.started).Seconds())
}
}
second := now.Unix()
for _, b := range q.drops {
if b.second <= second && second-b.second < int64(dropWindow/time.Second) {
s.RecentDrops += b.count
}
}
switch {
case s.LagSeconds >= q.maxLag.Seconds():
s.Reason = "backlog_lag"
case s.ProcessingSeconds >= q.maxLag.Seconds():
s.Reason = "processing_lag"
case !q.fullSince.IsZero() && clock.Sub(q.fullSince) >= fullWindow:
s.Reason = "queue_full"
case s.RecentDrops >= dropThreshold:
s.Reason = "dropped_work"
}
if s.Reason != "" {
s.Status = "degraded"
}
return s
}
// Package redisinfo wraps the go-redis client for the few read-only
// INFO calls CSM needs (memory metrics, keyspace counts). It replaces
// the `redis-cli info <section>` shell-outs the performance UI used
// to issue per-poll, eliminating libc/libpthread fork churn on hosts
// with a busy metrics dashboard.
//
// The client is a lazy package-level singleton: first call opens a
// connection-pooled client against the local redis socket / TCP, all
// subsequent calls reuse it. No Close path because the daemon runs
// for the host's lifetime; go-redis cleans up on process exit.
//
// Connection target matches redis-cli's default behaviour:
//
// 127.0.0.1:6379, no password, db 0
//
// When REDISCLI_AUTH is set, the client uses it as the password so
// daemon environments that previously made redis-cli work keep working.
// Hosts running redis on a non-default socket can override by calling
// SetAddr before the first MemoryUsage / Keyspace call. Absolute
// paths are treated as Unix sockets.
package redisinfo
import (
"context"
"fmt"
"os"
"strconv"
"strings"
"sync"
"time"
"github.com/redis/go-redis/v9"
)
const defaultAddr = "127.0.0.1:6379"
var (
mu sync.Mutex
defaultDB *redis.Client
clientBuilt bool
addr = defaultAddr
password = ""
)
// SetAddr overrides the connection target before the singleton opens.
// Calls after the first MemoryUsage / Keyspace are ignored: the
// singleton has already been built. Tests can also use SetClientForTest
// to swap the singleton wholesale.
func SetAddr(a, pwd string) {
mu.Lock()
defer mu.Unlock()
if clientBuilt {
return
}
addr = a
password = pwd
}
// SetClientForTest replaces the singleton with a caller-supplied
// client (typically pointed at miniredis or a real test instance).
// Pass nil to clear the override and let the next call rebuild the
// real singleton.
func SetClientForTest(c *redis.Client) {
mu.Lock()
defer mu.Unlock()
defaultDB = c
clientBuilt = c != nil
}
func client() *redis.Client {
mu.Lock()
defer mu.Unlock()
if defaultDB == nil {
defaultDB = redis.NewClient(redisOptions(addr, password))
clientBuilt = true
}
return defaultDB
}
func redisOptions(target, pwd string) *redis.Options {
network := "tcp"
if strings.HasPrefix(target, "/") {
network = "unix"
}
if pwd == "" {
pwd = os.Getenv("REDISCLI_AUTH")
}
return &redis.Options{
Network: network,
Addr: target,
Password: pwd,
DB: 0,
// Fail fast when redis is absent: this client only serves the
// metrics dashboard, never a hot path. Default MaxRetries=3 +
// ReadTimeout=3s would block the perfMetrics sampler for >9s on
// a host without redis (a normal config -- CSM does not require
// redis to run).
DialTimeout: 500 * time.Millisecond,
ReadTimeout: 500 * time.Millisecond,
WriteTimeout: 500 * time.Millisecond,
MaxRetries: -1, // disable retries entirely
PoolTimeout: 500 * time.Millisecond,
PoolSize: 2,
}
}
// MemoryUsage returns the redis used_memory and maxmemory values
// from `INFO memory`, in bytes. Either may be zero if the server
// omits the field. err non-nil only on connection / protocol error.
//
// Tests can intercept via SetMemoryUsageForTest.
func MemoryUsage(ctx context.Context) (used, max uint64, err error) {
if fn := getMemoryUsageMock(); fn != nil {
return fn(ctx)
}
c := client()
if c == nil {
return 0, 0, fmt.Errorf("redisinfo: client not initialised")
}
raw, err := c.Info(ctx, "memory").Result()
if err != nil {
return 0, 0, err
}
used, max = parseMemoryInfo(raw)
return used, max, nil
}
// MemoryUsageFunc is the signature SetMemoryUsageForTest accepts.
type MemoryUsageFunc func(ctx context.Context) (used, max uint64, err error)
// KeyspaceStatsFunc is the signature SetKeyspaceStatsForTest accepts.
type KeyspaceStatsFunc func(ctx context.Context) (KeyspaceStat, error)
// ConfigGetFunc is the signature SetConfigGetForTest accepts.
type ConfigGetFunc func(ctx context.Context, name string) (string, error)
var (
mockMu sync.RWMutex
memMock MemoryUsageFunc
keyMock KeyspaceStatsFunc
configMock ConfigGetFunc
)
// SetMemoryUsageForTest installs an interceptor for MemoryUsage. Pass
// nil to clear. Production code paths must NOT call this.
func SetMemoryUsageForTest(fn MemoryUsageFunc) {
mockMu.Lock()
defer mockMu.Unlock()
memMock = fn
}
// SetKeyspaceStatsForTest installs an interceptor for KeyspaceStats
// (and Keyspace, which proxies to it). Pass nil to clear.
func SetKeyspaceStatsForTest(fn KeyspaceStatsFunc) {
mockMu.Lock()
defer mockMu.Unlock()
keyMock = fn
}
// SetConfigGetForTest installs an interceptor for ConfigGet. Pass nil
// to clear.
func SetConfigGetForTest(fn ConfigGetFunc) {
mockMu.Lock()
defer mockMu.Unlock()
configMock = fn
}
func getMemoryUsageMock() MemoryUsageFunc {
mockMu.RLock()
defer mockMu.RUnlock()
return memMock
}
func getKeyspaceStatsMock() KeyspaceStatsFunc {
mockMu.RLock()
defer mockMu.RUnlock()
return keyMock
}
func getConfigGetMock() ConfigGetFunc {
mockMu.RLock()
defer mockMu.RUnlock()
return configMock
}
func parseMemoryInfo(raw string) (used, max uint64) {
for _, line := range strings.Split(raw, "\n") {
line = strings.TrimSpace(line)
switch {
case strings.HasPrefix(line, "used_memory:"):
used, _ = strconv.ParseUint(strings.TrimSpace(strings.TrimPrefix(line, "used_memory:")), 10, 64)
case strings.HasPrefix(line, "maxmemory:"):
max, _ = strconv.ParseUint(strings.TrimSpace(strings.TrimPrefix(line, "maxmemory:")), 10, 64)
}
}
return used, max
}
// Keyspace returns the sum of `keys=N` across every db<n> line in
// `INFO keyspace`. err non-nil only on connection / protocol error.
func Keyspace(ctx context.Context) (int64, error) {
stats, err := KeyspaceStats(ctx)
if err != nil {
return 0, err
}
return stats.TotalKeys, nil
}
// KeyspaceStat is the aggregated breakdown of `INFO keyspace`. Keys
// counts all keys across all dbs; Expires counts the subset with a
// TTL applied.
type KeyspaceStat struct {
TotalKeys int64
TotalExpires int64
}
// KeyspaceStats returns the per-db sums from `INFO keyspace`.
//
// Tests can intercept via SetKeyspaceStatsForTest.
func KeyspaceStats(ctx context.Context) (KeyspaceStat, error) {
if fn := getKeyspaceStatsMock(); fn != nil {
return fn(ctx)
}
c := client()
if c == nil {
return KeyspaceStat{}, fmt.Errorf("redisinfo: client not initialised")
}
raw, err := c.Info(ctx, "keyspace").Result()
if err != nil {
return KeyspaceStat{}, err
}
return parseKeyspaceStats(raw), nil
}
// ConfigGet returns the value of a single CONFIG GET parameter (e.g.
// "maxmemory", "save", "maxmemory-policy"). Empty result returns
// ("", nil) so callers can distinguish "unset" from "connection error".
//
// Tests can intercept via SetConfigGetForTest.
func ConfigGet(ctx context.Context, name string) (string, error) {
if fn := getConfigGetMock(); fn != nil {
return fn(ctx, name)
}
c := client()
if c == nil {
return "", fmt.Errorf("redisinfo: client not initialised")
}
m, err := c.ConfigGet(ctx, name).Result()
if err != nil {
return "", err
}
if v, ok := m[name]; ok {
return v, nil
}
return "", nil
}
func parseKeyspaceInfo(raw string) int64 {
return parseKeyspaceStats(raw).TotalKeys
}
func parseKeyspaceStats(raw string) KeyspaceStat {
var stat KeyspaceStat
for _, line := range strings.Split(raw, "\n") {
line = strings.TrimSpace(line)
if !strings.HasPrefix(line, "db") {
continue
}
parts := strings.SplitN(line, ":", 2)
if len(parts) < 2 {
continue
}
for _, kv := range strings.Split(parts[1], ",") {
kv = strings.TrimSpace(kv)
eq := strings.IndexByte(kv, '=')
if eq < 0 {
continue
}
val, perr := strconv.ParseInt(kv[eq+1:], 10, 64)
if perr != nil {
continue
}
switch kv[:eq] {
case "keys":
stat.TotalKeys += val
case "expires":
stat.TotalExpires += val
}
}
}
return stat
}
package reporting
import (
"context"
"sync/atomic"
)
// Action is the node's policy for acting on central scored-set data. It is
// deliberately conservative: central data never hard-blocks on its own.
type Action string
const (
// ActionOff consumes the set for visibility only; never acts.
ActionOff Action = "off"
// ActionChallenge elevates suspicion / serves a challenge for listed IPs.
ActionChallenge Action = "challenge"
// ActionBlockIfLocalCorroborated hard-blocks only when this node also saw
// abuse from the IP and the distributed score meets the threshold.
ActionBlockIfLocalCorroborated Action = "block_if_local_corroborated"
)
var validActions = [...]Action{
ActionOff,
ActionChallenge,
ActionBlockIfLocalCorroborated,
}
// ValidActions returns every config action ParseAction recognizes.
func ValidActions() []Action {
return append([]Action(nil), validActions[:]...)
}
// IsValidAction reports whether s names a config action ParseAction recognizes.
func IsValidAction(s string) bool {
action := Action(s)
for _, valid := range validActions {
if action == valid {
return true
}
}
return false
}
// ParseAction maps a config string to an Action, defaulting to challenge for
// any unknown value (the safe default; never silently block).
func ParseAction(s string) Action {
if IsValidAction(s) {
return Action(s)
}
return ActionChallenge
}
// Decision is what the node should do about an IP given central data.
type Decision int
const (
// DecisionIgnore takes no central-driven action.
DecisionIgnore Decision = iota
// DecisionChallenge elevates suspicion / serves a challenge.
DecisionChallenge
// DecisionBlock hard-blocks (only reachable with local corroboration).
DecisionBlock
)
// DecisionInput is the per-IP context for a central-data decision.
type DecisionInput struct {
Found bool // IP present in the central scored-set
Score int // distributed score 0-100
Protected bool // firebreak: infra/CF/crawler/RFC5737/allowlist
LocallyCorroborated bool // this node independently observed abuse from the IP
}
// Decide returns the node action for an IP. Firebreaks always win: a protected
// IP is never acted on from central data. A central-only signal can at most
// challenge; a hard block requires local corroboration, the action policy, and
// the score meeting blockThreshold.
func Decide(in DecisionInput, action Action, blockThreshold int) Decision {
if in.Protected || !in.Found || action == ActionOff {
return DecisionIgnore
}
if action == ActionBlockIfLocalCorroborated && in.LocallyCorroborated && in.Score >= blockThreshold {
return DecisionBlock
}
return DecisionChallenge
}
// centralState bundles the current snapshot and its derived lookup set so both
// are swapped in a single atomic store; two separate pointers could be read in
// a torn intermediate state.
type centralState struct {
snapshot ScoredSnapshot
set *Set
}
// CentralStore holds the current verified scored-set for concurrent lookups and
// is refreshed by a Puller. The state is swapped atomically so readers on the
// block/challenge path never block on a refresh and never see a torn snapshot.
type CentralStore struct {
puller *Puller
state atomic.Pointer[centralState]
}
// NewCentralStore builds an empty store backed by puller.
func NewCentralStore(puller *Puller) *CentralStore {
cs := &CentralStore{puller: puller}
empty := ScoredSnapshot{}
cs.state.Store(¢ralState{snapshot: empty, set: NewSet(empty)})
return cs
}
// Lookup returns the scored entry for ip from the current set.
func (cs *CentralStore) Lookup(ip string) (ScoredEntry, bool) {
return cs.state.Load().set.Lookup(ip)
}
// Version returns the current set version.
func (cs *CentralStore) Version() uint64 { return cs.state.Load().set.Version() }
// Refresh pulls an update and swaps in the new state on change. On a version
// gap it retries once from a cold pull (since=0) so a node that fell behind the
// diff window recovers with a full snapshot. A lower version is rejected so a
// rolled-back or hostile endpoint cannot regress the cache.
func (cs *CentralStore) Refresh(ctx context.Context) error {
cur := cs.state.Load().snapshot
next, changed, err := cs.puller.Refresh(ctx, cur)
if err == ErrSetVersionGap {
next, changed, err = cs.puller.Refresh(ctx, ScoredSnapshot{})
}
if err != nil {
return err
}
if changed {
nextState := ¢ralState{snapshot: next, set: NewSet(next)}
for {
latest := cs.state.Load()
if next.Version < latest.snapshot.Version {
return ErrSetVersionGap
}
if next.Version == latest.snapshot.Version {
return nil
}
if cs.state.CompareAndSwap(latest, nextState) {
return nil
}
}
}
return nil
}
package reporting
import (
"context"
"encoding/hex"
"errors"
"io"
"net/http"
"net/url"
"strconv"
"strings"
"time"
)
// maxScoredSetBytes caps a pulled scored-set payload so a hostile or broken
// endpoint cannot exhaust memory.
const maxScoredSetBytes = 64 << 20 // 64 MiB
var (
// ErrPullStatus means the endpoint returned an unexpected HTTP status.
ErrPullStatus = errors.New("reporting: scored-set pull bad status")
// ErrPullBodyTooLarge means the endpoint returned more bytes than a node
// will verify and cache.
ErrPullBodyTooLarge = errors.New("reporting: scored-set pull body too large")
)
// Puller fetches and verifies the signed scored-set from the central service.
// It pulls a full snapshot on a cold cache and a one-step diff thereafter,
// verifying the Ed25519 signature before applying anything.
type Puller struct {
client *http.Client
url string
pubHex string
}
// NewPuller builds a Puller for setURL, verifying against the central public
// key (hex). A nil client uses a default with a 30s timeout.
func NewPuller(client *http.Client, setURL, pubHex string) *Puller {
if client == nil {
client = &http.Client{Timeout: 30 * time.Second}
}
return &Puller{client: client, url: setURL, pubHex: pubHex}
}
// Refresh fetches an update relative to current and returns the new snapshot.
// When the endpoint reports no change (304), it returns current with
// changed=false. A diff that does not apply onto current (version gap) falls
// back by returning an error so the caller retries with a full pull (since=0).
func (p *Puller) Refresh(ctx context.Context, current ScoredSnapshot) (ScoredSnapshot, bool, error) {
reqURL := p.url
var err error
if current.Version > 0 {
reqURL, err = withSince(p.url, current.Version)
if err != nil {
return current, false, err
}
}
if validateErr := ValidateTargetURL(reqURL); validateErr != nil {
return current, false, validateErr
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, reqURL, nil)
if err != nil {
return current, false, err
}
// The scored-set URL is operator-configured (reputation.central.set_url) and
// the response is Ed25519-verified before use; not attacker-controlled.
// #nosec G704 -- central set URL is operator config, signature-verified.
resp, err := p.client.Do(req)
if err != nil {
return current, false, err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode == http.StatusNotModified {
return current, false, nil
}
if resp.StatusCode != http.StatusOK {
return current, false, ErrPullStatus
}
body, err := readScoredSetBody(resp.Body, maxScoredSetBytes)
if err != nil {
return current, false, err
}
sig, err := parseSetSignature(resp.Header.Get("X-CSM-Signature"))
if err != nil {
return current, false, err
}
switch resp.Header.Get("X-CSM-Kind") {
case "diff":
vd, err := OpenDiff(body, sig, p.pubHex)
if err != nil {
return current, false, err
}
next, err := ApplyDiff(current, vd)
if err != nil {
return current, false, err // version gap: caller retries full
}
return next, true, nil
default: // "snapshot" or unset
snap, err := OpenSnapshot(body, sig, p.pubHex)
if err != nil {
return current, false, err
}
return snap, true, nil
}
}
func readScoredSetBody(r io.Reader, limit int64) ([]byte, error) {
body, err := io.ReadAll(io.LimitReader(r, limit+1))
if err != nil {
return nil, err
}
if int64(len(body)) > limit {
return nil, ErrPullBodyTooLarge
}
return body, nil
}
func withSince(base string, version uint64) (string, error) {
u, err := url.Parse(base)
if err != nil {
return "", err
}
q, err := url.ParseQuery(u.RawQuery)
if err != nil {
return "", err
}
q.Set("since", strconv.FormatUint(version, 10))
u.RawQuery = q.Encode()
return u.String(), nil
}
func parseSetSignature(h string) ([]byte, error) {
scheme, hexSig, ok := strings.Cut(h, "=")
if !ok || scheme != "ed25519" {
return nil, ErrSetSignature
}
sig, err := hex.DecodeString(hexSig)
if err != nil {
return nil, ErrSetSignature
}
return sig, nil
}
package reporting
import (
"context"
"encoding/json"
"errors"
"fmt"
"log"
"time"
)
// Spooler is the production Reporter: it enqueues minimized reports to a durable
// spool and drains them to all configured targets on an interval, retrying from
// the spool when a collector is down. The daemon invokes persistence from its
// report worker so database writes do not block the alert path.
type Spooler struct {
spool *Spool
sender *Sender
targets map[string]Target
order []string
interval time.Duration
logf func(string, ...any)
}
// NewSpooler builds a Spooler over a spool, a sender, and the configured
// targets. A zero interval defaults to one minute.
func NewSpooler(spool *Spool, sender *Sender, targets []Target, interval time.Duration) *Spooler {
if interval <= 0 {
interval = time.Minute
}
m := make(map[string]Target, len(targets))
order := make([]string, 0, len(targets))
seen := make(map[string]bool, len(targets))
for _, t := range targets {
m[t.Name] = t
if !seen[t.Name] {
order = append(order, t.Name)
seen[t.Name] = true
}
}
return &Spooler{spool: spool, sender: sender, targets: m, order: order, interval: interval, logf: log.Printf}
}
// Enqueue persists r for delivery to every configured target. Dropped-count
// from spool overflow is logged so a sustained outage is visible.
func (s *Spooler) Enqueue(r Report) error {
body, err := json.Marshal(r)
if err != nil {
return err
}
var failures []error
for _, name := range s.order {
dropped, err := s.spool.Enqueue(name, body)
if err != nil {
failures = append(failures, fmt.Errorf("spool enqueue for %s: %w", name, err))
continue
}
if dropped > 0 {
s.logf("reporting: spool over capacity, dropped %d oldest reports for %s", dropped, name)
}
}
return errors.Join(failures...)
}
// DrainOnce attempts one delivery pass over the spool.
func (s *Spooler) DrainOnce(ctx context.Context) {
_, err := s.spool.Drain(func(target string, body []byte) error {
t, ok := s.targets[target]
if !ok {
// Target removed from config: drop the item by reporting success.
return nil
}
return s.sender.Send(ctx, t, body)
})
if err != nil {
s.logf("reporting: drain paused (will retry): %v", err)
}
}
// Run drains the spool every interval until ctx is cancelled.
func (s *Spooler) Run(ctx context.Context) {
t := time.NewTicker(s.interval)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
s.DrainOnce(ctx)
}
}
}
package reporting
import (
"net"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// Class is the public abuse classification sent on the wire. It is a closed set
// matching the central database's accepted classes.
type Class string
const (
ClassBruteforce Class = "bruteforce"
ClassPHPRelay Class = "php_relay"
// #nosec G101 -- abuse-class label, not a credential.
ClassCredentialStuffing Class = "credential_stuffing"
ClassBadASNEgress Class = "bad_asn_egress"
)
// checkClass maps a CSM finding check name to its public abuse class. Only
// confirmed-abuse checks that carry a source IP appear here; anything absent is
// never reported. This is the v1 reportable set (host_takeover is an incident
// Kind, not a finding check, and is added when the gate also taps incidents).
var checkClass = map[string]Class{
"pam_bruteforce": ClassBruteforce,
"wp_login_bruteforce": ClassBruteforce,
"xmlrpc_abuse": ClassBruteforce,
"ftp_bruteforce": ClassBruteforce,
"smtp_bruteforce": ClassBruteforce,
"mail_bruteforce": ClassBruteforce,
"admin_panel_bruteforce": ClassBruteforce,
"credential_stuffing": ClassCredentialStuffing,
"email_php_relay_abuse": ClassPHPRelay,
"bad_asn_outbound": ClassBadASNEgress,
}
// classMinSeverity is the least severe finding a class reports.
func classMinSeverity(class Class) alert.Severity {
if class == ClassBadASNEgress {
return alert.High
}
return alert.Critical
}
// Classify returns the abuse class for a check name, if it is reportable.
func Classify(check string) (Class, bool) {
c, ok := checkClass[check]
return c, ok
}
// Report is the minimized payload sent for a confirmed-abuse IP. It carries no
// hostnames, accounts, mailboxes, or paths. The JSON shape matches the central
// ingest contract exactly.
type Report struct {
IP string `json:"ip"`
Class Class `json:"class"`
Count int `json:"count"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
}
// Gate decides whether a finding is reportable and, if so, the minimized report
// to send. Only findings at or above their class minimum severity whose check
// is an enabled abuse class and that carry a usable source IP are reported.
type Gate struct {
// Enabled is the set of classes the operator has turned on. Empty means
// none are reported.
Enabled map[Class]bool
// Protected is the firebreak: addresses that must never be reported as
// attackers (private and documentation space, infrastructure, Cloudflare
// edges, verified crawlers). The daemon supplies the same predicate it
// uses before acting on central intelligence. Nil protects nothing.
Protected func(net.IP) bool
}
// Consider returns the minimized report for f, or ok=false when f must not be
// reported. The minimizer is deny-by-default: it copies only the IP, class,
// count, and timestamps, never tenant/domain/mailbox/path/process fields.
func (g Gate) Consider(f alert.Finding) (Report, bool) {
class, ok := Classify(f.Check)
if !ok || !g.Enabled[class] {
return Report{}, false
}
// Severity values grow more severe as they increase (Warning < High < Critical).
// Brute-force and relay classes need the Critical verdict; bad-ASN
// egress is only ever emitted as High, so demanding Critical made that
// class a dead option.
if f.Severity < classMinSeverity(class) {
return Report{}, false
}
ip := net.ParseIP(f.SourceIP)
if ip == nil {
return Report{}, false
}
if g.Protected != nil && g.Protected(ip) {
return Report{}, false
}
ts := f.Timestamp
if ts.IsZero() {
return Report{}, false
}
return Report{
IP: ip.String(),
Class: class,
Count: 1,
FirstSeen: ts.UTC(),
LastSeen: ts.UTC(),
}, true
}
// Reporter accepts minimized reports for later delivery, returning an error
// when persistence is incomplete. Call it outside the scan/alert hot path.
type Reporter interface {
Enqueue(Report) error
}
// Noop is the default Reporter; it discards reports. Used when reporting is
// disabled so call sites stay unconditional.
type Noop struct{}
// Enqueue discards r.
func (Noop) Enqueue(Report) error { return nil }
package reporting
import (
"bytes"
"crypto/ed25519"
"encoding/hex"
"encoding/json"
"errors"
"io"
"net"
"sort"
"time"
)
// This is the node consume side of the signed scored-set. The encoding here
// MUST stay byte-identical to the central publisher (csm-abuse-db
// internal/publish) so a node re-marshals a decoded payload to exactly the
// signed bytes. TestScoredSetGoldenCanonical pins the canonical form.
var (
// ErrSetSignature means the scored-set signature did not verify.
ErrSetSignature = errors.New("reporting: scored-set signature invalid")
// ErrSetInvalid means a decoded scored-set was malformed or noncanonical.
ErrSetInvalid = errors.New("reporting: scored-set invalid")
// ErrSetVersionGap means a diff does not apply onto the cached version.
ErrSetVersionGap = errors.New("reporting: scored-set version gap")
)
// ScoredEntry is one scored IP in the distributed set.
type ScoredEntry struct {
IP string `json:"ip"`
Score int `json:"score"`
Classes []Class `json:"classes"`
LastSeen time.Time `json:"last_seen"`
}
// ScoredSnapshot is the full scored-set at a version.
type ScoredSnapshot struct {
Version uint64 `json:"version"`
Entries []ScoredEntry `json:"entries"`
}
// ScoredDiff is an incremental update from FromVersion to ToVersion.
type ScoredDiff struct {
FromVersion uint64 `json:"from_version"`
ToVersion uint64 `json:"to_version"`
Added []ScoredEntry `json:"added"`
Removed []string `json:"removed"`
Changed []ScoredEntry `json:"changed"`
}
// VerifiedScoredDiff has passed Ed25519 verification and canonical decode.
// Its internals stay opaque so raw decoded bytes cannot be fed to ApplyDiff.
type VerifiedScoredDiff struct {
diff ScoredDiff
verified bool
}
func knownClass(c Class) bool {
switch c {
case ClassBruteforce, ClassPHPRelay, ClassCredentialStuffing, ClassBadASNEgress:
return true
}
return false
}
func canonicalScoredClasses(in []Class) ([]Class, bool) {
if len(in) == 0 {
return nil, false
}
out := append([]Class(nil), in...)
sort.Slice(out, func(i, j int) bool { return out[i] < out[j] })
for i, c := range out {
if !knownClass(c) {
return nil, false
}
if i > 0 && c == out[i-1] {
return nil, false
}
}
return out, true
}
func canonicalScoredEntry(e ScoredEntry) (ScoredEntry, bool) {
ip := net.ParseIP(e.IP)
if ip == nil || e.Score < 0 || e.Score > 100 || e.LastSeen.IsZero() {
return ScoredEntry{}, false
}
classes, ok := canonicalScoredClasses(e.Classes)
if !ok {
return ScoredEntry{}, false
}
return ScoredEntry{IP: ip.String(), Score: e.Score, Classes: classes, LastSeen: e.LastSeen.UTC()}, true
}
func canonicalScoredEntries(in []ScoredEntry) ([]ScoredEntry, bool) {
out := make([]ScoredEntry, len(in))
seen := make(map[string]struct{}, len(in))
for i, e := range in {
ce, ok := canonicalScoredEntry(e)
if !ok {
return nil, false
}
if _, dup := seen[ce.IP]; dup {
return nil, false
}
seen[ce.IP] = struct{}{}
out[i] = ce
}
sort.Slice(out, func(i, j int) bool { return out[i].IP < out[j].IP })
return out, true
}
func canonicalRemovedIPs(in []string) ([]string, bool) {
out := make([]string, len(in))
seen := make(map[string]struct{}, len(in))
for i, ip := range in {
parsed := net.ParseIP(ip)
if parsed == nil {
return nil, false
}
canonical := parsed.String()
if _, dup := seen[canonical]; dup {
return nil, false
}
seen[canonical] = struct{}{}
out[i] = canonical
}
sort.Strings(out)
return out, true
}
func canonicalScoredDiff(d ScoredDiff) (ScoredDiff, bool) {
if d.ToVersion <= d.FromVersion {
return ScoredDiff{}, false
}
added, ok := canonicalScoredEntries(d.Added)
if !ok {
return ScoredDiff{}, false
}
changed, ok := canonicalScoredEntries(d.Changed)
if !ok {
return ScoredDiff{}, false
}
removed, ok := canonicalRemovedIPs(d.Removed)
if !ok {
return ScoredDiff{}, false
}
if !scoredDiffOperationsDistinct(added, changed, removed) {
return ScoredDiff{}, false
}
return ScoredDiff{
FromVersion: d.FromVersion,
ToVersion: d.ToVersion,
Added: added,
Removed: removed,
Changed: changed,
}, true
}
func scoredDiffOperationsDistinct(added, changed []ScoredEntry, removed []string) bool {
seen := make(map[string]struct{}, len(added)+len(changed)+len(removed))
for _, e := range added {
seen[e.IP] = struct{}{}
}
for _, e := range changed {
if _, dup := seen[e.IP]; dup {
return false
}
seen[e.IP] = struct{}{}
}
for _, ip := range removed {
if _, dup := seen[ip]; dup {
return false
}
seen[ip] = struct{}{}
}
return true
}
func validateDiffAgainstBase(base ScoredSnapshot, d ScoredDiff) bool {
m := indexScoredEntries(base.Entries)
for _, ip := range d.Removed {
if _, ok := m[ip]; !ok {
return false
}
}
for _, e := range d.Added {
if _, ok := m[e.IP]; ok {
return false
}
}
for _, e := range d.Changed {
if _, ok := m[e.IP]; !ok {
return false
}
}
return true
}
func indexScoredEntries(entries []ScoredEntry) map[string]ScoredEntry {
m := make(map[string]ScoredEntry, len(entries))
for _, e := range entries {
m[e.IP] = e
}
return m
}
// MarshalScoredSnapshot deterministically encodes s (for signature/re-marshal).
func MarshalScoredSnapshot(s ScoredSnapshot) ([]byte, bool) {
entries, ok := canonicalScoredEntries(s.Entries)
if !ok {
return nil, false
}
b, err := json.Marshal(ScoredSnapshot{Version: s.Version, Entries: entries})
if err != nil {
return nil, false
}
return b, true
}
// MarshalScoredDiff deterministically encodes d (for signature/re-marshal).
func MarshalScoredDiff(d ScoredDiff) ([]byte, bool) {
d, ok := canonicalScoredDiff(d)
if !ok {
return nil, false
}
b, err := json.Marshal(d)
if err != nil {
return nil, false
}
return b, true
}
// VerifyScoredSet checks sig over payload under pubHex (the central public key).
func VerifyScoredSet(payload, sig []byte, pubHex string) error {
pub, err := hex.DecodeString(pubHex)
if err != nil || len(pub) != ed25519.PublicKeySize {
return ErrSetSignature
}
if !ed25519.Verify(ed25519.PublicKey(pub), payload, sig) {
return ErrSetSignature
}
return nil
}
func decodeStrict(b []byte, v any) error {
dec := json.NewDecoder(bytes.NewReader(b))
dec.DisallowUnknownFields()
if err := dec.Decode(v); err != nil {
return err
}
var extra struct{}
if err := dec.Decode(&extra); err != io.EOF {
if err == nil {
return ErrSetInvalid
}
return err
}
return nil
}
// OpenSnapshot verifies sig over payload, then decodes and canonical-checks the
// snapshot. Verification happens before any structural trust is placed in the
// bytes.
func OpenSnapshot(payload, sig []byte, pubHex string) (ScoredSnapshot, error) {
if err := VerifyScoredSet(payload, sig, pubHex); err != nil {
return ScoredSnapshot{}, err
}
var s ScoredSnapshot
if err := decodeStrict(payload, &s); err != nil {
return ScoredSnapshot{}, ErrSetInvalid
}
// A published snapshot is always version >= 1; version 0 is the node's
// empty-cache sentinel and must never be accepted from the wire (it would
// let a hostile endpoint pin the node to perpetual cold pulls).
if s.Version == 0 {
return ScoredSnapshot{}, ErrSetInvalid
}
entries, ok := canonicalScoredEntries(s.Entries)
if !ok {
return ScoredSnapshot{}, ErrSetInvalid
}
s.Entries = entries
canon, ok := MarshalScoredSnapshot(s)
if !ok || !bytes.Equal(canon, payload) {
return ScoredSnapshot{}, ErrSetInvalid
}
return s, nil
}
// OpenDiff verifies sig over payload then decodes the diff.
func OpenDiff(payload, sig []byte, pubHex string) (VerifiedScoredDiff, error) {
if err := VerifyScoredSet(payload, sig, pubHex); err != nil {
return VerifiedScoredDiff{}, err
}
var d ScoredDiff
if err := decodeStrict(payload, &d); err != nil {
return VerifiedScoredDiff{}, ErrSetInvalid
}
d, ok := canonicalScoredDiff(d)
if !ok {
return VerifiedScoredDiff{}, ErrSetInvalid
}
canon, ok := MarshalScoredDiff(d)
if !ok || !bytes.Equal(canon, payload) {
return VerifiedScoredDiff{}, ErrSetInvalid
}
return VerifiedScoredDiff{diff: d, verified: true}, nil
}
// ApplyDiff applies a verified diff onto base, returning the resulting snapshot.
func ApplyDiff(base ScoredSnapshot, d VerifiedScoredDiff) (ScoredSnapshot, error) {
if !d.verified {
return ScoredSnapshot{}, ErrSetSignature
}
entries, ok := canonicalScoredEntries(base.Entries)
if !ok {
return ScoredSnapshot{}, ErrSetInvalid
}
base.Entries = entries
diff := d.diff
if diff.FromVersion != base.Version {
return ScoredSnapshot{}, ErrSetVersionGap
}
if !validateDiffAgainstBase(base, diff) {
return ScoredSnapshot{}, ErrSetInvalid
}
m := indexScoredEntries(base.Entries)
for _, ip := range diff.Removed {
delete(m, ip)
}
for _, e := range diff.Added {
m[e.IP] = e
}
for _, e := range diff.Changed {
m[e.IP] = e
}
out := make([]ScoredEntry, 0, len(m))
for _, e := range m {
out = append(out, e)
}
entries, ok = canonicalScoredEntries(out)
if !ok {
return ScoredSnapshot{}, ErrSetInvalid
}
return ScoredSnapshot{Version: diff.ToVersion, Entries: entries}, nil
}
// Set is an in-memory lookup over the current scored-set.
type Set struct {
version uint64
byIP map[string]ScoredEntry
}
// NewSet builds a lookup set from a snapshot.
func NewSet(s ScoredSnapshot) *Set {
m := make(map[string]ScoredEntry, len(s.Entries))
for _, e := range s.Entries {
if ce, ok := canonicalScoredEntry(e); ok {
e = ce
}
e.Classes = append([]Class(nil), e.Classes...)
m[e.IP] = e
}
return &Set{version: s.Version, byIP: m}
}
// Version returns the set's version.
func (s *Set) Version() uint64 { return s.version }
// Lookup returns the scored entry for ip, normalizing the textual form.
func (s *Set) Lookup(ip string) (ScoredEntry, bool) {
p := net.ParseIP(ip)
if p == nil {
return ScoredEntry{}, false
}
e, ok := s.byIP[p.String()]
if ok {
e.Classes = append([]Class(nil), e.Classes...)
}
return e, ok
}
// Len returns the number of scored IPs.
func (s *Set) Len() int { return len(s.byIP) }
package reporting
import (
"bytes"
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"net"
"net/http"
"net/url"
"strings"
"time"
)
// Transport selects how a report is signed for a target.
type Transport string
const (
// TransportEd25519 signs with an Ed25519 node key (federation / central DB).
TransportEd25519 Transport = "ed25519"
// TransportHMAC signs with a shared HMAC secret (private collector).
TransportHMAC Transport = "hmac"
)
// Target is one reporting destination.
type Target struct {
Name string
URL string
Transport Transport
NodeID string
KeyID string
// Ed25519Key is the node's private key for TransportEd25519.
Ed25519Key ed25519.PrivateKey
// HMACSecret is the shared secret for TransportHMAC.
HMACSecret []byte
// BearerToken is an optional Authorization bearer for HMAC collectors.
BearerToken string
}
var (
// ErrInsecureURL means a non-HTTPS target was configured for a non-loopback
// host. Reports and their auth context must not cross the network in clear.
ErrInsecureURL = errors.New("reporting: target URL must be https")
// ErrRejected means the collector rejected the report (non-2xx, non-conflict).
ErrRejected = errors.New("reporting: report rejected")
)
// Sender delivers a signed report body to a target over HTTP.
type Sender struct {
client *http.Client
now func() time.Time
}
// NewSender builds a Sender. A nil client uses a default with a 15s timeout.
func NewSender(client *http.Client, now func() time.Time) *Sender {
if client == nil {
client = &http.Client{Timeout: 15 * time.Second}
}
if now == nil {
now = time.Now
}
return &Sender{client: client, now: now}
}
// ValidateTargetURL reports whether raw is allowed for report delivery without
// logging or returning the raw URL.
func ValidateTargetURL(raw string) error {
u, err := url.Parse(raw)
if err != nil || !secureURL(u) {
return ErrInsecureURL
}
return nil
}
// Send signs body for t and POSTs it. A 2xx is success; 409 Conflict (the
// collector already has this report) is treated as success. Other statuses and
// transport errors are failures the caller should retry from the spool.
func (s *Sender) Send(ctx context.Context, t Target, body []byte) error {
u, err := url.Parse(t.URL)
if err != nil {
return fmt.Errorf("reporting: bad target url: %w", err)
}
if !secureURL(u) {
return ErrInsecureURL
}
nonce, err := newNonce()
if err != nil {
return err
}
env := Envelope{
NodeID: t.NodeID,
KeyID: t.KeyID,
Method: http.MethodPost,
Path: requestPath(u),
BodyHash: HashBody(body),
Timestamp: s.now().UTC().Unix(),
Nonce: nonce,
}
sig, scheme, err := s.sign(t, env)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, t.URL, bytes.NewReader(body))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("X-CSM-Node", env.NodeID)
req.Header.Set("X-CSM-Key", env.KeyID)
req.Header.Set("X-CSM-Timestamp", fmt.Sprintf("%d", env.Timestamp))
req.Header.Set("X-CSM-Nonce", env.Nonce)
req.Header.Set("X-CSM-Signature", scheme+"="+hex.EncodeToString(sig))
if t.BearerToken != "" {
req.Header.Set("Authorization", "Bearer "+t.BearerToken)
}
resp, err := s.do(req)
if err != nil {
return err
}
defer func() { _ = resp.Body.Close() }()
switch {
case resp.StatusCode >= 200 && resp.StatusCode < 300:
return nil
case resp.StatusCode == http.StatusConflict:
return nil // collector already recorded this report (replay/dup)
default:
return fmt.Errorf("%w: status %d", ErrRejected, resp.StatusCode)
}
}
func (s *Sender) do(req *http.Request) (*http.Response, error) {
client := *s.client
// The signature is bound to the configured URL path, and auth headers must
// not be replayed to a redirected endpoint.
client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
}
// The destination is an operator-configured report target, validated to
// HTTPS (or loopback HTTP) by secureURL; it is not attacker-controlled.
// #nosec G704 -- report target URL is operator config, scheme-validated.
return client.Do(req)
}
func (s *Sender) sign(t Target, env Envelope) (sig []byte, scheme string, err error) {
switch t.Transport {
case TransportEd25519:
sig, err = SignEd25519(env, t.Ed25519Key)
return sig, "ed25519", err
case TransportHMAC:
sig, err = SignHMAC(env, t.HMACSecret)
return sig, "sha256", err
default:
return nil, "", fmt.Errorf("reporting: unknown transport %q", t.Transport)
}
}
// secureURL requires https, except for loopback hosts (local collectors / tests).
func secureURL(u *url.URL) bool {
host := u.Hostname()
if host == "" {
return false
}
if u.Scheme == "https" {
return true
}
if u.Scheme == "http" {
if ip := net.ParseIP(host); ip != nil {
return ip.IsLoopback()
}
return strings.EqualFold(host, "localhost")
}
return false
}
func requestPath(u *url.URL) string {
if u.Path == "" {
return "/"
}
return u.Path
}
func newNonce() (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", err
}
return hex.EncodeToString(b[:]), nil
}
package reporting
import (
"encoding/binary"
"encoding/json"
"sync"
"time"
bolt "go.etcd.io/bbolt"
)
// spoolItem is one queued report: the destination target name and the exact
// minimized body bytes to sign and send.
type spoolItem struct {
Target string `json:"t"`
Body []byte `json:"b"`
}
// Spool is a durable, bounded outbound queue for reports, backed by bbolt so a
// down collector or a daemon restart does not drop confirmed-abuse reports.
type Spool struct {
db *bolt.DB
bucket []byte
max int
drain sync.Mutex
mutation sync.Mutex
health spoolHealth
}
// NewSpool opens (or creates) a spool at path with a per-node entry cap.
func NewSpool(path, bucket string, max int) (*Spool, error) {
db, err := bolt.Open(path, 0o600, &bolt.Options{Timeout: 2 * time.Second})
if err != nil {
return nil, err
}
s := &Spool{db: db, bucket: []byte(bucket), max: max, health: newSpoolHealth(max)}
if err := db.Update(func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists(s.bucket)
if err != nil {
return err
}
now := time.Now()
c := b.Cursor()
for k, _ := c.First(); k != nil; k, _ = c.Next() {
s.health.pending[string(k)] = &spoolWork{ticket: s.health.stats.Begin(now)}
}
return nil
}); err != nil {
_ = db.Close()
return nil, err
}
return s, nil
}
// Close releases the underlying database.
func (s *Spool) Close() error {
s.mutation.Lock()
defer s.mutation.Unlock()
return s.db.Close()
}
// Enqueue appends a report body destined for target. When the queue exceeds its
// cap, the oldest entries are dropped (FIFO) and the dropped count is returned
// so the caller can surface it; reports are best-effort under sustained outage.
func (s *Spool) Enqueue(target string, body []byte) (dropped int, err error) {
item := spoolItem{Target: target, Body: body}
enc, err := json.Marshal(item)
if err != nil {
return 0, err
}
s.mutation.Lock()
defer s.mutation.Unlock()
ticket := s.health.stats.Begin(time.Now())
ticket.Start(time.Now())
var key [8]byte
var evicted []string
err = s.db.Update(func(tx *bolt.Tx) error {
b := tx.Bucket(s.bucket)
seq, seqErr := b.NextSequence()
if seqErr != nil {
return seqErr
}
binary.BigEndian.PutUint64(key[:], seq)
if e := b.Put(key[:], enc); e != nil {
return e
}
// Count current keys via the cursor; Bucket.Stats is not reliable for
// pending changes inside the same write transaction.
count := 0
c := b.Cursor()
for k, _ := c.First(); k != nil; k, _ = c.Next() {
count++
}
// Trim from the front (oldest keys sort first) until within cap.
for count > s.max {
tc := b.Cursor()
k, _ := tc.First()
if k == nil {
break
}
evicted = append(evicted, string(k))
if e := b.Delete(k); e != nil {
return e
}
count--
}
return nil
})
s.health.enqueueFailed.Store(err != nil)
if err != nil {
ticket.Reject(time.Now())
return 0, err
}
s.applyEnqueue(string(key[:]), ticket, evicted)
return len(evicted), nil
}
// Len returns the number of queued items.
func (s *Spool) Len() int {
n := 0
_ = s.db.View(func(tx *bolt.Tx) error {
n = tx.Bucket(s.bucket).Stats().KeyN
return nil
})
return n
}
// Drain processes queued items in FIFO order, calling send for each. An item is
// removed only when send returns nil; on the first send error Drain stops and
// leaves that item (and the rest) for a later retry, preserving order. It
// returns how many were delivered.
func (s *Spool) Drain(send func(target string, body []byte) error) (delivered int, err error) {
s.drain.Lock()
defer s.drain.Unlock()
for {
key, item, work, err := s.next()
if err != nil {
return delivered, err
}
if work == nil {
return delivered, nil
}
if err := s.deliver(key, item, work, send); err != nil {
return delivered, err
}
delivered++
}
}
func (s *Spool) next() (key []byte, item spoolItem, work *spoolWork, err error) {
s.mutation.Lock()
defer s.mutation.Unlock()
err = s.db.View(func(tx *bolt.Tx) error {
k, v := tx.Bucket(s.bucket).Cursor().First()
if k == nil {
return nil
}
key = append([]byte(nil), k...)
return json.Unmarshal(v, &item)
})
s.health.readFailed.Store(err != nil)
if err != nil {
return nil, item, nil, err
}
if key == nil {
s.health.sendFailed.Store(false)
s.health.removeFailed.Store(false)
return nil, item, nil, nil
}
work = s.health.pending[string(key)]
if work == nil {
// The spool outlives the process that wrote it. A record with no
// accounting still has to be delivered, so it is adopted here rather
// than taken as an invariant.
work = &spoolWork{ticket: s.health.stats.Begin(time.Now())}
s.health.pending[string(key)] = work
}
work.ticket.Start(time.Now())
s.health.active = work
return key, item, work, nil
}
func (s *Spool) deliver(key []byte, item spoolItem, work *spoolWork, send func(string, []byte) error) error {
sent, removed := false, false
defer func() { s.finishDelivery(work, sent, removed) }()
if err := send(item.Target, item.Body); err != nil {
return err
}
sent = true
s.health.sendFailed.Store(false)
s.mutation.Lock()
defer s.mutation.Unlock()
if !work.evicted {
err := s.db.Update(func(tx *bolt.Tx) error {
return tx.Bucket(s.bucket).Delete(key)
})
s.health.removeFailed.Store(err != nil)
if err != nil {
return err
}
delete(s.health.pending, string(key))
}
removed = true
return nil
}
package reporting
import (
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type spoolWork struct {
ticket queuehealth.Ticket
evicted bool
delivered bool
}
type spoolHealth struct {
stats *queuehealth.Tracker
pending map[string]*spoolWork // guarded by Spool.mutation
active *spoolWork
enqueueFailed atomic.Bool
readFailed atomic.Bool
removeFailed atomic.Bool
sendFailed atomic.Bool
}
func newSpoolHealth(capacity int) spoolHealth {
return spoolHealth{
// Delivery normally runs once a minute. Allow a complete retry interval
// before age alone declares the durable queue stalled.
stats: queuehealth.NewSharedCapacity(capacity, 2*time.Minute),
pending: make(map[string]*spoolWork),
}
}
// QueueStatuses uses memory only, including when a database write or sender
// stalls. Existing records are timed from open, not from an invented disk age.
func (s *Spool) QueueStatuses(now time.Time) map[string]queuehealth.Status {
status := s.health.stats.Snapshot(now)
status.LagBasis = "observed_age"
switch {
case s.health.enqueueFailed.Load() || s.health.readFailed.Load() || s.health.removeFailed.Load():
status.Status, status.Reason = "degraded", "spool_io"
case s.health.sendFailed.Load():
status.Status, status.Reason = "degraded", "delivery_failed"
}
return map[string]queuehealth.Status{"spool": status}
}
// applyEnqueue runs only after commit, under the mutation lock. An evicted
// record can still be owned by send; its outcome decides whether it was lost.
func (s *Spool) applyEnqueue(key string, ticket queuehealth.Ticket, evicted []string) {
now := time.Now()
ticket.Requeue(now)
s.health.pending[key] = &spoolWork{ticket: ticket}
for _, key := range evicted {
work := s.health.pending[key]
delete(s.health.pending, key)
switch work {
case nil:
// A record evicted from disk with no accounting cannot be
// attributed to a caller; count the report it carried as lost.
s.health.stats.Lose(now, 1)
case s.health.active:
work.evicted = true
default:
work.discard(now)
}
}
}
func (s *Spool) finishDelivery(work *spoolWork, sent, removed bool) {
s.mutation.Lock()
defer s.mutation.Unlock()
if !sent {
s.health.sendFailed.Store(true)
}
now := time.Now()
// Database removal can fail after receipt. Later eviction or a failed
// retry must not turn that earlier acknowledgement into a lost report.
work.delivered = work.delivered || sent
switch {
case removed:
work.ticket.Finish(now)
case work.evicted:
work.discard(now)
default:
work.ticket.Requeue(now)
}
s.health.active = nil
}
func (w *spoolWork) discard(now time.Time) {
if w.delivered {
w.ticket.Finish(now)
} else {
w.ticket.Reject(now)
}
}
// Package reporting is the node side of CSM abuse reporting (Layer A). It turns
// confirmed-abuse findings into minimized, signed reports for a central abuse
// database or a private collector.
//
// The signed-envelope wire format here MUST stay byte-identical to the central
// service's verifier (csm-abuse-db internal/envelope). It is duplicated rather
// than imported to keep this repo's build self-contained; the wire-format test
// pins the canonical bytes so any divergence fails the build.
package reporting
import (
"crypto/ed25519"
"crypto/hmac"
"crypto/sha256"
"encoding/binary"
"errors"
)
const (
bodyHashLen = sha256.Size
// maxFieldLen bounds each canonical field; matches the central verifier.
maxFieldLen = 1 << 16
)
// ErrInvalidEnvelope means the envelope is structurally unusable.
var ErrInvalidEnvelope = errors.New("reporting: invalid envelope")
// Envelope is the set of fields signed alongside a report body. Field order and
// encoding match the central verifier exactly.
type Envelope struct {
NodeID string
KeyID string
Method string
Path string
BodyHash []byte // SHA-256 of the report body, 32 bytes
Timestamp int64 // unix seconds, UTC
Nonce string
}
// HashBody returns the SHA-256 of body.
func HashBody(body []byte) []byte {
sum := sha256.Sum256(body)
return sum[:]
}
// canonical returns the deterministic, injective encoding signed over. Every
// variable-length field is length-prefixed with a 4-byte big-endian length, and
// the timestamp is appended as 8 big-endian bytes.
func (e Envelope) canonical() ([]byte, error) {
if len(e.BodyHash) != bodyHashLen {
return nil, ErrInvalidEnvelope
}
fields := [][]byte{
[]byte(e.NodeID),
[]byte(e.KeyID),
[]byte(e.Method),
[]byte(e.Path),
e.BodyHash,
[]byte(e.Nonce),
}
size := 8
for _, f := range fields {
if len(f) > maxFieldLen {
return nil, ErrInvalidEnvelope
}
size += 4 + len(f)
}
buf := make([]byte, 0, size)
var lenbuf [4]byte
for _, f := range fields {
// len(f) is bounded by maxFieldLen above, well under math.MaxUint32.
// #nosec G115 -- length validated <= maxFieldLen (64 KiB) above.
binary.BigEndian.PutUint32(lenbuf[:], uint32(len(f)))
buf = append(buf, lenbuf[:]...)
buf = append(buf, f...)
}
var ts [8]byte
// Lossless int64 -> uint64 bit reinterpretation for fixed-width encoding.
// #nosec G115 -- intentional bit-pattern reinterpret, not a narrowing cast.
binary.BigEndian.PutUint64(ts[:], uint64(e.Timestamp))
buf = append(buf, ts[:]...)
return buf, nil
}
// SignEd25519 signs the canonical envelope with an Ed25519 private key.
func SignEd25519(e Envelope, priv ed25519.PrivateKey) ([]byte, error) {
msg, err := e.canonical()
if err != nil {
return nil, err
}
if len(priv) != ed25519.PrivateKeySize {
return nil, ErrInvalidEnvelope
}
return ed25519.Sign(priv, msg), nil
}
// SignHMAC signs the canonical envelope with HMAC-SHA256 (private collector).
func SignHMAC(e Envelope, secret []byte) ([]byte, error) {
msg, err := e.canonical()
if err != nil {
return nil, err
}
if len(secret) == 0 {
return nil, ErrInvalidEnvelope
}
mac := hmac.New(sha256.New, secret)
_, _ = mac.Write(msg)
return mac.Sum(nil), nil
}
// Package safepath pins directories for operations on tenant-controlled names.
package safepath
import (
"crypto/rand"
"fmt"
"os"
"path/filepath"
"runtime"
"strings"
"golang.org/x/sys/unix"
)
// Dir owns an open directory. Operations accept single basenames only and
// never resolve a symlink, including in the final component.
type Dir struct {
file *os.File
}
// OpenDir opens an operator-controlled root. Its ancestors must be trusted;
// use OpenTarget to traverse anything beneath it controlled by an account.
func OpenDir(path string) (*Dir, error) {
// #nosec G304 -- this opens the operator-controlled anchor; all tenant components are traversed with descriptor-relative no-follow operations
f, err := os.OpenFile(path, os.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC, 0)
if err != nil {
return nil, err
}
return &Dir{file: f}, nil
}
// OpenDirNoFollow pins every component from the filesystem root. Configured
// content roots can belong to tenants, including their ancestor directories.
func OpenDirNoFollow(path string) (*Dir, error) {
if !filepath.IsAbs(path) || filepath.Clean(path) != path {
return nil, fmt.Errorf("directory root must be an absolute clean path: %q", path)
}
root, err := OpenDir("/")
if err != nil {
return nil, err
}
if path == "/" {
return root, nil
}
defer func() { _ = root.Close() }()
return root.walk(strings.TrimPrefix(path, "/"), false)
}
func (d *Dir) Close() error { return d.file.Close() }
func (d *Dir) Sync() error { return d.file.Sync() }
func validName(name string) bool {
return name != "" && name != "." && name != ".." && !strings.ContainsAny(name, "/\x00")
}
// fileFD returns f's integer descriptor for a descriptor-relative syscall.
// Callers must keep f alive across the call.
func fileFD(f *os.File) int {
// #nosec G115 -- an open descriptor is a small non-negative value; os.File only exposes it as uintptr
return int(f.Fd())
}
// adoptFD wraps a descriptor produced by a syscall whose error was already checked.
func adoptFD(fd int, name string) *os.File {
// #nosec G115 -- a syscall that reported success returns a non-negative descriptor
return os.NewFile(uintptr(fd), name)
}
func (d *Dir) OpenFile(name string, flags int, mode os.FileMode) (*os.File, error) {
if !validName(name) {
return nil, fmt.Errorf("invalid basename %q", name)
}
fd, err := unix.Openat(fileFD(d.file), name, flags|unix.O_NOFOLLOW|unix.O_CLOEXEC|unix.O_NONBLOCK, uint32(mode.Perm()))
// The integer descriptor does not keep os.File's finalizer alive.
runtime.KeepAlive(d)
if err != nil {
return nil, &os.PathError{Op: "openat", Path: name, Err: err}
}
return adoptFD(fd, name), nil
}
func (d *Dir) Stat(name string) (os.FileInfo, error) {
f, err := d.OpenFile(name, os.O_RDONLY, 0)
if err != nil {
return nil, err
}
defer f.Close()
return f.Stat()
}
func (d *Dir) CreateTemp() (*os.File, error) {
return d.OpenFile(".csm-restore-"+rand.Text(), os.O_RDWR|os.O_CREATE|os.O_EXCL, 0600)
}
// CreatePrivateTemp isolates transaction names from writers of the parent.
// The parent may rename this directory, so callers must keep using its handle.
func (d *Dir) CreatePrivateTemp() (*Dir, string, error) {
name := ".csm-restore-" + rand.Text()
err := unix.Mkdirat(fileFD(d.file), name, 0700)
runtime.KeepAlive(d)
if err != nil {
return nil, "", err
}
f, err := d.OpenFile(name, os.O_RDONLY|unix.O_DIRECTORY, 0)
if err != nil {
return nil, "", err
}
var stat unix.Stat_t
err = unix.Fstat(fileFD(f), &stat)
runtime.KeepAlive(f)
if err != nil {
_ = f.Close()
return nil, "", err
}
if int(stat.Uid) != os.Geteuid() || stat.Mode&0077 != 0 {
_ = f.Close()
return nil, "", fmt.Errorf("private restore directory changed while opening")
}
return &Dir{file: f}, name, nil
}
// RemoveDir removes only an empty directory; it never traverses its contents.
func (d *Dir) RemoveDir(name string) error {
return d.unlink(name, unix.AT_REMOVEDIR)
}
func (d *Dir) Remove(name string) error {
return d.unlink(name, 0)
}
func (d *Dir) unlink(name string, flags int) error {
if !validName(name) {
return fmt.Errorf("invalid basename %q", name)
}
err := unix.Unlinkat(fileFD(d.file), name, flags)
runtime.KeepAlive(d)
if err != nil {
return &os.PathError{Op: "unlinkat", Path: name, Err: err}
}
return nil
}
// RenameTo never replaces an existing destination. ExchangeTo atomically
// swaps two existing names. Neither operation follows either name's symlink.
func (d *Dir) RenameTo(name string, dest *Dir, destName string) error {
return d.rename(name, dest, destName, false)
}
func (d *Dir) ExchangeTo(name string, dest *Dir, destName string) error {
return d.rename(name, dest, destName, true)
}
func (d *Dir) rename(name string, dest *Dir, destName string, exchange bool) error {
if !validName(name) || !validName(destName) {
return fmt.Errorf("invalid rename basenames %q, %q", name, destName)
}
err := renameat(fileFD(d.file), name, fileFD(dest.file), destName, exchange)
runtime.KeepAlive(d)
runtime.KeepAlive(dest)
if err != nil {
return &os.LinkError{Op: "renameat", Old: name, New: destName, Err: err}
}
return nil
}
// Target keeps both the trusted root and the destination's parent open.
// Check detects a renamed parent; I/O uses Parent even if that check races.
type Target struct {
Parent *Dir
Name string
root *Dir
relDir string
}
func OpenTarget(rootPath, relative string, createParents bool) (*Target, error) {
if !filepath.IsLocal(relative) || filepath.Clean(relative) != relative || relative == "." {
return nil, fmt.Errorf("invalid relative restore path %q", relative)
}
root, err := openTargetRoot(rootPath)
if err != nil {
return nil, err
}
relDir := filepath.Dir(relative)
parent, err := root.walk(relDir, createParents)
if err != nil {
_ = root.Close()
return nil, err
}
return &Target{Parent: parent, Name: filepath.Base(relative), root: root, relDir: relDir}, nil
}
func (t *Target) Close() {
_ = t.Parent.Close()
_ = t.root.Close()
}
func (t *Target) Check() error {
current, err := t.root.walk(t.relDir, false)
if err != nil {
return fmt.Errorf("restore parent changed: %w", err)
}
defer func() { _ = current.Close() }()
want, err := t.Parent.file.Stat()
if err != nil {
return err
}
got, err := current.file.Stat()
if err != nil {
return err
}
if !os.SameFile(want, got) {
return fmt.Errorf("restore parent changed")
}
return nil
}
func (d *Dir) walk(relative string, create bool) (*Dir, error) {
fd, err := unix.Openat(fileFD(d.file), ".", unix.O_RDONLY|unix.O_DIRECTORY|unix.O_CLOEXEC, 0)
runtime.KeepAlive(d)
if err != nil {
return nil, err
}
current := &Dir{file: adoptFD(fd, ".")}
if relative == "." {
return current, nil
}
for _, name := range strings.Split(relative, string(filepath.Separator)) {
if !validName(name) {
_ = current.Close()
return nil, fmt.Errorf("invalid directory component %q", name)
}
next, openErr := current.OpenFile(name, os.O_RDONLY|unix.O_DIRECTORY, 0)
if os.IsNotExist(openErr) && create {
// Public document roots need traversable parents. A concurrent
// creator is harmless only if the no-follow open accepts its inode.
mkdirErr := unix.Mkdirat(fileFD(current.file), name, 0755)
runtime.KeepAlive(current)
if mkdirErr != nil && mkdirErr != unix.EEXIST {
_ = current.Close()
return nil, mkdirErr
}
next, openErr = current.OpenFile(name, os.O_RDONLY|unix.O_DIRECTORY, 0)
}
// A concurrent restore may have created the directory but not yet
// persisted its entry. Each successful restore needs its own sync.
if openErr == nil && create {
if syncErr := current.Sync(); syncErr != nil {
_ = next.Close()
openErr = syncErr
}
}
_ = current.Close()
if openErr != nil {
return nil, openErr
}
current = &Dir{file: next}
}
return current, nil
}
//go:build linux
package safepath
import (
"os"
"runtime"
"time"
"unsafe"
"golang.org/x/sys/unix"
)
// SetModTime changes only the pinned inode, preserving nanosecond precision.
func SetModTime(file *os.File, mtime time.Time) error {
ts, err := unix.TimeToTimespec(mtime)
if err != nil {
return err
}
times := [2]unix.Timespec{{Nsec: unix.UTIME_OMIT}, ts}
// This is the kernel's futimens ABI. A NULL pathname operates on the fd
// without the newer AT_EMPTY_PATH flag or resolving a tenant-owned name.
// #nosec G103 -- fixed timespec array passed directly to a synchronous fd-only syscall; no pointer arithmetic or user-controlled address.
_, _, errno := unix.Syscall6(unix.SYS_UTIMENSAT, file.Fd(), 0, uintptr(unsafe.Pointer(×[0])), 0, 0, 0)
runtime.KeepAlive(file)
if errno != 0 {
return &os.PathError{Op: "futimens", Path: file.Name(), Err: errno}
}
return nil
}
package safepath
import "golang.org/x/sys/unix"
func renameat(oldfd int, old string, newfd int, name string, exchange bool) error {
flags := uint(unix.RENAME_NOREPLACE)
if exchange {
flags = unix.RENAME_EXCHANGE
}
return unix.Renameat2(oldfd, old, newfd, name, flags)
}
//go:build linux
package safepath
func openTargetRoot(path string) (*Dir, error) { return OpenDirNoFollow(path) }
// Package sdnotify talks to the systemd notification socket. The daemon calls
// Ready when watchers are attached, Status to publish a one-line state visible
// in `systemctl status`, and Watchdog on a recurring ticker so systemd's
// WatchdogSec= keep-alive doesn't expire.
//
// Capture takes the systemd variables out of the process environment at
// startup and keeps them here. Everything the daemon executes afterwards --
// the PHP taint worker, every command a check shells out to -- would otherwise
// inherit NOTIFY_SOCKET and be able to write to it. The unit declares
// NotifyAccess=main, so each of those children produced a "notification
// message from PID ..., but reception only permitted for main PID" line in the
// journal. Scrubbing once at the source is what keeps a spawn site added later
// from reintroducing it.
//
// Every function is a no-op when no socket was captured (the daemon is not
// running under systemd; e.g. dev mode). That contract makes it safe to call
// these helpers unconditionally without runtime gates in the daemon code.
package sdnotify
import (
"net"
"os"
"strconv"
"sync"
"time"
)
// systemdEnv are the variables systemd passes to a service that must not reach
// any child process. Only the first two are used here; the rest describe
// socket activation this daemon does not use, and a child that read them would
// act on descriptors it was never given.
var systemdEnv = []string{
"NOTIFY_SOCKET",
"WATCHDOG_USEC",
"WATCHDOG_PID",
"LISTEN_FDS",
"LISTEN_PID",
"LISTEN_FDNAMES",
}
var (
mu sync.RWMutex
socket string
watchdogTimeout time.Duration
)
// Capture reads the systemd notification environment and removes it from the
// process environment. Called once, before the daemon spawns anything.
// Re-reading on a later call is intentional and lets tests drive it.
func Capture() {
addr := os.Getenv("NOTIFY_SOCKET")
timeout := parseWatchdogTimeout(os.Getenv("WATCHDOG_USEC"))
for _, name := range systemdEnv {
_ = os.Unsetenv(name)
}
mu.Lock()
socket, watchdogTimeout = addr, timeout
mu.Unlock()
}
func parseWatchdogTimeout(usec string) time.Duration {
if usec == "" {
return 0
}
parsed, err := strconv.ParseInt(usec, 10, 64)
if err != nil || parsed <= 0 {
return 0
}
return time.Duration(parsed) * time.Microsecond
}
// WatchdogTimeout reports the interval systemd expects keepalives within, and
// whether the unit configured one at all.
func WatchdogTimeout() (time.Duration, bool) {
mu.RLock()
defer mu.RUnlock()
return watchdogTimeout, watchdogTimeout > 0
}
// Enabled reports whether a notification socket was captured, i.e. whether
// this process is running under systemd with notifications configured.
func Enabled() bool {
mu.RLock()
defer mu.RUnlock()
return socket != ""
}
// notify sends one datagram to the captured socket. Returns (true, nil) when
// the notification was delivered, (false, nil) when no socket was captured, or
// (false, err) on a real I/O error.
func notify(state string) (bool, error) {
mu.RLock()
addr := socket
mu.RUnlock()
if addr == "" {
return false, nil
}
// A leading "@" marks an abstract socket name; the syscall layer turns it
// into the leading NUL byte the kernel expects.
conn, err := net.DialUnix("unixgram", nil, &net.UnixAddr{Name: addr, Net: "unixgram"})
if err != nil {
return false, err
}
defer func() { _ = conn.Close() }()
// A bound socket can stop draining its queue. Do not let a notification
// hold startup, status updates or watchdog shutdown indefinitely.
if err := conn.SetWriteDeadline(time.Now().Add(time.Second)); err != nil {
return false, err
}
if _, err := conn.Write([]byte(state)); err != nil {
return false, err
}
return true, nil
}
// Ready signals systemd that the daemon has finished startup.
func Ready() (bool, error) {
return notify("READY=1")
}
// Reloading signals systemd that the daemon is reloading its config.
func Reloading() (bool, error) {
return notify("RELOADING=1")
}
// Status sets a single-line status string visible in `systemctl status csm`.
func Status(msg string) (bool, error) {
return notify("STATUS=" + msg)
}
// Watchdog pings the systemd watchdog. Required when the unit declares
// WatchdogSec=; without periodic pings systemd will restart the daemon
// after WatchdogSec elapses.
func Watchdog() (bool, error) {
return notify("WATCHDOG=1")
}
// Package selftest runs CSM's detection over a small bundle of samples whose
// verdicts are known, so an operator can see what the shipped rules catch
// without pointing the scanner at a production account.
//
// The bundle carries adversarial samples and benign controls. The controls are
// the half that matters when judging a scanner: a rule set that flags an
// ordinary WordPress plugin is worse than one that misses a shell.
//
// Every sample is stored base64-encoded and only decoded in memory. Endpoint
// antivirus on a developer machine, or on the server an operator runs this on,
// deletes files that look like web shells; a decoded copy on disk would make
// the bundle disappear.
package selftest
import (
"encoding/base64"
"fmt"
)
// Sample is one file with a known verdict.
type Sample struct {
Name string
// Ext is the extension the scanner is told about, including the dot.
Ext string
// Malicious is what the sample is, not what the rules do with it.
Malicious bool
Description string
// RealtimeGap and YaraGap record that the shipped rule set does not fire
// on this sample today. They are measurements, not permissions: the gates
// fail when a gap closes as well as when one opens, so closing one is a
// deliberate edit here rather than a silent change in behaviour.
//
// A gap is a gap in the signature engines only. Taint analysis, the
// behavioural checks and PHP Shield are separate layers and are not
// measured by this bundle.
RealtimeGap bool
YaraGap bool
// Encoded is the sample content, base64-encoded at rest.
Encoded string
}
// Content decodes the sample.
func (s Sample) Content() ([]byte, error) {
data, err := base64.StdEncoding.DecodeString(s.Encoded)
if err != nil {
return nil, fmt.Errorf("sample %s: %w", s.Name, err)
}
if len(data) == 0 {
return nil, fmt.Errorf("sample %s: empty content", s.Name)
}
return data, nil
}
// Result is one sample's outcome.
type Result struct {
Name string `json:"name"`
Description string `json:"description"`
Malicious bool `json:"malicious"`
Detected bool `json:"detected"`
KnownGap bool `json:"known_gap"`
Rules []string `json:"rules,omitempty"`
Pass bool `json:"pass"`
Error string `json:"error,omitempty"`
}
// ScanFunc is the detection under test: it returns the names of the rules that
// fired on the content, or an error if the scan could not complete.
type ScanFunc func(content []byte, ext string) ([]string, error)
// Engine names the rule set being measured, so a sample's recorded gap is
// compared against the engine that has it.
type Engine string
const (
// Realtime is the YAML rule set the real-time watchers use.
Realtime Engine = "realtime"
// Yara is the YARA-X rule set used by scheduled and email scanning. It is
// present only in builds compiled with the yara tag.
Yara Engine = "yara"
)
func (s Sample) gap(engine Engine) bool {
if engine == Yara {
return s.YaraGap
}
return s.RealtimeGap
}
// Run scans every sample with one engine and reports whether the outcome
// matched what is recorded for it. A sample that cannot be decoded is a
// failure, not a skip: a bundle that silently shrinks proves nothing.
func Run(engine Engine, scan ScanFunc) []Result {
results := make([]Result, 0, len(samples))
for _, sample := range samples {
result := Result{
Name: sample.Name,
Description: sample.Description,
Malicious: sample.Malicious,
KnownGap: sample.Malicious && sample.gap(engine),
}
content, err := sample.Content()
if err != nil {
result.Error = err.Error()
results = append(results, result)
continue
}
result.Rules, err = scan(content, sample.Ext)
if err != nil {
result.Error = err.Error()
results = append(results, result)
continue
}
result.Detected = len(result.Rules) > 0
result.Pass = result.Detected == (sample.Malicious && !result.KnownGap)
results = append(results, result)
}
return results
}
// Summary counts outcomes. The three failure kinds mean different things and
// are never added together: a newly missed sample is a regression, a flagged
// control is a false positive, and a recorded gap that now fires is good news
// that needs the bundle updated.
type Summary struct {
Detected int `json:"detected"`
KnownGaps int `json:"known_gaps"`
Clean int `json:"clean"`
Missed int `json:"missed"`
FalsePositives int `json:"false_positives"`
ClosedGaps int `json:"closed_gaps"`
Errors int `json:"errors"`
}
// Failed reports whether the run found anything that needs attention.
func (s Summary) Failed() bool {
return s == (Summary{}) || s.Errors > 0 || s.Missed > 0 || s.FalsePositives > 0 || s.ClosedGaps > 0
}
// Summarize counts the results.
func Summarize(results []Result) Summary {
var out Summary
for _, r := range results {
switch {
case r.Error != "":
out.Errors++
case r.KnownGap && r.Detected:
out.ClosedGaps++
case r.KnownGap:
out.KnownGaps++
case r.Malicious && r.Detected:
out.Detected++
case r.Malicious:
out.Missed++
case r.Detected:
out.FalsePositives++
default:
out.Clean++
}
}
return out
}
// Samples returns the bundle.
func Samples() []Sample { return append([]Sample(nil), samples...) }
// Package session owns browser-session identities and their storage contract.
package session
import (
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"errors"
"time"
)
var (
ErrInvalid = errors.New("invalid or expired browser session")
ErrFull = errors.New("browser session capacity reached; revoke unused sessions")
)
// MaxSessions bounds persisted browser state. Active sessions are never evicted
// to admit a new login.
const MaxSessions = 128
// Record is internal state. HTTP adapters must expose a separate public view;
// neither the verifier nor the login credential fingerprint belongs in it.
type Record struct {
ID string `json:"id"`
Verifier string `json:"verifier"`
Credential string `json:"credential"`
Name string `json:"name"`
Created time.Time `json:"created"`
LastSeen time.Time `json:"last_seen"`
Expires time.Time `json:"expires"`
RemoteIP string `json:"remote_ip"`
UserAgent string `json:"user_agent"`
}
func (r Record) Valid(now time.Time, idle time.Duration) bool {
// Requests sample now before entering the repository. A newer request
// may have already committed activity by the time this one reads it.
return r.ID != "" && r.Verifier != "" && r.Credential != "" &&
!r.Created.IsZero() && !r.LastSeen.Before(r.Created) &&
now.Before(r.Expires) && now.Before(r.LastSeen.Add(idle))
}
// Repository methods are atomic relative to revocation. Access must never
// recreate a record deleted by a concurrent logout. A failed commit is an error.
type Repository interface {
ReplaceBrowserSession(Record, string, time.Time, time.Duration) error
AccessBrowserSession(string, time.Time, time.Duration, bool) (Record, error)
ListBrowserSessions(time.Time, time.Duration) ([]Record, error)
RevokeBrowserSession(string) error
ClearBrowserSessions() error
}
type Manager struct {
repo Repository
lifetime time.Duration
idle time.Duration
}
// New starts one browser-session authority. Restart invalidates every existing
// browser session; API credentials have an independent lifecycle.
func New(repo Repository, lifetime, idle time.Duration) (*Manager, error) {
if repo == nil || lifetime < time.Second || idle < time.Second || idle > lifetime {
return nil, errors.New("invalid browser session store or policy")
}
if err := repo.ClearBrowserSessions(); err != nil {
return nil, err
}
return &Manager{repo: repo, lifetime: lifetime, idle: idle}, nil
}
func Hash(secret string) string {
sum := sha256.Sum256([]byte(secret))
return hex.EncodeToString(sum[:])
}
func (m *Manager) Create(name, credential, previous, remoteIP, userAgent string, now time.Time) (string, Record, error) {
if name == "" || credential == "" {
return "", Record{}, ErrInvalid
}
var secretBytes [32]byte
if _, err := rand.Read(secretBytes[:]); err != nil {
return "", Record{}, err
}
var idBytes [16]byte
if _, err := rand.Read(idBytes[:]); err != nil {
return "", Record{}, err
}
secret := base64.RawURLEncoding.EncodeToString(secretBytes[:])
if len(userAgent) > 256 {
userAgent = userAgent[:256]
}
record := Record{ID: hex.EncodeToString(idBytes[:]), Verifier: Hash(secret), Credential: credential, Name: name,
Created: now, LastSeen: now, Expires: now.Add(m.lifetime), RemoteIP: remoteIP, UserAgent: userAgent}
previousHash := ""
if previous != "" {
previousHash = Hash(previous)
}
if err := m.repo.ReplaceBrowserSession(record, previousHash, now, m.idle); err != nil {
return "", Record{}, err
}
return secret, record, nil
}
func (m *Manager) Access(secret string, now time.Time, touch bool) (Record, error) {
if len(secret) != 43 {
return Record{}, ErrInvalid
}
return m.repo.AccessBrowserSession(Hash(secret), now, m.idle, touch)
}
func (m *Manager) List(now time.Time) ([]Record, error) {
return m.repo.ListBrowserSessions(now, m.idle)
}
func (m *Manager) Revoke(id string) error { return m.repo.RevokeBrowserSession(id) }
func (m *Manager) RevokeAll() error { return m.repo.ClearBrowserSessions() }
package signatures
import (
"archive/zip"
"bytes"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/yara"
)
const (
forgeReleasesURL = "https://api.github.com/repos/YARAHQ/yara-forge/releases/latest"
forgeHTTPTimeout = 30 * time.Second
forgeMaxZIPSize = 20 * 1024 * 1024
// forgeMaxYarSize caps the decompressed size of a single .yar entry. The
// compressed ZIP is bounded by forgeMaxZIPSize, but a zip bomb (or a
// compromised CDN / signing key) can encode a far larger decompressed
// payload; the full ruleset tier is a few MiB, so 64 MiB is generous.
forgeMaxYarSize = 64 * 1024 * 1024
)
var forgeTierAsset = map[string]string{
"core": "packages/core/yara-rules-core.yar",
"extended": "packages/extended/yara-rules-extended.yar",
"full": "packages/full/yara-rules-full.yar",
}
var forgeAtomicWrite = atomicio.AtomicWrite
// ForgeUpdate checks for a new YARA Forge release and downloads it if newer.
// A detached signature is fetched from the ZIP URL + ".sig" and verified
// against the raw ZIP content before extraction.
func ForgeUpdate(rulesDir, tier, currentVersion, signingKey string, disabledRules []string) (newVersion string, ruleCount int, err error) {
return ForgeUpdateFromURL(rulesDir, tier, currentVersion, signingKey, "", disabledRules)
}
// ForgeUpdateFromURL is ForgeUpdate with an explicit signed ZIP source. The
// downloadURL may contain {tier} and {version}; the signature is fetched from
// the resolved ZIP URL plus ".sig".
func ForgeUpdateFromURL(rulesDir, tier, currentVersion, signingKey, downloadURL string, disabledRules []string) (newVersion string, ruleCount int, err error) {
if _, ok := forgeTierAsset[tier]; !ok {
return "", 0, fmt.Errorf("unknown YARA Forge tier: %q (valid: core, extended, full)", tier)
}
if e := requireSigningKey(signingKey); e != nil {
return "", 0, e
}
if strings.TrimSpace(downloadURL) == "" {
return "", 0, fmt.Errorf("signatures.yara_forge.download_url is required: upstream YARA Forge does not publish CSM detached signatures")
}
latestTag, err := forgeResolveLatestTag(downloadURL, tier)
if err != nil {
return "", 0, fmt.Errorf("checking YARA Forge release: %w", err)
}
if latestTag == currentVersion {
return currentVersion, 0, nil
}
zipURL := forgeDownloadURL(downloadURL, tier, latestTag)
zipData, err := forgeDownload(zipURL)
if err != nil {
return "", 0, fmt.Errorf("downloading YARA Forge %s: %w", tier, err)
}
sig, err := fetchSignature(zipURL + ".sig")
if err != nil {
return "", 0, fmt.Errorf("YARA Forge signature verification required but failed: %w", err)
}
if e := VerifySignature(signingKey, zipData, sig); e != nil {
return "", 0, fmt.Errorf("YARA Forge signature invalid: %w", e)
}
assetPath := forgeTierAsset[tier]
yarContent, err := forgeExtractYar(zipData, assetPath)
if err != nil {
return "", 0, fmt.Errorf("extracting YARA Forge rules: %w", err)
}
yarContent = filterDisabledRules(yarContent, mergeDisabledRules(disabledRules))
ruleCount = countRules(yarContent)
outFile := filepath.Join(rulesDir, fmt.Sprintf("yara-forge-%s.yar", tier))
outFileExisted := false
if _, err := os.Stat(outFile); err == nil {
outFileExisted = true
} else if !os.IsNotExist(err) {
return "", 0, fmt.Errorf("checking existing Forge tier: %w", err)
}
if err := yara.TestCompile(string(yarContent)); err != nil {
return "", 0, fmt.Errorf("YARA compilation test failed (keeping existing rules): %w", err)
}
if err := os.MkdirAll(rulesDir, 0700); err != nil {
return "", 0, fmt.Errorf("creating rules dir: %w", err)
}
if outFileExisted {
if err := removeInactiveForgeTiers(rulesDir, tier); err != nil {
return "", 0, err
}
if err := forgeAtomicWrite(outFile, 0600, yarContent); err != nil {
return "", 0, fmt.Errorf("installing rules: %w", err)
}
return latestTag, ruleCount, nil
}
if err := forgeAtomicWrite(outFile, 0600, yarContent); err != nil {
return "", 0, fmt.Errorf("installing rules: %w", err)
}
if err := removeInactiveForgeTiers(rulesDir, tier); err != nil {
return "", 0, err
}
return latestTag, ruleCount, nil
}
func removeInactiveForgeTiers(rulesDir, activeTier string) error {
for t := range forgeTierAsset {
if t == activeTier {
continue
}
path := filepath.Join(rulesDir, fmt.Sprintf("yara-forge-%s.yar", t))
if err := os.Remove(path); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("removing inactive Forge tier %s: %w", t, err)
}
}
return nil
}
func forgeDownloadURL(tmpl, tier, version string) string {
url := strings.TrimSpace(tmpl)
url = strings.ReplaceAll(url, "{tier}", tier)
url = strings.ReplaceAll(url, "{version}", version)
return url
}
// errForgePointerAbsent signals that the mirror publishes no latest-version
// pointer (404 or a template without a {version} directory), so the caller
// should fall back to the upstream GitHub release tag.
var errForgePointerAbsent = errors.New("mirror latest pointer absent")
// forgeTagPattern constrains a version tag before it is interpolated into a
// download URL. YARA Forge tags are short date/version tokens (e.g. 20260705,
// v2026.04.11); restricting the charset stops a tampered pointer from
// redirecting the download via path traversal.
var forgeTagPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._-]{0,63}$`)
func forgeValidTag(tag string) bool {
return forgeTagPattern.MatchString(strings.TrimSpace(tag))
}
// forgeLatestPointerURL derives the mirror's latest-version pointer URL from
// the download template. The pointer lives in the directory that holds the
// per-{version} subdirectories:
//
// https://host/csm/yara-forge/{version}/yara-forge-rules-{tier}.zip
// -> https://host/csm/yara-forge/latest
//
// Returns ("", false) when the template has no {version} path segment, in
// which case there is no version-scoped directory to point into.
func forgeLatestPointerURL(tmpl, tier string) (string, bool) {
tmpl = strings.ReplaceAll(strings.TrimSpace(tmpl), "{tier}", tier)
schemeEnd := strings.Index(tmpl, "://")
if schemeEnd < 0 {
return "", false
}
authorityStart := schemeEnd + len("://")
authorityEndRel := strings.IndexAny(tmpl[authorityStart:], "/?#")
if authorityEndRel < 0 {
return "", false
}
pathStart := authorityStart + authorityEndRel
if tmpl[pathStart] != '/' {
return "", false
}
pathEnd := len(tmpl)
if pathEndRel := strings.IndexAny(tmpl[pathStart:], "?#"); pathEndRel >= 0 {
pathEnd = pathStart + pathEndRel
}
const versionPlaceholder = "{version}"
for search := pathStart; search < pathEnd; {
rel := strings.Index(tmpl[search:pathEnd], versionPlaceholder)
if rel < 0 {
return "", false
}
i := search + rel
after := i + len(versionPlaceholder)
standaloneSegment := tmpl[i-1] == '/' && (after == pathEnd || tmpl[after] == '/')
if standaloneSegment {
base := tmpl[:i]
if strings.Contains(base, versionPlaceholder) {
return "", false
}
return base + "latest", true
}
search = after
}
return "", false
}
// forgeLatestTagFromMirror fetches the mirror's latest-version pointer and
// returns the tag it names. It returns errForgePointerAbsent when the pointer
// is not published so the caller can fall back to the GitHub release API.
func forgeLatestTagFromMirror(downloadURL, tier string) (string, error) {
pointerURL, ok := forgeLatestPointerURL(downloadURL, tier)
if !ok {
return "", errForgePointerAbsent
}
client := &http.Client{Timeout: forgeHTTPTimeout}
resp, err := client.Get(pointerURL)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode == http.StatusNotFound {
return "", errForgePointerAbsent
}
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("mirror latest pointer returned %d", resp.StatusCode)
}
raw, err := io.ReadAll(io.LimitReader(resp.Body, 128))
if err != nil {
return "", fmt.Errorf("reading mirror latest pointer: %w", err)
}
tag := strings.TrimSpace(string(raw))
if !forgeValidTag(tag) {
return "", fmt.Errorf("mirror latest pointer has invalid tag %q", tag)
}
return tag, nil
}
// forgeResolveLatestTag resolves the newest downloadable version. It prefers
// the mirror's own latest pointer -- the mirror only holds versions it has
// signed and published, so resolving from it can never request a version the
// mirror lacks (the root cause of the release-day 404 gap when GitHub outran
// the weekly mirror sync). It falls back to the upstream GitHub release tag
// only when the mirror publishes no pointer, preserving prior behavior for
// older mirror layouts and download_url templates without a {version} segment.
func forgeResolveLatestTag(downloadURL, tier string) (string, error) {
tag, err := forgeLatestTagFromMirror(downloadURL, tier)
if err == nil {
return tag, nil
}
if !errors.Is(err, errForgePointerAbsent) {
return "", err
}
tag, err = forgeLatestTag()
if err != nil {
return "", err
}
if !forgeValidTag(tag) {
return "", fmt.Errorf("GitHub release has invalid tag %q", tag)
}
return tag, nil
}
func forgeLatestTag() (string, error) {
client := &http.Client{Timeout: forgeHTTPTimeout}
resp, err := client.Get(forgeReleasesURL)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", fmt.Errorf("GitHub API returned %d", resp.StatusCode)
}
var release struct {
TagName string `json:"tag_name"`
}
if err := json.NewDecoder(io.LimitReader(resp.Body, 1024*1024)).Decode(&release); err != nil {
return "", fmt.Errorf("parsing release JSON: %w", err)
}
if release.TagName == "" {
return "", fmt.Errorf("empty tag_name in release")
}
return release.TagName, nil
}
func forgeDownload(url string) ([]byte, error) {
client := &http.Client{Timeout: forgeHTTPTimeout}
resp, err := client.Get(url)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("download returned %d", resp.StatusCode)
}
data, err := io.ReadAll(io.LimitReader(resp.Body, forgeMaxZIPSize))
if err != nil {
return nil, fmt.Errorf("reading response: %w", err)
}
return data, nil
}
func forgeExtractYar(zipData []byte, assetPath string) ([]byte, error) {
reader, err := zip.NewReader(bytes.NewReader(zipData), int64(len(zipData)))
if err != nil {
return nil, fmt.Errorf("opening ZIP: %w", err)
}
for _, f := range reader.File {
if f.Name == assetPath {
rc, err := f.Open()
if err != nil {
return nil, fmt.Errorf("opening %s in ZIP: %w", assetPath, err)
}
defer rc.Close()
data, err := io.ReadAll(io.LimitReader(rc, forgeMaxYarSize+1))
if err != nil {
return nil, fmt.Errorf("reading %s: %w", assetPath, err)
}
if len(data) > forgeMaxYarSize {
return nil, fmt.Errorf("%s exceeds decompressed size limit of %d bytes", assetPath, forgeMaxYarSize)
}
return data, nil
}
}
return nil, fmt.Errorf("asset %s not found in ZIP", assetPath)
}
func filterDisabledRules(content []byte, disabled []string) []byte {
return yara.StripRules(content, disabled)
}
func extractRuleName(line string) string {
return yara.RuleNameFromLine(line)
}
func countRules(content []byte) int {
count := 0
for _, line := range strings.Split(string(content), "\n") {
trimmed := strings.TrimSpace(line)
if extractRuleName(trimmed) != "" {
count++
}
}
return count
}
package signatures
import "github.com/pidginhost/csm/internal/yara"
// forgeSuppressedRules returns the built-in Forge rule suppressions. The list
// and the stripping live in internal/yara so the compiler and the downloader
// agree on exactly which rules are active; see internal/yara/suppressed.go for
// the criteria an entry has to meet.
func forgeSuppressedRules() []string {
return yara.SuppressedRuleNames()
}
// mergeDisabledRules combines operator-configured disabled rule names with the
// built-in suppressions, dropping duplicates. Built-ins apply even when the
// operator configured nothing, so a clean install does not inherit a known
// false-positive flood.
func mergeDisabledRules(operator []string) []string {
builtin := forgeSuppressedRules()
seen := make(map[string]struct{}, len(operator)+len(builtin))
merged := make([]string, 0, len(operator)+len(builtin))
for _, name := range append(builtin, operator...) {
if name == "" {
continue
}
if _, dup := seen[name]; dup {
continue
}
seen[name] = struct{}{}
merged = append(merged, name)
}
return merged
}
package signatures
import "sync"
var (
globalScanner *Scanner
globalOnce sync.Once
)
// Init initializes the global scanner with rules from the given directory.
// Safe to call multiple times - only the first call takes effect.
// Call Reload() on the returned scanner to reload rules (e.g., on SIGHUP).
func Init(rulesDir string, disabled ...string) *Scanner {
globalOnce.Do(func() {
globalScanner = NewScanner(rulesDir, disabled...)
})
return globalScanner
}
// Global returns the global scanner, or nil if Init() hasn't been called.
func Global() *Scanner {
return globalScanner
}
// SetGlobal replaces the global scanner and returns the previous one.
//
// Init is guarded by a sync.Once, so only the first call in a process takes
// effect. That is right for the daemon, which initializes once, but it means
// a test helper that calls Init to install its own rules silently does
// nothing after the first test has run -- and its cleanup silently fails to
// restore anything. Tests that need to swap rules use this instead, and put
// the previous scanner back when they finish.
func SetGlobal(s *Scanner) *Scanner {
previous := globalScanner
globalScanner = s
return previous
}
package signatures
import (
"errors"
"fmt"
"io"
"os"
"path/filepath"
"regexp"
"strings"
"sync"
"github.com/pidginhost/csm/internal/contenttype"
"gopkg.in/yaml.v3"
)
// Rule represents a single malware detection rule loaded from an external file.
type Rule struct {
Name string `yaml:"name"`
Description string `yaml:"description"`
Severity string `yaml:"severity"` // "critical", "high", "warning"
Category string `yaml:"category"` // "webshell", "backdoor", "phishing", "dropper", "exploit"
FileTypes []string `yaml:"file_types"` // [".php", ".html", "*"] - which extensions to scan
Patterns []string `yaml:"patterns"` // literal string patterns (case-insensitive match)
Regexes []string `yaml:"regexes"` // regex patterns (for complex matching)
ExcludePatterns []string `yaml:"exclude_patterns"` // if any match, rule is skipped (false positive reduction)
ExcludeRegexes []string `yaml:"exclude_regexes"` // regex exclusions
MinMatch int `yaml:"min_match"` // minimum patterns that must match (default: 1)
RequireRegex bool `yaml:"require_regex"` // if true, at least one regex must match in addition to min_match
// MaxFileBytes skips the rule for content larger than this many bytes
// (0 = no limit). MaxFileBytesExemptRegexes retain high-confidence
// structural matches above the bound. This bounds weak heuristics by size
// without making padding an escape from stronger branches.
MaxFileBytes int `yaml:"max_file_bytes"`
MaxFileBytesExemptRegexes []string `yaml:"max_file_bytes_exempt_regexes"`
// Populated by compile().
compiledRegexes []*compiledRegex
compiledExcludeRegexes []*compiledRegex
compiledMaxFileBytesExemptRegexes []*compiledRegex
}
// RuleFile is the top-level structure of a rules YAML file.
type RuleFile struct {
Version int `yaml:"version"`
Updated string `yaml:"updated"`
Rules []Rule `yaml:"rules"`
}
// Scanner holds compiled rules and provides file scanning.
type Scanner struct {
mu sync.RWMutex
rules []Rule
version int
rulesDir string
loadErr error
// disabled holds the rule names the operator switched off, and
// disabledUnmatched the subset that matched nothing in the loaded
// ruleset. A name nobody recognises is almost always a typo, and a
// typo here reads as "the rule is off" while it keeps firing.
disabled []string
disabledUnmatched []string
disabledCount int
}
// NewScanner creates a scanner that loads rules from the given directory.
// Returns a scanner with no rules if the directory doesn't exist (not an error).
// Any load error is retained (see LoadError) so a best-effort init does not
// hide a corrupt rules directory that silently disabled all detection.
// Rule names in disabled are not loaded. This is the same operator setting
// that filters YARA-Forge downloads, applied to the rules CSM ships, so a
// misfiring signature can be switched off without editing rule files on a
// production host.
func NewScanner(rulesDir string, disabled ...string) *Scanner {
s := &Scanner{rulesDir: rulesDir, disabled: normalizeDisabled(disabled)}
s.disabledUnmatched = append([]string(nil), s.disabled...)
_ = s.Reload() // best-effort load on init; error retained via LoadError()
return s
}
// normalizeDisabled lowercases and de-duplicates the configured names, and
// drops empty entries. Rule names in the shipped files are lowercase, and an
// operator who types one in mixed case means the same rule.
func normalizeDisabled(names []string) []string {
if len(names) == 0 {
return nil
}
seen := make(map[string]struct{}, len(names))
out := make([]string, 0, len(names))
for _, name := range names {
trimmed := strings.ToLower(strings.TrimSpace(name))
if trimmed == "" {
continue
}
if _, dup := seen[trimmed]; dup {
continue
}
seen[trimmed] = struct{}{}
out = append(out, trimmed)
}
return out
}
// DisabledRules returns the rule names this scanner was told to switch off.
func (s *Scanner) DisabledRules() []string {
s.mu.RLock()
defer s.mu.RUnlock()
return append([]string(nil), s.disabled...)
}
// DisabledRulesWithoutMatch returns the configured names that matched no rule
// in the last load attempt. Config validation surfaces these: silently accepting
// a name nobody recognises is how an operator ends up believing a rule is off.
func (s *Scanner) DisabledRulesWithoutMatch() []string {
s.mu.RLock()
defer s.mu.RUnlock()
return append([]string(nil), s.disabledUnmatched...)
}
// DisabledRuleCount counts rules omitted from the installed ruleset by config.
func (s *Scanner) DisabledRuleCount() int {
s.mu.RLock()
defer s.mu.RUnlock()
return s.disabledCount
}
// LoadError returns the error from the most recent Reload, or nil if the last
// load was clean. A non-nil value means at least one rule file failed to load
// (corrupt YAML, unreadable file, bad regex); the rules that did load are still
// installed. Callers that can alert should surface this loudly at startup.
func (s *Scanner) LoadError() error {
s.mu.RLock()
defer s.mu.RUnlock()
return s.loadErr
}
// Reload loads/reloads all .yml and .yaml rule files from the rules directory.
func (s *Scanner) Reload() error {
if s.rulesDir == "" {
s.setLoadErr(nil)
return nil
}
entries, err := os.ReadDir(s.rulesDir)
if err != nil {
if os.IsNotExist(err) {
s.setLoadErr(nil)
return nil // no rules dir = no rules, not an error
}
e := fmt.Errorf("reading rules dir %s: %w", s.rulesDir, err)
s.setLoadErr(e)
return e
}
var allRules []Rule
maxVersion := 0
fileCount := 0
disabledCount := 0
disabled := make(map[string]struct{}, len(s.disabled))
for _, name := range s.disabled {
disabled[name] = struct{}{}
}
disabledSeen := make(map[string]struct{}, len(disabled))
shared := make(map[string]*compiledRegex)
// One corrupt or unreadable file must not abort the whole load: an
// attacker or a fat-fingered operator dropping one bad file would
// otherwise silently disable every other signature. Bad files/rules are
// skipped and logged; their errors are aggregated and returned so a
// caller that can alert still sees the failure.
var loadErrs []error
for _, entry := range entries {
name := entry.Name()
if entry.IsDir() {
continue
}
ext := strings.ToLower(filepath.Ext(name))
if ext != ".yml" && ext != ".yaml" {
continue
}
fileCount++
path := filepath.Join(s.rulesDir, name)
// #nosec G304 -- filepath.Join under operator-configured rulesDir.
data, err := os.ReadFile(path)
if err != nil {
loadErrs = append(loadErrs, fmt.Errorf("reading %s: %w", path, err))
fmt.Fprintf(os.Stderr, "signatures: skipping %s: %v\n", path, err)
continue
}
var rf RuleFile
if err := yaml.Unmarshal(data, &rf); err != nil {
loadErrs = append(loadErrs, fmt.Errorf("parsing %s: %w", path, err))
fmt.Fprintf(os.Stderr, "signatures: skipping %s: %v\n", path, err)
continue
}
// Compile rules
rulesBeforeFile := len(allRules)
for i := range rf.Rules {
rule := &rf.Rules[i]
if _, off := disabled[strings.ToLower(rule.Name)]; off {
disabledSeen[strings.ToLower(rule.Name)] = struct{}{}
disabledCount++
continue
}
if err := rule.compileShared(shared); err != nil {
loadErrs = append(loadErrs, fmt.Errorf("compiling rule %q in %s: %w", rule.Name, path, err))
fmt.Fprintf(os.Stderr, "signatures: skipping rule %q in %s: %v\n", rule.Name, path, err)
continue
}
if rule.MinMatch == 0 {
rule.MinMatch = 1
}
allRules = append(allRules, *rule)
}
if len(allRules) > rulesBeforeFile && rf.Version > maxVersion {
maxVersion = rf.Version
}
}
if fileCount == 0 {
s.mu.RLock()
hadRules := len(s.rules) > 0
s.mu.RUnlock()
if hadRules {
e := fmt.Errorf("no signature rule files found in %s", s.rulesDir)
s.setLoadErr(e)
return e
}
s.setLoadErr(nil)
return nil
}
var unmatched []string
for _, name := range s.disabled {
if _, seen := disabledSeen[name]; !seen {
unmatched = append(unmatched, name)
}
}
// Keep the old set on failure, but publish a clean load that config
// intentionally emptied. Retaining old rules in that case scans a set
// that is no longer on disk.
if len(allRules) == 0 && (disabledCount == 0 || len(loadErrs) > 0) {
err := errors.Join(append(loadErrs, fmt.Errorf("no signature rules loaded from %s", s.rulesDir))...)
s.mu.Lock()
s.loadErr = err
s.disabledUnmatched = unmatched
s.mu.Unlock()
return err
}
s.mu.Lock()
s.rules = allRules
s.version = maxVersion
s.loadErr = errors.Join(loadErrs...)
s.disabledUnmatched = unmatched
s.disabledCount = disabledCount
s.mu.Unlock()
if len(disabledSeen) > 0 {
fmt.Fprintf(os.Stderr, "signatures: %d rule(s) disabled by configuration\n", len(disabledSeen))
}
fmt.Fprintf(os.Stderr, "signatures: loaded %d rules (version %d) from %s\n", len(allRules), maxVersion, s.rulesDir)
// Rules installed, but report any skipped files so the caller can alert.
return errors.Join(loadErrs...)
}
// setLoadErr records the outcome of a load under the write lock.
func (s *Scanner) setLoadErr(err error) {
s.mu.Lock()
s.loadErr = err
s.mu.Unlock()
}
// compile pre-compiles regex patterns for a rule.
func (r *Rule) compile() error {
return r.compileShared(make(map[string]*compiledRegex))
}
// compileShared compiles the rule, reusing a regex already compiled from the
// same source for an earlier rule in shared, so a scan evaluates it once.
func (r *Rule) compileShared(shared map[string]*compiledRegex) error {
if r.MaxFileBytes < 0 {
return fmt.Errorf("max_file_bytes must be non-negative")
}
var err error
if r.compiledRegexes, err = compileRuleRegexes(shared, r.Regexes, "invalid regex"); err != nil {
return err
}
if r.compiledExcludeRegexes, err = compileRuleRegexes(shared, r.ExcludeRegexes, "invalid exclude regex"); err != nil {
return err
}
if r.compiledMaxFileBytesExemptRegexes, err = compileRuleRegexes(shared, r.MaxFileBytesExemptRegexes, "invalid max_file_bytes_exempt_regex"); err != nil {
return err
}
return nil
}
func compileRuleRegexes(shared map[string]*compiledRegex, patterns []string, errLabel string) ([]*compiledRegex, error) {
var out []*compiledRegex
for _, pattern := range patterns {
src := "(?i)" + pattern // rule regexes are case-insensitive
cr, ok := shared[src]
if !ok {
re, err := regexp.Compile(src)
if err != nil {
return nil, fmt.Errorf("%s '%s': %w", errLabel, pattern, err)
}
cr = &compiledRegex{Regexp: re, gate: gateFor(src)}
shared[src] = cr
}
out = append(out, cr)
}
return out, nil
}
// Match represents a rule that matched a file.
type Match struct {
RuleName string
Description string
Severity string
Category string
Matched []string // which patterns matched
}
// ScanContent scans file content against loaded rules.
// fileExt should include the dot (e.g., ".php").
func (s *Scanner) ScanContent(content []byte, fileExt string) []Match {
return s.ScanContentWithSize(content, fileExt, int64(len(content)))
}
// ScanContentWithSize scans content while using contentSize as the complete
// snapshot size for per-rule bounds. Prefix-scanning callers should pass the
// size of the open file represented by the prefix; ordinary callers should use
// ScanContent. A size smaller than the supplied bytes is raised to len(content)
// so a bad caller cannot turn a bounded rule back on for oversized content.
func (s *Scanner) ScanContentWithSize(content []byte, fileExt string, contentSize int64) []Match {
s.mu.RLock()
defer s.mu.RUnlock()
if len(s.rules) == 0 {
return nil
}
extLower := strings.ToLower(fileExt)
// Only a file that is an archive by name as well as by magic is left to the
// extraction-time scan; PHP executes past any leading bytes, so magic alone
// must never switch the rules off for an executable name.
if contenttype.IsArchiveExt(extLower) && contenttype.IsCompressedArchive(content) {
return nil
}
if contentSize < int64(len(content)) {
contentSize = int64(len(content))
}
contentLower := strings.ToLower(string(content))
eval := newRegexEval(content)
var matches []Match
for _, rule := range s.rules {
// Check if this rule applies to this file type
if !ruleMatchesExt(rule, extLower) {
continue
}
// Check exclusions first - if any exclude pattern matches, skip this rule
excluded := false
for _, pattern := range rule.ExcludePatterns {
if strings.Contains(contentLower, strings.ToLower(pattern)) {
excluded = true
break
}
}
if !excluded {
for _, re := range rule.compiledExcludeRegexes {
if eval.match(re) {
excluded = true
break
}
}
}
if excluded {
continue
}
if rule.MaxFileBytes > 0 && contentSize > int64(rule.MaxFileBytes) {
exempt := false
for _, re := range rule.compiledMaxFileBytesExemptRegexes {
if eval.match(re) {
exempt = true
break
}
}
if !exempt {
continue
}
}
// Count pattern matches
var matched []string
regexMatched := false
for _, pattern := range rule.Patterns {
if strings.Contains(contentLower, strings.ToLower(pattern)) {
matched = append(matched, pattern)
}
}
for _, re := range rule.compiledRegexes {
if eval.match(re) {
matched = append(matched, re.String())
regexMatched = true
}
}
if len(matched) >= rule.MinMatch && (!rule.RequireRegex || regexMatched) {
matches = append(matches, Match{
RuleName: rule.Name,
Description: rule.Description,
Severity: rule.Severity,
Category: rule.Category,
Matched: matched,
})
}
}
return matches
}
// ScanFile reads a file and scans it against loaded rules.
func (s *Scanner) ScanFile(path string, maxBytes int) []Match {
s.mu.RLock()
ruleCount := len(s.rules)
s.mu.RUnlock()
if ruleCount == 0 {
return nil
}
if maxBytes <= 0 {
return nil
}
// #nosec G304 -- ScanFile's whole purpose is to scan a file on disk;
// `path` comes from the daemon's file index walker or a fanotify event.
f, err := os.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
// ReadAll over a LimitReader, not a single Read into a pre-sized buffer:
// a bare f.Read can return a short count on the first call, which would
// hand only a prefix to the scanner and silently miss malware further
// into the file. LimitReader also makes a negative/huge maxBytes safe
// (no make([]byte, maxBytes) panic / over-allocation).
buf, err := io.ReadAll(io.LimitReader(f, int64(maxBytes)))
if err != nil || len(buf) == 0 {
return nil
}
contentSize := int64(len(buf))
if info, statErr := f.Stat(); statErr == nil && info.Size() > contentSize {
contentSize = info.Size()
}
ext := filepath.Ext(path)
return s.ScanContentWithSize(buf, ext, contentSize)
}
// RuleNames returns the names of the loaded rules.
func (s *Scanner) RuleNames() []string {
s.mu.RLock()
defer s.mu.RUnlock()
names := make([]string, 0, len(s.rules))
for _, r := range s.rules {
names = append(names, r.Name)
}
return names
}
// RuleCount returns the number of loaded rules.
func (s *Scanner) RuleCount() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.rules)
}
// Version returns the highest version number across loaded rule files.
func (s *Scanner) Version() int {
s.mu.RLock()
defer s.mu.RUnlock()
return s.version
}
// canonicalScanExt folds extensions that carry PHP source but are not the
// extension rules are written against. Every extension a stock PHP handler
// executes (.phtml, .pht, .php5 ...) must meet the same rules as .php, and so
// must ".phps": it is PHP source by definition -- the extension exists so a
// server can display it -- so a payload staged under it is still matched.
// Without this fold such a file is read and then compared against nothing,
// because every PHP rule declares file_types [".php"].
func canonicalScanExt(ext string) string {
ext = strings.ToLower(ext)
if ext == ".phps" || contenttype.IsExecutablePHPExt(ext) {
return ".php"
}
return ext
}
func ruleMatchesExt(rule Rule, ext string) bool {
if len(rule.FileTypes) == 0 {
return true // no filter = match all
}
ext = strings.ToLower(ext)
canonicalExt := canonicalScanExt(ext)
for _, ft := range rule.FileTypes {
ft = strings.ToLower(ft)
// The alias is one-way: PHP rules also inspect .phps source, while a
// deliberately .phps-only rule must not broaden to executable .php.
if ft == "*" || ft == ext || ft == canonicalExt {
return true
}
}
return false
}
package signatures
import (
"regexp"
"strings"
"github.com/pidginhost/csm/internal/contenttype"
)
// maxReferencedPayloadPaths bounds how many payload paths one finding carries.
// Findings travel through alert mail, the audit log and webhooks, so a file
// holding hundreds of include statements must not turn one alert into an
// unbounded message.
const maxReferencedPayloadPaths = 5
// referencedPayloadWindow is how far a quoted path may sit from the
// include/require keyword that consumes it. The keyword and the literal are
// usually in the same statement; a local assigned one line earlier and
// included on the next is the shape the 2026-09-17 loader used.
const referencedPayloadWindow = 300
// nonExecutablePayloadLiteral matches a quoted path whose extension belongs to
// an image, archive or opaque data file. The extension list mirrors the one in
// the backdoor_include_nonexecutable rules; source partials (.php, .html,
// .tpl, .txt, .svg) are deliberately absent, because including those is templating.
var nonExecutablePayloadLiteral = regexp.MustCompile(
`(?i)['"]([^'"\r\n]{1,240}\.(?:png|jp(?:eg?|g)|gif|bmp|ico|cur|webp|tiff?|zip|rar|tar|gz|bz2|7z|log|dat|bin|cache|bak|old|csv|pdf|woff2?|ttf|eot|otf))['"]`,
)
// includeKeyword matches a PHP inclusion keyword in statement position.
var includeKeyword = regexp.MustCompile(`(?i)(?:^|[^A-Za-z0-9_$>])(?:include|require)(?:_once)?[\s('$"]`)
// ReferencedPayloadPaths returns the non-executable files that PHP content
// pulls in through include or require, in the order they appear and without
// duplicates.
//
// A loader finding names the PHP file that changed. The payload lives
// somewhere else -- in the incident that prompted this, a picture in a plugin
// asset directory two levels away -- and survives a clean-up of the PHP alone,
// so the finding has to name it too.
//
// Association is positional. RE2 cannot prove that the local assigned a path
// is the same local an include consumes, so a quoted payload path counts when
// an inclusion keyword sits within referencedPayloadWindow bytes of it. The
// output is remediation context attached to a finding that already fired, not
// a detection signal.
func ReferencedPayloadPaths(content []byte) []string {
// Module loaders spell require and include too, and a JavaScript bundle
// pulling in a sprite is not a loader for a hidden payload.
if !contenttype.HasPHPOpenTag(content) {
return nil
}
keyword := includeKeyword.FindIndex(content)
if keyword == nil {
return nil
}
var paths []string
seen := make(map[string]bool)
// Advance both cursors monotonically. Materializing every match before
// enforcing the output cap wastes memory; restarting the keyword search
// for each repeated literal makes enrichment quadratic on hostile input.
for offset := 0; offset < len(content); {
literal := nonExecutablePayloadLiteral.FindSubmatchIndex(content[offset:])
if literal == nil {
break
}
for i := range literal {
literal[i] += offset
}
offset = literal[1]
for keyword[1] <= literal[0]-referencedPayloadWindow {
next := keyword[1]
keyword = includeKeyword.FindIndex(content[next:])
if keyword == nil {
return paths
}
keyword[0] += next
keyword[1] += next
}
if keyword[0] >= literal[1]+referencedPayloadWindow {
continue
}
path := string(content[literal[2]:literal[3]])
if seen[path] {
continue
}
seen[path] = true
paths = append(paths, path)
if len(paths) == maxReferencedPayloadPaths {
break
}
}
return paths
}
// ReferencedPayloadDetail renders ReferencedPayloadPaths as a line to append
// to a finding's details, or an empty string when the content pulls in no
// non-executable file. Callers concatenate it unconditionally.
func ReferencedPayloadDetail(content []byte) string {
paths := ReferencedPayloadPaths(content)
if len(paths) == 0 {
return ""
}
return "\nIncluded payload files: " + strings.Join(paths, ", ")
}
package signatures
import (
"bytes"
"errors"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"time"
"gopkg.in/yaml.v3"
"github.com/pidginhost/csm/internal/atomicio"
)
// ErrUpdateRollback marks a signed update refused by rollback protection.
var ErrUpdateRollback = errors.New("signed rules update refused")
// UpdateOptions controls operator-approved exceptions to rollback protection.
type UpdateOptions struct {
AllowRuleCountDecrease bool
}
// Update downloads the latest rules from the configured URL.
// Validates the downloaded rules before installing.
// A detached ed25519 signature is fetched from url+".sig" and verified
// before the rules are installed.
// Returns the number of rules loaded, or error.
func Update(rulesDir, url, signingKey string, options UpdateOptions) (int, error) {
if url == "" {
return 0, fmt.Errorf("no update URL configured (set signatures.update_url in csm.yaml)")
}
if err := requireSigningKey(signingKey); err != nil {
return 0, err
}
// Download
client := &http.Client{Timeout: 30 * time.Second}
resp, err := client.Get(url)
if err != nil {
return 0, fmt.Errorf("downloading rules from %s: %w", url, err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != 200 {
return 0, fmt.Errorf("download failed: HTTP %d from %s", resp.StatusCode, url)
}
data, err := io.ReadAll(io.LimitReader(resp.Body, 10*1024*1024)) // max 10MB
if err != nil {
return 0, fmt.Errorf("reading response: %w", err)
}
sig, err := fetchSignature(url + ".sig")
if err != nil {
return 0, fmt.Errorf("signature verification required but failed: %w", err)
}
if err := VerifySignature(signingKey, data, sig); err != nil {
return 0, fmt.Errorf("rules signature invalid: %w", err)
}
// Validate: must parse as valid YAML rules
var rf RuleFile
if err := yaml.Unmarshal(data, &rf); err != nil {
return 0, fmt.Errorf("invalid rules file: %w", err)
}
if len(rf.Rules) == 0 {
return 0, fmt.Errorf("rules file contains no rules")
}
// Validate each rule compiles
for _, rule := range rf.Rules {
if err := rule.compile(); err != nil {
return 0, fmt.Errorf("rule '%s' failed validation: %w", rule.Name, err)
}
}
destPath := filepath.Join(rulesDir, "malware.yml")
// The daemon rescans every file on the host when a rules file changes on
// disk, so re-downloading the installed ruleset must not rewrite it.
// #nosec G304 -- destPath is under the operator-configured rules dir.
if installed, err := os.ReadFile(destPath); err == nil && bytes.Equal(installed, data) {
return len(rf.Rules), nil
}
if err := refuseRollback(destPath, rf, options.AllowRuleCountDecrease); err != nil {
return 0, err
}
// Ensure rules directory exists
if err := os.MkdirAll(rulesDir, 0700); err != nil {
return 0, fmt.Errorf("creating rules dir: %w", err)
}
// Atomic write: write-temp, fsync, rename, dir-fsync. The daemon reloads
// these rules on the next tick, so a torn write must never be observable.
if err := atomicio.AtomicWrite(destPath, 0600, data); err != nil {
return 0, fmt.Errorf("installing rules: %w", err)
}
return len(rf.Rules), nil
}
// refuseRollback rejects a validly signed update that would move the
// installed ruleset backwards: an older version number is a replayed release,
// and a rule count that collapses to under half of what is installed is a
// stale mirror or a truncated publish rather than ordinary churn. Either
// would silently strip detection while reporting a successful update. A
// missing or unparsable installed file gives nothing to compare against and
// is not protected: the signed update is the recovery path out of that state.
func refuseRollback(destPath string, next RuleFile, allowRuleCountDecrease bool) error {
current, err := os.ReadFile(destPath) // #nosec G304 -- operator-configured rules dir.
if errors.Is(err, os.ErrNotExist) {
return nil
}
if err != nil {
return fmt.Errorf("reading installed rules: %w", err)
}
installed, ok := parseInstalledRules(current)
if !ok {
return nil
}
if next.Version < installed.Version {
return fmt.Errorf("%w: refusing rules downgrade: update is version %d, installed rules are version %d", ErrUpdateRollback, next.Version, installed.Version)
}
if !allowRuleCountDecrease && len(next.Rules)*2 < len(installed.Rules) {
return fmt.Errorf("%w: refusing rules rollback: update carries %d rules, installed rules carry %d", ErrUpdateRollback, len(next.Rules), len(installed.Rules))
}
return nil
}
// parseInstalledRules reports false for an unparsable or empty installed
// file: the states a signed update must be allowed to repair.
func parseInstalledRules(data []byte) (RuleFile, bool) {
var installed RuleFile
if yaml.Unmarshal(data, &installed) != nil || len(installed.Rules) == 0 {
return RuleFile{}, false
}
return installed, true
}
package signatures
import (
"crypto/ed25519"
"encoding/hex"
"fmt"
"io"
"net/http"
"time"
)
func requireSigningKey(signingKey string) error {
if signingKey == "" {
return fmt.Errorf("signatures.signing_key is required for remote rule updates")
}
return nil
}
// VerifySignature checks an ed25519 signature over data using a hex-encoded public key.
// Returns nil if the signature is valid.
func VerifySignature(pubKeyHex string, data, signature []byte) error {
pubKeyBytes, err := hex.DecodeString(pubKeyHex)
if err != nil {
return fmt.Errorf("invalid signing key (bad hex): %w", err)
}
if len(pubKeyBytes) != ed25519.PublicKeySize {
return fmt.Errorf("invalid signing key length: got %d bytes, want %d", len(pubKeyBytes), ed25519.PublicKeySize)
}
if len(signature) != ed25519.SignatureSize {
return fmt.Errorf("invalid signature length: got %d bytes, want %d", len(signature), ed25519.SignatureSize)
}
pubKey := ed25519.PublicKey(pubKeyBytes)
if !ed25519.Verify(pubKey, data, signature) {
return fmt.Errorf("signature verification failed")
}
return nil
}
// fetchSignature downloads a detached signature from url + ".sig".
func fetchSignature(sigURL string) ([]byte, error) {
client := &http.Client{Timeout: 10 * time.Second}
resp, err := client.Get(sigURL)
if err != nil {
return nil, fmt.Errorf("downloading signature from %s: %w", sigURL, err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("signature download returned HTTP %d from %s", resp.StatusCode, sigURL)
}
sig, err := io.ReadAll(io.LimitReader(resp.Body, 1024)) // ed25519 sig is 64 bytes
if err != nil {
return nil, fmt.Errorf("reading signature: %w", err)
}
return sig, nil
}
// Package sshdconf reads the effective sshd_config the way sshd itself does:
// Include directives are followed, Match blocks are ignored because their
// directives are connection-scoped, and most keywords keep first-match-wins
// semantics.
//
// Port is the exception. sshd accumulates every Port directive, so a config
// can name several listening ports; collapsing them into a single value hides
// ports and makes any firewall guard built on top of it wrong.
package sshdconf
import (
"bufio"
"crypto/sha256"
"encoding/hex"
"io"
"net"
"os"
"path/filepath"
"sort"
"strconv"
"strings"
)
// DefaultPath is where sshd reads its configuration from.
const DefaultPath = "/etc/ssh/sshd_config"
// defaultPort is OpenSSH's compiled-in listen port, used when nothing in the
// config names one.
const defaultPort = 22
// maxIncludeDepth mirrors OpenSSH's own nesting limit. Without it a config
// that includes itself recurses until the stack overflows, which Go cannot
// recover from.
const maxIncludeDepth = 16
// maxLineBytes caps how much of one line is kept. No real directive comes
// close; the cap exists so a corrupt file cannot be read into the daemon's
// memory whole. The remainder of an over-long line is discarded and parsing
// continues with the next one.
const maxLineBytes = 64 * 1024
// compiledDefaults are the OpenSSH defaults for the keywords CSM reads.
var compiledDefaults = map[string]string{
"addressfamily": "any",
"protocol": "2",
"passwordauthentication": "yes",
"permitrootlogin": "prohibit-password",
"maxauthtries": "6",
"x11forwarding": "no",
"usedns": "no",
}
// FS is the file access parsing needs. It matches the corresponding subset of
// the internal/checks OS injector, so that package can pass its own hook and
// keep its tests fake-driven.
type FS interface {
Open(name string) (*os.File, error)
Glob(pattern string) ([]string, error)
}
// OSFS reads the real filesystem.
type OSFS struct{}
// #nosec G304 -- callers pass the operator-owned sshd config path.
func (OSFS) Open(name string) (*os.File, error) { return os.Open(name) }
// Glob expands an Include pattern.
func (OSFS) Glob(pattern string) ([]string, error) { return filepath.Glob(pattern) }
// Config is a parsed sshd_config.
type Config struct {
present bool
values map[string]string
files []string
digests [][sha256.Size]byte
ports []int
listenAddresses []listenAddress
}
type listenAddress struct {
host string
port int
}
// Parse reads path and every file it includes. A missing or unreadable root
// file yields a Config of pure compiled defaults with Present reporting false,
// so callers can tell "sshd runs on defaults" from "there is no sshd config
// here at all".
func Parse(fsys FS, path string) *Config {
c := &Config{values: make(map[string]string)}
c.present = c.parseFile(fsys, path, filepath.Dir(path), 0, make(map[string]struct{}), false)
return c
}
// Present reports whether the root config file was readable.
func (c *Config) Present() bool { return c.present }
// Files lists every file the parse read, root first then Includes in the
// order sshd would read them. Change detection has to cover all of them:
// a drop-in can flip a setting without touching the root file.
func (c *Config) Files() []string {
return append([]string(nil), c.files...)
}
// Digest identifies the file paths and exact bytes consumed by Parse, in
// read order. Hashing through the parser keeps the effective values and the
// change-detection digest on one coherent filesystem snapshot.
func (c *Config) Digest() string {
h := sha256.New()
for i, file := range c.files {
h.Write([]byte(file))
h.Write([]byte{0})
h.Write(c.digests[i][:])
h.Write([]byte{0})
}
return hex.EncodeToString(h.Sum(nil))
}
// Value returns the effective value of keyword, lowercased, falling back to
// the OpenSSH compiled default. Keywords with no shipped default return "".
func (c *Config) Value(keyword string) string {
key := strings.ToLower(keyword)
if v, ok := c.values[key]; ok {
return strings.ToLower(v)
}
return compiledDefaults[key]
}
// ListenPorts returns the sorted TCP ports sshd accepts connections on.
//
// Port applies only to ListenAddress entries that omit a port, so a config
// where every ListenAddress carries its own port never binds the Port value.
func (c *Config) ListenPorts() []int {
v4, v6 := c.listenPorts(false)
return sortedUnique(append(v4, v6...))
}
// RemoteListenPorts returns the IPv4 and IPv6 ports reachable beyond the
// local host. It honors AddressFamily and explicit ListenAddress directives,
// and excludes loopback-only listeners that an inbound firewall cannot cut
// off.
func (c *Config) RemoteListenPorts() (ipv4, ipv6 []int) {
return c.listenPorts(true)
}
// parseFile reads one config file, following Include directives relative to
// rootDir. It reports whether the file was readable.
func (c *Config) parseFile(fsys FS, path, rootDir string, depth int, seen map[string]struct{}, collectOnly bool) bool {
if depth > maxIncludeDepth {
return false
}
path = filepath.Clean(path)
// A glob may match its containing file, and mutually recursive globs can
// branch exponentially before the depth limit. Re-reading a file cannot
// change this parser's first-value or set-like accumulated results.
if _, parsed := seen[path]; parsed {
return false
}
seen[path] = struct{}{}
f, err := fsys.Open(path)
if err != nil {
return false
}
defer func() { _ = f.Close() }()
c.files = append(c.files, path)
c.digests = append(c.digests, [sha256.Size]byte{})
digestIndex := len(c.digests) - 1
contentHash := sha256.New()
defer func() { copy(c.digests[digestIndex][:], contentHash.Sum(nil)) }()
inMatch := false
reader := bufio.NewReader(io.TeeReader(f, contentHash))
for {
rawLine, more := readLine(reader)
if !more {
break
}
line := strings.TrimSpace(rawLine)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
keyword, args, ok := splitDirective(line)
if !ok {
continue
}
// A Match block runs to EOF in the effective global view. Its
// directives do not apply globally, but Include targets are still read
// by sshd and must be recorded for change detection.
if keyword == "match" {
inMatch = true
continue
}
if keyword == "include" {
// sshd resolves a relative Include against its config directory,
// not against the including file, so nested drop-ins agree with
// the top-level file.
for _, pattern := range args {
if !filepath.IsAbs(pattern) {
pattern = filepath.Join(rootDir, pattern)
}
matches, globErr := fsys.Glob(pattern)
if globErr != nil {
continue
}
for _, m := range matches {
c.parseFile(fsys, m, rootDir, depth+1, seen, collectOnly || inMatch)
}
}
continue
}
if inMatch || collectOnly {
continue
}
switch keyword {
case "port":
if port, valid := parsePort(args[0]); valid {
c.ports = append(c.ports, port)
}
case "listenaddress":
c.recordListenAddress(args[0])
default:
if _, exists := c.values[keyword]; !exists {
c.values[keyword] = args[0]
}
}
}
return true
}
// readLine returns the next line without its terminator, keeping at most
// maxLineBytes of it, and reports whether one was read. bufio.Reader.ReadLine
// hands back the line in buffer-sized pieces, so the discarded tail of an
// over-long line never accumulates.
func readLine(r *bufio.Reader) (string, bool) {
var line strings.Builder
for {
chunk, isPrefix, err := r.ReadLine()
if err != nil {
return line.String(), line.Len() > 0
}
if remaining := maxLineBytes - line.Len(); remaining > 0 {
if len(chunk) > remaining {
chunk = chunk[:remaining]
}
line.Write(chunk)
}
if !isPrefix {
return line.String(), true
}
}
}
// recordListenAddress records an address and its optional explicit port.
// Forms without a port inherit every Port directive.
func (c *Config) recordListenAddress(value string) {
host, portStr, err := net.SplitHostPort(value)
if err != nil {
c.listenAddresses = append(c.listenAddresses, listenAddress{host: strings.Trim(value, "[]")})
return
}
port, valid := parsePort(portStr)
if !valid {
c.listenAddresses = append(c.listenAddresses, listenAddress{host: host})
return
}
c.listenAddresses = append(c.listenAddresses, listenAddress{host: host, port: port})
}
// splitDirective splits a keyword from its OpenSSH-style argument vector.
// Quoting, basic escapes, and comments follow argv_split in OpenSSH.
func splitDirective(line string) (keyword string, args []string, ok bool) {
sep := strings.IndexFunc(line, func(r rune) bool {
return r == ' ' || r == '\t' || r == '='
})
if sep <= 0 {
return "", nil, false
}
keyword = strings.ToLower(line[:sep])
rest := strings.TrimLeft(line[sep:], " \t")
if strings.HasPrefix(rest, "=") {
rest = strings.TrimLeft(rest[1:], " \t")
}
args, ok = splitArguments(rest)
if !ok || len(args) == 0 {
return "", nil, false
}
return keyword, args, true
}
func splitArguments(s string) ([]string, bool) {
var args []string
for i := 0; i < len(s); {
for i < len(s) && (s[i] == ' ' || s[i] == '\t') {
i++
}
if i == len(s) || s[i] == '#' {
break
}
var arg strings.Builder
var quote byte
for i < len(s) {
ch := s[i]
switch {
case ch == '\\' && i+1 < len(s) &&
(s[i+1] == '\'' || s[i+1] == '"' || s[i+1] == '\\' || (quote == 0 && s[i+1] == ' ')):
i++
arg.WriteByte(s[i])
case quote == 0 && (ch == ' ' || ch == '\t'):
i++
goto argumentDone
case quote == 0 && (ch == '\'' || ch == '"'):
quote = ch
case quote != 0 && ch == quote:
quote = 0
default:
arg.WriteByte(ch)
}
i++
}
argumentDone:
if quote != 0 {
return nil, false
}
args = append(args, arg.String())
}
return args, true
}
func parsePort(s string) (int, bool) {
port, err := strconv.Atoi(strings.TrimSpace(s))
if err == nil {
return port, port >= 1 && port <= 65535
}
port, err = net.LookupPort("tcp", s)
if err != nil || port < 1 || port > 65535 {
return 0, false
}
return port, true
}
func (c *Config) listenPorts(remoteOnly bool) (ipv4, ipv6 []int) {
base := c.ports
if len(base) == 0 {
base = []int{defaultPort}
}
allowV4 := c.Value("addressfamily") != "inet6"
allowV6 := c.Value("addressfamily") != "inet"
if len(c.listenAddresses) == 0 {
if allowV4 {
ipv4 = append(ipv4, base...)
}
if allowV6 {
ipv6 = append(ipv6, base...)
}
return sortedUnique(ipv4), sortedUnique(ipv6)
}
for _, listener := range c.listenAddresses {
isV4, isV6, loopback := addressFamilies(listener.host)
if remoteOnly && loopback {
continue
}
ports := base
if listener.port != 0 {
ports = []int{listener.port}
}
if allowV4 && isV4 {
ipv4 = append(ipv4, ports...)
}
if allowV6 && isV6 {
ipv6 = append(ipv6, ports...)
}
}
return sortedUnique(ipv4), sortedUnique(ipv6)
}
func addressFamilies(host string) (ipv4, ipv6, loopback bool) {
if host == "*" {
return true, true, false
}
canonicalHost := strings.TrimSuffix(strings.ToLower(host), ".")
if canonicalHost == "localhost" || strings.HasSuffix(canonicalHost, ".localhost") {
return true, true, true
}
if zone := strings.LastIndex(host, "%"); zone >= 0 {
host = host[:zone]
}
ip := net.ParseIP(host)
if ip == nil {
return true, true, false
}
if ip.To4() != nil {
return true, false, ip.IsLoopback()
}
return false, true, ip.IsLoopback()
}
func sortedUnique(ports []int) []int {
seen := make(map[int]struct{}, len(ports))
out := make([]int, 0, len(ports))
for _, p := range ports {
if _, dup := seen[p]; dup {
continue
}
seen[p] = struct{}{}
out = append(out, p)
}
sort.Ints(out)
return out
}
package state
import (
"fmt"
"os"
"path/filepath"
"syscall"
)
const lockFileName = "csm.lock"
// LockFile provides file-based locking to prevent concurrent CSM runs.
type LockFile struct {
path string
file *os.File
}
// AcquireLock creates an exclusive lock. Returns error if already locked.
func AcquireLock(stateDir string) (*LockFile, error) {
lockPath := filepath.Join(stateDir, lockFileName)
// #nosec G304 -- filepath.Join under operator-configured stateDir.
f, err := os.OpenFile(lockPath, os.O_CREATE|os.O_RDWR, 0600)
if err != nil {
return nil, fmt.Errorf("opening lock file: %w", err)
}
// Try non-blocking exclusive lock
// #nosec G115 -- os.File.Fd returns uintptr but POSIX file descriptors
// are small non-negative ints (rlimit ~1024); int conversion is lossless.
if err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX|syscall.LOCK_NB); err != nil {
_ = f.Close()
return nil, fmt.Errorf("another CSM instance is already running")
}
// Write PID for debugging
_ = f.Truncate(0)
fmt.Fprintf(f, "%d\n", os.Getpid())
return &LockFile{path: lockPath, file: f}, nil
}
// Release releases the lock.
func (l *LockFile) Release() {
if l.file != nil {
// #nosec G115 -- see AcquireLock: POSIX fd fits in int.
_ = syscall.Flock(int(l.file.Fd()), syscall.LOCK_UN)
_ = l.file.Close()
os.Remove(l.path)
}
}
package state
import (
"errors"
"fmt"
"io"
"io/fs"
"os"
"path/filepath"
)
// MigrateStateDir copies the contents of oldPath into newPath if and only if:
// (1) oldPath exists and is non-empty,
// (2) newPath does not exist or is empty,
// (3) the two paths are not the same.
// Returns (true, nil) when a copy occurred, (false, nil) when no migration
// was needed. Any I/O error is fatal — the caller (daemon startup) should
// abort rather than silently start with mixed state.
func MigrateStateDir(oldPath, newPath string) (bool, error) {
if oldPath == "" || newPath == "" || oldPath == newPath {
return false, nil
}
oldInfo, err := os.Stat(oldPath)
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
return false, nil
}
return false, fmt.Errorf("stat %s: %w", oldPath, err)
}
if !oldInfo.IsDir() {
return false, nil
}
oldEntries, err := os.ReadDir(oldPath)
if err != nil {
return false, fmt.Errorf("reading %s: %w", oldPath, err)
}
if len(oldEntries) == 0 {
return false, nil
}
if newEntries, err := os.ReadDir(newPath); err == nil && hasStateContent(newEntries) {
return false, nil // new dir non-empty: no migration
} else if err != nil && !errors.Is(err, fs.ErrNotExist) {
return false, fmt.Errorf("reading %s: %w", newPath, err)
}
if err := os.MkdirAll(newPath, 0o700); err != nil {
return false, fmt.Errorf("mkdir %s: %w", newPath, err)
}
copied := false
for _, e := range oldEntries {
if e.Name() == lockFileName {
continue
}
src := filepath.Join(oldPath, e.Name())
dst := filepath.Join(newPath, e.Name())
if err := copyEntry(src, dst); err != nil {
return false, fmt.Errorf("copying %s: %w", src, err)
}
copied = true
}
return copied, nil
}
func hasStateContent(entries []os.DirEntry) bool {
for _, entry := range entries {
if entry.Name() != lockFileName {
return true
}
}
return false
}
func copyEntry(src, dst string) error {
info, err := os.Stat(src)
if err != nil {
return err
}
if info.IsDir() {
if mkdirErr := os.MkdirAll(dst, info.Mode().Perm()); mkdirErr != nil {
return mkdirErr
}
entries, readDirErr := os.ReadDir(src)
if readDirErr != nil {
return readDirErr
}
for _, e := range entries {
if copyErr := copyEntry(filepath.Join(src, e.Name()), filepath.Join(dst, e.Name())); copyErr != nil {
return copyErr
}
}
return nil
}
// #nosec G304 -- src derived from os.ReadDir of a trusted state directory.
in, err := os.Open(src)
if err != nil {
return err
}
defer in.Close()
// #nosec G304 -- dst is new state directory plus an entry name from os.ReadDir.
out, err := os.OpenFile(dst, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, info.Mode().Perm())
if err != nil {
return err
}
if _, err := io.Copy(out, in); err != nil {
_ = out.Close()
return err
}
if err := out.Sync(); err != nil {
_ = out.Close()
return err
}
return out.Close()
}
package state
import (
"encoding/json"
"fmt"
"os"
"path/filepath"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/atomicio"
)
const pendingFindingsFile = "pending_findings.json"
var removePendingFindingsFile = os.Remove
// pendingFindingsMax bounds the parked batch; a stop during a flood keeps the
// newest findings rather than growing the file without limit.
const pendingFindingsMax = 10000
// AppendPendingFindings parks findings that were still queued for dispatch
// when the daemon stopped. The next start takes them back and runs them
// through the full dispatch pipeline, so a realtime-only finding raised in the
// last seconds before a restart still triggers its auto-response instead of
// surviving only as a history line nothing re-detects. Shutdown drains the
// channel twice, so this appends rather than replaces.
func (s *Store) AppendPendingFindings(findings []alert.Finding) error {
if len(findings) == 0 {
return nil
}
call := s.pendingHealth().begin(len(findings))
defer call.finish()
s.mu.Lock()
defer s.mu.Unlock()
defer call.settleIO()
call.start()
pending, err := s.readPendingLocked(call)
if err != nil {
return err
}
old := pendingIdentity(pending)
call.observe(old)
total := len(pending) + len(findings)
pending = append(pending, findings...)
if len(pending) > pendingFindingsMax {
pending = pending[len(pending)-pendingFindingsMax:]
}
next := pendingIdentity(pending)
ages := call.appendedAges(old.count, len(pending))
write := s.writePendingFile
if write == nil {
write = atomicio.AtomicWriteJSON
}
knownLoss := max(0, len(findings)-pendingFindingsMax)
if !next.valid {
knownLoss = len(findings)
}
call.offer(knownLoss)
err = write(filepath.Join(s.path, pendingFindingsFile), 0o600, pending)
if err == nil {
call.complete(next, total-len(pending), false, ages)
return nil
}
call.ioFailed()
actual, readErr := s.readPendingLocked(call)
if readErr != nil {
return err
}
image := pendingIdentity(actual)
switch {
case image.valid && image == old:
call.complete(image, len(findings), true, nil)
case image.valid && next.valid && image == next:
call.complete(image, total-len(pending), true, ages)
default:
call.unreadable(call.lost)
call.observe(image)
}
return err
}
// TakePendingFindings returns the parked findings and clears them, so a
// replay that itself gets interrupted cannot double-dispatch on the next start.
func (s *Store) TakePendingFindings() ([]alert.Finding, error) {
call := s.pendingHealth().begin(0)
defer call.finish()
findings, err := s.takePendingFindings(call)
if err == nil {
call.finishReplay()
}
return findings, err
}
// ReplayPendingFindings keeps the cleared batch owned through dispatch. The
// callback runs without the state lock; later shutdown appends stay independent.
func (s *Store) ReplayPendingFindings(consume func([]alert.Finding)) error {
call := s.pendingHealth().begin(0)
defer call.finish()
pending, err := s.takePendingFindings(call)
if err != nil {
return err
}
if len(pending) > 0 {
consume(pending)
}
call.finishReplay()
return nil
}
func (s *Store) takePendingFindings(call *pendingCall) ([]alert.Finding, error) {
s.mu.Lock()
defer s.mu.Unlock()
defer call.settleIO()
call.start()
pending, err := s.readPendingLocked(call)
if err != nil {
return nil, err
}
old := pendingIdentity(pending)
call.observe(old)
call.offer(0)
if err := removePendingFindingsFile(filepath.Join(s.path, pendingFindingsFile)); err != nil {
if os.IsNotExist(err) && len(pending) == 0 {
call.complete(pendingImage{valid: true}, 0, false, nil)
return nil, nil
}
call.ioFailed()
actual, readErr := s.readPendingLocked(call)
if readErr == nil {
image := pendingIdentity(actual)
switch {
case image.valid && image == old:
call.complete(image, 0, true, nil)
case len(actual) == 0:
call.complete(image, len(pending), true, nil)
default:
call.unreadable(0)
call.observe(image)
}
}
return nil, fmt.Errorf("clear pending findings: %w", err)
}
call.detach(len(pending))
return pending, nil
}
func (s *Store) readPendingLocked(call *pendingCall) ([]alert.Finding, error) {
read := s.readPendingFile
if read == nil {
read = os.ReadFile
}
data, err := read(filepath.Join(s.path, pendingFindingsFile))
if err != nil {
if os.IsNotExist(err) {
return nil, nil
}
call.failedRead()
return nil, fmt.Errorf("read pending findings: %w", err)
}
var pending []alert.Finding
if err := json.Unmarshal(data, &pending); err != nil {
call.failedRead()
return nil, fmt.Errorf("decode pending findings: %w", err)
}
return pending, nil
}
package state
import (
"crypto/sha256"
"encoding/json"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/queuehealth"
)
type pendingImage struct {
count int
digest [sha256.Size]byte
valid bool
}
// Compare complete serialized batches, including duplicate occurrences. A
// finding's dedup key does not identify its persisted payload.
func pendingIdentity(findings []alert.Finding) pendingImage {
if len(findings) == 0 {
return pendingImage{valid: true}
}
data, err := json.Marshal(findings)
if err != nil {
return pendingImage{count: len(findings)}
}
// JSON repairs invalid UTF-8, including log text cut inside a character.
// Hash the recovered payload so a successful readback has the same identity.
var persisted []alert.Finding
if decodeErr := json.Unmarshal(data, &persisted); decodeErr != nil {
return pendingImage{count: len(findings)}
}
data, err = json.Marshal(persisted)
return pendingImage{count: len(findings), digest: sha256.Sum256(data), valid: err == nil}
}
type pendingQueue struct {
mu sync.Mutex
disk pendingImage
known bool
arrivals []time.Time
calls map[*pendingCall]struct{}
losses *queuehealth.Tracker
stateIO bool
lowerBound bool
uncertainAt time.Time
}
type pendingCall struct {
queue *pendingQueue
queued, progress time.Time
incoming, replay, lost int
offered, settled bool
}
func (s *Store) pendingHealth() *pendingQueue {
s.pendingHealthOnce.Do(func() {
s.pendingQueue = &pendingQueue{calls: make(map[*pendingCall]struct{}), losses: queuehealth.New(pendingFindingsMax, time.Minute)}
})
return s.pendingQueue
}
func (s *Store) observePendingQueue() {
pending, err := s.readPendingLocked(nil)
image := pendingIdentity(pending)
q := s.pendingHealth()
q.mu.Lock()
defer q.mu.Unlock()
if err != nil {
q.unknown(true)
return
}
q.setDisk(image)
}
func (q *pendingQueue) begin(count int) *pendingCall {
c := &pendingCall{queue: q, queued: time.Now(), incoming: count}
q.mu.Lock()
q.calls[c] = struct{}{}
q.mu.Unlock()
return c
}
func (c *pendingCall) start() {
c.queue.mu.Lock()
c.progress = time.Now()
c.queue.mu.Unlock()
}
func (q *pendingQueue) setDisk(next pendingImage) {
if !q.known || !next.valid || q.disk != next {
q.arrivals = make([]time.Time, next.count)
now := time.Now()
for i := range q.arrivals {
q.arrivals[i] = now
}
}
q.disk = next
q.known = true
}
func (c *pendingCall) appendedAges(oldCount, kept int) []time.Time {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
ages := make([]time.Time, kept)
start := oldCount + c.incoming - kept
for i := range ages {
if index := start + i; index < oldCount {
ages[i] = q.arrivals[index]
} else {
ages[i] = c.queued
}
}
return ages
}
func (q *pendingQueue) unknown(ioFailure bool) {
q.known = false
q.lowerBound = true
q.stateIO = ioFailure
q.uncertainAt = time.Now()
}
func (c *pendingCall) observe(image pendingImage) {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
if q.known && q.disk != image {
q.lowerBound = true
q.uncertainAt = time.Now()
}
q.setDisk(image)
c.progress = time.Now()
}
func (c *pendingCall) lose(count int) {
if count > c.lost {
c.queue.losses.Lose(time.Now(), uint64(count-c.lost)) // #nosec G115 -- the guarded positive difference is bounded by the finding batch lengths.
c.lost = count
}
}
// Incoming findings outside the retained suffix cannot survive either the old
// or the replacement file. Publish that known loss before attempting I/O.
func (c *pendingCall) offer(knownLoss int) {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
c.lose(knownLoss)
c.offered = true
c.progress = time.Now()
}
func (c *pendingCall) ioFailed() {
c.queue.mu.Lock()
c.queue.stateIO = true
c.progress = time.Now()
c.queue.mu.Unlock()
}
func (c *pendingCall) complete(image pendingImage, loss int, failed bool, arrivals []time.Time) {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
q.setDisk(image)
if arrivals != nil {
q.arrivals = arrivals
}
q.stateIO = failed
c.lose(loss)
c.incoming = 0
c.settled = true
c.progress = time.Now()
}
func (c *pendingCall) unreadable(loss int) {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
q.unknown(true)
c.lose(loss)
c.incoming = 0
c.settled = true
c.progress = time.Now()
}
func (c *pendingCall) failedRead() {
if c == nil {
return
}
loss := c.lost
if !c.offered {
loss = c.incoming
}
c.unreadable(loss)
}
// The state lock still belongs to this call. Settle an unreturned write or
// clear before a newer operation can publish its own confirmed disk state.
func (c *pendingCall) settleIO() {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
if c.offered && !c.settled {
q.unknown(false)
c.incoming = 0
c.settled = true
}
}
func (c *pendingCall) detach(count int) {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
q.setDisk(pendingImage{valid: true})
q.stateIO = false
c.offered = false
c.replay = count
c.progress = time.Now()
}
func (c *pendingCall) finishReplay() {
c.queue.mu.Lock()
c.settled = true
c.queue.mu.Unlock()
}
func (c *pendingCall) finish() {
q := c.queue
q.mu.Lock()
defer q.mu.Unlock()
if !c.settled {
switch {
case c.replay > 0:
// Dispatch may have partially completed. Do not invent an exact loss or
// requeue findings that the existing at-most-once policy already cleared.
q.lowerBound = true
q.uncertainAt = time.Now()
default:
c.lose(c.incoming)
}
}
delete(q.calls, c)
}
// QueueStatuses uses metadata only, independently of state locks, file I/O and
// the replay callback. Waiting for the next restart has no processing deadline.
func (s *Store) QueueStatuses(now time.Time) map[string]queuehealth.Status {
q := s.pendingHealth()
q.mu.Lock()
defer q.mu.Unlock()
row := q.losses.Snapshot(now)
row.DepthUnit = "findings"
row.LagBasis = "deferred_checkpoint"
row.DepthUnavailable = !q.known
if q.known {
row.Depth = q.disk.count
}
row.DroppedLowerBound = q.lowerBound
if row.Depth > 0 {
row.LagSeconds = max(0, now.Sub(q.arrivals[0]).Seconds())
}
op := queuehealth.Status{Status: "ok", CapacityUnavailable: true, DepthUnit: "operations", LagBasis: "operation_progress"}
for c := range q.calls {
row.InFlight += c.incoming + c.replay
if c.progress.IsZero() {
op.Depth++
op.LagSeconds = max(op.LagSeconds, now.Sub(c.queued).Seconds())
} else {
op.InFlight++
op.ProcessingSeconds = max(op.ProcessingSeconds, now.Sub(c.progress).Seconds())
}
}
switch {
case q.stateIO:
row.Status, row.Reason = "degraded", "state_io"
case !q.known:
row.Status, row.Reason = "degraded", "persistence_uncertain"
case !q.uncertainAt.IsZero() && now.Sub(q.uncertainAt) < time.Minute:
row.Status, row.Reason = "degraded", "persistence_uncertain"
}
switch {
case op.LagSeconds >= time.Minute.Seconds():
op.Status, op.Reason = "degraded", "backlog_lag"
case op.ProcessingSeconds >= time.Minute.Seconds():
op.Status, op.Reason = "degraded", "processing_lag"
}
return map[string]queuehealth.Status{"pending": row, "pending_operations": op}
}
package state
import "github.com/pidginhost/csm/internal/alert"
// ScanCoverage describes only units proved by a completed scan. Scope keys
// are scanner-issued identities, never file paths or parsed display text.
type ScanCoverage struct {
PreservePaths map[string]map[string]bool
CompletedScopes map[string]map[string]bool
// IncompleteChecks preserves unexamined findings at the active-set cap,
// including when the scanner could not complete any database scope.
IncompleteChecks map[string]bool
}
func (c *ScanCoverage) completed(f alert.Finding) bool {
return c != nil && f.CoverageScope != "" && c.CompletedScopes[f.Check][f.CoverageScope]
}
// unexamined identifies retained state that must not be evicted to make room
// for new findings from a completed scope in the same scanner.
func (c *ScanCoverage) unexamined(f alert.Finding) bool {
return c != nil && (c.IncompleteChecks[f.Check] || c.CompletedScopes[f.Check] != nil) && !c.completed(f)
}
package state
import (
"fmt"
"os"
"strings"
"github.com/pidginhost/csm/internal/alert"
)
// RearmAbsentDedupFindings forgets dismissals of resolved conditions, only for
// the finding names whose owner supplied a replacement scan result. Update
// cannot do this: its input may be an unrelated tier or realtime batch. An
// alerted, undismissed entry is kept: deep scans do not cover every file each
// cycle, so a condition that is absent for one run and back the next keeps its
// daily reminder instead of alerting on every return.
func (s *Store) RearmAbsentDedupFindings(checks []string, findings []alert.Finding) {
if len(checks) == 0 {
return
}
owners := make(map[string]bool, len(checks))
for _, check := range checks {
owners[check] = true
}
seen := make(map[string]bool, len(findings))
for _, f := range findings {
seen[f.Key()] = true
}
s.mu.Lock()
defer s.mu.Unlock()
changed := false
for key := range s.entries {
identity, pinned := strings.CutPrefix(key, "dedup:")
check, _, _ := strings.Cut(identity, ":")
if pinned && owners[check] && !seen[key] && s.entries[key].IsBaseline {
delete(s.entries, key)
changed = true
}
}
if changed {
s.dirty = true
if err := s.save(); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving resolved scan conditions: %v\n", err)
}
}
}
// RearmFindings resets specific conditions after their producer proves recovery.
// Other producers sharing a finding name keep their own alert history.
func (s *Store) RearmFindings(keys []string) {
s.mu.Lock()
defer s.mu.Unlock()
changed := false
for _, key := range keys {
if s.deleteRawLocked(key) {
changed = true
}
}
if changed {
if err := s.save(); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving recovered scan conditions: %v\n", err)
}
}
}
// RearmDismissedFindings re-arms specific conditions only where the operator
// dismissed them. An alerted, undismissed condition keeps its daily reminder.
func (s *Store) RearmDismissedFindings(keys []string) {
s.mu.Lock()
defer s.mu.Unlock()
changed := false
for _, key := range keys {
if entry, ok := s.entries[key]; ok && entry.IsBaseline {
delete(s.entries, key)
changed = true
}
}
if changed {
s.dirty = true
if err := s.save(); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving resolved scan conditions: %v\n", err)
}
}
}
package state
import (
"bytes"
"crypto/rand"
"crypto/sha256"
"encoding/json"
"fmt"
"os"
"path/filepath"
"sort"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
"github.com/pidginhost/csm/internal/store"
)
// findingsTotal counts every finding CSM records, partitioned by
// severity. Registered lazily (via sync.Once) the first time a finding
// lands so tests that open Stores without running a full daemon do not
// panic on duplicate registration.
var (
findingsTotal *metrics.CounterVec
findingsTotalOnce sync.Once
)
func ensureFindingsMetric() {
findingsTotalOnce.Do(func() {
findingsTotal = metrics.NewCounterVec(
"csm_findings_total",
"Findings recorded by CSM, partitioned by severity.",
[]string{"severity"},
)
metrics.MustRegister("csm_findings_total", findingsTotal)
})
}
func recordFindings(findings []alert.Finding) {
if len(findings) == 0 {
return
}
ensureFindingsMetric()
for _, f := range findings {
findingsTotal.With(f.Severity.String()).Inc()
}
}
type Store struct {
mu sync.RWMutex
path string
entries map[string]*Entry
dirty bool // true if state changed since last save
savedHash string // hash of last saved state
throttleReservations map[string]struct{}
pendingHealthOnce sync.Once
pendingQueue *pendingQueue
readPendingFile func(string) ([]byte, error)
writePendingFile func(string, os.FileMode, any) error
// LatestFindings holds the full output of the most recent scan cycle.
// This is what the Findings page shows - "what's wrong right now" -
// separate from the alert dedup state above which controls "what to email."
latestMu sync.RWMutex
latestFindings []alert.Finding
latestDigest [sha256.Size]byte // digest of the last persisted latest_findings.json
latestDigestSet bool
latestScanTime time.Time
// suppressMu makes each change to the suppression rules one
// read-modify-write, so concurrent edits do not overwrite each other.
suppressMu sync.Mutex
}
type Entry struct {
Hash string `json:"hash"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
AlertSent time.Time `json:"alert_sent"`
IsBaseline bool `json:"is_baseline"`
DismissalID string `json:"dismissal_id,omitempty"`
}
func Open(path string) (*Store, error) {
if err := os.MkdirAll(path, 0700); err != nil {
return nil, fmt.Errorf("creating state dir: %w", err)
}
s := &Store{
path: path,
entries: make(map[string]*Entry),
throttleReservations: make(map[string]struct{}),
}
stateFile := filepath.Join(path, "state.json")
loadJSONWithBackup(stateFile, &s.entries)
if s.entries == nil {
s.entries = make(map[string]*Entry)
}
// Load latest findings from disk (survives restart)
latestFile := filepath.Join(path, "latest_findings.json")
var findings []alert.Finding
if loadJSONWithBackup(latestFile, &findings) {
s.latestFindings = findings
}
s.observePendingQueue()
return s, nil
}
// loadJSONWithBackup parses file into v. The backup copy is refreshed only
// after a successful parse, so a corrupt file can never overwrite the last
// good one; on a parse failure the backup is parsed instead. Copying the
// file over its backup before parsing, as this used to, destroyed the
// backup exactly when it was needed and silently reset the alert dedup
// state, re-alerting every known finding. Returns true when v was loaded
// from either source.
func loadJSONWithBackup(file string, v any) bool {
bakFile := file + ".bak"
// #nosec G304 -- operator-configured statePath + fixed filename.
data, err := os.ReadFile(file)
if err == nil {
if unmarshalErr := json.Unmarshal(data, v); unmarshalErr == nil {
// Refresh atomically so a crash while copying a good primary cannot
// destroy the last usable backup.
_ = atomicio.AtomicWrite(bakFile, 0o600, data)
return true
} else {
fmt.Fprintf(os.Stderr, "warning: failed to parse %s: %v (trying %s)\n", file, unmarshalErr, bakFile)
}
}
// #nosec G304 -- same fixed name with a .bak suffix.
bak, err := os.ReadFile(bakFile)
if err != nil {
return false
}
if unmarshalErr := json.Unmarshal(bak, v); unmarshalErr != nil {
fmt.Fprintf(os.Stderr, "warning: failed to parse %s: %v\n", bakFile, unmarshalErr)
return false
}
fmt.Fprintf(os.Stderr, "warning: restored %s from %s\n", file, bakFile)
return true
}
func (s *Store) Close() error {
s.mu.Lock()
defer s.mu.Unlock()
if !s.dirty {
return nil
}
return s.save()
}
func (s *Store) save() error {
data, err := json.MarshalIndent(s.entries, "", " ")
if err != nil {
return err
}
// Skip write if content hasn't changed
newHash := fmt.Sprintf("%x", sha256.Sum256(data))
if newHash == s.savedHash {
s.dirty = false
return nil
}
// Atomic write with fsync. Power loss between rename and a dir
// fsync can otherwise leave the file truncated; the savedHash gate
// above then "skips" the next attempt because the in-memory hash is
// unchanged, so corruption persists across restarts.
stateFile := filepath.Join(s.path, "state.json")
if err := atomicio.AtomicWriteJSON(stateFile, 0o600, s.entries); err != nil {
return err
}
s.savedHash = newHash
s.dirty = false
return nil
}
func findingKey(f alert.Finding) string {
return f.Key()
}
func findingHash(f alert.Finding) string {
return f.Fingerprint()
}
func (s *Store) FilterNew(findings []alert.Finding) []alert.Finding {
s.mu.RLock()
defer s.mu.RUnlock()
var newFindings []alert.Finding
for _, f := range findings {
key := findingKey(f)
hash := findingHash(f)
entry, exists := s.entries[key]
if !exists {
newFindings = append(newFindings, f)
continue
}
if entry.IsBaseline && entry.Hash == hash {
continue
}
if entry.Hash != hash {
newFindings = append(newFindings, f)
continue
}
// Same finding, check if we should re-alert (state expiry)
if !entry.AlertSent.IsZero() && time.Since(entry.AlertSent) > 24*time.Hour {
newFindings = append(newFindings, f)
}
}
return newFindings
}
func (s *Store) Update(findings []alert.Finding) {
s.mu.Lock()
defer s.mu.Unlock()
s.dirty = true
now := time.Now()
seen := make(map[string]bool)
for _, f := range findings {
key := findingKey(f)
hash := findingHash(f)
seen[key] = true
entry, exists := s.entries[key]
if !exists {
s.entries[key] = &Entry{
Hash: hash,
FirstSeen: now,
LastSeen: now,
AlertSent: now,
}
} else {
entry.Hash = hash
entry.LastSeen = now
if entry.AlertSent.IsZero() {
entry.AlertSent = now
}
}
}
// Clean up entries that are no longer found. Keys with a leading
// underscore are internal housekeeping written via SetRaw (throttles,
// per-file content-hash baselines, sentinel flags) and never appear in
// the findings stream, so the !seen branch would always evict them and
// silently re-arm one-shot detectors that gate on their presence.
for key, entry := range s.entries {
if strings.HasPrefix(key, "_") {
continue
}
if !seen[key] && !entry.IsBaseline {
if time.Since(entry.LastSeen) > 24*time.Hour {
delete(s.entries, key)
}
}
}
if err := s.save(); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving after update: %v\n", err)
}
}
// MarkAlerted refreshes AlertSent on each finding's entry so the 24-hour
// dedup window restarts. Call after dispatch with the slice that came back
// from FilterNew. Without this, any finding that survives past the 24-hour
// expiry branch in FilterNew re-emits on every subsequent tick because
// Update only sets AlertSent when an entry is first created.
func (s *Store) MarkAlerted(findings []alert.Finding) {
if len(findings) == 0 {
return
}
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
changed := false
for _, f := range findings {
key := findingKey(f)
if entry, ok := s.entries[key]; ok {
entry.AlertSent = now
changed = true
}
}
if !changed {
return
}
s.dirty = true
if err := s.save(); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving after MarkAlerted: %v\n", err)
}
}
func (s *Store) SetBaseline(findings []alert.Finding) {
s.mu.Lock()
defer s.mu.Unlock()
s.dirty = true
var baselineAt *Entry
if entry, ok := s.entries[baselineAtMetaKey]; ok {
entryCopy := *entry
baselineAt = &entryCopy
}
s.entries = make(map[string]*Entry)
now := time.Now()
for _, f := range findings {
key := findingKey(f)
hash := findingHash(f)
s.entries[key] = &Entry{
Hash: hash,
FirstSeen: now,
LastSeen: now,
IsBaseline: true,
}
}
if baselineAt != nil {
s.entries[baselineAtMetaKey] = baselineAt
}
if err := s.save(); err != nil {
fmt.Fprintf(os.Stderr, "state: error saving baseline: %v\n", err)
}
}
// ShouldRunThrottled reports whether the throttle window for checkName has
// elapsed, consuming the slot when it allows. Use only when the throttled
// work cannot fail or time out after scheduling; otherwise pair the
// read-only ThrottleAllows probe with MarkThrottledRan on completion so a
// failed run does not forfeit its slot.
func (s *Store) ShouldRunThrottled(checkName string, intervalMin int) bool {
s.mu.Lock()
defer s.mu.Unlock()
key := throttleKey(checkName)
entry, exists := s.entries[key]
if !exists {
s.entries[key] = &Entry{LastSeen: time.Now()}
s.dirty = true
return true
}
if time.Since(entry.LastSeen) >= time.Duration(intervalMin)*time.Minute {
entry.LastSeen = time.Now()
s.dirty = true
return true
}
return false
}
// ThrottleAllows reports whether the throttle window for checkName has
// elapsed WITHOUT consuming the slot. Callers stamp the slot with
// MarkThrottledRan only after the work completes, so a check that times
// out keeps its slot and may retry on the next cycle instead of waiting
// out the full window with nothing stored.
func (s *Store) ThrottleAllows(checkName string, intervalMin int) bool {
s.mu.RLock()
defer s.mu.RUnlock()
entry, exists := s.entries[throttleKey(checkName)]
if !exists {
return true
}
return time.Since(entry.LastSeen) >= time.Duration(intervalMin)*time.Minute
}
// ReserveThrottle reports whether the throttle window allows checkName to
// start now and reserves the slot for this process until the caller either
// marks the run complete or releases the reservation. The reservation keeps
// concurrent scans from launching the same expensive check before the first
// one has had a chance to stamp its successful run.
func (s *Store) ReserveThrottle(checkName string, intervalMin int) bool {
s.mu.Lock()
defer s.mu.Unlock()
key := throttleKey(checkName)
s.ensureThrottleReservationsLocked()
if _, running := s.throttleReservations[key]; running {
return false
}
entry, exists := s.entries[key]
if exists && time.Since(entry.LastSeen) < time.Duration(intervalMin)*time.Minute {
return false
}
s.throttleReservations[key] = struct{}{}
return true
}
// MarkThrottledRan stamps the throttle slot for checkName. Call only after
// the throttled work actually completed, never at scheduling time.
func (s *Store) MarkThrottledRan(checkName string) {
s.mu.Lock()
defer s.mu.Unlock()
key := throttleKey(checkName)
if s.throttleReservations != nil {
delete(s.throttleReservations, key)
}
if entry, exists := s.entries[key]; exists {
entry.LastSeen = time.Now()
} else {
s.entries[key] = &Entry{LastSeen: time.Now()}
}
s.dirty = true
}
// ReleaseThrottle clears an in-process throttle reservation without stamping
// a successful run. Use when the reserved work times out or is interrupted.
func (s *Store) ReleaseThrottle(checkName string) {
s.mu.Lock()
defer s.mu.Unlock()
if s.throttleReservations != nil {
delete(s.throttleReservations, throttleKey(checkName))
}
}
func (s *Store) ensureThrottleReservationsLocked() {
if s.throttleReservations == nil {
s.throttleReservations = make(map[string]struct{})
}
}
func throttleKey(checkName string) string {
return "_throttle:" + checkName
}
func (s *Store) GetRaw(key string) (string, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
entry, ok := s.entries[key]
if !ok {
return "", false
}
return entry.Hash, true
}
func (s *Store) SetRaw(key, value string) {
s.mu.Lock()
defer s.mu.Unlock()
s.setRawLocked(key, value)
}
// DeleteRaw drops a raw housekeeping entry. Underscore-prefixed keys are exempt
// from the sweeper, so a caller that no longer wants one has to say so.
func (s *Store) DeleteRaw(key string) {
s.mu.Lock()
defer s.mu.Unlock()
s.deleteRawLocked(key)
}
// DeleteRawAndSave drops a raw housekeeping entry and persists the deletion
// before returning.
func (s *Store) DeleteRawAndSave(key string) error {
s.mu.Lock()
defer s.mu.Unlock()
changed := s.deleteRawLocked(key)
if !changed && !s.dirty {
return nil
}
return s.save()
}
func (s *Store) deleteRawLocked(key string) bool {
if _, ok := s.entries[key]; !ok {
return false
}
delete(s.entries, key)
s.dirty = true
return true
}
// SetRawAndSave stores a raw housekeeping value and immediately persists the
// state file. Use it for cursors where a completed scan must survive restart
// even when no finding is emitted later in the cycle.
func (s *Store) SetRawAndSave(key, value string) error {
s.mu.Lock()
defer s.mu.Unlock()
changed := s.setRawLocked(key, value)
if !changed && !s.dirty {
return nil
}
return s.save()
}
// ClaimRawTimestamp atomically claims a persisted time window. It returns
// true for the first claim and after interval has elapsed. A future stored
// timestamp can result from a backward wall-clock step; rebase it to now so
// the window does not stay suppressed until the old clock value catches up.
func (s *Store) ClaimRawTimestamp(key string, now time.Time, interval time.Duration) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
now = now.Round(0)
if entry, ok := s.entries[key]; ok {
if previous, err := time.Parse(time.RFC3339Nano, entry.Hash); err == nil {
previous = previous.Round(0)
if previous.After(now) {
entry.Hash = now.Format(time.RFC3339Nano)
entry.LastSeen = now
s.dirty = true
return false, s.save()
}
if now.Sub(previous) < interval {
return false, nil
}
}
}
s.setRawLocked(key, now.Format(time.RFC3339Nano))
return true, s.save()
}
func (s *Store) setRawLocked(key, value string) bool {
if s.entries == nil {
// A zero-value Store (tests, ad-hoc callers) must not panic on its
// first write; it simply has nothing persisted yet.
s.entries = make(map[string]*Entry)
}
entry, exists := s.entries[key]
if !exists {
s.entries[key] = &Entry{
Hash: value,
FirstSeen: time.Now(),
LastSeen: time.Now(),
}
s.dirty = true
return true
}
if entry.Hash != value {
entry.Hash = value
entry.LastSeen = time.Now()
s.dirty = true
return true
}
return false
}
// AppendHistory writes findings to the bbolt store (if available) or
// falls back to the append-only JSONL history file.
// The JSONL fallback is deprecated and will be removed in a future release.
func (s *Store) AppendHistory(findings []alert.Finding) {
if len(findings) == 0 {
return
}
recordFindings(findings)
// Use bbolt store when available; skip JSONL writes entirely.
if db := store.Global(); db != nil {
if err := db.AppendHistory(findings); err != nil {
fmt.Fprintf(os.Stderr, "store: append history: %v\n", err)
}
return
}
// Deprecated: flat-file JSONL fallback has a truncation race condition.
// This path is kept only for installations that have not yet migrated to bbolt.
fmt.Fprintf(os.Stderr, "DEPRECATION: using JSONL history fallback; migrate to bbolt store\n")
s.appendHistoryFile(findings)
}
// appendHistoryFile writes findings to the append-only JSONL history file.
// Caps file at 10MB by truncating the oldest half.
func (s *Store) appendHistoryFile(findings []alert.Finding) {
histPath := filepath.Join(s.path, "history.jsonl")
// Check size, truncate if over 10MB
if info, err := os.Stat(histPath); err == nil && info.Size() > 10*1024*1024 {
// #nosec G304 -- histPath is {s.path}/history.jsonl; s.path is the
// operator-configured statePath set at Store creation.
data, err := os.ReadFile(histPath)
if err == nil {
// Keep the second half
half := len(data) / 2
for half < len(data) && data[half] != '\n' {
half++
}
if half < len(data) {
// #nosec G703 -- histPath comes from state.historyPath, a
// filepath.Join under the operator-configured statePath.
_ = os.WriteFile(histPath, data[half+1:], 0600)
}
}
}
// #nosec G304 -- see histPath derivation above.
f, err := os.OpenFile(histPath, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600)
if err != nil {
return
}
defer func() { _ = f.Close() }()
for _, finding := range findings {
line, err := json.Marshal(alert.SanitizeFinding(finding))
if err != nil {
continue
}
_, _ = f.Write(line)
_, _ = f.Write([]byte("\n"))
}
}
func (s *Store) PrintStatus() {
if len(s.entries) == 0 {
fmt.Println("No state entries. Run 'csm baseline' first.")
return
}
baselineCount := 0
activeCount := 0
for key, entry := range s.entries {
if entry.IsBaseline {
baselineCount++
continue
}
if key[0] == '_' {
continue
}
activeCount++
fmt.Printf(" [ACTIVE] %s (first: %s, last: %s)\n",
key,
entry.FirstSeen.Format("2006-01-02 15:04"),
entry.LastSeen.Format("2006-01-02 15:04"),
)
}
fmt.Printf("\nBaseline entries: %d, Active findings: %d\n", baselineCount, activeCount)
}
// Entries returns a snapshot copy of current state entries (thread-safe).
func (s *Store) Entries() map[string]*Entry {
s.mu.RLock()
defer s.mu.RUnlock()
copy := make(map[string]*Entry, len(s.entries))
for k, v := range s.entries {
if k[0] == '_' {
continue // skip internal keys
}
entryCopy := *v
copy[k] = &entryCopy
}
return copy
}
// EntryForKey returns the dedup entry for a finding key (check:message).
func (s *Store) EntryForKey(key string) (Entry, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
e, ok := s.entries[key]
if !ok {
return Entry{}, false
}
return *e, true
}
// ReadHistory reads the last limit entries, starting at offset.
// Returns the findings (newest first) and total count.
// Uses bbolt store when available, falls back to flat-file JSONL.
func (s *Store) ReadHistory(limit, offset int) ([]alert.Finding, int) {
// Use bbolt store when available.
if db := store.Global(); db != nil {
return db.ReadHistory(limit, offset)
}
// Fallback: flat-file JSONL.
historyPath := filepath.Join(s.path, "history.jsonl")
// #nosec G304 -- {s.path}/history.jsonl; s.path from operator config.
data, err := os.ReadFile(historyPath)
if err != nil {
return nil, 0
}
var all []alert.Finding
for _, line := range splitLines(data) {
if len(line) == 0 {
continue
}
var f alert.Finding
if err := json.Unmarshal(line, &f); err != nil {
continue
}
all = append(all, f)
}
// Reverse (newest first)
for i, j := 0, len(all)-1; i < j; i, j = i+1, j-1 {
all[i], all[j] = all[j], all[i]
}
total := len(all)
if offset >= total {
return nil, total
}
end := offset + limit
if end > total {
end = total
}
return all[offset:end], total
}
// ReadHistoryFiltered reads matching history entries newest-first.
// Bounds accept server-local calendar days or RFC 3339 instants. Invalid
// bounds are ignored before they reach the lexicographic bbolt key filter.
func (s *Store) ReadHistoryFiltered(limit, offset int, from, to string, severity int, search string) ([]alert.Finding, int) {
return s.ReadHistoryFilteredWithChecks(limit, offset, from, to, severity, search, nil)
}
// ReadHistoryFilteredWithChecks reads matching history entries newest-first,
// optionally constrained to an exact check-name set.
func (s *Store) ReadHistoryFilteredWithChecks(
limit, offset int,
from, to string,
severity int,
search string,
checks map[string]bool,
) ([]alert.Finding, int) {
// An unreadable bound is dropped; the web UI rejects one before it gets
// here. toEnd is exclusive.
fromDate, fromErr := store.ParseHistoryBound(from, false)
toEnd, toErr := store.ParseHistoryBound(to, true)
fromFilter := ""
toFilter := ""
if fromErr == nil {
fromFilter = from
}
if toErr == nil {
toFilter = to
}
if db := store.Global(); db != nil {
return db.ReadHistoryFilteredWithChecks(limit, offset, fromFilter, toFilter, severity, search, checks)
}
all, _ := s.ReadHistory(1<<30, 0)
searchLower := strings.ToLower(search)
var results []alert.Finding
matched := 0
for _, f := range all {
if !fromDate.IsZero() && f.Timestamp.Before(fromDate) {
continue
}
if !toEnd.IsZero() && !f.Timestamp.Before(toEnd) {
continue
}
if severity >= 0 && int(f.Severity) != severity {
continue
}
if checks != nil && !checks[f.Check] {
continue
}
if search != "" {
if !strings.Contains(strings.ToLower(f.Check), searchLower) &&
!strings.Contains(strings.ToLower(f.Message), searchLower) &&
!strings.Contains(strings.ToLower(f.Details), searchLower) {
continue
}
}
matched++
if matched > offset && len(results) < limit {
results = append(results, f)
}
}
return results, matched
}
// ReadHistorySince returns all findings since the given time.
// Uses bbolt cursor seeking for efficiency. Results are newest-first.
// Falls back to the JSONL store with a linear cutoff filter when bbolt
// is unavailable so test wiring and migration-pending hosts still get
// time-bounded search results.
func (s *Store) ReadHistorySince(since time.Time) []alert.Finding {
if db := store.Global(); db != nil {
return db.ReadHistorySince(since)
}
all, _ := s.ReadHistory(1<<30, 0)
out := all[:0]
for _, f := range all {
if !f.Timestamp.Before(since) {
out = append(out, f)
}
}
return out
}
// HistoryMark changes whenever history does. The JSONL fallback uses the
// history file's size and modification time.
func (s *Store) HistoryMark() string {
if db := store.Global(); db != nil {
return db.HistoryMark()
}
info, err := os.Stat(filepath.Join(s.path, "history.jsonl"))
if err != nil {
return ""
}
return fmt.Sprintf("%d|%d", info.Size(), info.ModTime().UnixNano())
}
// LatestByCheck returns the timestamp of the newest history entry of every
// check. The bbolt store keeps an index; the JSONL fallback reads history.
func (s *Store) LatestByCheck() map[string]time.Time {
if db := store.Global(); db != nil {
return db.LatestByCheck()
}
out := map[string]time.Time{}
all, _ := s.ReadHistory(1<<30, 0)
for _, f := range all {
if f.Check != "" && f.Timestamp.After(out[f.Check]) {
out[f.Check] = f.Timestamp
}
}
return out
}
// SearchHistorySince returns up to limit matching findings since the given
// time, newest-first.
func (s *Store) SearchHistorySince(since time.Time, limit int, match func(alert.Finding) bool) []alert.Finding {
if limit <= 0 {
return nil
}
if db := store.Global(); db != nil {
return db.SearchHistorySince(since, limit, match)
}
return s.searchHistoryFileSince(since, limit, match)
}
func (s *Store) searchHistoryFileSince(since time.Time, limit int, match func(alert.Finding) bool) []alert.Finding {
historyPath := filepath.Join(s.path, "history.jsonl")
// #nosec G304 -- {s.path}/history.jsonl; s.path from operator config.
data, err := os.ReadFile(historyPath)
if err != nil {
return nil
}
var results []alert.Finding
end := len(data)
for end > 0 && len(results) < limit {
for end > 0 && (data[end-1] == '\n' || data[end-1] == '\r') {
end--
}
if end == 0 {
break
}
start := bytes.LastIndexByte(data[:end], '\n') + 1
line := data[start:end]
if start == 0 {
end = 0
} else {
end = start - 1
}
var f alert.Finding
if err := json.Unmarshal(line, &f); err != nil {
continue
}
// The JSONL fallback appends findings in chronological order, so
// once a reverse scan reaches an old row the remaining rows are older.
if f.Timestamp.Before(since) {
break
}
if match != nil && !match(f) {
continue
}
results = append(results, f)
}
return results
}
// AggregateByHour returns 24 hourly severity buckets for the last 24 hours.
func (s *Store) AggregateByHour() []store.HourBucket {
if db := store.Global(); db != nil {
return db.AggregateByHour()
}
return nil
}
// AggregateByDay returns 30 daily severity buckets for the last 30 days.
func (s *Store) AggregateByDay() []store.DayBucket {
if db := store.Global(); db != nil {
return db.AggregateByDay()
}
return nil
}
// AggregateByDayN returns `days` daily severity buckets (oldest first),
// clamped by the underlying store's retention window.
func (s *Store) AggregateByDayN(days int) []store.DayBucket {
if db := store.Global(); db != nil {
return db.AggregateByDayN(days)
}
return nil
}
func splitLines(data []byte) [][]byte {
var lines [][]byte
start := 0
for i, b := range data {
if b == '\n' {
if i > start {
lines = append(lines, data[start:i])
}
start = i + 1
}
}
if start < len(data) {
lines = append(lines, data[start:])
}
return lines
}
// SetLatestFindings merges scan results into the current findings set.
// Called by the daemon after each periodic scan completes. Merges rather
// than replaces - critical scan results coexist with deep scan results.
// Use ClearLatestFindings() + SetLatestFindings() for a full replace.
func (s *Store) SetLatestFindings(findings []alert.Finding) {
s.latestMu.Lock()
defer s.latestMu.Unlock()
// Build map of existing findings by key
existing := make(map[string]alert.Finding)
for _, f := range s.latestFindings {
existing[f.Key()] = f
}
// Merge new findings (update existing, add new)
for _, f := range findings {
existing[f.Key()] = f // newer overwrites older
}
s.latestFindings = orderAndCapLatest(existing)
s.latestScanTime = time.Now()
s.persistLatestLocked()
}
// PurgeFindingsByChecks removes all findings whose Check field matches
// any of the given check names and persists the result to disk.
// Used to clear stale performance findings before merging fresh results
// from a scan tier.
func (s *Store) PurgeFindingsByChecks(checks []string) {
if len(checks) == 0 {
return
}
s.latestMu.Lock()
defer s.latestMu.Unlock()
remove := make(map[string]bool, len(checks))
for _, c := range checks {
remove[c] = true
}
n := 0
for _, f := range s.latestFindings {
if !remove[f.Check] {
s.latestFindings[n] = f
n++
}
}
s.latestFindings = s.latestFindings[:n]
// Persist to disk so purged findings don't reappear after restart.
s.persistLatestLocked()
}
// PurgeAndMergeFindings atomically removes findings matching the given check
// names and then merges the new findings. This prevents a race window where
// concurrent readers could see findings with perf checks missing.
func (s *Store) PurgeAndMergeFindings(purgeChecks []string, findings []alert.Finding) {
s.PurgeAndMergeFindingsDerived(purgeChecks, findings, nil, nil)
}
// PurgeAndMergeFindingsDerived is PurgeAndMergeFindings followed, under the
// same lock and before the single persist, by a second purge-and-merge of
// derivedChecks with derive(merged): the correlation findings a tier cycle
// rebuilds from the merged set. One file write per cycle instead of two.
func (s *Store) PurgeAndMergeFindingsDerived(purgeChecks []string, findings []alert.Finding, derivedChecks []string, derive func([]alert.Finding) []alert.Finding) {
s.PurgeAndMergeFindingsDerivedWithCoverage(purgeChecks, findings, nil, derivedChecks, derive)
}
// PurgeAndMergeFindingsDerivedWithGaps preserves current findings for path
// aliases captured when the completed scan observed a coverage gap. Callers
// must supply both lexical and resolved aliases they accepted at scan time;
// preservation is evaluated under latestMu without resolving them again.
func (s *Store) PurgeAndMergeFindingsDerivedWithGaps(purgeChecks []string, findings []alert.Finding, preserveAliasesByCheck map[string]map[string]bool, derivedChecks []string, derive func([]alert.Finding) []alert.Finding) {
s.PurgeAndMergeFindingsDerivedWithCoverage(purgeChecks, findings, &ScanCoverage{PreservePaths: preserveAliasesByCheck}, derivedChecks, derive)
}
// PurgeAndMergeFindingsDerivedWithCoverage applies file gaps and completed
// scanner scopes against current state in the same transaction as correlation.
func (s *Store) PurgeAndMergeFindingsDerivedWithCoverage(purgeChecks []string, findings []alert.Finding, coverage *ScanCoverage, derivedChecks []string, derive func([]alert.Finding) []alert.Finding) {
s.latestMu.Lock()
defer s.latestMu.Unlock()
merged := purgeAndMergeLatestWithCoverage(s.latestFindings, purgeChecks, findings, coverage, true)
if derive != nil {
merged = purgeAndMergeLatestWithCoverage(merged, derivedChecks, derive(append([]alert.Finding(nil), merged...)), coverage, false)
}
s.latestFindings = merged
s.latestScanTime = time.Now()
s.persistLatestLocked()
}
// earliestObservation resolves the first-seen time a re-reported finding
// keeps. The stored row wins because it was there first; a row written before
// FirstSeen existed contributes its Timestamp instead, so upgrading does not
// reset the history of every long-lived finding. A finding nobody stored
// before starts from its own timestamp.
func earliestObservation(stored, reported alert.Finding) time.Time {
candidates := []time.Time{stored.FirstSeen, stored.Timestamp, reported.FirstSeen, reported.Timestamp}
var earliest time.Time
for _, t := range candidates {
if t.IsZero() {
continue
}
if earliest.IsZero() || t.Before(earliest) {
earliest = t
}
}
return earliest
}
// purgeAndMergeLatest drops findings owned by purgeChecks (and the timeout
// findings those runners produced), merges findings by key, and returns the
// ordered, capped result.
func purgeAndMergeLatest(current []alert.Finding, purgeChecks []string, findings []alert.Finding, preservePathsByCheck map[string]map[string]bool) []alert.Finding {
return purgeAndMergeLatestWithCoverage(current, purgeChecks, findings, &ScanCoverage{PreservePaths: preservePathsByCheck}, true)
}
func purgeAndMergeLatestWithCoverage(current []alert.Finding, purgeChecks []string, findings []alert.Finding, coverage *ScanCoverage, retireScopes bool) []alert.Finding {
if coverage == nil {
coverage = &ScanCoverage{}
}
preservePathsByCheck := coverage.PreservePaths
remove := make(map[string]bool, len(purgeChecks))
for _, c := range purgeChecks {
remove[c] = true
}
existing := make(map[string]alert.Finding, len(current)+len(findings))
// Owner replacement removes old rows from the output, but a key reported
// again in this merge must inherit its pre-purge observation. This history
// lasts only for this merge; resolved keys leave no tombstone behind.
observations := make(map[string]time.Time, len(current)+len(findings))
preserveAliases := normalizedPreservePathAliases(preservePathsByCheck)
preservationActive := preservePathsByCheck != nil
mismatchedCarryKeys := make(map[string]struct{})
holdChecks := make(map[string]bool)
if preservationActive {
for _, f := range findings {
if !f.ScanCarryForward {
continue
}
if pathMatchesPreservedAliases(f.FilePath, preserveAliases[f.Check]) {
continue
}
// Carry-forward and preservation metadata must describe the same
// scope. If they ever disagree, retain the whole owner rather than
// choosing between a purge and a stale resurrection.
holdChecks[f.Check] = true
mismatchedCarryKeys[f.Key()] = struct{}{}
}
}
protectedKeys := make(map[string]struct{})
scopeProtectedKeys := make(map[string]struct{})
for _, f := range current {
key := f.Key()
observations[key] = earliestObservation(alert.Finding{}, f)
// A finding this package demoted is waiting on the re-verifier, which
// reads the file, not on a scan that merely did not raise it again.
// Purging it here would discard the demotion state and let the same
// file come back at full severity on the next detection, so the
// demotion has to outlive a negative scan. A fresh finding normally
// replaces it below; a protected snapshot stays authoritative against a
// colliding finding from another path.
preserved := holdChecks[f.Check] || pathMatchesPreservedAliases(f.FilePath, preserveAliases[f.Check])
retire := shouldPurgeLatestFinding(f, remove) || (retireScopes && coverage.completed(f))
if preserved || isAutomaticallyDemotedFinding(f) || !retire {
existing[key] = f
if coverage.unexamined(f) {
scopeProtectedKeys[key] = struct{}{}
}
if preserved {
protectedKeys[key] = struct{}{}
}
}
}
for _, f := range findings {
// Marked carry-forward state came from the scanner's earlier snapshot.
// The current set under latestMu is authoritative: overwriting it would
// undo a concurrent update, and inserting it when absent would resurrect
// a concurrent dismissal. A genuinely fresh detection has no marker.
if preservationActive && f.ScanCarryForward {
continue
}
key := f.Key()
if _, staleSnapshot := mismatchedCarryKeys[key]; staleSnapshot {
continue
}
if _, currentIsAuthoritative := protectedKeys[key]; currentIsAuthoritative &&
!pathMatchesPreservedAliases(f.FilePath, preserveAliases[f.Check]) {
continue
}
f.ScanCarryForward = false
f.FirstSeen = earliestObservation(alert.Finding{FirstSeen: observations[key]}, f)
observations[key] = f.FirstSeen
existing[key] = f
if pathMatchesPreservedAliases(f.FilePath, preserveAliases[f.Check]) {
protectedKeys[key] = struct{}{}
}
}
// Protect only previously admitted unexamined findings. New partial
// detections compete for the remaining slots; protecting them too would
// let a persistently incomplete scanner grow the active set without bound.
// Refreshes of retained keys keep their protection and original first-seen.
for key := range scopeProtectedKeys {
protectedKeys[key] = struct{}{}
}
return orderAndCapLatestPreserving(existing, protectedKeys)
}
func normalizedPreservePathAliases(pathsByCheck map[string]map[string]bool) map[string]map[string]struct{} {
if len(pathsByCheck) == 0 {
return nil
}
aliasesByCheck := make(map[string]map[string]struct{}, len(pathsByCheck))
for check, paths := range pathsByCheck {
if len(paths) == 0 {
continue
}
aliases := make(map[string]struct{}, len(paths)*2)
for path := range paths {
// The scanner already captured both lexical and symlink-resolved
// identities at gap time. Re-resolving here would let a symlink
// retarget between scan and merge change which file is protected.
if alias := normalizedLatestFindingPath(path); alias != "" {
aliases[alias] = struct{}{}
}
}
aliasesByCheck[check] = aliases
}
return aliasesByCheck
}
func pathMatchesPreservedAliases(path string, preserved map[string]struct{}) bool {
if path == "" || len(preserved) == 0 {
return false
}
for _, alias := range latestFindingPathAliases(path) {
if _, ok := preserved[alias]; ok {
return true
}
}
return false
}
func latestFindingPathAliases(path string) []string {
if path == "" {
return nil
}
lexical := filepath.Clean(path)
if absolute, err := filepath.Abs(lexical); err == nil {
lexical = filepath.Clean(absolute)
}
aliases := []string{lexical}
if real, err := filepath.EvalSymlinks(lexical); err == nil {
real = filepath.Clean(real)
if real != lexical {
aliases = append(aliases, real)
}
}
return aliases
}
func normalizedLatestFindingPath(path string) string {
if path == "" {
return ""
}
lexical := filepath.Clean(path)
if absolute, err := filepath.Abs(lexical); err == nil {
lexical = filepath.Clean(absolute)
}
return lexical
}
// latestFindingsCap bounds the active set to keep memory and the persisted
// file bounded; the ordering below decides what the cap keeps.
const latestFindingsCap = 15000
// orderAndCapLatest flattens a keyed set into a deterministic order:
// severity first, then most recent, then key. Map iteration order used to
// decide both the file bytes and which findings a full set dropped.
func orderAndCapLatest(existing map[string]alert.Finding) []alert.Finding {
return orderAndCapLatestPreserving(existing, nil)
}
func orderAndCapLatestPreserving(existing map[string]alert.Finding, protected map[string]struct{}) []alert.Finding {
merged := make([]alert.Finding, 0, len(existing))
for _, f := range existing {
merged = append(merged, f)
}
sort.Slice(merged, func(i, j int) bool {
if merged[i].Severity != merged[j].Severity {
return merged[i].Severity > merged[j].Severity
}
if !merged[i].Timestamp.Equal(merged[j].Timestamp) {
return merged[i].Timestamp.After(merged[j].Timestamp)
}
return merged[i].Key() < merged[j].Key()
})
if len(merged) <= latestFindingsCap {
return merged
}
protectedCount := 0
for _, f := range merged {
if _, ok := protected[f.Key()]; ok {
protectedCount++
}
}
if protectedCount == 0 {
return merged[:latestFindingsCap]
}
// A preservation transaction must not retire an unexamined finding merely
// because newly merged, higher-severity findings reached the normal cap.
// Keep every protected identity, then fill the remaining slots in the same
// deterministic priority order. In the pathological case that protected
// state alone exceeds the cap, the coverage invariant wins temporarily.
unprotectedSlots := latestFindingsCap - protectedCount
if unprotectedSlots < 0 {
unprotectedSlots = 0
}
capped := make([]alert.Finding, 0, max(latestFindingsCap, protectedCount))
for _, f := range merged {
if _, ok := protected[f.Key()]; ok {
capped = append(capped, f)
continue
}
if unprotectedSlots > 0 {
capped = append(capped, f)
unprotectedSlots--
}
}
return capped
}
// latestFindingsWriter writes the persisted file; a seam for tests.
var latestFindingsWriter = atomicio.AtomicWrite
// persistLatestLocked writes latest_findings.json when its content changed
// since the last write. Callers hold latestMu. Compact JSON: the file is
// read back by this process only.
func (s *Store) persistLatestLocked() {
data, err := json.Marshal(s.latestFindings)
if err != nil {
return
}
digest := sha256.Sum256(data)
if s.latestDigestSet && digest == s.latestDigest {
return
}
path := filepath.Join(s.path, "latest_findings.json")
// Older releases wrote through a fixed <path>.tmp; clear one left by a
// crash so it cannot be mistaken for live state.
if err := os.Remove(path + ".tmp"); err != nil && !os.IsNotExist(err) {
return
}
if err := latestFindingsWriter(path, 0o600, data); err != nil {
return
}
s.latestDigest = digest
s.latestDigestSet = true
}
func shouldPurgeLatestFinding(f alert.Finding, remove map[string]bool) bool {
if remove[f.Check] {
return true
}
if f.Check != "check_timeout" {
return false
}
runner, ok := timeoutFindingRunner(f.Message)
return ok && remove[runner]
}
func timeoutFindingRunner(msg string) (string, bool) {
rest, ok := strings.CutPrefix(msg, "Check '")
if !ok {
return "", false
}
runner, _, ok := strings.Cut(rest, "'")
return runner, ok && runner != ""
}
// ClearLatestFindings removes all findings from the latest set.
// Use before SetLatestFindings for a full replace (e.g. initial scan).
func (s *Store) ClearLatestFindings() {
s.latestMu.Lock()
defer s.latestMu.Unlock()
s.latestFindings = nil
}
// LatestFindings returns the full results of the most recent scan.
// This is what the Findings page shows - "what's wrong right now."
func (s *Store) LatestFindings() []alert.Finding {
s.latestMu.RLock()
defer s.latestMu.RUnlock()
result := make([]alert.Finding, len(s.latestFindings))
copy(result, s.latestFindings)
return result
}
// LatestScanTime returns when the last scan completed.
func (s *Store) LatestScanTime() time.Time {
s.latestMu.RLock()
defer s.latestMu.RUnlock()
return s.latestScanTime
}
const baselineAtMetaKey = "__baseline_at"
// EnsureBaseline records the first-start timestamp the first time it is
// called against a fresh state directory. Subsequent calls preserve the
// original value so reinstalls / upgrades do not reset the baseline. Safe
// to call from the daemon boot path on every start.
func (s *Store) EnsureBaseline(now time.Time) {
if _, ok := s.GetRaw(baselineAtMetaKey); ok {
return
}
s.SetRaw(baselineAtMetaKey, now.UTC().Format(time.RFC3339Nano))
}
// BaselineAt returns the persisted baseline timestamp, or the zero time
// when EnsureBaseline has not been called yet.
func (s *Store) BaselineAt() time.Time {
raw, ok := s.GetRaw(baselineAtMetaKey)
if !ok {
return time.Time{}
}
t, err := time.Parse(time.RFC3339Nano, raw)
if err != nil {
return time.Time{}
}
return t
}
// DismissLatestFinding removes a finding from the latest scan results.
func (s *Store) DismissLatestFinding(key string) {
s.latestMu.Lock()
defer s.latestMu.Unlock()
var filtered []alert.Finding
for _, f := range s.latestFindings {
if f.Key() != key {
filtered = append(filtered, f)
}
}
s.latestFindings = filtered
s.persistLatestLocked()
}
// DemoteLatestFinding conditionally lowers a finding's severity in the latest
// scan results. The expected snapshot prevents a completed scan or realtime
// alert from being overwritten by an older re-verification result.
//
// Check, Message and Details never change. Finding.Key() hashes those fields,
// so an explanation written into the finding would orphan every dismissal,
// suppression and alert-dedup entry already keyed to it. DemotedFrom records
// only the severity needed to reverse the operation.
func (s *Store) DemoteLatestFinding(expected alert.Finding, severity alert.Severity) bool {
if severity != alert.Warning || expected.Severity < alert.High || expected.Severity > alert.Critical {
return false
}
s.latestMu.Lock()
defer s.latestMu.Unlock()
for i := range s.latestFindings {
current := &s.latestFindings[i]
if current.Key() != expected.Key() || current.Severity <= severity {
continue
}
if !sameLatestFindingSnapshot(*current, expected) {
return false
}
current.DemotedFrom = current.Severity
current.Severity = severity
s.latestFindings = orderAndCapLatest(findingsByKey(s.latestFindings))
s.persistLatestLocked()
return true
}
return false
}
// DismissFindingIfLatest clears only the finding snapshot that was actually
// verified. A realtime alert or completed scan may refresh the same key while
// verification is in flight; that newer evidence must remain active.
func (s *Store) DismissFindingIfLatest(expected alert.Finding) bool {
s.latestMu.Lock()
defer s.latestMu.Unlock()
for i := range s.latestFindings {
if !sameLatestFindingSnapshot(s.latestFindings[i], expected) {
continue
}
s.mu.Lock()
if entry, exists := s.entries[expected.Key()]; exists {
entry.IsBaseline = true
entry.DismissalID = ""
s.dirty = true
}
s.mu.Unlock()
s.latestFindings = append(s.latestFindings[:i], s.latestFindings[i+1:]...)
s.persistLatestLocked()
return true
}
return false
}
// RestoreLatestFindingSeverity reverses an automatic demotion after the exact
// verifier that owns the finding reports the content as live again.
func (s *Store) RestoreLatestFindingSeverity(expected alert.Finding) bool {
s.latestMu.Lock()
defer s.latestMu.Unlock()
for i := range s.latestFindings {
current := &s.latestFindings[i]
if current.Key() != expected.Key() || !isAutomaticallyDemotedFinding(*current) {
continue
}
if !sameLatestFindingSnapshot(*current, expected) {
return false
}
current.Severity = current.DemotedFrom
current.DemotedFrom = alert.Warning
s.latestFindings = orderAndCapLatest(findingsByKey(s.latestFindings))
s.persistLatestLocked()
return true
}
return false
}
func isAutomaticallyDemotedFinding(f alert.Finding) bool {
return f.Severity == alert.Warning &&
f.DemotedFrom >= alert.High && f.DemotedFrom <= alert.Critical
}
// sameLatestFindingSnapshot reports whether the stored finding is still the one
// verification looked at. Any field the verifier's decision rested on must
// match, or a newer detection would be silently overwritten by a stale verdict.
func sameLatestFindingSnapshot(current, expected alert.Finding) bool {
if current.Key() != expected.Key() ||
current.Check != expected.Check ||
current.Message != expected.Message ||
current.Details != expected.Details ||
current.Severity != expected.Severity ||
!current.Timestamp.Equal(expected.Timestamp) ||
current.FilePath != expected.FilePath ||
current.ContentSHA256 != expected.ContentSHA256 ||
current.DetectLogic != expected.DetectLogic {
return false
}
return current.DemotedFrom == expected.DemotedFrom
}
func findingsByKey(findings []alert.Finding) map[string]alert.Finding {
keyed := make(map[string]alert.Finding, len(findings))
for _, finding := range findings {
keyed[finding.Key()] = finding
}
return keyed
}
// DismissFinding marks a finding as baseline (acknowledged/dismissed).
// It will no longer appear in active findings or trigger new alerts.
func (s *Store) DismissFinding(key string) {
s.mu.Lock()
defer s.mu.Unlock()
if entry, exists := s.entries[key]; exists {
entry.IsBaseline = true
entry.DismissalID = ""
s.dirty = true
}
}
// DismissUndo records what DismissFindingWithUndo changed, so UndoDismiss can
// return the finding to the state it had before the operator dismissed it.
type DismissUndo struct {
Key string `json:"key"`
// Identity and hash prevent an old undo from reversing a later decision
// or changing alert state for newer evidence under the same key.
DismissalID string `json:"dismissal_id"`
Hash string `json:"hash"`
ClearBaseline bool `json:"clear_baseline,omitempty"`
CreatedEntry bool `json:"created_entry,omitempty"`
Removed []alert.Finding `json:"removed,omitempty"`
}
// DismissFindingWithUndo marks a finding as baseline and removes it from the
// latest list as one operation. Latest-only findings also need alert state:
// realtime findings can reach the UI before the next scan updates dedup state.
func (s *Store) DismissFindingWithUndo(key string) DismissUndo {
s.latestMu.Lock()
defer s.latestMu.Unlock()
s.mu.Lock()
defer s.mu.Unlock()
u := DismissUndo{Key: key}
var kept []alert.Finding
for _, f := range s.latestFindings {
if f.Key() == key {
u.Removed = append(u.Removed, f)
continue
}
kept = append(kept, f)
}
entry := s.entries[key]
if entry == nil && len(u.Removed) > 0 {
now := time.Now()
entry = &Entry{Hash: findingHash(u.Removed[0]), FirstSeen: now, LastSeen: now}
s.entries[key] = entry
u.CreatedEntry = true
}
if entry != nil {
u.ClearBaseline = !entry.IsBaseline
entry.IsBaseline = true
entry.DismissalID = rand.Text()
u.DismissalID, u.Hash = entry.DismissalID, entry.Hash
s.dirty = true
}
s.latestFindings = kept
s.persistLatestLocked()
return u
}
// UndoDismiss reverses only the dismissal that still owns the entry. A copy
// reported by a scan in the meantime is newer evidence and remains listed.
func (s *Store) UndoDismiss(u DismissUndo) bool {
s.latestMu.Lock()
defer s.latestMu.Unlock()
s.mu.Lock()
defer s.mu.Unlock()
entry := s.entries[u.Key]
if u.DismissalID == "" || entry == nil || entry.DismissalID != u.DismissalID || entry.Hash != u.Hash {
return false
}
entry.DismissalID = ""
if u.CreatedEntry {
delete(s.entries, u.Key)
} else if u.ClearBaseline {
entry.IsBaseline = false
}
s.dirty = true
if len(u.Removed) > 0 {
merged := findingsByKey(s.latestFindings)
for _, f := range u.Removed {
if _, present := merged[f.Key()]; !present {
merged[f.Key()] = f
}
}
s.latestFindings = orderAndCapLatest(merged)
s.persistLatestLocked()
}
return true
}
// ParseKey splits a state key "check:message" into its components.
func ParseKey(key string) (check, message string) {
for i := 0; i < len(key); i++ {
if key[i] == ':' {
return key[:i], key[i+1:]
}
}
return key, ""
}
// --- Suppression rules ---
// SuppressionRule defines a rule for suppressing specific findings.
type SuppressionRule struct {
ID string `json:"id"`
Check string `json:"check"`
PathPattern string `json:"path_pattern,omitempty"`
Reason string `json:"reason"`
CreatedAt time.Time `json:"created_at"`
}
// LoadSuppressions reads suppression rules from disk.
func (s *Store) LoadSuppressions() []SuppressionRule {
data, err := os.ReadFile(filepath.Join(s.path, "suppressions.json"))
if err != nil {
return nil
}
var rules []SuppressionRule
if err := json.Unmarshal(data, &rules); err != nil {
return nil
}
return rules
}
// SaveSuppressions writes suppression rules to disk atomically with fsync.
// Callers that change the existing rules use UpdateSuppressions instead.
func (s *Store) SaveSuppressions(rules []SuppressionRule) error {
return atomicio.AtomicWriteJSON(filepath.Join(s.path, "suppressions.json"), 0o600, rules)
}
// UpdateSuppressions applies fn to the stored rules and saves what it returns
// as one step. An error from fn leaves the stored rules unchanged.
func (s *Store) UpdateSuppressions(fn func([]SuppressionRule) ([]SuppressionRule, error)) error {
s.suppressMu.Lock()
defer s.suppressMu.Unlock()
rules, err := fn(s.LoadSuppressions())
if err != nil {
return err
}
return s.SaveSuppressions(rules)
}
// IsSuppressed checks if a finding matches any loaded suppression rule.
// Load rules once with LoadSuppressions() and pass them in to avoid
// re-reading the file for every finding.
func (s *Store) IsSuppressed(f alert.Finding, rules []SuppressionRule) bool {
for _, rule := range rules {
if config.CanonicalCheckName(f.Check) != config.CanonicalCheckName(rule.Check) {
continue
}
// If no path pattern, suppress all findings for this check type
if rule.PathPattern == "" {
return true
}
// Match against the finding's FilePath
if f.FilePath != "" {
if matched, _ := filepath.Match(rule.PathPattern, f.FilePath); matched {
return true
}
}
for _, candidate := range suppressionPathCandidates(f) {
if matched, _ := filepath.Match(rule.PathPattern, candidate); matched {
return true
}
}
}
return false
}
func suppressionPathCandidates(f alert.Finding) []string {
if f.FilePath != "" {
return []string{f.FilePath}
}
fields := strings.Fields(f.Message + " " + f.Details)
seen := make(map[string]bool)
var paths []string
for _, field := range fields {
field = strings.Trim(field, `"'():,;[]{}<>`)
if !strings.HasPrefix(field, "/") {
continue
}
candidate := filepath.Clean(field)
if candidate == "." || candidate == "/" || seen[candidate] {
continue
}
seen[candidate] = true
paths = append(paths, candidate)
}
return paths
}
package store
import (
"bytes"
"encoding/json"
"strings"
"time"
bolt "go.etcd.io/bbolt"
)
// AdminEmailEntry records that a given email was observed as a WordPress
// administrator on (Account, Schema). One email may carry multiple
// entries when the same person administers several customer sites --
// surfacing that overlap is the whole point of the bucket.
type AdminEmailEntry struct {
Account string `json:"account"`
Schema string `json:"schema"`
LastSeen time.Time `json:"last_seen"`
}
const adminEmailsBucket = "admin:emails"
// RecordAdminEmail upserts an observation that `email` is administrator
// on (account, schema) at `now`. Re-recording the same triple updates
// LastSeen without creating a duplicate row; recording the same email
// across a different (account, schema) appends to the owner list. The
// email is lowercased before storage so case-mismatched recordings
// collapse to a single key.
func (db *DB) RecordAdminEmail(email, account, schema string, now time.Time) error {
email = strings.ToLower(strings.TrimSpace(email))
if email == "" || account == "" {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(adminEmailsBucket))
var entries []AdminEmailEntry
if v := b.Get([]byte(email)); v != nil {
if err := json.Unmarshal(v, &entries); err != nil {
// Corrupt entry: restart the list rather than fail the
// whole write -- we'd rather lose one stale observation
// than block detection on a malformed payload.
entries = nil
}
}
updated := false
for i := range entries {
if entries[i].Account == account && entries[i].Schema == schema {
entries[i].LastSeen = now
updated = true
break
}
}
if !updated {
entries = append(entries, AdminEmailEntry{
Account: account,
Schema: schema,
LastSeen: now,
})
}
payload, err := json.Marshal(entries)
if err != nil {
return err
}
return b.Put([]byte(email), payload)
})
}
// AdminEmailOwners returns the full list of (account, schema, last_seen)
// triples recorded for `email`. Returns an empty slice when the email
// is unknown.
func (db *DB) AdminEmailOwners(email string) ([]AdminEmailEntry, error) {
email = strings.ToLower(strings.TrimSpace(email))
var out []AdminEmailEntry
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(adminEmailsBucket))
v := b.Get([]byte(email))
if v == nil {
return nil
}
return json.Unmarshal(v, &out)
})
return out, err
}
// OverlappingAdminEmails returns every email whose owner list has at
// least `minAccounts` distinct accounts after stale entries (older than
// `retention`) are pruned.
//
// `minAccounts` smaller than 2 is clamped to 2 because a single-account
// observation is never an overlap by definition.
func (db *DB) OverlappingAdminEmails(minAccounts int, retention time.Duration) (map[string][]AdminEmailEntry, error) {
if minAccounts < 2 {
minAccounts = 2
}
cutoff := time.Now().Add(-retention)
out := map[string][]AdminEmailEntry{}
type pruneItem struct {
key []byte
original []byte
fresh []byte
}
var toPrune []pruneItem
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(adminEmailsBucket))
return b.ForEach(func(k, v []byte) error {
var entries []AdminEmailEntry
if err := json.Unmarshal(v, &entries); err != nil {
return nil //nolint:nilerr // skip corrupt entry
}
fresh := entries[:0]
seen := map[string]struct{}{}
for _, e := range entries {
if e.LastSeen.Before(cutoff) {
continue
}
key := e.Account + "|" + e.Schema
if _, dup := seen[key]; dup {
continue
}
seen[key] = struct{}{}
fresh = append(fresh, e)
}
distinct := map[string]struct{}{}
for _, e := range fresh {
distinct[e.Account] = struct{}{}
}
if len(distinct) >= minAccounts {
out[string(k)] = append([]AdminEmailEntry(nil), fresh...)
}
if len(fresh) == 0 {
toPrune = append(toPrune, pruneItem{key: append([]byte(nil), k...), original: append([]byte(nil), v...)})
} else if len(fresh) != len(entries) {
payload, err := json.Marshal(fresh)
if err != nil {
return err
}
toPrune = append(toPrune, pruneItem{
key: append([]byte(nil), k...),
original: append([]byte(nil), v...),
fresh: payload,
})
}
return nil
})
})
if err != nil || len(toPrune) == 0 {
return out, err
}
err = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(adminEmailsBucket))
for _, item := range toPrune {
if !bytes.Equal(b.Get(item.key), item.original) {
continue
}
if len(item.fresh) == 0 {
if deleteErr := b.Delete(item.key); deleteErr != nil {
return deleteErr
}
continue
}
if putErr := b.Put(item.key, item.fresh); putErr != nil {
return putErr
}
}
return nil
})
return out, err
}
package store
import (
"encoding/json"
"time"
"github.com/pidginhost/csm/internal/alert"
bolt "go.etcd.io/bbolt"
)
// SeverityBucket holds aggregated counts by severity.
type SeverityBucket struct {
Critical int `json:"critical"`
High int `json:"high"`
Warning int `json:"warning"`
Total int `json:"total"`
}
// HourBucket is a SeverityBucket for the hour that begins at Start.
type HourBucket struct {
Start time.Time `json:"start"`
SeverityBucket
}
// DayBucket is a SeverityBucket keyed by date.
type DayBucket struct {
Date string `json:"date"`
SeverityBucket
}
// AggregateByHour returns 24 hourly buckets (oldest first) for the last 24 hours.
// It seeks directly to the start key in bbolt, scanning only the relevant range.
func (db *DB) AggregateByHour() []HourBucket {
now := time.Now()
currentHour := now.Truncate(time.Hour)
cutoff := currentHour.Add(-23 * time.Hour)
// Map: hours-ago (0=current, 23=oldest) → counts
counts := make(map[int]*SeverityBucket, 24)
for i := 0; i < 24; i++ {
counts[i] = &SeverityBucket{}
}
seekPrefix := timeKeyLowerBound(cutoff)
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
if b == nil {
return nil
}
c := b.Cursor()
// Seek to the earliest key that could be in our 24h window.
for k, v := c.Seek([]byte(seekPrefix)); k != nil; k, v = c.Next() {
var f alert.Finding
if err := json.Unmarshal(v, &f); err != nil {
continue
}
if f.Timestamp.Before(cutoff) {
continue
}
if f.Timestamp.After(now) {
continue
}
fHour := f.Timestamp.Truncate(time.Hour)
hoursAgo := int(currentHour.Sub(fHour).Hours())
if hoursAgo < 0 || hoursAgo >= 24 {
continue
}
bucket := counts[hoursAgo]
bucket.Total++
switch f.Severity {
case alert.Critical:
bucket.Critical++
case alert.High:
bucket.High++
case alert.Warning:
bucket.Warning++
}
}
return nil
})
// Build result oldest→newest (23h ago → 0h ago)
result := make([]HourBucket, 24)
for i := 0; i < 24; i++ {
hoursAgo := 23 - i
t := currentHour.Add(-time.Duration(hoursAgo) * time.Hour)
result[i] = HourBucket{
Start: t,
SeverityBucket: *counts[hoursAgo],
}
}
return result
}
// ReadHistorySince returns all findings since the given time, using bbolt cursor
// seeking for efficiency. Results are newest-first.
func (db *DB) ReadHistorySince(since time.Time) []alert.Finding {
seekPrefix := timeKeyLowerBound(since)
var results []alert.Finding
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
if b == nil {
return nil
}
c := b.Cursor()
for k, v := c.Seek([]byte(seekPrefix)); k != nil; k, v = c.Next() {
var f alert.Finding
if err := json.Unmarshal(v, &f); err != nil {
continue
}
results = append(results, f)
}
return nil
})
// Reverse to newest-first
for i, j := 0, len(results)-1; i < j; i, j = i+1, j-1 {
results[i], results[j] = results[j], results[i]
}
return results
}
// SearchHistorySince returns up to limit findings since the given time,
// newest-first. The matcher runs while the bbolt cursor walks backward so
// callers that only need a bounded result set do not materialize the whole
// time window first.
func (db *DB) SearchHistorySince(since time.Time, limit int, match func(alert.Finding) bool) []alert.Finding {
if limit <= 0 {
return nil
}
cutoffKey := timeKeyLowerBound(since)
var results []alert.Finding
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
if b == nil {
return nil
}
c := b.Cursor()
for k, v := c.Last(); k != nil && len(results) < limit; k, v = c.Prev() {
if string(k) < cutoffKey {
break
}
var f alert.Finding
if err := json.Unmarshal(v, &f); err != nil {
continue
}
if match != nil && !match(f) {
continue
}
results = append(results, f)
}
return nil
})
return results
}
// AggregateByDay returns 30 daily buckets (oldest first) for the last 30 days.
// Reads from the pre-aggregated stats:daily bucket so the trend chart is
// not affected by history pruning.
func (db *DB) AggregateByDay() []DayBucket {
return db.AggregateByDayN(30)
}
// AggregateByDayN returns `days` daily buckets (oldest first) ending today.
// Days outside [1, dailyRetentionDays] are clamped to that range. Days with
// no recorded findings are returned as zero-value buckets.
func (db *DB) AggregateByDayN(days int) []DayBucket {
if days < 1 {
days = 1
}
if days > dailyRetentionDays {
days = dailyRetentionDays
}
// Hard cap so a future bump to dailyRetentionDays cannot accidentally
// drive a huge slice allocation here. Daily aggregation over more than
// a decade is not a real workload for this surface.
const aggregateBucketHardCap = 4096
if days > aggregateBucketHardCap {
days = aggregateBucketHardCap
}
now := time.Now()
today := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.Local)
cutoff := today.AddDate(0, 0, -(days - 1))
buckets := make([]DayBucket, days)
for i := 0; i < days; i++ {
d := cutoff.AddDate(0, 0, i)
buckets[i] = DayBucket{Date: d.Format("2006-01-02")}
}
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucketStatsDaily))
if b == nil {
return nil
}
for i := range buckets {
v := b.Get([]byte(buckets[i].Date))
if v == nil {
continue
}
var sb SeverityBucket
if err := json.Unmarshal(v, &sb); err != nil {
continue
}
buckets[i].SeverityBucket = sb
}
return nil
})
return buckets
}
package store
import (
"archive/tar"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"os"
"path/filepath"
"sort"
"strings"
"syscall"
"time"
"github.com/klauspost/compress/zstd"
bolt "go.etcd.io/bbolt"
)
// writeFileNoFollow is os.WriteFile that refuses to write through a
// symlink standing at path.
//
// #nosec G304 G703 -- path is the export destination the operator named.
func writeFileNoFollow(path string, data []byte, perm os.FileMode) error {
f, err := os.OpenFile(path, os.O_CREATE|os.O_WRONLY|os.O_TRUNC|syscall.O_NOFOLLOW, perm)
if err != nil {
return err
}
if _, err := f.Write(data); err != nil {
_ = f.Close()
return err
}
return f.Close()
}
// ArchiveSchemaVersion is the on-wire schema for backup archives. Bump
// when the manifest layout or contents shape changes incompatibly. Old
// CSM binaries refuse archives newer than the version they understand.
const ArchiveSchemaVersion = 1
// Standard entry names inside the tar.
const (
manifestEntry = "manifest.json"
bboltSnapshotEntry = "bbolt.snapshot"
stateEntryPrefix = "state/"
rulesEntryPrefix = "rules/"
stateLockFileName = "csm.lock"
)
var transientStatePaths = []string{
"firewall/confirm_pending",
"firewall/rollback.nft",
"firewall/rollback.nft.state.json",
"firewall/rollback.nft.config.json",
"firewall/rollback.sh",
}
// Sentinel errors so callers can branch on the failure mode instead of
// matching strings.
var (
ErrSchemaVersionTooNew = errors.New("archive schema version is newer than this binary supports")
ErrPlatformMismatch = errors.New("archive source platform does not match current host")
ErrManifestMissing = errors.New("archive does not contain manifest.json")
ErrCorruptArchive = errors.New("archive is corrupt or not a CSM backup")
)
// Manifest is the JSON header at the top of every archive.
type Manifest struct {
SchemaVersion int `json:"schema_version"`
CSMVersion string `json:"csm_version"`
SourceHostname string `json:"source_hostname"`
SourcePlatform map[string]string `json:"source_platform"`
ExportTS time.Time `json:"export_ts"`
Contents []string `json:"contents"`
BboltBuckets []string `json:"bbolt_buckets,omitempty"`
BboltSHA256 string `json:"bbolt_sha256,omitempty"`
}
// ExportOptions configures Export.
type ExportOptions struct {
StatePath string // /var/lib/csm/state, source for state JSON files
RulesPath string // /opt/csm/rules, source for signature cache (empty -> skip)
DstPath string // .csmbak file to create
Manifest Manifest // caller fills CSMVersion/SourceHostname/SourcePlatform; rest filled here
}
// ExportResult summarises a successful export.
type ExportResult struct {
Path string
Bytes int64
ArchiveSHA256 string
BboltSHA256 string
}
// ImportOptions configures Import.
type ImportOptions struct {
SrcPath string
StatePath string
RulesPath string
Only string // "all" | "baseline" | "firewall"
ForcePlatformMismatch bool
CurrentPlatform map[string]string // for the mismatch check
}
// ImportResult summarises a successful import.
type ImportResult struct {
Manifest Manifest
BucketsRestored []string
StateFiles int
RulesFiles int
}
// Export writes a tar+zstd archive containing a bbolt snapshot, the
// state directory, and (optionally) the signature-rules directory. The
// daemon is the single source of truth for paths; the caller fills the
// manifest with hostname/version/platform.
func (db *DB) Export(opts ExportOptions) (*ExportResult, error) {
if opts.DstPath == "" {
return nil, errors.New("DstPath is empty")
}
if opts.StatePath == "" {
return nil, errors.New("StatePath is empty")
}
man := opts.Manifest
if man.SchemaVersion == 0 {
man.SchemaVersion = ArchiveSchemaVersion
}
if man.ExportTS.IsZero() {
man.ExportTS = time.Now().UTC()
}
man.Contents = []string{"bbolt", "state"}
if opts.RulesPath != "" {
man.Contents = append(man.Contents, "rules")
}
// Snapshot bbolt to a temp file in the same directory so we can hash it
// and stream it into the tar without holding a long bolt transaction.
snapDir := filepath.Dir(opts.DstPath)
snap, err := os.CreateTemp(snapDir, "csm-export-bbolt-*.snap")
if err != nil {
return nil, fmt.Errorf("creating bbolt snapshot temp: %w", err)
}
snapPath := snap.Name()
defer os.Remove(snapPath)
err = db.bolt.View(func(tx *bolt.Tx) error {
_, werr := tx.WriteTo(snap)
return werr
})
if err != nil {
_ = snap.Close()
return nil, fmt.Errorf("bbolt snapshot: %w", err)
}
if err = snap.Close(); err != nil {
return nil, fmt.Errorf("closing bbolt snapshot: %w", err)
}
if err = DisarmBrowserSessionsSnapshot(snapPath); err != nil {
return nil, err
}
if err = DisarmFirewallRollbackSnapshot(snapPath); err != nil {
return nil, err
}
// Describe the sanitized snapshot, not the live database: a full import
// reports these buckets as restored.
man.BboltBuckets, err = listSnapshotBuckets(snapPath)
if err != nil {
return nil, err
}
man.BboltSHA256, err = sha256File(snapPath)
if err != nil {
return nil, fmt.Errorf("hashing bbolt snapshot: %w", err)
}
// Build the archive on disk; hash it as we write. Close is called
// explicitly below so any close error after fsync is surfaced --
// silently dropping it would mean the operator gets "export
// succeeded" for a file that may not be fully persisted.
// O_NOFOLLOW: a local account that gets to create the destination name
// first must not have the daemon write the archive through a symlink
// into a file it can read.
out, err := os.OpenFile(opts.DstPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC|syscall.O_NOFOLLOW, 0600)
if err != nil {
return nil, fmt.Errorf("creating archive: %w", err)
}
// closed gates the deferred Close so the explicit Close below
// can return its error without a double-close. success gates the
// cleanup of the partial archive + companion file: any
// error-return path between OpenFile and the final return removes
// the half-written files so the operator does not mistake them
// for a usable backup.
closed := false
success := false
defer func() {
if !closed {
_ = out.Close()
}
if !success {
_ = os.Remove(opts.DstPath)
_ = os.Remove(opts.DstPath + ".sha256")
}
}()
archHash := sha256.New()
mw := io.MultiWriter(out, archHash)
zw, err := zstd.NewWriter(mw, zstd.WithEncoderLevel(zstd.SpeedDefault))
if err != nil {
return nil, fmt.Errorf("zstd writer: %w", err)
}
tw := tar.NewWriter(zw)
// 1. manifest first so a streaming reader sees schema info before payload.
manBytes, err := json.MarshalIndent(man, "", " ")
if err != nil {
return nil, fmt.Errorf("marshal manifest: %w", err)
}
if err = writeTarFile(tw, manifestEntry, manBytes, man.ExportTS); err != nil {
return nil, err
}
// 2. bbolt snapshot (already on disk).
if err = streamFileToTar(tw, bboltSnapshotEntry, snapPath, man.ExportTS); err != nil {
return nil, err
}
// 3. state files (skip runtime-owned files captured separately or not at all).
stateSkip := append([]string{"csm.db", stateLockFileName, "exports"}, transientStatePaths...)
if _, err = walkDirIntoTar(tw, opts.StatePath, stateEntryPrefix, stateSkip, man.ExportTS); err != nil {
return nil, err
}
// 4. rules files (optional).
if opts.RulesPath != "" {
if _, err = walkDirIntoTar(tw, opts.RulesPath, rulesEntryPrefix, nil, man.ExportTS); err != nil {
return nil, err
}
}
if err = tw.Close(); err != nil {
return nil, fmt.Errorf("closing tar: %w", err)
}
if err = zw.Close(); err != nil {
return nil, fmt.Errorf("closing zstd: %w", err)
}
if err = out.Sync(); err != nil {
return nil, fmt.Errorf("fsync archive: %w", err)
}
if err = out.Close(); err != nil {
return nil, fmt.Errorf("closing archive: %w", err)
}
closed = true
info, err := os.Stat(opts.DstPath)
if err != nil {
return nil, fmt.Errorf("stat archive: %w", err)
}
archiveSHA := hex.EncodeToString(archHash.Sum(nil))
// Write companion .sha256 file alongside for operator verification.
// O_NOFOLLOW for the same reason as the archive itself: the companion
// can truncate whatever a planted symlink points at.
companion := opts.DstPath + ".sha256"
companionLine := fmt.Sprintf("%s %s\n", archiveSHA, filepath.Base(opts.DstPath))
if err = writeFileNoFollow(companion, []byte(companionLine), 0600); err != nil {
return nil, fmt.Errorf("writing companion sha256: %w", err)
}
success = true
return &ExportResult{
Path: opts.DstPath,
Bytes: info.Size(),
ArchiveSHA256: archiveSHA,
BboltSHA256: man.BboltSHA256,
}, nil
}
// Import unpacks an archive into the target state and rules paths. Live
// daemons must be stopped first; callers enforce that before invoking.
func Import(opts ImportOptions) (*ImportResult, error) {
if opts.SrcPath == "" {
return nil, errors.New("SrcPath is empty")
}
if opts.StatePath == "" {
return nil, errors.New("StatePath is empty")
}
only := opts.Only
if only == "" {
only = "all"
}
switch only {
case "all", "baseline", "firewall":
default:
return nil, fmt.Errorf("invalid Only value %q (want all|baseline|firewall)", only)
}
in, err := os.Open(opts.SrcPath)
if err != nil {
return nil, fmt.Errorf("opening archive: %w", err)
}
defer in.Close()
zr, err := zstd.NewReader(in)
if err != nil {
return nil, fmt.Errorf("%w: %v", ErrCorruptArchive, err)
}
defer zr.Close()
tr := tar.NewReader(zr)
// Manifest must come first.
hdr, err := tr.Next()
if err != nil {
return nil, fmt.Errorf("%w: reading first entry: %v", ErrCorruptArchive, err)
}
if hdr.Name != manifestEntry {
return nil, fmt.Errorf("%w: first entry is %q, want %q", ErrManifestMissing, hdr.Name, manifestEntry)
}
manBytes, err := io.ReadAll(tr)
if err != nil {
return nil, fmt.Errorf("%w: reading manifest: %v", ErrCorruptArchive, err)
}
var man Manifest
if err = json.Unmarshal(manBytes, &man); err != nil {
return nil, fmt.Errorf("%w: parsing manifest: %v", ErrCorruptArchive, err)
}
if man.SchemaVersion > ArchiveSchemaVersion {
return nil, fmt.Errorf("%w: archive=%d binary=%d", ErrSchemaVersionTooNew, man.SchemaVersion, ArchiveSchemaVersion)
}
if !opts.ForcePlatformMismatch && !platformMatches(man.SourcePlatform, opts.CurrentPlatform) {
return nil, fmt.Errorf("%w: archive=%v current=%v (use --force-platform-mismatch to override)", ErrPlatformMismatch, man.SourcePlatform, opts.CurrentPlatform)
}
// Stage every payload into a temp dir; commit only after a complete
// successful read so a half-imported state is impossible.
stage, err := os.MkdirTemp(filepath.Dir(opts.StatePath), "csm-import-stage-*")
if err != nil {
return nil, fmt.Errorf("creating staging dir: %w", err)
}
defer os.RemoveAll(stage)
stagedBbolt := ""
stagedState := []string{}
stagedRules := []string{}
for {
nextHdr, nextErr := tr.Next()
if nextErr == io.EOF {
break
}
if nextErr != nil {
return nil, fmt.Errorf("%w: reading entry: %v", ErrCorruptArchive, nextErr)
}
if nextHdr.Typeflag != tar.TypeReg {
continue
}
clean := filepath.Clean(nextHdr.Name)
if strings.HasPrefix(clean, "..") || filepath.IsAbs(clean) {
return nil, fmt.Errorf("%w: unsafe entry name %q", ErrCorruptArchive, nextHdr.Name)
}
dst := filepath.Join(stage, clean)
// Defence in depth against zip-slip: confirm the joined path still
// resolves inside the staging dir even after symlinks / unusual
// path elements in the archive header.
rel, relErr := filepath.Rel(stage, dst)
if relErr != nil || rel == ".." || strings.HasPrefix(rel, ".."+string(filepath.Separator)) || filepath.IsAbs(rel) {
return nil, fmt.Errorf("%w: unsafe entry name %q", ErrCorruptArchive, nextHdr.Name)
}
if isTransientStateArchiveEntry(clean) {
continue
}
if mkErr := os.MkdirAll(filepath.Dir(dst), 0700); mkErr != nil {
return nil, fmt.Errorf("staging dir: %w", mkErr)
}
// #nosec G304 -- dst is filepath.Join(stage, clean) where stage is a freshly created MkdirTemp under StatePath's parent and clean has been validated three lines above (no ".." prefix, not absolute). Cannot escape the staging dir.
f, openErr := os.OpenFile(dst, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0600)
if openErr != nil {
return nil, fmt.Errorf("staging file: %w", openErr)
}
// nextHdr.Size bound caps bytes copied so a hostile archive can't
// drive the stage dir to fill the filesystem.
if _, copyErr := io.CopyN(f, tr, nextHdr.Size); copyErr != nil {
_ = f.Close()
return nil, fmt.Errorf("staging copy %s: %w", clean, copyErr)
}
if closeErr := f.Close(); closeErr != nil {
return nil, fmt.Errorf("closing staged file: %w", closeErr)
}
switch {
case clean == bboltSnapshotEntry:
stagedBbolt = dst
case strings.HasPrefix(clean, stateEntryPrefix):
stagedState = append(stagedState, clean)
case strings.HasPrefix(clean, rulesEntryPrefix):
stagedRules = append(stagedRules, clean)
}
}
if (only == "all" || only == "firewall") && stagedBbolt == "" {
return nil, fmt.Errorf("%w: %s import needs bbolt snapshot in archive", ErrCorruptArchive, only)
}
// Verify the staged bbolt snapshot against the hash the manifest recorded
// at export time before any bbolt-consuming import applies payloads.
// zstd's frame CRC catches transport corruption, but a snapshot that was
// altered before compression would otherwise be promoted over csm.db
// unchecked. Empty hash means a pre-hash archive; skip rather than reject.
if (only == "all" || only == "firewall") && stagedBbolt != "" && man.BboltSHA256 != "" {
gotHash, hashErr := sha256File(stagedBbolt)
if hashErr != nil {
return nil, fmt.Errorf("hashing staged bbolt snapshot: %w", hashErr)
}
if gotHash != man.BboltSHA256 {
return nil, fmt.Errorf("%w: bbolt snapshot hash mismatch (archive manifest %s, staged %s)", ErrCorruptArchive, man.BboltSHA256, gotHash)
}
}
if (only == "all" || only == "firewall") && stagedBbolt != "" {
if err := DisarmBrowserSessionsSnapshot(stagedBbolt); err != nil {
return nil, err
}
if err := DisarmFirewallRollbackSnapshot(stagedBbolt); err != nil {
return nil, err
}
}
res := &ImportResult{Manifest: man}
if only == "all" {
// Older archives can list session buckets removed during sanitization.
// Read the staged snapshot before applying any files to the destination.
restored, listErr := listSnapshotBuckets(stagedBbolt)
if listErr != nil {
return nil, listErr
}
res.BucketsRestored = restored
}
// Apply state files (always, unless caller filtered everything out).
if only == "all" || only == "baseline" {
if err := os.MkdirAll(opts.StatePath, 0700); err != nil {
return nil, fmt.Errorf("creating state path: %w", err)
}
for _, rel := range stagedState {
src := filepath.Join(stage, rel)
dst := filepath.Join(opts.StatePath, strings.TrimPrefix(rel, stateEntryPrefix))
if err := atomicReplace(src, dst); err != nil {
return nil, fmt.Errorf("restoring %s: %w", rel, err)
}
res.StateFiles++
}
}
// Apply rules files (only=all only; baseline and firewall skip rules).
if only == "all" && opts.RulesPath != "" {
if err := os.MkdirAll(opts.RulesPath, 0700); err != nil {
return nil, fmt.Errorf("creating rules path: %w", err)
}
for _, rel := range stagedRules {
src := filepath.Join(stage, rel)
dst := filepath.Join(opts.RulesPath, strings.TrimPrefix(rel, rulesEntryPrefix))
if err := atomicReplace(src, dst); err != nil {
return nil, fmt.Errorf("restoring %s: %w", rel, err)
}
res.RulesFiles++
}
}
// Apply bbolt:
// only=all -> wholesale rename the snapshot over csm.db
// only=firewall -> open snapshot read-only and copy fw:* buckets
// into the target bbolt
// only=baseline -> skip bbolt entirely
switch only {
case "all":
target := filepath.Join(opts.StatePath, "csm.db")
if err := atomicReplace(stagedBbolt, target); err != nil {
return nil, fmt.Errorf("restoring csm.db: %w", err)
}
case "firewall":
restored, err := mergeBucketsFromSnapshot(stagedBbolt, opts.StatePath, isFirewallBucket)
if err != nil {
return nil, fmt.Errorf("merging firewall buckets: %w", err)
}
// The merge deliberately keeps destination-only firewall keys, but a
// tentative config rollback is process-lifetime state, not durable
// firewall data. Clear any local pending record as well as the one
// removed from the imported snapshot.
if err := DisarmFirewallRollbackSnapshot(filepath.Join(opts.StatePath, "csm.db")); err != nil {
return nil, fmt.Errorf("disarming target firewall rollback: %w", err)
}
res.BucketsRestored = restored
case "baseline":
// no bbolt work
}
return res, nil
}
func isTransientStateArchiveEntry(name string) bool {
rel, ok := strings.CutPrefix(name, stateEntryPrefix)
if !ok {
return false
}
if rel == stateLockFileName || rel == "exports" || strings.HasPrefix(rel, "exports/") {
return true
}
for _, transient := range transientStatePaths {
if rel == transient {
return true
}
}
return false
}
// platformMatches compares the archive's stored platform map against the
// current host. Empty maps match anything (used when a caller hasn't
// supplied detection results, e.g., in some test paths).
func platformMatches(a, b map[string]string) bool {
if len(a) == 0 || len(b) == 0 {
return true
}
keys := []string{"os", "panel", "webserver"}
for _, k := range keys {
if a[k] != b[k] {
return false
}
}
return true
}
// listSnapshotBuckets returns the bucket names actually present in a
// snapshot file (not the static bucketNames slice; migrations may have
// removed some, and export sanitization removes others).
func listSnapshotBuckets(path string) ([]string, error) {
snapshot, err := bolt.Open(path, 0600, &bolt.Options{Timeout: time.Second, ReadOnly: true})
if err != nil {
return nil, fmt.Errorf("opening bbolt snapshot: %w", err)
}
out := []string{}
err = snapshot.View(func(tx *bolt.Tx) error {
return tx.ForEach(func(name []byte, _ *bolt.Bucket) error {
out = append(out, string(name))
return nil
})
})
if closeErr := snapshot.Close(); err == nil && closeErr != nil {
err = fmt.Errorf("closing bbolt snapshot: %w", closeErr)
}
if err != nil {
return nil, fmt.Errorf("listing bbolt snapshot buckets: %w", err)
}
sort.Strings(out)
return out, nil
}
func writeTarFile(tw *tar.Writer, name string, data []byte, modTime time.Time) error {
hdr := &tar.Header{
Name: name,
Mode: 0600,
Size: int64(len(data)),
ModTime: modTime,
}
if err := tw.WriteHeader(hdr); err != nil {
return fmt.Errorf("tar header %s: %w", name, err)
}
if _, err := tw.Write(data); err != nil {
return fmt.Errorf("tar write %s: %w", name, err)
}
return nil
}
// streamFileToTar copies srcPath into the tar under name. There is a
// theoretical Stat-then-Copy race when the source file is replaced
// between f.Stat() and io.Copy() -- the tar writer would then see a
// mismatched byte count. In practice the only files this function
// reads on a live daemon are state JSON written via atomic-replace
// (so a mid-write read sees either the old or the new full content)
// and the bbolt snapshot (already a frozen copy on disk). The risk
// is bounded enough to live with for v1.
// sha256File returns the lowercase hex SHA-256 of a file, matching the
// encoding Export records in Manifest.BboltSHA256.
func sha256File(path string) (string, error) {
// #nosec G304 -- path is the staged bbolt snapshot inside the import
// staging dir (MkdirTemp under StatePath), not attacker-controlled.
f, err := os.Open(path)
if err != nil {
return "", err
}
defer f.Close()
h := sha256.New()
if _, err := io.Copy(h, f); err != nil {
return "", err
}
return hex.EncodeToString(h.Sum(nil)), nil
}
func streamFileToTar(tw *tar.Writer, name, srcPath string, modTime time.Time) error {
// #nosec G304 -- srcPath is supplied by Export's caller (the daemon's control handler, sourced from cfg.StatePath / cfg.Signatures.RulesDir) or is the bbolt snapshot path created via os.CreateTemp earlier in Export. Both are root-controlled; csm runs as root via systemd. Not attacker-controlled.
f, err := os.Open(srcPath)
if err != nil {
return fmt.Errorf("open %s: %w", srcPath, err)
}
defer f.Close()
info, err := f.Stat()
if err != nil {
return fmt.Errorf("stat %s: %w", srcPath, err)
}
hdr := &tar.Header{
Name: name,
Mode: 0600,
Size: info.Size(),
ModTime: modTime,
}
if err := tw.WriteHeader(hdr); err != nil {
return fmt.Errorf("tar header %s: %w", name, err)
}
if _, err := io.Copy(tw, f); err != nil {
return fmt.Errorf("tar copy %s: %w", name, err)
}
return nil
}
// walkDirIntoTar streams every regular file under srcDir into the tar
// under entryPrefix, skipping relative paths in the skip list. A skipped
// directory excludes its complete tree. Returns
// how many files were written.
func walkDirIntoTar(tw *tar.Writer, srcDir, entryPrefix string, skip []string, modTime time.Time) (int, error) {
skipSet := map[string]bool{}
for _, s := range skip {
skipSet[s] = true
}
count := 0
err := filepath.Walk(srcDir, func(path string, info os.FileInfo, err error) error {
if err != nil {
return err
}
rel, err := filepath.Rel(srcDir, path)
if err != nil {
return err
}
rel = filepath.ToSlash(rel)
if skipSet[rel] {
if info.IsDir() {
return filepath.SkipDir
}
return nil
}
if info.IsDir() {
return nil
}
entry := entryPrefix + filepath.ToSlash(rel)
if err := streamFileToTar(tw, entry, path, modTime); err != nil {
return err
}
count++
return nil
})
if err != nil && !os.IsNotExist(err) {
return count, err
}
return count, nil
}
// atomicReplace renames src over dst, ensuring the parent directory is
// fsync'd so the rename survives a crash.
func atomicReplace(src, dst string) error {
if err := os.MkdirAll(filepath.Dir(dst), 0700); err != nil {
return err
}
if err := os.Rename(src, dst); err != nil {
return err
}
parent, err := os.Open(filepath.Dir(dst))
if err != nil {
return err
}
defer parent.Close()
return parent.Sync()
}
// mergeBucketsFromSnapshot opens the snapshot bbolt read-only, iterates
// matching buckets, and copies their key/value pairs into the target
// bbolt at statePath/csm.db. The target may or may not exist; Open
// creates it. Returns the bucket names actually merged.
func mergeBucketsFromSnapshot(snapshotPath, statePath string, match func(string) bool) ([]string, error) {
src, err := bolt.Open(snapshotPath, 0600, &bolt.Options{Timeout: 5 * time.Second, ReadOnly: true})
if err != nil {
return nil, fmt.Errorf("opening snapshot: %w", err)
}
defer func() { _ = src.Close() }()
dst, err := Open(statePath)
if err != nil {
return nil, fmt.Errorf("opening target: %w", err)
}
defer func() { _ = dst.Close() }()
merged := []string{}
err = src.View(func(stx *bolt.Tx) error {
return stx.ForEach(func(name []byte, sb *bolt.Bucket) error {
if !match(string(name)) {
return nil
}
if upErr := dst.bolt.Update(func(dtx *bolt.Tx) error {
db, berr := dtx.CreateBucketIfNotExists(name)
if berr != nil {
return berr
}
return sb.ForEach(func(k, v []byte) error {
return db.Put(append([]byte(nil), k...), append([]byte(nil), v...))
})
}); upErr != nil {
return upErr
}
merged = append(merged, string(name))
return nil
})
})
if err != nil {
return merged, err
}
sort.Strings(merged)
return merged, nil
}
func isFirewallBucket(name string) bool {
return strings.HasPrefix(name, "fw:")
}
package store
import (
"bytes"
"encoding/json"
"fmt"
"time"
bolt "go.etcd.io/bbolt"
)
// maxAttackEvents is the maximum number of attack events to retain.
// It is a var (not const) so tests can override it.
var maxAttackEvents = 100_000
// AttackEvent is the store-layer representation of an attack event.
type AttackEvent struct {
Timestamp time.Time `json:"timestamp"`
IP string `json:"ip"`
AttackType string `json:"attack_type"`
CheckName string `json:"check_name"`
Severity int `json:"severity"`
Account string `json:"account,omitempty"`
Message string `json:"message,omitempty"`
}
// IPRecord is the store-layer representation of an IP attack record.
type IPRecord struct {
IP string `json:"ip"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
EventCount int `json:"event_count"`
AttackCounts map[string]int `json:"attack_counts,omitempty"`
Accounts map[string]int `json:"accounts,omitempty"`
AuthSuccessAccounts map[string]int `json:"auth_success_accounts,omitempty"`
ThreatScore int `json:"threat_score"`
AutoBlocked bool `json:"auto_blocked,omitempty"`
BruteForceWindowStart time.Time `json:"brute_force_window_start,omitempty"`
BruteForceWindowCount int `json:"brute_force_window_count,omitempty"`
BruteForceSustainedAt time.Time `json:"brute_force_sustained_at,omitempty"`
}
// RecordAttackEvent inserts an attack event into both the primary bucket
// (attacks:events, keyed by TimeKey) and the secondary index bucket
// (attacks:events:ip, keyed by IP/TimeKey). It increments the event counter
// and prunes oldest entries if the count exceeds maxAttackEvents.
func (db *DB) RecordAttackEvent(event AttackEvent, counter int) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
primary := tx.Bucket([]byte("attacks:events"))
writer := newTimeKeyWriter(primary)
secondary := tx.Bucket([]byte("attacks:events:ip"))
key := TimeKey(event.Timestamp, counter)
val, err := json.Marshal(event)
if err != nil {
return err
}
// Callers restart the counter for each batch. Reserve an unused
// primary key in this write transaction so an index never loses its event.
for primary.Get([]byte(key)) != nil {
counter++
key = TimeKey(event.Timestamp, counter)
}
if err := writer.put([]byte(key), val); err != nil {
return err
}
writer.settle()
// The index only needs the key; the event itself lives in the
// primary bucket.
secondaryKey := event.IP + "/" + key
if err := secondary.Put([]byte(secondaryKey), []byte{}); err != nil {
return err
}
if err := incrCounter(tx, "attacks:events:count", 1); err != nil {
return err
}
// Prune oldest entries if count exceeds maxAttackEvents.
meta := tx.Bucket([]byte("meta"))
var count int
if v := meta.Get([]byte("attacks:events:count")); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", &count)
}
if count > maxAttackEvents {
excess := count - maxAttackEvents
c := primary.Cursor()
k, v := c.First()
for ; k != nil && excess > 0; excess-- {
// Unmarshal to get the IP for secondary index cleanup.
var ev AttackEvent
if err := json.Unmarshal(v, &ev); err != nil {
return err
}
secKey := ev.IP + "/" + string(k)
if err := secondary.Delete([]byte(secKey)); err != nil {
return err
}
if err := c.Delete(); err != nil {
return err
}
// Re-seek after delete (bbolt cursor behavior).
k, v = c.First()
}
if err := setCounter(tx, "attacks:events:count", maxAttackEvents); err != nil {
return err
}
}
return nil
})
}
// QueryAttackEvents returns up to limit attack events for the given IP,
// newest-first. It walks the secondary index backwards from the end of the
// IP's key range and resolves each entry in the primary bucket.
func (db *DB) QueryAttackEvents(ip string, limit int) []AttackEvent {
var results []AttackEvent
prefix := []byte(ip + "/")
_ = db.bolt.View(func(tx *bolt.Tx) error {
primary := tx.Bucket([]byte("attacks:events"))
c := tx.Bucket([]byte("attacks:events:ip")).Cursor()
var v []byte
k, _ := c.Seek(append(append([]byte(nil), prefix...), 0xff))
if k == nil {
k, v = c.Last()
} else {
k, v = c.Prev()
}
for ; k != nil && bytes.HasPrefix(k, prefix) && len(results) < limit; k, v = c.Prev() {
if ev, ok := resolveIndexedAttackEvent(primary, ip, k[len(prefix):], v); ok {
results = append(results, ev)
}
}
return nil
})
return results
}
// resolveIndexedAttackEvent reads the event an index entry points at. Index
// entries written by earlier builds carry their own copy of the event, which
// is used when the primary row is gone or now holds another address's event.
func resolveIndexedAttackEvent(primary *bolt.Bucket, ip string, timeKey, indexValue []byte) (AttackEvent, bool) {
if raw := primary.Get(timeKey); raw != nil {
var ev AttackEvent
if json.Unmarshal(raw, &ev) == nil && ev.IP == ip {
return ev, true
}
}
if len(indexValue) > 0 {
var ev AttackEvent
if json.Unmarshal(indexValue, &ev) == nil && ev.IP == ip {
return ev, true
}
}
return AttackEvent{}, false
}
// SaveIPRecord stores an IP record in the attacks:records bucket, keyed by IP.
func (db *DB) SaveIPRecord(record IPRecord) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("attacks:records"))
val, err := json.Marshal(record)
if err != nil {
return err
}
return b.Put([]byte(record.IP), val)
})
}
// LoadIPRecord retrieves an IP record from the attacks:records bucket.
// Returns the record and true if found, or a zero value and false if not.
func (db *DB) LoadIPRecord(ip string) (IPRecord, bool) {
var record IPRecord
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("attacks:records"))
v := b.Get([]byte(ip))
if v == nil {
return nil
}
if json.Unmarshal(v, &record) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
found = true
return nil
})
return record, found
}
// LoadAllIPRecords returns all IP records from the attacks:records bucket.
func (db *DB) LoadAllIPRecords() map[string]*IPRecord {
records := make(map[string]*IPRecord)
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("attacks:records"))
return b.ForEach(func(k, v []byte) error {
var record IPRecord
if json.Unmarshal(v, &record) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
records[string(k)] = &record
return nil
})
})
return records
}
// DeleteIPRecord removes an IP record from the attacks:records bucket.
func (db *DB) DeleteIPRecord(ip string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("attacks:records"))
return b.Delete([]byte(ip))
})
}
// ReadAllAttackEvents returns all attack events from the primary bucket.
// Used for stats computation (hourly/daily bucketing).
func (db *DB) ReadAllAttackEvents() []AttackEvent {
var events []AttackEvent
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("attacks:events"))
return b.ForEach(func(k, v []byte) error {
var ev AttackEvent
if json.Unmarshal(v, &ev) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
events = append(events, ev)
return nil
})
})
return events
}
package store
import (
"bytes"
"encoding/binary"
"net"
"time"
bolt "go.etcd.io/bbolt"
)
// botVerifyEntry layout:
//
// byte 0 verified flag (1 byte)
// bytes 1..8 expiry as unix nanos (8 bytes big-endian)
//
// Key format:
//
// <bot-bytes> 0x00 <ip-bytes>
//
// IP is stored in its 16-byte form (IPv4 promoted via To16). The
// 0x00 separator allows the bot name to be any non-null string while
// keeping keys byte-sortable within a single bucket.
func botVerifyKey(ip net.IP, bot string) []byte {
ipBytes := ip.To16()
if ipBytes == nil {
return nil
}
out := make([]byte, 0, len(bot)+1+16)
out = append(out, []byte(bot)...)
out = append(out, 0x00)
out = append(out, ipBytes...)
return out
}
// PutBotVerify stores a PTR+forward-A verification result with an
// explicit expiry. A verified=false entry means the IP failed rDNS
// and will emit http_ua_spoof on the next scan that sees it with the
// same bot UA.
func (db *DB) PutBotVerify(ip net.IP, bot string, verified bool, expiresAt time.Time) error {
key := botVerifyKey(ip, bot)
if key == nil {
return nil
}
var val [9]byte
if verified {
val[0] = 1
}
binary.BigEndian.PutUint64(val[1:], uint64(expiresAt.UnixNano())) // #nosec G115 -- unix nano stored as bit pattern; sign is irrelevant for expiry comparison
return db.bolt.Update(func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists([]byte("botverify"))
if err != nil {
return err
}
if records := tx.Bucket([]byte(botVerifyUnverifiableBucket)); records != nil {
if err := records.Delete(key); err != nil {
return err
}
}
return b.Put(key, val[:])
})
}
// botVerifyUnverifiableBucket holds sources whose claimed bot identity had no
// PTR record, keyed like the verdict cache. The value is the observation time
// as unix nanos; the verifier derives retry suppression and history from it
// and sweeps records that can no longer matter. It lives apart from the
// verdict bucket because a missing PTR proves nothing about identity, and a
// build that predates this bucket must never read it as a spoof verdict.
const botVerifyUnverifiableBucket = "botverify_unverifiable"
// PutBotVerifyUnverifiable records that ip had no PTR for bot at observedAt.
func (db *DB) PutBotVerifyUnverifiable(ip net.IP, bot string, observedAt time.Time) error {
key := botVerifyKey(ip, bot)
if key == nil {
return nil
}
var val [8]byte
binary.BigEndian.PutUint64(val[:], uint64(observedAt.UnixNano())) // #nosec G115 -- unix nano stored as bit pattern; sign is irrelevant for time comparison
return db.bolt.Update(func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists([]byte(botVerifyUnverifiableBucket))
if err != nil {
return err
}
return b.Put(key, val[:])
})
}
// BotVerifyUnverifiable returns when ip was recorded without a PTR for bot.
// Reads never write, so they cannot extend a record or race a bucket reset.
func (db *DB) BotVerifyUnverifiable(ip net.IP, bot string) (observedAt time.Time, ok bool) {
key := botVerifyKey(ip, bot)
if key == nil {
return time.Time{}, false
}
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(botVerifyUnverifiableBucket))
if b == nil {
return nil
}
if val := b.Get(key); len(val) == 8 {
observedAt, ok = unixNanoBits(val), true
}
return nil
})
return observedAt, ok
}
// SweepBotVerifyUnverifiable removes records observed before cutoff and
// returns how many it removed.
func (db *DB) SweepBotVerifyUnverifiable(cutoff time.Time) (int, error) {
removed := 0
err := db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(botVerifyUnverifiableBucket))
if b == nil {
return nil
}
var stale [][]byte
if err := b.ForEach(func(k, v []byte) error {
if len(v) != 8 || unixNanoBits(v).Before(cutoff) {
stale = append(stale, append([]byte(nil), k...))
}
return nil
}); err != nil {
return err
}
for _, k := range stale {
if err := b.Delete(k); err != nil {
return err
}
}
removed = len(stale)
return nil
})
if err != nil {
return 0, err
}
return removed, nil
}
func unixNanoBits(val []byte) time.Time {
return time.Unix(0, int64(binary.BigEndian.Uint64(val))) // #nosec G115 -- reinterpret stored bit pattern as signed nanos
}
// EnsureBotVerifyLogicVersion compares the stored cache logic version
// with version and, on mismatch (or when no marker exists yet), drops
// the botverify bucket and its missing-PTR records, then records the new
// version. The marker lives in the "meta" bucket under
// botverify:logic_version. Returns true when the bucket was dropped.
//
// Use this from daemon startup so that any change to the verifier
// logic (BotDomains suffix list, ClaimedBotFromUA mapping, etc.)
// automatically invalidates entries written under the old rules.
// Operators do not need to know about the cache.
func (db *DB) EnsureBotVerifyLogicVersion(version int) (bool, error) {
var dropped bool
err := db.bolt.Update(func(tx *bolt.Tx) error {
meta, mErr := tx.CreateBucketIfNotExists([]byte("meta"))
if mErr != nil {
return mErr
}
current := uint64(version) // #nosec G115 -- logic version is a small positive internal constant
stored := ^uint64(0)
if raw := meta.Get([]byte("botverify:logic_version")); len(raw) == 8 {
stored = binary.BigEndian.Uint64(raw)
}
if stored == current {
return nil
}
if _, dErr := resetBotVerifyBuckets(tx); dErr != nil {
return dErr
}
var buf [8]byte
binary.BigEndian.PutUint64(buf[:], current)
if pErr := meta.Put([]byte("botverify:logic_version"), buf[:]); pErr != nil {
return pErr
}
dropped = true
return nil
})
if err != nil {
return false, err
}
return dropped, nil
}
// ResetBotVerify drops every cached PTR+forward-A result and no-PTR record.
// Returns the number of entries cleared. Use after a verifier-logic upgrade that
// would invalidate prior negative cache entries (e.g., a domain suffix
// fix that turns prior false-spoof entries into positives). Safe to
// call when the bucket is missing or empty.
func (db *DB) ResetBotVerify() (int, error) {
var cleared int
err := db.bolt.Update(func(tx *bolt.Tx) error {
var err error
cleared, err = resetBotVerifyBuckets(tx)
return err
})
if err != nil {
return 0, err
}
return cleared, nil
}
// resetBotVerifyBuckets empties the verdict cache and its no-PTR records
// together, so no retry suppression outlives the rules that produced it.
func resetBotVerifyBuckets(tx *bolt.Tx) (int, error) {
cleared := 0
for _, name := range []string{"botverify", botVerifyUnverifiableBucket} {
if b := tx.Bucket([]byte(name)); b != nil {
cleared += b.Stats().KeyN
if err := tx.DeleteBucket([]byte(name)); err != nil {
return 0, err
}
}
}
for _, name := range []string{"botverify", botVerifyUnverifiableBucket} {
if _, err := tx.CreateBucket([]byte(name)); err != nil {
return 0, err
}
}
return cleared, nil
}
// GetBotVerify returns (verified, valid). valid=false means no
// non-expired entry exists; the caller should treat the IP as
// unverified and (optionally) enqueue an async verify job.
func (db *DB) GetBotVerify(ip net.IP, bot string) (verified, valid bool) {
key := botVerifyKey(ip, bot)
if key == nil {
return false, false
}
var stored []byte
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("botverify"))
if b == nil {
return nil
}
val := b.Get(key)
if len(val) != 9 {
return nil
}
stored = append([]byte(nil), val...)
exp := time.Unix(0, int64(binary.BigEndian.Uint64(val[1:]))) // #nosec G115 -- reinterpret stored bit pattern as signed nanos
if time.Now().After(exp) {
return nil
}
verified = val[0] == 1
valid = true
return nil
})
if valid || len(stored) != 9 {
return verified, valid
}
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("botverify"))
current := b.Get(key)
if !bytes.Equal(current, stored) {
return nil
}
exp := time.Unix(0, int64(binary.BigEndian.Uint64(current[1:]))) // #nosec G115 -- reinterpret stored bit pattern as signed nanos
if time.Now().After(exp) {
return b.Delete(key)
}
return nil
})
return verified, valid
}
package store
import (
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"sort"
"time"
"github.com/pidginhost/csm/internal/session"
bolt "go.etcd.io/bbolt"
)
const browserSessionsBucket = "browser_sessions"
func (db *DB) ReplaceBrowserSession(rec session.Record, previous string, now time.Time, idle time.Duration) error {
if !rec.Valid(now, idle) {
return session.ErrInvalid
}
raw, err := json.Marshal(rec)
if err != nil {
return err
}
return boltUpdate(db.bolt, func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists([]byte(browserSessionsBucket))
if err != nil {
return err
}
if previous != "" {
if _, readErr := readBrowserSession(tx, previous, now, idle); readErr != nil {
return readErr
}
}
c := b.Cursor()
count := 0
for k, v := c.First(); k != nil; k, v = c.Next() {
var old session.Record
if err := json.Unmarshal(v, &old); err != nil {
return fmt.Errorf("decode browser session: %w", err)
}
if string(k) == previous || !old.Valid(now, idle) {
if err := c.Delete(); err != nil {
return err
}
continue
}
count++
}
if count >= session.MaxSessions {
return session.ErrFull
}
if b.Get([]byte(rec.Verifier)) != nil {
return errors.New("browser session identity collision")
}
return b.Put([]byte(rec.Verifier), raw)
})
}
func readBrowserSession(tx *bolt.Tx, key string, now time.Time, idle time.Duration) (session.Record, error) {
b := tx.Bucket([]byte(browserSessionsBucket))
if b == nil {
return session.Record{}, session.ErrInvalid
}
raw := b.Get([]byte(key))
if raw == nil {
return session.Record{}, session.ErrInvalid
}
var rec session.Record
if err := json.Unmarshal(raw, &rec); err != nil {
return session.Record{}, fmt.Errorf("decode browser session: %w", err)
}
if rec.Verifier != key || !rec.Valid(now, idle) {
return session.Record{}, session.ErrInvalid
}
return rec, nil
}
func (db *DB) AccessBrowserSession(key string, now time.Time, idle time.Duration, touch bool) (session.Record, error) {
var rec session.Record
err := db.bolt.View(func(tx *bolt.Tx) error {
var readErr error
rec, readErr = readBrowserSession(tx, key, now, idle)
return readErr
})
if err != nil || !touch {
return rec, err
}
// Bound write amplification from dashboard polling. Expiry remains based on
// the last committed activity, so this can only shorten the idle window.
interval := min(30*time.Second, idle/4)
if now.Sub(rec.LastSeen) < interval {
return rec, nil
}
err = boltUpdate(db.bolt, func(tx *bolt.Tx) error {
current, readErr := readBrowserSession(tx, key, now, idle)
if readErr != nil {
return readErr
}
// Another request may have touched the session since our read.
// Preserve its newer activity and the existing write interval.
if now.Sub(current.LastSeen) < interval {
rec = current
return nil
}
current.LastSeen = now
raw, marshalErr := json.Marshal(current)
if marshalErr != nil {
return marshalErr
}
if putErr := tx.Bucket([]byte(browserSessionsBucket)).Put([]byte(key), raw); putErr != nil {
return putErr
}
rec = current
return nil
})
if err != nil {
return session.Record{}, err
}
return rec, nil
}
func (db *DB) ListBrowserSessions(now time.Time, idle time.Duration) ([]session.Record, error) {
records := make([]session.Record, 0)
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(browserSessionsBucket))
if b == nil {
return nil
}
return b.ForEach(func(k, v []byte) error {
var rec session.Record
if err := json.Unmarshal(v, &rec); err != nil {
return fmt.Errorf("decode browser session: %w", err)
}
if rec.Verifier != string(k) {
return errors.New("browser session identity mismatch")
}
if rec.Valid(now, idle) {
records = append(records, rec)
}
return nil
})
})
if err != nil {
return nil, err
}
sort.Slice(records, func(i, j int) bool { return records[i].Created.After(records[j].Created) })
return records, nil
}
func (db *DB) RevokeBrowserSession(id string) error {
return boltUpdate(db.bolt, func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(browserSessionsBucket))
if b == nil {
return nil
}
c := b.Cursor()
for k, v := c.First(); k != nil; k, v = c.Next() {
var rec session.Record
if err := json.Unmarshal(v, &rec); err != nil {
return fmt.Errorf("decode browser session: %w", err)
}
if rec.ID == id {
return c.Delete()
}
}
return nil
})
}
func (db *DB) ClearBrowserSessions() error {
return boltUpdate(db.bolt, func(tx *bolt.Tx) error {
if tx.Bucket([]byte(browserSessionsBucket)) == nil {
return nil
}
return tx.DeleteBucket([]byte(browserSessionsBucket))
})
}
// DisarmBrowserSessionsSnapshot strips verifiers from a private, stopped
// snapshot. It must never be called on the daemon's live database.
func DisarmBrowserSessionsSnapshot(path string) error {
file, err := os.CreateTemp(filepath.Dir(path), ".session-free-*")
if err != nil {
return err
}
cleanPath := file.Name()
defer func() { _ = os.Remove(cleanPath) }()
if err = file.Close(); err != nil {
return err
}
err = func() (result error) {
snapshot, openErr := bolt.Open(path, 0600, &bolt.Options{Timeout: time.Second})
if openErr != nil {
return openErr
}
defer func() { result = errors.Join(result, snapshot.Close()) }()
db := &DB{bolt: snapshot, path: path}
if clearErr := db.ClearBrowserSessions(); clearErr != nil {
return clearErr
}
// Compact even if the bucket was already absent: earlier revocations
// may have left metadata in free pages. Only live records are copied.
_, _, compactErr := db.CompactInto(cleanPath, 16*1024*1024)
return compactErr
}()
if err != nil {
return err
}
return os.Rename(cleanPath, path)
}
package store
import bolt "go.etcd.io/bbolt"
const contentLogicVersionKey = "content:logic_version"
// ContentLogicVersionChanged reports whether token differs from the last
// completed finding re-verification sweep. Reading and committing the marker
// are separate so a crash or shutdown during the sweep causes a retry.
func (db *DB) ContentLogicVersionChanged(token string) (bool, error) {
var changed bool
err := db.bolt.View(func(tx *bolt.Tx) error {
meta := tx.Bucket([]byte("meta"))
if meta == nil {
changed = true
return nil
}
changed = string(meta.Get([]byte(contentLogicVersionKey))) != token
return nil
})
if err != nil {
return false, err
}
return changed, nil
}
// SetContentLogicVersion records token after a finding re-verification sweep
// completes, preventing another run until one of its logic versions changes.
func (db *DB) SetContentLogicVersion(token string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
meta, err := tx.CreateBucketIfNotExists([]byte("meta"))
if err != nil {
return err
}
return meta.Put([]byte(contentLogicVersionKey), []byte(token))
})
}
package store
import (
"bytes"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"time"
bolt "go.etcd.io/bbolt"
)
// Bucket names - all buckets are created on Open().
var bucketNames = []string{
"history",
"attacks:records",
"attacks:events",
"attacks:events:ip",
"threats",
"threats:whitelist",
"fw:blocked",
"fw:allowed",
"fw:subnets",
"fw:port_allowed",
"reputation",
mailGoodSourceBucket,
"plugins",
"plugins:sites",
"meta",
"email:geo",
"email:fwd",
"db_object_backups",
"sig_watch",
bucketStatsDaily,
bucketLatestByCheck,
"phprelay:meta",
"phprelay:msgindex",
"phprelay:ignore",
"phprelay:settings",
"incidents",
"fw:rollback",
adminEmailsBucket,
"botverify",
botVerifyUnverifiableBucket,
prefsBucket,
"scan_jobs",
"scan_job_findings",
scanCursorBucket,
}
// DB wraps a bbolt database.
type DB struct {
bolt *bolt.DB
path string
}
var (
globalDB *DB
globalMu sync.Mutex
ensureOnce sync.Once
ensureErr error
)
// Global returns the singleton DB instance.
func Global() *DB {
globalMu.Lock()
defer globalMu.Unlock()
return globalDB
}
// SetGlobal sets the singleton DB instance.
func SetGlobal(db *DB) {
globalMu.Lock()
globalDB = db
globalMu.Unlock()
}
// WriteTxID returns the id of the most recently committed write
// transaction. bbolt bumps the id once per committed write transaction,
// so two readings bracket a code path and their difference counts its
// write commits.
func (db *DB) WriteTxID() int {
var id int
_ = db.bolt.View(func(tx *bolt.Tx) error {
id = tx.ID()
return nil
})
return id
}
// EnsureOpen opens the store if not already open. Safe to call from any CLI path.
// First call opens the DB; subsequent calls return immediately.
func EnsureOpen(statePath string) error {
ensureOnce.Do(func() {
db, err := Open(statePath)
if err != nil {
ensureErr = err
return
}
SetGlobal(db)
})
return ensureErr
}
// Open opens or creates the bbolt database at {statePath}/csm.db.
// Creates all buckets if they don't exist. Runs migration if needed.
func Open(statePath string) (*DB, error) {
dbPath := filepath.Join(statePath, "csm.db")
if err := os.MkdirAll(statePath, 0700); err != nil {
return nil, fmt.Errorf("creating state dir: %w", err)
}
bdb, err := bolt.Open(dbPath, 0600, &bolt.Options{Timeout: 5 * time.Second})
if err != nil {
return nil, fmt.Errorf("opening bbolt: %w", err)
}
// Create all buckets
err = bdb.Update(func(tx *bolt.Tx) error {
// No time-index buckets means this database cannot contain legacy
// wall-clock keys. Mark it canonical in the same transaction that
// creates the buckets, avoiding an extra startup write on new stores.
freshTimeKeyStore := tx.Bucket([]byte("history")) == nil &&
tx.Bucket([]byte("attacks:events")) == nil &&
tx.Bucket([]byte("attacks:events:ip")) == nil
for _, name := range bucketNames {
if _, berr := tx.CreateBucketIfNotExists([]byte(name)); berr != nil {
return fmt.Errorf("creating bucket %s: %w", name, berr)
}
}
// Initialise phprelay schema_version on first open. Stored as
// 8-byte big-endian uint64 so future migrations can read/compare
// it consistently.
meta := tx.Bucket([]byte("phprelay:meta"))
if meta.Get([]byte("schema_version")) == nil {
if perr := meta.Put([]byte("schema_version"), []byte{0, 0, 0, 0, 0, 0, 0, 1}); perr != nil {
return fmt.Errorf("init phprelay schema_version: %w", perr)
}
}
if freshTimeKeyStore {
if markerErr := tx.Bucket([]byte("meta")).Put(
[]byte(timeKeyUTCMarker), []byte(timeKeyMigrationDone),
); markerErr != nil {
return fmt.Errorf("init UTC time-key marker: %w", markerErr)
}
}
return nil
})
if err != nil {
_ = bdb.Close()
return nil, err
}
db := &DB{bolt: bdb, path: dbPath}
// Run migration if needed. A failed migration does not set the "migrated"
// sentinel, so proceeding would boot the daemon on partial security state
// and retry the same broken migration every restart. Fail loud instead.
if err := db.migrateIfNeeded(statePath); err != nil {
_ = bdb.Close()
return nil, fmt.Errorf("store migration: %w", err)
}
// Older releases formatted keys in whatever zone each producer supplied.
// Canonicalize those persisted keys before any key-based reads or pruning.
if err := db.migrateTimeKeysToUTC(); err != nil {
_ = bdb.Close()
return nil, fmt.Errorf("time-key migration: %w", err)
}
// One-time backfill of stats:daily from existing history. Runs on
// hosts upgrading from a build that pre-dates the stats:daily bucket;
// no-op afterwards thanks to a meta sentinel.
if err := db.BackfillStatsDaily(); err != nil {
fmt.Fprintf(os.Stderr, "store: stats:daily backfill warning: %v\n", err)
}
if err := db.BackfillLatestByCheck(); err != nil {
fmt.Fprintf(os.Stderr, "store: stats:latest_by_check backfill warning: %v\n", err)
}
if err := db.seedDefaultModSecNoEscalateRules(); err != nil {
fmt.Fprintf(os.Stderr, "store: ModSecurity no-escalate seed warning: %v\n", err)
}
return db, nil
}
const (
modsecNoEscalateSeededKey = "modsec:no_escalate_seeded"
defaultModSecNoEscalateWPEnumerationID = 900112
)
func (db *DB) seedDefaultModSecNoEscalateRules() error {
return db.bolt.Update(func(tx *bolt.Tx) error {
meta := tx.Bucket([]byte("meta"))
if meta.Get([]byte(modsecNoEscalateSeededKey)) != nil {
return nil
}
if meta.Get([]byte(modsecNoEscalateKey)) != nil {
return meta.Put([]byte(modsecNoEscalateSeededKey), []byte("1"))
}
// WordPress user enumeration is blocked at the HTTP layer only.
val, err := json.Marshal([]int{defaultModSecNoEscalateWPEnumerationID})
if err != nil {
return err
}
if err := meta.Put([]byte(modsecNoEscalateKey), val); err != nil {
return err
}
return meta.Put([]byte(modsecNoEscalateSeededKey), []byte("1"))
})
}
// Close closes the bbolt database.
func (db *DB) Close() error {
if db.bolt == nil {
return nil
}
return db.bolt.Close()
}
// Path returns the on-disk path of the bbolt database file.
func (db *DB) Path() string {
return db.path
}
// HasBucket reports whether a top-level bucket named name exists in db.
func (db *DB) HasBucket(name string) bool {
found := false
_ = db.bolt.View(func(tx *bolt.Tx) error {
if tx.Bucket([]byte(name)) != nil {
found = true
}
return nil
})
return found
}
// timeKeyFillPercent packs mostly ordered writes to buckets keyed by TimeKey.
// Their split tail pages are rarely written again, so bbolt's default of 0.5
// would leave them half empty.
const timeKeyFillPercent = 0.9
// timeKeyWriter picks the split point for one write transaction. bbolt applies
// a bucket's FillPercent to every page it splits at commit, so the choice is
// per transaction rather than per key. A few delayed keys cost less than half
// empty tail pages; when most keys land before the stored tail, the balanced
// default leaves room in the pages they split. Also balance when delayed
// records carry at least half the inserted bytes: a minority of large records
// can otherwise leave sparse pages throughout the backfill.
type timeKeyWriter struct {
bucket *bolt.Bucket
tail []byte
older int
total int
olderBytes int
totalBytes int
}
func newTimeKeyWriter(b *bolt.Bucket) *timeKeyWriter {
tail, _ := b.Cursor().Last()
return &timeKeyWriter{bucket: b, tail: bytes.Clone(tail)}
}
func (w *timeKeyWriter) put(key, value []byte) error {
size := len(key) + len(value)
w.total++
w.totalBytes += size
if w.tail != nil && bytes.Compare(key, w.tail) <= 0 {
w.older++
w.olderBytes += size
}
return w.bucket.Put(key, value)
}
// settle sets the split point once every key of the transaction is known. It
// must run before the transaction commits.
func (w *timeKeyWriter) settle() {
w.bucket.FillPercent = timeKeyFillPercent
if w.older*2 > w.total || (w.olderBytes > 0 && w.olderBytes >= w.totalBytes-w.olderBytes) {
w.bucket.FillPercent = bolt.DefaultFillPercent
}
}
// TimeKey produces a fixed-width 28-byte key for chronological ordering.
// Format: YYYYMMDDHHmmssNNNNNNNNN-CCCC
// Lexicographic order equals chronological order.
//
// Keys use UTC so ordering is stable across producer zones, host timezone
// changes, and repeated local hours during daylight-saving transitions.
func TimeKey(t time.Time, counter int) string {
t = t.UTC()
return fmt.Sprintf("%04d%02d%02d%02d%02d%02d%09d-%04d",
t.Year(), t.Month(), t.Day(),
t.Hour(), t.Minute(), t.Second(),
t.Nanosecond(), counter)
}
// ParseTimeKeyPrefix converts a date string "YYYY-MM-DD" to a seek prefix "YYYYMMDD".
func ParseTimeKeyPrefix(date string) string {
if len(date) == 10 && date[4] == '-' && date[7] == '-' {
return date[:4] + date[5:7] + date[8:10]
}
return date
}
// getCounter reads a counter from the meta bucket. Returns 0 if not found.
// HistoryCount returns the number of findings in the history bucket.
func (db *DB) HistoryCount() int {
return db.getCounter("history:count")
}
func (db *DB) getCounter(key string) int {
var count int
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
if v := b.Get([]byte(key)); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", &count)
}
return nil
})
return count
}
// setCounter writes a counter to the meta bucket within an existing transaction.
func setCounter(tx *bolt.Tx, key string, count int) error {
b := tx.Bucket([]byte("meta"))
return b.Put([]byte(key), []byte(fmt.Sprintf("%d", count)))
}
// IsHealthy returns true if the bbolt file is open and all required buckets exist.
func (db *DB) IsHealthy() bool {
if db == nil || db.bolt == nil {
return false
}
err := db.bolt.View(func(tx *bolt.Tx) error {
for _, name := range []string{"history", "fw:blocked", "meta"} {
if tx.Bucket([]byte(name)) == nil {
return fmt.Errorf("bucket missing: %s", name)
}
}
return nil
})
return err == nil
}
// SizeBytes returns the on-disk size of the bbolt database file. Returns 0 if unavailable.
func (db *DB) SizeBytes() int64 {
if db == nil || db.path == "" {
return 0
}
info, err := os.Stat(db.path)
if err != nil {
return 0
}
return info.Size()
}
// incrCounter increments a counter within an existing transaction.
func incrCounter(tx *bolt.Tx, key string, delta int) error {
b := tx.Bucket([]byte("meta"))
var current int
if v := b.Get([]byte(key)); v != nil {
fmt.Sscanf(string(v), "%d", ¤t)
}
return b.Put([]byte(key), []byte(fmt.Sprintf("%d", current+delta)))
}
type dryRunBlockRecord struct {
IP string `json:"ip"`
Reason string `json:"reason"`
TimeoutSec int `json:"timeout_sec"`
}
// RecordDryRunBlock appends a dry-run-block record to the "dry_run_blocks"
// bucket. Called by the firewall engine when auto_response.dry_run is active
// so operators can review "what would have been blocked" before going live.
func (db *DB) RecordDryRunBlock(ip, reason string, timeout time.Duration) {
if db == nil || db.bolt == nil {
return
}
// Log-derived reasons can carry raw control bytes; the JSON encoder
// emits escape forms that dry-run readers can decode.
payload := dryRunBlockRecord{
IP: ip,
Reason: reason,
TimeoutSec: int(timeout.Seconds()),
}
val, err := json.Marshal(payload)
if err != nil {
return
}
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists([]byte("dry_run_blocks"))
if err != nil {
return err
}
key := []byte(time.Now().UTC().Format(time.RFC3339Nano) + ":" + ip)
return b.Put(key, val)
})
}
// PurgeAllDryRunBlocks deletes every record from the dry_run_blocks
// bucket and returns the number removed. Called when the operator
// flips auto_response.dry_run from true to false so /api/v1/status no
// longer reports a stale count from the previous dry-run window. A
// later periodic prune handles the slow accumulation case via
// PurgeDryRunBlocksOlderThan.
func (db *DB) PurgeAllDryRunBlocks() int {
if db == nil || db.bolt == nil {
return 0
}
removed := 0
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("dry_run_blocks"))
if b == nil {
return nil
}
removed = b.Stats().KeyN
return tx.DeleteBucket([]byte("dry_run_blocks"))
})
return removed
}
// PurgeDryRunBlocksOlderThan removes every dry_run_blocks record
// whose timestamp prefix is strictly older than cutoff. Returns the
// number removed. Key format is "<RFC3339Nano>:<ip>"; entries with a
// key that does not parse as a timestamp are left in place so a
// future key-format change does not silently drop records.
func (db *DB) PurgeDryRunBlocksOlderThan(cutoff time.Time) int {
if db == nil || db.bolt == nil {
return 0
}
removed := 0
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("dry_run_blocks"))
if b == nil {
return nil
}
var stale [][]byte
_ = b.ForEach(func(k, _ []byte) error {
s := string(k)
// RecordDryRunBlock writes UTC timestamps which always
// end in `Z`, so the first `Z:` reliably separates the
// timestamp from the IP (the IP itself can contain
// colons in v6 form, ruling out a naive first-colon
// split).
idx := strings.Index(s, "Z:")
if idx < 0 {
return nil
}
ts, err := time.Parse(time.RFC3339Nano, s[:idx+1])
if err != nil {
// Forward-compat: unrecognised key format is left
// in place rather than treated as stale, so a
// future schema change does not silently drop
// records during the rolling upgrade window.
return nil //nolint:nilerr // intentional skip
}
if ts.Before(cutoff) {
keyCopy := append([]byte(nil), k...)
stale = append(stale, keyCopy)
}
return nil
})
for _, k := range stale {
if err := b.Delete(k); err == nil {
removed++
}
}
return nil
})
return removed
}
// DryRunBlocksCount returns the number of recorded dry-run block entries.
func (db *DB) DryRunBlocksCount() int {
if db == nil || db.bolt == nil {
return 0
}
count := 0
_ = db.bolt.View(func(tx *bolt.Tx) error {
if b := tx.Bucket([]byte("dry_run_blocks")); b != nil {
count = b.Stats().KeyN
}
return nil
})
return count
}
// migrateIfNeeded checks for the meta:migrated key and runs migration if absent.
func (db *DB) migrateIfNeeded(statePath string) error {
var migrated bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
if b.Get([]byte("migrated")) != nil {
migrated = true
}
return nil
})
if migrated {
return nil
}
return db.runMigration(statePath)
}
package store
import (
"encoding/json"
"fmt"
"time"
bolt "go.etcd.io/bbolt"
)
// DBObjectBackup is the persisted record of a SHOW CREATE captured
// before a manual `csm db-clean drop-object`. The CREATE SQL is the
// backup -- replaying it restores the object verbatim. Fields are
// public so the cleanup-history UI can render them without a
// separate API.
type DBObjectBackup struct {
Account string `json:"account"`
Schema string `json:"schema"`
Kind string `json:"kind"` // trigger | event | procedure | function
Name string `json:"name"`
CreateSQL string `json:"create_sql"`
DroppedAt time.Time `json:"dropped_at"`
DroppedBy string `json:"dropped_by"` // operator login or "csm" for daemon-driven
FindingID string `json:"finding_id,omitempty"`
RestoredAt time.Time `json:"restored_at,omitempty"`
}
// PutDBObjectBackup writes one backup record. Key shape:
// `<account>:<schema>:<kind>:<name>:<unix_nanos>` so multiple drops of
// the same object name (e.g., re-creates by an attacker) each get
// their own record.
func (db *DB) PutDBObjectBackup(b DBObjectBackup) error {
if b.Account == "" || b.Schema == "" || b.Kind == "" || b.Name == "" {
return fmt.Errorf("PutDBObjectBackup: account/schema/kind/name all required")
}
if b.DroppedAt.IsZero() {
b.DroppedAt = time.Now().UTC()
}
key := fmt.Sprintf("%s:%s:%s:%s:%d",
b.Account, b.Schema, b.Kind, b.Name, b.DroppedAt.UnixNano())
payload, err := json.Marshal(b)
if err != nil {
return fmt.Errorf("marshal backup: %w", err)
}
return db.bolt.Update(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte("db_object_backups"))
if bucket == nil {
return fmt.Errorf("db_object_backups bucket missing (store not migrated)")
}
return bucket.Put([]byte(key), payload)
})
}
// ListDBObjectBackups returns every record for the given account, in
// insertion order. Used by the CLI's listing path and cleanup-history UI.
func (db *DB) ListDBObjectBackups(account string) ([]DBObjectBackup, error) {
var out []DBObjectBackup
prefix := []byte(account + ":")
err := db.bolt.View(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte("db_object_backups"))
if bucket == nil {
return nil
}
c := bucket.Cursor()
for k, v := c.Seek(prefix); k != nil && hasPrefix(k, prefix); k, v = c.Next() {
var b DBObjectBackup
if err := json.Unmarshal(v, &b); err != nil {
continue
}
out = append(out, b)
}
return nil
})
return out, err
}
// GetDBObjectBackupByKey fetches a single record by its exact bbolt
// key. Returns ok=false (not an error) when the key is missing,
// matching the lookup-then-act flow callers use.
func (db *DB) GetDBObjectBackupByKey(key string) (DBObjectBackup, bool, error) {
var rec DBObjectBackup
var found bool
err := db.bolt.View(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte("db_object_backups"))
if bucket == nil {
return nil
}
raw := bucket.Get([]byte(key))
if raw == nil {
return nil
}
if err := json.Unmarshal(raw, &rec); err != nil {
return err
}
found = true
return nil
})
return rec, found, err
}
// MarkDBObjectBackupRestored records that a backup has been replayed. The
// backup row stays in place for audit and future manual inspection, but the
// WebUI can stop offering repeat restore actions for that exact archive.
func (db *DB) MarkDBObjectBackupRestored(key string, restoredAt time.Time) error {
if restoredAt.IsZero() {
restoredAt = time.Now().UTC()
}
return db.bolt.Update(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte("db_object_backups"))
if bucket == nil {
return fmt.Errorf("db_object_backups bucket missing (store not migrated)")
}
raw := bucket.Get([]byte(key))
if raw == nil {
return nil
}
var rec DBObjectBackup
if err := json.Unmarshal(raw, &rec); err != nil {
return err
}
rec.RestoredAt = restoredAt.UTC()
payload, err := json.Marshal(rec)
if err != nil {
return fmt.Errorf("marshal backup: %w", err)
}
return bucket.Put([]byte(key), payload)
})
}
// ListDBObjectBackupsAll returns every record in the bucket,
// regardless of account, in insertion order. Used by the webui
// cleanup-history listing where the operator browses across all
// accounts at once.
func (db *DB) ListDBObjectBackupsAll() ([]DBObjectBackup, []string, error) {
var records []DBObjectBackup
var keys []string
err := db.bolt.View(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte("db_object_backups"))
if bucket == nil {
return nil
}
return bucket.ForEach(func(k, v []byte) error {
var b DBObjectBackup
// Skip malformed rows silently; returning the unmarshal
// error from ForEach would abort the entire iteration,
// which is the wrong choice when one bad row shouldn't
// hide every other operator's history.
if json.Unmarshal(v, &b) == nil {
records = append(records, b)
keys = append(keys, string(k))
}
return nil
})
})
return records, keys, err
}
func hasPrefix(b, prefix []byte) bool {
if len(b) < len(prefix) {
return false
}
for i, c := range prefix {
if b[i] != c {
return false
}
}
return true
}
package store
import (
"encoding/json"
"time"
bolt "go.etcd.io/bbolt"
)
// GeoHistory tracks the countries from which a mailbox has logged in.
type GeoHistory struct {
Countries map[string]int64 `json:"countries"`
LoginCount int `json:"login_count"`
}
// SetGeoHistory stores geo login history for a mailbox in the email:geo bucket.
func (db *DB) SetGeoHistory(mailbox string, h GeoHistory) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("email:geo"))
val, err := json.Marshal(h)
if err != nil {
return err
}
return b.Put([]byte(mailbox), val)
})
}
// GetGeoHistory retrieves geo login history for a mailbox.
// Returns the entry and true if found, or a zero value and false if not.
func (db *DB) GetGeoHistory(mailbox string) (GeoHistory, bool) {
var h GeoHistory
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("email:geo"))
v := b.Get([]byte(mailbox))
if v == nil {
return nil
}
if json.Unmarshal(v, &h) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
found = true
return nil
})
return h, found
}
// SetForwarderHash stores a forwarder config hash in the email:fwd bucket.
func (db *DB) SetForwarderHash(key, hash string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("email:fwd"))
return b.Put([]byte(key), []byte(hash))
})
}
// GetForwarderHash retrieves a forwarder config hash.
// Returns the hash and true if found, or an empty string and false if not.
func (db *DB) GetForwarderHash(key string) (string, bool) {
var hash string
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("email:fwd"))
v := b.Get([]byte(key))
if v == nil {
return nil
}
hash = string(v)
found = true
return nil
})
return hash, found
}
// GetEmailPWLastRefresh reads the last email password-check refresh timestamp
// from the meta bucket. Returns the zero time if not set.
func (db *DB) GetEmailPWLastRefresh() time.Time {
var t time.Time
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
v := b.Get([]byte("email:pw_last_refresh"))
if v == nil {
return nil
}
parsed, err := time.Parse(time.RFC3339, string(v))
if err != nil {
return nil //nolint:nilerr // skip corrupt entry
}
t = parsed
return nil
})
return t
}
// SetEmailPWLastRefresh writes the email password-check refresh timestamp
// to the meta bucket.
func (db *DB) SetEmailPWLastRefresh(t time.Time) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
return b.Put([]byte("email:pw_last_refresh"), []byte(t.Format(time.RFC3339)))
})
}
// GetMetaString reads a string value from the meta bucket.
// Returns an empty string if the key is not found.
func (db *DB) GetMetaString(key string) string {
val, _ := db.ReadMetaString(key)
return val
}
// ReadMetaString distinguishes an absent key from a failed read for callers
// that must not treat unavailable metadata as permission to repeat an action.
func (db *DB) ReadMetaString(key string) (string, error) {
var val string
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
v := b.Get([]byte(key))
if v == nil {
return nil
}
val = string(v)
return nil
})
return val, err
}
// SetMetaString writes a string value to the meta bucket.
func (db *DB) SetMetaString(key, val string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
return b.Put([]byte(key), []byte(val))
})
}
package store
import (
"time"
bolt "go.etcd.io/bbolt"
)
// externalScriptsBucket remembers every external <script src> host seen in
// a site's wp_options, keyed "<site>|<option>|<host>" -> RFC3339 first-seen
// time, plus "<site>|" -> baseline completion time. The structural
// classifier only flags attacker markers, so a loader on an unremarkable
// HTTPS host is invisible to it; remembering hosts lets the scan report a
// host the first time it appears after the site's baseline, once.
const externalScriptsBucket = "db:external_scripts"
func externalScriptKey(site, option, host string) []byte {
return []byte(site + "|" + option + "|" + host)
}
func externalScriptBaselineKey(site string) []byte {
return []byte(site + "|")
}
// MarkExternalScriptSeen records host for (site, option) and reports whether
// it is new. Nothing is new until FinishExternalScriptBaseline has run for
// the site: the first scan records what is already there without reporting
// it, the way file_index treats its first pass.
func (db *DB) MarkExternalScriptSeen(site, option, host string, now time.Time) (bool, error) {
if db == nil || db.bolt == nil {
return false, nil
}
isNew := false
err := db.bolt.Update(func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists([]byte(externalScriptsBucket))
if err != nil {
return err
}
key := externalScriptKey(site, option, host)
if b.Get(key) != nil {
return nil
}
isNew = b.Get(externalScriptBaselineKey(site)) != nil
return b.Put(key, []byte(now.UTC().Format(time.RFC3339)))
})
return isNew, err
}
// FinishExternalScriptBaseline marks the site's first scan as complete, so
// hosts recorded from now on count as new.
func (db *DB) FinishExternalScriptBaseline(site string, now time.Time) error {
if db == nil || db.bolt == nil {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b, err := tx.CreateBucketIfNotExists([]byte(externalScriptsBucket))
if err != nil {
return err
}
key := externalScriptBaselineKey(site)
if b.Get(key) != nil {
return nil
}
return b.Put(key, []byte(now.UTC().Format(time.RFC3339)))
})
}
package store
import (
"encoding/json"
"fmt"
"time"
"github.com/pidginhost/csm/internal/firewall"
bolt "go.etcd.io/bbolt"
)
// FWBlockedEntry represents an IP blocked by the firewall.
type FWBlockedEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
BlockedAt time.Time `json:"blocked_at"`
ExpiresAt time.Time `json:"expires_at"` // zero = permanent
}
// FWAllowedEntry represents an IP explicitly allowed through the firewall.
type FWAllowedEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
Port int `json:"port"` // 0 = all ports
ExpiresAt time.Time `json:"expires_at"` // zero = permanent
}
// FWSubnetEntry represents a subnet added to the firewall.
type FWSubnetEntry struct {
CIDR string `json:"cidr"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
AddedAt time.Time `json:"added_at"`
}
// FWPortAllowEntry represents a per-IP port allow rule.
type FWPortAllowEntry struct {
Key string `json:"key"` // IP:port/proto
IP string `json:"ip"`
Port int `json:"port"`
Proto string `json:"proto"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
}
// FirewallState holds the full state across all 4 firewall buckets.
type FirewallState struct {
Blocked []FWBlockedEntry
Allowed []FWAllowedEntry
Subnets []FWSubnetEntry
PortAllowed []FWPortAllowEntry
}
// portAllowKey returns the composite key "IP:port/proto".
func portAllowKey(ip string, port int, proto string) string {
return fmt.Sprintf("%s:%d/%s", ip, port, proto)
}
// BlockIP adds an IP to the fw:blocked bucket.
func (db *DB) BlockIP(ip, reason string, expiresAt time.Time) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:blocked"))
entry := FWBlockedEntry{
IP: ip,
Reason: reason,
Source: firewall.InferProvenance("block", reason),
BlockedAt: time.Now(),
ExpiresAt: expiresAt,
}
val, err := json.Marshal(entry)
if err != nil {
return err
}
return b.Put([]byte(ip), val)
})
}
// UnblockIP removes an IP from the fw:blocked bucket.
func (db *DB) UnblockIP(ip string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:blocked"))
return b.Delete([]byte(ip))
})
}
// GetBlockedIP looks up a blocked IP. Returns false if not found or expired.
func (db *DB) GetBlockedIP(ip string) (FWBlockedEntry, bool) {
var entry FWBlockedEntry
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:blocked"))
v := b.Get([]byte(ip))
if v == nil {
return nil
}
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
// Filter expired entries (zero ExpiresAt = permanent).
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(time.Now()) {
return nil
}
found = true
return nil
})
return entry, found
}
// AllowIP adds an IP to the fw:allowed bucket.
func (db *DB) AllowIP(ip, reason string, expiresAt time.Time) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:allowed"))
entry := FWAllowedEntry{
IP: ip,
Reason: reason,
Source: firewall.InferProvenance("allow", reason),
ExpiresAt: expiresAt,
}
val, err := json.Marshal(entry)
if err != nil {
return err
}
return b.Put([]byte(ip), val)
})
}
// RemoveAllow removes an IP from the fw:allowed bucket.
func (db *DB) RemoveAllow(ip string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:allowed"))
return b.Delete([]byte(ip))
})
}
// AddSubnet adds a CIDR to the fw:subnets bucket.
func (db *DB) AddSubnet(cidr, reason string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:subnets"))
entry := FWSubnetEntry{
CIDR: cidr,
Reason: reason,
Source: firewall.InferProvenance("block_subnet", reason),
AddedAt: time.Now(),
}
val, err := json.Marshal(entry)
if err != nil {
return err
}
return b.Put([]byte(cidr), val)
})
}
// RemoveSubnet removes a CIDR from the fw:subnets bucket.
func (db *DB) RemoveSubnet(cidr string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:subnets"))
return b.Delete([]byte(cidr))
})
}
// AddPortAllow adds a per-IP port allow rule to the fw:port_allowed bucket.
func (db *DB) AddPortAllow(ip string, port int, proto, reason string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:port_allowed"))
key := portAllowKey(ip, port, proto)
entry := FWPortAllowEntry{
Key: key,
IP: ip,
Port: port,
Proto: proto,
Reason: reason,
Source: firewall.InferProvenance("allow_port", reason),
}
val, err := json.Marshal(entry)
if err != nil {
return err
}
return b.Put([]byte(key), val)
})
}
// RemovePortAllow removes a per-IP port allow rule from the fw:port_allowed bucket.
func (db *DB) RemovePortAllow(ip string, port int, proto string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:port_allowed"))
key := portAllowKey(ip, port, proto)
return b.Delete([]byte(key))
})
}
// ListPortAllows returns all entries in the fw:port_allowed bucket.
func (db *DB) ListPortAllows() []FWPortAllowEntry {
var entries []FWPortAllowEntry
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("fw:port_allowed"))
return b.ForEach(func(k, v []byte) error {
var entry FWPortAllowEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
entries = append(entries, entry)
return nil
})
})
return entries
}
// LoadFirewallState reads all 4 firewall buckets and assembles a FirewallState.
// Expired blocked and allowed entries are filtered out.
func (db *DB) LoadFirewallState() FirewallState {
var state FirewallState
now := time.Now()
_ = db.bolt.View(func(tx *bolt.Tx) error {
// fw:blocked - filter expired
blocked := tx.Bucket([]byte("fw:blocked"))
_ = blocked.ForEach(func(k, v []byte) error {
var entry FWBlockedEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
return nil // expired
}
state.Blocked = append(state.Blocked, entry)
return nil
})
// fw:allowed - filter expired
allowed := tx.Bucket([]byte("fw:allowed"))
_ = allowed.ForEach(func(k, v []byte) error {
var entry FWAllowedEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if !entry.ExpiresAt.IsZero() && !entry.ExpiresAt.After(now) {
return nil // expired
}
state.Allowed = append(state.Allowed, entry)
return nil
})
// fw:subnets
subnets := tx.Bucket([]byte("fw:subnets"))
_ = subnets.ForEach(func(k, v []byte) error {
var entry FWSubnetEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
state.Subnets = append(state.Subnets, entry)
return nil
})
// fw:port_allowed
portAllowed := tx.Bucket([]byte("fw:port_allowed"))
_ = portAllowed.ForEach(func(k, v []byte) error {
var entry FWPortAllowEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
state.PortAllowed = append(state.PortAllowed, entry)
return nil
})
return nil
})
return state
}
package store
import (
"bytes"
"encoding/binary"
"fmt"
"time"
"github.com/pidginhost/csm/internal/firewall"
bolt "go.etcd.io/bbolt"
)
// Terminal actions are indexed by outcome time so retention reads an ordered
// list without decoding the evidence each record carries.
const firewallActionHistoryBucket = "fw:action_history"
type firewallActionRetentionCaps struct {
// Actions and Bytes bound the retained outcome history. Bytes matters
// because one record keeps two complete states plus kernel evidence, so a
// host with a large blocked set writes far more per action than a small one.
Actions int
Bytes uint64
// BudgetWindows bounds the hourly scan counters. Only the current window is
// ever read, and the inventory is validated on every admission.
BudgetWindows int
}
// Retention sweeps are opt-in, so on a default install these caps are the only
// bound on journal growth. They apply on the write path, not on a timer.
var firewallActionRetention = firewallActionRetentionCaps{Actions: 1000, Bytes: 64 << 20, BudgetWindows: 48}
func firewallActionHistoryKey(at time.Time, id string) []byte {
nanoseconds := at.UnixNano()
if nanoseconds < 0 {
nanoseconds = 0
}
return []byte(fmt.Sprintf("%020d\x00%s", nanoseconds, id))
}
func firewallActionHistoryID(key []byte) (string, bool) {
separator := bytes.IndexByte(key, 0)
if separator < 0 || separator+1 >= len(key) {
return "", false
}
return string(key[separator+1:]), true
}
type firewallActionHistoryEntry struct {
key []byte
id string
size uint64
}
// readFirewallActionHistory returns retained outcomes oldest first. A history
// key without its action record is corruption: both are written in one
// transaction, so recovery must not silently accept a half-deleted record.
func readFirewallActionHistory(tx *bolt.Tx) ([]firewallActionHistoryEntry, uint64, error) {
b := tx.Bucket([]byte(firewallActionHistoryBucket))
if b == nil {
return nil, 0, nil
}
actions := tx.Bucket([]byte(firewallActionsBucket))
var entries []firewallActionHistoryEntry
var total uint64
cursor := b.Cursor()
for key, raw := cursor.First(); key != nil; key, raw = cursor.Next() {
id, ok := firewallActionHistoryID(key)
if !ok || len(raw) != 8 {
return nil, 0, fmt.Errorf("%w: firewall action history entry", firewall.ErrStateCorrupt)
}
if actions == nil || actions.Get([]byte(id)) == nil {
return nil, 0, fmt.Errorf("%w: firewall action history without record", firewall.ErrStateCorrupt)
}
// The index stores the action row size. Audit rows are counted live
// so existing indexes also include every retained evidence copy.
size := binary.BigEndian.Uint64(raw) + firewallActionAuditSize(tx, id)
total += size
entries = append(entries, firewallActionHistoryEntry{key: bytes.Clone(key), id: id, size: size})
}
return entries, total, nil
}
// Older journals predate the outcome index. Build it atomically before the
// first retention operation, including undelivered outcomes for later pruning.
func initializeFirewallActionHistory(tx *bolt.Tx) error {
if tx.Bucket([]byte(firewallActionHistoryBucket)) != nil {
return nil
}
actions := tx.Bucket([]byte(firewallActionsBucket))
if actions == nil {
return nil
}
if _, err := tx.CreateBucketIfNotExists([]byte(firewallActionHistoryBucket)); err != nil {
return err
}
return actions.ForEach(func(key, raw []byte) error {
a, err := readFirewallAction(tx, string(key))
if err != nil {
return err
}
if firewallActionPending(a.Phase) {
return nil
}
return recordFirewallActionHistory(tx, a, len(raw))
})
}
// Audit keys have a fixed-width version suffix. A prefix alone also matches
// other valid request IDs containing NUL, so match the complete key length.
func firewallActionAuditSize(tx *bolt.Tx, id string) uint64 {
var size uint64
if b := tx.Bucket([]byte(firewallAuditBucket)); b != nil {
prefix := append([]byte(id), 0)
cursor := b.Cursor()
for key, raw := cursor.Seek(prefix); key != nil && bytes.HasPrefix(key, prefix); key, raw = cursor.Next() {
if len(key) == len(prefix)+20 {
size += uint64(len(raw))
}
}
}
return size
}
func recordFirewallActionHistory(tx *bolt.Tx, a firewall.FirewallAction, size int) error {
if err := initializeFirewallActionHistory(tx); err != nil {
return err
}
b := tx.Bucket([]byte(firewallActionHistoryBucket))
var encoded [8]byte
binary.BigEndian.PutUint64(encoded[:], uint64(size)) // #nosec G115 -- a stored record length is never negative.
return b.Put(firewallActionHistoryKey(a.UpdatedAt, a.Request.ID), encoded[:])
}
// updateFirewallActionHistorySize keeps byte accounting honest after an
// acknowledgement rewrites a retained record.
func updateFirewallActionHistorySize(tx *bolt.Tx, a firewall.FirewallAction, size int) error {
if err := initializeFirewallActionHistory(tx); err != nil {
return err
}
b := tx.Bucket([]byte(firewallActionHistoryBucket))
if b == nil {
return nil
}
key := firewallActionHistoryKey(a.UpdatedAt, a.Request.ID)
if b.Get(key) == nil {
return nil
}
var encoded [8]byte
binary.BigEndian.PutUint64(encoded[:], uint64(size)) // #nosec G115 -- a stored record length is never negative.
return b.Put(key, encoded[:])
}
func undeliveredFirewallAuditIDs(index firewallJournalIndex) map[string]bool {
undelivered := make(map[string]bool, len(index.Audit))
for _, ref := range index.Audit {
undelivered[ref.ID] = true
}
return undelivered
}
// deleteFirewallAction removes one retained outcome with every event that
// refers to it. Undelivered outcomes are never passed here.
func deleteFirewallAction(tx *bolt.Tx, entry firewallActionHistoryEntry) error {
if b := tx.Bucket([]byte(firewallAuditBucket)); b != nil {
prefix := append([]byte(entry.id), 0)
// Collect first: acknowledgement rewrites an audit leaf in this
// transaction, and a bbolt cursor that deletes and then steps with
// Next skips keys in a bucket already written by the transaction.
var keys [][]byte
cursor := b.Cursor()
for key, _ := cursor.Seek(prefix); key != nil && bytes.HasPrefix(key, prefix); key, _ = cursor.Next() {
if len(key) == len(prefix)+20 {
keys = append(keys, append([]byte(nil), key...))
}
}
for _, key := range keys {
if err := b.Delete(key); err != nil {
return err
}
}
}
if b := tx.Bucket([]byte(firewallActionsBucket)); b != nil {
if err := b.Delete([]byte(entry.id)); err != nil {
return err
}
}
return tx.Bucket([]byte(firewallActionHistoryBucket)).Delete(entry.key)
}
// pruneFirewallActionHistory enforces the retained-outcome caps inside the
// transaction that recorded the newest outcome. The newest record and any
// record with undelivered audit are kept whatever the caps say.
func pruneFirewallActionHistory(tx *bolt.Tx, index firewallJournalIndex) error {
if err := initializeFirewallActionHistory(tx); err != nil {
return err
}
entries, total, err := readFirewallActionHistory(tx)
if err != nil {
return err
}
undelivered := undeliveredFirewallAuditIDs(index)
count := len(entries)
for _, entry := range entries[:max(len(entries)-1, 0)] {
if count <= firewallActionRetention.Actions && total <= firewallActionRetention.Bytes {
return nil
}
if undelivered[entry.id] {
continue
}
if err := deleteFirewallAction(tx, entry); err != nil {
return err
}
count--
total -= entry.size
}
return nil
}
// SweepFirewallActionsOlderThan deletes proven outcomes whose audit has been
// delivered and whose result is older than cutoff. Pending actions and
// undelivered outcomes are retained regardless of age: recovery still needs
// them. Deleting an outcome ends the undo window for that action.
func (db *DB) SweepFirewallActionsOlderThan(cutoff time.Time) (int, error) {
var deleted int
err := boltUpdate(db.bolt, func(tx *bolt.Tx) error {
deleted = 0
index, err := readFirewallJournalIndex(tx)
if err != nil {
return err
}
if initErr := initializeFirewallActionHistory(tx); initErr != nil {
return initErr
}
entries, _, err := readFirewallActionHistory(tx)
if err != nil {
return err
}
undelivered := undeliveredFirewallAuditIDs(index)
boundary := firewallActionHistoryKey(cutoff, "")
for _, entry := range entries {
if bytes.Compare(entry.key, boundary) >= 0 {
break
}
if undelivered[entry.id] {
continue
}
if err := deleteFirewallAction(tx, entry); err != nil {
return err
}
deleted++
}
return sweepFirewallScanBudget(tx, cutoff)
})
if err != nil {
return 0, err
}
return deleted, nil
}
// sweepFirewallScanBudget drops hourly counters the admission path can no
// longer charge. The inventory and its rows are updated together so a missing
// counter keeps meaning corruption.
func sweepFirewallScanBudget(tx *bolt.Tx, cutoff time.Time) error {
inventory, err := readFirewallBudgetInventory(tx)
if err != nil {
return err
}
boundary := cutoff.UTC().Format(firewallScanWindowLayout)
keep := inventory.Windows[:0:0]
for _, window := range inventory.Windows {
if window >= boundary {
keep = append(keep, window)
continue
}
if err := deleteFirewallScanBudgetWindow(tx, window); err != nil {
return err
}
inventory.PrunedThrough = max(inventory.PrunedThrough, window)
}
if len(keep) == len(inventory.Windows) {
return nil
}
inventory.Windows = keep
return writeFirewallBudgetInventory(tx, inventory)
}
func deleteFirewallScanBudgetWindow(tx *bolt.Tx, window string) error {
b := tx.Bucket([]byte(firewallBudgetBucket))
if b == nil {
return nil
}
return b.Delete([]byte(window))
}
// pruneFirewallScanBudgetWindows keeps the newest windows only. The current
// window may move backwards after a clock correction. Keep a durable boundary
// so admission refuses discarded windows instead of resetting their charges.
func pruneFirewallScanBudgetWindows(tx *bolt.Tx, inventory firewallBudgetInventory) (firewallBudgetInventory, error) {
if len(inventory.Windows) <= firewallActionRetention.BudgetWindows {
return inventory, nil
}
excess := len(inventory.Windows) - firewallActionRetention.BudgetWindows
for _, window := range inventory.Windows[:excess] {
if err := deleteFirewallScanBudgetWindow(tx, window); err != nil {
return inventory, err
}
}
inventory.PrunedThrough = max(inventory.PrunedThrough, inventory.Windows[excess-1])
inventory.Windows = append(inventory.Windows[:0:0], inventory.Windows[excess:]...)
return inventory, writeFirewallBudgetInventory(tx, inventory)
}
package store
import (
"bytes"
"crypto/sha256"
"encoding/json"
"errors"
"fmt"
"math"
"slices"
"time"
"unicode/utf8"
"github.com/pidginhost/csm/internal/firewall"
bolt "go.etcd.io/bbolt"
)
const firewallActionsBucket = "fw:actions"
const firewallBudgetBucket = "fw:scan_budget"
const firewallAuditBucket = "fw:action_audit"
const firewallActionIndexBucket = "fw:action_index"
const firewallBudgetIndexBucket = "fw:budget_index"
const firewallScanWindowLayout = "2006-01-02T15"
var _ firewall.ActionStore = (*DB)(nil)
// ErrFirewallActionMissing distinguishes an unknown request from a damaged
// journal. Neither is permission to execute an unrecorded mutation.
var ErrFirewallActionMissing = firewall.ErrActionMissing
func firewallActionPending(phase string) bool {
return phase == "planned" || phase == "executing" || phase == "applied" || phase == "unknown"
}
func validateFirewallAdmission(a firewall.FirewallAction) error {
if a.Request.ID == "" || len(a.Request.ID) > 128 || a.Request.Operation == "" || a.Request.Actor == "" || a.CreatedAt.IsZero() || a.Revision == 0 {
return errors.New("invalid firewall action admission")
}
for _, value := range []string{a.Request.ID, a.Request.Operation, a.Request.Target, a.Request.Reason, a.Request.Actor, a.Request.ActorDetail, a.Request.Source, a.Request.FindingID, a.Request.IncidentID, a.Request.UndoOf, a.Detail} {
if !utf8.ValidString(value) {
return errors.New("firewall action contains invalid Unicode")
}
}
for _, at := range []time.Time{a.CreatedAt, a.UpdatedAt} {
_, offset := at.Zone()
if offset%60 != 0 {
return errors.New("firewall action time cannot round-trip")
}
}
if a.Budget != nil {
if a.Request.Source != "scan" || a.Budget.Limit <= 0 {
return errors.New("invalid firewall scan admission")
}
if _, err := time.Parse(firewallScanWindowLayout, a.Budget.Window); err != nil {
return errors.New("invalid firewall scan window")
}
}
return nil
}
// Journal envelopes detect valid-JSON corruption before recovery or typed undo
// interprets historical evidence. Checksums cover the exact stored payload.
type firewallJournalEnvelope struct {
Version uint64 `json:"version"`
Payload json.RawMessage `json:"payload"`
SHA256 string `json:"sha256"`
}
func encodeFirewallJournal(value any) ([]byte, error) {
payload, err := json.Marshal(value)
if err != nil {
return nil, err
}
digest := sha256.Sum256(payload)
return json.Marshal(firewallJournalEnvelope{Version: 1, Payload: payload, SHA256: fmt.Sprintf("%x", digest)})
}
func decodeFirewallJournal(raw []byte) ([]byte, error) {
var envelope firewallJournalEnvelope
if !validFirewallJSONText(raw) || json.Unmarshal(raw, &envelope) != nil || envelope.Version != 1 || !bytes.HasPrefix(bytes.TrimSpace(envelope.Payload), []byte("{")) {
return nil, firewall.ErrStateCorrupt
}
digest := sha256.Sum256(envelope.Payload)
if envelope.SHA256 != fmt.Sprintf("%x", digest) {
return nil, firewall.ErrStateCorrupt
}
return envelope.Payload, nil
}
func decodeFirewallAction(raw []byte) (firewall.FirewallAction, error) {
var a firewall.FirewallAction
if len(raw) == 0 {
return a, ErrFirewallActionMissing
}
var err error
raw, err = decodeFirewallJournal(raw)
if err != nil {
return firewall.FirewallAction{}, err
}
var evidence struct {
Before json.RawMessage `json:"before"`
After json.RawMessage `json:"after"`
}
if json.Unmarshal(raw, &evidence) != nil || !bytes.HasPrefix(bytes.TrimSpace(evidence.Before), []byte("{")) || !bytes.HasPrefix(bytes.TrimSpace(evidence.After), []byte("{")) {
return firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
if !validFirewallJSONText(raw) || json.Unmarshal(raw, &a) != nil || validateFirewallAdmission(a) != nil || a.UpdatedAt.IsZero() || a.AuditVersion == 0 || a.AuditAck > a.AuditVersion {
return firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
if !firewallActionPending(a.Phase) && a.Phase != "verified" && a.Phase != "failed" {
return firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
if _, _, err := encodeFirewallSnapshot(a.Before); err != nil {
return firewall.FirewallAction{}, fmt.Errorf("%w: action before state: %v", firewall.ErrStateCorrupt, err)
}
if _, _, err := encodeFirewallSnapshot(a.After); err != nil {
return firewall.FirewallAction{}, fmt.Errorf("%w: action after state: %v", firewall.ErrStateCorrupt, err)
}
return a, nil
}
func readFirewallAction(tx *bolt.Tx, id string) (firewall.FirewallAction, error) {
b := tx.Bucket([]byte(firewallActionsBucket))
if b == nil {
return firewall.FirewallAction{}, ErrFirewallActionMissing
}
a, err := decodeFirewallAction(b.Get([]byte(id)))
if err == nil && a.Request.ID != id {
return firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
return a, err
}
func writeFirewallAction(tx *bolt.Tx, a firewall.FirewallAction) (int, error) {
raw, err := encodeFirewallJournal(a)
if err != nil {
return 0, err
}
b, err := tx.CreateBucketIfNotExists([]byte(firewallActionsBucket))
if err != nil {
return 0, err
}
return len(raw), b.Put([]byte(a.Request.ID), raw)
}
// Only outstanding work is indexed. Retained history is validated when read for
// inspection or undo, so admission does not decode every historical snapshot.
type firewallAuditReference struct {
ID string `json:"id"`
Version uint64 `json:"version"`
}
type firewallJournalIndex struct {
Initialized bool `json:"initialized"`
PendingID string `json:"pending_id"`
Audit []firewallAuditReference `json:"audit"`
}
func readFirewallJournalIndex(tx *bolt.Tx) (firewallJournalIndex, error) {
var index firewallJournalIndex
b := tx.Bucket([]byte(firewallActionIndexBucket))
if b == nil {
for _, name := range []string{firewallActionsBucket, firewallAuditBucket, firewallBudgetBucket, firewallBudgetIndexBucket, firewallActionHistoryBucket} {
if tx.Bucket([]byte(name)) != nil {
return index, firewall.ErrStateCorrupt
}
}
return index, nil
}
payload, err := decodeFirewallJournal(b.Get([]byte("index")))
if err != nil {
return index, err
}
var required struct {
PendingID *string `json:"pending_id"`
Audit json.RawMessage `json:"audit"`
}
if json.Unmarshal(payload, &required) != nil || required.PendingID == nil || !bytes.HasPrefix(bytes.TrimSpace(required.Audit), []byte("[")) {
return firewallJournalIndex{}, firewall.ErrStateCorrupt
}
if json.Unmarshal(payload, &index) != nil || !index.Initialized || tx.Bucket([]byte(firewallActionsBucket)) == nil {
return firewallJournalIndex{}, firewall.ErrStateCorrupt
}
seen := make(map[firewallAuditReference]bool, len(index.Audit))
for _, ref := range index.Audit {
if ref.ID == "" || ref.Version == 0 || seen[ref] {
return firewallJournalIndex{}, firewall.ErrStateCorrupt
}
seen[ref] = true
}
return index, nil
}
func writeFirewallJournalIndex(tx *bolt.Tx, index firewallJournalIndex) error {
index.Initialized = true
if index.Audit == nil {
index.Audit = []firewallAuditReference{}
}
raw, err := encodeFirewallJournal(index)
if err != nil {
return err
}
b, err := tx.CreateBucketIfNotExists([]byte(firewallActionIndexBucket))
if err != nil {
return err
}
return b.Put([]byte("index"), raw)
}
func indexedPendingFirewallAction(tx *bolt.Tx, index firewallJournalIndex) (firewall.FirewallAction, error) {
if index.PendingID == "" {
return firewall.FirewallAction{}, nil
}
a, err := readFirewallAction(tx, index.PendingID)
if err != nil || !firewallActionPending(a.Phase) {
return firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
return a, nil
}
func refusePendingFirewallActions(tx *bolt.Tx) error {
index, err := readFirewallJournalIndex(tx)
if err != nil {
return err
}
a, err := indexedPendingFirewallAction(tx, index)
if err != nil {
return err
}
if a.Request.ID != "" {
return fmt.Errorf("%w: action %s requires recovery", firewall.ErrStateConflict, a.Request.ID)
}
return nil
}
type firewallScanBudget struct {
Window string `json:"window"`
Count *int `json:"count"`
}
// Keep the inventory separate from pending action metadata: recovery need not
// decode previously used budget windows. Missing charged counters are corruption.
type firewallBudgetInventory struct {
Initialized bool `json:"initialized"`
Windows []string `json:"windows"`
// PrunedThrough prevents discarded charges from reopening after a clock
// correction. Omitted in journals written before budget retention.
PrunedThrough string `json:"pruned_through,omitempty"`
}
func readFirewallBudgetInventory(tx *bolt.Tx) (firewallBudgetInventory, error) {
var inventory firewallBudgetInventory
b := tx.Bucket([]byte(firewallBudgetIndexBucket))
if b == nil {
for _, name := range []string{firewallActionsBucket, firewallAuditBucket, firewallBudgetBucket, firewallActionIndexBucket, firewallActionHistoryBucket} {
if tx.Bucket([]byte(name)) != nil {
return inventory, firewall.ErrStateCorrupt
}
}
return inventory, nil
}
payload, err := decodeFirewallJournal(b.Get([]byte("index")))
if err != nil {
return inventory, err
}
if json.Unmarshal(payload, &inventory) != nil || !inventory.Initialized || inventory.Windows == nil {
return firewallBudgetInventory{}, firewall.ErrStateCorrupt
}
if inventory.PrunedThrough != "" {
if _, err := time.Parse(firewallScanWindowLayout, inventory.PrunedThrough); err != nil {
return firewallBudgetInventory{}, firewall.ErrStateCorrupt
}
}
for i, window := range inventory.Windows {
if window <= inventory.PrunedThrough {
return firewallBudgetInventory{}, firewall.ErrStateCorrupt
}
if _, err := time.Parse(firewallScanWindowLayout, window); err != nil {
return firewallBudgetInventory{}, firewall.ErrStateCorrupt
}
if i > 0 && inventory.Windows[i-1] >= window {
return firewallBudgetInventory{}, firewall.ErrStateCorrupt
}
}
return inventory, nil
}
func writeFirewallBudgetInventory(tx *bolt.Tx, inventory firewallBudgetInventory) error {
inventory.Initialized = true
if inventory.Windows == nil {
inventory.Windows = []string{}
}
raw, err := encodeFirewallJournal(inventory)
if err != nil {
return err
}
b, err := tx.CreateBucketIfNotExists([]byte(firewallBudgetIndexBucket))
if err != nil {
return err
}
return b.Put([]byte("index"), raw)
}
func readFirewallScanBudget(tx *bolt.Tx, window string) (int, error) {
inventory, err := readFirewallBudgetInventory(tx)
if err != nil {
return 0, err
}
return readFirewallScanBudgetCount(tx, window, inventory)
}
func readFirewallScanBudgetCount(tx *bolt.Tx, window string, inventory firewallBudgetInventory) (int, error) {
_, known := slices.BinarySearch(inventory.Windows, window)
b := tx.Bucket([]byte(firewallBudgetBucket))
if b == nil && len(inventory.Windows) > 0 {
return 0, firewall.ErrStateCorrupt
}
if b != nil && b.Bucket([]byte(window)) != nil {
return 0, firewall.ErrStateCorrupt
}
if b == nil || b.Get([]byte(window)) == nil {
if known {
return 0, firewall.ErrStateCorrupt
}
return 0, nil
}
if !known {
return 0, firewall.ErrStateCorrupt
}
payload, err := decodeFirewallJournal(b.Get([]byte(window)))
if err != nil {
return 0, err
}
var budget firewallScanBudget
if json.Unmarshal(payload, &budget) != nil || budget.Window != window || budget.Count == nil || *budget.Count <= 0 {
return 0, firewall.ErrStateCorrupt
}
if _, err := time.Parse(firewallScanWindowLayout, budget.Window); err != nil {
return 0, firewall.ErrStateCorrupt
}
return *budget.Count, nil
}
func (db *DB) ReadFirewallScanBudget(window string) (int, error) {
var count int
err := db.bolt.View(func(tx *bolt.Tx) error {
var err error
count, err = readFirewallScanBudget(tx, window)
return err
})
return count, err
}
// AdmitFirewallAction commits the recovery evidence and accepted scan charge
// together. The complete committed state stays at Before until verification.
func (db *DB) AdmitFirewallAction(in firewall.FirewallAction) (result firewall.FirewallAction, fresh bool, err error) {
defer func() { recordFirewallWriteError(err) }()
if validationErr := validateFirewallAdmission(in); validationErr != nil {
return result, false, validationErr
}
if _, _, err = encodeFirewallSnapshot(in.Before); err != nil {
return result, false, err
}
if _, _, err = encodeFirewallSnapshot(in.After); err != nil {
return result, false, err
}
before, err := json.Marshal(in.Before)
if err != nil {
return result, false, err
}
err = db.updateFirewallSnapshot(1, func(tx *bolt.Tx) error {
index, indexErr := readFirewallJournalIndex(tx)
if indexErr != nil {
return indexErr
}
pending, pendingErr := indexedPendingFirewallAction(tx, index)
if pendingErr != nil {
return pendingErr
}
existing, readErr := readFirewallAction(tx, in.Request.ID)
if readErr == nil {
if existing.Request != in.Request {
return fmt.Errorf("%w: request ID reused", firewall.ErrStateConflict)
}
result = existing
return nil
}
if !errors.Is(readErr, ErrFirewallActionMissing) {
return readErr
}
if pending.Request.ID != "" {
return fmt.Errorf("%w: action %s requires recovery", firewall.ErrStateConflict, pending.Request.ID)
}
meta, txErr := readFirewallSnapshotMeta(tx)
if txErr != nil {
return txErr
}
if meta.Revision != in.Revision {
return firewall.ErrStateConflict
}
rows, txErr := readFirewallSnapshotRows(tx, meta)
if txErr != nil {
return txErr
}
state, txErr := decodeFirewallSnapshot(rows, meta)
if txErr != nil {
return txErr
}
current, txErr := json.Marshal(state)
if txErr != nil {
return txErr
}
if !bytes.Equal(current, before) {
return fmt.Errorf("%w: action before state differs", firewall.ErrStateConflict)
}
if !index.Initialized {
if inventoryErr := writeFirewallBudgetInventory(tx, firewallBudgetInventory{}); inventoryErr != nil {
return inventoryErr
}
}
if in.Budget != nil {
inventory, inventoryErr := readFirewallBudgetInventory(tx)
if inventoryErr != nil {
return inventoryErr
}
if in.Budget.Window <= inventory.PrunedThrough {
return firewall.ErrScanBudget
}
count, budgetErr := readFirewallScanBudgetCount(tx, in.Budget.Window, inventory)
if budgetErr != nil {
return budgetErr
}
if count >= in.Budget.Limit {
return firewall.ErrScanBudget
}
b, createErr := tx.CreateBucketIfNotExists([]byte(firewallBudgetBucket))
if createErr != nil {
return createErr
}
count++
budgetRaw, encodeErr := encodeFirewallJournal(firewallScanBudget{Window: in.Budget.Window, Count: &count})
if encodeErr != nil {
return encodeErr
}
if putErr := b.Put([]byte(in.Budget.Window), budgetRaw); putErr != nil {
return putErr
}
if position, known := slices.BinarySearch(inventory.Windows, in.Budget.Window); !known {
inventory.Windows = slices.Insert(inventory.Windows, position, in.Budget.Window)
if inventoryErr := writeFirewallBudgetInventory(tx, inventory); inventoryErr != nil {
return inventoryErr
}
if _, pruneErr := pruneFirewallScanBudgetWindows(tx, inventory); pruneErr != nil {
return pruneErr
}
}
}
in.Phase = "planned"
in.UpdatedAt = in.CreatedAt
in.Detail = ""
in.AuditVersion, in.AuditAck = 1, 0
if _, writeErr := writeFirewallAction(tx, in); writeErr != nil {
return writeErr
}
index.PendingID = in.Request.ID
if writeErr := writeFirewallJournalIndex(tx, index); writeErr != nil {
return writeErr
}
// Decode to detach every slice from caller-owned input.
result, txErr = readFirewallAction(tx, in.Request.ID)
fresh = txErr == nil
return txErr
})
if err != nil {
return firewall.FirewallAction{}, false, err
}
return result, fresh, nil
}
func (db *DB) ReadFirewallAction(id string) (firewall.FirewallAction, error) {
var a firewall.FirewallAction
err := db.bolt.View(func(tx *bolt.Tx) error {
var err error
a, err = readFirewallAction(tx, id)
return err
})
if err != nil {
return firewall.FirewallAction{}, err
}
return a, nil
}
func (db *DB) PendingFirewallActions() ([]firewall.FirewallAction, error) {
var pending []firewall.FirewallAction
err := db.bolt.View(func(tx *bolt.Tx) error {
index, err := readFirewallJournalIndex(tx)
if err != nil {
return err
}
a, err := indexedPendingFirewallAction(tx, index)
if err != nil {
return err
}
if a.Request.ID != "" {
pending = append(pending, a)
}
return nil
})
if err != nil {
return nil, err
}
return pending, nil
}
func firewallAuditPhase(phase string) bool {
return phase == "unknown" || phase == "verified" || phase == "failed"
}
// Each outcome event keeps its original payload. Acknowledgement changes only
// the envelope; retaining it permits exact, idempotent acknowledgements.
type firewallAuditEvent struct {
Action json.RawMessage `json:"action"`
Acknowledged bool `json:"acknowledged"`
}
func firewallAuditKey(id string, version uint64) []byte {
return []byte(fmt.Sprintf("%s\x00%020d", id, version))
}
func decodeFirewallAuditEvent(key, raw []byte) (firewallAuditEvent, firewall.FirewallAction, error) {
var event firewallAuditEvent
var err error
raw, err = decodeFirewallJournal(raw)
if err != nil {
return event, firewall.FirewallAction{}, err
}
if !validFirewallJSONText(raw) || json.Unmarshal(raw, &event) != nil {
return event, firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
a, err := decodeFirewallAction(event.Action)
if err != nil || !firewallAuditPhase(a.Phase) || !bytes.Equal(key, firewallAuditKey(a.Request.ID, a.AuditVersion)) {
return event, firewall.FirewallAction{}, firewall.ErrStateCorrupt
}
return event, a, nil
}
func writeFirewallAuditEvent(tx *bolt.Tx, a firewall.FirewallAction) error {
if !firewallAuditPhase(a.Phase) {
return nil
}
action, err := encodeFirewallJournal(a)
if err != nil {
return err
}
raw, err := encodeFirewallJournal(firewallAuditEvent{Action: action})
if err != nil {
return err
}
b, err := tx.CreateBucketIfNotExists([]byte(firewallAuditBucket))
if err != nil {
return err
}
key := firewallAuditKey(a.Request.ID, a.AuditVersion)
if b.Get(key) != nil || b.Bucket(key) != nil {
return firewall.ErrStateConflict
}
return b.Put(key, raw)
}
func (db *DB) FirewallAuditPending() ([]firewall.FirewallAction, error) {
var pending []firewall.FirewallAction
err := db.bolt.View(func(tx *bolt.Tx) error {
index, err := readFirewallJournalIndex(tx)
if err != nil {
return err
}
b := tx.Bucket([]byte(firewallAuditBucket))
for _, ref := range index.Audit {
if b == nil {
return firewall.ErrStateCorrupt
}
key := firewallAuditKey(ref.ID, ref.Version)
event, a, decodeErr := decodeFirewallAuditEvent(key, b.Get(key))
if decodeErr != nil || event.Acknowledged {
return firewall.ErrStateCorrupt
}
current, readErr := readFirewallAction(tx, ref.ID)
if readErr != nil || current.Request != a.Request || a.AuditVersion > current.AuditVersion {
return firewall.ErrStateCorrupt
}
pending = append(pending, a)
}
return nil
})
if err != nil {
return nil, err
}
return pending, nil
}
func validFirewallActionTransition(from, to string) bool {
if !firewallActionPending(from) {
return false
}
switch to {
case "executing":
return from == "planned"
case "applied":
return from == "executing"
case "unknown", "verified", "failed":
return true
}
return false
}
func validateFirewallActionBase(tx *bolt.Tx, a firewall.FirewallAction) error {
meta, err := readFirewallSnapshotMeta(tx)
if err != nil {
return err
}
if meta.Revision != a.Revision {
return firewall.ErrStateConflict
}
rows, err := readFirewallSnapshotRows(tx, meta)
if err != nil {
return err
}
state, err := decodeFirewallSnapshot(rows, meta)
if err != nil {
return err
}
current, err := json.Marshal(state)
if err != nil {
return err
}
before, err := json.Marshal(a.Before)
if err != nil {
return err
}
if !bytes.Equal(current, before) {
return fmt.Errorf("%w: action before evidence differs from committed state", firewall.ErrStateCorrupt)
}
return nil
}
func (db *DB) TransitionFirewallAction(id, phase, detail string, at time.Time) (result firewall.FirewallAction, resultErr error) {
defer func() { recordFirewallWriteError(resultErr) }()
resultErr = db.updateFirewallSnapshot(1, func(tx *bolt.Tx) error {
index, indexErr := readFirewallJournalIndex(tx)
if indexErr != nil {
return indexErr
}
a, err := readFirewallAction(tx, id)
if err != nil {
return err
}
if firewallActionPending(a.Phase) {
if index.PendingID != id {
return firewall.ErrStateCorrupt
}
if err := validateFirewallActionBase(tx, a); err != nil {
return err
}
}
if a.Phase == phase && a.Detail == detail {
result = a
return nil
}
if !validFirewallActionTransition(a.Phase, phase) || at.IsZero() || a.AuditVersion == math.MaxUint64 {
return fmt.Errorf("%w: invalid action transition", firewall.ErrStateConflict)
}
if phase == "verified" || phase == "failed" {
state := a.Before
if phase == "verified" {
state = a.After
}
meta, rows, err := encodeFirewallSnapshot(state)
if err != nil {
return err
}
if a.Revision == math.MaxUint64 {
return firewall.ErrStateConflict
}
meta.Revision = a.Revision + 1
raw, err := json.Marshal(meta)
if err != nil {
return err
}
if err := replaceFirewallSnapshot(tx, a.Revision, meta, rows, raw); err != nil {
return err
}
}
a.Phase, a.Detail, a.UpdatedAt = phase, detail, at
if err := validateFirewallAdmission(a); err != nil {
return err
}
a.AuditVersion++
size, writeErr := writeFirewallAction(tx, a)
if writeErr != nil {
return writeErr
}
if err := writeFirewallAuditEvent(tx, a); err != nil {
return err
}
if !firewallActionPending(a.Phase) {
index.PendingID = ""
}
if firewallAuditPhase(a.Phase) {
index.Audit = append(index.Audit, firewallAuditReference{ID: id, Version: a.AuditVersion})
}
if err := writeFirewallJournalIndex(tx, index); err != nil {
return err
}
if !firewallActionPending(a.Phase) {
if err := recordFirewallActionHistory(tx, a, size); err != nil {
return err
}
if err := pruneFirewallActionHistory(tx, index); err != nil {
return err
}
}
result = a
return nil
})
if resultErr != nil {
return firewall.FirewallAction{}, resultErr
}
return result, nil
}
func (db *DB) AcknowledgeFirewallAudit(id string, version uint64) (err error) {
defer func() { recordFirewallWriteError(err) }()
return db.updateFirewallSnapshot(1, func(tx *bolt.Tx) error {
index, indexErr := readFirewallJournalIndex(tx)
if indexErr != nil {
return indexErr
}
position := -1
for i, ref := range index.Audit {
if ref.ID == id && ref.Version == version {
position = i
break
}
}
a, err := readFirewallAction(tx, id)
if err != nil {
return err
}
if version == 0 || version > a.AuditVersion {
return firewall.ErrStateConflict
}
b := tx.Bucket([]byte(firewallAuditBucket))
key := firewallAuditKey(id, version)
if b == nil || b.Get(key) == nil {
if position >= 0 {
return firewall.ErrStateCorrupt
}
return firewall.ErrStateConflict
}
event, eventAction, err := decodeFirewallAuditEvent(key, b.Get(key))
if err != nil {
return err
}
if eventAction.Request != a.Request {
return firewall.ErrStateCorrupt
}
if event.Acknowledged {
if position >= 0 {
return firewall.ErrStateCorrupt
}
return nil
}
if position < 0 {
return firewall.ErrStateCorrupt
}
event.Acknowledged = true
raw, err := encodeFirewallJournal(event)
if err != nil {
return err
}
if err := b.Put(key, raw); err != nil {
return err
}
a.AuditAck = max(a.AuditAck, version)
size, writeErr := writeFirewallAction(tx, a)
if writeErr != nil {
return writeErr
}
if err := updateFirewallActionHistorySize(tx, a, size); err != nil {
return err
}
index.Audit = append(index.Audit[:position], index.Audit[position+1:]...)
if err := writeFirewallJournalIndex(tx, index); err != nil {
return err
}
return pruneFirewallActionHistory(tx, index)
})
}
package store
import (
"encoding/json"
"fmt"
"os"
"time"
bolt "go.etcd.io/bbolt"
)
// FirewallRollback is a pending tentative-apply record. The previous
// csm.yaml bytes are stashed verbatim so a recovery path can restore the
// file byte-for-byte without re-rendering through any encoder. Hashes are
// recorded so the daemon can sanity-check the on-disk file matches what
// was applied before deciding to revert.
type FirewallRollback struct {
PrevYAML []byte `json:"prev_yaml"`
PrevHash string `json:"prev_hash"`
NewHash string `json:"new_hash"`
AppliedAt time.Time `json:"applied_at"`
ExpiresAt time.Time `json:"expires_at"`
AppliedBy string `json:"applied_by"`
}
const fwRollbackBucket = "fw:rollback"
const fwRollbackKey = "pending"
// SaveFirewallRollback writes a pending rollback record. Overwrites any
// existing pending entry; callers must clear or revert the previous one
// first if that matters for their flow.
func (db *DB) SaveFirewallRollback(rb FirewallRollback) error {
val, err := json.Marshal(rb)
if err != nil {
return err
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(fwRollbackBucket))
return b.Put([]byte(fwRollbackKey), val)
})
}
// GetFirewallRollback returns the pending rollback or (zero, false) if
// none. The bool distinguishes "no record" from a zero-valued record.
// A bbolt unmarshal failure is treated as "no usable record" so the
// daemon can skip a corrupt entry instead of refusing to start.
func (db *DB) GetFirewallRollback() (FirewallRollback, bool) {
var rb FirewallRollback
found := false
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(fwRollbackBucket))
val := b.Get([]byte(fwRollbackKey))
if val == nil {
return nil
}
if uerr := json.Unmarshal(val, &rb); uerr != nil {
// Corrupt record: leave found=false so the caller treats
// it as "no pending rollback" rather than panicking.
return nil //nolint:nilerr // swallowing is the intent; see comment.
}
found = true
return nil
})
return rb, found
}
// ClearFirewallRollback drops the pending rollback. Idempotent: deleting
// a non-existent key is not an error in bbolt.
func (db *DB) ClearFirewallRollback() error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(fwRollbackBucket))
return b.Delete([]byte(fwRollbackKey))
})
}
// DisarmFirewallRollbackSnapshot removes process-lifetime rollback intent
// from a stopped, private bbolt snapshot. Backup and restore artifacts keep
// durable firewall state, but must not make an old tentative configuration
// write active again on a later daemon start.
func DisarmFirewallRollbackSnapshot(path string) error {
if _, err := os.Stat(path); err != nil {
return err
}
snapshot, err := bolt.Open(path, 0o600, nil)
if err != nil {
return fmt.Errorf("opening state snapshot: %w", err)
}
pending := false
if err := snapshot.View(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte(fwRollbackBucket))
pending = bucket != nil && bucket.Get([]byte(fwRollbackKey)) != nil
return nil
}); err != nil {
_ = snapshot.Close()
return fmt.Errorf("checking firewall rollback snapshot: %w", err)
}
if !pending {
if err := snapshot.Close(); err != nil {
return fmt.Errorf("closing state snapshot: %w", err)
}
return nil
}
if err := snapshot.Update(func(tx *bolt.Tx) error {
bucket := tx.Bucket([]byte(fwRollbackBucket))
if bucket == nil {
return nil
}
return bucket.Delete([]byte(fwRollbackKey))
}); err != nil {
_ = snapshot.Close()
return fmt.Errorf("disarming firewall rollback snapshot: %w", err)
}
if err := snapshot.Close(); err != nil {
return fmt.Errorf("closing state snapshot: %w", err)
}
return nil
}
package store
import (
"bytes"
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"encoding/json"
"fmt"
"math"
"strconv"
"time"
"unicode/utf8"
"github.com/pidginhost/csm/internal/firewall"
bolt "go.etcd.io/bbolt"
)
var _ firewall.StateStore = (*DB)(nil)
const firewallSnapshotBucket = "fw:state"
const firewallSnapshotKey = "snapshot"
const firewallSnapshotVersion = 1
// Fixed collection order is part of the versioned storage format. Metadata
// lives under fw:* so a future atomic firewall restore can include it.
var firewallSnapshotBuckets = [...]string{"fw:blocked", "fw:subnets", "fw:allowed", "fw:port_allowed"}
type firewallCollection struct {
// Keys preserve engine order, including distinct entries for the same IP.
// Nil and empty keys preserve nil and empty domain slices respectively.
Keys []string `json:"keys"`
Digest string `json:"sha256"`
}
type firewallSnapshotMeta struct {
Version int `json:"version"`
Revision uint64 `json:"revision"`
Collections []firewallCollection `json:"collections"`
}
type firewallSnapshotRows [4][][]byte
// ReadFirewallState copies a single consistent snapshot before decoding domain
// values. No bbolt-owned bytes or partially decoded state escape this method.
func (db *DB) ReadFirewallState() (result firewall.FirewallState, revision uint64, err error) {
started := time.Now()
defer func() {
firewallReadDuration.Observe(time.Since(started).Seconds())
if err != nil {
firewallReadFailures.Inc()
}
}()
var meta firewallSnapshotMeta
var rows firewallSnapshotRows
err = db.bolt.View(func(tx *bolt.Tx) error {
var readErr error
meta, readErr = readFirewallSnapshotMeta(tx)
if readErr != nil {
return readErr
}
rows, readErr = readFirewallSnapshotRows(tx, meta)
return readErr
})
if err != nil {
return firewall.FirewallState{}, 0, err
}
state, err := decodeFirewallSnapshot(rows, meta)
if err != nil {
return firewall.FirewallState{}, 0, err
}
return state, meta.Revision, nil
}
// ReplaceFirewallState encodes outside the writer lock and publishes state and
// revision together. Existing runtime callers continue using their current
// backend; this API is not an implicit migration or a cutover marker.
func (db *DB) ReplaceFirewallState(expectedRevision uint64, state firewall.FirewallState) (revision uint64, err error) {
defer func() { recordFirewallWriteError(err) }()
if expectedRevision == math.MaxUint64 {
return 0, fmt.Errorf("%w: revision exhausted", firewall.ErrStateConflict)
}
meta, rows, err := encodeFirewallSnapshot(state)
if err != nil {
return 0, err
}
meta.Revision = expectedRevision + 1
encoded, err := json.Marshal(meta)
if err != nil {
return 0, err
}
err = db.updateFirewallSnapshot(len(state.Blocked)+len(state.BlockedNet)+len(state.Allowed)+len(state.PortAllowed), func(tx *bolt.Tx) error {
if pendingErr := refusePendingFirewallActions(tx); pendingErr != nil {
return pendingErr
}
return replaceFirewallSnapshot(tx, expectedRevision, meta, rows, encoded)
})
if err != nil {
return 0, err
}
return meta.Revision, nil
}
// replaceFirewallSnapshot is private so future action admission can share this
// transaction without exporting CRUD or transaction callbacks to domain code.
func replaceFirewallSnapshot(tx *bolt.Tx, expected uint64, next firewallSnapshotMeta, rows firewallSnapshotRows, encoded []byte) error {
current, metaErr := readFirewallSnapshotMeta(tx)
if metaErr == firewall.ErrStateUninitialized {
if expected != 0 {
return firewall.ErrStateConflict
}
} else {
if metaErr != nil {
return metaErr
}
if current.Revision != expected {
return firewall.ErrStateConflict
}
// An out-of-band legacy bucket write cannot silently bypass the revision.
currentRows, readErr := readFirewallSnapshotRows(tx, current)
if readErr != nil {
return readErr
}
if _, err := decodeFirewallSnapshot(currentRows, current); err != nil {
return err
}
}
for i, name := range firewallSnapshotBuckets {
if tx.Bucket([]byte(name)) != nil {
if err := tx.DeleteBucket([]byte(name)); err != nil {
return err
}
}
b, err := tx.CreateBucket([]byte(name))
if err != nil {
return err
}
for j, key := range next.Collections[i].Keys {
if err := b.Put([]byte(key), rows[i][j]); err != nil {
return err
}
}
}
b, err := tx.CreateBucketIfNotExists([]byte(firewallSnapshotBucket))
if err != nil {
return err
}
return b.Put([]byte(firewallSnapshotKey), encoded)
}
func readFirewallSnapshotMeta(tx *bolt.Tx) (firewallSnapshotMeta, error) {
var meta firewallSnapshotMeta
b := tx.Bucket([]byte(firewallSnapshotBucket))
if b == nil {
return meta, firewall.ErrStateUninitialized
}
raw := b.Get([]byte(firewallSnapshotKey))
if raw == nil {
return meta, fmt.Errorf("%w: missing metadata", firewall.ErrStateCorrupt)
}
if err := json.Unmarshal(raw, &meta); err != nil {
return meta, fmt.Errorf("%w: metadata decode: %v", firewall.ErrStateCorrupt, err)
}
if meta.Version != firewallSnapshotVersion || meta.Revision == 0 || len(meta.Collections) != len(firewallSnapshotBuckets) {
return meta, fmt.Errorf("%w: metadata version, revision or collections", firewall.ErrStateCorrupt)
}
return meta, nil
}
func readFirewallSnapshotRows(tx *bolt.Tx, meta firewallSnapshotMeta) (firewallSnapshotRows, error) {
var rows firewallSnapshotRows
for i, name := range firewallSnapshotBuckets {
b := tx.Bucket([]byte(name))
if b == nil {
return rows, fmt.Errorf("%w: missing collection %s", firewall.ErrStateCorrupt, name)
}
collection := meta.Collections[i]
expected := make(map[string]struct{}, len(collection.Keys))
for _, key := range collection.Keys {
if _, exists := expected[key]; exists || key == "" {
return rows, fmt.Errorf("%w: invalid collection keys", firewall.ErrStateCorrupt)
}
expected[key] = struct{}{}
raw := b.Get([]byte(key))
if raw == nil {
return rows, fmt.Errorf("%w: missing row in %s", firewall.ErrStateCorrupt, name)
}
rows[i] = append(rows[i], bytes.Clone(raw))
}
if err := b.ForEach(func(k, v []byte) error {
if _, ok := expected[string(k)]; !ok || v == nil {
return fmt.Errorf("%w: unexpected row in %s", firewall.ErrStateCorrupt, name)
}
return nil
}); err != nil {
return rows, err
}
if digestFirewallCollection(collection.Keys, rows[i]) != collection.Digest {
return rows, fmt.Errorf("%w: changed collection %s", firewall.ErrStateCorrupt, name)
}
}
return rows, nil
}
func digestFirewallCollection(keys []string, rows [][]byte) string {
hash := sha256.New()
var size [8]byte
for i, key := range keys {
binary.BigEndian.PutUint64(size[:], uint64(len(key)))
_, _ = hash.Write(size[:])
_, _ = hash.Write([]byte(key))
binary.BigEndian.PutUint64(size[:], uint64(len(rows[i])))
_, _ = hash.Write(size[:])
_, _ = hash.Write(rows[i])
}
return hex.EncodeToString(hash.Sum(nil))
}
func encodeFirewallSnapshot(state firewall.FirewallState) (firewallSnapshotMeta, firewallSnapshotRows, error) {
meta := firewallSnapshotMeta{Version: firewallSnapshotVersion, Collections: make([]firewallCollection, 4)}
var rows firewallSnapshotRows
var err error
meta.Collections[0], rows[0], err = encodeFirewallCollection(state.Blocked, func(e firewall.BlockedEntry) string { return e.IP })
if err != nil {
return meta, rows, err
}
meta.Collections[1], rows[1], err = encodeFirewallCollection(state.BlockedNet, func(e firewall.SubnetEntry) string { return e.CIDR })
if err != nil {
return meta, rows, err
}
meta.Collections[2], rows[2], err = encodeFirewallCollection(state.Allowed, func(e firewall.AllowedEntry) string { return e.IP })
if err != nil {
return meta, rows, err
}
meta.Collections[3], rows[3], err = encodeFirewallCollection(state.PortAllowed, func(e firewall.PortAllowEntry) string { return portAllowKey(e.IP, e.Port, e.Proto) })
return meta, rows, err
}
func encodeFirewallCollection[T any](entries []T, identity func(T) string) (firewallCollection, [][]byte, error) {
var collection firewallCollection
var rows [][]byte
if entries != nil {
collection.Keys = make([]string, 0, len(entries))
}
used := make(map[string]bool, len(entries))
for i, entry := range entries {
key := identity(entry)
if key == "" || used[key] {
key = fmt.Sprintf("\x00%d", i)
}
if used[key] {
return collection, nil, fmt.Errorf("invalid firewall row identity")
}
used[key] = true
if !validFirewallRowEncoding(entry) {
return collection, nil, fmt.Errorf("firewall row cannot be encoded losslessly")
}
row, err := json.Marshal(entry)
if err != nil {
return collection, nil, fmt.Errorf("encode firewall row: %w", err)
}
collection.Keys = append(collection.Keys, key)
rows = append(rows, row)
}
collection.Digest = digestFirewallCollection(collection.Keys, rows)
return collection, rows, nil
}
func decodeFirewallSnapshot(rows firewallSnapshotRows, meta firewallSnapshotMeta) (firewall.FirewallState, error) {
var state firewall.FirewallState
var err error
state.Blocked, err = decodeFirewallCollection[firewall.BlockedEntry](rows[0], meta.Collections[0].Keys != nil)
if err != nil {
return firewall.FirewallState{}, err
}
state.BlockedNet, err = decodeFirewallCollection[firewall.SubnetEntry](rows[1], meta.Collections[1].Keys != nil)
if err != nil {
return firewall.FirewallState{}, err
}
state.Allowed, err = decodeFirewallCollection[firewall.AllowedEntry](rows[2], meta.Collections[2].Keys != nil)
if err != nil {
return firewall.FirewallState{}, err
}
state.PortAllowed, err = decodeFirewallCollection[firewall.PortAllowEntry](rows[3], meta.Collections[3].Keys != nil)
if err != nil {
return firewall.FirewallState{}, err
}
return state, nil
}
func decodeFirewallCollection[T any](rows [][]byte, nonNil bool) ([]T, error) {
var entries []T
if nonNil {
entries = make([]T, 0, len(rows))
}
for _, row := range rows {
var entry T
// A matching checksum does not make malformed text decodable. JSON
// silently replaces invalid Unicode, which would alter stored evidence.
if !validFirewallJSONText(row) {
return nil, fmt.Errorf("%w: row is not valid Unicode", firewall.ErrStateCorrupt)
}
trimmed := bytes.TrimSpace(row)
if len(trimmed) == 0 || trimmed[0] != '{' {
return nil, fmt.Errorf("%w: row is not an object", firewall.ErrStateCorrupt)
}
if err := json.Unmarshal(row, &entry); err != nil {
return nil, fmt.Errorf("%w: row decode: %v", firewall.ErrStateCorrupt, err)
}
entries = append(entries, entry)
}
return entries, nil
}
// The JSON decoder checks syntax but accepts invalid UTF-8 and unpaired UTF-16
// surrogate escapes. Check those first without rejecting literal backslashes
// or valid surrogate pairs. Non-Unicode escapes are left to the JSON decoder.
func validFirewallJSONText(raw []byte) bool {
if !utf8.Valid(raw) {
return false
}
for i := 0; i < len(raw); i++ {
if raw[i] != '\\' {
continue
}
i++
if i >= len(raw) {
return false
}
if raw[i] != 'u' {
continue
}
if i+4 >= len(raw) {
return false
}
code, err := strconv.ParseUint(string(raw[i+1:i+5]), 16, 16)
if err != nil || code >= 0xdc00 && code <= 0xdfff {
return false
}
i += 4
if code < 0xd800 || code > 0xdbff {
continue
}
if i+6 >= len(raw) || raw[i+1] != '\\' || raw[i+2] != 'u' {
return false
}
low, err := strconv.ParseUint(string(raw[i+3:i+7]), 16, 16)
if err != nil || low < 0xdc00 || low > 0xdfff {
return false
}
i += 6
}
return true
}
// encoding/json replaces invalid UTF-8 and rounds timezone offsets to minutes.
// Reject either loss before a successful write can acknowledge altered evidence.
func validFirewallRowEncoding(entry any) bool {
var fields []string
var times []time.Time
switch e := entry.(type) {
case firewall.BlockedEntry:
times = []time.Time{e.BlockedAt, e.ExpiresAt}
fields = []string{e.IP, e.Reason, e.Source}
case firewall.SubnetEntry:
times = []time.Time{e.BlockedAt, e.ExpiresAt}
fields = []string{e.CIDR, e.Reason, e.Source}
case firewall.AllowedEntry:
times = []time.Time{e.ExpiresAt}
fields = []string{e.IP, e.Reason, e.Source}
case firewall.PortAllowEntry:
fields = []string{e.IP, e.Proto, e.Reason, e.Source}
default:
return false
}
for _, timestamp := range times {
_, offset := timestamp.Zone()
if offset%60 != 0 {
return false
}
}
for _, field := range fields {
if !utf8.ValidString(field) {
return false
}
}
return true
}
package store
import (
"errors"
"fmt"
"time"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/metrics"
bolt "go.etcd.io/bbolt"
)
// These collectors have no labels: paths, addresses and caller identities never
// allocate time series. They cover only the new firewall snapshot contract.
var (
firewallWriteWait = metrics.NewHistogram("csm_storage_firewall_write_wait_seconds", "Firewall writer lock wait.", []float64{.001, .01, .1, 1, 10})
firewallWriteTransaction = metrics.NewHistogram("csm_storage_firewall_write_transaction_seconds", "Firewall write transaction through commit or rollback, excluding writer wait.", []float64{.001, .01, .1, 1, 10})
firewallReadDuration = metrics.NewHistogram("csm_storage_firewall_read_seconds", "Firewall snapshot read including copying and decoding.", []float64{.001, .01, .1, 1, 10})
firewallBatchRows = metrics.NewHistogram("csm_storage_firewall_batch_rows", "Rows in attempted firewall snapshot transactions.", []float64{0, 100, 1000, 10000, 100000})
firewallPendingWrites = metrics.NewGauge("csm_storage_firewall_pending_writes", "Firewall writes waiting or executing.")
firewallWriteFailures = metrics.NewCounter("csm_storage_firewall_write_failures_total", "Firewall replacements returning errors, including admission refusals.")
firewallReadFailures = metrics.NewCounter("csm_storage_firewall_read_failures_total", "Firewall reads returning errors.")
firewallCommitFailures = metrics.NewCounter("csm_storage_firewall_commit_failures_total", "Firewall transactions accepted by the callback without a confirmed commit.")
firewallConflicts = metrics.NewCounter("csm_storage_firewall_conflicts_total", "Firewall replacements refused due to revision conflicts.")
)
func init() {
metrics.MustRegister("csm_storage_firewall_write_wait_seconds", firewallWriteWait)
metrics.MustRegister("csm_storage_firewall_write_transaction_seconds", firewallWriteTransaction)
metrics.MustRegister("csm_storage_firewall_read_seconds", firewallReadDuration)
metrics.MustRegister("csm_storage_firewall_batch_rows", firewallBatchRows)
metrics.MustRegister("csm_storage_firewall_pending_writes", firewallPendingWrites)
metrics.MustRegister("csm_storage_firewall_write_failures_total", firewallWriteFailures)
metrics.MustRegister("csm_storage_firewall_read_failures_total", firewallReadFailures)
metrics.MustRegister("csm_storage_firewall_commit_failures_total", firewallCommitFailures)
metrics.MustRegister("csm_storage_firewall_conflicts_total", firewallConflicts)
}
func recordFirewallWriteError(err error) {
if err != nil {
firewallWriteFailures.Inc()
}
if errors.Is(err, firewall.ErrStateConflict) {
firewallConflicts.Inc()
}
}
func (db *DB) updateFirewallSnapshot(rowCount int, fn func(*bolt.Tx) error) error {
firewallBatchRows.Observe(float64(rowCount))
firewallPendingWrites.Inc()
defer firewallPendingWrites.Dec()
waiting := time.Now()
var acquired time.Time
accepted := false
err := boltUpdate(db.bolt, func(tx *bolt.Tx) error {
acquired = time.Now()
firewallWriteWait.Observe(acquired.Sub(waiting).Seconds())
callbackErr := fn(tx)
accepted = callbackErr == nil
return callbackErr
})
if acquired.IsZero() {
// Begin failed: elapsed time belongs to acquisition, not a transaction.
firewallWriteWait.Observe(time.Since(waiting).Seconds())
} else {
firewallWriteTransaction.Observe(time.Since(acquired).Seconds())
}
if err != nil && accepted {
firewallCommitFailures.Inc()
// bbolt can return a sync error after publishing its new metadata.
// An accepted callback plus an error is not proof of rollback.
return fmt.Errorf("%w: %w", firewall.ErrStateCommitUncertain, err)
}
return err
}
package store
import (
"encoding/json"
"time"
bolt "go.etcd.io/bbolt"
)
// AuditResult represents a single hardening check result.
type AuditResult struct {
Category string `json:"category"`
Name string `json:"name"`
Title string `json:"title"`
Status string `json:"status"`
Message string `json:"message"`
Fix string `json:"fix,omitempty"`
}
// AuditReport is the full result of a hardening audit run.
type AuditReport struct {
Timestamp time.Time `json:"timestamp,omitzero"`
ServerType string `json:"server_type"`
Results []AuditResult `json:"results"`
Score int `json:"score"`
Total int `json:"total"`
}
const hardeningReportKey = "hardening:report"
// SaveHardeningReport persists the latest audit report in the meta bucket.
func (db *DB) SaveHardeningReport(report *AuditReport) error {
data, err := json.Marshal(report)
if err != nil {
return err
}
return db.bolt.Update(func(tx *bolt.Tx) error {
return tx.Bucket([]byte("meta")).Put([]byte(hardeningReportKey), data)
})
}
// LoadHardeningReport retrieves the latest audit report from the meta bucket.
// Returns a zero-value report (nil Results) if no report has been saved yet.
func (db *DB) LoadHardeningReport() (*AuditReport, error) {
var report AuditReport
err := db.bolt.View(func(tx *bolt.Tx) error {
v := tx.Bucket([]byte("meta")).Get([]byte(hardeningReportKey))
if v == nil {
return nil
}
return json.Unmarshal(v, &report)
})
return &report, err
}
package store
import (
"bytes"
"encoding/json"
"fmt"
"strings"
"time"
"unicode/utf8"
"github.com/pidginhost/csm/internal/alert"
bolt "go.etcd.io/bbolt"
)
// maxHistoryEntries is the maximum number of history entries to retain.
// It is a var (not const) so tests can override it.
var maxHistoryEntries = 100_000
// AppendHistory inserts findings into the history bucket with TimeKey keys.
// It increments the history:count counter and prunes oldest entries if the
// count exceeds maxHistoryEntries.
func (db *DB) AppendHistory(findings []alert.Finding) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
writer := newTimeKeyWriter(b)
for i, f := range findings {
val, err := json.Marshal(alert.SanitizeFinding(f))
if err != nil {
return err
}
key := nextHistoryKey(b, f.Timestamp, i)
if err := writer.put([]byte(key), val); err != nil {
return err
}
// Same transaction as the history insert: either both land or
// neither does, so the daily aggregate can never drift.
if err := incrStatsDaily(tx, f.Timestamp, f.Severity); err != nil {
return err
}
if err := bumpLatestByCheck(tx, f.Check, f.Timestamp); err != nil {
return err
}
}
writer.settle()
if err := incrCounter(tx, "history:count", len(findings)); err != nil {
return err
}
// Cheap sweep against the bounded stats:daily bucket. Done here
// (rather than on a timer) so the daily-aggregate path has a
// single owner.
if len(findings) > 0 {
if err := bumpHistoryRevision(tx); err != nil {
return err
}
if err := pruneStatsDaily(tx, time.Now()); err != nil {
return err
}
}
// Prune oldest entries if count exceeds maxHistoryEntries.
meta := tx.Bucket([]byte("meta"))
var count int
if v := meta.Get([]byte("history:count")); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", &count)
}
if count > maxHistoryEntries {
excess := count - maxHistoryEntries
c := b.Cursor()
k, _ := c.First()
for ; k != nil && excess > 0; excess-- {
// Delete() moves the cursor to the next item, so we
// must NOT call c.Next() after it.
if err := c.Delete(); err != nil {
return err
}
k, _ = c.First()
}
if err := setCounter(tx, "history:count", maxHistoryEntries); err != nil {
return err
}
}
return nil
})
}
func nextHistoryKey(b *bolt.Bucket, timestamp time.Time, start int) string {
for counter := start; ; counter++ {
key := TimeKey(timestamp, counter)
// Shutdown drains can persist separate batches whose findings carry
// the same detector timestamp. Probe instead of overwriting history.
if b.Get([]byte(key)) == nil {
return key
}
}
}
func bumpHistoryRevision(tx *bolt.Tx) error {
return incrCounter(tx, "history:revision", 1)
}
// HistoryMark changes with history mutations, even when retention and a later
// append reuse the same keys or a migration changes only interior keys.
func (db *DB) HistoryMark() string {
var mark string
_ = db.bolt.View(func(tx *bolt.Tx) error {
mark = string(tx.Bucket([]byte("meta")).Get([]byte("history:revision")))
return nil
})
return mark
}
// ReadHistory reads findings from the history bucket, newest-first.
// It returns up to limit findings starting at offset, plus the total count.
func (db *DB) ReadHistory(limit, offset int) ([]alert.Finding, int) {
total := db.getCounter("history:count")
var results []alert.Finding
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
c := b.Cursor()
// Skip offset entries from the end (newest first).
skipped := 0
k, v := c.Last()
for ; k != nil && skipped < offset; k, v = c.Prev() {
skipped++
}
// Collect up to limit entries.
for ; k != nil && len(results) < limit; k, v = c.Prev() {
var f alert.Finding
if err := json.Unmarshal(v, &f); err == nil {
results = append(results, f)
}
}
return nil
})
return results, total
}
// ReadHistoryFiltered reads findings with optional filtering.
// Parameters:
// - from, to: calendar dates or RFC 3339 instants (empty to skip)
// - severity: filter by severity level (-1 for no filter)
// - search: case-insensitive substring match on check/message/details (empty to skip)
func (db *DB) ReadHistoryFiltered(limit, offset int, from, to string, severity int, search string) ([]alert.Finding, int) {
return db.ReadHistoryFilteredWithChecks(limit, offset, from, to, severity, search, nil)
}
// ParseHistoryBound reads one end of a history date range. A calendar date
// (YYYY-MM-DD) names a server-local day: as a start it is that day's
// midnight, as an end the next day's, so the whole day is included. An RFC
// 3339 instant is used as given; as an end it is exclusive. An empty bound
// is the zero time.
func ParseHistoryBound(s string, end bool) (time.Time, error) {
return parseHistoryBoundIn(s, end, time.Local)
}
func parseHistoryBoundIn(s string, end bool, loc *time.Location) (time.Time, error) {
s = strings.TrimSpace(s)
if s == "" {
return time.Time{}, nil
}
if t, err := time.Parse(time.RFC3339, s); err == nil {
return t, nil
}
// Parse the calendar in UTC so a missing local midnight cannot normalize
// the date into the preceding day before we choose the exclusive end.
day, err := time.Parse("2006-01-02", s)
if err != nil {
return time.Time{}, fmt.Errorf("date %q is neither YYYY-MM-DD nor RFC 3339", s)
}
if end {
day = day.AddDate(0, 0, 1)
}
// Walk the surrounding zone intervals to find the earliest instant of
// the day. A repeated midnight uses its first occurrence; a skipped
// midnight (or date) starts at the transition into the next valid time.
for at := day.Add(-48 * time.Hour).In(loc); ; {
_, offset := at.Zone()
candidate := day.Add(-time.Duration(offset) * time.Second).In(loc)
if candidate.Before(at) {
candidate = at
}
_, zoneEnd := at.ZoneBounds()
if zoneEnd.IsZero() || candidate.Before(zoneEnd) {
return candidate, nil
}
at = zoneEnd
}
}
// decodeHistoryEntry decodes one stored finding; a var so tests can count
// decodes.
var decodeHistoryEntry = func(v []byte, f *alert.Finding) error { return json.Unmarshal(v, f) }
// ReadHistoryFilteredWithChecks reads findings with optional filters, including
// an exact check-name set when checks is non-nil.
func (db *DB) ReadHistoryFilteredWithChecks(
limit, offset int,
from, to string,
severity int,
search string,
checks map[string]bool,
) ([]alert.Finding, int) {
var results []alert.Finding
matched := 0
searchLower := strings.ToLower(search)
from, to = strings.TrimSpace(from), strings.TrimSpace(to)
mayMatch := historyPrefilter(severity, searchLower, checks)
var fromPrefix, toPrefix string
if from != "" {
if fromTime, err := ParseHistoryBound(from, false); err == nil {
fromPrefix = timeKeyLowerBound(fromTime)
} else {
fromPrefix = ParseTimeKeyPrefix(from)
}
}
if to != "" {
if toTime, err := ParseHistoryBound(to, true); err == nil {
// An exclusive upper bound: the next local midnight for a
// date, the instant itself for an RFC 3339 end.
toPrefix = timeKeyLowerBound(toTime)
} else {
toPrefix = ParseTimeKeyPrefix(to) + "99"
}
}
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
c := b.Cursor()
// Seek to the first key at the exclusive upper bound, then step back so the
// descending walk starts at the newest in-range entry. Without this the
// loop walked (and skipped) every entry newer than `to`, which is O(N)
// of the whole bucket when querying an old range on a large history.
k, v := c.Last()
if toPrefix != "" {
if sk, _ := c.Seek([]byte(toPrefix)); sk != nil {
// Seek lands on the first key >= toPrefix (out of range);
// the previous key is the newest in-range entry.
k, v = c.Prev()
} else {
// All keys are below toPrefix; Last() is already in range.
k, v = c.Last()
}
}
for ; k != nil; k, v = c.Prev() {
key := string(k)
// Defensive: anything still above the upper bound is out of range.
if toPrefix != "" && key >= toPrefix {
continue
}
// Time-range: if key is below fromPrefix, all remaining are older - stop.
if fromPrefix != "" && key < fromPrefix {
break
}
if !mayMatch(v) {
continue
}
var f alert.Finding
if err := decodeHistoryEntry(v, &f); err != nil {
continue
}
// Severity filter.
if severity >= 0 && int(f.Severity) != severity {
continue
}
// Exact check-name filter.
if checks != nil && !checks[f.Check] {
continue
}
// Search filter.
if search != "" && !containsLower(f.Check, searchLower) &&
!containsLower(f.Message, searchLower) &&
!containsLower(f.Details, searchLower) {
continue
}
matched++
if matched > offset && len(results) < limit {
results = append(results, f)
}
}
return nil
})
return results, matched
}
// historyPrefilter returns a test on a stored entry's JSON that is false only
// when the entry cannot pass the severity, check or search filter, so the
// walk that counts matches decodes only plausible entries. Each part applies
// only when the text it looks for is stored verbatim: JSON escapes quotes,
// backslashes, control characters, <, > and &, so a search for those is
// left to the decoded check.
func historyPrefilter(severity int, searchLower string, checks map[string]bool) func([]byte) bool {
var sevNeedles [][]byte
if severity >= 0 {
sevNeedles = [][]byte{
[]byte(fmt.Sprintf(`"severity":%d,`, severity)),
[]byte(fmt.Sprintf(`"severity":%d}`, severity)),
}
}
var checkNeedles [][]byte
activeChecks := 0
for _, enabled := range checks {
if enabled {
activeChecks++
}
}
if checks != nil && activeChecks == 0 {
return func([]byte) bool { return false }
}
for name, enabled := range checks {
if !enabled {
continue
}
// The empty check also matches an absent or null field.
if name == "" || !storedVerbatim(name) {
checkNeedles = nil
break
}
checkNeedles = append(checkNeedles, []byte(`"check":"`+name+`"`))
}
var searchNeedle []byte
if searchLower != "" && storedVerbatim(searchLower) {
searchNeedle = []byte(searchLower)
}
return func(v []byte) bool {
if sevNeedles != nil && !bytes.Contains(v, sevNeedles[0]) && !bytes.Contains(v, sevNeedles[1]) {
// Zero also accepts omitted/null fields and negative zero. Keep
// the fast rejection for ordinary nonzero severity values.
if severity == 0 && (!bytes.Contains(v, []byte(`"severity":`)) ||
bytes.Contains(v, []byte(`"severity":null`)) || bytes.Contains(v, []byte(`"severity":-0`))) {
return true
}
return !plainHistoryJSON(v)
}
if checks != nil && checkNeedles != nil && !containsAnyBytes(v, checkNeedles) {
return !plainHistoryJSON(v)
}
if searchNeedle != nil && !bytes.Contains(bytes.ToLower(v), searchNeedle) {
return !plainHistoryJSON(v)
}
return true
}
}
// Raw needles are conclusive only for unescaped compact JSON with canonical
// field names. Other valid encodings (including case-insensitive JSON keys)
// are left to the decoder. This is an eligibility check, not a JSON validator.
func plainHistoryJSON(v []byte) bool {
if len(v) < 2 || v[0] != '{' || v[len(v)-1] != '}' {
return false
}
start := -1
for i, b := range v {
if b == '\\' {
return false
}
if b == '"' {
if start < 0 {
start = i + 1
} else {
if i+1 < len(v) && v[i+1] == ':' {
for _, c := range v[start:i] {
if (c < 'a' || c > 'z') && (c < '0' || c > '9') && c != '_' {
return false
}
}
}
start = -1
}
} else if start < 0 && (b == ' ' || b == '\t' || b == '\n' || b == '\r') {
return false
}
}
return true
}
// storedVerbatim reports whether encoding/json writes s unchanged inside a
// JSON string.
func storedVerbatim(s string) bool {
for _, r := range s {
if r < 0x20 || r == '"' || r == '\\' || r == '<' || r == '>' || r == '&' || r == '\u2028' || r == '\u2029' || r == utf8.RuneError {
return false
}
}
return true
}
func containsAnyBytes(v []byte, needles [][]byte) bool {
for _, n := range needles {
if bytes.Contains(v, n) {
return true
}
}
return false
}
// substr must already be lowercase.
func containsLower(s, substr string) bool {
return strings.Contains(strings.ToLower(s), substr)
}
package store
import (
"encoding/json"
"errors"
"fmt"
"sort"
"time"
bolt "go.etcd.io/bbolt"
"github.com/pidginhost/csm/internal/incident"
csmlog "github.com/pidginhost/csm/internal/log"
)
const incidentsBucket = "incidents"
// Bound rows inspected as well as deleted: retained or corrupt rows must
// not turn a sweep into one long transaction on bbolt's shared writer.
const incidentCompactionBatchSize = 256
// SaveIncident persists an incident, overwriting any prior record with
// the same ID. Caller is responsible for setting UpdatedAt before
// invoking; this method just writes.
func (db *DB) SaveIncident(inc incident.Incident) error {
data, err := json.Marshal(inc)
if err != nil {
return err
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(incidentsBucket))
return b.Put([]byte(inc.ID), data)
})
}
// GetIncident returns (incident, true, nil) if found, (zero, false, nil)
// if not, (zero, false, err) on store error.
func (db *DB) GetIncident(id string) (incident.Incident, bool, error) {
var (
inc incident.Incident
found bool
)
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(incidentsBucket))
v := b.Get([]byte(id))
if v == nil {
return nil
}
if err := json.Unmarshal(v, &inc); err != nil {
return err
}
if err := validateIncidentRow(id, inc); err != nil {
return err
}
found = true
return nil
})
return inc, found, err
}
// ListIncidents returns every stored incident, newest UpdatedAt first.
// Rows that fail JSON decode or storage invariants are skipped with a
// warn log so the rest of the bucket is still restorable. Aborting on
// the first bad row would leave the daemon with no restored incidents
// at startup.
func (db *DB) ListIncidents() ([]incident.Incident, error) {
var out []incident.Incident
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(incidentsBucket))
return b.ForEach(func(k, v []byte) error {
inc, ok := decodeIncidentRow(k, v)
if !ok {
return nil
}
out = append(out, inc)
return nil
})
})
if err != nil {
return nil, err
}
sort.Slice(out, func(i, j int) bool {
return out[i].UpdatedAt.After(out[j].UpdatedAt)
})
return out, nil
}
// ListIncidentsByStatus returns incidents matching the requested status,
// newest UpdatedAt first.
func (db *DB) ListIncidentsByStatus(status incident.Status) ([]incident.Incident, error) {
all, err := db.ListIncidents()
if err != nil {
return nil, err
}
out := all[:0]
for _, inc := range all {
if inc.Status == status {
out = append(out, inc)
}
}
return out, nil
}
// CompactIncidents removes resolved/dismissed incidents that have outlived
// their retention. Open and Contained incidents are never pruned regardless
// of age. Each transaction inspects a bounded batch, letting other store
// writers run during a large backlog. On error the count includes only
// committed deletions from preceding batches.
func (db *DB) CompactIncidents(now time.Time, retention incident.ClosedRetention) (int, error) {
pruned := 0
var next []byte
for {
n, resume, err := db.compactIncidentsBatch(now, retention, next)
pruned += n
if err != nil || resume == nil {
return pruned, err
}
next = resume
}
}
// The resume key is inclusive and copied before the transaction ends.
// Selection and deletion share a transaction so a concurrent operator
// update cannot be deleted using a stale retention decision.
func (db *DB) compactIncidentsBatch(now time.Time, retention incident.ClosedRetention, start []byte) (int, []byte, error) {
var next []byte
var toDelete [][]byte
err := db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(incidentsBucket))
cursor := b.Cursor()
k, v := cursor.First()
if start != nil {
k, v = cursor.Seek(start)
}
for scanned := 0; k != nil && scanned < incidentCompactionBatchSize; scanned++ {
inc, ok := decodeIncidentRow(k, v)
if ok && retention.Expired(inc, now) {
toDelete = append(toDelete, append([]byte(nil), k...))
}
k, v = cursor.Next()
}
if k != nil {
next = append([]byte(nil), k...)
}
for _, k := range toDelete {
if err := b.Delete(k); err != nil {
return err
}
}
return nil
})
if err != nil {
return 0, nil, err
}
return len(toDelete), next, nil
}
func decodeIncidentRow(k, v []byte) (incident.Incident, bool) {
rowID := string(k)
var inc incident.Incident
if err := json.Unmarshal(v, &inc); err != nil {
warnSkippedIncidentRow(rowID, err)
return incident.Incident{}, false
}
if err := validateIncidentRow(rowID, inc); err != nil {
warnSkippedIncidentRow(rowID, err)
return incident.Incident{}, false
}
return inc, true
}
func validateIncidentRow(rowID string, inc incident.Incident) error {
if rowID == "" {
return errors.New("empty row key")
}
if inc.ID == "" {
return errors.New("empty incident id")
}
if inc.ID != rowID {
return fmt.Errorf("incident id %q does not match row key %q", inc.ID, rowID)
}
if !incidentStatusValid(inc.Status) {
return fmt.Errorf("invalid status %q", inc.Status)
}
return nil
}
func incidentStatusValid(status incident.Status) bool {
switch status {
case incident.StatusOpen, incident.StatusContained, incident.StatusResolved, incident.StatusDismissed:
return true
default:
return false
}
}
func warnSkippedIncidentRow(rowID string, err error) {
csmlog.Warn("store: skipping corrupt incident row",
"bucket", incidentsBucket, "id", rowID, "err", err)
}
package store
import (
"encoding/json"
"errors"
"time"
bolt "go.etcd.io/bbolt"
bolterrors "go.etcd.io/bbolt/errors"
)
const mailGoodSourceBucket = "mail:good_source"
// GoodSourcePair is the persisted established-sender window for one (IP,
// mailbox): the earliest and most recent successful auth. Persisting it lets
// established good standing survive a daemon restart so the post-restart
// cold-start does not re-open the mail brute-force false-positive window.
type GoodSourcePair struct {
First time.Time `json:"first"`
Last time.Time `json:"last"`
}
// SaveMailGoodSource replaces the persisted mail good-source snapshot with data.
// Snapshot semantics: the bucket is fully rewritten so IPs no longer present are
// dropped, keeping the persisted set in step with the live tracker.
func (db *DB) SaveMailGoodSource(data map[string]map[string]GoodSourcePair) error {
encoded := make(map[string][]byte, len(data))
for ip, accts := range data {
if ip == "" || len(accts) == 0 {
continue
}
val, err := json.Marshal(accts)
if err != nil {
return err
}
encoded[ip] = val
}
return db.bolt.Update(func(tx *bolt.Tx) error {
if err := tx.DeleteBucket([]byte(mailGoodSourceBucket)); err != nil && !errors.Is(err, bolterrors.ErrBucketNotFound) {
return err
}
b, err := tx.CreateBucket([]byte(mailGoodSourceBucket))
if err != nil {
return err
}
for ip, val := range encoded {
if perr := b.Put([]byte(ip), val); perr != nil {
return perr
}
}
return nil
})
}
// LoadMailGoodSource returns the persisted mail good-source snapshot keyed by
// IP. Returns an empty map when nothing has been persisted yet.
func (db *DB) LoadMailGoodSource() (map[string]map[string]GoodSourcePair, error) {
out := make(map[string]map[string]GoodSourcePair)
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(mailGoodSourceBucket))
if b == nil {
return nil
}
return b.ForEach(func(k, v []byte) error {
var accts map[string]GoodSourcePair
if json.Unmarshal(v, &accts) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
out[string(k)] = accts
return nil
})
})
return out, err
}
package store
import (
"bufio"
"encoding/json"
"fmt"
"os"
"path/filepath"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
bolt "go.etcd.io/bbolt"
)
// All os.Open / os.ReadFile calls below take paths derived from the
// operator-configured statePath (root-owned /opt/csm or /var/lib/csm
// by default) joined with fixed filenames. gosec G304 suppressions on
// each line refer back to this package-level trust model.
func (db *DB) runMigration(statePath string) error {
fmt.Fprintf(os.Stderr, "store: migrating flat files to bbolt...\n")
var errs []string
if err := db.migrateHistory(statePath); err != nil {
errs = append(errs, fmt.Sprintf("history: %v", err))
}
if err := db.migrateAttackDB(statePath); err != nil {
errs = append(errs, fmt.Sprintf("attackdb: %v", err))
}
if err := db.migrateThreatDB(statePath); err != nil {
errs = append(errs, fmt.Sprintf("threatdb: %v", err))
}
if err := db.migrateReputation(statePath); err != nil {
errs = append(errs, fmt.Sprintf("reputation: %v", err))
}
if len(errs) > 0 {
return fmt.Errorf("partial migration: %s", strings.Join(errs, "; "))
}
_ = db.bolt.Update(func(tx *bolt.Tx) error {
return tx.Bucket([]byte("meta")).Put([]byte("migrated"), []byte(time.Now().Format(time.RFC3339)))
})
fmt.Fprintf(os.Stderr, "store: migration complete\n")
return nil
}
func (db *DB) migrateHistory(statePath string) error {
path := filepath.Join(statePath, "history.jsonl")
if _, err := os.Stat(path); err != nil {
return nil
}
// #nosec G304 -- see runMigration trust note.
f, err := os.Open(path)
if err != nil {
return err
}
defer f.Close()
var findings []alert.Finding
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
for scanner.Scan() {
var finding alert.Finding
if err := json.Unmarshal(scanner.Bytes(), &finding); err != nil {
continue
}
findings = append(findings, finding)
}
if len(findings) > 0 {
if err := db.AppendHistory(findings); err != nil {
return err
}
}
renameToBackup(path)
fmt.Fprintf(os.Stderr, "store: migrated %d history entries\n", len(findings))
return nil
}
func (db *DB) migrateAttackDB(statePath string) error {
dbDir := filepath.Join(statePath, "attack_db")
// records.json
recordsPath := filepath.Join(dbDir, "records.json")
if _, err := os.Stat(recordsPath); err == nil {
// #nosec G304 -- see runMigration trust note.
data, err := os.ReadFile(recordsPath)
if err != nil {
return err
}
var records map[string]*IPRecord
if err := json.Unmarshal(data, &records); err != nil {
return err
}
for _, r := range records {
if err := db.SaveIPRecord(*r); err != nil {
return err
}
}
renameToBackup(recordsPath)
fmt.Fprintf(os.Stderr, "store: migrated %d attack records\n", len(records))
}
// events.jsonl
eventsPath := filepath.Join(dbDir, "events.jsonl")
if _, err := os.Stat(eventsPath); err == nil {
// #nosec G304 -- see runMigration trust note.
f, err := os.Open(eventsPath)
if err != nil {
return err
}
defer f.Close()
counter := 0
scanner := bufio.NewScanner(f)
scanner.Buffer(make([]byte, 1024*1024), 1024*1024)
for scanner.Scan() {
var event AttackEvent
if err := json.Unmarshal(scanner.Bytes(), &event); err != nil {
continue
}
if err := db.RecordAttackEvent(event, counter); err != nil {
return err
}
counter++
}
renameToBackup(eventsPath)
fmt.Fprintf(os.Stderr, "store: migrated %d attack events\n", counter)
}
return nil
}
func (db *DB) migrateThreatDB(statePath string) error {
dbDir := filepath.Join(statePath, "threat_db")
// permanent.txt
permPath := filepath.Join(dbDir, "permanent.txt")
if _, err := os.Stat(permPath); err == nil {
// #nosec G304 -- see runMigration trust note.
data, err := os.ReadFile(permPath)
if err != nil {
return err
}
count := 0
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
parts := strings.SplitN(line, " # ", 2)
ip := strings.TrimSpace(parts[0])
reason := ""
if len(parts) > 1 {
reason = strings.TrimSpace(parts[1])
}
if ip != "" {
// Migrated rows keep an empty Source so Expired can still
// classify them by reason text: the flat file mixes operator
// adds with old temp auto-block writes, and stamping them all
// as operator would make the auto-block ones permanent.
if err := db.putThreatEntry(PermanentBlockEntry{
IP: ip,
Reason: reason,
BlockedAt: time.Now(),
}); err != nil {
return err
}
count++
}
}
renameToBackup(permPath)
fmt.Fprintf(os.Stderr, "store: migrated %d permanent blocks\n", count)
}
// whitelist.txt
wlPath := filepath.Join(dbDir, "whitelist.txt")
if _, err := os.Stat(wlPath); err == nil {
// #nosec G304 -- see runMigration trust note.
data, err := os.ReadFile(wlPath)
if err != nil {
return err
}
count := 0
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
parts := strings.Fields(line)
ip := parts[0]
permanent := true
var expiresAt time.Time
for _, p := range parts[1:] {
if strings.HasPrefix(p, "expires=") {
t, err := time.Parse(time.RFC3339, strings.TrimPrefix(p, "expires="))
if err == nil {
expiresAt = t
permanent = false
}
}
if p == "permanent" {
permanent = true
}
}
if err := db.AddWhitelistEntry(ip, expiresAt, permanent); err != nil {
return err
}
count++
}
renameToBackup(wlPath)
fmt.Fprintf(os.Stderr, "store: migrated %d whitelist entries\n", count)
}
return nil
}
func (db *DB) migrateReputation(statePath string) error {
path := filepath.Join(statePath, "reputation_cache.json")
if _, serr := os.Stat(path); serr != nil {
return nil //nolint:nilerr // file does not exist, skip migration
}
// #nosec G304 -- see runMigration trust note.
data, err := os.ReadFile(path)
if err != nil {
return err
}
type rawCache struct {
Entries map[string]*ReputationEntry `json:"entries"`
}
var cache rawCache
if err := json.Unmarshal(data, &cache); err != nil {
return err
}
count := 0
for ip, entry := range cache.Entries {
if err := db.SetReputation(ip, *entry); err != nil {
return err
}
count++
}
renameToBackup(path)
fmt.Fprintf(os.Stderr, "store: migrated %d reputation entries\n", count)
return nil
}
func renameToBackup(path string) {
_ = os.Rename(path, path+".bak")
}
package store
import (
"bytes"
"encoding/json"
"fmt"
"strconv"
"time"
bolt "go.etcd.io/bbolt"
)
const modsecNoEscalateKey = "modsec:no_escalate_rules"
// GetModSecNoEscalateRules returns the set of ModSecurity rule IDs that should
// NOT escalate to nftables firewall blocks. Stored in the meta bucket.
func (db *DB) GetModSecNoEscalateRules() map[int]bool {
rules := make(map[int]bool)
_ = db.bolt.View(func(tx *bolt.Tx) error {
v := tx.Bucket([]byte("meta")).Get([]byte(modsecNoEscalateKey))
if v == nil {
return nil
}
var ids []int
if json.Unmarshal(v, &ids) != nil {
return nil //nolint:nilerr // skip corrupt data
}
for _, id := range ids {
rules[id] = true
}
return nil
})
return rules
}
// SetModSecNoEscalateRules stores the set of rule IDs that should not escalate.
func (db *DB) SetModSecNoEscalateRules(rules map[int]bool) error {
var ids []int
for id := range rules {
ids = append(ids, id)
}
return db.bolt.Update(func(tx *bolt.Tx) error {
val, err := json.Marshal(ids)
if err != nil {
return err
}
return tx.Bucket([]byte("meta")).Put([]byte(modsecNoEscalateKey), val)
})
}
// AddModSecNoEscalateRule atomically adds a single rule ID to the no-escalate set.
// Read-modify-write happens in a single bbolt Update transaction to prevent races.
func (db *DB) AddModSecNoEscalateRule(ruleID int) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
var ids []int
if v := b.Get([]byte(modsecNoEscalateKey)); v != nil {
_ = json.Unmarshal(v, &ids)
}
for _, id := range ids {
if id == ruleID {
return nil // already present
}
}
ids = append(ids, ruleID)
val, err := json.Marshal(ids)
if err != nil {
return err
}
return b.Put([]byte(modsecNoEscalateKey), val)
})
}
// RemoveModSecNoEscalateRule atomically removes a single rule ID from the no-escalate set.
func (db *DB) RemoveModSecNoEscalateRule(ruleID int) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
var ids []int
if v := b.Get([]byte(modsecNoEscalateKey)); v != nil {
_ = json.Unmarshal(v, &ids)
}
var filtered []int
for _, id := range ids {
if id != ruleID {
filtered = append(filtered, id)
}
}
val, err := json.Marshal(filtered)
if err != nil {
return err
}
return b.Put([]byte(modsecNoEscalateKey), val)
})
}
// RuleHitStats holds hit count and last-hit time for a ModSecurity rule.
type RuleHitStats struct {
Hits int `json:"hits"`
LastHit time.Time `json:"last_hit"`
}
type ruleHitData struct {
Buckets map[string]int `json:"buckets"` // key: "YYYYMMDDHH" → count
LastHit time.Time `json:"last_hit"`
}
type modSecPruneItem struct {
key []byte
original []byte
data ruleHitData
remove bool
}
func modsecHitKey(ruleID int) string {
return fmt.Sprintf("modsec:hits:%d", ruleID)
}
func hourBucket(t time.Time) string {
return fmt.Sprintf("%04d%02d%02d%02d", t.Year(), t.Month(), t.Day(), t.Hour())
}
// IncrModSecRuleHit increments the hit counter for a rule ID in the current hour bucket.
func (db *DB) IncrModSecRuleHit(ruleID int, timestamp time.Time) {
key := modsecHitKey(ruleID)
bucket := hourBucket(timestamp)
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
var data ruleHitData
if v := b.Get([]byte(key)); v != nil {
if json.Unmarshal(v, &data) != nil {
data = ruleHitData{Buckets: make(map[string]int)}
}
} else {
data = ruleHitData{Buckets: make(map[string]int)}
}
data.Buckets[bucket]++
data.LastHit = timestamp
val, err := json.Marshal(data)
if err != nil {
return err
}
return b.Put([]byte(key), val)
})
}
// GetModSecRuleHits returns hit counts and last-hit timestamps for all rules
// within the last 24 hours. Prunes buckets older than 24h.
// Note: hourly bucket granularity means the window is 24h +/- 1h at boundaries.
func (db *DB) GetModSecRuleHits() map[int]RuleHitStats {
result := make(map[int]RuleHitStats)
cutoff := time.Now().Add(-24 * time.Hour)
cutoffBucket := hourBucket(cutoff)
prefix := []byte("modsec:hits:")
var toPrune []modSecPruneItem
// Read pass stays read-only so the common (nothing-to-prune) call never
// commits a write transaction. Pruning runs in a separate Update only when
// stale buckets were found.
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
c := b.Cursor()
for k, v := c.Seek(prefix); k != nil && len(k) >= len(prefix) && string(k[:len(prefix)]) == string(prefix); k, v = c.Next() {
idStr := string(k[len(prefix):])
ruleID, err := strconv.Atoi(idStr)
if err != nil {
continue
}
var data ruleHitData
if json.Unmarshal(v, &data) != nil {
continue
}
total := 0
needsPrune := false
for bk, count := range data.Buckets {
if bk >= cutoffBucket {
total += count
} else {
needsPrune = true
}
}
if needsPrune {
// bbolt keys/values are only valid inside the tx; deep copy both.
toPrune = append(toPrune, modSecPruneItem{
key: append([]byte(nil), k...),
original: append([]byte(nil), v...),
data: data,
remove: total == 0,
})
}
if total > 0 {
result[ruleID] = RuleHitStats{
Hits: total,
LastHit: data.LastHit,
}
}
}
return nil
})
if len(toPrune) == 0 {
return result
}
_ = db.bolt.Update(func(tx *bolt.Tx) error {
return pruneModSecRuleHits(tx.Bucket([]byte("meta")), toPrune, cutoffBucket)
})
return result
}
func pruneModSecRuleHits(bucket *bolt.Bucket, items []modSecPruneItem, cutoffBucket string) error {
for _, item := range items {
// A concurrent hit may have rewritten the row after the read pass.
// Leave it for the next read rather than clobbering the fresh count.
if !bytes.Equal(bucket.Get(item.key), item.original) {
continue
}
if item.remove {
if err := bucket.Delete(item.key); err != nil {
return err
}
continue
}
for bk := range item.data.Buckets {
if bk < cutoffBucket {
delete(item.data.Buckets, bk)
}
}
val, err := json.Marshal(item.data)
if err != nil {
return err
}
if err := bucket.Put(item.key, val); err != nil {
return err
}
}
return nil
}
package store
import (
"errors"
"fmt"
bolt "go.etcd.io/bbolt"
)
// PHPRelayKV is a single key/value pair for batched writes.
type PHPRelayKV struct {
Key []byte
Value []byte
}
// allowed php_relay buckets. Validated on every helper to prevent the
// daemon from accidentally writing into other buckets through these
// generic wrappers.
var phpRelayBucketAllowlist = map[string]struct{}{
"phprelay:meta": {},
"phprelay:msgindex": {},
"phprelay:ignore": {},
"phprelay:settings": {},
"phprelay:baseline": {}, // Stage 3
}
func phpRelayCheckBucket(name string) error {
if _, ok := phpRelayBucketAllowlist[name]; !ok {
return fmt.Errorf("php_relay: bucket %q is not in the allowlist", name)
}
return nil
}
// PHPRelayPut writes a single key/value into the named php_relay bucket.
func (db *DB) PHPRelayPut(bucket, key string, value []byte) error {
if err := phpRelayCheckBucket(bucket); err != nil {
return err
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucket))
if b == nil {
return fmt.Errorf("bucket %q missing", bucket)
}
return b.Put([]byte(key), append([]byte(nil), value...))
})
}
// PHPRelayGet reads a single value. ok=false when the key is absent.
func (db *DB) PHPRelayGet(bucket, key string) ([]byte, bool, error) {
if err := phpRelayCheckBucket(bucket); err != nil {
return nil, false, err
}
var out []byte
var found bool
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucket))
if b == nil {
return fmt.Errorf("bucket %q missing", bucket)
}
v := b.Get([]byte(key))
if v == nil {
return nil
}
out = append([]byte(nil), v...)
found = true
return nil
})
return out, found, err
}
// PHPRelayDelete removes a key. Missing keys are not an error.
func (db *DB) PHPRelayDelete(bucket, key string) error {
if err := phpRelayCheckBucket(bucket); err != nil {
return err
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucket))
if b == nil {
return fmt.Errorf("bucket %q missing", bucket)
}
return b.Delete([]byte(key))
})
}
// PHPRelayPutBatch writes many key/value pairs in a single bbolt
// transaction. Used by the msgIndexPersister to keep IOPS bounded.
// Returns on the first encode/put error; partial commits are visible
// only at transaction boundary.
func (db *DB) PHPRelayPutBatch(bucket string, ops []PHPRelayKV) error {
if err := phpRelayCheckBucket(bucket); err != nil {
return err
}
if len(ops) == 0 {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucket))
if b == nil {
return fmt.Errorf("bucket %q missing", bucket)
}
for _, kv := range ops {
if len(kv.Key) == 0 {
return errors.New("php_relay: empty key")
}
if err := b.Put(append([]byte(nil), kv.Key...), append([]byte(nil), kv.Value...)); err != nil {
return err
}
}
return nil
})
}
// PHPRelaySweep iterates the bucket and deletes every key for which
// shouldDelete returns true. Decoding is the caller's responsibility.
// Returns the number of deletions.
func (db *DB) PHPRelaySweep(bucket string, shouldDelete func(key, value []byte) bool) (int, error) {
if err := phpRelayCheckBucket(bucket); err != nil {
return 0, err
}
n := 0
err := db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucket))
if b == nil {
return fmt.Errorf("bucket %q missing", bucket)
}
var toDelete [][]byte
c := b.Cursor()
for k, v := c.First(); k != nil; k, v = c.Next() {
if shouldDelete(k, v) {
toDelete = append(toDelete, append([]byte(nil), k...))
}
}
for _, k := range toDelete {
if err := b.Delete(k); err != nil {
return err
}
}
n = len(toDelete)
return nil
})
return n, err
}
// PHPRelayList returns a copy of every key/value in the bucket. Used at
// daemon start to restore the in-memory ignoreList.
func (db *DB) PHPRelayList(bucket string) (map[string][]byte, error) {
if err := phpRelayCheckBucket(bucket); err != nil {
return nil, err
}
out := make(map[string][]byte)
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(bucket))
if b == nil {
return fmt.Errorf("bucket %q missing", bucket)
}
return b.ForEach(func(k, v []byte) error {
out[string(k)] = append([]byte(nil), v...)
return nil
})
})
return out, err
}
package store
import (
"encoding/json"
"time"
bolt "go.etcd.io/bbolt"
)
// PluginInfo holds cached metadata for a WordPress plugin from the API.
type PluginInfo struct {
LatestVersion string `json:"latest_version"`
TestedUpTo string `json:"tested_up_to"`
LastChecked int64 `json:"last_checked_unix"`
}
// SitePluginEntry describes a single plugin installed on a WordPress site.
type SitePluginEntry struct {
Slug string `json:"slug"`
Name string `json:"name"`
Status string `json:"status"`
InstalledVersion string `json:"installed_version"`
UpdateVersion string `json:"update_version"`
}
// SitePlugins holds the full plugin inventory for a WordPress installation.
type SitePlugins struct {
Account string `json:"account"`
Domain string `json:"domain"`
Plugins []SitePluginEntry `json:"plugins"`
}
// SetPluginInfo stores plugin metadata keyed by slug in the plugins bucket.
func (db *DB) SetPluginInfo(slug string, info PluginInfo) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("plugins"))
val, err := json.Marshal(info)
if err != nil {
return err
}
return b.Put([]byte(slug), val)
})
}
// GetPluginInfo retrieves plugin metadata for the given slug.
// Returns the entry and true if found, or a zero value and false if not.
func (db *DB) GetPluginInfo(slug string) (PluginInfo, bool) {
var info PluginInfo
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("plugins"))
v := b.Get([]byte(slug))
if v == nil {
return nil
}
if json.Unmarshal(v, &info) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
found = true
return nil
})
return info, found
}
// SetSitePlugins stores the plugin inventory for a WordPress installation
// keyed by its filesystem path in the plugins:sites bucket.
func (db *DB) SetSitePlugins(wpPath string, site SitePlugins) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("plugins:sites"))
val, err := json.Marshal(site)
if err != nil {
return err
}
return b.Put([]byte(wpPath), val)
})
}
// GetSitePlugins retrieves the plugin inventory for a WordPress installation.
// Returns the entry and true if found, or a zero value and false if not.
func (db *DB) GetSitePlugins(wpPath string) (SitePlugins, bool) {
var site SitePlugins
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("plugins:sites"))
v := b.Get([]byte(wpPath))
if v == nil {
return nil
}
if json.Unmarshal(v, &site) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
found = true
return nil
})
return site, found
}
// DeleteSitePlugins removes the plugin inventory for a WordPress installation.
func (db *DB) DeleteSitePlugins(wpPath string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("plugins:sites"))
return b.Delete([]byte(wpPath))
})
}
// AllSitePlugins returns all site plugin inventories keyed by WordPress path.
func (db *DB) AllSitePlugins() map[string]SitePlugins {
entries := make(map[string]SitePlugins)
_ = db.bolt.View(func(tx *bolt.Tx) error {
return tx.Bucket([]byte("plugins:sites")).ForEach(func(k, v []byte) error {
var s SitePlugins
if json.Unmarshal(v, &s) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
entries[string(k)] = s
return nil
})
})
return entries
}
// GetPluginRefreshTime reads the last plugin refresh timestamp from the meta bucket.
// Returns the zero time if not set.
func (db *DB) GetPluginRefreshTime() time.Time {
var t time.Time
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
v := b.Get([]byte("plugins:last_refresh"))
if v == nil {
return nil
}
parsed, err := time.Parse(time.RFC3339, string(v))
if err != nil {
return nil //nolint:nilerr // skip corrupt entry
}
t = parsed
return nil
})
return t
}
// SetPluginRefreshTime writes the plugin refresh timestamp to the meta bucket.
func (db *DB) SetPluginRefreshTime(t time.Time) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
return b.Put([]byte("plugins:last_refresh"), []byte(t.Format(time.RFC3339)))
})
}
package store
import (
"bytes"
"encoding/binary"
"errors"
"fmt"
"time"
bolt "go.etcd.io/bbolt"
)
// prefsBucket holds per-operator preference blobs:
// - "<opkey>:user" -> user settings JSON (density, timezone, etc.)
// - "<opkey>:views:<page>" -> saved filter views JSON
// - "<opkey>:table:<tableID>" -> per-table column visibility JSON
// - "<opkey>:undo:<seq>" -> bulk-action undo entry (sequence is
// big-endian uint64 of unix nano, so prefix iteration is chronological)
//
// opkey is a SHA-256 hex of the operator's auth token, computed at request
// time by the webui layer. The store never sees the token itself.
const prefsBucket = "prefs:operator"
// MaxPrefBlobSize caps the size of any single preference blob. Large enough
// for saved views with dozens of params, small enough to keep abuse bounded.
const MaxPrefBlobSize = 64 * 1024
// MaxUndoEntries caps how many undo entries one operator may have queued at
// once. Older entries fall out as new ones are recorded.
const MaxUndoEntries = 32
// UndoTTL is how long an undo entry remains valid. Matches the banner timeout
// the UI advertises so an operator can never undo an action whose banner has
// already disappeared.
const UndoTTL = 30 * time.Second
// ErrPrefBlobTooLarge is returned when a preference blob exceeds MaxPrefBlobSize.
var ErrPrefBlobTooLarge = errors.New("preference blob too large")
func prefsKey(opkey, ns string) []byte {
return []byte(opkey + ":" + ns)
}
func undoKey(opkey string, seq uint64) []byte {
out := make([]byte, 0, len(opkey)+6+8)
out = append(out, opkey...)
out = append(out, ':', 'u', 'n', 'd', 'o', ':')
var seqBE [8]byte
binary.BigEndian.PutUint64(seqBE[:], seq)
return append(out, seqBE[:]...)
}
func undoKeyPrefix(opkey string) []byte {
return []byte(opkey + ":undo:")
}
// GetOperatorPref returns the raw JSON blob stored at (opkey, ns). Returns
// nil with nil error when no entry exists.
func (db *DB) GetOperatorPref(opkey, ns string) ([]byte, error) {
if db == nil || db.bolt == nil {
return nil, errors.New("store unavailable")
}
if opkey == "" || ns == "" {
return nil, errors.New("opkey and namespace required")
}
var out []byte
err := db.bolt.View(func(tx *bolt.Tx) error {
b, err := bucketOrCreate(tx, prefsBucket)
if err != nil {
return err
}
v := b.Get(prefsKey(opkey, ns))
if v == nil {
return nil
}
// Bolt invalidates the slice after the tx ends; copy out.
out = append([]byte(nil), v...)
return nil
})
return out, err
}
// PutOperatorPref writes the JSON blob to (opkey, ns). Rejects payloads above
// MaxPrefBlobSize.
func (db *DB) PutOperatorPref(opkey, ns string, data []byte) error {
if db == nil || db.bolt == nil {
return errors.New("store unavailable")
}
if opkey == "" || ns == "" {
return errors.New("opkey and namespace required")
}
if len(data) > MaxPrefBlobSize {
return ErrPrefBlobTooLarge
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b, err := bucketOrCreate(tx, prefsBucket)
if err != nil {
return err
}
return b.Put(prefsKey(opkey, ns), data)
})
}
// DeleteOperatorPref removes the blob at (opkey, ns). No error if absent.
func (db *DB) DeleteOperatorPref(opkey, ns string) error {
if db == nil || db.bolt == nil {
return errors.New("store unavailable")
}
if opkey == "" || ns == "" {
return errors.New("opkey and namespace required")
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b, err := bucketOrCreate(tx, prefsBucket)
if err != nil {
return err
}
return b.Delete(prefsKey(opkey, ns))
})
}
// UndoEntry is one record in the bulk-action undo queue.
type UndoEntry struct {
Targets []string `json:"targets,omitempty"` // IPs whose later edits invalidate this action
ID string `json:"id"` // matches the seq encoded in the bbolt key
RecordedAt time.Time `json:"recorded_at"` // wall time the entry was written
Action string `json:"action"` // e.g. "threat_bulk_block"
Inverse string `json:"inverse"` // inverse action key the runner will dispatch
Payload []byte `json:"payload"` // opaque JSON the runner understands
Summary string `json:"summary"` // human-readable label for the banner
}
// AppendUndoEntry queues an undo entry for opkey. The entry's ID and
// RecordedAt are filled in. Prunes expired entries and trims the queue to
// MaxUndoEntries before writing. Returns the saved entry.
func (db *DB) AppendUndoEntry(opkey string, e UndoEntry) (UndoEntry, error) {
if db == nil || db.bolt == nil {
return UndoEntry{}, errors.New("store unavailable")
}
if opkey == "" {
return UndoEntry{}, errors.New("opkey required")
}
if e.Inverse == "" {
return UndoEntry{}, errors.New("inverse action required")
}
now := time.Now().UTC()
seq := uint64(now.UnixNano())
e.ID = fmt.Sprintf("%016x", seq)
e.RecordedAt = now
raw, err := encodeUndoEntry(e)
if err != nil {
return UndoEntry{}, err
}
if len(raw) > MaxPrefBlobSize {
return UndoEntry{}, ErrPrefBlobTooLarge
}
err = db.bolt.Update(func(tx *bolt.Tx) error {
b, berr := bucketOrCreate(tx, prefsBucket)
if berr != nil {
return berr
}
if perr := pruneOperatorUndo(b, opkey, now); perr != nil {
return perr
}
return b.Put(undoKey(opkey, seq), raw)
})
if err != nil {
return UndoEntry{}, err
}
return e, nil
}
// LatestUndoEntry returns the most recent non-expired undo entry for opkey,
// or (zero, false, nil) if none exists.
func (db *DB) LatestUndoEntry(opkey string) (UndoEntry, bool, error) {
if db == nil || db.bolt == nil {
return UndoEntry{}, false, errors.New("store unavailable")
}
if opkey == "" {
return UndoEntry{}, false, errors.New("opkey required")
}
now := time.Now().UTC()
var entry UndoEntry
var found bool
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(prefsBucket))
if b == nil {
return nil
}
c := b.Cursor()
prefix := undoKeyPrefix(opkey)
// Seek to first key strictly greater than the prefix, then walk
// backwards. Newest entries sort last because the suffix is the
// big-endian nano timestamp.
for k, v := c.Last(); k != nil; k, v = c.Prev() {
if !bytes.HasPrefix(k, prefix) {
if bytes.Compare(k, prefix) < 0 {
return nil
}
continue
}
e, err := decodeUndoEntry(v)
if err != nil {
continue
}
if now.Sub(e.RecordedAt) > UndoTTL {
return nil
}
entry = e
found = true
return nil
}
return nil
})
return entry, found, err
}
// ConsumeUndoEntry removes the entry identified by id from opkey's queue and
// returns the decoded value. Returns (zero, false, nil) if the entry has
// already expired or never existed.
func (db *DB) ConsumeUndoEntry(opkey, id string) (UndoEntry, bool, error) {
if db == nil || db.bolt == nil {
return UndoEntry{}, false, errors.New("store unavailable")
}
if opkey == "" || id == "" {
return UndoEntry{}, false, errors.New("opkey and id required")
}
now := time.Now().UTC()
var entry UndoEntry
var found bool
err := db.bolt.Update(func(tx *bolt.Tx) error {
b, err := bucketOrCreate(tx, prefsBucket)
if err != nil {
return err
}
if perr := pruneOperatorUndo(b, opkey, now); perr != nil {
return perr
}
c := b.Cursor()
prefix := undoKeyPrefix(opkey)
for k, v := c.Seek(prefix); k != nil && bytes.HasPrefix(k, prefix); k, v = c.Next() {
e, derr := decodeUndoEntry(v)
if derr != nil {
continue
}
if e.ID != id {
continue
}
if now.Sub(e.RecordedAt) > UndoTTL {
return b.Delete(k)
}
entry = e
found = true
return b.Delete(k)
}
return nil
})
return entry, found, err
}
// PurgeOperatorUndo drops every undo entry for opkey (used when an operator
// logs out or when tests need to reset state).
func (db *DB) PurgeOperatorUndo(opkey string) error {
if db == nil || db.bolt == nil {
return errors.New("store unavailable")
}
if opkey == "" {
return errors.New("opkey required")
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(prefsBucket))
if b == nil {
return nil
}
var keys [][]byte
prefix := undoKeyPrefix(opkey)
c := b.Cursor()
for k, _ := c.Seek(prefix); k != nil && bytes.HasPrefix(k, prefix); k, _ = c.Next() {
keys = append(keys, append([]byte(nil), k...))
}
for _, k := range keys {
if err := b.Delete(k); err != nil {
return err
}
}
return nil
})
}
// pruneOperatorUndo drops expired entries and trims the queue to
// MaxUndoEntries. Must run inside a writable transaction.
func pruneOperatorUndo(b *bolt.Bucket, opkey string, now time.Time) error {
prefix := undoKeyPrefix(opkey)
type kv struct {
key []byte
entry UndoEntry
}
var kept []kv
var expired [][]byte
c := b.Cursor()
for k, v := c.Seek(prefix); k != nil && bytes.HasPrefix(k, prefix); k, v = c.Next() {
e, err := decodeUndoEntry(v)
if err != nil {
expired = append(expired, append([]byte(nil), k...))
continue
}
if now.Sub(e.RecordedAt) > UndoTTL {
expired = append(expired, append([]byte(nil), k...))
continue
}
kept = append(kept, kv{key: append([]byte(nil), k...), entry: e})
}
for _, k := range expired {
if err := b.Delete(k); err != nil {
return err
}
}
// kept is already in chronological order because keys sort by big-endian
// nano suffix. Trim oldest first.
if len(kept) >= MaxUndoEntries {
drop := len(kept) - MaxUndoEntries + 1
for i := 0; i < drop; i++ {
if err := b.Delete(kept[i].key); err != nil {
return err
}
}
}
return nil
}
func bucketOrCreate(tx *bolt.Tx, name string) (*bolt.Bucket, error) {
if tx.Writable() {
return tx.CreateBucketIfNotExists([]byte(name))
}
if b := tx.Bucket([]byte(name)); b != nil {
return b, nil
}
return nil, fmt.Errorf("bucket %s not initialised", name)
}
// InvalidateUndoTargets retires actions superseded by a later decision about
// any of their IPs, across operators. Invalidate the whole action so a stale
// bulk undo cannot silently restore only part of an operator's decision.
func (db *DB) InvalidateUndoTargets(ips []string) error {
return db.bolt.Update(func(tx *bolt.Tx) error { return invalidateUndoTargets(tx, ips) })
}
func invalidateUndoTargets(tx *bolt.Tx, ips []string) error {
b := tx.Bucket([]byte(prefsBucket))
if b == nil {
return nil
}
targets := make(map[string]bool, len(ips))
for _, ip := range ips {
targets[ip] = true
}
var keys [][]byte
if err := b.ForEach(func(k, v []byte) error {
if !bytes.Contains(k, []byte(":undo:")) {
return nil
}
entry, err := decodeUndoEntry(v)
if err != nil {
return nil //nolint:nilerr // Corrupt entries cannot execute.
}
for _, ip := range entry.Targets {
if targets[ip] {
keys = append(keys, append([]byte(nil), k...))
break
}
}
return nil
}); err != nil {
return err
}
for _, key := range keys {
if err := b.Delete(key); err != nil {
return err
}
}
return nil
}
package store
import "encoding/json"
func encodeUndoEntry(e UndoEntry) ([]byte, error) {
return json.Marshal(e)
}
func decodeUndoEntry(raw []byte) (UndoEntry, error) {
var e UndoEntry
if err := json.Unmarshal(raw, &e); err != nil {
return UndoEntry{}, err
}
return e, nil
}
package store
import (
"bytes"
"fmt"
"time"
bolt "go.etcd.io/bbolt"
)
// RewriteUndoEntryRecordedAt rewrites the RecordedAt field of an undo
// entry identified by id. Used by tests that need to age an entry past the
// TTL window without sleeping for real.
func RewriteUndoEntryRecordedAt(db *DB, id string, at time.Time) error {
if db == nil || db.bolt == nil {
return fmt.Errorf("store unavailable")
}
if id == "" {
return fmt.Errorf("id required")
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(prefsBucket))
if b == nil {
return fmt.Errorf("prefs bucket missing")
}
c := b.Cursor()
needle := []byte(":undo:")
for k, v := c.First(); k != nil; k, v = c.Next() {
if !bytes.Contains(k, needle) {
continue
}
e, err := decodeUndoEntry(v)
if err != nil {
continue
}
if e.ID != id {
continue
}
e.RecordedAt = at
raw, err := encodeUndoEntry(e)
if err != nil {
return err
}
return b.Put(k, raw)
}
return fmt.Errorf("entry %s not found", id)
})
}
package store
import (
"encoding/json"
"fmt"
"sort"
"strings"
"time"
bolt "go.etcd.io/bbolt"
)
// Meta-bucket keys for AbuseIPDB quota accounting. Persisted so enforcement
// survives daemon restarts and spans across 10-minute scan cycles.
const (
abuseQuotaExhaustedKey = "abuse:quota_exhausted_until"
abuseDailyCountPrefix = "abuse:daily_count:" // + YYYY-MM-DD in UTC
)
// ReputationEntry holds the cached reputation data for an IP address.
type ReputationEntry struct {
Score int `json:"score"`
Category string `json:"category"`
CheckedAt time.Time `json:"checked_at"`
}
// SetReputation stores a reputation entry for the given IP.
func (db *DB) SetReputation(ip string, entry ReputationEntry) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
val, err := json.Marshal(entry)
if err != nil {
return err
}
return b.Put([]byte(ip), val)
})
}
// SetReputationBatch stores multiple reputation entries in one write
// transaction. Every bbolt commit fsyncs, so per-entry SetReputation
// calls in a loop cost one disk flush each; a cycle's worth of cache
// updates must land in a single commit.
func (db *DB) SetReputationBatch(entries map[string]ReputationEntry) error {
if len(entries) == 0 {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
for ip, entry := range entries {
val, err := json.Marshal(entry)
if err != nil {
return err
}
if err := b.Put([]byte(ip), val); err != nil {
return err
}
}
return nil
})
}
// ApplyReputationChanges stores and deletes reputation entries in one write
// transaction.
func (db *DB) ApplyReputationChanges(upserts map[string]ReputationEntry, deletes map[string]bool) error {
if len(upserts) == 0 && len(deletes) == 0 {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
for ip := range deletes {
if err := b.Delete([]byte(ip)); err != nil {
return err
}
}
for ip, entry := range upserts {
val, err := json.Marshal(entry)
if err != nil {
return err
}
if err := b.Put([]byte(ip), val); err != nil {
return err
}
}
return nil
})
}
// GetReputation retrieves a reputation entry for the given IP.
// Returns the entry and true if found, or a zero value and false if not.
func (db *DB) GetReputation(ip string) (ReputationEntry, bool) {
var entry ReputationEntry
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
v := b.Get([]byte(ip))
if v == nil {
return nil
}
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
found = true
return nil
})
return entry, found
}
// CleanExpiredReputation deletes entries older than maxAge.
// Uses a collect-then-delete pattern because bbolt does not allow mutation
// during ForEach iteration. Returns the count of entries removed.
func (db *DB) CleanExpiredReputation(maxAge time.Duration) int {
var removed int
cutoff := time.Now().Add(-maxAge)
updateErr := boltUpdate(db.bolt, func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
// Collect keys to delete.
var toDelete [][]byte
_ = b.ForEach(func(k, v []byte) error {
var entry ReputationEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if entry.CheckedAt.Before(cutoff) {
keyCopy := make([]byte, len(k))
copy(keyCopy, k)
toDelete = append(toDelete, keyCopy)
}
return nil
})
// Delete collected keys.
for _, k := range toDelete {
if err := b.Delete(k); err != nil {
return err
}
removed++
}
return nil
})
return committedCount("reputation", updateErr, removed)
}
// AllReputation returns all reputation entries keyed by IP.
func (db *DB) AllReputation() map[string]ReputationEntry {
entries := make(map[string]ReputationEntry)
_ = db.bolt.View(func(tx *bolt.Tx) error {
return tx.Bucket([]byte("reputation")).ForEach(func(k, v []byte) error {
var e ReputationEntry
if json.Unmarshal(v, &e) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
entries[string(k)] = e
return nil
})
})
return entries
}
// SetAbuseQuotaExhaustedUntil records the time at which the AbuseIPDB
// quota is expected to reset. While now < t, callers should skip API
// queries. The daemon re-reads this on every cycle so the flag survives
// restarts and multi-hour backoffs.
func (db *DB) SetAbuseQuotaExhaustedUntil(t time.Time) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
return tx.Bucket([]byte("meta")).Put(
[]byte(abuseQuotaExhaustedKey),
[]byte(t.UTC().Format(time.RFC3339)),
)
})
}
// AbuseQuotaExhaustedUntil returns the persisted quota-reset timestamp,
// or zero time if none is recorded (or the stored value is unparseable).
func (db *DB) AbuseQuotaExhaustedUntil() time.Time {
var ts time.Time
_ = db.bolt.View(func(tx *bolt.Tx) error {
v := tx.Bucket([]byte("meta")).Get([]byte(abuseQuotaExhaustedKey))
if v == nil {
return nil
}
parsed, err := time.Parse(time.RFC3339, string(v))
if err != nil {
return nil //nolint:nilerr // skip corrupt entry
}
ts = parsed
return nil
})
return ts
}
// IncrementAbuseQueryCount bumps and returns the AbuseIPDB query counter
// for the given UTC date (YYYY-MM-DD). Used as a daily circuit breaker.
func (db *DB) IncrementAbuseQueryCount(utcDate string) int {
var count int
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
pruneAbuseDailyCounts(b, utcDate)
key := []byte(abuseDailyCountPrefix + utcDate)
if v := b.Get(key); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", &count)
}
count++
return b.Put(key, []byte(fmt.Sprintf("%d", count)))
})
return count
}
// ReserveAbuseQuerySlots atomically reserves up to requested AbuseIPDB
// query slots for utcDate without increasing the daily counter beyond max.
func (db *DB) ReserveAbuseQuerySlots(utcDate string, requested, max int) int {
if requested <= 0 || max <= 0 {
return 0
}
var reserved int
_ = db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
pruneAbuseDailyCounts(b, utcDate)
key := []byte(abuseDailyCountPrefix + utcDate)
count := 0
if v := b.Get(key); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", &count)
}
if count >= max {
return nil
}
remaining := max - count
reserved = requested
if reserved > remaining {
reserved = remaining
}
count += reserved
return b.Put(key, []byte(fmt.Sprintf("%d", count)))
})
return reserved
}
func pruneAbuseDailyCounts(bucket *bolt.Bucket, keepDate string) {
prefix := []byte(abuseDailyCountPrefix)
cutoff, err := time.Parse("2006-01-02", keepDate)
if err != nil {
return
}
oldest := cutoff.AddDate(0, 0, -1)
cursor := bucket.Cursor()
for key, _ := cursor.Seek(prefix); key != nil && strings.HasPrefix(string(key), abuseDailyCountPrefix); key, _ = cursor.Next() {
date := strings.TrimPrefix(string(key), abuseDailyCountPrefix)
parsed, parseErr := time.Parse("2006-01-02", date)
if parseErr != nil || parsed.Before(oldest) {
_ = cursor.Delete()
}
}
}
// AbuseQueryCount returns the AbuseIPDB query count for the given UTC date.
func (db *DB) AbuseQueryCount(utcDate string) int {
var count int
_ = db.bolt.View(func(tx *bolt.Tx) error {
v := tx.Bucket([]byte("meta")).Get([]byte(abuseDailyCountPrefix + utcDate))
if v != nil {
_, _ = fmt.Sscanf(string(v), "%d", &count)
}
return nil
})
return count
}
// EnforceReputationCap ensures the reputation bucket has at most max entries.
// If the count exceeds max, the oldest entries (by CheckedAt) are deleted.
// Returns the count of entries removed.
func (db *DB) EnforceReputationCap(max int) int {
var removed int
updateErr := boltUpdate(db.bolt, func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
// Collect all entries with their keys.
type keyed struct {
key []byte
checkedAt time.Time
}
var all []keyed
_ = b.ForEach(func(k, v []byte) error {
var entry ReputationEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
keyCopy := make([]byte, len(k))
copy(keyCopy, k)
all = append(all, keyed{key: keyCopy, checkedAt: entry.CheckedAt})
return nil
})
if len(all) <= max {
return nil
}
// Sort by CheckedAt ascending (oldest first).
sort.Slice(all, func(i, j int) bool {
return all[i].checkedAt.Before(all[j].checkedAt)
})
// Delete the oldest entries beyond the cap.
excess := len(all) - max
for i := 0; i < excess; i++ {
if err := b.Delete(all[i].key); err != nil {
return err
}
removed++
}
return nil
})
return committedCount("reputation cap", updateErr, removed)
}
package store
import (
"encoding/json"
"errors"
"fmt"
"os"
"time"
bolt "go.etcd.io/bbolt"
)
// timeKeyLowerBound computes the lexicographic lower bound of any TimeKey
// produced for timestamp t. Any TimeKey whose stored time is strictly
// earlier than t sorts before this string; any TimeKey at or after t sorts
// at or above it. Matches the format in TimeKey().
func timeKeyLowerBound(t time.Time) string {
t = t.UTC()
return fmt.Sprintf("%04d%02d%02d%02d%02d%02d%09d-0000",
t.Year(), t.Month(), t.Day(),
t.Hour(), t.Minute(), t.Second(),
t.Nanosecond())
}
// SweepHistoryOlderThan deletes history entries whose TimeKey is strictly
// older than cutoff. Returns the number of entries deleted. All work runs
// in a single bbolt transaction so the UI never sees a half-swept state;
// callers pick cutoffs that keep the batch bounded.
func (db *DB) SweepHistoryOlderThan(cutoff time.Time) (int, error) {
cutoffKey := timeKeyLowerBound(cutoff)
var deleted int
err := db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
if b == nil {
return nil
}
c := b.Cursor()
for k, _ := c.First(); k != nil && string(k) < cutoffKey; k, _ = c.First() {
if err := c.Delete(); err != nil {
return err
}
deleted++
}
if deleted == 0 {
return nil
}
if err := bumpHistoryRevision(tx); err != nil {
return err
}
// Decrement the history:count counter without letting it underflow.
current := 0
if v := tx.Bucket([]byte("meta")).Get([]byte("history:count")); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", ¤t)
}
newCount := current - deleted
if newCount < 0 {
newCount = 0
}
return setCounter(tx, "history:count", newCount)
})
return deleted, err
}
// SweepAttackEventsOlderThan deletes attacks:events entries older than
// cutoff and the matching entries from the attacks:events:ip secondary
// index. Returns the number of primary-bucket entries deleted.
func (db *DB) SweepAttackEventsOlderThan(cutoff time.Time) (int, error) {
cutoffKey := timeKeyLowerBound(cutoff)
var deleted int
err := db.bolt.Update(func(tx *bolt.Tx) error {
primary := tx.Bucket([]byte("attacks:events"))
secondary := tx.Bucket([]byte("attacks:events:ip"))
if primary == nil {
return nil
}
c := primary.Cursor()
for k, v := c.First(); k != nil && string(k) < cutoffKey; k, v = c.First() {
// The secondary index is keyed "<ip>/<TimeKey>", so we need
// the event's IP to prune it.
var ev AttackEvent
if err := json.Unmarshal(v, &ev); err == nil && secondary != nil {
secKey := []byte(ev.IP + "/" + string(k))
if err := secondary.Delete(secKey); err != nil {
return err
}
}
if err := c.Delete(); err != nil {
return err
}
deleted++
}
if deleted == 0 {
return nil
}
current := 0
if v := tx.Bucket([]byte("meta")).Get([]byte("attacks:events:count")); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", ¤t)
}
newCount := current - deleted
if newCount < 0 {
newCount = 0
}
return setCounter(tx, "attacks:events:count", newCount)
})
return deleted, err
}
// Size returns the on-disk size of the bbolt file in bytes. bbolt does
// not shrink the file on delete; compare Size() before and after a
// CompactInto call to see how much space would be reclaimed.
func (db *DB) Size() (int64, error) {
info, err := os.Stat(db.path)
if err != nil {
return 0, err
}
return info.Size(), nil
}
// CompactionDue reports whether the state db is worth compacting: large enough
// that the slack matters AND fragmented enough that a compaction would reclaim
// a meaningful fraction. minSizeMB and fillRatio come from Retention config;
// non-positive values disable the check. freeBytes above sizeBytes is clamped
// (used=0) rather than producing a negative fill. Startup compaction and the
// daemon's compaction hint share it so the hint never promises a compaction
// the next start will skip.
func CompactionDue(sizeBytes, freeBytes int64, minSizeMB int, fillRatio float64) bool {
if minSizeMB <= 0 || fillRatio <= 0 || sizeBytes <= 0 {
return false
}
if sizeBytes < int64(minSizeMB)*1024*1024 {
return false
}
used := sizeBytes - freeBytes
if used < 0 {
used = 0
}
fill := float64(used) / float64(sizeBytes)
return fill < fillRatio
}
// FreeBytes returns the number of bytes held by free and pending pages in the
// bbolt freelist -- the space a compaction would reclaim. bbolt never shrinks
// the file on delete, so a large FreeBytes relative to Size means the on-disk
// file is mostly slack and is worth compacting.
func (db *DB) FreeBytes() (int64, error) {
// FreeAlloc includes free and pending pages using the database's page
// size. Stats holds bbolt's statistics lock; Info reads the mmap without
// locking and can race with remapping while the live database grows.
return int64(db.bolt.Stats().FreeAlloc), nil
}
// CompactInto snapshots the live DB into a fresh bbolt file at dstPath
// using bolt.Compact. Returns the source size and the compacted size
// (both in bytes).
//
// Correctness: bolt.Compact runs a View transaction on src for the
// duration of the walk, so concurrent Update calls on src will either
// land before the walk begins (captured in the snapshot) or after it
// completes (not in the snapshot). It is the caller's job to quiesce
// writers between the CompactInto call and the file rename+reopen that
// promotes the new file; otherwise post-snapshot writes are silently
// dropped during the swap.
//
// txMaxSize caps per-transaction bytes written to the destination (see
// bolt.Compact docs). Zero means "one transaction for the whole copy",
// which is the fastest path for DBs that comfortably fit in memory.
func (db *DB) CompactInto(dstPath string, txMaxSize int64) (srcSize, dstSize int64, err error) {
if dstPath == "" {
return 0, 0, errors.New("dst path is empty")
}
// Snapshot the src size up front; if bolt.Compact mutates src in ways
// we didn't anticipate, a concurrent reader still sees consistent
// numbers.
srcInfo, statErr := os.Stat(db.path)
if statErr != nil {
return 0, 0, fmt.Errorf("stat src: %w", statErr)
}
srcSize = srcInfo.Size()
dst, err := bolt.Open(dstPath, 0600, &bolt.Options{Timeout: 5 * time.Second})
if err != nil {
// bolt.Open may have created a zero-byte file before failing; clean it up.
_ = os.Remove(dstPath)
return srcSize, 0, fmt.Errorf("opening dst: %w", err)
}
compactErr := bolt.Compact(dst, db.bolt, txMaxSize)
if closeErr := dst.Close(); closeErr != nil && compactErr == nil {
compactErr = fmt.Errorf("closing dst: %w", closeErr)
}
if compactErr != nil {
_ = os.Remove(dstPath)
return srcSize, 0, compactErr
}
dstInfo, err := os.Stat(dstPath)
if err != nil {
return srcSize, 0, fmt.Errorf("stat dst: %w", err)
}
return srcSize, dstInfo.Size(), nil
}
// SweepReputationOlderThan deletes reputation entries whose CheckedAt is
// strictly older than cutoff. The bucket is keyed by IP and not by time,
// so the sweep inspects each value; malformed rows are skipped rather than
// aborting the sweep.
func (db *DB) SweepReputationOlderThan(cutoff time.Time) (int, error) {
var deleted int
err := db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("reputation"))
if b == nil {
return nil
}
var stale [][]byte
if err := b.ForEach(func(k, v []byte) error {
var e ReputationEntry
// Malformed rows are skipped: a bad Unmarshal here must not
// abort the sweep for the rest of the bucket. Expressed as
// "proceed only when the row parses and is stale" so the
// happy path stays on the left of the guard.
if err := json.Unmarshal(v, &e); err == nil && e.CheckedAt.Before(cutoff) {
// Copy k because the slice is only valid for the
// duration of the callback.
stale = append(stale, append([]byte(nil), k...))
}
return nil
}); err != nil {
return err
}
for _, k := range stale {
if err := b.Delete(k); err != nil {
return err
}
deleted++
}
return nil
})
return deleted, err
}
package store
import (
"encoding/json"
"fmt"
"time"
bolt "go.etcd.io/bbolt"
)
const scanCursorBucket = "scan_cursor"
// ScanCursorRecord is the rolling-coverage cursor for one (account, check).
type ScanCursorRecord struct {
Account string `json:"account"`
Check string `json:"check"`
LastPath string `json:"last_path"` // last path-sorted candidate scanned
WrappedAt time.Time `json:"wrapped_at"` // when the cursor last wrapped to start
LastFullCycleTS time.Time `json:"last_full_cycle_ts"` // when a full sweep last completed
}
// scanCursorKey builds the bucket key for a (account, check) pair.
// Format: "<account>/<check>".
func scanCursorKey(account, check string) []byte {
return []byte(account + "/" + check)
}
// GetScanCursor retrieves the cursor record for (account, check).
// ok=false when absent (no error).
func (db *DB) GetScanCursor(account, check string) (ScanCursorRecord, bool, error) {
var rec ScanCursorRecord
var found bool
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanCursorBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanCursorBucket)
}
v := b.Get(scanCursorKey(account, check))
if v == nil {
return nil
}
found = true
return json.Unmarshal(v, &rec)
})
if err != nil {
return ScanCursorRecord{}, false, err
}
return rec, found, nil
}
// PutScanCursor creates or replaces the cursor record for rec.Account/rec.Check.
func (db *DB) PutScanCursor(rec ScanCursorRecord) error {
val, err := json.Marshal(rec)
if err != nil {
return fmt.Errorf("scancursor: marshal: %w", err)
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanCursorBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanCursorBucket)
}
return b.Put(scanCursorKey(rec.Account, rec.Check), val)
})
}
package store
import (
"bytes"
"encoding/json"
"fmt"
"sort"
"time"
bolt "go.etcd.io/bbolt"
"github.com/pidginhost/csm/internal/alert"
)
const (
scanJobsBucket = "scan_jobs"
scanJobFindingsBucket = "scan_job_findings"
)
// ScanJobRecord holds the metadata for a single full-scan job.
type ScanJobRecord struct {
ID string `json:"id"`
Scope string `json:"scope"`
Target string `json:"target"`
State string `json:"state"`
Created time.Time `json:"created"`
Started time.Time `json:"started,omitzero"`
Finished time.Time `json:"finished,omitzero"`
FilesScanned int `json:"files_scanned,omitempty"`
FilesEst int `json:"files_est,omitempty"`
FindingCount int `json:"finding_count,omitempty"`
// FindingsStored is how many findings were actually persisted. It equals
// FindingCount unless the per-job cap truncated the tail, in which case
// FindingsTruncated is set and the UI shows "showing first N of M".
FindingsStored int `json:"findings_stored,omitempty"`
FindingsTruncated bool `json:"findings_truncated,omitempty"`
// Progress fields for scope="all" jobs. Zero/empty for account-scope jobs.
AccountsTotal int `json:"accounts_total,omitempty"`
AccountsDone int `json:"accounts_done,omitempty"`
CurrentAccount string `json:"current_account,omitempty"`
Options map[string]any `json:"options,omitempty"`
Error string `json:"error,omitempty"`
}
// PutScanJob creates or replaces a scan-job record.
func (db *DB) PutScanJob(rec ScanJobRecord) error {
val, err := json.Marshal(rec)
if err != nil {
return fmt.Errorf("scanjobs: marshal: %w", err)
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanJobsBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanJobsBucket)
}
return b.Put([]byte(rec.ID), val)
})
}
// GetScanJob retrieves a single scan-job record. ok=false when the ID is absent.
func (db *DB) GetScanJob(id string) (ScanJobRecord, bool, error) {
var rec ScanJobRecord
var found bool
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanJobsBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanJobsBucket)
}
v := b.Get([]byte(id))
if v == nil {
return nil
}
found = true
return json.Unmarshal(v, &rec)
})
if err != nil {
return ScanJobRecord{}, false, err
}
return rec, found, nil
}
// ListScanJobs returns all scan-job records ordered newest-first (by Created,
// with ID as a deterministic tiebreaker for equal timestamps).
func (db *DB) ListScanJobs() ([]ScanJobRecord, error) {
var jobs []ScanJobRecord
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanJobsBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanJobsBucket)
}
return b.ForEach(func(_, v []byte) error {
var rec ScanJobRecord
if err := json.Unmarshal(v, &rec); err != nil {
return err
}
jobs = append(jobs, rec)
return nil
})
})
if err != nil {
return nil, err
}
sort.Slice(jobs, func(i, j int) bool {
ci, cj := jobs[i].Created, jobs[j].Created
if ci.Equal(cj) {
return jobs[i].ID > jobs[j].ID
}
return ci.After(cj)
})
return jobs, nil
}
// findingKey builds the bucket key for a scan-job finding.
// Format: "<job_id>/<zero-padded-seq>" using 8 decimal digits for seq so
// lexicographic order equals insertion order and prefix scans work cleanly.
func findingKey(jobID string, seq int) []byte {
return []byte(fmt.Sprintf("%s/%08d", jobID, seq))
}
// findingPrefix returns the prefix used to address all findings for a job.
func findingPrefix(jobID string) []byte {
return []byte(jobID + "/")
}
// AppendScanJobFinding persists a single finding for the given job.
// seq must be unique within the job (caller supplies a monotonic counter).
// Using many small values rather than one growing blob enables pagination
// via cursor seeks without loading the whole finding list into memory.
func (db *DB) AppendScanJobFinding(id string, seq int, f alert.Finding) error {
return db.AppendScanJobFindings(id, seq, []alert.Finding{f})
}
// AppendScanJobFindings persists a batch of findings for the given job in a
// single write transaction. Keys are seq, seq+1, ... seq+len-1. Batching
// amortizes bbolt's per-commit fsync across the whole slice instead of paying
// one fsync per finding, which dominated the write cost of a large scan.
func (db *DB) AppendScanJobFindings(id string, seq int, findings []alert.Finding) error {
if len(findings) == 0 {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanJobFindingsBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanJobFindingsBucket)
}
for i, f := range findings {
val, err := json.Marshal(f)
if err != nil {
return fmt.Errorf("scanjobs: marshal finding: %w", err)
}
if err := b.Put(findingKey(id, seq+i), val); err != nil {
return err
}
}
return nil
})
}
// ListScanJobFindings returns a paginated slice of findings for a job.
// total is the count of ALL findings for the job (ignoring offset/limit).
// offset and limit follow the usual slice semantics; limit=0 returns all.
func (db *DB) ListScanJobFindings(id string, offset, limit int) ([]alert.Finding, int, error) {
prefix := findingPrefix(id)
var findings []alert.Finding
total := 0
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte(scanJobFindingsBucket))
if b == nil {
return fmt.Errorf("bucket %q missing", scanJobFindingsBucket)
}
c := b.Cursor()
pos := 0
for k, v := c.Seek(prefix); k != nil && bytes.HasPrefix(k, prefix); k, v = c.Next() {
total++
if pos >= offset && (limit == 0 || len(findings) < limit) {
var f alert.Finding
if err := json.Unmarshal(v, &f); err != nil {
return err
}
findings = append(findings, f)
}
pos++
}
return nil
})
if err != nil {
return nil, 0, err
}
return findings, total, nil
}
// PruneScanJobs enforces two independent retention limits and removes the
// finding rows of every pruned job. Returns the number of jobs pruned.
//
// - keepJobs bounds how many job records are retained.
// - maxTotalFindings bounds the cumulative finding rows across retained jobs
// (0 disables the volume cap). A single scan can emit tens of thousands of
// findings, so a job-count cap alone lets retained findings grow without
// bound; the volume cap keeps the state file in check.
//
// Jobs are ranked newest-first (Created desc, ID tiebreaker). The newest job is
// always kept even if it alone exceeds the volume cap, so retention never wipes
// the most recent result. The whole operation runs in one bolt.Update so a
// concurrent insert cannot make the decision stale.
func (db *DB) PruneScanJobs(keepJobs, maxTotalFindings int) (int, error) {
if keepJobs < 0 {
keepJobs = 0
}
pruned := 0
err := db.bolt.Update(func(tx *bolt.Tx) error {
jb := tx.Bucket([]byte(scanJobsBucket))
if jb == nil {
return fmt.Errorf("bucket %q missing", scanJobsBucket)
}
fb := tx.Bucket([]byte(scanJobFindingsBucket))
if fb == nil {
return fmt.Errorf("bucket %q missing", scanJobFindingsBucket)
}
// Collect all job records from the bucket (reads are allowed inside Update).
var jobs []ScanJobRecord
if err := jb.ForEach(func(_, v []byte) error {
var rec ScanJobRecord
if err := json.Unmarshal(v, &rec); err != nil {
return err
}
jobs = append(jobs, rec)
return nil
}); err != nil {
return err
}
// Sort newest-first: same ordering as ListScanJobs.
sort.Slice(jobs, func(i, j int) bool {
ci, cj := jobs[i].Created, jobs[j].Created
if ci.Equal(cj) {
return jobs[i].ID > jobs[j].ID
}
return ci.After(cj)
})
kept := 0
keptFindings := 0
var delErr error
for _, rec := range jobs {
// Retention applies to finished jobs only. A queued or running
// job is older than the job that just completed, so counting or
// deleting it here would drop work the scheduler still owns.
if rec.State == "queued" || rec.State == "running" {
continue
}
overJobCount := kept >= keepJobs
if overJobCount {
if delErr = deleteScanJobAndFindings(jb, fb, rec.ID); delErr != nil {
return delErr
}
pruned++
continue
}
jobFindings := 0
if maxTotalFindings > 0 {
// Count only far enough to decide whether this job fits. For
// the newest job, which is always kept, a saturated count is
// enough to force later non-empty jobs over the volume cap.
remaining := maxTotalFindings - keptFindings
if kept == 0 {
remaining = maxTotalFindings
}
if remaining < 0 {
remaining = 0
}
jobFindings = countFindingRowsUpTo(fb, rec.ID, remaining+1)
if kept >= 1 && jobFindings > 0 && keptFindings+jobFindings > maxTotalFindings {
if delErr = deleteScanJobAndFindings(jb, fb, rec.ID); delErr != nil {
return delErr
}
pruned++
continue
}
}
kept++
keptFindings += jobFindings
}
return nil
})
if err != nil {
return 0, err
}
return pruned, nil
}
// countFindingRowsUpTo counts finding rows stored for jobID, stopping early
// after limit rows when limit > 0. The volume-cap retention path only needs to
// know whether a job crosses a threshold, not its exact size above that point.
func countFindingRowsUpTo(fb *bolt.Bucket, jobID string, limit int) int {
prefix := findingPrefix(jobID)
n := 0
c := fb.Cursor()
for k, _ := c.Seek(prefix); k != nil && bytes.HasPrefix(k, prefix); k, _ = c.Next() {
n++
if limit > 0 && n >= limit {
break
}
}
return n
}
// deleteScanJobAndFindings removes a job record and all of its finding rows.
// Keys are collected before deletion because the cursor cannot be mutated mid-iteration.
func deleteScanJobAndFindings(jb, fb *bolt.Bucket, jobID string) error {
if err := jb.Delete([]byte(jobID)); err != nil {
return err
}
prefix := findingPrefix(jobID)
c := fb.Cursor()
var stale [][]byte
for k, _ := c.Seek(prefix); k != nil && bytes.HasPrefix(k, prefix); k, _ = c.Next() {
stale = append(stale, append([]byte(nil), k...))
}
for _, k := range stale {
if err := fb.Delete(k); err != nil {
return err
}
}
return nil
}
package store
import (
"encoding/json"
"errors"
"time"
bolt "go.etcd.io/bbolt"
)
// Persistence helpers for the signature-update watcher. The daemon stores
// what it last saw of each signature file in bbolt so a restart does not
// trigger a phantom rescan -- without this the in-memory map starts empty
// after every restart and every file looks new.
const sigWatchKey = "last_mtimes"
// SignatureFileState is what the watcher last saw of one rules file. SHA256
// is empty when the content was never hashed, and Size is -1 when unknown:
// both hold for entries recorded before content hashes were tracked.
type SignatureFileState struct {
Mtime time.Time `json:"mtime"`
Size int64 `json:"size"`
SHA256 string `json:"sha256,omitempty"`
}
// GetSignatureFiles returns the persisted watcher state. Empty (not nil) when
// nothing has been persisted yet. A map written before content hashes were
// tracked holds bare mtimes; its entries come back with Size -1 and no hash.
func (db *DB) GetSignatureFiles() (map[string]SignatureFileState, error) {
out := map[string]SignatureFileState{}
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("sig_watch"))
if b == nil {
return nil
}
raw := b.Get([]byte(sigWatchKey))
if len(raw) == 0 {
return nil
}
if err := json.Unmarshal(raw, &out); err == nil {
return nil
}
var legacy map[string]time.Time
if err := json.Unmarshal(raw, &legacy); err != nil {
return err
}
out = make(map[string]SignatureFileState, len(legacy))
for path, mtime := range legacy {
out[path] = SignatureFileState{Mtime: mtime, Size: -1}
}
return nil
})
return out, err
}
// PutSignatureFiles overwrites the persisted watcher state. Removed files
// must disappear from the store, so the whole map is replaced.
func (db *DB) PutSignatureFiles(m map[string]SignatureFileState) error {
payload, err := json.Marshal(m)
if err != nil {
return err
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("sig_watch"))
if b == nil {
return errors.New("sig_watch bucket missing (store not migrated)")
}
return b.Put([]byte(sigWatchKey), payload)
})
}
package store
import (
"encoding/json"
"fmt"
"time"
"github.com/pidginhost/csm/internal/alert"
bolt "go.etcd.io/bbolt"
)
// stats:daily holds pre-aggregated SeverityBucket counters keyed by
// "YYYY-MM-DD" date. It is updated atomically with every history insert
// so the 30-day trend chart is decoupled from history pruning. The
// bucket has at most dailyRetentionDays rows and grows by ~50 bytes/day.
const (
bucketStatsDaily = "stats:daily"
metaStatsDailyBackfilled = "stats:daily:backfilled"
)
// dailyRetentionDays caps how far back stats:daily keeps per-day rows.
// Var (not const) so tests can override.
var dailyRetentionDays = 365
// incrStatsDaily increments the per-severity counters for a single
// finding's date inside an existing bbolt write transaction. The caller
// owns the transaction; this helper does not commit.
func incrStatsDaily(tx *bolt.Tx, t time.Time, sev alert.Severity) error {
b := tx.Bucket([]byte(bucketStatsDaily))
if b == nil {
return fmt.Errorf("bucket %s missing", bucketStatsDaily)
}
key := []byte(t.Format("2006-01-02"))
var sb SeverityBucket
if v := b.Get(key); v != nil {
if err := json.Unmarshal(v, &sb); err != nil {
// Corrupted entry - reset rather than refusing to record.
sb = SeverityBucket{}
}
}
sb.Total++
switch sev {
case alert.Critical:
sb.Critical++
case alert.High:
sb.High++
case alert.Warning:
sb.Warning++
}
val, err := json.Marshal(sb)
if err != nil {
return err
}
return b.Put(key, val)
}
// pruneStatsDaily deletes stats:daily rows older than dailyRetentionDays.
// Cheap because the bucket is bounded to ~dailyRetentionDays entries and
// keys sort lexicographically as YYYY-MM-DD.
func pruneStatsDaily(tx *bolt.Tx, now time.Time) error {
b := tx.Bucket([]byte(bucketStatsDaily))
if b == nil {
return nil
}
cutoff := now.AddDate(0, 0, -(dailyRetentionDays - 1)).Format("2006-01-02")
// Collect first: the caller has just written today's row in this
// transaction, and a bbolt cursor that deletes and then steps with Next
// skips keys in a bucket already written by the transaction.
var stale [][]byte
c := b.Cursor()
for k, _ := c.First(); k != nil && string(k) < cutoff; k, _ = c.Next() {
stale = append(stale, append([]byte(nil), k...))
}
for _, k := range stale {
if err := b.Delete(k); err != nil {
return err
}
}
return nil
}
// BackfillStatsDaily seeds stats:daily from the history bucket on first
// run after upgrade. Idempotent: a meta sentinel ensures it only runs
// once. Safe on hosts where the meta:migrated sentinel was set before
// stats:daily existed.
func (db *DB) BackfillStatsDaily() error {
var alreadyDone bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
if v := tx.Bucket([]byte("meta")).Get([]byte(metaStatsDailyBackfilled)); v != nil {
alreadyDone = true
}
return nil
})
if alreadyDone {
return nil
}
// Read history in a View transaction and aggregate in memory so we
// don't hold the write lock while scanning potentially large history.
type counts struct {
c, h, w, total int
}
perDay := make(map[string]*counts)
err := db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("history"))
if b == nil {
return nil
}
c := b.Cursor()
for k, v := c.First(); k != nil; k, v = c.Next() {
var f alert.Finding
if err := json.Unmarshal(v, &f); err != nil {
continue
}
key := f.Timestamp.Format("2006-01-02")
cnt, ok := perDay[key]
if !ok {
cnt = &counts{}
perDay[key] = cnt
}
cnt.total++
switch f.Severity {
case alert.Critical:
cnt.c++
case alert.High:
cnt.h++
case alert.Warning:
cnt.w++
}
}
return nil
})
if err != nil {
return err
}
// Apply the aggregated counts and set the sentinel atomically. We
// merge into existing rows (additive) so the operation stays
// idempotent if the sentinel write somehow gets lost mid-flight.
return db.bolt.Update(func(tx *bolt.Tx) error {
// Re-check sentinel inside the write transaction in case another
// process raced ahead of us.
if v := tx.Bucket([]byte("meta")).Get([]byte(metaStatsDailyBackfilled)); v != nil {
return nil
}
b := tx.Bucket([]byte(bucketStatsDaily))
if b == nil {
return fmt.Errorf("bucket %s missing", bucketStatsDaily)
}
for key, cnt := range perDay {
var sb SeverityBucket
if v := b.Get([]byte(key)); v != nil {
if err := json.Unmarshal(v, &sb); err != nil {
sb = SeverityBucket{}
}
}
sb.Critical += cnt.c
sb.High += cnt.h
sb.Warning += cnt.w
sb.Total += cnt.total
val, mErr := json.Marshal(sb)
if mErr != nil {
return mErr
}
if pErr := b.Put([]byte(key), val); pErr != nil {
return pErr
}
}
return tx.Bucket([]byte("meta")).Put(
[]byte(metaStatsDailyBackfilled),
[]byte(time.Now().Format(time.RFC3339)),
)
})
}
package store
import (
"encoding/json"
"fmt"
"strings"
"time"
bolt "go.etcd.io/bbolt"
)
// Threat entry sources. Rows written before source tagging existed carry
// an empty Source; Expired classifies those by their reason text.
const (
ThreatSourceOperator = "operator"
ThreatSourceAutoBlock = "autoblock"
)
// legacyOperatorReasonPrefixes identifies pre-source-tagging rows added by a
// human. The Web UI manual/bulk block handlers are the only operator-facing
// writers this bucket ever had and they always used these exact reason
// strings; flat-file migration appended a "[YYYY-MM-DD]" suffix, hence the
// prefix match. Every other legacy row came from the temporary auto-block
// path, which historically wrote no-expiry rows that re-flagged the IP on
// every access after the block lapsed (permablock loop).
var legacyOperatorReasonPrefixes = []string{
"Manually blocked via CSM Web UI",
"Bulk blocked via CSM Web UI",
"Permanently blocked via CSM Web UI",
"Bulk permanently blocked via CSM Web UI",
}
// PermanentBlockEntry represents an IP blocked by the threat system.
// A zero ExpiresAt on an operator row means the row never expires. Auto-block
// rows must carry a real expiry; a zero expiry is treated as lapsed so it
// cannot become permanent reputation evidence.
type PermanentBlockEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
BlockedAt time.Time `json:"blocked_at"`
Source string `json:"source,omitempty"`
ExpiresAt time.Time `json:"expires_at,omitzero"`
}
// Expired reports whether the entry should no longer count as a live
// threat. Legacy rows (no source, no expiry) that do not match a known
// operator reason are treated as expired: they were written by the old
// temporary auto-block path and keeping them alive re-creates the
// permablock loop on upgraded hosts.
func (e PermanentBlockEntry) Expired(now time.Time) bool {
if e.Source == ThreatSourceAutoBlock && e.ExpiresAt.IsZero() {
return true
}
if !e.ExpiresAt.IsZero() {
return !e.ExpiresAt.After(now)
}
if e.Source != "" {
return false
}
for _, prefix := range legacyOperatorReasonPrefixes {
if strings.HasPrefix(e.Reason, prefix) {
return false
}
}
return true
}
// WhitelistEntry represents an IP that should bypass threat checks.
type WhitelistEntry struct {
IP string `json:"ip"`
ExpiresAt time.Time `json:"expires_at"`
Permanent bool `json:"permanent"`
}
// AddPermanentBlock adds an IP to the block list as a never-expiring
// operator entry. Only increments threats:count if the key is new.
// Overwrites any temp entry for the same IP: an explicit operator block
// upgrades it to permanent.
func (db *DB) AddPermanentBlock(ip, reason string) error {
return db.putThreatEntry(PermanentBlockEntry{
IP: ip,
Reason: reason,
BlockedAt: time.Now(),
Source: ThreatSourceOperator,
})
}
// AddTempBlock records an auto-blocked IP with an expiry matching the
// firewall block, so the entry lapses together with the block instead of
// flagging the IP forever. A zero expiresAt is deliberately ignored: an
// auto-block-sourced threat row must never become permanent evidence.
// Never downgrades an existing permanent row, and never shortens a longer
// temp expiry already on file. Only increments threats:count if the key is
// new.
func (db *DB) AddTempBlock(ip, reason string, expiresAt time.Time) error {
return db.addExpiringThreat(ip, reason, ThreatSourceAutoBlock, expiresAt)
}
// AddOperatorTempBlock records a timed operator block (for example the Web
// UI 24h block) with the same expiry as the firewall block. The row carries
// the operator source but lapses with the block, so a mistaken timed block
// cannot turn into permanent reputation evidence. Same guards as
// AddTempBlock: zero expiry ignored, permanent rows never downgraded, longer
// live expiries never shortened.
func (db *DB) AddOperatorTempBlock(ip, reason string, expiresAt time.Time) error {
return db.addExpiringThreat(ip, reason, ThreatSourceOperator, expiresAt)
}
func (db *DB) addExpiringThreat(ip, reason, source string, expiresAt time.Time) error {
if expiresAt.IsZero() {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
if err := invalidateUndoTargets(tx, []string{ip}); err != nil {
return err
}
b := tx.Bucket([]byte("threats"))
entry := PermanentBlockEntry{
IP: ip,
Reason: reason,
BlockedAt: time.Now(),
Source: source,
ExpiresAt: expiresAt,
}
existing := b.Get([]byte(ip))
if existing != nil {
var cur PermanentBlockEntry
if json.Unmarshal(existing, &cur) == nil && !cur.Expired(time.Now()) {
if cur.ExpiresAt.IsZero() {
return nil
}
if cur.ExpiresAt.After(expiresAt) {
return nil
}
}
}
val, err := json.Marshal(entry)
if err != nil {
return err
}
if err := b.Put([]byte(ip), val); err != nil {
return err
}
if existing == nil {
return incrCounter(tx, "threats:count", 1)
}
return nil
})
}
// putThreatEntry writes an entry as-is. Only increments threats:count if
// the key is new. Migration uses it directly so flat-file rows keep their
// empty Source and stay subject to legacy classification in Expired;
// routing them through AddPermanentBlock would stamp them as operator rows
// and bless historical auto-block poison as permanent.
func (db *DB) putThreatEntry(entry PermanentBlockEntry) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
if err := invalidateUndoTargets(tx, []string{entry.IP}); err != nil {
return err
}
b := tx.Bucket([]byte("threats"))
isNew := b.Get([]byte(entry.IP)) == nil
val, err := json.Marshal(entry)
if err != nil {
return err
}
if err := b.Put([]byte(entry.IP), val); err != nil {
return err
}
if isNew {
return incrCounter(tx, "threats:count", 1)
}
return nil
})
}
// PruneExpiredThreats deletes threat entries whose lifetime has lapsed,
// including legacy no-source auto-block rows (see Expired). Returns the
// count removed. Collect-then-delete because bbolt forbids mutation during
// ForEach iteration.
func (db *DB) PruneExpiredThreats() int {
var removed int
now := time.Now()
updateErr := boltUpdate(db.bolt, func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("threats"))
var toDelete [][]byte
_ = b.ForEach(func(k, v []byte) error {
var entry PermanentBlockEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if entry.Expired(now) {
keyCopy := make([]byte, len(k))
copy(keyCopy, k)
toDelete = append(toDelete, keyCopy)
}
return nil
})
for _, k := range toDelete {
if err := b.Delete(k); err != nil {
return err
}
removed++
}
if removed == 0 {
return nil
}
// Clamp instead of blind decrement: a bulk prune of legacy rows
// surfaces any historical counter drift at scale, and the counter
// must never go negative.
current := 0
if v := tx.Bucket([]byte("meta")).Get([]byte("threats:count")); v != nil {
_, _ = fmt.Sscanf(string(v), "%d", ¤t)
}
newCount := current - removed
if newCount < 0 {
newCount = 0
}
return setCounter(tx, "threats:count", newCount)
})
return committedCount("threats", updateErr, removed)
}
// RemovePermanentBlock removes an IP from the permanent block list and decrements the count.
func (db *DB) RemovePermanentBlock(ip string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
if err := invalidateUndoTargets(tx, []string{ip}); err != nil {
return err
}
b := tx.Bucket([]byte("threats"))
if b.Get([]byte(ip)) == nil {
return nil
}
if err := b.Delete([]byte(ip)); err != nil {
return err
}
return incrCounter(tx, "threats:count", -1)
})
}
// TiedToFirewallBlock reports whether the entry only lives as long as a
// firewall block: auto-block rows and any row carrying an expiry (timed
// operator blocks). Never-expiring operator rows and legacy no-source rows
// are standalone evidence.
func (e PermanentBlockEntry) TiedToFirewallBlock() bool {
return e.Source == ThreatSourceAutoBlock || !e.ExpiresAt.IsZero()
}
// RemoveTemporaryBlock deletes the threat row for ip only when it is tied to
// a firewall block (see TiedToFirewallBlock). Never-expiring operator rows
// and legacy no-source rows are left untouched, so a firewall-only unblock
// never silently clears an operator's deliberate permanent block. Stale
// timed rows would otherwise keep re-flagging the IP via ip_reputation.
// Returns whether a row was removed. The read and delete run in one
// transaction so a concurrent upgrade to a permanent row cannot be clobbered.
func (db *DB) RemoveTemporaryBlock(ip string) (bool, error) {
removed := false
err := db.bolt.Update(func(tx *bolt.Tx) error {
if err := invalidateUndoTargets(tx, []string{ip}); err != nil {
return err
}
b := tx.Bucket([]byte("threats"))
v := b.Get([]byte(ip))
if v == nil {
return nil
}
var entry PermanentBlockEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if !entry.TiedToFirewallBlock() {
return nil
}
if err := b.Delete([]byte(ip)); err != nil {
return err
}
removed = true
return incrCounter(tx, "threats:count", -1)
})
return removed, err
}
// GetPermanentBlock looks up a permanent block entry by IP.
// Returns the entry and true if found, or a zero value and false if not.
func (db *DB) GetPermanentBlock(ip string) (PermanentBlockEntry, bool) {
var entry PermanentBlockEntry
var found bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("threats"))
v := b.Get([]byte(ip))
if v == nil {
return nil
}
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
found = true
return nil
})
return entry, found
}
// AllPermanentBlocks returns all entries in the permanent block list
// (including expired ones - callers filter via Expired).
func (db *DB) AllPermanentBlocks() []PermanentBlockEntry {
var entries []PermanentBlockEntry
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("threats"))
return b.ForEach(func(k, v []byte) error {
var entry PermanentBlockEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
entries = append(entries, entry)
return nil
})
})
return entries
}
// AddWhitelistEntry adds an IP to the whitelist.
func (db *DB) AddWhitelistEntry(ip string, expiresAt time.Time, permanent bool) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
if err := invalidateUndoTargets(tx, []string{ip}); err != nil {
return err
}
b := tx.Bucket([]byte("threats:whitelist"))
entry := WhitelistEntry{
IP: ip,
ExpiresAt: expiresAt,
Permanent: permanent,
}
val, err := json.Marshal(entry)
if err != nil {
return err
}
return b.Put([]byte(ip), val)
})
}
// RemoveWhitelistEntry removes an IP from the whitelist.
func (db *DB) RemoveWhitelistEntry(ip string) error {
return db.bolt.Update(func(tx *bolt.Tx) error {
if err := invalidateUndoTargets(tx, []string{ip}); err != nil {
return err
}
b := tx.Bucket([]byte("threats:whitelist"))
return b.Delete([]byte(ip))
})
}
// IsWhitelisted checks if an IP is whitelisted and not expired.
func (db *DB) IsWhitelisted(ip string) bool {
var whitelisted bool
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("threats:whitelist"))
v := b.Get([]byte(ip))
if v == nil {
return nil
}
var entry WhitelistEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if entry.Permanent || entry.ExpiresAt.After(time.Now()) {
whitelisted = true
}
return nil
})
return whitelisted
}
// ListWhitelist returns all whitelist entries (including expired - caller filters).
func (db *DB) ListWhitelist() []WhitelistEntry {
var entries []WhitelistEntry
_ = db.bolt.View(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("threats:whitelist"))
return b.ForEach(func(k, v []byte) error {
var entry WhitelistEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
entries = append(entries, entry)
return nil
})
})
return entries
}
// PruneExpiredWhitelist deletes expired non-permanent whitelist entries.
// Returns the count of entries removed. Uses a collect-then-delete pattern
// because bbolt does not allow mutation during ForEach iteration.
func (db *DB) PruneExpiredWhitelist() int {
var removed int
now := time.Now()
updateErr := boltUpdate(db.bolt, func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("threats:whitelist"))
// Collect keys to delete.
var toDelete [][]byte
_ = b.ForEach(func(k, v []byte) error {
var entry WhitelistEntry
if json.Unmarshal(v, &entry) != nil {
return nil //nolint:nilerr // skip corrupt entry
}
if !entry.Permanent && !entry.ExpiresAt.After(now) {
keyCopy := make([]byte, len(k))
copy(keyCopy, k)
toDelete = append(toDelete, keyCopy)
}
return nil
})
// Delete collected keys.
for _, k := range toDelete {
if err := b.Delete(k); err != nil {
return err
}
removed++
}
return nil
})
return committedCount("whitelist", updateErr, removed)
}
package store
import (
"bytes"
"encoding/json"
"fmt"
"strconv"
"time"
"github.com/pidginhost/csm/internal/alert"
bolt "go.etcd.io/bbolt"
)
const (
timeKeyUTCMarker = "time_keys:utc_v1"
timeKeyHistoryTemp = "time_keys:migrate:history"
timeKeyAttacksTemp = "time_keys:migrate:attacks"
timeKeyAttacksIPTemp = "time_keys:migrate:attacks_ip"
timeKeyMigrationDone = "1"
historyBucketName = "history"
attackEventsBucketName = "attacks:events"
attackEventsIPBucket = "attacks:events:ip"
)
// migrateTimeKeysToUTC rewrites keys created by releases that formatted the
// wall clock from each timestamp's own location. Values retain the authoritative
// instant, so the rewrite can safely canonicalize mixed local/UTC stores. The
// buckets and marker are replaced in one transaction: a failed migration leaves
// the original indexes intact and is retried on the next Open.
func (db *DB) migrateTimeKeysToUTC() error {
var done bool
if err := db.bolt.View(func(tx *bolt.Tx) error {
meta := tx.Bucket([]byte("meta"))
done = meta != nil && string(meta.Get([]byte(timeKeyUTCMarker))) == timeKeyMigrationDone
return nil
}); err != nil {
return err
}
if done {
return nil
}
return db.bolt.Update(func(tx *bolt.Tx) error {
meta := tx.Bucket([]byte("meta"))
if string(meta.Get([]byte(timeKeyUTCMarker))) == timeKeyMigrationDone {
return nil
}
if err := migrateHistoryTimeKeys(tx); err != nil {
return fmt.Errorf("history: %w", err)
}
if err := migrateAttackTimeKeys(tx); err != nil {
return fmt.Errorf("attack events: %w", err)
}
return meta.Put([]byte(timeKeyUTCMarker), []byte(timeKeyMigrationDone))
})
}
func migrateHistoryTimeKeys(tx *bolt.Tx) error {
src := tx.Bucket([]byte(historyBucketName))
if bucketIsEmpty(src) {
return nil
}
tmp, err := freshMigrationBucket(tx, timeKeyHistoryTemp)
if err != nil {
return err
}
if src != nil {
c := src.Cursor()
for k, v := c.First(); k != nil; k, v = c.Next() {
var finding alert.Finding
if decodeErr := json.Unmarshal(v, &finding); decodeErr != nil {
if _, putErr := putPreservedValue(tmp, k, v); putErr != nil {
return putErr
}
continue
}
if _, putErr := putMigratedTimeValue(tmp, finding.Timestamp, timeKeyCounter(k), v); putErr != nil {
return putErr
}
}
}
if err := replaceBucketFromTemp(tx, historyBucketName, timeKeyHistoryTemp); err != nil {
return err
}
return bumpHistoryRevision(tx)
}
func migrateAttackTimeKeys(tx *bolt.Tx) error {
primary := tx.Bucket([]byte(attackEventsBucketName))
secondary := tx.Bucket([]byte(attackEventsIPBucket))
if bucketIsEmpty(primary) && bucketIsEmpty(secondary) {
return nil
}
tmpPrimary, err := freshMigrationBucket(tx, timeKeyAttacksTemp)
if err != nil {
return err
}
keyMap := make(map[string]string)
if primary != nil {
c := primary.Cursor()
for k, v := c.First(); k != nil; k, v = c.Next() {
oldKey := string(k)
var event AttackEvent
var newKey string
var putErr error
if decodeErr := json.Unmarshal(v, &event); decodeErr != nil {
newKey, putErr = putPreservedValue(tmpPrimary, k, v)
} else {
newKey, putErr = putMigratedTimeValue(tmpPrimary, event.Timestamp, timeKeyCounter(k), v)
}
if putErr != nil {
return putErr
}
keyMap[oldKey] = newKey
}
}
if replaceErr := replaceBucketFromTemp(tx, attackEventsBucketName, timeKeyAttacksTemp); replaceErr != nil {
return replaceErr
}
tmpSecondary, err := freshMigrationBucket(tx, timeKeyAttacksIPTemp)
if err != nil {
return err
}
if secondary != nil {
c := secondary.Cursor()
for k, v := c.First(); k != nil; k, v = c.Next() {
sep := bytes.LastIndexByte(k, '/')
if sep >= 0 {
if newTimeKey, ok := keyMap[string(k[sep+1:])]; ok {
newKey := string(k[:sep+1]) + newTimeKey
if _, err := putPreservedValue(tmpSecondary, []byte(newKey), v); err != nil {
return err
}
continue
}
}
if _, err := putPreservedValue(tmpSecondary, k, v); err != nil {
return err
}
}
}
return replaceBucketFromTemp(tx, attackEventsIPBucket, timeKeyAttacksIPTemp)
}
func bucketIsEmpty(b *bolt.Bucket) bool {
if b == nil {
return true
}
k, _ := b.Cursor().First()
return k == nil
}
func freshMigrationBucket(tx *bolt.Tx, name string) (*bolt.Bucket, error) {
if tx.Bucket([]byte(name)) != nil {
if err := tx.DeleteBucket([]byte(name)); err != nil {
return nil, err
}
}
return tx.CreateBucket([]byte(name))
}
func replaceBucketFromTemp(tx *bolt.Tx, name, tempName string) error {
if tx.Bucket([]byte(name)) != nil {
if err := tx.DeleteBucket([]byte(name)); err != nil {
return err
}
}
dst, err := tx.CreateBucket([]byte(name))
if err != nil {
return err
}
if name == historyBucketName || name == attackEventsBucketName {
// The temporary bucket's cursor yields sorted keys even when
// canonicalizing mixed time zones changed their original order.
dst.FillPercent = timeKeyFillPercent
}
tmp := tx.Bucket([]byte(tempName))
if tmp != nil {
if err := tmp.ForEach(func(k, v []byte) error {
return dst.Put(k, v)
}); err != nil {
return err
}
}
return tx.DeleteBucket([]byte(tempName))
}
func putMigratedTimeValue(b *bolt.Bucket, timestamp time.Time, start int, value []byte) (string, error) {
for counter := start; ; counter++ {
key := TimeKey(timestamp, counter)
if b.Get([]byte(key)) == nil {
return key, b.Put([]byte(key), value)
}
}
}
func putPreservedValue(b *bolt.Bucket, key, value []byte) (string, error) {
candidate := string(key)
for suffix := 0; b.Get([]byte(candidate)) != nil; suffix++ {
candidate = fmt.Sprintf("%s~%04d", key, suffix)
}
return candidate, b.Put([]byte(candidate), value)
}
func timeKeyCounter(key []byte) int {
sep := bytes.LastIndexByte(key, '-')
if sep < 0 || sep == len(key)-1 {
return 0
}
counter, err := strconv.Atoi(string(key[sep+1:]))
if err != nil || counter < 0 {
return 0
}
return counter
}
package store
import (
"fmt"
"os"
bolt "go.etcd.io/bbolt"
)
// boltUpdate runs fn in a read-write transaction. A variable so tests can
// make the commit fail after fn has run, which is the failure the prune
// helpers must report honestly.
var boltUpdate = func(b *bolt.DB, fn func(*bolt.Tx) error) error {
return b.Update(fn)
}
// committedCount turns a per-transaction deletion count into what actually
// happened: a transaction that failed to commit removed nothing, whatever
// the loop inside it counted.
func committedCount(what string, err error, removed int) int {
if err != nil {
fmt.Fprintf(os.Stderr, "store: %s prune failed to commit, %d row(s) kept: %v\n", what, removed, err)
return 0
}
return removed
}
package store
import (
"encoding/json"
"fmt"
"time"
bolt "go.etcd.io/bbolt"
)
// WPVerificationResult describes a completed attempt without storing command
// output, which may contain account credentials or arbitrary PHP output.
type WPVerificationResult struct {
State string `json:"state"`
Reason string `json:"reason,omitempty"`
}
// WPVerificationRecord retains the last result and a bounded failure streak.
type WPVerificationRecord struct {
WPVerificationResult
Account string `json:"account,omitempty"`
ObservedAt time.Time `json:"observed_at"`
AttemptAt time.Time `json:"attempt_at"`
Failures int `json:"failures"`
// Keep the preceding attempt so overlapping scans can finish out of order
// without inventing or losing a consecutive failure.
PreviousAttemptAt time.Time `json:"previous_attempt_at,omitzero"`
PreviousState string `json:"previous_state,omitempty"`
// AbsentAt fences off attempts from before this installation was removed.
AbsentAt time.Time `json:"absent_at,omitzero"`
}
type wpVerificationState struct {
Rows map[string]WPVerificationRecord `json:"rows"`
// Completed discovery watermarks prevent a late scan from resurrecting
// installations removed by a newer full or account-scoped discovery.
Discovery map[string]time.Time `json:"discovery"`
}
func wpVerificationKey(kind string) ([]byte, error) {
if kind != "core" && kind != "plugins" {
return nil, fmt.Errorf("unknown WordPress verification kind %q", kind)
}
return []byte("wp_verification:" + kind), nil
}
func readWPVerification(b *bolt.Bucket, key []byte) (wpVerificationState, error) {
s := wpVerificationState{}
if data := b.Get(key); data != nil {
if err := json.Unmarshal(data, &s); err != nil {
return s, fmt.Errorf("decode WordPress verification state: %w", err)
}
}
if s.Rows == nil {
s.Rows = make(map[string]WPVerificationRecord)
}
if s.Discovery == nil {
s.Discovery = make(map[string]time.Time)
}
return s, nil
}
// UpdateWPVerification atomically merges one scan's actual attempts. Discovery
// alone never increments failures. Repeated consumers of the same scan share
// at, so they cannot turn one failed attempt into a persistent failure.
func (db *DB) UpdateWPVerification(kind string, at time.Time, scope string, discovered map[string]string, results map[string]WPVerificationResult, complete bool) error {
key, err := wpVerificationKey(kind)
if err != nil {
return err
}
if at.IsZero() {
return fmt.Errorf("WordPress verification requires a scan time")
}
return db.bolt.Update(func(tx *bolt.Tx) error {
b := tx.Bucket([]byte("meta"))
s, err := readWPVerification(b, key)
if err != nil {
return err
}
for path, account := range discovered {
if scope != "" && account != scope {
continue
}
row, exists := s.Rows[path]
if !exists {
row.AbsentAt = s.Discovery[""]
if s.Discovery[account].After(row.AbsentAt) {
row.AbsentAt = s.Discovery[account]
}
}
if row.AbsentAt.After(at) || (row.Account != account && row.ObservedAt.After(at)) {
continue
}
if !row.ObservedAt.After(at) {
row.Account, row.ObservedAt = account, at
}
if result, ok := results[path]; ok {
if attemptErr := row.recordAttempt(at, result); attemptErr != nil {
return attemptErr
}
}
s.Rows[path] = row
}
if complete && !s.Discovery[""].After(at) && !s.Discovery[scope].After(at) {
for path, row := range s.Rows {
if _, found := discovered[path]; !found && (scope == "" || row.Account == scope) && !row.ObservedAt.After(at) {
delete(s.Rows, path)
}
}
if scope == "" {
for account, prior := range s.Discovery {
if !prior.After(at) {
delete(s.Discovery, account)
}
}
}
s.Discovery[scope] = at
}
data, err := json.Marshal(s)
if err != nil {
return err
}
return b.Put(key, data)
})
}
func (row *WPVerificationRecord) recordAttempt(at time.Time, result WPVerificationResult) error {
if at.Equal(row.AttemptAt) || !at.After(row.PreviousAttemptAt) {
return nil
}
switch result.State {
case "verified", "modified", "not_wordpress", "unverified":
default:
return fmt.Errorf("invalid WordPress verification result %q", result.State)
}
if at.After(row.AttemptAt) {
row.PreviousAttemptAt, row.PreviousState = row.AttemptAt, row.State
row.WPVerificationResult, row.AttemptAt = result, at
} else {
row.PreviousAttemptAt, row.PreviousState = at, result.State
}
row.Failures = 0
if row.State == "unverified" {
row.Failures = 1
if row.PreviousState == "unverified" {
row.Failures = 2
}
}
return nil
}
// WPVerification reports persisted coverage, returning read errors rather than
// presenting missing evidence as a clean empty inventory.
func (db *DB) WPVerification(kind string) (map[string]WPVerificationRecord, error) {
key, err := wpVerificationKey(kind)
if err != nil {
return nil, err
}
var rows map[string]WPVerificationRecord
err = db.bolt.View(func(tx *bolt.Tx) error {
s, readErr := readWPVerification(tx.Bucket([]byte("meta")), key)
rows = s.Rows
return readErr
})
return rows, err
}
package systemdrun
import "context"
// LookPathFunc resolves an executable, returning a non-nil error when it is not
// installed. Callers inject their own so tests do not depend on the host.
type LookPathFunc func(file string) (string, error)
// RunnerFunc executes a command and returns its output. Callers inject the
// runner their package already uses, which keeps command execution behind one
// abstraction per package instead of two.
type RunnerFunc func(ctx context.Context, name string, args ...string) ([]byte, error)
// Run executes name/args as a transient unit forked by PID 1, so the command
// escapes the caller's service sandbox, and returns its output.
//
// When systemd-run cannot be used -- not installed, or a harmless probe fails
// while the caller's context remains active -- the command runs directly
// instead. That is no worse than not having the wrapper at all, and it is the
// only honest option: giving up would disable the command entirely on a host
// where it might still work.
//
// The caller's command is attempted at most once. Its exit status and output
// cannot be told apart from systemd-run's own (--wait and --pipe propagate the
// unit's status and output verbatim), so a failure is never retried: running it
// again would repeat side effects like freezing or removing queued mail, and
// would mask the real error.
func Run(ctx context.Context, lookPath LookPathFunc, run RunnerFunc, opt Options, name string, args ...string) ([]byte, error) {
systemdRunPath, _ := lookPath("systemd-run")
if ctxErr := ctx.Err(); ctxErr != nil {
return nil, ctxErr
}
if systemdRunPath == "" {
return run(ctx, name, args...)
}
// Probe with a command that has no side effects, so the decision to fall
// back is made before the caller's command has had a chance to run.
probeName, probeArgs := Argv(systemdRunPath, opt, probeCommand)
_, probeErr := run(ctx, probeName, probeArgs...)
if ctxErr := ctx.Err(); ctxErr != nil {
return nil, ctxErr
}
if probeErr != nil {
return run(ctx, name, args...)
}
wrappedName, wrappedArgs := Argv(systemdRunPath, opt, name, args...)
return run(ctx, wrappedName, wrappedArgs...)
}
// probeCommand is the no-op used to test whether systemd-run can start a unit.
// /bin/true is part of coreutils, which systemd itself depends on.
const probeCommand = "/bin/true"
// Package systemdrun builds argv for running a command outside the calling
// service's systemd sandbox.
//
// CSM's daemon runs under ProtectSystem=strict with a narrow ReadWritePaths
// allow-list. Heavyweight system tools write paths that allow-list cannot
// reasonably enumerate, and some refuse to start at all when a path they need
// is read-only. Handing such a command to systemd-run makes PID 1 fork it as a
// transient unit, outside csm.service's mount namespace and seccomp filters.
//
// Never use --scope for this: scope mode wraps a child of the sandboxed
// process, so it inherits exactly the restrictions the wrapper exists to
// escape.
package systemdrun
import (
"fmt"
"strings"
"time"
)
// Options tunes how the transient unit is run.
type Options struct {
// Pipe connects the unit's stdio to the caller so its output can be
// captured. Without it the unit's stdout goes to the journal and the
// caller reads nothing. --pipe already implies waiting for the unit.
Pipe bool
// RuntimeMax bounds the unit's lifetime via RuntimeMaxSec. Zero omits the
// property, leaving the unit unbounded.
RuntimeMax time.Duration
}
// Argv returns the argv that runs name/args as a transient unit. When
// systemdRunPath is empty -- systemd-run is not installed, so there is no
// sandbox to escape either -- the command's own argv is returned unchanged.
func Argv(systemdRunPath string, opt Options, name string, args ...string) (string, []string) {
if systemdRunPath == "" {
return name, args
}
flags := []string{"--quiet", "--collect"}
if opt.Pipe {
flags = append(flags, "--pipe")
} else {
flags = append(flags, "--wait")
}
if opt.RuntimeMax > 0 {
flags = append(flags, "--property=RuntimeMaxSec="+formatRuntimeMax(opt.RuntimeMax))
}
flags = append(flags, "--")
flags = append(flags, name)
return systemdRunPath, append(flags, args...)
}
func formatRuntimeMax(d time.Duration) string {
// systemd time spans have microsecond granularity. Round a positive
// sub-microsecond duration up so it never becomes the special zero value,
// then render fractional seconds without float rounding.
microseconds := d / time.Microsecond
if d%time.Microsecond != 0 {
microseconds++
}
seconds := microseconds / 1_000_000
fraction := microseconds % 1_000_000
if fraction == 0 {
return fmt.Sprintf("%ds", seconds)
}
fractionText := strings.TrimRight(fmt.Sprintf("%06d", fraction), "0")
return fmt.Sprintf("%d.%ss", seconds, fractionText)
}
package threat
import (
"encoding/json"
"os"
"path/filepath"
"time"
"github.com/pidginhost/csm/internal/attackdb"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/store"
)
// IPIntelligence is the complete picture of an IP from all sources.
type IPIntelligence struct {
IP string `json:"ip"`
// Internal attack history
AttackRecord *attackdb.IPRecord `json:"attack_record,omitempty"`
LocalScore int `json:"local_score"`
// ThreatDB (feeds + permanent blocklist)
InThreatDB bool `json:"in_threat_db"`
ThreatDBSource string `json:"threat_db_source,omitempty"`
// ThreatDBPermanent marks local evidence that never lapses, so an
// unblocked IP keeps scoring as malicious and is flagged again on its
// next sighting. ThreatDBExpiresAt is set when the evidence lapses
// with the block that recorded it.
ThreatDBPermanent bool `json:"threat_db_permanent,omitempty"`
ThreatDBExpiresAt *time.Time `json:"threat_db_expires_at,omitempty"`
// AbuseIPDB cache
AbuseScore int `json:"abuse_score"`
AbuseCategory string `json:"abuse_category,omitempty"`
// Firewall state
CurrentlyBlocked bool `json:"currently_blocked"`
BlockReason string `json:"block_reason,omitempty"`
BlockedAt *time.Time `json:"blocked_at,omitempty"`
BlockExpiresAt *time.Time `json:"block_expires_at,omitempty"`
BlockPermanent bool `json:"block_permanent,omitempty"`
// GeoIP (populated by API layer, not by Lookup)
Country string `json:"country,omitempty"`
CountryName string `json:"country_name,omitempty"`
City string `json:"city,omitempty"`
ASN uint `json:"asn,omitempty"`
ASOrg string `json:"as_org,omitempty"`
Network string `json:"network,omitempty"`
// Composite
UnifiedScore int `json:"unified_score"`
Verdict string `json:"verdict"` // "clean", "suspicious", "malicious", "blocked"
}
// Lookup returns the full intelligence picture for an IP.
// All reads are local (no network calls). Pre-loads shared state files
// once (same as LookupBatch) to avoid re-reading per field.
func Lookup(ip, statePath string) *IPIntelligence {
abuseCache := loadFullAbuseCache(statePath)
blockMap := loadFullBlockState(statePath)
intel := &IPIntelligence{
IP: ip,
AbuseScore: -1, // not cached
}
// 1. Attack DB
if adb := attackdb.Global(); adb != nil {
if rec := adb.LookupIP(ip); rec != nil {
intel.AttackRecord = rec
intel.LocalScore = rec.ThreatScore
}
}
// 2. ThreatDB (feeds + permanent)
if tdb := checks.GetThreatDB(); tdb != nil {
applyThreatDBMatch(intel, tdb)
}
// 3. AbuseIPDB from pre-loaded cache
if entry, ok := abuseCache[ip]; ok {
intel.AbuseScore = entry.Score
intel.AbuseCategory = entry.Category
}
// 4. Block state from pre-loaded map
applyBlockState(intel, blockMap)
computeVerdict(intel)
return intel
}
// LookupBatch returns intelligence for multiple IPs efficiently.
// Pre-loads shared state files once instead of per-IP.
func LookupBatch(ips []string, statePath string) []*IPIntelligence {
results := make([]*IPIntelligence, len(ips))
abuseCache := loadFullAbuseCache(statePath)
blockMap := loadFullBlockState(statePath)
for i, ip := range ips {
intel := &IPIntelligence{
IP: ip,
AbuseScore: -1,
}
// Attack DB
if adb := attackdb.Global(); adb != nil {
if rec := adb.LookupIP(ip); rec != nil {
intel.AttackRecord = rec
intel.LocalScore = rec.ThreatScore
}
}
// ThreatDB
if tdb := checks.GetThreatDB(); tdb != nil {
applyThreatDBMatch(intel, tdb)
}
// AbuseIPDB from pre-loaded cache
if entry, ok := abuseCache[ip]; ok {
intel.AbuseScore = entry.Score
intel.AbuseCategory = entry.Category
}
// Block state from pre-loaded map
applyBlockState(intel, blockMap)
computeVerdict(intel)
results[i] = intel
}
return results
}
// applyThreatDBMatch copies the threat-DB match and its lifetime onto intel.
func applyThreatDBMatch(intel *IPIntelligence, tdb *checks.ThreatDB) {
match, found := tdb.LookupMatch(intel.IP)
if !found {
return
}
intel.InThreatDB = true
intel.ThreatDBSource = match.Source
intel.ThreatDBPermanent = match.Permanent
if !match.ExpiresAt.IsZero() {
t := match.ExpiresAt
intel.ThreatDBExpiresAt = &t
}
}
func computeVerdict(intel *IPIntelligence) {
intel.UnifiedScore = intel.LocalScore
if intel.AbuseScore > intel.UnifiedScore {
intel.UnifiedScore = intel.AbuseScore
}
if intel.InThreatDB && intel.UnifiedScore < 100 {
intel.UnifiedScore = 100
}
switch {
case intel.CurrentlyBlocked:
intel.Verdict = "blocked"
case intel.UnifiedScore >= 80:
intel.Verdict = "malicious"
case intel.UnifiedScore >= 40:
intel.Verdict = "suspicious"
default:
intel.Verdict = "clean"
}
}
func applyBlockState(intel *IPIntelligence, blockMap map[string]*blockEntry) {
bs, ok := blockMap[intel.IP]
if !ok {
return
}
intel.CurrentlyBlocked = true
intel.BlockReason = bs.reason
intel.BlockPermanent = bs.permanent
if !bs.blockedAt.IsZero() {
t := bs.blockedAt
intel.BlockedAt = &t
}
if !bs.expiresAt.IsZero() && bs.expiresAt.Year() > 1 {
t := bs.expiresAt
intel.BlockExpiresAt = &t
}
}
// --- AbuseIPDB cache reader ---
type abuseEntry struct {
Score int `json:"score"`
Category string `json:"category"`
}
func loadFullAbuseCache(statePath string) map[string]*abuseEntry {
result := make(map[string]*abuseEntry)
sixHoursAgo := time.Now().Add(-6 * time.Hour)
// Try bbolt store first - after migration the flat file is renamed to .bak.
if sdb := store.Global(); sdb != nil {
for ip, entry := range sdb.AllReputation() {
if entry.CheckedAt.Before(sixHoursAgo) || entry.Score < 0 {
continue // expired or error sentinel
}
result[ip] = &abuseEntry{Score: entry.Score, Category: entry.Category}
}
return result
}
// Fallback: flat-file JSON (pre-migration).
type cacheEntry struct {
Score int `json:"score"`
Category string `json:"category"`
CheckedAt time.Time `json:"checked_at"`
}
type cacheFile struct {
Entries map[string]*cacheEntry `json:"entries"`
}
// #nosec G304 -- filepath.Join under operator-configured statePath.
data, err := os.ReadFile(filepath.Join(statePath, "reputation_cache.json"))
if err != nil {
return result
}
var cf cacheFile
if json.Unmarshal(data, &cf) != nil || cf.Entries == nil {
return result
}
for ip, entry := range cf.Entries {
if entry.CheckedAt.Before(sixHoursAgo) || entry.Score < 0 {
continue // expired or error sentinel
}
result[ip] = &abuseEntry{Score: entry.Score, Category: entry.Category}
}
return result
}
// --- Firewall block state reader ---
type blockEntry struct {
reason string
blockedAt time.Time
expiresAt time.Time
permanent bool
}
func loadFullBlockState(statePath string) map[string]*blockEntry {
result := make(map[string]*blockEntry)
now := time.Now()
// Read the authoritative firewall engine state (flat-file state.json).
// The bbolt fw:blocked bucket is written only at migration, so it would
// return a stale snapshot rather than the live block set.
if fwState, err := firewall.LoadState(statePath); err == nil && fwState != nil {
for _, entry := range fwState.Blocked {
perm := entry.ExpiresAt.IsZero() || entry.ExpiresAt.Year() <= 1
result[entry.IP] = &blockEntry{
reason: entry.Reason,
blockedAt: entry.BlockedAt,
expiresAt: entry.ExpiresAt,
permanent: perm,
}
}
}
// CSM blocked_ips.json (legacy)
type csmEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
BlockedAt time.Time `json:"blocked_at"`
ExpiresAt time.Time `json:"expires_at"`
}
type csmFile struct {
IPs []csmEntry `json:"ips"`
}
// #nosec G304 -- filepath.Join under operator-configured statePath.
if data, err := os.ReadFile(filepath.Join(statePath, "blocked_ips.json")); err == nil {
var cf csmFile
if json.Unmarshal(data, &cf) == nil {
for _, entry := range cf.IPs {
if entry.ExpiresAt.IsZero() || now.Before(entry.ExpiresAt) {
if _, exists := result[entry.IP]; !exists {
perm := entry.ExpiresAt.IsZero() || entry.ExpiresAt.Year() <= 1
result[entry.IP] = &blockEntry{
reason: entry.Reason,
blockedAt: entry.BlockedAt,
expiresAt: entry.ExpiresAt,
permanent: perm,
}
}
}
}
}
}
return result
}
package threatintel
import "context"
// AbuseIPDBSource adapts the existing reputation lookup function to the
// Source interface. The underlying function lives in internal/checks/
// and is injected at construction (avoids an import cycle).
type AbuseIPDBSource struct {
lookup func(ctx context.Context, ip string) (int, error)
}
// NewAbuseIPDBSource wraps the provided lookup function. The caller
// (typically internal/checks/reputation.go) supplies a closure over its
// existing AbuseIPDB query path.
func NewAbuseIPDBSource(lookup func(context.Context, string) (int, error)) *AbuseIPDBSource {
return &AbuseIPDBSource{lookup: lookup}
}
func (a *AbuseIPDBSource) Name() string { return "abuseipdb" }
func (a *AbuseIPDBSource) Score(ctx context.Context, ip string) (int, error) {
return a.lookup(ctx, ip)
}
// Package threatintel -- bot allowlist + verification.
//
// botallowlist.go owns the embedded static IP CIDR ranges and the UA
// substring -> claimed-bot mapping. The static snapshots are a fast positive
// allow path: if the source IP falls inside a published bot range, skip the
// request without rDNS. The embedded snapshots are a trusted fallback; the
// bot-ranges auto-updater (botranges_update.go) refreshes them at runtime.
package threatintel
import (
_ "embed"
"encoding/json"
"net"
"strings"
"sync"
)
//go:embed embed/googlebot.json
var googlebotJSON []byte
//go:embed embed/bingbot.json
var bingbotJSON []byte
//go:embed embed/applebot.json
var applebotJSON []byte
// AI crawlers publish IP ranges rather than crawler reverse DNS, so they ship
// as embedded snapshots and are kept current by the auto-updater. OpenAI's
// three feeds (GPTBot, ChatGPT-User, OAI-SearchBot) all verify "gptbot".
//
//go:embed embed/openai-gptbot.json
var openaiGPTBotJSON []byte
//go:embed embed/openai-chatgpt-user.json
var openaiChatGPTUserJSON []byte
//go:embed embed/openai-searchbot.json
var openaiSearchBotJSON []byte
//go:embed embed/perplexitybot.json
var perplexitybotJSON []byte
// Anthropic publishes one combined crawler feed (ClaudeBot, Claude-User,
// Claude-SearchBot) at claude.com/crawling/bots.json. It documents IP-list
// verification, not reverse DNS, so the snapshot is the authoritative source.
//
//go:embed embed/claudebot.json
var claudebotJSON []byte
// BotRanges holds the parsed allowlist data, indexed by claimed-bot
// identity ("googlebot", "bingbot", "applebot").
type BotRanges struct {
byBot map[string][]*net.IPNet
}
type embedFile struct {
Prefixes []struct {
IPv4 string `json:"ipv4Prefix"`
IPv6 string `json:"ipv6Prefix"`
} `json:"prefixes"`
}
var (
defaultRanges *BotRanges
rangesOnce sync.Once
)
// DefaultRanges parses the embedded snapshots once and returns the
// global BotRanges. Safe to call concurrently.
func DefaultRanges() *BotRanges {
rangesOnce.Do(func() {
defaultRanges = &BotRanges{byBot: map[string][]*net.IPNet{}}
// A slice (not a map) so several feeds can share one identity:
// OpenAI's three crawler feeds all populate "gptbot".
for _, src := range []struct {
bot string
raw []byte
}{
{"googlebot", googlebotJSON},
{"bingbot", bingbotJSON},
{"applebot", applebotJSON},
{"gptbot", openaiGPTBotJSON},
{"gptbot", openaiChatGPTUserJSON},
{"gptbot", openaiSearchBotJSON},
{"perplexitybot", perplexitybotJSON},
{"claudebot", claudebotJSON},
} {
bot, raw := src.bot, src.raw
var f embedFile
if err := json.Unmarshal(raw, &f); err != nil {
continue
}
for _, p := range f.Prefixes {
if cidr := p.IPv4; cidr != "" {
if _, n, err := net.ParseCIDR(cidr); err == nil {
defaultRanges.byBot[bot] = append(defaultRanges.byBot[bot], n)
}
}
if cidr := p.IPv6; cidr != "" {
if _, n, err := net.ParseCIDR(cidr); err == nil {
defaultRanges.byBot[bot] = append(defaultRanges.byBot[bot], n)
}
}
}
}
})
return defaultRanges
}
// IPInBot reports whether the given IP falls inside the static range
// of the given bot identity.
func (r *BotRanges) IPInBot(ip net.IP, bot string) bool {
if ip == nil {
return false
}
for _, n := range r.byBot[bot] {
if n.Contains(ip) {
return true
}
}
for _, n := range fetchedRangesFor(bot) {
if n.Contains(ip) {
return true
}
}
return false
}
// IPInAnyBot reports whether the IP falls inside any published crawler
// range (Googlebot/Bingbot/Applebot snapshots). Unlike IPInBot it needs
// no claimed-UA, so callers that only have an IP (e.g. the incident
// correlator's whitelist backstop) can recognise a verified-crawler
// address. Deliberately covers only crawlers that publish authoritative
// IP ranges -- CDN edge ranges are NOT included, because legitimate and
// malicious traffic share a CDN's egress IPs and whitelisting them would
// hide attacks proxied through the CDN.
func (r *BotRanges) IPInAnyBot(ip net.IP) bool {
if ip == nil {
return false
}
for _, nets := range r.byBot {
for _, n := range nets {
if n.Contains(ip) {
return true
}
}
}
for _, nets := range fetchedRangesSnapshot() {
for _, n := range nets {
if n.Contains(ip) {
return true
}
}
}
return false
}
// AICrawlerRangePrefixCounts returns the active prefix count for each built-in
// AI crawler that has a vendor range feed. Embedded snapshots are counted along
// with any fetched overlay so status views never report "none" while the
// packaged fallback ranges are still active.
func AICrawlerRangePrefixCounts() map[string]int {
sourceBots := map[string]struct{}{}
for _, src := range DefaultRangeSources() {
sourceBots[src.Bot] = struct{}{}
}
seen := map[string]map[string]struct{}{}
add := func(bot string, nets []*net.IPNet) {
if len(nets) == 0 {
return
}
if seen[bot] == nil {
seen[bot] = map[string]struct{}{}
}
for _, n := range nets {
if n != nil {
seen[bot][n.String()] = struct{}{}
}
}
}
defaults := DefaultRanges()
for bot := range sourceBots {
add(bot, defaults.byBot[bot])
}
for bot, nets := range fetchedRangesSnapshot() {
add(bot, nets)
}
out := make(map[string]int, len(seen))
for bot, prefixes := range seen {
out[bot] = len(prefixes)
}
return out
}
// ClaimedBotFromUA returns the lower-case bot identity if the UA looks
// like a known bot. Empty string otherwise. Identities match BotDomains
// keys in botverify.go so the async verifier can look up the right
// DNS suffix list.
func ClaimedBotFromUA(ua string) string {
low := strings.ToLower(ua)
switch {
case strings.Contains(low, "googlebot"):
return "googlebot"
case strings.Contains(low, "bingbot"):
return "bingbot"
case strings.Contains(low, "applebot"):
return "applebot"
// Appendix A bots: no published static IP range.
case strings.Contains(low, "duckduckbot"):
return "duckduckbot"
case strings.Contains(low, "amazonbot"):
return "amazonbot"
case strings.Contains(low, "gptbot"),
strings.Contains(low, "chatgpt-user"),
strings.Contains(low, "oai-searchbot"):
return "gptbot"
case strings.Contains(low, "claudebot"),
strings.Contains(low, "claude-user"),
strings.Contains(low, "claude-searchbot"):
return "claudebot"
case strings.Contains(low, "perplexitybot"):
return "perplexitybot"
case strings.Contains(low, "meta-externalagent"),
strings.Contains(low, "meta-webindexer"),
strings.Contains(low, "facebookexternalhit"):
return "facebookbot"
case strings.Contains(low, "bravebot"):
return "bravebot"
// SEO backlink crawlers: no published static IP range, rDNS-verified.
case strings.Contains(low, "seranking"):
return "seranking"
default:
// Operator-configured bots (reputation.verified_bots) extend the
// built-in set without a code change.
return OperatorBotFromUA(low)
}
}
package threatintel
import (
"context"
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"os"
"sort"
"strings"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/atomicio"
"github.com/pidginhost/csm/internal/netutil"
)
// fetchedRanges is the runtime-updatable overlay of vendor-published IP ranges,
// keyed by bot identity. It augments the embedded snapshots in DefaultRanges so
// crawler allowlists refresh without a new release. Swapped wholesale; scan-path
// readers load the pointer lock-free.
var fetchedRanges atomic.Pointer[map[string][]*net.IPNet]
// PublishFetchedRanges installs the runtime range overlay. nil clears it.
func PublishFetchedRanges(m map[string][]*net.IPNet) {
if len(m) == 0 {
fetchedRanges.Store(nil)
setLastRefreshUnix(0)
return
}
cp := make(map[string][]*net.IPNet, len(m))
for k, v := range m {
cp[k] = append([]*net.IPNet(nil), v...)
}
fetchedRanges.Store(&cp)
}
func fetchedRangesFor(bot string) []*net.IPNet {
p := fetchedRanges.Load()
if p == nil {
return nil
}
return (*p)[bot]
}
func fetchedRangesSnapshot() map[string][]*net.IPNet {
p := fetchedRanges.Load()
if p == nil {
return nil
}
return *p
}
// FetchedRangesSnapshot returns a shallow copy of the current overlay for
// inspection or merge by the updater.
func FetchedRangesSnapshot() map[string][]*net.IPNet {
snap := fetchedRangesSnapshot()
out := make(map[string][]*net.IPNet, len(snap))
for k, v := range snap {
out[k] = append([]*net.IPNet(nil), v...)
}
return out
}
// lastRefreshUnix is the Unix time of the last successful vendor fetch. It is
// persisted in the cache file so a restart restores the genuine fetch time
// rather than resetting it. Zero means never refreshed.
var lastRefreshUnix atomic.Int64
func setLastRefreshUnix(ts int64) { lastRefreshUnix.Store(ts) }
// LastFetchedRangesRefresh returns when the AI-crawler ranges were last fetched
// from the vendor feeds, or the zero time if they never have been. The web UI
// surfaces this so operators can see whether the auto-updater is keeping the
// snapshot current.
func LastFetchedRangesRefresh() time.Time {
ts := lastRefreshUnix.Load()
if ts == 0 {
return time.Time{}
}
return time.Unix(ts, 0)
}
// RefreshFetchedRanges fetches every source, persists the merged overlay to
// cachePath (skipped when empty), and publishes it. The first successful feed
// for a bot identity replaces that bot's previous overlay; later successful
// feeds for the same identity append. A bot whose every feed fails keeps its
// previous overlay, so a full vendor outage never narrows the allowlist.
// Returns the number of bot identities refreshed.
func RefreshFetchedRanges(ctx context.Context, client *http.Client, sources []RangeSource, cachePath string) (int, error) {
merged := FetchedRangesSnapshot()
updated := map[string]struct{}{}
var lastErr error
for _, src := range sources {
nets, err := FetchRange(ctx, client, src.URL)
if err != nil {
lastErr = err
continue
}
if _, ok := updated[src.Bot]; !ok {
merged[src.Bot] = nil
updated[src.Bot] = struct{}{}
}
merged[src.Bot] = append(merged[src.Bot], nets...)
}
if len(updated) == 0 {
return 0, lastErr
}
refreshedAt := time.Now().Unix()
if cachePath != "" {
if err := saveFetchedRanges(cachePath, merged, refreshedAt); err != nil {
return 0, err
}
}
PublishFetchedRanges(merged)
setLastRefreshUnix(refreshedAt)
return len(updated), nil
}
// RangeSource is one vendor IP-range feed mapped to a bot identity. Multiple
// sources may share an identity (OpenAI publishes GPTBot, ChatGPT-User and
// OAI-SearchBot separately; all verify the "gptbot" identity).
type RangeSource struct {
Bot string
URL string
}
// DefaultRangeSources are the vendor feeds the auto-updater refreshes. URLs are
// stable, vendor-published JSON in the {prefixes:[{ipv4Prefix|ipv6Prefix}]}
// shape. Anthropic's feed is one combined list for ClaudeBot, Claude-User, and
// Claude-SearchBot, all of which verify the "claudebot" identity.
func DefaultRangeSources() []RangeSource {
return []RangeSource{
{Bot: "gptbot", URL: "https://openai.com/gptbot.json"},
{Bot: "gptbot", URL: "https://openai.com/chatgpt-user.json"},
{Bot: "gptbot", URL: "https://openai.com/searchbot.json"},
{Bot: "perplexitybot", URL: "https://www.perplexity.ai/perplexitybot.json"},
{Bot: "claudebot", URL: "https://claude.com/crawling/bots.json"},
}
}
const (
maxRangeBytes = 4 << 20 // 4 MiB; vendor feeds are well under this
maxPrefixesPerFeed = 100000 // sanity cap against a runaway feed
)
// ParseRangeJSON parses a vendor range feed and returns valid, public,
// suitably-narrow CIDRs. Unparseable, over-broad, or non-public entries are
// dropped (not errors) so one bad row cannot poison the feed, and the same
// guards as operator-configured ranges stop a compromised or mistaken feed from
// allowlisting the whole internet.
func ParseRangeJSON(data []byte) ([]*net.IPNet, error) {
var f embedFile
if err := json.Unmarshal(data, &f); err != nil {
return nil, fmt.Errorf("bot range feed: %w", err)
}
var out []*net.IPNet
for _, p := range f.Prefixes {
for _, cidr := range []string{p.IPv4, p.IPv6} {
if cidr == "" {
continue
}
_, n, err := net.ParseCIDR(strings.TrimSpace(cidr))
if err != nil {
continue
}
n = netutil.NormalizeIPNet(n)
if n == nil || !operatorBotIPRangeAllowed(n) {
continue
}
out = append(out, n)
if len(out) >= maxPrefixesPerFeed {
return out, nil
}
}
}
return out, nil
}
// FetchRange downloads and parses one vendor range feed.
func FetchRange(ctx context.Context, client *http.Client, url string) ([]*net.IPNet, error) {
if client == nil {
client = http.DefaultClient
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return nil, err
}
resp, err := client.Do(req)
if err != nil {
return nil, err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("bot range feed %s: HTTP %d", url, resp.StatusCode)
}
data, err := io.ReadAll(io.LimitReader(resp.Body, maxRangeBytes+1))
if err != nil {
return nil, err
}
if len(data) > maxRangeBytes {
return nil, fmt.Errorf("bot range feed %s: exceeds %d bytes", url, maxRangeBytes)
}
nets, err := ParseRangeJSON(data)
if err != nil {
return nil, err
}
if len(nets) == 0 {
return nil, fmt.Errorf("bot range feed %s: no valid prefixes", url)
}
return nets, nil
}
// rangeCacheFile is the on-disk persistence shape for the fetched overlay.
type rangeCacheFile struct {
Bots map[string][]string `json:"bots"`
// RefreshedAt is the Unix time the cache was written, i.e. when the ranges
// were last fetched. Persisted so the last-refresh time survives a restart.
RefreshedAt int64 `json:"refreshed_at,omitempty"`
}
// SaveFetchedRanges persists the overlay so a restart keeps the last-good feed
// until the next refresh completes. The write time is stamped as the
// last-refresh time.
func SaveFetchedRanges(path string, m map[string][]*net.IPNet) error {
return saveFetchedRanges(path, m, time.Now().Unix())
}
func saveFetchedRanges(path string, m map[string][]*net.IPNet, refreshedAt int64) error {
c := rangeCacheFile{Bots: map[string][]string{}, RefreshedAt: refreshedAt}
for bot, nets := range m {
strs := make([]string, 0, len(nets))
for _, n := range nets {
strs = append(strs, n.String())
}
sort.Strings(strs)
c.Bots[bot] = strs
}
data, err := json.Marshal(c)
if err != nil {
return err
}
return atomicio.AtomicWrite(path, 0o600, data)
}
// LoadFetchedRanges reads a previously saved overlay and publishes it. A
// missing file is not an error (first run). Entries are re-validated on load.
func LoadFetchedRanges(path string) error {
return loadFetchedRanges(path, true)
}
// LoadFetchedRangesRequired is the control-command variant: by the time the
// daemon is asked to reload, the CLI must already have written a cache file.
func LoadFetchedRangesRequired(path string) error {
return loadFetchedRanges(path, false)
}
func loadFetchedRanges(path string, allowMissing bool) error {
data, err := os.ReadFile(path) // #nosec G304 -- daemon-owned state path
if err != nil {
if allowMissing && os.IsNotExist(err) {
return nil
}
return err
}
var c rangeCacheFile
if err := json.Unmarshal(data, &c); err != nil {
return err
}
m := make(map[string][]*net.IPNet, len(c.Bots))
for bot, strs := range c.Bots {
for _, s := range strs {
_, n, err := net.ParseCIDR(s)
if err != nil {
continue
}
if n = netutil.NormalizeIPNet(n); n != nil && operatorBotIPRangeAllowed(n) {
m[bot] = append(m[bot], n)
}
}
}
PublishFetchedRanges(m)
setLastRefreshUnix(c.RefreshedAt)
return nil
}
package threatintel
import (
"context"
"errors"
"net"
"slices"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/queuehealth"
)
type resolver interface {
LookupAddr(ctx context.Context, ip string) ([]string, error)
LookupIP(ctx context.Context, network, host string) ([]net.IP, error)
}
// verifier owns one resolver + a domain suffix list per bot identity.
// One verifier per bot identity in practice; tests construct directly.
type verifier struct {
res resolver
domains []string // lower-case suffix list, e.g. "googlebot.com"
}
func newVerifier(r resolver, domains []string) *verifier {
low := make([]string, len(domains))
for i, d := range domains {
low[i] = strings.ToLower(d)
}
return &verifier{res: r, domains: low}
}
// LogicVersion identifies the current shape of the bot-verifier logic
// (BotDomains suffix list, ClaimedBotFromUA mapping, no-PTR semantics).
// Bump this whenever a change here would invalidate cache entries
// written by an older build -- for example, adding a new domain suffix
// that turns prior negatives into positives, or adding a new UA -> bot
// identity mapping. The daemon calls store.DB.EnsureBotVerifyLogicVersion
// at startup with this value; a mismatch wipes the botverify bucket so
// the next scan re-verifies every IP under the new rules.
const LogicVersion = 5
// ErrUnverifiable signals that the resolver returned no usable PTR for
// the source IP, so the verifier cannot prove or disprove the claimed
// bot identity. Callers treat this as fail-open: do not cache a verdict or
// flag as spoof. Genuine spoof signals -- PTR present but outside the
// bot's domain suffix list, or forward-confirm mismatch -- still return
// (false, nil).
var ErrUnverifiable = errors.New("bot verify: no PTR record for source IP")
// verify performs Google's official PTR + forward-A method. Returns
// (true, nil) on success, (false, nil) on a definitive negative
// (PTR resolves but does not belong to the claimed bot's domain, or
// forward-A fails to round-trip the IP), (false, ErrUnverifiable) when
// the IP has no PTR at all, and (false, err) on context cancellation
// or transient resolver failure. Both error paths cause the async
// worker to skip the verdict write so unverifiable IPs do not get pinned
// as spoof for the TTL window.
func (v *verifier) verify(ctx context.Context, ip net.IP, bot string) (bool, error) {
names, err := v.res.LookupAddr(ctx, ip.String())
if err != nil {
if ctxErr := ctx.Err(); ctxErr != nil {
return false, ctxErr
}
if isDNSNotFound(err) {
return false, ErrUnverifiable
}
return false, err
}
if len(names) == 0 {
return false, ErrUnverifiable
}
matched := ""
for _, n := range names {
ln := strings.ToLower(strings.TrimSuffix(n, "."))
for _, suf := range v.domains {
if strings.HasSuffix(ln, "."+suf) || ln == suf {
matched = ln
break
}
}
if matched != "" {
break
}
}
if matched == "" {
return false, nil
}
addrs, err := v.res.LookupIP(ctx, "ip", matched)
if err != nil {
if ctxErr := ctx.Err(); ctxErr != nil {
return false, ctxErr
}
if isDNSNotFound(err) {
return false, nil
}
return false, err
}
for _, a := range addrs {
if a.Equal(ip) {
return true, nil
}
}
return false, nil
}
func isDNSNotFound(err error) bool {
var dnsErr *net.DNSError
return errors.As(err, &dnsErr) && dnsErr.IsNotFound && !dnsErr.IsTemporary && !dnsErr.IsTimeout
}
// AsyncBotVerifier runs PTR+forward-A verify in a single background
// goroutine, deduplicating in-flight jobs. Result writes through the
// put callback (store.DB.PutBotVerify); reads happen from the scan
// hot path via store.DB.GetBotVerify with no goroutine.
type AsyncBotVerifier struct {
mu sync.Mutex
inflight map[string]time.Time
attempts map[string]botVerifyAttempt
ch chan verifyJob
v map[string]*verifier // bot identity -> verifier; guarded by mu
res resolver // retained so SetOperatorEntries can rebuild v
put func(net.IP, string, bool, time.Time) error
// unverifiable holds no-PTR results. A missing PTR rarely changes between
// scans, so repeat claims wait for the record to lapse instead of queuing
// the same lookup on every scan and crowding out crawlers not yet checked.
unverifiable UnverifiableRecords
// unverifiableSwept tracks sweep attempts, including failures, and is
// owned by the single worker that writes records.
unverifiableSwept time.Time
stats *queuehealth.Tracker
stop <-chan struct{}
closed bool
}
// Attempt history prevents each unresolved retry from renewing the initial
// grace. It is bounded by queue capacity. New sources can replace the oldest
// completed attempt after its initial cooldown, so repeated failures cannot
// reserve every slot for the full cache TTL. A live job or newly granted grace
// is never evicted; retries cannot slide that reservation indefinitely.
type botVerifyAttempt struct {
retryAfter time.Time
retainUntil time.Time
expiresAt time.Time
}
type verifyJob struct {
IP net.IP
Bot string
ticket queuehealth.Ticket
}
// BotDomains maps each claimed-bot identity to the DNS suffix list
// used for PTR + forward-A verification. Covers all bots that appear
// frequently in production traffic and have no published static IP
// range (Task 4 handles static-range bots via embedded JSON).
var BotDomains = map[string][]string{
"googlebot": {"googlebot.com", "google.com"},
"bingbot": {"search.msn.com"},
"applebot": {"applebot.apple.com", "apple.com"},
"duckduckbot": {"duckduckgo.com"},
"amazonbot": {"amazonbot.amazon", "amazon.com", "developer.amazon.com"},
"gptbot": {"openai.com"},
"claudebot": {"anthropic.com"},
"perplexitybot": {"perplexity.ai"},
"facebookbot": {"fbsv.net", "tfbnw.net", "facebook.com"},
"bravebot": {"brave.com"},
"seranking": {"seranking.com"},
}
// UnverifiableRecords persists when a claimed bot identity had no PTR. The
// verifier derives retry suppression and attempt history from that time and
// sweeps records older than the history. store.DB implements it.
type UnverifiableRecords interface {
PutBotVerifyUnverifiable(ip net.IP, bot string, observedAt time.Time) error
BotVerifyUnverifiable(ip net.IP, bot string) (observedAt time.Time, ok bool)
SweepBotVerifyUnverifiable(cutoff time.Time) (int, error)
}
// NewAsyncBotVerifier constructs an async verifier backed by the
// system resolver. put is store.DB.PutBotVerify or a test seam; records is
// the store, or nil to keep no-PTR results in retry history only.
func NewAsyncBotVerifier(put func(net.IP, string, bool, time.Time) error, records UnverifiableRecords) *AsyncBotVerifier {
res := net.DefaultResolver
a := &AsyncBotVerifier{
inflight: make(map[string]time.Time),
ch: make(chan verifyJob, 256),
v: make(map[string]*verifier),
res: res,
put: put,
unverifiable: records,
stats: queuehealth.New(256, time.Minute),
}
for bot, domains := range BotDomains {
a.v[bot] = newVerifier(res, domains)
}
return a
}
// SetOperatorEntries rebuilds the per-bot verifier set from the built-in
// BotDomains plus operator-configured entries. An operator entry naming a
// built-in extends that bot's suffix list; a new name adds its own verifier.
// Safe to call after Run has started (SIGHUP reload): v is swapped under mu,
// which the worker also holds when reading it.
func (a *AsyncBotVerifier) SetOperatorEntries(entries []BotEntry) {
entries = normalizeBotEntries(entries, false)
m := make(map[string]*verifier, len(BotDomains)+len(entries))
for bot, domains := range BotDomains {
m[bot] = newVerifier(a.res, domains)
}
for _, e := range entries {
if len(e.RDNSSuffixes) == 0 {
continue
}
if existing, ok := m[e.Name]; ok {
merged := append(append([]string(nil), existing.domains...), e.RDNSSuffixes...)
m[e.Name] = newVerifier(a.res, merged)
} else {
m[e.Name] = newVerifier(a.res, e.RDNSSuffixes)
}
}
a.mu.Lock()
a.v = m
a.mu.Unlock()
}
// Enqueue reports whether a job is queued or already in flight. Unsupported
// identities and unavailable capacity never receive pending treatment, and a
// source with a live no-PTR record is not queued again until it lapses.
func (a *AsyncBotVerifier) Enqueue(ip net.IP, bot string) bool {
key := bot + "|" + ip.String()
a.mu.Lock()
defer a.mu.Unlock()
// Serialize the record read with finish: a worker that writes after this
// read must still be in flight when we decide whether to admit a retry.
var recorded bool
if a.unverifiable != nil {
if observed, ok := a.unverifiable.BotVerifyUnverifiable(ip, bot); ok {
now := time.Now()
if !now.After(observed.Add(botVerifyUnverifiableTTL)) {
return false
}
// A lapsed record still counts as an attempt for as long as
// in-memory history would, so a restart or eviction cannot
// grant fresh pending treatment.
recorded = now.Before(observed.Add(botVerifyCacheTTL))
}
}
if a.closed {
a.stats.Lose(time.Now(), 1)
return false
}
select {
case <-a.stop:
a.stats.Lose(time.Now(), 1)
return false
default:
}
if _, ok := a.inflight[key]; ok {
return true
}
if ip == nil || a.v[bot] == nil {
return false
}
now := time.Now()
for attemptKey, previous := range a.attempts {
_, live := a.inflight[attemptKey]
if !live && !now.Before(previous.expiresAt) && !now.Before(previous.retryAfter) {
delete(a.attempts, attemptKey)
}
}
attempt, attempted := a.attempts[key]
if attempted && now.Before(attempt.retryAfter) {
return false
}
var replaceKey string
if !attempted && len(a.attempts) >= cap(a.ch) {
var oldest time.Time
for attemptKey, previous := range a.attempts {
if _, live := a.inflight[attemptKey]; live || now.Before(previous.retainUntil) {
continue
}
if replaceKey == "" || previous.expiresAt.Before(oldest) {
replaceKey = attemptKey
oldest = previous.expiresAt
}
}
if replaceKey == "" {
a.stats.Lose(now, 1)
return false
}
}
var pendingUntil time.Time
if !attempted && !recorded {
pendingUntil = now.Add(botVerifyTimeout)
}
a.inflight[key] = pendingUntil
// The caller may reuse its IP buffer as soon as admission returns. The
// queued lookup and its dedup key must retain the same address.
job := verifyJob{IP: slices.Clone(ip), Bot: bot, ticket: a.stats.Begin(time.Now())}
select {
case a.ch <- job:
if !attempted {
if a.attempts == nil {
a.attempts = make(map[string]botVerifyAttempt)
}
delete(a.attempts, replaceKey)
a.attempts[key] = botVerifyAttempt{
retainUntil: now.Add(botVerifyRetryDelay),
expiresAt: now.Add(botVerifyCacheTTL),
}
}
return true
default:
job.ticket.Reject(time.Now())
delete(a.inflight, key)
return false
}
}
// Pending is true only while an admitted job is live and its initial grace
// has not expired. Queue wait counts against the same bound as a DNS lookup.
func (a *AsyncBotVerifier) Pending(ip net.IP, bot string) bool {
a.mu.Lock()
defer a.mu.Unlock()
if a.closed || a.v[bot] == nil {
return false
}
select {
case <-a.stop:
return false
default:
}
until, ok := a.inflight[bot+"|"+ip.String()]
return ok && time.Now().Before(until)
}
const (
botVerifyTimeout = 5 * time.Second
botVerifyRetryDelay = time.Minute
botVerifyCacheTTL = 24 * time.Hour
// Shorter than a verdict: a crawler that gains a PTR is verified within
// the hour, while a stable no-PTR source costs one lookup per hour.
botVerifyUnverifiableTTL = time.Hour
)
func (a *AsyncBotVerifier) QueueStatuses(now time.Time) map[string]queuehealth.Status {
return map[string]queuehealth.Status{"requests": a.stats.Snapshot(now)}
}
// Run processes the queue until stopCh closes. Runs as a single
// goroutine so DNS calls are serialised; volume is bounded by the
// inflight dedup map so bursts do not launch unbounded goroutines.
//
// Closing stopCh cancels the parent context, so any in-flight verify
// returns from its DNS lookup immediately rather than holding the Run
// goroutine for the per-job 5s timeout.
func (a *AsyncBotVerifier) Run(stopCh <-chan struct{}) {
a.mu.Lock()
a.stop = stopCh
a.mu.Unlock()
ctx, cancel := context.WithCancel(context.Background())
bridge := make(chan struct{})
go func() {
defer close(bridge)
select {
case <-stopCh:
cancel()
case <-ctx.Done():
}
}()
defer func() {
cancel()
<-bridge
a.mu.Lock()
a.closed = true
close(a.ch)
a.mu.Unlock()
for job := range a.ch {
a.finish(job, false, false)
}
}()
for {
select {
case <-stopCh:
return
default:
}
select {
case <-ctx.Done():
return
case job := <-a.ch:
select {
case <-stopCh:
a.finish(job, false, false)
return
default:
}
a.processWithContext(ctx, job)
}
}
}
func (a *AsyncBotVerifier) process(job verifyJob) {
a.processWithContext(context.Background(), job)
}
func (a *AsyncBotVerifier) processWithContext(parent context.Context, job verifyJob) {
job.ticket.Start(time.Now())
completed, cached := false, false
defer func() { a.finish(job, completed, cached) }()
a.mu.Lock()
v, ok := a.v[job.Bot]
a.mu.Unlock()
if !ok {
completed = true
return
}
ctx, cancel := context.WithTimeout(parent, botVerifyTimeout)
defer cancel()
result, err := v.verify(ctx, job.IP, job.Bot)
cancel()
if err != nil {
completed = errors.Is(err, ErrUnverifiable) && a.recordUnverifiable(job)
return
}
if a.put == nil {
completed = true
return
}
cached = a.put(job.IP, job.Bot, result, time.Now().Add(botVerifyCacheTTL)) == nil
completed = cached
}
// recordUnverifiable reports whether a no-PTR result is settled. It is not a
// verdict, so a lapsed record prevents fresh pending grace while within the
// history window. An unwritten record leaves the work unaccounted for.
func (a *AsyncBotVerifier) recordUnverifiable(job verifyJob) bool {
if a.unverifiable == nil {
return true
}
now := time.Now()
if a.unverifiable.PutBotVerifyUnverifiable(job.IP, job.Bot, now) != nil {
return false
}
// Records past the attempt history affect nothing. Sweeping after a write
// bounds the bucket by recent no-PTR volume; once an hour keeps the scan
// off the per-result path. Failures also wait an hour, so a failing large
// transaction cannot hold up every subsequent result write.
if now.Sub(a.unverifiableSwept) >= botVerifyUnverifiableTTL {
a.unverifiableSwept = now
_, _ = a.unverifiable.SweepBotVerifyUnverifiable(now.Add(-botVerifyCacheTTL))
}
return true
}
func (a *AsyncBotVerifier) finish(job verifyJob, completed, cached bool) {
a.mu.Lock()
defer a.mu.Unlock()
key := job.Bot + "|" + job.IP.String()
delete(a.inflight, key)
if cached {
delete(a.attempts, key)
} else if attempt, tracked := a.attempts[key]; tracked {
attempt.retryAfter = time.Now().Add(botVerifyRetryDelay)
a.attempts[key] = attempt
}
if completed {
job.ticket.Finish(time.Now())
} else {
job.ticket.Reject(time.Now())
}
}
package threatintel
import (
"sync"
"sync/atomic"
"github.com/pidginhost/csm/internal/metrics"
)
var upstreamMetrics = struct {
mu sync.Mutex
registered map[*metrics.Registry]struct{}
active atomic.Pointer[UpstreamSource]
cacheHitsTotal atomic.Int64
cacheMissesTotal atomic.Int64
backendFailuresTotal atomic.Int64
}{
registered: make(map[*metrics.Registry]struct{}),
}
// RegisterUpstreamMetrics binds upstream counters to reg so operators
// can observe cache effectiveness and upstream health.
// Production callers pass metrics.Default(); tests pass an isolated
// registry. Idempotent per registry because reputation checks rebuild
// the upstream source every cycle.
func RegisterUpstreamMetrics(reg *metrics.Registry, src *UpstreamSource) {
if reg == nil || src == nil {
return
}
upstreamMetrics.active.Store(src)
upstreamMetrics.mu.Lock()
if _, ok := upstreamMetrics.registered[reg]; ok {
upstreamMetrics.mu.Unlock()
return
}
upstreamMetrics.registered[reg] = struct{}{}
upstreamMetrics.mu.Unlock()
registerUpstreamMetricsLocked(reg)
}
// ClearUpstreamMetricsSource clears the source used by the breaker
// gauge when upstream reputation is disabled by a hot-reloaded config.
func ClearUpstreamMetricsSource() {
upstreamMetrics.active.Store(nil)
}
func registerUpstreamMetricsLocked(reg *metrics.Registry) {
reg.RegisterCounterFunc(
"csm_threatintel_cache_hits_total",
"Upstream threat-intel cache hits.",
func() float64 {
return float64(upstreamMetrics.cacheHitsTotal.Load())
},
)
reg.RegisterCounterFunc(
"csm_threatintel_cache_misses_total",
"Upstream threat-intel lookups not served from the local cache.",
func() float64 {
return float64(upstreamMetrics.cacheMissesTotal.Load())
},
)
reg.RegisterCounterFunc(
"csm_threatintel_backend_failures_total",
"Upstream threat-intel backend failures (network, 4xx, 5xx, malformed body).",
func() float64 {
return float64(upstreamMetrics.backendFailuresTotal.Load())
},
)
reg.RegisterGaugeFunc(
"csm_threatintel_breaker_open",
"Circuit breaker for the upstream source; 1 when open (calls refused), 0 when closed or half-open.",
func() float64 {
src := activeUpstreamMetricsSource()
if src != nil && src.BreakerOpen() {
return 1
}
return 0
},
)
}
func activeUpstreamMetricsSource() *UpstreamSource {
return upstreamMetrics.active.Load()
}
func resetUpstreamMetricsForTest() {
upstreamMetrics.active.Store(nil)
upstreamMetrics.cacheHitsTotal.Store(0)
upstreamMetrics.cacheMissesTotal.Store(0)
upstreamMetrics.backendFailuresTotal.Store(0)
}
package threatintel
import (
"hash/fnv"
"net"
"sort"
"strconv"
"strings"
"sync/atomic"
"github.com/pidginhost/csm/internal/netutil"
)
// BotEntry is an operator-configured verified bot: claimed-UA substrings
// mapped to the rDNS suffixes or IP ranges that confirm it. These are additive
// on top of the built-in allowlist (BotDomains / the ClaimedBotFromUA switch /
// the embedded static IP ranges); operators extend coverage for crawlers CSM
// does not ship without a code change.
type BotEntry struct {
Name string
UASubstrings []string
RDNSSuffixes []string
// IPRanges are CIDRs (or single IPs) for bots that verify by address
// rather than reverse DNS -- AI agents that publish ranges and have no
// crawler-domain rDNS. Membership is checked synchronously on the scan
// path, so these need no async PTR lookup.
IPRanges []string
}
// operatorBots holds the active operator list. Swapped wholesale on reload;
// readers run on the scan hot path so the pointer load must stay lock-free.
var operatorBots atomic.Pointer[[]BotEntry]
// operatorNets is the parsed IP-range index (bot name -> CIDRs), kept in sync
// with operatorBots. Stored before operatorBots on update so a reader that
// sees a new bot name already finds its ranges.
var operatorNets atomic.Pointer[map[string][]*net.IPNet]
// SetOperatorBots installs the operator-configured bot list. Names,
// substrings, and suffixes are lower-cased and trimmed (leading dots on
// suffixes dropped) so matching is case-insensitive and consistent with the
// built-in tables. Entries without a name or any UA substring are skipped:
// the UA substring is what links a request to the identity.
func SetOperatorBots(entries []BotEntry) {
norm := normalizeBotEntries(entries, true)
nets := make(map[string][]*net.IPNet)
for _, e := range norm {
for _, r := range e.IPRanges {
if n := netutil.ParseCIDROrIP(r); n != nil && operatorBotIPRangeAllowed(n) {
nets[e.Name] = append(nets[e.Name], n)
}
}
}
operatorNets.Store(&nets)
operatorBots.Store(&norm)
}
// operatorBotIPRangeAllowed rejects ranges too broad to be a real crawler fleet
// and any non-public address space, so the allowlist cannot be turned into a
// blanket detection bypass. The public-space test is shared with the config
// validator via internal/netutil.
func operatorBotIPRangeAllowed(n *net.IPNet) bool {
ones, bits := n.Mask.Size()
if bits == 32 && ones < 16 {
return false
}
if bits == 128 && ones < 32 {
return false
}
return netutil.IsPublicIP(n.IP)
}
// IPInAnyVerifiedBotRange reports whether ip belongs to any verified-bot IP
// range: the built-in/auto-updated crawler snapshots plus operator
// reputation.verified_bots IP-range entries. Unlike the UA-keyed lookups it
// needs only an address, so the firewall auto-block guard can recognise a
// published crawler from the finding IP alone. rDNS-only bots have no range
// here and are handled on the request path, not by this address check.
func IPInAnyVerifiedBotRange(ip net.IP) bool {
if ip == nil {
return false
}
if DefaultRanges().IPInAnyBot(ip) {
return true
}
p := operatorNets.Load()
if p == nil {
return false
}
for _, nets := range *p {
for _, n := range nets {
if n.Contains(ip) {
return true
}
}
}
return false
}
// OperatorBotIPVerified reports whether ip falls in any IP range configured
// for the named operator bot. Synchronous; no DNS. This is how AI agents that
// publish address ranges (PerplexityBot, GPTBot, ClaudeBot) are confirmed.
func OperatorBotIPVerified(name string, ip net.IP) bool {
if name == "" || ip == nil {
return false
}
p := operatorNets.Load()
if p == nil {
return false
}
for _, n := range (*p)[name] {
if n.Contains(ip) {
return true
}
}
return false
}
func normalizeBotEntries(entries []BotEntry, requireUA bool) []BotEntry {
norm := make([]BotEntry, 0, len(entries))
for _, e := range entries {
ne := BotEntry{Name: normalizeBotName(e.Name)}
seenUA := map[string]struct{}{}
for _, raw := range e.UASubstrings {
s := normalizeUASubstring(raw)
if s == "" {
continue
}
if _, ok := seenUA[s]; ok {
continue
}
seenUA[s] = struct{}{}
ne.UASubstrings = append(ne.UASubstrings, s)
}
seenSuffix := map[string]struct{}{}
for _, raw := range e.RDNSSuffixes {
d := normalizeSuffix(raw)
if d == "" {
continue
}
if _, ok := seenSuffix[d]; ok {
continue
}
seenSuffix[d] = struct{}{}
ne.RDNSSuffixes = append(ne.RDNSSuffixes, d)
}
seenRange := map[string]struct{}{}
for _, raw := range e.IPRanges {
n := netutil.ParseCIDROrIP(raw)
if n == nil {
continue
}
r := n.String()
if _, ok := seenRange[r]; ok {
continue
}
seenRange[r] = struct{}{}
ne.IPRanges = append(ne.IPRanges, r)
}
if ne.Name == "" || (requireUA && len(ne.UASubstrings) == 0) {
continue
}
norm = append(norm, ne)
}
return norm
}
func normalizeBotName(name string) string {
return strings.ToLower(strings.TrimSpace(name))
}
func normalizeUASubstring(s string) string {
return strings.ToLower(strings.TrimSpace(s))
}
func normalizeSuffix(d string) string {
return strings.TrimPrefix(strings.ToLower(strings.TrimSpace(d)), ".")
}
func operatorBotEntries() []BotEntry {
p := operatorBots.Load()
if p == nil {
return nil
}
return *p
}
// OperatorBotFromUA returns the operator bot identity whose UA substring
// matches ua, or "" if none. Built-in identities are checked first by
// ClaimedBotFromUA; this is the fallback.
func OperatorBotFromUA(lowUA string) string {
lowUA = strings.ToLower(lowUA)
for _, e := range operatorBotEntries() {
for _, s := range e.UASubstrings {
if strings.Contains(lowUA, s) {
return e.Name
}
}
}
return ""
}
func operatorDomains(name string) []string {
for _, e := range operatorBotEntries() {
if e.Name == name {
return e.RDNSSuffixes
}
}
return nil
}
// OperatorBotsCacheVersion folds the operator bot list into the base cache
// logic version. Changing verified_bots therefore changes the stamp the
// daemon hands EnsureBotVerifyLogicVersion, which drops the PTR-verdict cache
// so a previously-spoofed IP is re-checked under the new suffixes instead of
// staying pinned for the cache TTL. The hash is order-independent so the same
// set in a different file order yields the same stamp.
func OperatorBotsCacheVersion(base int, entries []BotEntry) int {
entries = normalizeBotEntries(entries, false)
lines := make([]string, 0, len(entries))
for _, e := range entries {
subs := append([]string(nil), e.UASubstrings...)
sufs := append([]string(nil), e.RDNSSuffixes...)
rngs := append([]string(nil), e.IPRanges...)
sort.Strings(subs)
sort.Strings(sufs)
sort.Strings(rngs)
lines = append(lines, e.Name+
"\x00"+strings.Join(subs, ",")+
"\x00"+strings.Join(sufs, ",")+
"\x00"+strings.Join(rngs, ","))
}
sort.Strings(lines)
h := fnv.New64a()
_, _ = h.Write([]byte(strconv.Itoa(base)))
for _, l := range lines {
_, _ = h.Write([]byte{0x01})
_, _ = h.Write([]byte(l))
}
return int(h.Sum64() & 0x7fffffff)
}
package threatintel
import (
"context"
"encoding/json"
"fmt"
"io"
"math"
"net"
"net/http"
"net/url"
"os"
"strings"
"time"
)
const rspamdMaxHistoryBytes = 2 << 20
const (
// rspamdPriorMass is Laplace-style smoothing added to the history
// mass so small samples cannot saturate the score: one reject alone
// scores 33 and it takes two fresh rejects with zero delivered ham
// to reach the 50 auto-block threshold used by the reputation check.
rspamdPriorMass = 2.0
// rspamdDecayHalfLife halves a row's influence per week so the score
// tracks recent behaviour instead of lifetime accumulation; rspamd's
// rolling history can span months on quiet servers.
rspamdDecayHalfLife = 7 * 24 * time.Hour
)
// RspamdSource queries rspamd's rolling history and returns a score
// 0..100 derived only from rows matching the requested IP.
//
// Token resolution reads the process environment at Score time. External
// environment changes require a daemon restart.
type RspamdSource struct {
url string
token string // static token from config (may be empty)
tokenEnv string // env var name to consult at query time
client *http.Client
}
func NewRspamdSource(url, token, tokenEnv string) *RspamdSource {
return &RspamdSource{
url: url,
token: token,
tokenEnv: tokenEnv,
client: &http.Client{Timeout: 5 * time.Second},
}
}
func (s *RspamdSource) Name() string { return "rspamd" }
// resolveToken reads the env var (if set) at query time, falling back to
// the static token. External environment changes require a daemon restart.
func (s *RspamdSource) resolveToken() string {
if s.tokenEnv != "" {
if v := os.Getenv(s.tokenEnv); v != "" {
return v
}
}
return s.token
}
type rspamdHistoryResp struct {
Rows []rspamdHistoryRow `json:"rows"`
History []rspamdHistoryRow `json:"history"`
Data []rspamdHistoryRow `json:"data"`
}
func (r rspamdHistoryResp) entries() []rspamdHistoryRow {
out := make([]rspamdHistoryRow, 0, len(r.Rows)+len(r.History)+len(r.Data))
out = append(out, r.Rows...)
out = append(out, r.History...)
out = append(out, r.Data...)
return out
}
type rspamdHistoryRow struct {
IP string
Action string
Score float64
UnixTime float64
}
func (r *rspamdHistoryRow) UnmarshalJSON(data []byte) error {
var raw map[string]json.RawMessage
if err := json.Unmarshal(data, &raw); err != nil {
return err
}
r.IP = firstJSONString(raw, "ip", "sender_ip", "client_ip")
r.Action = firstJSONString(raw, "action", "metric_action")
r.Score = firstJSONFloat(raw, "score")
r.UnixTime = firstJSONFloat(raw, "unix_time")
return nil
}
func firstJSONString(raw map[string]json.RawMessage, names ...string) string {
for _, name := range names {
v, ok := raw[name]
if !ok {
continue
}
var s string
if err := json.Unmarshal(v, &s); err == nil {
return s
}
}
return ""
}
func firstJSONFloat(raw map[string]json.RawMessage, names ...string) float64 {
for _, name := range names {
v, ok := raw[name]
if !ok {
continue
}
var f float64
if err := json.Unmarshal(v, &f); err == nil {
return f
}
}
return 0
}
// Score sends a GET to <url>/history and scores only history rows for ip.
func (s *RspamdSource) Score(ctx context.Context, ip string) (int, error) {
endpoint, err := s.historyEndpoint()
if err != nil {
return 0, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return 0, err
}
if tok := s.resolveToken(); tok != "" {
req.Header.Set("Password", tok)
}
resp, err := s.client.Do(req)
if err != nil {
return 0, fmt.Errorf("rspamd: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return 0, fmt.Errorf("rspamd HTTP %d", resp.StatusCode)
}
rows, err := decodeRspamdHistory(resp.Body)
if err != nil {
return 0, fmt.Errorf("rspamd decode: %w", err)
}
return scoreRspamdHistory(rows, ip, time.Now()), nil
}
func (s *RspamdSource) historyEndpoint() (string, error) {
u, err := url.Parse(s.url)
if err != nil {
return "", err
}
if u.Scheme == "" || u.Host == "" {
return "", fmt.Errorf("rspamd URL must include scheme and host")
}
u.Path = strings.TrimRight(u.Path, "/") + "/history"
u.RawQuery = ""
return u.String(), nil
}
func decodeRspamdHistory(r io.Reader) ([]rspamdHistoryRow, error) {
var raw json.RawMessage
if err := json.NewDecoder(io.LimitReader(r, rspamdMaxHistoryBytes)).Decode(&raw); err != nil {
return nil, err
}
trimmed := strings.TrimSpace(string(raw))
if strings.HasPrefix(trimmed, "[") {
var rows []rspamdHistoryRow
if err := json.Unmarshal(raw, &rows); err != nil {
return nil, err
}
return rows, nil
}
var history rspamdHistoryResp
if err := json.Unmarshal(raw, &history); err != nil {
return nil, err
}
return history.entries(), nil
}
// scoreRspamdHistory converts an IP's history rows into a 0..100
// confidence that the IP is a spam source. The score is the recency-
// weighted proportion of definitive spam verdicts among classifiable
// delivered/spam verdicts, smoothed by rspamdPriorMass. Delivered ham
// therefore dilutes the score toward 0 instead of accumulating it: a
// correspondent MTA that sends mostly legitimate mail stays near 0 no
// matter how much it sends, while an IP whose recent traffic is
// predominantly rejected climbs toward 100.
func scoreRspamdHistory(rows []rspamdHistoryRow, ip string, now time.Time) int {
want := normalizeIP(ip)
var spamMass, totalMass float64
for _, row := range rows {
if normalizeIP(row.IP) != want {
continue
}
spamWeight, counts := rspamdActionSignal(row.Action)
if !counts {
continue
}
recency := rspamdRecencyFactor(row.UnixTime, now)
spamMass += spamWeight * recency
totalMass += recency
}
if spamMass == 0 || totalMass == 0 {
return 0
}
return int(math.Round(100 * spamMass / (totalMass + rspamdPriorMass)))
}
// rspamdActionSignal maps an rspamd action onto per-message spam mass
// in [0,1] and reports whether the row is classifiable as ham or spam.
// "no action" is delivered ham; greylist and soft reject are tempfail
// flow control that fires on first contact from every unknown sender
// and on rate limits, so they are neutral instead of ham. Genuinely
// spammy retries earn reject, quarantine, discard, add-header, or
// rewrite-subject rows, which do count. The numeric rspamd score is
// ignored on purpose: the action already is rspamd's calibrated
// thresholding of that score.
func rspamdActionSignal(action string) (float64, bool) {
switch strings.ToLower(strings.TrimSpace(action)) {
case "reject", "discard", "quarantine", "spam":
return 1.0, true
case "add header", "rewrite subject", "probable spam":
return 0.7, true
case "no action", "clean":
return 0, true
default:
return 0, false
}
}
// rspamdRecencyFactor halves a row's weight per rspamdDecayHalfLife.
// Rows without a parseable unix_time count at full weight rather than
// being dropped: rspamd's rolling history is bounded, so undated rows
// are treated as current.
func rspamdRecencyFactor(unixTime float64, now time.Time) float64 {
if unixTime <= 0 {
return 1
}
age := now.Sub(time.Unix(int64(unixTime), 0))
if age <= 0 {
return 1
}
return math.Exp2(-age.Hours() / rspamdDecayHalfLife.Hours())
}
func normalizeIP(ip string) string {
ip = strings.TrimSpace(ip)
parsed := net.ParseIP(ip)
if parsed == nil {
return ip
}
return parsed.String()
}
// Package threatintel defines a pluggable interface for IP reputation
// providers and an Aggregator that combines their scores. CSM uses this
// to consult AbuseIPDB plus optional rspamd / upstream sources without
// hardcoding multiple lookup paths in the reputation check.
package threatintel
import "context"
// Source is a single scoring provider. Score returns 0..100 (higher is
// worse). A source that has no opinion on the IP must return 0, nil -
// that score is excluded from the aggregator's average. Errors are
// per-source and non-fatal at the aggregator level (other sources still
// run).
type Source interface {
Name() string
Score(ctx context.Context, ip string) (int, error)
}
// Aggregator runs every registered source and averages their non-zero scores.
type Aggregator struct {
sources []Source
}
// NewAggregator constructs an empty Aggregator.
func NewAggregator() *Aggregator { return &Aggregator{} }
// Register adds a Source to the aggregator. Order matters only for the
// `Sources` field of Result (which preserves registration order); the
// aggregated score is order-independent.
func (a *Aggregator) Register(s Source) { a.sources = append(a.sources, s) }
// Result holds the aggregated value and per-source breakdown.
type Result struct {
AggregatedScore int `json:"aggregated_score"`
Sources map[string]int `json:"sources"`
}
// Score queries every registered source. Per-source errors are swallowed
// (the source contributes "no signal"). The aggregated score is the mean
// of non-zero scores; if every source returned 0 (or errored), the
// aggregated score is 0.
func (a *Aggregator) Score(ctx context.Context, ip string) (Result, error) {
out := Result{Sources: map[string]int{}}
sum, n := 0, 0
for _, s := range a.sources {
score, err := s.Score(ctx, ip)
if err != nil {
continue
}
out.Sources[s.Name()] = score
if score > 0 {
sum += score
n++
}
}
if n > 0 {
out.AggregatedScore = sum / n
}
return out, nil
}
package threatintel
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strings"
"sync"
"sync/atomic"
"time"
)
const (
upstreamMaxResponseBytes = 1 << 20
// defaultMaxCacheEntries caps the per-IP score cache so a sustained
// flood of unique attacker IPs cannot grow the map without bound.
// Once the cap is reached, expired entries are pruned first; if the
// cap is still exceeded the oldest entry (by expires) is evicted.
defaultMaxCacheEntries = 10000
// defaultBreakerTrip is the number of consecutive upstream failures
// after which the source short-circuits subsequent Score calls.
defaultBreakerTrip = 5
// defaultBreakerCooldown is how long the breaker stays open before
// allowing a probe call through.
defaultBreakerCooldown = 60 * time.Second
)
// UpstreamConfig configures the HTTP threat-intel client. TokenEnv reads
// the process environment before each HTTP request. External environment
// changes require a daemon restart.
type UpstreamConfig struct {
URL string
Token string
TokenEnv string
CacheTTL time.Duration
Timeout time.Duration
}
// UpstreamSource queries a panel-side TI cache. The wire contract is
// documented in docs/upstream-threat-intel-contract.md.
//
// GET <URL>/lookup?ip=<ip>
// Authorization: Bearer <token> (omitted if no token resolved)
//
// 200 OK
// {"ip":"1.2.3.4","score":75,"source":"upstream","ttl_sec":900}
//
// Errors of any flavour (network, 4xx, 5xx, malformed JSON) propagate
// up - the aggregator treats them as "no signal" rather than fatal.
type UpstreamSource struct {
cfg UpstreamConfig
client *http.Client
mu sync.RWMutex
cache map[string]upstreamEntry
// maxCacheEntries caps the in-memory score cache. Exposed for tests
// that want to validate eviction without staging 10k entries.
maxCacheEntries int
// breakerMu guards the circuit-breaker state below.
breakerMu sync.Mutex
consecutiveErrors int
breakerOpenedAt time.Time
breakerProbe bool
breakerTrip int
breakerCooldown time.Duration
// Per-source counters back MetricsSnapshot. Package-level metrics use
// process-wide counters because reputation checks rebuild this source.
cacheHitsTotal atomic.Int64
cacheMissesTotal atomic.Int64
backendFailuresTotal atomic.Int64
}
type upstreamEntry struct {
score int
expires time.Time
}
// upstreamResponse mirrors the documented panel response shape.
type upstreamResponse struct {
IP string `json:"ip"`
Score int `json:"score"`
Source string `json:"source,omitempty"`
TTLSec int `json:"ttl_sec,omitempty"`
}
func NewUpstreamSource(cfg UpstreamConfig) *UpstreamSource {
if cfg.Timeout == 0 {
cfg.Timeout = 5 * time.Second
}
if cfg.CacheTTL == 0 {
cfg.CacheTTL = 15 * time.Minute
}
return &UpstreamSource{
cfg: cfg,
client: &http.Client{Timeout: cfg.Timeout},
cache: make(map[string]upstreamEntry),
maxCacheEntries: defaultMaxCacheEntries,
breakerTrip: defaultBreakerTrip,
breakerCooldown: defaultBreakerCooldown,
}
}
func (u *UpstreamSource) Name() string { return "upstream" }
// resolveToken reads TokenEnv (if set) at query time, falling back to the
// static token. External environment changes require a daemon restart.
func (u *UpstreamSource) resolveToken() string {
if u.cfg.TokenEnv != "" {
if v := os.Getenv(u.cfg.TokenEnv); v != "" {
return v
}
}
return u.cfg.Token
}
func (u *UpstreamSource) Score(ctx context.Context, ip string) (int, error) {
if v, ok := u.cacheGet(ip); ok {
u.cacheHitsTotal.Add(1)
upstreamMetrics.cacheHitsTotal.Add(1)
return v, nil
}
u.cacheMissesTotal.Add(1)
upstreamMetrics.cacheMissesTotal.Add(1)
if open, until := u.breakerOpen(); open {
if until.IsZero() {
return 0, fmt.Errorf("upstream breaker probe already running")
}
return 0, fmt.Errorf("upstream breaker open for %s", time.Until(until).Round(time.Second))
}
endpoint, err := u.lookupEndpoint(ip)
if err != nil {
return 0, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return 0, err
}
if tok := u.resolveToken(); tok != "" {
req.Header.Set("Authorization", "Bearer "+tok)
}
req.Header.Set("Accept", "application/json")
resp, err := u.client.Do(req)
if err != nil {
u.breakerObserve(false)
return 0, fmt.Errorf("upstream request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
u.breakerObserve(false)
fmt.Fprintf(os.Stderr, "upstream threat-intel: HTTP %d for %s\n", resp.StatusCode, ip)
return 0, fmt.Errorf("upstream HTTP %d", resp.StatusCode)
}
var body upstreamResponse
if err := json.NewDecoder(io.LimitReader(resp.Body, upstreamMaxResponseBytes)).Decode(&body); err != nil {
u.breakerObserve(false)
return 0, fmt.Errorf("upstream decode: %w", err)
}
if normalizeIP(body.IP) != normalizeIP(ip) {
u.breakerObserve(false)
return 0, fmt.Errorf("upstream response ip mismatch: got %q want %q", body.IP, ip)
}
if body.Score < 0 || body.Score > 100 {
u.breakerObserve(false)
return 0, fmt.Errorf("upstream score out of range: %d", body.Score)
}
u.breakerObserve(true)
// The response may shorten the cache lifetime, never extend it past the
// operator's cache_ttl: an unbounded ttl_sec pinned a score for as long
// as the daemon ran.
ttl := u.cfg.CacheTTL
if body.TTLSec > 0 {
if responseTTL := time.Duration(body.TTLSec) * time.Second; responseTTL < ttl {
ttl = responseTTL
}
}
u.cachePut(ip, body.Score, ttl)
return body.Score, nil
}
// breakerOpen reports whether the circuit breaker is currently open.
// When open and the cooldown has elapsed, the breaker transitions to a
// half-open state by clearing the timestamp so one probe call may pass.
func (u *UpstreamSource) breakerOpen() (bool, time.Time) {
u.breakerMu.Lock()
defer u.breakerMu.Unlock()
if u.breakerOpenedAt.IsZero() {
return false, time.Time{}
}
until := u.breakerOpenedAt.Add(u.breakerCooldown)
if time.Now().Before(until) {
return true, until
}
if u.breakerProbe {
return true, time.Time{}
}
u.breakerProbe = true
return false, time.Time{}
}
func (u *UpstreamSource) resetBreakerLocked() {
u.consecutiveErrors = 0
u.breakerProbe = false
u.breakerOpenedAt = time.Time{}
}
func (u *UpstreamSource) recordBreakerFailureLocked(now time.Time) {
u.breakerProbe = false
u.consecutiveErrors++
if u.breakerTrip > 0 && u.consecutiveErrors >= u.breakerTrip {
u.breakerOpenedAt = now
}
}
func (u *UpstreamSource) closeExpiredBreakerLocked(now time.Time) {
if !u.breakerOpenedAt.IsZero() && now.Sub(u.breakerOpenedAt) >= u.breakerCooldown && !u.breakerProbe {
u.breakerOpenedAt = time.Time{}
}
}
// breakerObserve records the outcome of an upstream call so the
// breaker can trip after enough consecutive failures or reset on a
// successful response.
func (u *UpstreamSource) breakerObserve(success bool) {
u.breakerMu.Lock()
defer u.breakerMu.Unlock()
if success {
u.resetBreakerLocked()
return
}
u.backendFailuresTotal.Add(1)
upstreamMetrics.backendFailuresTotal.Add(1)
now := time.Now()
u.closeExpiredBreakerLocked(now)
u.recordBreakerFailureLocked(now)
}
// MetricsSnapshot returns this source's cache-hit, cache-miss, and
// backend-failure counters. Safe to call from any goroutine.
func (u *UpstreamSource) MetricsSnapshot() (cacheHits, cacheMisses, backendFailures int64) {
return u.cacheHitsTotal.Load(), u.cacheMissesTotal.Load(), u.backendFailuresTotal.Load()
}
// BreakerOpen reports whether the circuit breaker is currently in the
// open state (refusing calls). Read-only.
func (u *UpstreamSource) BreakerOpen() bool {
u.breakerMu.Lock()
defer u.breakerMu.Unlock()
if u.breakerOpenedAt.IsZero() {
return false
}
return time.Now().Before(u.breakerOpenedAt.Add(u.breakerCooldown))
}
func (u *UpstreamSource) lookupEndpoint(ip string) (string, error) {
endpoint, err := url.Parse(u.cfg.URL)
if err != nil {
return "", fmt.Errorf("parsing upstream URL: %w", err)
}
if endpoint.Scheme != "http" && endpoint.Scheme != "https" {
return "", fmt.Errorf("upstream URL must use http or https")
}
if endpoint.Host == "" {
return "", fmt.Errorf("upstream URL must include host")
}
endpoint.Path = strings.TrimRight(endpoint.Path, "/") + "/lookup"
endpoint.Fragment = ""
q := endpoint.Query()
q.Set("ip", normalizeIP(ip))
endpoint.RawQuery = q.Encode()
return endpoint.String(), nil
}
func (u *UpstreamSource) cacheGet(ip string) (int, bool) {
u.mu.RLock()
defer u.mu.RUnlock()
e, ok := u.cache[ip]
if !ok || time.Now().After(e.expires) {
return 0, false
}
return e.score, true
}
func (u *UpstreamSource) cachePut(ip string, score int, ttl time.Duration) {
u.mu.Lock()
defer u.mu.Unlock()
u.cache[ip] = upstreamEntry{score: score, expires: time.Now().Add(ttl)}
u.evictLocked()
}
// evictLocked drops expired entries first, then evicts the oldest
// (smallest expires) until size is within maxCacheEntries. Caller
// must hold u.mu.
func (u *UpstreamSource) evictLocked() {
if u.maxCacheEntries <= 0 || len(u.cache) <= u.maxCacheEntries {
return
}
now := time.Now()
for k, e := range u.cache {
if now.After(e.expires) {
delete(u.cache, k)
}
}
for len(u.cache) > u.maxCacheEntries {
var oldestKey string
var oldestAt time.Time
first := true
for k, e := range u.cache {
if first || e.expires.Before(oldestAt) {
oldestKey = k
oldestAt = e.expires
first = false
}
}
delete(u.cache, oldestKey)
}
}
// cacheLen is exposed for tests to inspect cache size without exposing
// the underlying map.
func (u *UpstreamSource) cacheLen() int {
u.mu.RLock()
defer u.mu.RUnlock()
return len(u.cache)
}
package updatecheck
import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
)
const defaultGitHubReleasesURL = "https://api.github.com/repos/pidginhost/csm/releases/latest"
// fetchGitHubLatest returns the newest tagged release version, with
// any leading "v" stripped. Pre-releases are skipped via the standard
// /releases/latest endpoint, which already excludes them.
func fetchGitHubLatest(ctx context.Context, hc *http.Client, url string) (string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, url, nil)
if err != nil {
return "", err
}
req.Header.Set("Accept", "application/vnd.github+json")
req.Header.Set("User-Agent", "csm-update-check")
resp, err := hc.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
return "", fmt.Errorf("github releases: status %d", resp.StatusCode)
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 1<<20))
if err != nil {
return "", err
}
var payload struct {
TagName string `json:"tag_name"`
}
if err := json.Unmarshal(body, &payload); err != nil {
return "", err
}
tag := strings.TrimSpace(payload.TagName)
if tag == "" {
return "", fmt.Errorf("github releases: empty tag_name")
}
return strings.TrimPrefix(tag, "v"), nil
}
package updatecheck
import (
"context"
"errors"
"fmt"
"os/exec"
"reflect"
"runtime"
"sort"
"strings"
"time"
)
// AptProbe queries `apt-cache policy <pkg>` and returns the candidate
// version. Returns an error when apt-cache is missing, the package is
// unknown, or the candidate is "(none)".
func AptProbe(packageName string) PackageProbe {
return func(ctx context.Context) (string, error) {
ctx, cancel := context.WithTimeout(ctx, 30*time.Second)
defer cancel()
cmd := exec.CommandContext(ctx, "apt-cache", "policy", packageName) // #nosec G204 -- packageName is operator-controlled config, not attacker input
out, err := cmd.Output()
if err != nil {
return "", fmt.Errorf("apt-cache policy: %w", err)
}
return parseAptPolicy(string(out))
}
}
// DnfProbe queries `dnf --quiet repoquery --queryformat=%{version}
// <pkg>` and returns the highest version line. Returns an error when
// dnf is missing or returns no rows.
func DnfProbe(packageName string) PackageProbe {
return func(ctx context.Context) (string, error) {
ctx, cancel := context.WithTimeout(ctx, 60*time.Second)
defer cancel()
cmd := exec.CommandContext(ctx, "dnf", "--quiet", "repoquery", "--queryformat=%{version}\n", packageName) // #nosec G204 -- packageName is operator-controlled config, not attacker input
out, err := cmd.Output()
if err != nil {
return "", fmt.Errorf("dnf repoquery: %w", err)
}
return parseDnfRepoquery(string(out))
}
}
func parseAptPolicy(out string) (string, error) {
for _, line := range strings.Split(out, "\n") {
line = strings.TrimSpace(line)
if !strings.HasPrefix(line, "Candidate:") {
continue
}
v := strings.TrimSpace(strings.TrimPrefix(line, "Candidate:"))
if v == "" || v == "(none)" {
return "", errors.New("apt-cache policy: candidate (none)")
}
return aptStripEpochRevision(v), nil
}
return "", errors.New("apt-cache policy: no candidate line")
}
func parseDnfRepoquery(out string) (string, error) {
versions := []string{}
for _, line := range strings.Split(out, "\n") {
v := strings.TrimSpace(line)
if v == "" {
continue
}
versions = append(versions, v)
}
if len(versions) == 0 {
return "", errors.New("dnf repoquery: no versions")
}
sort.Slice(versions, func(i, j int) bool { return isNewer(versions[i], versions[j]) })
return versions[0], nil
}
// aptStripEpochRevision drops the optional "EPOCH:" prefix and
// "-DEBIAN_REVISION" suffix added by the apt versioning scheme so the
// returned string can be compared to a plain semver tag.
func aptStripEpochRevision(v string) string {
if i := strings.Index(v, ":"); i >= 0 {
v = v[i+1:]
}
if i := strings.Index(v, "-"); i >= 0 {
v = v[:i]
}
return v
}
// pkgSourceLabel best-effort labels a probe as "apt" or "dnf" by name.
// Unknown probes get "package".
func pkgSourceLabel(p PackageProbe) string {
if p == nil {
return "package"
}
name := runtime.FuncForPC(reflect.ValueOf(p).Pointer()).Name()
switch {
case strings.Contains(name, "AptProbe"):
return "apt"
case strings.Contains(name, "DnfProbe"):
return "dnf"
default:
return "package"
}
}
package updatecheck
import (
"strconv"
"strings"
)
// isNewer returns true when a is strictly greater than b under
// dot-separated numeric ordering. "dev" or empty current always
// loses (a real release is always newer than "dev"). Current-version
// strings produced by git describe, such as 3.0.0-12-gabcdef0, compare
// as newer than their base tag but older than the next tagged release.
func isNewer(a, b string) bool {
a = strings.TrimPrefix(strings.TrimSpace(a), "v")
b = strings.TrimPrefix(strings.TrimSpace(b), "v")
if a == "" {
return false
}
if b == "" || b == "dev" {
return true
}
if base, ok := gitDescribeBase(b); ok {
return isNewer(a, base)
}
ap := strings.Split(a, ".")
bp := strings.Split(b, ".")
for i := 0; i < len(ap) || i < len(bp); i++ {
var av, bv string
if i < len(ap) {
av = ap[i]
}
if i < len(bp) {
bv = bp[i]
}
ai, aErr := strconv.Atoi(av)
bi, bErr := strconv.Atoi(bv)
switch {
case aErr == nil && bErr == nil:
if ai != bi {
return ai > bi
}
case aErr == nil && bErr != nil:
return true
case aErr != nil && bErr == nil:
return false
default:
if av != bv {
return av > bv
}
}
}
return false
}
func gitDescribeBase(v string) (string, bool) {
parts := strings.Split(v, "-")
if len(parts) < 3 {
return "", false
}
if _, err := strconv.Atoi(parts[1]); err != nil {
return "", false
}
if !strings.HasPrefix(parts[2], "g") || !isNumericDotted(parts[0]) {
return "", false
}
return parts[0], true
}
func isNumericDotted(v string) bool {
parts := strings.Split(v, ".")
if len(parts) < 2 {
return false
}
for _, p := range parts {
if p == "" {
return false
}
if _, err := strconv.Atoi(p); err != nil {
return false
}
}
return true
}
// Package updatecheck polls upstream release channels and tells the
// daemon whether a newer CSM version is available so the Web UI can
// surface a banner. It never fetches binaries or modifies the running
// install. Operators upgrade through their normal channel
// (apt, dnf, install.sh, deploy pipeline).
//
// Two sources, tried in order:
//
// 1. GitHub Releases API ("https://api.github.com/repos/pidginhost/csm/releases/latest").
// 2. apt-cache policy or dnf repoquery against the OS package, used
// when the GitHub call fails (network blocked, rate-limited, etc.).
//
// The package never panics on a transient network or exec failure --
// it records the error in Info.Err and keeps the previous successful
// result in place so the banner does not flicker on a single bad poll.
package updatecheck
import (
"context"
"net/http"
"sync/atomic"
"time"
)
// Info is the cached result surfaced to /api/v1/status. Zero value
// means "no check has completed yet"; CheckedAt.IsZero() is the
// canonical signal.
type Info struct {
LatestVersion string `json:"latest_version,omitempty"`
Available bool `json:"available"`
Source string `json:"source,omitempty"` // "github" | "apt" | "dnf"
CheckedAt time.Time `json:"checked_at,omitempty"`
Err string `json:"err,omitempty"`
}
// PackageProbe queries the OS package manager for the highest
// available version of the configured package. Implementations
// must respect ctx for cancellation and timeout.
type PackageProbe func(ctx context.Context) (string, error)
// Options configures a Checker. Zero-valued fields take safe defaults.
type Options struct {
// CurrentVersion is the running daemon's version string ("dev" is
// treated as "always older than any tagged release").
CurrentVersion string
// Interval is how often the checker polls upstream. Clamped to a
// minimum of 1h to avoid hammering the GitHub API.
Interval time.Duration
// GitHubAPIURL overrides the default GitHub releases URL. Used by
// tests; production should leave this empty.
GitHubAPIURL string
// HTTPClient is the HTTP client used for the GitHub request. nil
// gets a sane default with a 15s timeout.
HTTPClient *http.Client
// PackageProbe is the apt/dnf fallback. nil disables fallback.
PackageProbe PackageProbe
// Now lets tests inject a clock. Defaults to time.Now.
Now func() time.Time
// LogErr receives non-fatal probe errors. Optional; nil silences.
LogErr func(source string, err error)
}
// Checker holds polling state. Safe for concurrent reads of Latest()
// while the goroutine started by Run is updating the cache.
type Checker struct {
opts Options
cache atomic.Pointer[Info]
}
// New builds a Checker. Validate the options here so Run can rely on
// invariants without re-checking each tick.
func New(opts Options) *Checker {
if opts.Interval <= 0 {
opts.Interval = 24 * time.Hour
}
if opts.Interval < time.Hour {
opts.Interval = time.Hour
}
if opts.GitHubAPIURL == "" {
opts.GitHubAPIURL = defaultGitHubReleasesURL
}
if opts.HTTPClient == nil {
opts.HTTPClient = &http.Client{Timeout: 15 * time.Second}
}
if opts.Now == nil {
opts.Now = time.Now
}
c := &Checker{opts: opts}
c.cache.Store(&Info{})
return c
}
// Latest returns the most recent successful poll plus any error from
// the most recent attempt. The returned value is safe to mutate.
func (c *Checker) Latest() Info {
if v := c.cache.Load(); v != nil {
return *v
}
return Info{}
}
// CheckOnce runs a single poll synchronously and returns the result.
// Run uses this internally; tests call it directly to avoid the ticker.
func (c *Checker) CheckOnce(ctx context.Context) Info {
now := c.opts.Now()
latest, err := fetchGitHubLatest(ctx, c.opts.HTTPClient, c.opts.GitHubAPIURL)
source := "github"
if err != nil {
if c.opts.LogErr != nil {
c.opts.LogErr("github", err)
}
if c.opts.PackageProbe != nil {
pkgVer, pkgErr := c.opts.PackageProbe(ctx)
if pkgErr == nil {
latest = pkgVer
source = pkgSourceLabel(c.opts.PackageProbe)
err = nil
} else {
if c.opts.LogErr != nil {
c.opts.LogErr("package", pkgErr)
}
err = pkgErr
}
}
}
info := Info{CheckedAt: now}
if err != nil {
// Preserve the last good LatestVersion so the banner does
// not flicker on a single bad poll.
prev := c.Latest()
info.LatestVersion = prev.LatestVersion
info.Available = prev.Available
info.Source = prev.Source
info.Err = err.Error()
} else {
info.LatestVersion = latest
info.Source = source
info.Available = isNewer(latest, c.opts.CurrentVersion)
}
c.cache.Store(&info)
return info
}
// Run polls on the configured interval until ctx is cancelled. It
// performs an initial check after a 5-minute warm-up so daemon
// startup is not blocked on outbound HTTP.
func (c *Checker) Run(ctx context.Context) {
select {
case <-ctx.Done():
return
case <-time.After(5 * time.Minute):
}
c.CheckOnce(ctx)
t := time.NewTicker(c.opts.Interval)
defer t.Stop()
for {
select {
case <-ctx.Done():
return
case <-t.C:
c.CheckOnce(ctx)
}
}
}
// Package verdict implements an HMAC-signed HTTP client for the
// auto_response.verdict_callback hook. CSM POSTs a Request to the panel
// URL before each automatic block; the panel's Response is advisory
// (block / allow / "" -> block default; tenant_id is logged). Errors are
// fail-open: the caller (firewall.Engine.BlockIP) proceeds with the
// default block on any callback failure.
package verdict
import (
"bytes"
"context"
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/hex"
"encoding/json"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strings"
"time"
)
const (
verdictMaxResponseBytes = 64 << 10 // 64 KB - verdict response is small JSON
// verdictMaxResponseSkew bounds how far the panel's reply timestamp
// may drift from CSM's clock. Anything older or further into the
// future is treated as a replayed or forged response.
verdictMaxResponseSkew = 5 * time.Minute
)
// Config configures the verdict callback client. HMACSecretEnv reads the
// process environment once per exchange. External environment changes
// require a daemon restart.
//
// RequireResponseSignature controls whether the panel must sign its
// response body with the same HMAC scheme used on the request
// (X-CSM-Signature header) and echo the request nonce + timestamp.
// Default is true: when a secret is configured, CSM rejects unsigned
// or forged responses to prevent an on-path attacker from silently
// downgrading a block to "allow". Set the pointer to a false value
// only during phpanel-side rollouts that have not yet implemented
// response signing. Even on the opt-out path, replay protection is
// best-effort enforced: if the panel does echo nonce or timestamp,
// they must match; a panel that echoes neither still works (legacy
// shape). When no HMAC secret is configured at all, signature and
// replay checks are skipped because there is no key to verify against.
//
// AllowUnsigned is the runtime form of
// auto_response.verdict_callback.allow_unsigned. By default, an unsigned
// "allow" response is rejected because it would disable the block. Set true
// only for configs that explicitly opted in to an unsigned rollout.
type Config struct {
URL string
HMACSecret string
HMACSecretEnv string
RequireResponseSignature *bool
AllowUnsigned bool
Timeout time.Duration
}
// requireResponseSig returns the effective response-signature requirement,
// defaulting to true (secure by default) when the operator did not set
// an explicit value.
func (c Config) requireResponseSig() bool {
if c.RequireResponseSignature == nil {
return true
}
return *c.RequireResponseSignature
}
// Request is what CSM asks the panel about. Ask sets Nonce and Timestamp
// for every exchange. They bind each request to its own reply so an
// attacker cannot replay an old "allow" verdict.
type Request struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Severity string `json:"severity,omitempty"`
Source string `json:"source,omitempty"` // "auto_response" | "manual"
Nonce string `json:"nonce,omitempty"`
Timestamp int64 `json:"timestamp,omitempty"`
}
// Response is what the panel may answer.
//
// verdict = "block" - CSM proceeds with its default action.
// verdict = "allow" - CSM logs the verdict but does NOT block.
// verdict = "" or missing - equivalent to "block" (default).
// tenant_id = optional attribution string CSM logs alongside the decision.
// nonce = MUST equal the Request.Nonce when present. Required when
// response signing is in effect; defeats replay of captured replies.
// timestamp = unix seconds the panel produced the reply. MUST be within
// verdictMaxResponseSkew of CSM's clock when present.
type Response struct {
Verdict string `json:"verdict,omitempty"`
TenantID string `json:"tenant_id,omitempty"`
Note string `json:"note,omitempty"`
Nonce string `json:"nonce,omitempty"`
Timestamp int64 `json:"timestamp,omitempty"`
}
// Client posts each block decision to the configured URL and reads the
// (advisory) response. Timeouts and 5xx are returned as errors - the
// caller decides whether to fail open (allow CSM to proceed) or closed.
type Client struct {
cfg Config
client *http.Client
}
func New(cfg Config) *Client {
if cfg.Timeout == 0 {
cfg.Timeout = 2 * time.Second
}
return &Client{cfg: cfg, client: &http.Client{Timeout: cfg.Timeout}}
}
// resolveSecret reads HMACSecretEnv at call time, falling back to the
// static secret. External environment changes require a daemon restart.
func (c *Client) resolveSecret() string {
if c.cfg.HMACSecretEnv != "" {
if v := os.Getenv(c.cfg.HMACSecretEnv); v != "" {
return v
}
}
return c.cfg.HMACSecret
}
// verifyResponseSignature checks that header carries a well-formed
// X-CSM-Signature header (sha256=<hex>) over body, computed with secret.
// Returns a descriptive error on missing, malformed, or mismatched values
// using a constant-time comparison.
func verifyResponseSignature(secret string, body []byte, header string) error {
if header == "" {
return fmt.Errorf("verdict callback response missing signature header (X-CSM-Signature)")
}
const prefix = "sha256="
if !strings.HasPrefix(header, prefix) {
return fmt.Errorf("verdict callback response signature has unsupported algorithm")
}
gotHex := strings.TrimPrefix(header, prefix)
got, err := hex.DecodeString(gotHex)
if err != nil {
return fmt.Errorf("verdict callback response signature is not valid hex")
}
mac := hmac.New(sha256.New, []byte(secret))
mac.Write(body)
want := mac.Sum(nil)
if !hmac.Equal(got, want) {
return fmt.Errorf("verdict callback response signature mismatch")
}
return nil
}
// Ask POSTs the request, returns the panel's response or an error.
func (c *Client) Ask(ctx context.Context, req Request) (Response, error) {
// Defense-in-depth URL check (config validation already ran at load,
// this re-check defends against misconfiguration via cfg corruption).
rawURL := strings.TrimSpace(c.cfg.URL)
if rawURL == "" {
return Response{}, fmt.Errorf("verdict callback URL not configured")
}
parsed, err := url.Parse(rawURL)
if err != nil {
return Response{}, fmt.Errorf("verdict callback URL parse: %w", err)
}
if parsed.Scheme != "http" && parsed.Scheme != "https" {
return Response{}, fmt.Errorf("verdict callback URL must use http or https")
}
if parsed.Host == "" {
return Response{}, fmt.Errorf("verdict callback URL must include host")
}
nonce, err := newNonce()
if err != nil {
return Response{}, fmt.Errorf("verdict callback nonce generation: %w", err)
}
req.Nonce = nonce
req.Timestamp = time.Now().Unix()
body, err := json.Marshal(req)
if err != nil {
return Response{}, err
}
secret := c.resolveSecret()
r, err := http.NewRequestWithContext(ctx, http.MethodPost, rawURL, bytes.NewReader(body))
if err != nil {
return Response{}, err
}
r.Header.Set("Content-Type", "application/json")
r.Header.Set("Accept", "application/json")
if secret != "" {
mac := hmac.New(sha256.New, []byte(secret))
mac.Write(body)
r.Header.Set("X-CSM-Signature", "sha256="+hex.EncodeToString(mac.Sum(nil)))
}
resp, err := c.client.Do(r)
if err != nil {
return Response{}, fmt.Errorf("verdict callback: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
// Return the IP in the error so the caller logs it through the
// structured audit path instead of leaking it to raw stderr.
return Response{}, fmt.Errorf("verdict callback HTTP %d for %s", resp.StatusCode, req.IP)
}
data, err := io.ReadAll(io.LimitReader(resp.Body, verdictMaxResponseBytes+1))
if err != nil {
return Response{}, fmt.Errorf("verdict callback read: %w", err)
}
if int64(len(data)) > verdictMaxResponseBytes {
return Response{}, fmt.Errorf("verdict callback response exceeds %d bytes", verdictMaxResponseBytes)
}
if secret != "" && c.cfg.requireResponseSig() {
// Verify the panel signed its response with the same secret used
// on the request. Without this check a network attacker could
// downgrade block to allow on every call. Rejecting the reply keeps
// the engine on its default block path.
if err := verifyResponseSignature(secret, data, resp.Header.Get("X-CSM-Signature")); err != nil {
return Response{}, err
}
}
strictReplay := secret != "" && c.cfg.requireResponseSig()
if strings.TrimSpace(string(data)) == "" {
return Response{}, nil
}
var out Response
dec := json.NewDecoder(bytes.NewReader(data))
if err := dec.Decode(&out); err != nil {
return Response{}, fmt.Errorf("verdict callback decode: %w", err)
}
var trailing struct{}
if err := dec.Decode(&trailing); err != io.EOF {
return Response{}, fmt.Errorf("verdict callback decode: trailing JSON")
}
// Validate response shape. Unknown verdict strings are rejected
// defensively rather than silently treated as "block".
if out.Verdict != "" && out.Verdict != "block" && out.Verdict != "allow" {
return Response{}, fmt.Errorf("verdict callback returned unknown verdict %q", out.Verdict)
}
// Fail closed on an unsigned allow. Without a secret there is no replay or
// signature protection (the checks below are gated on secret != ""), so an
// on-path attacker could return "allow" on every call and silently disable
// auto-blocking. Returning an error keeps the engine on its default block
// path unless the config explicitly opted in to unsigned callback verdicts.
if secret == "" && out.Verdict == "allow" && !c.cfg.AllowUnsigned {
return Response{}, fmt.Errorf("verdict callback: refusing unsigned allow (no HMAC secret configured)")
}
// Replay protection runs whenever a secret is configured. Strict
// mode (response signing required) demands the panel echo nonce and
// timestamp. Best-effort mode (signing opt-out) still enforces what
// the panel did echo, so a captured stale reply with a wrong nonce
// or a long-expired timestamp is rejected even when the operator
// has not yet enabled response signing on the panel side.
if secret != "" {
if out.Nonce != "" {
if subtle.ConstantTimeCompare([]byte(out.Nonce), []byte(req.Nonce)) != 1 {
return Response{}, fmt.Errorf("verdict callback response nonce mismatch")
}
} else if strictReplay {
return Response{}, fmt.Errorf("verdict callback response missing nonce")
}
if out.Timestamp != 0 {
drift := time.Since(time.Unix(out.Timestamp, 0))
if drift < 0 {
drift = -drift
}
if drift > verdictMaxResponseSkew {
return Response{}, fmt.Errorf("verdict callback response timestamp drift %s exceeds %s", drift, verdictMaxResponseSkew)
}
} else if strictReplay {
return Response{}, fmt.Errorf("verdict callback response missing timestamp")
}
// An "allow" is the only verdict an on-path attacker gains from
// forging. In best-effort mode (signing not required) the replay
// checks above only fire when the panel echoed a nonce or timestamp,
// so an attacker can strip both to slip an unbound allow through.
// Require an allow to carry at least one replay binding; otherwise
// treat it like the no-secret case and refuse, keeping the engine on
// its default block path.
if out.Verdict == "allow" && out.Nonce == "" && out.Timestamp == 0 {
return Response{}, fmt.Errorf("verdict callback: refusing allow with no replay binding (nonce/timestamp)")
}
}
return out, nil
}
// newNonce returns a fresh 128-bit hex nonce. crypto/rand is the only
// acceptable source: math/rand would let an attacker who observes one
// nonce predict the next and craft a replay reply in advance.
func newNonce() (string, error) {
var buf [16]byte
if _, err := io.ReadFull(rand.Reader, buf[:]); err != nil {
return "", err
}
return hex.EncodeToString(buf[:]), nil
}
package webui
import (
"net/http"
"path/filepath"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
)
func (s *Server) handleAccount(w http.ResponseWriter, r *http.Request) {
name := r.URL.Query().Get("name")
if err := validateAccountName(name); err != nil {
http.Redirect(w, r, "/findings", http.StatusFound)
return
}
s.renderTemplate(w, r, "account.html", map[string]string{
"Hostname": s.cfg.Hostname,
"AccountName": name,
"HomeDir": checks.AccountHomeDirIn(s.accountRoots(), name),
})
}
// accountPathPrefixes returns "<root>/<name>/" for every account root, the
// prefixes a path inside the account's home starts with.
func (s *Server) accountPathPrefixes(name string) []string {
roots := s.accountRoots()
out := make([]string, 0, len(roots))
for _, root := range roots {
out = append(out, filepath.Join(root, name)+"/")
// Findings and quarantine metadata can carry the resolved path even
// after the file has been moved away. Resolve the root, not the file.
if resolved, err := filepath.EvalSymlinks(root); err == nil && resolved != filepath.Clean(root) {
out = append(out, filepath.Join(resolved, name)+"/")
}
}
return out
}
func pathHasAnyPrefix(path string, prefixes []string) bool {
for _, prefix := range prefixes {
if strings.HasPrefix(path, prefix) {
return true
}
}
return false
}
func containsAny(s string, subs []string) bool {
for _, sub := range subs {
if strings.Contains(s, sub) {
return true
}
}
return false
}
func accountFindingMatches(f alert.Finding, name string, prefixes []string) bool {
for _, owner := range []string{f.TenantID, f.CPUser} {
if owner = strings.TrimSpace(owner); owner != "" {
return owner == name
}
}
return containsAny(f.Message, prefixes) || containsAny(f.Details, prefixes) || pathHasAnyPrefix(f.FilePath, prefixes)
}
func (s *Server) apiAccountDetail(w http.ResponseWriter, r *http.Request) {
name := r.URL.Query().Get("name")
if err := validateAccountName(name); err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
homePrefixes := s.accountPathPrefixes(name)
// Current findings for this account
type findingView struct {
Severity string `json:"severity"`
Check string `json:"check"`
Message string `json:"message"`
HasFix bool `json:"has_fix"`
}
var accountFindings []findingView
latest := s.store.LatestFindings()
for _, f := range latest {
if !operatorFacingCheck(f.Check) {
continue
}
if accountFindingMatches(f, name, homePrefixes) {
accountFindings = append(accountFindings, findingView{
Severity: f.Severity.String(),
Check: f.Check,
Message: f.Message,
HasFix: checks.HasFix(f.Check),
})
}
}
// Quarantined files for this account
type qEntry struct {
ID string `json:"id"`
OriginalPath string `json:"original_path"`
Size int64 `json:"size"`
Reason string `json:"reason"`
}
var quarantined []qEntry
rootMetas := listMetaFiles(quarantineDir)
preCleanMetas := listMetaFiles(filepath.Join(quarantineDir, "pre_clean"))
metas := rootMetas
metas = append(metas, preCleanMetas...)
for _, metaPath := range metas {
meta, err := readQuarantineMeta(metaPath)
if err != nil {
continue
}
if pathHasAnyPrefix(meta.OriginalPath, homePrefixes) {
id := strings.TrimSuffix(filepath.Base(metaPath), ".meta")
quarantined = append(quarantined, qEntry{
ID: id, OriginalPath: meta.OriginalPath, Size: meta.Size, Reason: meta.Reason,
})
}
}
// Recent history for this account (last 100 matching entries)
allHistory, _ := s.store.ReadHistory(2000, 0)
type histEntry struct {
Severity string `json:"severity"`
Check string `json:"check"`
Message string `json:"message"`
Timestamp time.Time `json:"timestamp"`
}
var history []histEntry
for _, f := range allHistory {
if len(history) >= 100 {
break
}
if accountFindingMatches(f, name, homePrefixes) {
history = append(history, histEntry{
Severity: f.Severity.String(), Check: f.Check, Message: f.Message,
Timestamp: f.Timestamp.UTC(),
})
}
}
writeJSON(w, map[string]interface{}{
"account": name,
"findings": accountFindings,
"quarantined": quarantined,
"history": history,
"whm_url": "https://" + s.cfg.Hostname + ":2087/scripts/domainsdata?user=" + name,
})
}
package webui
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"fmt"
"log"
"net"
"net/http"
"net/netip"
"os"
"os/exec"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"time"
"unicode"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
var reIPReputation = regexp.MustCompile(`Known malicious IP accessing server: (\S+) \((.+)\)`)
// apiStatus returns the daemon's full health snapshot as JSON. Backward
// compatible with prior callers: every field they consumed (hostname,
// uptime, started_at, rules_loaded, scan_running, last_scan_time) is
// still present, with new fields added alongside.
func (s *Server) apiStatus(w http.ResponseWriter, _ *http.Request) {
provider := s.provider
s.scanMu.Lock()
scanning := s.scanRunning
s.scanMu.Unlock()
if s.scanInProgress != nil && s.scanInProgress() {
scanning = true
}
if provider == nil {
// No daemon-side provider installed (test harness). Fall back to
// the legacy minimal payload so existing UI code keeps working.
resp := map[string]interface{}{
"hostname": s.cfg.Hostname,
"uptime_seconds": int64(time.Since(s.startTime).Seconds()),
"started_at": s.startTime.UTC(),
"started_at_token": daemonStartToken(s.startTime),
"rules_loaded": s.signatureCount(),
"scan_running": scanning,
"status": "down",
}
if s.store != nil {
if last := s.store.LatestScanTime(); !last.IsZero() {
resp["last_scan_time"] = last.UTC()
}
}
writeJSON(w, resp)
return
}
snap := health.Build(provider, s.version, health.Capabilities())
resp := map[string]interface{}{
"hostname": snap.Hostname,
"version": snap.Version,
"uptime_seconds": snap.UptimeSec,
"started_at": snap.StartedAt.UTC(),
"started_at_token": daemonStartToken(snap.StartedAt),
"rules_loaded": s.signatureCount(),
"scan_running": scanning,
"blocklist_size": snap.BlocklistSize,
"incidents_open": snap.IncidentsOpen,
"bpf_enforcement_active": snap.BPFEnforcementActive,
"history_count": snap.HistoryCount,
"severities": snap.Severities,
"watchers": snap.Watchers,
"store_healthy": snap.StoreHealthy,
"store_size_mb": snap.StoreSizeMB,
"config_hash": snap.ConfigHash,
"binary_hash": snap.BinaryHash,
"capabilities": snap.Capabilities,
"dry_run_blocks": snap.DryRunBlocks,
"automation": snap.Automation,
"mode": snap.Mode,
"status": snap.OverallStatus(),
}
// latest_scan mirrors the health.Snapshot JSON tag and is the canonical
// name; last_scan_time is the legacy key kept for older clients (the
// cPHulk dashboard). A time that is not set is left out.
if !snap.LatestScan.IsZero() {
resp["latest_scan"] = snap.LatestScan.UTC()
resp["last_scan_time"] = snap.LatestScan.UTC()
}
if !snap.BaselineAt.IsZero() {
resp["baseline_at"] = snap.BaselineAt.UTC()
}
// security_posture is the threat-aware badge signal, distinct from
// "status" (which stays operational: is the daemon alive and attached).
// A daemon can be operationally "ok" while sitting on open critical
// incidents, so the dashboard pill must fold in incident severity rather
// than report "Healthy" next to thousands of criticals.
openBySev := map[string]int{}
if s.incidentCorrelator != nil {
openBySev = s.incidentCorrelator.OpenCountsBySeverity()
}
resp["incidents_open_by_severity"] = openBySev
resp["security_posture"] = securityPosture(operationalProblems(s.signatureCount(), snap), openBySev["critical"], openBySev["high"])
if !snap.Update.CheckedAt.IsZero() {
resp["update"] = snap.Update
}
// Present only after the daemon has merged an active set; absence means
// "not observed yet", not "clean".
if snap.CorrelationAttribution != nil {
resp["correlation_attribution"] = snap.CorrelationAttribution
}
if len(snap.Queues) != 0 {
resp["queues"] = snap.Queues
}
if len(snap.WordPressVerification) != 0 {
resp["wordpress_verification"] = snap.WordPressVerification
}
writeJSON(w, resp)
}
// operationalProblems counts daemon faults that should prevent a healthy
// posture even when there are no active high-severity incidents.
func operationalProblems(sigCount int, snap health.Snapshot) int {
problems := 0
if sigCount == 0 {
problems++
}
if !snap.StoreHealthy {
problems++
}
if !snap.AllWatchersAttached() {
problems++
}
for _, q := range snap.Queues {
if q.Status == "degraded" && !q.Advisory {
problems++
break
}
}
return problems
}
// securityPosture collapses operational faults and active-incident severity
// into the dashboard badge tier:
//
// - "critical" if any active critical incident exists, or the daemon has
// multiple operational problems.
// - "warning" if any active high incident exists, or exactly one operational
// problem.
// - "healthy" otherwise.
//
// Incident severity takes precedence over operational faults so a fully
// attached daemon still reads "critical" while critical incidents are active.
func securityPosture(opProblems, openCritical, openHigh int) string {
if openCritical > 0 || opProblems >= 2 {
return "critical"
}
if openHigh > 0 || opProblems == 1 {
return "warning"
}
return "healthy"
}
// apiCapabilities returns the static feature-flag list for this build.
func (s *Server) apiCapabilities(w http.ResponseWriter, _ *http.Request) {
writeJSON(w, map[string]interface{}{
"capabilities": health.Capabilities(),
"version": s.version,
})
}
// apiFindings returns current scan results - "what's wrong right now."
func (s *Server) apiFindings(w http.ResponseWriter, _ *http.Request) {
latest := s.store.LatestFindings()
type entryView struct {
Severity string `json:"severity"`
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details,omitempty"`
Time time.Time `json:"time"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
HasFix bool `json:"has_fix"`
}
suppressions := s.store.LoadSuppressions()
var result []entryView
for _, f := range latest {
if !operatorFacingCheck(f.Check) {
continue
}
// Skip suppressed findings
if s.store.IsSuppressed(f, suppressions) {
continue
}
firstSeen := f.Timestamp
lastSeen := f.Timestamp
if entry, ok := s.store.EntryForKey(f.Key()); ok {
firstSeen = entry.FirstSeen
lastSeen = entry.LastSeen
}
result = append(result, entryView{
Severity: f.Severity.String(),
Check: f.Check,
Message: f.Message,
Details: f.Details,
Time: f.Timestamp.UTC(),
FirstSeen: firstSeen.UTC(),
LastSeen: lastSeen.UTC(),
HasFix: checks.HasFix(f.Check),
})
}
writeAll(w, result)
}
// enrichedFinding is the JSON response type for the enriched findings endpoint.
type enrichedFinding struct {
Key string `json:"key"`
Severity string `json:"severity"`
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details,omitempty"`
FilePath string `json:"file_path,omitempty"`
Account string `json:"account,omitempty"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
HasFix bool `json:"has_fix"`
HasVerify bool `json:"has_verify"`
FixDesc string `json:"fix_desc,omitempty"`
ContentSHA256 string `json:"content_sha256,omitempty"`
// BlockIP is the attacker address an operator may block from this
// finding; empty for checks that do not report one.
BlockIP string `json:"block_ip,omitempty"`
}
// dedupIPReputation groups ip_reputation findings by IP, merging sources and
// promoting to the highest severity. Non-ip_reputation findings pass through unchanged.
func dedupIPReputation(items []enrichedFinding) []enrichedFinding {
type ipGroup struct {
entry enrichedFinding
sources []string
}
ipGroups := make(map[string]*ipGroup)
var ipOrder []string
var result []enrichedFinding
for _, item := range items {
if item.Check != "ip_reputation" {
result = append(result, item)
continue
}
m := reIPReputation.FindStringSubmatch(item.Message)
if m == nil {
result = append(result, item)
continue
}
ip, source := m[1], m[2]
if g, ok := ipGroups[ip]; ok {
g.sources = append(g.sources, source)
if item.FirstSeen.Before(g.entry.FirstSeen) {
g.entry.FirstSeen = item.FirstSeen
}
if item.LastSeen.After(g.entry.LastSeen) {
g.entry.LastSeen = item.LastSeen
}
if severityRank(item.Severity) > severityRank(g.entry.Severity) {
g.entry.Severity = item.Severity
}
} else {
ipGroups[ip] = &ipGroup{
entry: item,
sources: []string{source},
}
ipOrder = append(ipOrder, ip)
}
}
for _, ip := range ipOrder {
g := ipGroups[ip]
sort.Strings(g.sources)
g.entry.Message = fmt.Sprintf("Known malicious IP accessing server: %s (%s)", ip, strings.Join(g.sources, ", "))
result = append(result, g.entry)
}
return result
}
// apiFindingsEnriched returns findings with IP dedup, account extraction, and severity counts.
func (s *Server) apiFindingsEnriched(w http.ResponseWriter, r *http.Request) {
latest := s.store.LatestFindings()
suppressions := s.store.LoadSuppressions()
items := make([]enrichedFinding, 0)
for _, f := range latest {
if !operatorFacingCheck(f.Check) {
continue
}
if s.store.IsSuppressed(f, suppressions) {
continue
}
firstSeen := f.Timestamp
lastSeen := f.Timestamp
if entry, ok := s.store.EntryForKey(f.Key()); ok {
firstSeen = entry.FirstSeen
lastSeen = entry.LastSeen
}
items = append(items, enrichedFinding{
Key: f.Key(),
Severity: severityLabel(f.Severity),
Check: f.Check,
Message: f.Message,
Details: f.Details,
FilePath: f.FilePath,
Account: extractAccountFromFinding(f),
FirstSeen: firstSeen.UTC(),
LastSeen: lastSeen.UTC(),
HasFix: checks.HasFix(f.Check),
HasVerify: checks.CanVerify(f.Check),
FixDesc: checks.FixDescription(f.Check, f.Message, f.FilePath),
ContentSHA256: f.ContentSHA256,
BlockIP: checks.ManualBlockIP(f),
})
}
items = dedupIPReputation(items)
version := enrichedFindingsVersion(items)
if r.URL.Query().Get("fields") == "version" {
writeJSON(w, map[string]interface{}{"version": version, "total": len(items)})
return
}
var critCount, highCount, warnCount int
for _, item := range items {
switch item.Severity {
case "CRITICAL":
critCount++
case "HIGH":
highCount++
default:
warnCount++
}
}
checkTypeSet := make(map[string]bool)
accountSet := make(map[string]bool)
for _, item := range items {
checkTypeSet[item.Check] = true
if item.Account != "" {
accountSet[item.Account] = true
}
}
checkTypes := make([]string, 0, len(checkTypeSet))
for ct := range checkTypeSet {
checkTypes = append(checkTypes, ct)
}
sort.Strings(checkTypes)
accounts := make([]string, 0, len(accountSet))
for a := range accountSet {
accounts = append(accounts, a)
}
sort.Strings(accounts)
extra := map[string]interface{}{
"check_types": checkTypes,
"accounts": accounts,
"critical_count": critCount,
"high_count": highCount,
"warning_count": warnCount,
"version": version,
}
if limit := queryInt(r, "limit", 0); limit > 0 {
sortEnrichedBySeverity(items)
writeCapped(w, items, len(items), limit, extra)
return
}
extra["total"] = len(items)
writeItems(w, items, extra)
}
// sortEnrichedBySeverity orders findings most severe first, newest first
// within a severity, so a limited list keeps the ones that matter.
func sortEnrichedBySeverity(items []enrichedFinding) {
rank := map[string]int{"CRITICAL": 3, "HIGH": 2}
sort.SliceStable(items, func(i, j int) bool {
if ri, rj := rank[items[i].Severity], rank[items[j].Severity]; ri != rj {
return ri > rj
}
return items[i].LastSeen.After(items[j].LastSeen)
})
}
// enrichedFindingsVersion changes when the listed findings or their
// severities do. A client that polls only to learn whether the list changed
// asks for ?fields=version and compares. ip_reputation rows are identified by
// their message, which carries the merged sources.
func enrichedFindingsVersion(items []enrichedFinding) string {
ids := make([]string, 0, len(items))
for _, f := range items {
id := f.Key
if f.Check == "ip_reputation" {
id = f.Check + ":" + f.Message
}
ids = append(ids, id+"|"+f.Severity)
}
sort.Strings(ids)
h := sha256.New()
for _, id := range ids {
h.Write([]byte(id))
h.Write([]byte{0})
}
return hex.EncodeToString(h.Sum(nil))[:16]
}
// apiHistory returns paginated finding history.
// Supports optional filtering via "from", "to" (YYYY-MM-DD or RFC 3339), and "severity" (a label or 0/1/2) query params.
func (s *Server) apiHistory(w http.ResponseWriter, r *http.Request) {
limit := queryInt(r, "limit", 50)
if limit > 5000 {
limit = 5000
}
offset := queryInt(r, "offset", 0)
q, ok := parseHistoryQuery(w, r)
if !ok {
return
}
findings, total := s.readHistoryPage(q, limit, offset)
writeItems(w, withAccountIP(findings), map[string]interface{}{
"total": total,
"limit": limit,
"offset": offset,
"truncated": historyPageTruncated(total, offset, len(findings)),
})
}
// historyQuery is the filter set /api/v1/history and its CSV export share.
type historyQuery struct {
from, to, search string
severity int // -1 for any
checks map[string]bool
}
func (q historyQuery) filtered() bool {
return q.from != "" || q.to != "" || q.severity >= 0 || q.search != "" || q.checks != nil
}
// parseHistoryQuery reads the history filters. An unreadable date is a 400,
// written here, and ok is false.
func parseHistoryQuery(w http.ResponseWriter, r *http.Request) (historyQuery, bool) {
v := r.URL.Query()
if _, _, ok := historyRangeQuery(w, v, time.Time{}, time.Time{}); !ok {
return historyQuery{}, false
}
q := historyQuery{from: v.Get("from"), to: v.Get("to"), search: v.Get("search"), severity: -1}
if sev, ok := parseSeverity(v.Get("severity")); ok {
q.severity = int(sev)
}
if checksStr := v.Get("checks"); checksStr != "" {
q.checks = make(map[string]bool)
for _, c := range strings.Split(checksStr, ",") {
if c = strings.TrimSpace(c); c != "" {
q.checks[c] = true
}
}
}
return q, true
}
func (s *Server) readHistoryPage(q historyQuery, limit, offset int) ([]alert.Finding, int) {
if !q.filtered() {
return s.store.ReadHistory(limit, offset)
}
return s.store.ReadHistoryFilteredWithChecks(limit, offset, q.from, q.to, q.severity, q.search, q.checks)
}
// historyPageTruncated reports whether matches exist past the returned page.
// total counts every match, so a client can page on through offset.
func historyPageTruncated(total, offset, returned int) bool {
return total > offset+returned
}
// apiFinding is a stored finding as the API sends it: its severity, and the
// one it was demoted from, are labels. The embedded Finding promotes its
// other JSON fields; these shallower fields take the severity keys.
type apiFinding struct {
alert.Finding
Severity string `json:"severity"`
DemotedFrom string `json:"demoted_from,omitempty"`
}
func toAPIFinding(f alert.Finding) apiFinding {
a := apiFinding{Finding: f, Severity: f.Severity.String()}
// Demotion only lowers a severity, so WARNING, the zero level, is never
// the one a finding was demoted from.
if f.DemotedFrom > alert.Warning {
a.DemotedFrom = f.DemotedFrom.String()
}
return a
}
func toAPIFindings(findings []alert.Finding) []apiFinding {
out := make([]apiFinding, len(findings))
for i, f := range findings {
out[i] = toAPIFinding(f)
}
return out
}
// historyFinding is an apiFinding with the normalized account and remote IP,
// so clients render structured fields instead of regex-scraping the
// human-readable message. It embeds the exported Finding, not apiFinding:
// apiValue cannot reach into an unexported embedded struct.
type historyFinding struct {
alert.Finding
Severity string `json:"severity"`
DemotedFrom string `json:"demoted_from,omitempty"`
Account string `json:"account,omitempty"`
IP string `json:"ip,omitempty"`
}
func withAccountIP(findings []alert.Finding) []historyFinding {
out := make([]historyFinding, len(findings))
for i, f := range findings {
a := toAPIFinding(f)
out[i] = historyFinding{Finding: f, Severity: a.Severity, DemotedFrom: a.DemotedFrom,
Account: findingAccount(f), IP: findingIP(f)}
}
return out
}
// parseSeverity reads a severity filter: a label in any case, or the 0/1/2
// level older callers send. ok is false for anything else.
func parseSeverity(text string) (alert.Severity, bool) {
switch strings.ToUpper(strings.TrimSpace(text)) {
case "WARNING", "0":
return alert.Warning, true
case "HIGH", "1":
return alert.High, true
case "CRITICAL", "2":
return alert.Critical, true
}
return 0, false
}
// findingAccount returns the account/mailbox attribution for a finding,
// preferring structured fields over legacy text extraction so the value
// survives message wording changes.
func findingAccount(f alert.Finding) string {
for _, account := range []string{f.Mailbox, f.TenantID, f.CPUser} {
if account = strings.TrimSpace(account); account != "" {
return account
}
}
if acct := legacyEmailAccount(f); acct != "" {
return acct
}
if acct := extractAccountFromFinding(f); acct != "" {
return acct
}
return strings.TrimSpace(f.Domain)
}
func findingIP(f alert.Finding) string {
if ip := strings.TrimSpace(f.SourceIP); ip != "" {
return ip
}
if !isEmailHistoryCheck(f.Check) {
return ""
}
for _, s := range []string{f.Message, f.Details} {
if ip := firstIPToken(s); ip != "" {
return ip
}
}
return ""
}
func legacyEmailAccount(f alert.Finding) string {
if !isEmailHistoryCheck(f.Check) {
return ""
}
for _, source := range []string{f.Message, f.Details} {
for _, prefix := range []string{"for ", "Account ", "account ", "Sender "} {
if token := tokenAfter(source, prefix); strings.Contains(token, "@") {
return token
}
}
if token := tokenAfter(source, "set_id="); token != "" {
return token
}
for _, prefix := range []string{"Domain ", "Domain: ", "domain ", "domain: "} {
if token := tokenAfter(source, prefix); token != "" {
return token
}
}
}
return ""
}
func isEmailHistoryCheck(check string) bool {
return emailKindForCheck(check) != "" ||
strings.HasPrefix(check, "email_") ||
strings.HasPrefix(check, "mail_") ||
strings.HasPrefix(check, "smtp_")
}
func tokenAfter(s, prefix string) string {
idx := strings.Index(s, prefix)
if idx < 0 {
return ""
}
rest := s[idx+len(prefix):]
if end := strings.IndexAny(rest, " \n\t,"); end >= 0 {
rest = rest[:end]
}
return cleanHistoryToken(rest)
}
func cleanHistoryToken(s string) string {
s = strings.TrimSpace(s)
s = strings.Trim(s, `"'<>[]()`)
s = strings.TrimRight(s, ".,;:")
return s
}
func firstIPToken(s string) string {
for _, raw := range strings.FieldsFunc(s, func(r rune) bool {
if unicode.IsSpace(r) {
return true
}
switch r {
case ',', ';', '(', ')':
return true
default:
return false
}
}) {
if ip := normalizeHistoryIPToken(raw); ip != "" {
return ip
}
}
return ""
}
func normalizeHistoryIPToken(raw string) string {
raw = strings.Trim(raw, `"'<>`)
if ip := parseHistoryIP(raw); ip != "" {
return ip
}
if eq := strings.LastIndexByte(raw, '='); eq >= 0 && eq < len(raw)-1 {
if ip := normalizeHistoryIPToken(raw[eq+1:]); ip != "" {
return ip
}
}
candidate := strings.TrimRight(raw, ".,;:")
if host, _, err := net.SplitHostPort(candidate); err == nil {
if ip := parseHistoryIP(host); ip != "" {
return ip
}
}
if strings.HasPrefix(candidate, "[") {
if end := strings.IndexByte(candidate, ']'); end > 1 {
if ip := parseHistoryIP(candidate[1:end]); ip != "" {
return ip
}
}
}
return parseHistoryIP(candidate)
}
func parseHistoryIP(raw string) string {
raw = strings.TrimSpace(strings.Trim(raw, "[]"))
if raw == "" {
return ""
}
if addr, err := netip.ParseAddr(raw); err == nil {
return addr.String()
}
if prefix, err := netip.ParsePrefix(raw); err == nil {
return prefix.String()
}
return ""
}
// var (not const) so tests can redirect to t.TempDir(). Production
// callers must not mutate at runtime.
var quarantineDir = "/opt/csm/quarantine"
// apiQuarantine lists quarantined files with metadata.
func (s *Server) apiQuarantine(w http.ResponseWriter, _ *http.Request) {
type quarantineEntry struct {
ID string `json:"id"`
Kind string `json:"kind"`
OriginalPath string `json:"original_path"`
Size int64 `json:"size"`
QuarantineAt time.Time `json:"quarantined_at,omitzero"`
Reason string `json:"reason"`
LiveState string `json:"live_state"`
OriginalModTime time.Time `json:"original_mtime,omitzero"`
}
var entries []quarantineEntry
// Scan both root quarantine dir and pre_clean subdirectory
rootMetas := listMetaFiles(quarantineDir)
preCleanMetas := listMetaFiles(filepath.Join(quarantineDir, "pre_clean"))
metaFiles := rootMetas
metaFiles = append(metaFiles, preCleanMetas...)
for _, metaFile := range metaFiles {
meta, err := readQuarantineMeta(metaFile)
if err != nil {
continue
}
// Hide entries whose original has been restored byte-identical to
// the archive: the archive is redundant and the UI should reflect
// the live filesystem, not the quarantine history. Divergence
// (missing, different size, different content) keeps the entry
// visible -- the operator still has to reconcile it.
archivePath := strings.TrimSuffix(metaFile, ".meta")
liveState := quarantineLiveState(archivePath, meta.OriginalPath)
if liveState == "restored_identical" {
continue
}
kind := "quarantine"
if strings.HasPrefix(quarantineEntryID(metaFile), preCleanQuarantineIDPrefix) {
kind = "pre_clean"
}
entries = append(entries, quarantineEntry{
ID: quarantineEntryID(metaFile),
Kind: kind,
OriginalPath: meta.OriginalPath,
Size: meta.Size,
QuarantineAt: meta.QuarantineAt.UTC(),
Reason: meta.Reason,
LiveState: liveState,
OriginalModTime: meta.OriginalModTime.UTC(),
})
}
// Sort newest first
sort.Slice(entries, func(i, j int) bool {
if entries[i].QuarantineAt.Equal(entries[j].QuarantineAt) {
return entries[i].ID < entries[j].ID
}
return entries[i].QuarantineAt.After(entries[j].QuarantineAt)
})
writeAll(w, entries)
}
// apiStats returns severity counts and per-check breakdown for the last 24
// hours. The summary is shared with the dashboard page and recomputed only
// when history changes.
func (s *Server) apiStats(w http.ResponseWriter, _ *http.Request) {
sum := s.statsSummary24h()
result := map[string]interface{}{
"last_24h": map[string]interface{}{
"critical": sum.critical,
"high": sum.high,
"warning": sum.warning,
"total": sum.critical + sum.high + sum.warning,
},
"by_check": sum.byCheck,
"accounts_at_risk": sum.atRisk,
"auto_response": map[string]int{
"blocked": sum.autoBlocked,
"quarantined": sum.autoQuarantined,
"killed": sum.autoKilled,
},
"top_accounts": sum.topAccounts,
"brute_force": sum.bruteForce,
}
if !sum.lastCritical.IsZero() {
result["last_critical"] = sum.lastCritical.UTC()
}
writeJSON(w, result)
}
func buildBruteForceSummary(ips map[string]int, types map[string]int) map[string]interface{} {
// Top attacker IPs
type ipCount struct {
IP string `json:"ip"`
Count int `json:"count"`
}
var topIPs []ipCount
for ip, count := range ips {
topIPs = append(topIPs, ipCount{ip, count})
}
sort.Slice(topIPs, func(i, j int) bool {
return topIPs[i].Count > topIPs[j].Count
})
if len(topIPs) > 10 {
topIPs = topIPs[:10]
}
total := 0
for _, v := range types {
total += v
}
return map[string]interface{}{
"total_attacks": total,
"unique_ips": len(ips),
"wp_login_count": types["wp-login"],
"xmlrpc_count": types["xmlrpc"] + types["xmlrpc-modsec"],
"top_ips": topIPs,
}
}
// apiStatsTrend returns daily finding counts by severity for the trend
// chart. Accepts optional ?days=N (default 30, clamped to the store's
// retention window).
func (s *Server) apiStatsTrend(w http.ResponseWriter, r *http.Request) {
days := 30
if v := r.URL.Query().Get("days"); v != "" {
if n, err := strconv.Atoi(v); err == nil && n > 0 {
days = n
}
}
writeAll(w, s.store.AggregateByDayN(days))
}
// apiStatsTimeline returns 24 hourly buckets for the findings timeline chart.
// Uses efficient bbolt cursor seeking instead of loading all findings into memory.
func (s *Server) apiStatsTimeline(w http.ResponseWriter, _ *http.Request) {
buckets, _ := s.timelineMemo.get(s.store.HistoryMark(), func() any {
return s.store.AggregateByHour()
}).([]store.HourBucket)
writeAll(w, buckets)
}
// apiHealth returns daemon health status.
func (s *Server) apiHealth(w http.ResponseWriter, _ *http.Request) {
health := map[string]interface{}{
"daemon_mode": true,
"uptime_seconds": int(time.Since(s.startTime).Seconds()),
"rules_loaded": s.signatureCount(),
"fanotify": s.fanotifyRunning(),
"log_watchers": s.logWatchersRunning(),
}
writeJSON(w, health)
}
// historyCSVMax bounds one CSV export; the history filters narrow it to reach
// older entries.
const historyCSVMax = 5000
// apiHistoryCSV exports the newest history entries matching the History
// filters as a CSV download.
func (s *Server) apiHistoryCSV(w http.ResponseWriter, r *http.Request) {
q, ok := parseHistoryQuery(w, r)
if !ok {
return
}
findings, _ := s.readHistoryPage(q, historyCSVMax, 0)
w.Header().Set("Content-Type", "text/csv")
w.Header().Set("Content-Disposition", "attachment; filename=csm-history.csv")
// CSV header
fmt.Fprintf(w, "Timestamp,Severity,Check,Message,Details\n")
for _, f := range findings {
sev := "WARNING"
switch f.Severity {
case alert.Critical:
sev = "CRITICAL"
case alert.High:
sev = "HIGH"
}
// Escape CSV fields
msg := csvEscape(f.Message)
details := csvEscape(f.Details)
fmt.Fprintf(w, "%s,%s,%s,%s,%s\n",
f.Timestamp.Format(time.RFC3339), sev, f.Check, msg, details)
}
}
// csvEscape quotes a field for CSV and neutralises spreadsheet formula
// triggers. Finding text is attacker-chosen (a filename, a User-Agent, a
// mailbox): a field starting with "=", "+", "-", "@", a tab or a carriage
// return became a live formula when the export was opened in a spreadsheet.
// Such fields are prefixed with a single quote, the convention spreadsheets
// use to force text, and quoted so the prefix survives.
func csvEscape(s string) string {
if s != "" && strings.ContainsAny(s[:1], "=+-@\t\r") {
s = "'" + s
}
if strings.ContainsAny(s, ",\"\n\r'") {
return "\"" + strings.ReplaceAll(s, "\"", "\"\"") + "\""
}
return s
}
// safeLogString renders a caller-controlled string into a log entry without
// allowing embedded CR/LF/control bytes to forge a separate log line.
func safeLogString(s string) string { return strconv.Quote(s) }
func daemonStartToken(start time.Time) string {
return start.UTC().Format(time.RFC3339Nano)
}
func (s *Server) daemonStartToken() string {
if s.provider != nil {
return daemonStartToken(s.provider.StartedAt())
}
return daemonStartToken(s.startTime)
}
// apiFix applies a known remediation action for a finding.
// POST /api/v1/fix body: {"check": "check_type", "message": "...", "details": "..."}
func (s *Server) apiFix(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details"`
FilePath string `json:"file_path"`
Key string `json:"key"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.Check == "" || req.Message == "" {
writeJSONError(w, "check and message are required", http.StatusBadRequest)
return
}
if !checks.HasFix(req.Check) {
writeJSONError(w, "no automated fix available for this check type", http.StatusBadRequest)
return
}
message, details, filePath, dismissKey, err := s.fixTargetFromStore(req.Key, req.Check, req.Message, req.Details, req.FilePath)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
result := s.applyFix(r.Context(), req.Check, message, details, filePath)
// If fix succeeded, dismiss from both alert state and latest findings.
if result.Success {
s.store.DismissFinding(dismissKey)
s.store.DismissLatestFinding(dismissKey)
s.auditLog(r, "fix", req.Check, result.Action)
}
writeRemediation(w, result)
}
// apiVerifyFinding re-checks whether a finding's condition still holds against
// the live filesystem. It dismisses a resolved finding, lowers an inert
// replacement to Warning, or restores an earlier automatic demotion. This lets
// an operator confirm a manual fix immediately instead of waiting for the next
// scan, and is the "Re-check" action behind a finding row.
// POST /api/v1/verify-finding body: {"check":"...","message":"...","details":"...","file_path":"...","key":"..."}
func (s *Server) apiVerifyFinding(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req verifyFindingRequest
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.Check == "" || req.Message == "" {
writeJSONError(w, "check and message are required", http.StatusBadRequest)
return
}
in, key, stored, found := s.verifyFindingInput(req)
in.Context = r.Context()
response := verifyFindingResponse{OK: true, VerifyResult: s.verifyFinding(in)}
switch {
case response.Checked && response.Resolved:
if key == "" {
key = req.Check + ":" + req.Message
}
s.store.DismissFinding(key)
s.store.DismissLatestFinding(key)
s.auditLog(r, "verify-resolved", req.Check, response.Detail)
// A severity change rewrites the stored finding, so it needs the exact
// snapshot the verifier read; a request that could not be matched to one
// leaves the finding alone rather than guessing which it meant.
case found && checks.ShouldRestoreSeverity(stored, response.VerifyResult):
if s.store.RestoreLatestFindingSeverity(stored) {
response.SeverityChange = "restored"
s.auditLog(r, "verify-restored", req.Check, response.Detail)
}
case found && checks.ShouldDemoteSeverity(stored, response.VerifyResult):
if s.store.DemoteLatestFinding(stored, alert.Warning) {
response.SeverityChange = "demoted"
s.auditLog(r, "verify-demoted", req.Check, response.Detail)
}
}
writeJSON(w, response)
}
// verifyFindingResponse distinguishes the verifier's recommendation from a
// state change the store actually accepted. Demote remains a verdict: it can be
// true for an already-demoted finding or after a concurrent scan replaced the
// snapshot, neither of which means this request changed the stored severity.
type verifyFindingResponse struct {
OK bool `json:"ok"`
checks.VerifyResult
SeverityChange string `json:"severity_change,omitempty"`
}
type verifyFindingRequest struct {
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details"`
FilePath string `json:"file_path"`
ContentSHA256 string `json:"content_sha256"`
Key string `json:"key"`
}
// verifyFindingInput builds the verifier input, and returns the stored finding
// it was built from so a caller applying a severity change can pass the exact
// snapshot the verifier saw.
func (s *Server) verifyFindingInput(req verifyFindingRequest) (checks.VerifyInput, string, alert.Finding, bool) {
in := checks.VerifyInput{
Check: req.Check, Message: req.Message, Details: req.Details,
Path: req.FilePath,
}
f, ok := s.latestFindingForVerify(req.Key, req.Check, req.Message)
if !ok {
return in, req.Key, alert.Finding{}, false
}
in.Message = f.Message
in.Details = f.Details
in.Path = f.FilePath
in.ContentSHA256 = f.ContentSHA256
in.DetectLogic = f.DetectLogic
return in, f.Key(), f, true
}
func (s *Server) latestFindingForVerify(key, check, message string) (alert.Finding, bool) {
var matched alert.Finding
found := false
for _, f := range s.store.LatestFindings() {
if f.Check != check || f.Message != message {
continue
}
if key != "" {
if f.Key() == key {
return f, true
}
continue
}
if found {
return alert.Finding{}, false
}
matched = f
found = true
}
return matched, found
}
const bulkFixBodyMax = 64 * 1024
// apiBulkFix applies fixes to multiple findings at once.
// POST /api/v1/fix-bulk body: [{"check":"...", "message":"...", "details":"..."}, ...]
func (s *Server) apiBulkFix(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var reqs []struct {
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details"`
FilePath string `json:"file_path"`
Key string `json:"key"`
}
if err := decodeJSONBodyLimited(w, r, bulkFixBodyMax, &reqs); err != nil {
writeJSONError(w, "invalid request body", http.StatusBadRequest)
return
}
if len(reqs) == 0 {
writeJSONError(w, "At least one fix is required", http.StatusBadRequest)
return
}
results := make([]bulkFixItem, 0, len(reqs))
for _, req := range reqs {
if !checks.HasFix(req.Check) {
results = append(results, bulkFixItem{Check: req.Check, Error: fmt.Sprintf("no fix for %s", req.Check)})
continue
}
message, details, filePath, dismissKey, err := s.fixTargetFromStore(req.Key, req.Check, req.Message, req.Details, req.FilePath)
if err != nil {
results = append(results, bulkFixItem{Check: req.Check, Error: err.Error()})
continue
}
result := s.applyFix(r.Context(), req.Check, message, details, filePath)
if result.Success {
s.store.DismissFinding(dismissKey)
s.store.DismissLatestFinding(dismissKey)
s.auditLog(r, "fix", req.Check, result.Action)
}
results = append(results, bulkFixItem{
Check: req.Check, OK: result.Success, Action: result.Action,
Description: result.Description, Error: result.Error, Reverted: result.Reverted,
})
}
succeeded := 0
for _, item := range results {
if item.OK {
succeeded++
}
}
fields := map[string]interface{}{
"results": results,
"total": len(results),
"succeeded": succeeded,
"failed": len(results) - succeeded,
}
if succeeded == 0 {
fields["error"] = "No fix applied"
writeJSONStatus(w, http.StatusUnprocessableEntity, fields)
return
}
writeOK(w, fields)
}
// bulkFixItem is one fix of a bulk request: which check, whether it applied,
// and what it did or why it did not.
type bulkFixItem struct {
Check string `json:"check"`
OK bool `json:"ok"`
Action string `json:"action,omitempty"`
Description string `json:"description,omitempty"`
Error string `json:"error,omitempty"`
Reverted bool `json:"reverted,omitempty"`
}
// bulkItemFailure names one item of a batch that did not apply and why.
type bulkItemFailure struct {
Item string `json:"item"`
Error string `json:"error"`
}
// apiAccounts returns the account names for the scan dropdown: the accounts a
// server-wide scan covers.
//
//nolint:unused // registered via mux.Handle in server.go
func (s *Server) apiAccounts(w http.ResponseWriter, _ *http.Request) {
accounts, err := s.scanAccounts(s.liveCfg())
if err != nil {
writeJSONError(w, "Could not list accounts", http.StatusInternalServerError)
return
}
writeAll(w, accounts)
}
// --- Action endpoints ---
// apiBlockIP blocks an IP via the firewall engine.
// POST /api/v1/block-ip body: {"ip": "1.2.3.4", "reason": "..."}
func (s *Server) apiBlockIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Duration string `json:"duration"`
// IncidentID, when set, notes the block on that incident so the
// timeline shows an operator acted. Optional: the firewall action is
// the point, the note is bookkeeping.
IncidentID json.RawMessage `json:"incident_id"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
if req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Audit, incident and threat records key on the canonical spelling.
req.IP = parsedIP.String()
if req.Reason == "" {
req.Reason = "Blocked via CSM Web UI"
}
dur, err := parseDuration(req.Duration)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
if s.blocker == nil {
writeJSONError(w, "Firewall engine not available", http.StatusServiceUnavailable)
return
}
// Operator-initiated: bypass auto_response.dry_run gate.
if err := blockIPForOperator(s.blocker, req.IP, req.Reason, dur); err != nil {
writeJSONError(w, fmt.Sprintf("Block failed: %v", err), http.StatusInternalServerError)
return
}
s.auditLog(r, "block_ip", req.IP, req.Reason)
// A failure to annotate must not turn a successful block into an error:
// the address is blocked either way, and a stale incident id is the
// operator's tab being out of date, not a fault worth refusing.
var incidentID string
if json.Unmarshal(req.IncidentID, &incidentID) == nil && incidentID != "" && s.incidentCorrelator != nil {
if err := s.incidentCorrelator.RecordOperatorBlock(incidentID, req.IP, dur); err != nil {
log.Printf("webui: could not note an operator block on incident %s: %v",
safeLogString(incidentID), err)
}
}
resp := map[string]interface{}{"ip": req.IP}
// The input chain accepts Cloudflare edges on 80/443 before the blocked
// drop, so a block of a covered IP does not stop its web traffic.
if cc, ok := s.blocker.(cloudflareChecker); ok && cc.CloudflareCovers(req.IP) {
resp["warning"] = firewall.CloudflareCoverageWarning
}
writeOK(w, resp)
}
// apiUnblockIP removes an IP from the firewall + cphulk.
// POST /api/v1/unblock-ip body: {"ip": "1.2.3.4"}
func (s *Server) apiUnblockIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
req.IP = parsedIP.String()
if s.blocker == nil {
writeJSONError(w, "Firewall engine not available", http.StatusServiceUnavailable)
return
}
if err := s.blocker.UnblockIP(req.IP); err != nil {
writeJSONError(w, fmt.Sprintf("Unblock failed: %v", err), http.StatusInternalServerError)
return
}
dropAutoBlockThreatRow(req.IP)
// Also flush from cphulk (cPanel brute force detector); best effort,
// the unblock is what was asked.
_ = flushCphulk(req.IP)
s.auditLog(r, "unblock_ip", req.IP, "manual unblock via UI")
writeOK(w, map[string]interface{}{"ip": req.IP})
}
// apiUnblockBulk unblocks multiple IPs at once.
func (s *Server) apiUnblockBulk(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IPs []string `json:"ips"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || len(req.IPs) == 0 {
writeJSONError(w, "IPs array is required", http.StatusBadRequest)
return
}
// 500 fits the largest Blocks-table page size (250 plus All) so one UI
// selection round-trips as a single request and a single undo token.
if len(req.IPs) > 500 {
writeJSONError(w, "IPs must be 1-500 items", http.StatusBadRequest)
return
}
if s.blocker == nil {
writeJSONError(w, "Firewall engine not available", http.StatusServiceUnavailable)
return
}
priorBlocks := make(map[string]firewall.BlockedEntry)
seen := make(map[string]bool, len(req.IPs))
succeeded := 0
unblocked := make([]string, 0, len(req.IPs))
failed := []bulkItemFailure{}
removedThreats := make([]undoThreatRow, 0, len(req.IPs))
for _, ip := range req.IPs {
parsed, err := parseAndValidateIP(ip)
if err != nil {
failed = append(failed, bulkItemFailure{Item: ip, Error: err.Error()})
continue
}
ip = parsed.String()
if seen[ip] {
continue
}
seen[ip] = true
before, err := s.unblockIPForUndo(ip)
if err != nil {
failed = append(failed, bulkItemFailure{Item: ip, Error: err.Error()})
continue
}
if before != nil {
priorBlocks[ip] = *before
}
if row, ok := captureUndoThreatRow(ip, true); ok {
removedThreats = append(removedThreats, row)
}
dropAutoBlockThreatRow(ip)
s.auditLog(r, "unblock_ip", ip, "bulk unblock via UI")
unblocked = append(unblocked, ip)
succeeded++
}
_ = flushCphulkIPs(unblocked) // best effort: the unblocks are what was asked
var undoToken string
if succeeded > 0 {
undoToken = s.recordUndoEntry(r, "firewall_bulk_unblock", undoInverseFirewallUnblock,
fmt.Sprintf("Unblocked %d IPs", succeeded),
undoPayloadIPs{
IPs: unblocked,
RestoreThreats: removedThreats,
BlockSnapshot: true,
RestoreBlocks: priorBlocks,
})
}
fields := map[string]interface{}{
"total": len(req.IPs),
"succeeded": succeeded,
"failed": failed,
}
if succeeded == 0 {
fields["error"] = "No address was unblocked"
writeJSONStatus(w, http.StatusUnprocessableEntity, fields)
return
}
if undoToken != "" {
fields["undo_token"] = undoToken
}
writeOK(w, fields)
}
// blockedEntry is a raw blocked IP record from firewall state.
type blockedEntry struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
BlockedAt time.Time `json:"blocked_at"`
ExpiresAt time.Time `json:"expires_at"`
}
type blockedView struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source"`
BlockedAt time.Time `json:"blocked_at,omitzero"`
// ExpiresAt is left out for a permanent block.
ExpiresAt time.Time `json:"expires_at,omitzero"`
}
func formatBlockedView(b blockedEntry) (blockedView, bool) {
if !b.ExpiresAt.IsZero() && time.Now().After(b.ExpiresAt) {
return blockedView{}, false // expired
}
view := blockedView{
IP: b.IP,
Reason: b.Reason,
Source: b.Source,
BlockedAt: b.BlockedAt.UTC(),
ExpiresAt: b.ExpiresAt.UTC(),
}
if view.Source == "" {
view.Source = firewall.InferProvenance("block", b.Reason)
}
return view, true
}
// apiBlockedIPs returns the list of currently blocked IPs.
func (s *Server) apiBlockedIPs(w http.ResponseWriter, _ *http.Request) {
result := []blockedView{}
fwFile := filepath.Join(s.cfg.StatePath, "firewall", "state.json")
_, fwStatErr := os.Stat(fwFile) // #nosec G304 -- filepath.Join under operator-configured StatePath.
fwState, fwErr := firewall.LoadState(s.cfg.StatePath)
if fwErr != nil && fwStatErr == nil {
// The engine state exists but cannot be read: an empty list would
// tell the operator nothing is blocked.
writeJSONError(w, "Firewall state unavailable", http.StatusInternalServerError)
return
}
if fwErr == nil && fwState != nil {
for _, entry := range fwState.Blocked {
b := blockedEntry{
IP: entry.IP,
Reason: entry.Reason,
Source: entry.Source,
BlockedAt: entry.BlockedAt,
ExpiresAt: entry.ExpiresAt,
}
if view, ok := formatBlockedView(b); ok {
result = append(result, view)
}
}
// A present engine state file wins even when empty. blocked_ips.json
// is only a legacy fallback when the engine file does not exist.
if fwStatErr == nil || len(fwState.Blocked) > 0 {
writeAll(w, result)
return
}
}
// Fall back to blocked_ips.json (legacy)
stateFile := filepath.Join(s.cfg.StatePath, "blocked_ips.json")
// #nosec G304 -- filepath.Join under operator-configured StatePath.
data, err := os.ReadFile(stateFile)
if os.IsNotExist(err) {
writeAll(w, result)
return
}
if err != nil {
writeJSONError(w, "Firewall state unavailable", http.StatusInternalServerError)
return
}
var blockState struct {
IPs []blockedEntry `json:"ips"`
}
if err := json.Unmarshal(data, &blockState); err != nil {
writeJSONError(w, "Firewall state unavailable", http.StatusInternalServerError)
return
}
for _, b := range blockState.IPs {
if view, ok := formatBlockedView(b); ok {
result = append(result, view)
}
}
writeAll(w, result)
}
// apiDismissFinding marks a finding as baseline (acknowledged/dismissed).
// POST /api/v1/dismiss body: {"key": "check:message"}
// dismissBulkMax bounds one dismiss request. The whole request is one undo
// entry, so a larger selection must be narrowed rather than split.
const dismissBulkMax = 500
func (s *Server) apiDismissFinding(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Key string `json:"key"`
Keys []string `json:"keys"`
}
if err := decodeJSONBodyLimited(w, r, 1<<20, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
keys := req.Keys
switch {
case req.Key != "" && len(req.Keys) > 0:
writeJSONError(w, "Send key or keys, not both", http.StatusBadRequest)
return
case req.Key != "":
keys = []string{req.Key}
case len(keys) == 0:
writeJSONError(w, "Key is required", http.StatusBadRequest)
return
case len(keys) > dismissBulkMax:
writeJSONError(w, fmt.Sprintf("At most %d findings per request", dismissBulkMax), http.StatusBadRequest)
return
}
uniqueKeys := make([]string, 0, len(keys))
seenKeys := make(map[string]bool, len(keys))
for _, key := range keys {
if key == "" {
writeJSONError(w, "Key is required", http.StatusBadRequest)
return
}
if !seenKeys[key] {
seenKeys[key] = true
uniqueKeys = append(uniqueKeys, key)
}
}
keys = uniqueKeys
undos := make([]state.DismissUndo, 0, len(keys))
for _, key := range keys {
undos = append(undos, s.store.DismissFindingWithUndo(key))
s.auditLog(r, "dismiss", key, "")
}
resp := map[string]interface{}{"count": len(keys)}
summary := fmt.Sprintf("Dismissed %d findings", len(keys))
if len(keys) == 1 {
resp["key"] = keys[0]
check, _ := state.ParseKey(keys[0])
summary = "Dismissed " + check + " finding"
}
if token := s.recordUndoEntry(r, "dismiss", undoInverseFindingUndismiss, summary,
undoPayloadIPs{Dismissals: undos}); token != "" {
resp["undo_token"] = token
}
writeOK(w, resp)
}
// apiQuarantinePreview returns the first 8KB of a quarantined file for inspection.
func (s *Server) apiQuarantinePreview(w http.ResponseWriter, r *http.Request) {
entry, err := resolveQuarantineEntry(r.URL.Query().Get("id"))
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
info, err := os.Stat(entry.ItemPath)
if err != nil {
writeJSONError(w, "not found", http.StatusNotFound)
return
}
if info.IsDir() {
writeJSON(w, map[string]interface{}{
"id": entry.ID, "is_dir": true,
"preview": "[directory - content preview not available]",
})
return
}
f, err := os.Open(entry.ItemPath)
if err != nil {
writeJSONError(w, "cannot read file", http.StatusInternalServerError)
return
}
defer f.Close()
buf := make([]byte, 8192)
n, _ := f.Read(buf)
writeJSON(w, map[string]interface{}{
"id": entry.ID,
"preview": string(buf[:n]),
"truncated": info.Size() > 8192,
"total_size": info.Size(),
})
}
// quarantineBulkDeleteMax bounds the files one bulk-delete request removes.
// The UI sends larger selections as several requests of this size.
const quarantineBulkDeleteMax = 100
// apiQuarantineBulkDelete permanently removes quarantined files and their metadata.
// removeQuarantineItem deletes one quarantined file or directory. Tests
// replace it to exercise a deletion the filesystem refuses.
var removeQuarantineItem = os.RemoveAll
func (s *Server) apiQuarantineBulkDelete(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IDs []string `json:"ids"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
if len(req.IDs) == 0 || len(req.IDs) > quarantineBulkDeleteMax {
writeJSONError(w, fmt.Sprintf("IDs must be 1-%d items", quarantineBulkDeleteMax), http.StatusBadRequest)
return
}
count := 0
deleted := []string{}
failed := []string{}
for _, id := range req.IDs {
entry, err := resolveQuarantineEntry(id)
if err != nil || !quarantineEntryDeletable(entry) {
failed = append(failed, id)
continue
}
if _, statErr := os.Lstat(entry.ItemPath); statErr == nil {
if err := removeQuarantineItem(entry.ItemPath); err != nil {
// Keep the sidecar: the list is built from sidecars, so the
// archive stays visible and the delete can be retried.
log.Printf("webui: failed to delete quarantined %s: %v", safeLogString(entry.ItemPath), err)
failed = append(failed, id)
continue
}
count++
} else if !os.IsNotExist(statErr) {
failed = append(failed, id)
continue
}
if err := os.Remove(entry.MetaPath); err != nil && !os.IsNotExist(err) {
log.Printf("webui: failed to remove quarantine meta %s: %v", safeLogString(entry.MetaPath), err)
}
deleted = append(deleted, id)
}
details := "deleted: " + strings.Join(deleted, ", ")
if len(failed) > 0 {
details += "; failed: " + strings.Join(failed, ", ")
}
s.auditLog(r, "quarantine_bulk_delete", fmt.Sprintf("%d files", count), details)
if len(deleted) == 0 {
writeJSONStatus(w, http.StatusUnprocessableEntity, map[string]interface{}{
"error": "No file was deleted", "count": 0, "failed": failed,
})
return
}
writeOK(w, map[string]interface{}{"count": count, "failed": failed})
}
// apiTestAlert sends a test finding through all configured alert channels.
func (s *Server) apiTestAlert(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
testFinding := []alert.Finding{{
Severity: alert.Warning,
Check: "test_alert",
Message: "Test alert from CSM Web UI",
Details: fmt.Sprintf("Sent by admin at %s", time.Now().Format("2006-01-02 15:04:05")),
Timestamp: time.Now(),
}}
err := alert.Dispatch(s.liveCfg(), testFinding)
if err != nil {
// The request was fine; the alert channel behind the daemon failed.
writeJSONError(w, "Alert delivery failed: "+err.Error(), http.StatusBadGateway)
return
}
s.auditLog(r, "test_alert", "notification", "sent test alert")
writeOK(w, nil)
}
// apiScanAccount runs an on-demand scan for a single cPanel account.
// POST /api/v1/scan-account body: {"account": "username"}
func (s *Server) apiScanAccount(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Account string `json:"account"`
}
if err := decodeJSONBodyLimited(w, r, 32*1024, &req); err != nil || req.Account == "" {
writeJSONError(w, "Account name is required", http.StatusBadRequest)
return
}
if err := validateAccountName(req.Account); err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Rate limit: only one scan at a time
if !s.acquireScan() {
writeJSONError(w, "A scan is already in progress. Please wait.", http.StatusTooManyRequests)
return
}
defer s.releaseScan()
// Extend the write deadline for this long-running request.
// Account scans can take several minutes; the default WriteTimeout
// causes ERR_HTTP2_PROTOCOL_ERROR in browsers when it fires mid-stream.
rc := http.NewResponseController(w)
_ = rc.SetWriteDeadline(time.Now().Add(longRequestTimeout))
start := time.Now()
findings := checks.RunAccountScan(s.liveCfg(), s.store, req.Account)
elapsed := time.Since(start).Round(time.Millisecond)
s.auditLog(r, "scan_account", req.Account, fmt.Sprintf("%d findings in %s", len(findings), elapsed))
writeOK(w, map[string]interface{}{
"account": req.Account,
"count": len(findings),
"elapsed_seconds": elapsed.Seconds(),
})
}
// parseModeString converts a permission string like "-rw-r--r--" to os.FileMode.
func parseModeString(s string) os.FileMode {
if len(s) < 10 {
return 0644
}
var mode os.FileMode
perms := s[len(s)-9:] // last 9 chars: "rwxr-xr-x"
bits := []os.FileMode{
0400, 0200, 0100, // owner r/w/x
0040, 0020, 0010, // group r/w/x
0004, 0002, 0001, // other r/w/x
}
for i, b := range bits {
if i < len(perms) && perms[i] != '-' {
mode |= b
}
}
for _, flag := range s[:len(s)-9] {
switch flag {
case 'u':
mode |= os.ModeSetuid
case 'g':
mode |= os.ModeSetgid
case 't':
mode |= os.ModeSticky
}
}
return mode
}
// flushCphulk removes brute-force login history for an IP from cPanel's cphulk.
// Callers should pre-validate `ip` with parseAndValidateIP. This function
// re-validates as defense-in-depth so a future caller that forgets cannot
// expose a shell-execution surface even if exec.Command itself does not
// invoke a shell.
func flushCphulk(ip string) error {
return flushCphulkIPs([]string{ip})
}
// flushCphulkIPs uses the WHM API's indexed array arguments so a bulk
// firewall action starts one whmapi1 process instead of one per address.
// It returns whmapi1's failure, including whmapi1 not being installed, so
// a caller reports the flush only when it ran.
func flushCphulkIPs(ips []string) error {
args := []string{"flush_cphulk_login_history_for_ips"}
valid := 0
for _, ip := range ips {
parsed, err := parseAndValidateIP(ip)
if err != nil {
continue
}
param := "ip"
if valid > 0 {
param = fmt.Sprintf("ip-%d", valid)
}
args = append(args, param+"="+parsed.String())
valid++
}
if valid == 0 {
return nil
}
// #nosec G204 -- whmapi1 is fixed and every argument value is parsed as
// an IP above. exec.Command passes arguments directly without a shell.
if _, err := exec.Command("whmapi1", args...).Output(); err != nil {
return fmt.Errorf("whmapi1: %w", err)
}
return nil
}
// apiExport returns a JSON bundle of exportable state.
func (s *Server) apiExport(w http.ResponseWriter, _ *http.Request) {
// Collect suppressions
suppressions := s.store.LoadSuppressions()
if suppressions == nil {
suppressions = []state.SuppressionRule{}
}
// Collect whitelist
var whitelist []checks.WhitelistIP
if tdb := checks.GetThreatDB(); tdb != nil {
whitelist = tdb.WhitelistedIPs()
}
if whitelist == nil {
whitelist = []checks.WhitelistIP{}
}
bundle := map[string]interface{}{
"exported_at": time.Now().UTC(),
"hostname": s.cfg.Hostname,
"suppressions": suppressions,
"whitelist": whitelist,
}
w.Header().Set("Content-Disposition", "attachment; filename=csm-state-export.json")
writeJSON(w, bundle)
}
// apiImport merges an exported state bundle into the current state.
func (s *Server) apiImport(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// The bundle is what /api/v1/export writes; the decoder refuses unknown
// fields, so every field the export carries is named here.
var bundle struct {
ExportedAt string `json:"exported_at"`
Hostname string `json:"hostname"`
Suppressions []state.SuppressionRule `json:"suppressions"`
Whitelist []checks.WhitelistIP `json:"whitelist"`
}
if err := decodeJSONBodyLimited(w, r, 512*1024, &bundle); err != nil {
writeJSONError(w, "invalid JSON body", http.StatusBadRequest)
return
}
imported, skipped := 0, 0
warning := ""
// Merge suppressions (dedup by ID)
if len(bundle.Suppressions) > 0 {
err := s.store.UpdateSuppressions(func(existing []state.SuppressionRule) ([]state.SuppressionRule, error) {
existingIDs := make(map[string]bool)
for _, rule := range existing {
existingIDs[rule.ID] = true
}
for _, rule := range bundle.Suppressions {
// Same contract as a rule added through the UI: a check name
// and a valid glob are required (otherwise the rule
// suppresses nothing and only clutters the list), and every
// rule needs an ID or it can never be deleted from the UI.
if !suppressionCheckName.MatchString(rule.Check) {
skipped++
continue
}
if rule.PathPattern != "" {
if _, err := filepath.Match(rule.PathPattern, ""); err != nil {
skipped++
continue
}
}
if rule.ID == "" {
rule.ID = newSuppressionID()
}
if rule.CreatedAt.IsZero() {
rule.CreatedAt = time.Now()
}
if !existingIDs[rule.ID] {
existingIDs[rule.ID] = true
existing = append(existing, rule)
imported++
}
}
return existing, nil
})
if err != nil {
writeJSONError(w, fmt.Sprintf("failed to save suppressions: %v", err), http.StatusInternalServerError)
return
}
}
// Merge whitelist IPs
if len(bundle.Whitelist) > 0 {
tdb := checks.GetThreatDB()
if tdb == nil {
skipped += len(bundle.Whitelist)
warning = "The threat database is not available; whitelist entries were not imported."
} else {
existingSet := make(map[string]bool)
for _, w := range tdb.WhitelistedIPs() {
existingSet[w.IP] = true
}
now := time.Now()
for _, entry := range bundle.Whitelist {
// Validate imported IPs like every interactive route does: an
// unvalidated bundle could otherwise poison the threat DB /
// firewall allow-list with malformed or attacker-chosen entries
// (whitelisting bypasses blocking). Use the canonical form.
ip, err := parseAndValidateIP(entry.IP)
// Entries from the configuration file are managed there, and an
// expired temporary entry has nothing left to import.
if err != nil || entry.Configured || (entry.ExpiresAt != nil && !entry.ExpiresAt.After(now)) {
skipped++
continue
}
canonical := ip.String()
if existingSet[canonical] {
continue
}
// A temporary entry stays temporary, with its remaining time.
if entry.ExpiresAt != nil {
tdb.TempWhitelist(canonical, entry.ExpiresAt.Sub(now))
} else {
tdb.AddWhitelist(canonical)
}
existingSet[canonical] = true
imported++
}
}
}
s.auditLog(r, "import", "state", fmt.Sprintf("imported %d items, skipped %d", imported, skipped))
resp := map[string]interface{}{"imported": imported, "skipped": skipped}
if warning != "" {
resp["warning"] = warning
}
writeOK(w, resp)
}
// apiFindingDetail returns detail about a specific finding including related actions.
func (s *Server) apiFindingDetail(w http.ResponseWriter, r *http.Request) {
check := r.URL.Query().Get("check")
message := r.URL.Query().Get("message")
if check == "" {
writeJSONError(w, "check is required", http.StatusBadRequest)
return
}
// Alert state is keyed by Finding.Key(), which folds in a hash of the
// details (and the source IP for IP-keyed checks); "check:message" is
// only the key of a finding with neither. Resolve the stored finding
// first so findings with details get their first/last-seen times.
key := check + ":" + message
if f, ok := s.latestFindingForVerify(r.URL.Query().Get("key"), check, message); ok {
key = f.Key()
}
// Get state entry for this finding (first/last seen)
var firstSeen, lastSeen time.Time
if entry, ok := s.store.EntryForKey(key); ok {
firstSeen = entry.FirstSeen.UTC()
lastSeen = entry.LastSeen.UTC()
}
// Search audit log for related actions
actions := s.searchAuditEntries(check, 20)
// Search history for related findings (same check type, last 50)
allHistory, _ := s.store.ReadHistory(2000, 0)
type histEntry struct {
Severity string `json:"severity"`
Check string `json:"check"`
Message string `json:"message"`
Timestamp time.Time `json:"timestamp"`
}
var related []histEntry
for _, f := range allHistory {
if len(related) >= 50 {
break
}
if f.Check == check {
related = append(related, histEntry{
Severity: f.Severity.String(),
Check: f.Check,
Message: f.Message,
Timestamp: f.Timestamp.UTC(),
})
}
}
detail := map[string]interface{}{
"check": check,
"message": message,
"actions": actions,
"related": related,
}
if !firstSeen.IsZero() {
detail["first_seen"] = firstSeen
detail["last_seen"] = lastSeen
}
writeJSON(w, detail)
}
// extractAccountFromFinding returns the cPanel account a finding belongs to:
// the owner the check recorded (TenantID, or CPUser for mail relay), else a
// /home/{user}/ path in the message, details or file path, else "Account: "
// / "user: " in the details field (used by login checks).
func extractAccountFromFinding(f alert.Finding) string {
for _, owner := range []string{f.TenantID, f.CPUser} {
if owner = strings.TrimSpace(owner); owner != "" {
return owner
}
}
if f.FilePath == "" && (f.Check == "wp_core_unverified" || f.Check == "wp_plugin_inventory_unverified") {
// Collapsed coverage warnings carry an account only when every
// installation shares it. Their bounded path sample cannot establish
// ownership, even when it happens to show just one account.
return ""
}
for _, s := range []string{f.Message, f.Details, f.FilePath} {
if idx := strings.Index(s, "/home/"); idx >= 0 {
rest := s[idx+6:]
if end := strings.IndexByte(rest, '/'); end > 0 {
return rest[:end]
}
}
}
for _, prefix := range []string{"Account: ", "user: "} {
if idx := strings.Index(f.Details, prefix); idx >= 0 {
rest := f.Details[idx+len(prefix):]
end := strings.IndexAny(rest, " \n\t,")
if end > 0 {
return rest[:end]
}
if len(rest) > 0 {
return rest
}
}
}
return ""
}
func writeJSONError(w http.ResponseWriter, message string, code int) {
writeJSONStatus(w, code, map[string]string{"error": message})
}
// writeJSON sends compact JSON: indentation added about a third to large
// lists such as findings and history, and nothing reads it but code.
func writeJSON(w http.ResponseWriter, data interface{}) {
writeJSONStatus(w, http.StatusOK, data)
}
// writeJSONStatus sends data as JSON with the given status code.
func writeJSONStatus(w http.ResponseWriter, code int, data interface{}) {
body, err := apiValue(data)
var encoded []byte
if err == nil {
encoded, err = json.Marshal(body)
}
if err != nil {
code = http.StatusInternalServerError
encoded, _ = json.Marshal(map[string]string{"error": err.Error()})
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(code)
_, _ = w.Write(append(encoded, '\n'))
}
// durationSeconds reads Go duration text such as "24h" as seconds. ok is
// false when the text is empty or not a duration.
func durationSeconds(text string) (float64, bool) {
d, err := time.ParseDuration(text)
if err != nil {
return 0, false
}
return d.Seconds(), true
}
// writeItems answers a collection: {"items": [...]} plus extra, which holds
// "total" when the handler counted the matches, the paging keys and side
// data. A nil list goes out as [] so an empty collection is never null.
func writeItems[T any](w http.ResponseWriter, items []T, extra map[string]interface{}) {
if items == nil {
items = []T{}
}
body := make(map[string]interface{}, len(extra)+1)
for k, v := range extra {
body[k] = v
}
body["items"] = items
writeJSON(w, body)
}
// writeAll answers a collection that holds every match, with their count.
func writeAll[T any](w http.ResponseWriter, items []T) {
writeItems(w, items, map[string]interface{}{"total": len(items)})
}
// writeCapped answers a collection cut to limit items out of total matches.
func writeCapped[T any](w http.ResponseWriter, items []T, total, limit int, extra map[string]interface{}) {
if extra == nil {
extra = map[string]interface{}{}
}
if len(items) > limit {
items = items[:limit]
}
extra["total"] = total
extra["offset"] = 0
extra["limit"] = limit
extra["truncated"] = total > len(items) || extra["truncated"] == true
writeItems(w, items, extra)
}
// writeOK answers a successful action: "ok": true plus the action's fields.
func writeOK(w http.ResponseWriter, fields map[string]interface{}) {
writeOKStatus(w, http.StatusOK, fields)
}
// writeOKStatus is writeOK with another 2xx status, such as 202 for work
// that continues after the response.
func writeOKStatus(w http.ResponseWriter, code int, fields map[string]interface{}) {
body := make(map[string]interface{}, len(fields)+1)
for k, v := range fields {
body[k] = v
}
body["ok"] = true
writeJSONStatus(w, code, body)
}
// writeRemediation answers one fix. A fix that did not apply is an error:
// 422 when the target was not eligible and left unchanged, 500 when applying
// it failed.
func writeRemediation(w http.ResponseWriter, res checks.RemediationResult) {
if !res.Success {
msg := res.Error
if msg == "" {
msg = "The fix did not apply"
}
code := http.StatusInternalServerError
if res.Refused {
code = http.StatusUnprocessableEntity
}
body := map[string]interface{}{"error": msg}
if res.Action != "" {
body["action"] = res.Action
}
writeJSONStatus(w, code, body)
return
}
fields := map[string]interface{}{"action": res.Action, "description": res.Description}
if res.Reverted {
fields["reverted"] = true
}
writeOK(w, fields)
}
// writeRequestError answers a failure from middleware that guards both API
// and page routes: JSON under /api/, plain text elsewhere.
func writeRequestError(w http.ResponseWriter, r *http.Request, msg string, code int) {
if strings.HasPrefix(r.URL.Path, "/api/") {
writeJSONError(w, msg, code)
return
}
http.Error(w, msg, code)
}
// apiNotFound answers every /api/ path no route matches. Without it the
// page catch-all served the dashboard HTML with 200 to API clients.
// Unauthenticated callers get 401, as for a real route.
func (s *Server) apiNotFound(w http.ResponseWriter, r *http.Request) {
if !s.tokenHasScope(r, "read") {
writeJSONError(w, "Unauthorized", http.StatusUnauthorized)
return
}
writeJSONError(w, "Not found", http.StatusNotFound)
}
// queryInt reads a non-negative integer query parameter. A missing, negative
// or non-numeric value gives defaultVal.
func queryInt(r *http.Request, key string, defaultVal int) int {
val := r.URL.Query().Get(key)
if val == "" {
return defaultVal
}
n, err := strconv.Atoi(val)
if err != nil || n < 0 {
return defaultVal
}
return n
}
package webui
import (
"encoding/json"
"fmt"
"net/http"
"time"
"github.com/pidginhost/csm/internal/alert"
)
// sseWriteTimeout caps how long each SSE write is allowed to block on a
// slow or stuck client. It must stay below the daemon's WebUI shutdown
// budget so an in-flight flush cannot outlive graceful shutdown.
const sseWriteTimeout = 3 * time.Second
// apiEvents streams findings to the client over Server-Sent Events. A
// subscriber connects once and receives a `data: {...}\n\n` block per
// finding plus a periodic `: keepalive\n\n` comment line every 25s so
// intermediate proxies don't time the connection out. Auth is checked
// by the upstream requireRead middleware.
func (s *Server) apiEvents(w http.ResponseWriter, r *http.Request) {
s.mu.RLock()
bus := s.findingBus
s.mu.RUnlock()
if bus == nil {
writeJSONError(w, "event bus not available", http.StatusServiceUnavailable)
return
}
if _, ok := w.(http.Flusher); !ok {
writeJSONError(w, "streaming unsupported", http.StatusInternalServerError)
return
}
// Reserve a subscriber slot before sending any stream headers so a flood of
// connections cannot exhaust daemon memory. Done up front so the cap can be
// reported as a clean 503 rather than mid-stream.
sub, ok := bus.TrySubscribe()
if !ok {
writeJSONError(w, "too many event stream subscribers", http.StatusServiceUnavailable)
return
}
shutdownDone := s.pruneDone
// The upstream middleware authenticates all subscribers. Only cookie
// sessions have revocation/idle state to recheck during a stream.
_, cookieErr := r.Cookie("csm_auth")
cookieStream := cookieErr == nil && !s.isBearerAuth(r)
streamStopped := func() bool {
if r.Context().Err() != nil {
return true
}
// Recheck session revocation/expiry before every event and heartbeat.
// A passive stream must not keep an idle browser session alive.
if cookieStream {
if _, ok := s.cookieSessionToken(r, "read", false); !ok {
return true
}
}
select {
case <-shutdownDone:
return true
default:
return false
}
}
defer func() {
if streamStopped() {
bus.Unsubscribe(sub)
} else {
bus.Abort(sub)
}
}()
rc := http.NewResponseController(w)
setWriteDeadline := func() error {
return rc.SetWriteDeadline(time.Now().Add(sseWriteTimeout))
}
if err := setWriteDeadline(); err != nil {
writeJSONError(w, "streaming write deadlines unsupported", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "text/event-stream")
w.Header().Set("Cache-Control", "no-cache")
w.Header().Set("Connection", "keep-alive")
w.Header().Set("X-Accel-Buffering", "no") // nginx won't buffer
writeFrame := func(format string, args ...any) error {
if err := setWriteDeadline(); err != nil {
return err
}
if _, err := fmt.Fprintf(w, format, args...); err != nil {
return err
}
return rc.Flush()
}
// Initial flush establishes the connection and proxies see the headers
// before the first event. Bound it by the same write deadline.
if err := writeFrame(""); err != nil {
return
}
keepalive := time.NewTicker(25 * time.Second)
defer keepalive.Stop()
for {
if streamStopped() {
return
}
select {
case <-r.Context().Done():
return
case <-shutdownDone:
return
case <-keepalive.C:
if streamStopped() {
return
}
if err := writeFrame(": keepalive\n\n"); err != nil {
return
}
case delivery, ok := <-sub.Events():
if !ok {
return
}
encodingFailed := false
err := delivery.Process(func(f alert.Finding) error {
if streamStopped() {
return nil
}
wire, err := apiValue(toAPIFinding(f))
if err != nil {
encodingFailed = true
return err
}
body, err := json.Marshal(wire)
if err != nil {
encodingFailed = true
return err
}
err = writeFrame("data: %s\n\n", body)
// Closing a tab can interrupt its write. Once demand is withdrawn,
// that cancellation must not become a host delivery failure.
if err != nil && streamStopped() {
return nil
}
return err
})
if err != nil && !encodingFailed {
return
}
}
}
}
package webui
import (
"fmt"
"path/filepath"
"github.com/pidginhost/csm/internal/alert"
)
// fixTargetFromStore pins a fix request to the stored finding it names. The
// returned key is also the key that must be dismissed after a successful fix.
// Only when the finding is no longer in the latest set does the client input
// stand on its own, bounded by the remediation roots as before.
func (s *Server) fixTargetFromStore(key, check, message, details, filePath string) (string, string, string, string, error) {
latest := s.store.LatestFindings()
if key != "" {
for _, f := range latest {
if f.Key() != key {
continue
}
if f.Check != check {
return "", "", "", "", fmt.Errorf("check does not match the stored finding")
}
return storedFixTarget(f.Message, f.Details, f.FilePath, f.Key(), filePath)
}
}
var matched alert.Finding
found := false
for _, f := range latest {
if f.Check != check || f.Message != message {
continue
}
if found {
return "", "", "", "", fmt.Errorf("finding key is required for an ambiguous fix target")
}
matched = f
found = true
}
if !found {
if key != "" {
return message, details, filePath, key, nil
}
return message, details, filePath, check + ":" + message, nil
}
return storedFixTarget(matched.Message, matched.Details, matched.FilePath, matched.Key(), filePath)
}
func storedFixTarget(message, details, storedPath, key, requestedPath string) (string, string, string, string, error) {
// Older findings may not have FilePath populated. In that case the stored
// message remains the authority and the remediation extracts its path from
// there; a caller-supplied path must not replace it.
if storedPath != "" && requestedPath != "" && filepath.Clean(requestedPath) != filepath.Clean(storedPath) {
return "", "", "", "", fmt.Errorf("file_path does not match the stored finding")
}
return message, details, storedPath, key, nil
}
package webui
import (
"encoding/json"
"errors"
"fmt"
"io"
"log"
"net/http"
"os"
"path/filepath"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/quarantinefs"
"github.com/pidginhost/csm/internal/safepath"
)
// apiQuarantineRestore restores a quarantined file to its original location.
// POST /api/v1/quarantine-restore body: {"id": "filename"}
func (s *Server) apiQuarantineRestore(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
ID string `json:"id"`
}
if err := decodeJSONBodyLimited(w, r, 16*1024, &req); err != nil || req.ID == "" {
writeJSONError(w, "ID is required", http.StatusBadRequest)
return
}
entry, err := resolveQuarantineEntry(req.ID)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
if !quarantineEntryDeletable(entry) {
writeJSONError(w, "Quarantine entry not found", http.StatusNotFound)
return
}
metaData, err := os.ReadFile(entry.MetaPath)
if err != nil {
writeJSONError(w, "Quarantine entry not found", http.StatusNotFound)
return
}
var meta checks.QuarantineMeta
if unmarshalErr := json.Unmarshal(metaData, &meta); unmarshalErr != nil {
writeJSONError(w, "Invalid metadata", http.StatusInternalServerError)
return
}
roots, rootErr := quarantineRootsForConfig(s.cfg)
restorePath, err := validateQuarantineRestorePath(meta.OriginalPath, roots)
if err != nil {
writeJSONError(w, errors.Join(err, rootErr).Error(), http.StatusBadRequest)
return
}
if quarantineRestoreAfterValidateForTest != nil {
quarantineRestoreAfterValidateForTest(restorePath)
}
target, err := openQuarantineRestoreTarget(restorePath, roots, meta.RestoreAction == "")
if err != nil {
writeJSONError(w, fmt.Sprintf("Cannot open restore destination: %v", err), http.StatusConflict)
return
}
defer target.Close()
// Check if quarantined item is a directory or file
quarInfo, err := os.Lstat(entry.ItemPath)
if err != nil {
writeJSONError(w, fmt.Sprintf("Cannot stat quarantined file: %v", err), http.StatusInternalServerError)
return
}
if quarInfo.Mode()&os.ModeSymlink != 0 {
writeJSONError(w, "Cannot restore symlink quarantine entry", http.StatusInternalServerError)
return
}
// Parse original mode from metadata (format: "-rw-r--r--" or "drwxr-xr-x")
restoredMode := os.FileMode(0644)
if meta.Mode != "" && len(meta.Mode) >= 10 {
restoredMode = parseModeString(meta.Mode)
}
if meta.RestoreAction != "" {
if filepath.Clean(filepath.Dir(entry.ItemPath)) != filepath.Join(quarantineDir, "pre_clean") {
writeJSONError(w, "Virtual-patch backups must come from pre_clean", http.StatusBadRequest)
return
}
if quarInfo.IsDir() {
writeJSONError(w, "Invalid virtual-patch backup", http.StatusInternalServerError)
return
}
if err := checks.RestoreVirtualPatchBackup(entry.ItemPath, target, meta); err != nil {
if errors.Is(err, checks.ErrVirtualPatchRestoreConflict) {
writeJSONError(w, err.Error(), http.StatusConflict)
return
}
writeJSONError(w, fmt.Sprintf("Cannot restore virtual-patch backup: %v", err), http.StatusInternalServerError)
return
}
if err := removeRestoredQuarantineEvidence(entry.ItemPath, entry.MetaPath); err != nil {
writeJSONError(w, err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "restore", restorePath, "virtual-patch restore")
writeOK(w, map[string]interface{}{
"path": restorePath,
"warning": "Virtual-patch reverted. Re-scan recommended.",
})
return
}
if err := target.Check(); err != nil {
writeJSONError(w, err.Error(), http.StatusConflict)
return
}
if quarInfo.IsDir() {
if err := restoreQuarantineDirectory(entry.ItemPath, target, restoredMode, meta); err != nil {
writeJSONError(w, fmt.Sprintf("Cannot restore directory: %v", err), http.StatusConflict)
return
}
} else {
// File restore: use O_EXCL to prevent overwriting an existing file
src, readErr := os.Open(entry.ItemPath)
if readErr != nil {
writeJSONError(w, fmt.Sprintf("Cannot read quarantined file: %v", readErr), http.StatusInternalServerError)
return
}
defer src.Close()
// Allocate cleanup space before creating the destination. A copy may
// fail because the filesystem is full, when mkdir can fail as well.
stage, stageName, stageErr := target.Parent.CreatePrivateTemp()
if stageErr != nil {
writeJSONError(w, fmt.Sprintf("Cannot stage restored file: %v", stageErr), http.StatusInternalServerError)
return
}
defer func() {
_ = stage.Close()
_ = target.Parent.RemoveDir(stageName)
}()
const stagedName = "restore"
dst, createErr := stage.OpenFile(stagedName, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0600)
if createErr != nil {
writeJSONError(w, fmt.Sprintf("Cannot create restored file: %v", createErr), http.StatusInternalServerError)
return
}
// Even a failed first stat can be cleaned safely inside this private
// directory; no tenant can substitute a different file under its name.
defer func() { _ = stage.Remove(stagedName) }()
defer dst.Close()
createdInfo, statErr := statQuarantineCreatedFile(dst)
if statErr != nil {
writeJSONError(w, fmt.Sprintf("Cannot stat restored file: %v", statErr), http.StatusInternalServerError)
return
}
// Keep the inode pinned after dst.Close, preventing inode reuse from
// making a foreign file pass cleanup's identity check.
guard, guardErr := stage.OpenFile(stagedName, os.O_RDONLY, 0)
if guardErr != nil {
writeJSONError(w, "Cannot pin restored file; quarantine retained", http.StatusInternalServerError)
return
}
defer guard.Close()
if err := target.Check(); err != nil {
writeJSONError(w, "Cannot restore - destination changed during restore", http.StatusConflict)
return
}
if err := stage.RenameTo(stagedName, target.Parent, target.Name); err != nil {
writeJSONError(w, fmt.Sprintf("Cannot restore - file already exists at original path: %v", err), http.StatusConflict)
return
}
keepDestination := false
defer func() {
if !keepDestination {
if err := discardQuarantineRestore(target, createdInfo, stage, stageName); err != nil {
log.Printf("webui: restore cleanup failed: %v", err)
}
}
}()
if quarantineRestoreAfterCreateForTest != nil {
quarantineRestoreAfterCreateForTest(restorePath)
}
if _, err := ensureOpenFileStillAtTarget(dst, target); err != nil {
_ = src.Close()
writeJSONError(w, "Cannot restore - destination changed during restore", http.StatusConflict)
return
}
_, copyErr := io.Copy(dst, src)
if closeErr := src.Close(); copyErr == nil && closeErr != nil {
copyErr = closeErr
}
if copyErr != nil {
writeJSONError(w, fmt.Sprintf("Cannot write restored file: %v", copyErr), http.StatusInternalServerError)
return
}
if quarantineRestoreBeforeFinalizeForTest != nil {
quarantineRestoreBeforeFinalizeForTest(restorePath)
}
if _, err := ensureOpenFileStillAtTarget(dst, target); err != nil {
writeJSONError(w, "Cannot restore - destination changed during restore", http.StatusConflict)
return
}
if err := dst.Chown(meta.Owner, meta.Group); err != nil {
writeJSONError(w, fmt.Sprintf("Cannot restore file ownership; quarantine retained: %v", err), http.StatusInternalServerError)
return
}
if err := dst.Chmod(restoredMode); err != nil {
writeJSONError(w, fmt.Sprintf("Cannot restore file mode: %v", err), http.StatusInternalServerError)
return
}
if !meta.OriginalModTime.IsZero() {
if err := restoreQuarantineModTime(dst, meta.OriginalModTime); err != nil {
writeJSONError(w, fmt.Sprintf("Cannot restore modification time; quarantine retained: %v", err), http.StatusInternalServerError)
return
}
}
restoredInfo, err := ensureOpenFileStillAtTarget(dst, target)
if err != nil {
writeJSONError(w, "Cannot restore - destination changed during restore", http.StatusConflict)
return
}
if err := syncQuarantineRestoredFile(dst); err != nil {
writeJSONError(w, fmt.Sprintf("Restored file could not be synced; quarantine retained: %v", err), http.StatusInternalServerError)
return
}
if err := dst.Close(); err != nil {
writeJSONError(w, fmt.Sprintf("Cannot write restored file: %v", err), http.StatusInternalServerError)
return
}
if err := ensureTargetStillNamesInfo(target, restoredInfo); err != nil {
writeJSONError(w, "Cannot restore - destination changed during restore", http.StatusConflict)
return
}
if err := syncQuarantineRestoredParent(target.Parent); err != nil {
writeJSONError(w, fmt.Sprintf("Restored directory could not be synced; quarantine retained: %v", err), http.StatusInternalServerError)
return
}
if err := ensureTargetStillNamesInfo(target, restoredInfo); err != nil {
writeJSONError(w, "Cannot restore - destination changed during restore", http.StatusConflict)
return
}
keepDestination = true
}
if err := removeRestoredQuarantineEvidence(entry.ItemPath, entry.MetaPath); err != nil {
writeJSONError(w, err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "restore", restorePath, "quarantine restore")
writeOK(w, map[string]interface{}{
"path": restorePath,
"warning": "File restored to original location. Re-scan recommended.",
})
}
// quarantineRestoreAfterCreateForTest lets race tests replace the path
// after O_EXCL creation; nil in production.
var quarantineRestoreAfterCreateForTest func(string)
var quarantineRestoreAfterValidateForTest func(string)
var quarantineRestoreBeforeFinalizeForTest func(string)
var quarantineRestoreBeforeDiscardForTest func()
var syncQuarantineRestoredFile = (*os.File).Sync
var statQuarantineCreatedFile = (*os.File).Stat
var syncQuarantineRestoredParent = (*safepath.Dir).Sync
var removeRestoredQuarantineEvidence = quarantinefs.RemoveEvidence
var restoreQuarantineModTime = safepath.SetModTime
// Isolate the name before testing its inode: checking and then unlinking in
// a tenant-writable parent would allow a replacement between those operations.
// Cleanup uses the pinned parent even if its original pathname was renamed.
func discardQuarantineRestore(target *safepath.Target, want os.FileInfo, stage *safepath.Dir, stageName string) error {
if quarantineRestoreBeforeDiscardForTest != nil {
quarantineRestoreBeforeDiscardForTest()
}
const name = "discard"
if err := target.Parent.RenameTo(target.Name, stage, name); err != nil {
if os.IsNotExist(err) {
return nil
}
return err
}
got, statErr := stage.Stat(name)
if statErr != nil || !os.SameFile(want, got) {
if err := stage.RenameTo(name, target.Parent, target.Name); err != nil {
return fmt.Errorf("destination changed; displaced entry retained in %s: %w", stageName, err)
}
return nil
}
if err := stage.Remove(name); err != nil {
return err
}
return target.Parent.Sync()
}
func ensureOpenFileStillAtTarget(f *os.File, target *safepath.Target) (os.FileInfo, error) {
fileInfo, err := f.Stat()
if err != nil {
return nil, fmt.Errorf("cannot stat restored file handle: %w", err)
}
if err := ensureTargetStillNamesInfo(target, fileInfo); err != nil {
return nil, err
}
return fileInfo, nil
}
func ensureTargetStillNamesInfo(target *safepath.Target, fileInfo os.FileInfo) error {
if err := target.Check(); err != nil {
return err
}
pathInfo, err := target.Parent.Stat(target.Name)
if err != nil {
return fmt.Errorf("cannot stat restored file path: %w", err)
}
if !os.SameFile(fileInfo, pathInfo) {
return fmt.Errorf("restore destination changed during restore")
}
return nil
}
func openQuarantineRestoreTarget(path string, roots []string, createParents bool) (*safepath.Target, error) {
var root string
for _, base := range roots {
if isPathUnder(path, base) && len(base) > len(root) {
root = base
}
}
if root == "" {
return nil, fmt.Errorf("restore path is outside the allowed restore roots")
}
relative, err := filepath.Rel(root, path)
if err != nil {
return nil, err
}
return safepath.OpenTarget(root, relative, createParents)
}
func restoreQuarantineDirectory(path string, target *safepath.Target, mode os.FileMode, meta checks.QuarantineMeta) error {
// Quarantine is daemon-owned. Both sides of the rename still use pinned
// parents so destination swaps cannot redirect the transaction.
source, err := safepath.OpenDir(filepath.Dir(path))
if err != nil {
return err
}
defer func() { _ = source.Close() }()
name := filepath.Base(path)
dir, err := source.OpenFile(name, os.O_RDONLY, 0)
if err != nil {
return err
}
defer dir.Close()
info, err := dir.Stat()
if err != nil {
return err
}
if !info.IsDir() {
return fmt.Errorf("quarantine entry is no longer a directory")
}
if err := quarantinefs.SyncTree(path, info); err != nil {
return err
}
if err := dir.Chown(meta.Owner, meta.Group); err != nil {
return err
}
if err := dir.Chmod(mode); err != nil {
return err
}
if !meta.OriginalModTime.IsZero() {
if err := restoreQuarantineModTime(dir, meta.OriginalModTime); err != nil {
return err
}
}
if err := syncQuarantineRestoredFile(dir); err != nil {
return err
}
if err := dir.Close(); err != nil {
return err
}
if err := target.Check(); err != nil {
return err
}
if err := source.RenameTo(name, target.Parent, target.Name); err != nil {
return err
}
if quarantineRestoreAfterDirectoryMoveForTest != nil {
quarantineRestoreAfterDirectoryMoveForTest()
}
if err := ensureTargetStillNamesInfo(target, info); err != nil {
if rollbackErr := rollbackQuarantineDirectory(source, name, target, info); rollbackErr != nil {
return fmt.Errorf("%w; restoring quarantine entry failed: %v", err, rollbackErr)
}
return err
}
if err := syncQuarantineRestoredParent(target.Parent); err != nil {
return fmt.Errorf("directory moved to restore destination but sync failed; quarantine metadata retained: %w", err)
}
if err := source.Sync(); err != nil {
return fmt.Errorf("directory restored but quarantine removal is not durable; metadata retained: %w", err)
}
return ensureTargetStillNamesInfo(target, info)
}
var quarantineRestoreAfterDirectoryMoveForTest func()
// Rollback must identify the isolated inode, not a name a tenant can replace
// between validation and rename. Foreign entries go back without overwriting.
func rollbackQuarantineDirectory(source *safepath.Dir, name string, target *safepath.Target, want os.FileInfo) error {
stage, stageName, err := source.CreatePrivateTemp()
if err != nil {
return err
}
defer func() {
_ = stage.Close()
_ = source.RemoveDir(stageName)
}()
const displaced = "displaced"
if err := target.Parent.RenameTo(target.Name, stage, displaced); err != nil {
return err
}
got, statErr := stage.Stat(displaced)
if statErr != nil || !os.SameFile(want, got) {
if err := stage.RenameTo(displaced, target.Parent, target.Name); err != nil {
return fmt.Errorf("destination changed; displaced entry retained in %s: %w", stageName, err)
}
return fmt.Errorf("destination changed; quarantine directory was moved by another writer")
}
if err := stage.RenameTo(displaced, source, name); err != nil {
return fmt.Errorf("quarantine directory retained in %s: %w", stageName, err)
}
return nil
}
package webui
import (
"bytes"
"context"
"encoding/json"
"io"
"log"
"net/http"
"os"
"path/filepath"
"strings"
"time"
)
const (
uiAuditFile = "ui_audit.jsonl"
maxUIAuditSize = 10 * 1024 * 1024 // 10 MB
)
// UIAuditEntry records a UI action for compliance and accountability.
type UIAuditEntry struct {
Timestamp time.Time `json:"timestamp"`
Action string `json:"action"` // block, unblock, dismiss, fix, whitelist, etc.
Target string `json:"target"` // IP, finding key, file path
Details string `json:"details,omitempty"` // extra context
SourceIP string `json:"source_ip,omitempty"` // admin's IP
// Actor is the name of the credential that acted, and Via says whether
// it came as an API token or a browser login.
Actor string `json:"actor,omitempty"`
Via string `json:"via,omitempty"`
}
// auditLog records a UI action to the audit log, attributed to the
// credential behind r.
func (s *Server) auditLog(r *http.Request, action, target, details string) {
actor, via := s.requestActor(r)
s.auditLogAs(r, actor, via, action, target, details)
}
// auditLogAs records a UI action for an actor resolved by the caller. Login
// has no session yet, and logout or revocation ends the session that made
// the request, so those handlers name the actor themselves.
func (s *Server) auditLogAs(r *http.Request, actor, via, action, target, details string) {
entry := UIAuditEntry{
Timestamp: time.Now(),
Action: action,
Target: target,
Details: details,
SourceIP: extractClientIP(r),
Actor: actor,
Via: via,
}
path := filepath.Join(s.cfg.StatePath, uiAuditFile)
data, err := json.Marshal(entry)
if err != nil {
return
}
data = append(data, '\n')
// Rotation and append are one step: two writers that both saw an
// oversized log would otherwise rotate twice and rename the fresh file
// over the archived history.
s.auditMu.Lock()
defer s.auditMu.Unlock()
if info, statErr := os.Stat(path); statErr == nil && info.Size() > maxUIAuditSize {
if renameErr := os.Rename(path, path+".1"); renameErr != nil {
log.Printf("webui: audit rotation failed for %s: %v", path, renameErr)
}
}
// #nosec G304 -- filepath.Join under operator-configured StatePath.
f, err := os.OpenFile(path, os.O_APPEND|os.O_CREATE|os.O_WRONLY, 0600)
if err != nil {
log.Printf("webui: audit open failed for %s: %v", path, err)
return
}
if _, err := f.Write(data); err != nil {
_ = f.Close()
log.Printf("webui: audit write failed for %s: %v", path, err)
return
}
if err := f.Close(); err != nil {
log.Printf("webui: audit close failed for %s: %v", path, err)
}
}
type auditActorKey struct{}
type auditActor struct{ name, via string }
func withAuditActor(r *http.Request, actor, via string) *http.Request {
return r.WithContext(context.WithValue(r.Context(), auditActorKey{}, auditActor{actor, via}))
}
// requestActor uses the identity captured at authorization so an action
// finishing after logout or expiry keeps its actor. Direct callers resolve
// against startup credentials without refreshing session activity.
func (s *Server) requestActor(r *http.Request) (actor, via string) {
if r == nil {
return "", ""
}
if actor, ok := r.Context().Value(auditActorKey{}).(auditActor); ok {
return actor.name, actor.via
}
// Match authorization's cookie-first order, including credential binding.
if tok, ok := s.cookieSessionCredential(r, "admin", false); ok {
return tok.Name, "browser"
}
if tok, ok := s.bearerCredentialWithScope(r, "admin"); ok {
return tok.Name, "api"
}
return "", ""
}
func extractClientIP(r *http.Request) string {
// Use RemoteAddr directly - XFF is trivially spoofable and this is
// a security audit log, so we only trust the TCP connection source.
return clientIPKey(r.RemoteAddr)
}
// readUIAuditLog returns the last N audit entries, newest first (all when
// limit is 0). It reads the log from the end, so asking for the newest few
// does not parse up to maxUIAuditSize of older entries.
func readUIAuditLog(statePath string, limit int) []UIAuditEntry {
path := filepath.Join(statePath, uiAuditFile)
// #nosec G304 -- filepath.Join under operator-configured statePath.
f, err := os.Open(path)
if err != nil {
return nil
}
defer func() { _ = f.Close() }()
info, err := f.Stat()
if err != nil {
return nil
}
return tailAuditEntries(f, info.Size(), limit)
}
// tailAuditEntries parses the audit lines in r[0:size] newest first and stops
// once limit entries are found (0 means all). Lines that are blank or not an
// entry are skipped; one line may be as long as the log (bulk undo targets).
func tailAuditEntries(r io.ReaderAt, size int64, limit int) []UIAuditEntry {
var out []UIAuditEntry
done := func(line []byte) bool {
line = bytes.TrimRight(line, "\r")
if len(line) == 0 {
return false
}
var entry UIAuditEntry
if json.Unmarshal(line, &entry) == nil {
out = append(out, entry)
}
return limit > 0 && len(out) >= limit
}
// head holds the bytes before the earliest newline seen so far: the
// unfinished start of a line. Chunks grow with it, so a long line is
// read in a logarithmic number of steps.
var head []byte
end := size
for end > 0 {
n := int64(64 * 1024)
if int64(len(head)) > n {
n = int64(len(head))
}
start := max(end-n, 0)
data := make([]byte, end-start, end-start+int64(len(head)))
if _, err := r.ReadAt(data, start); err != nil && err != io.EOF {
return out
}
data = append(data, head...)
for {
i := bytes.LastIndexByte(data, '\n')
if i < 0 {
break
}
if done(data[i+1:]) {
return out
}
data = data[:i]
}
head = data
end = start
}
done(head)
return out
}
// searchAuditEntries returns audit entries whose target or details contain the search string.
func (s *Server) searchAuditEntries(search string, limit int) []UIAuditEntry {
if search == "" || limit <= 0 {
return nil
}
readLimit := limit * 10
if readLimit > 5000 {
readLimit = 5000
}
all := readUIAuditLog(s.cfg.StatePath, readLimit)
searchLower := strings.ToLower(search)
var matched []UIAuditEntry
for _, e := range all {
if strings.Contains(strings.ToLower(e.Target), searchLower) ||
strings.Contains(strings.ToLower(e.Details), searchLower) ||
strings.Contains(strings.ToLower(e.Action), searchLower) {
matched = append(matched, e)
if len(matched) >= limit {
break
}
}
}
return matched
}
func (s *Server) handleAudit(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "audit.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
// uiAuditPageLimit is how many of the newest UI audit entries the API returns.
const uiAuditPageLimit = 200
// GET /api/v1/audit - return the newest UI audit log entries
func (s *Server) apiUIAudit(w http.ResponseWriter, r *http.Request) {
// One entry past the limit tells whether older entries were left out.
entries := readUIAuditLog(s.cfg.StatePath, uiAuditPageLimit+1)
truncated := len(entries) > uiAuditPageLimit
if truncated {
entries = entries[:uiAuditPageLimit]
}
writeItems(w, entries, map[string]interface{}{"offset": 0, "limit": uiAuditPageLimit, "truncated": truncated})
}
package webui
import (
"crypto/subtle"
"errors"
"log"
"net"
"net/http"
"net/netip"
"strings"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/session"
)
// --- Authentication ---
// tokenHasScope reports whether the credentials in r grant at least the
// requested scope. "read" is granted by any token; "admin" is granted only
// by admin-scope tokens. Constant-time compare against every configured
// token. Browser sessions bind to the named administrator credential that
// created them; the API credential itself is never accepted from a cookie.
func (s *Server) tokenHasScope(r *http.Request, want string) bool {
// Browser cookie session
if _, ok := s.cookieTokenWithScope(r, want); ok {
return true
}
// Bearer token
_, ok := s.bearerTokenWithScope(r, want)
return ok
}
func (s *Server) cookieTokenWithScope(r *http.Request, want string) (string, bool) {
return s.cookieSessionToken(r, want, sessionActivity(r))
}
// sessionActivity reports whether a request is the operator's own activity,
// which extends an idle browser session: a page load, or an API call the UI
// marks with X-CSM-Active because it followed the operator's input. Timer
// polls and the event stream do not, so a page left open still reaches the
// idle timeout.
func sessionActivity(r *http.Request) bool {
if r.URL.Path == "/metrics" || r.URL.Path == "/api/v1/events" {
return false
}
if strings.HasPrefix(r.URL.Path, "/api/") {
return r.Header.Get("X-CSM-Active") == "1"
}
if r.Method != http.MethodGet {
return false
}
// Fetches can poll HTML pages too. Only navigation counts as a page
// load; older browsers identify it by Accept instead of Fetch Metadata.
if mode := r.Header.Get("Sec-Fetch-Mode"); mode != "" {
return mode == "navigate"
}
return strings.Contains(r.Header.Get("Accept"), "text/html")
}
func (s *Server) cookieSessionToken(r *http.Request, want string, touch bool) (string, bool) {
tok, ok := s.cookieSessionCredential(r, want, touch)
return tok.Token, ok
}
func (s *Server) cookieSessionCredential(r *http.Request, want string, touch bool) (config.WebUIToken, bool) {
c, err := r.Cookie("csm_auth")
if err != nil || s.sessions == nil {
return config.WebUIToken{}, false
}
rec, err := s.sessions.Access(c.Value, s.sessionNow(), touch)
if err != nil {
return config.WebUIToken{}, false
}
for _, tok := range s.cfg.WebUI.Tokens {
if tok.Name == rec.Name && session.Hash(tok.Token) == rec.Credential && tok.Scope == "admin" && webUITokenAllows(tok, want) {
return tok, true
}
}
// A removed, rotated or downgraded login credential cannot leave a
// browser session active, even if a caller changes config in place.
_ = s.sessions.Revoke(rec.ID)
return config.WebUIToken{}, false
}
func (s *Server) bearerTokenWithScope(r *http.Request, want string) (string, bool) {
tok, ok := s.bearerCredentialWithScope(r, want)
return tok.Token, ok
}
func (s *Server) bearerCredentialWithScope(r *http.Request, want string) (config.WebUIToken, bool) {
auth := r.Header.Get("Authorization")
if !strings.HasPrefix(auth, "Bearer ") {
return config.WebUIToken{}, false
}
supplied := strings.TrimPrefix(auth, "Bearer ")
if supplied == "" {
return config.WebUIToken{}, false
}
for _, tok := range s.cfg.WebUI.Tokens {
if webUITokenMatches(supplied, tok) && webUITokenAllows(tok, want) {
return tok, true
}
}
return config.WebUIToken{}, false
}
func webUITokenMatches(supplied string, tok config.WebUIToken) bool {
return supplied != "" &&
tok.Token != "" &&
subtle.ConstantTimeCompare([]byte(supplied), []byte(tok.Token)) == 1
}
func webUITokenAllows(tok config.WebUIToken, want string) bool {
switch want {
case "read":
return tok.Scope == "read" || tok.Scope == "admin"
case "admin":
return tok.Scope == "admin"
default:
return false
}
}
func (s *Server) requireAuth(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if tok, ok := s.cookieSessionCredential(r, "admin", sessionActivity(r)); ok {
next.ServeHTTP(w, withAuditActor(r, tok.Name, "browser"))
return
}
if tok, ok := s.bearerCredentialWithScope(r, "admin"); ok {
next.ServeHTTP(w, withAuditActor(r, tok.Name, "api"))
return
}
// API calls get 401 JSON; browser requests get redirect to login
if strings.HasPrefix(r.URL.Path, "/api/") {
writeJSONError(w, "Unauthorized", http.StatusUnauthorized)
return
}
http.Redirect(w, r, "/login", http.StatusFound)
})
}
func (s *Server) requireRead(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if s.tokenHasScope(r, "read") {
if r.Method != http.MethodGet {
w.Header().Set("Allow", http.MethodGet)
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
next.ServeHTTP(w, r)
return
}
// API calls get 401 JSON; browser requests get redirect to login
if strings.HasPrefix(r.URL.Path, "/api/") {
writeJSONError(w, "Unauthorized", http.StatusUnauthorized)
return
}
http.Redirect(w, r, "/login", http.StatusFound)
})
}
// isAuthenticated is a thin shim used by handleLogin and metrics_api.
// New callers should prefer tokenHasScope directly.
func (s *Server) isAuthenticated(r *http.Request) bool {
return s.tokenHasScope(r, "admin")
}
// clientIPKey strips the port from a net/http RemoteAddr for use as a
// per-client rate-limit key, handling bracketed IPv6 ([::1]:443 -> ::1).
// Falls back to the raw value when there is no host:port to split, so a
// missing port never collapses distinct clients onto one key.
func clientIPKey(remoteAddr string) string {
if host, _, err := net.SplitHostPort(remoteAddr); err == nil {
return host
}
return remoteAddr
}
// rateLimitKey is the client a rate limit counts: the IPv4 address, or the
// /64 of an IPv6 address, since one IPv6 client is routed a whole /64 and can
// rotate addresses inside it.
func rateLimitKey(remoteAddr string) string {
host := clientIPKey(remoteAddr)
ip, err := netip.ParseAddr(host)
if err != nil {
return host
}
ip = ip.WithZone("").Unmap()
if ip.Is4() {
return ip.String()
}
return netip.PrefixFrom(ip, 64).Masked().String()
}
func (s *Server) handleLogin(w http.ResponseWriter, r *http.Request) {
// Redirect already-authenticated users to dashboard
if r.Method == http.MethodGet && s.isAuthenticated(r) {
http.Redirect(w, r, "/dashboard", http.StatusFound)
return
}
if r.Method == http.MethodGet {
s.renderTemplate(w, r, "login.html", nil)
return
}
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Rate limit: 5 attempts per minute per client (IPv4 address or IPv6 /64)
ip := rateLimitKey(r.RemoteAddr)
s.loginMu.Lock()
now := time.Now()
attempts := s.loginAttempts[ip]
var recent []time.Time
for _, t := range attempts {
if now.Sub(t) < time.Minute {
recent = append(recent, t)
}
}
if len(recent) >= 5 {
s.loginMu.Unlock()
http.Error(w, "Too many login attempts", http.StatusTooManyRequests)
return
}
if _, tracked := s.loginAttempts[ip]; !tracked {
boundRateLimitMap(s.loginAttempts, now.Add(-time.Minute))
}
s.loginAttempts[ip] = append(recent, now)
s.loginMu.Unlock()
r.Body = http.MaxBytesReader(w, r.Body, 4096)
if err := r.ParseForm(); err != nil {
http.Error(w, "Invalid login request", http.StatusBadRequest)
return
}
token := r.PostForm.Get("token")
// Only admin-scope tokens may log in via the browser form.
var loginName string
if token != "" {
for _, tok := range s.cfg.WebUI.Tokens {
if tok.Scope == "admin" && webUITokenMatches(token, tok) {
loginName = tok.Name
break
}
}
}
if loginName == "" {
// Failed logins go to the daemon log, not the UI audit trail: they
// are unauthenticated, and addresses from a whole network could
// otherwise rotate operator history out of the audit file.
// #nosec G706 -- RemoteAddr is the TCP peer address net/http sets, not request content.
log.Printf("webui: failed browser login from %s", safeLogString(clientIPKey(r.RemoteAddr)))
s.renderTemplate(w, r, "login.html", map[string]string{"Error": "Invalid token"})
return
}
if s.sessions == nil {
http.Error(w, "Session store unavailable", http.StatusServiceUnavailable)
return
}
previous := ""
if c, err := r.Cookie("csm_auth"); err == nil {
// Reauthentication already proved the new login credential. Look up
// the old session without touching it, and retain storage errors so
// a failed lookup cannot silently turn rotation into a fresh login.
_, err := s.sessions.Access(c.Value, s.sessionNow(), false)
if err != nil && !errors.Is(err, session.ErrInvalid) {
http.Error(w, "Session store unavailable", http.StatusServiceUnavailable)
return
}
if err == nil {
previous = c.Value
}
}
secret, record, err := s.sessions.Create(loginName, session.Hash(token), previous, clientIPKey(r.RemoteAddr), r.UserAgent(), s.sessionNow())
if errors.Is(err, session.ErrFull) {
// The Sessions page needs a login, so name the recovery paths here.
http.Error(w, "Browser session limit reached. Wait for idle sessions to expire, or revoke sessions with an admin API token.", http.StatusServiceUnavailable)
return
}
if err != nil {
http.Error(w, "Cannot create browser session", http.StatusServiceUnavailable)
return
}
http.SetCookie(w, &http.Cookie{
Name: "csm_auth",
Value: secret,
Path: "/",
HttpOnly: true,
Secure: true,
SameSite: http.SameSiteStrictMode,
MaxAge: int(record.Expires.Sub(record.Created).Seconds()),
Expires: record.Expires,
})
s.auditLogAs(r, loginName, "browser", "login", loginName, "browser session started")
http.Redirect(w, r, "/dashboard", http.StatusFound)
}
// --- Logout ---
func (s *Server) handleLogout(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if c, err := r.Cookie("csm_auth"); err == nil && s.sessions != nil {
rec, err := s.sessions.Access(c.Value, s.sessionNow(), false)
if err != nil && !errors.Is(err, session.ErrInvalid) {
http.Error(w, "Cannot revoke browser session", http.StatusServiceUnavailable)
return
}
if err == nil {
if err = s.sessions.Revoke(rec.ID); err != nil {
http.Error(w, "Cannot revoke browser session", http.StatusServiceUnavailable)
return
}
s.auditLogAs(r, rec.Name, "browser", "logout", rec.Name, "browser session ended")
}
}
clearBrowserCookie(w)
http.Redirect(w, r, "/login", http.StatusFound)
}
func clearBrowserCookie(w http.ResponseWriter) {
http.SetCookie(w, &http.Cookie{
Name: "csm_auth",
Value: "",
Path: "/",
HttpOnly: true,
Secure: true,
SameSite: http.SameSiteStrictMode,
MaxAge: -1, // delete cookie
})
}
package webui
import (
"fmt"
"log"
"net"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/store"
)
func (s *Server) blockIPPreservingLifetime(ip, reason string, ttl time.Duration) error {
if ttl > 0 {
if guarded, ok := s.blocker.(lifetimeKeepingBlocker); ok {
err := guarded.BlockIPForcePreserveLifetime(ip, reason, ttl)
checks.ObserveOperatorBlock(err, checks.BlockSourceWebUI)
return err
}
// Test blockers and integrations without the engine extension still
// honor the persisted lifetime. The live engine checks under its lock.
blocks, err := snapshotFirewallBlocks(s.cfg.StatePath)
if err != nil {
return err
}
if entry, ok := blocks[ip]; ok {
if entry.ExpiresAt.IsZero() {
return firewall.ErrPermanentBlock
}
if time.Until(entry.ExpiresAt) > ttl {
return firewall.ErrLongerBlock
}
}
}
return blockIPForOperator(s.blocker, ip, reason, ttl)
}
func snapshotFirewallBlocks(statePath string) (map[string]firewall.BlockedEntry, error) {
state, err := firewall.LoadState(statePath)
if err != nil {
return nil, err
}
blocks := make(map[string]firewall.BlockedEntry, len(state.Blocked))
for _, entry := range state.Blocked {
if ip := net.ParseIP(entry.IP); ip != nil {
entry.IP = ip.String()
blocks[entry.IP] = entry
}
}
return blocks, nil
}
func invalidateIPUndo(ip string) {
if db := store.Global(); db != nil {
if err := db.InvalidateUndoTargets([]string{ip}); err != nil {
log.Printf("webui: invalidate IP undo: %v", err)
}
}
}
func threatRowsForIP(rows []undoThreatRow, ip string) []undoThreatRow {
for _, row := range rows {
if row.IP == ip {
return []undoThreatRow{row}
}
}
return nil
}
// undoSnapshotBlocks restores each original deadline, never a fresh 24h
// window. Later firewall decisions invalidate the snapshot even if they
// originated outside the Web UI and did not touch the threat database.
func (s *Server) undoSnapshotBlocks(payload undoPayloadIPs, clearEvidence bool) (int, error) {
if s.blocker == nil {
return 0, fmt.Errorf("firewall engine not available")
}
current, err := snapshotFirewallBlocks(s.cfg.StatePath)
if err != nil {
return 0, err
}
for _, ip := range payload.IPs {
got, exists := current[ip]
want, expected := payload.ExpectedBlocks[ip]
if exists != expected || (exists && !firewall.SameBlockedEntry(got, want)) {
return 0, firewall.ErrBlockChanged
}
}
count := 0
for _, ip := range payload.IPs {
if _, err := parseAndValidateIP(ip); err != nil {
continue
}
prior, hadPrior := payload.RestoreBlocks[ip]
ttl := time.Duration(0)
if hadPrior && !prior.ExpiresAt.IsZero() {
ttl = time.Until(prior.ExpiresAt)
if ttl <= 0 {
hadPrior = false
}
}
if !hadPrior && !clearEvidence {
continue
}
if restorer, ok := s.blocker.(blockRestorer); ok {
var expected, restore *firewall.BlockedEntry
if entry, exists := payload.ExpectedBlocks[ip]; exists {
expected = &entry
}
if hadPrior {
restore = &prior
}
err := restorer.RestoreBlockIfUnchanged(ip, expected, restore)
if hadPrior {
checks.ObserveOperatorBlock(err, checks.BlockSourceWebUI)
}
if err != nil {
continue
}
} else if hadPrior {
if err := blockIPForOperator(s.blocker, ip, prior.Reason, ttl); err != nil {
continue
}
} else if err := s.blocker.UnblockIP(ip); err != nil {
continue
}
if clearEvidence {
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.RemovePermanent(ip)
}
}
restoreUndoThreatRows(threatRowsForIP(payload.RestoreThreats, ip))
if clearEvidence && !hadPrior {
_ = flushCphulk(ip) // best effort
}
count++
}
return count, nil
}
func (s *Server) blockIPForUndo(ip, reason string, ttl time.Duration) (*firewall.BlockedEntry, *firewall.BlockedEntry, error) {
if blocker, ok := s.blocker.(undoableBlocker); ok {
before, after, err := blocker.BlockIPForUndo(ip, reason, ttl)
checks.ObserveOperatorBlock(err, checks.BlockSourceWebUI)
return before, after, err
}
before, err := snapshotFirewallBlocks(s.cfg.StatePath)
if err != nil {
return nil, nil, err
}
if err = s.blockIPPreservingLifetime(ip, reason, ttl); err != nil {
return nil, nil, err
}
after, err := snapshotFirewallBlocks(s.cfg.StatePath)
return blockPointer(before, ip), blockPointer(after, ip), err
}
func (s *Server) unblockIPForUndo(ip string) (*firewall.BlockedEntry, error) {
if blocker, ok := s.blocker.(undoableUnblocker); ok {
return blocker.UnblockIPForUndo(ip)
}
before, err := snapshotFirewallBlocks(s.cfg.StatePath)
if err != nil {
return nil, err
}
if err := s.blocker.UnblockIP(ip); err != nil {
return nil, err
}
return blockPointer(before, ip), nil
}
func blockPointer(blocks map[string]firewall.BlockedEntry, ip string) *firewall.BlockedEntry {
if entry, exists := blocks[ip]; exists {
return &entry
}
return nil
}
package webui
import (
"net/http"
"github.com/pidginhost/csm/internal/checks"
)
// apiChallengeStats returns challenge-routing activity for the web UI: the live
// pending count, how many challenge timeouts escalated to a hard block, the
// cumulative routes per source check since daemon start, and the most recent
// routes. Read-only; safe for read-scope tokens.
func (s *Server) apiChallengeStats(w http.ResponseWriter, _ *http.Request) {
stats := checks.ChallengeUIStats()
resp := map[string]interface{}{
"routed_by_check": stats.RoutedByCheck,
"recent": stats.Recent,
"pending": 0,
"escalated": 0,
}
if s.provider != nil {
automation := s.provider.AutomationStatus()
resp["pending"] = automation.ChallengePending
resp["escalated"] = automation.ChallengeEscalated
}
writeJSON(w, resp)
}
package webui
import (
"net/http"
"sort"
"time"
"github.com/pidginhost/csm/internal/health"
)
// componentsProvider is the optional capability surface the daemon
// exposes for /api/v1/components. Tests and the API-only fallback path
// can omit it; the handler degrades to attached/unknown.
type componentsProvider interface {
WatcherStatuses() map[string]bool
WatcherChangedAt() map[string]time.Time
}
// componentsUpstreamProvider is the optional capability surface the
// daemon adds when it has per-watcher upstream probes wired. Returning
// nil / absent for a watcher means "no probe, do not flag deaf".
type componentsUpstreamProvider interface {
WatcherUpstream() map[string]health.UpstreamResult
}
// componentRow is the JSON shape returned per watcher.
type componentRow struct {
Name string `json:"name"`
Label string `json:"label"`
Status string `json:"status"` // "ok" | "degraded" | "deaf" | "idle" | "unknown"
Attached bool `json:"attached"`
ChangedAt time.Time `json:"changed_at,omitzero"`
LastEventAt time.Time `json:"last_event_at,omitzero"`
LastEventCheck string `json:"last_event_check,omitempty"`
UpstreamFresh *bool `json:"upstream_fresh,omitempty"`
UpstreamReason string `json:"upstream_reason,omitempty"`
UpstreamSeenAt time.Time `json:"upstream_seen_at,omitzero"`
}
// componentLabels maps the short watcher name to the operator-facing label.
// Watchers not in the map render with their raw key.
var componentLabels = map[string]string{
"fanotify": "Fanotify (filesystem)",
"audit": "Auditd",
"modsec": "ModSecurity audit",
"afalg": "AF_ALG kernel monitor",
"phprelay": "PHP relay watcher",
"maillog": "Mail log",
"email_av_spool": "Email AV spool",
"forwarder": "Forwarder watcher",
"pamlistener": "PAM listener",
"php_shield": "PHP Shield event log",
"connection": "Connection tracker",
"exec": "Exec monitor",
"sensitive": "Sensitive file monitor",
"accesslog": "Access log",
"dovecot_log": "Dovecot log",
"exim_mainlog": "Exim mainlog",
"cpanel_access_log": "cPanel access log",
"yara_worker": "YARA-X worker",
}
// componentCheckOrigin maps a finding Check name back to the watcher that
// emits it. Only the entries with a clear single-source origin are listed;
// finding names reused by periodic or retroactive scans intentionally have
// no entry so they do not advance a watcher's "last event" clock.
var componentCheckOrigin = map[string]string{
"cgi_backdoor_realtime": "fanotify",
"cgi_suspicious_location_realtime": "fanotify",
"credential_log_realtime": "fanotify",
"email_auth_failure_realtime": "maillog",
"email_av_degraded": "email_av_spool",
"email_av_encrypted_archive": "email_av_spool",
"email_av_parse_error": "email_av_spool",
"email_av_quarantine_error": "email_av_spool",
"email_av_timeout": "email_av_spool",
"email_compromised_account": "maillog",
"email_credential_leak": "maillog",
"email_dkim_failure": "maillog",
"email_malware": "email_av_spool",
"email_php_relay_action_dry_run": "phprelay",
"email_php_relay_action_failed": "phprelay",
"email_php_relay_action_skipped": "phprelay",
"email_php_relay_abuse": "phprelay",
"email_php_relay_account_volume_capped": "phprelay",
"email_php_relay_cpanel_limit_unreadable": "phprelay",
"email_php_relay_disabled": "phprelay",
"email_php_relay_inotify_overflow": "phprelay",
"email_php_relay_inotify_overflow_recovered": "phprelay",
"email_php_relay_msgindex_persist_failed": "phprelay",
"email_php_relay_no_exim": "phprelay",
"email_php_relay_overflow_scan_truncated": "phprelay",
"email_php_relay_path2b_disabled": "phprelay",
"email_php_relay_policies_reload": "phprelay",
"email_php_relay_rate_limit_hit": "phprelay",
"email_php_relay_sweep_failed": "phprelay",
"email_php_relay_watcher_failed": "phprelay",
"email_defer_fail_governor": "maillog",
"email_rate_critical": "maillog",
"email_rate_warning": "maillog",
"email_spam_outbreak": "maillog",
"email_spf_rejection": "maillog",
"executable_in_config_realtime": "fanotify",
"executable_in_tmp_realtime": "fanotify",
"exim_frozen_realtime": "maillog",
"fanotify_overflow": "fanotify",
"htaccess_injection_realtime": "fanotify",
"self_deleting_dropper_realtime": "fanotify",
"self_deleting_dropper_overflow": "fanotify",
"mail_account_compromised": "maillog",
"mail_account_spray": "maillog",
"mail_auth_backend_degraded": "maillog",
"mail_bruteforce": "maillog",
"mail_bruteforce_suspected": "maillog",
"mail_log_source_unavailable": "maillog",
"mail_subnet_spray": "maillog",
"modsec_block_escalation": "modsec",
"modsec_block_realtime": "modsec",
"modsec_classifier_gap": "modsec",
"modsec_csm_block_escalation": "modsec",
"modsec_low_confidence_burst": "modsec",
"modsec_warning_realtime": "modsec",
"obfuscated_php_realtime": "fanotify",
"credential_stuffing": "pamlistener",
"pam_bruteforce": "pamlistener",
"pam_login": "pamlistener",
"phishing_kit_realtime": "fanotify",
"phishing_realtime": "fanotify",
"php_shield_block": "php_shield",
"php_shield_eval": "php_shield",
"php_shield_webshell": "php_shield",
"php_config_realtime": "fanotify",
"php_dropper_realtime": "fanotify",
"php_in_sensitive_dir_realtime": "fanotify",
"php_in_uploads_realtime": "fanotify",
"signature_match_realtime": "fanotify",
"smtp_account_spray": "maillog",
"smtp_bruteforce": "maillog",
"smtp_probe_abuse": "maillog",
"smtp_subnet_spray": "maillog",
"webshell_content_realtime": "fanotify",
"webshell_realtime": "fanotify",
"yara_match_realtime": "fanotify",
"yara_match_scheduled": "scheduled",
"yara_realtime_scan_error": "fanotify",
}
// apiComponents returns one row per registered watcher with its live
// state, time since last state change, and the most recent finding it
// emitted within a 7-day lookback. Drives the dashboard component
// matrix.
func (s *Server) apiComponents(w http.ResponseWriter, _ *http.Request) {
cp, _ := s.provider.(componentsProvider)
if cp == nil {
writeAll(w, []componentRow{})
return
}
statuses := cp.WatcherStatuses()
changed := cp.WatcherChangedAt()
lastEvents := s.lastEventByWatcher(7 * 24 * time.Hour)
var upstream map[string]health.UpstreamResult
if up, ok := s.provider.(componentsUpstreamProvider); ok {
upstream = up.WatcherUpstream()
}
rows := make([]componentRow, 0, len(statuses))
for name, attached := range statuses {
row := componentRow{
Name: name,
Label: componentLabel(name),
Attached: attached,
}
if t, ok := changed[name]; ok {
row.ChangedAt = t.UTC()
}
if ev, ok := lastEvents[name]; ok && !ev.at.IsZero() {
row.LastEventAt = ev.at.UTC()
row.LastEventCheck = ev.check
}
var upstreamFresh *bool
if up, ok := upstream[name]; ok {
fresh := up.Fresh
upstreamFresh = &fresh
row.UpstreamFresh = upstreamFresh
row.UpstreamReason = up.Reason
row.UpstreamSeenAt = up.LastActivity.UTC()
}
row.Status = componentStatus(attached, lastEvents[name].at, upstreamFresh)
rows = append(rows, row)
}
sort.Slice(rows, func(i, j int) bool {
if rows[i].Status != rows[j].Status {
return componentStatusRank(rows[i].Status) < componentStatusRank(rows[j].Status)
}
return rows[i].Label < rows[j].Label
})
writeAll(w, rows)
}
type watcherEvent struct {
at time.Time
check string
}
// lastEventByWatcher returns the most recent finding per known watcher key
// within the lookback window. It reads the store's per-check index of newest
// timestamps rather than decoding the window's history on every poll.
// Checks not in componentCheckOrigin are skipped so periodic-scan output
// does not get attributed to a real-time watcher.
func (s *Server) lastEventByWatcher(window time.Duration) map[string]watcherEvent {
out := map[string]watcherEvent{}
if s.store == nil {
return out
}
since := time.Now().Add(-window)
for check, at := range s.store.LatestByCheck() {
watcher, ok := componentCheckOrigin[check]
if !ok || at.Before(since) {
continue
}
// Map order is random; equal times resolve by check name.
if cur, exists := out[watcher]; exists && (cur.at.After(at) || cur.at.Equal(at) && cur.check < check) {
continue
}
out[watcher] = watcherEvent{at: at, check: check}
}
// Also fold in the latest scan set so freshly-emitted findings appear
// before they have rolled into history.
for _, f := range s.store.LatestFindings() {
if f.Timestamp.Before(since) {
continue
}
watcher, ok := componentCheckOrigin[f.Check]
if !ok {
continue
}
if cur, exists := out[watcher]; exists && !cur.at.Before(f.Timestamp) {
continue
}
out[watcher] = watcherEvent{at: f.Timestamp, check: f.Check}
}
return out
}
func componentLabel(name string) string {
if l, ok := componentLabels[name]; ok {
return l
}
return name
}
// componentStatus collapses the per-row state into a UI bucket.
// - degraded: watcher detached (attempted setup, failed or fell off)
// - deaf: attached but the upstream feeding it has gone silent
// (probe registered, returned Fresh=false). Operator action needed
// before this watcher will ever produce events again.
// - ok: attached AND has produced at least one event recently
// - idle: attached, no events in window, and either no probe is
// wired or the probe still confirms the upstream is alive
// - unknown: not attached and no record either way (reserved)
func componentStatus(attached bool, lastEvent time.Time, upstreamFresh *bool) string {
if !attached {
return "degraded"
}
if upstreamFresh != nil && !*upstreamFresh {
return "deaf"
}
if lastEvent.IsZero() {
return "idle"
}
return "ok"
}
func componentStatusRank(status string) int {
switch status {
case "degraded":
return 0
case "deaf":
return 1
case "idle":
return 2
case "ok":
return 3
default:
return 4
}
}
package webui
import (
"net/http"
"sort"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/store"
)
// Cleanup-history handlers for the db_object_backups bbolt bucket.
// htaccess pre_clean backups already surface through the existing
// /api/v1/quarantine listing (their .meta sidecars now match the
// JSON QuarantineMeta shape). The bbolt-backed db_object_backups
// bucket needs its own list + restore endpoints because the data
// lives outside the filesystem-quarantine flow.
//
// All handlers are registered behind requireAuth in server.go;
// the restore handler additionally requires CSRF.
// dbObjectBackupEntry is the JSON shape returned to the cleanup-
// history UI. The Key field is opaque to the UI -- it round-trips
// to the restore endpoint as-is so the lookup is a single bbolt
// Get, not a multi-field reconstruction.
type dbObjectBackupEntry struct {
Key string `json:"key"`
Account string `json:"account"`
Schema string `json:"schema"`
Kind string `json:"kind"`
Name string `json:"name"`
DroppedAt time.Time `json:"dropped_at"`
DroppedBy string `json:"dropped_by"`
FindingID string `json:"finding_id,omitempty"`
BodyBytes int `json:"body_bytes"` // length of CreateSQL; surfaced for size hint
RestoredAt time.Time `json:"restored_at,omitzero"`
Restored bool `json:"restored"`
}
const dbObjectBackupPreviewBytes = 8 * 1024
// apiDBObjectBackups returns every record in the bucket, newest
// first by DroppedAt. The full CreateSQL is intentionally NOT
// returned in the listing -- those payloads can be large and the
// listing is meant for browse-and-pick. The preview endpoint returns
// one bounded CREATE SQL payload on demand.
func (s *Server) apiDBObjectBackups(w http.ResponseWriter, _ *http.Request) {
sdb := store.Global()
if sdb == nil {
writeAll(w, []dbObjectBackupEntry{})
return
}
records, keys, err := sdb.ListDBObjectBackupsAll()
if err != nil {
writeJSONError(w, "failed to list backups: "+err.Error(), http.StatusInternalServerError)
return
}
out := make([]dbObjectBackupEntry, 0, len(records))
for i, r := range records {
entry := dbObjectBackupEntry{
Key: keys[i],
Account: r.Account,
Schema: r.Schema,
Kind: r.Kind,
Name: r.Name,
DroppedAt: r.DroppedAt.UTC(),
DroppedBy: r.DroppedBy,
FindingID: r.FindingID,
BodyBytes: len(r.CreateSQL),
}
if !r.RestoredAt.IsZero() {
entry.Restored = true
entry.RestoredAt = r.RestoredAt.UTC()
}
out = append(out, entry)
}
// Newest first -- the bbolt key embeds unix-nanos so a string
// sort already produces chronological-by-drop-time when
// reversed. Doing it in Go keeps the contract explicit.
sortDBObjectBackupsNewestFirst(out)
writeAll(w, out)
}
// apiDBObjectBackupPreview returns a bounded CREATE SQL preview for one
// backup. Full restore still round-trips only the opaque key.
func (s *Server) apiDBObjectBackupPreview(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
key := r.URL.Query().Get("key")
if key == "" {
writeJSONError(w, "key is required", http.StatusBadRequest)
return
}
sdb := store.Global()
if sdb == nil {
writeJSONError(w, "bbolt store not available", http.StatusServiceUnavailable)
return
}
rec, ok, err := sdb.GetDBObjectBackupByKey(key)
if err != nil {
writeJSONError(w, "failed to read backup: "+err.Error(), http.StatusInternalServerError)
return
}
if !ok {
writeJSONError(w, "backup not found", http.StatusNotFound)
return
}
preview := rec.CreateSQL
truncated := false
if len(preview) > dbObjectBackupPreviewBytes {
preview = preview[:dbObjectBackupPreviewBytes]
truncated = true
}
writeJSON(w, map[string]any{
"key": key,
"account": rec.Account,
"schema": rec.Schema,
"kind": rec.Kind,
"name": rec.Name,
"preview": preview,
"truncated": truncated,
"total_size": len(rec.CreateSQL),
})
}
// apiDBObjectBackupRestore re-executes the captured CREATE SQL.
// POST body: {"key": "<bbolt key>"}. The handler delegates to
// checks.RestoreDBObjectBackup; CSRF is enforced upstream in
// server.go's requireCSRF wrapper.
func (s *Server) apiDBObjectBackupRestore(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Key string `json:"key"`
}
if err := decodeJSONBodyLimited(w, r, 16*1024, &req); err != nil || req.Key == "" {
writeJSONError(w, "key is required", http.StatusBadRequest)
return
}
// Tell a missing store or backup apart from a restore that failed.
sdb := store.Global()
if sdb == nil {
writeJSONError(w, "bbolt store not available", http.StatusServiceUnavailable)
return
}
if _, found, err := sdb.GetDBObjectBackupByKey(req.Key); err != nil {
writeJSONError(w, "looking up backup: "+err.Error(), http.StatusInternalServerError)
return
} else if !found {
writeJSONError(w, "backup not found (may have been pruned)", http.StatusNotFound)
return
}
result := checks.RestoreDBObjectBackup(req.Key)
if !result.Success {
writeJSONError(w, result.Message, http.StatusInternalServerError)
return
}
s.auditLog(r, "db_object_restore", req.Key, result.Message)
details := result.Details
if details == nil {
details = []string{}
}
writeOK(w, map[string]interface{}{
"message": result.Message,
"details": details,
})
}
// sortDBObjectBackupsNewestFirst sorts in place by DroppedAt
// descending. Local helper rather than relying on sort.Slice so
// the comparator is unambiguous in code review.
func sortDBObjectBackupsNewestFirst(entries []dbObjectBackupEntry) {
// Stable, so backups dropped at the same instant keep their store order.
sort.SliceStable(entries, func(i, j int) bool {
return entries[i].DroppedAt.After(entries[j].DroppedAt)
})
}
package webui
import (
"net/http"
"github.com/pidginhost/csm/internal/mailfwd/intel"
"github.com/pidginhost/csm/internal/platform"
)
// selectDeferralReporter picks the deferral-intel source for the host. Only
// cPanel/exim is wired; other platforms get the empty reporter until their
// adapters land (Phase 3).
func selectDeferralReporter() intel.Reporter {
if platform.Detect().IsCPanel() {
return intel.NewEximSource()
}
return intel.EmptyReporter{}
}
// apiEmailDeferrals handles GET /api/v1/email/deferrals and returns the
// outbound-deferral picture parsed from exim_mainlog: per-provider deferral
// rollup and per-outbound-IP reputation with stated reason codes.
func (s *Server) apiEmailDeferrals(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if s.deferralReporter == nil {
writeJSON(w, intel.Report{
Providers: []intel.ProviderRollup{},
OutboundIPs: []intel.OutboundIPRollup{},
})
return
}
rep, err := s.deferralReporter.Report()
if err != nil {
writeJSONError(w, "Failed to read deferral log", http.StatusInternalServerError)
return
}
writeJSON(w, rep)
}
package webui
import (
"context"
"io"
"net/http"
"os"
"os/exec"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailav"
"github.com/pidginhost/csm/internal/mailfwd/intel"
"github.com/pidginhost/csm/internal/systemdrun"
"github.com/pidginhost/csm/internal/yara"
)
type emailStatsResponse struct {
QueueSize int `json:"queue_size"`
// QueueUnavailable marks queue_size as meaningless because the depth could
// not be read. Without it a failed probe looks like an empty queue.
QueueUnavailable bool `json:"queue_unavailable"`
QueueWarn int `json:"queue_warn"`
QueueCrit int `json:"queue_crit"`
FrozenCount int `json:"frozen_count"`
// OldestAgeSeconds is the age of the oldest queued message; left out
// when the queue is empty or could not be read.
OldestAgeSeconds *int `json:"oldest_age_seconds,omitempty"`
SMTPBlock bool `json:"smtp_block"`
SMTPAllowUsers []string `json:"smtp_allow_users"`
SMTPPorts []int `json:"smtp_ports"`
PortFlood []portFloodEntry `json:"port_flood"`
TopSenders []senderEntry `json:"top_senders"`
}
type portFloodEntry struct {
Port int `json:"port"`
Proto string `json:"proto"`
Hits int `json:"hits"`
Seconds int `json:"window_seconds"`
}
type senderEntry struct {
Domain string `json:"domain"`
Count int `json:"count"`
}
func (s *Server) apiEmailStats(w http.ResponseWriter, _ *http.Request) {
cfg := s.liveCfg()
resp := emailStatsResponse{
QueueWarn: cfg.Thresholds.MailQueueWarn,
QueueCrit: cfg.Thresholds.MailQueueCrit,
}
// Live queue size and frozen/oldest via exim
var queueKnown bool
resp.QueueSize, queueKnown = eximQueueSize()
resp.QueueUnavailable = !queueKnown
var oldestAge string
resp.FrozenCount, oldestAge = eximQueueDetails()
if oldestAge != "" {
secs := intel.AgeToSeconds(oldestAge)
resp.OldestAgeSeconds = &secs
}
// Firewall config
fw := cfg.Firewall
resp.SMTPBlock = fw.SMTPBlock
resp.SMTPAllowUsers = fw.SMTPAllowUsers
if resp.SMTPAllowUsers == nil {
resp.SMTPAllowUsers = []string{}
}
resp.SMTPPorts = fw.SMTPPorts
if resp.SMTPPorts == nil {
resp.SMTPPorts = []int{}
}
// Port flood rules for SMTP ports only
smtpPorts := map[int]bool{25: true, 465: true, 587: true}
for _, pf := range fw.PortFlood {
if smtpPorts[pf.Port] {
resp.PortFlood = append(resp.PortFlood, portFloodEntry{
Port: pf.Port,
Proto: pf.Proto,
Hits: pf.Hits,
Seconds: pf.Seconds,
})
}
}
if resp.PortFlood == nil {
resp.PortFlood = []portFloodEntry{}
}
// Top senders from exim_mainlog
resp.TopSenders = topMailSenders(500, 10)
if resp.TopSenders == nil {
resp.TopSenders = []senderEntry{}
}
writeJSON(w, resp)
}
// apiEmailQuarantineList handles GET /api/v1/email/quarantine and returns all
// quarantined email messages, or an empty array if the quarantine is not configured.
func (s *Server) apiEmailQuarantineList(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
quarantine := s.emailQuarantineHandle()
if quarantine == nil {
writeAll(w, []emailav.QuarantineMetadata{})
return
}
msgs, err := quarantine.ListMessages()
if err != nil {
writeJSONError(w, "Failed to list quarantine", http.StatusInternalServerError)
return
}
writeAll(w, msgs)
}
// apiEmailQuarantineAction handles GET, POST (release), and DELETE operations on
// individual quarantined messages at /api/v1/email/quarantine/{msgID}.
func (s *Server) apiEmailQuarantineAction(w http.ResponseWriter, r *http.Request) {
// Extract everything after the prefix, e.g. "abc123" or "abc123/release".
tail := strings.TrimPrefix(r.URL.Path, "/api/v1/email/quarantine/")
if tail == "" {
writeJSONError(w, "Missing message ID", http.StatusBadRequest)
return
}
parts := strings.SplitN(tail, "/", 2)
msgID := parts[0]
action := ""
if len(parts) == 2 {
action = parts[1]
}
if err := validateEximMessageID(msgID); err != nil {
writeJSONError(w, "Invalid message ID: "+err.Error(), http.StatusBadRequest)
return
}
quarantine := s.emailQuarantineHandle()
if quarantine == nil {
writeJSONError(w, "Email quarantine not configured", http.StatusServiceUnavailable)
return
}
switch r.Method {
case http.MethodGet:
if action != "" {
writeJSONError(w, "Unknown action", http.StatusBadRequest)
return
}
meta, err := quarantine.GetMessage(msgID)
if err != nil {
writeJSONError(w, "Message not found", http.StatusNotFound)
return
}
writeJSON(w, meta)
case http.MethodPost:
if action != "release" {
writeJSONError(w, "Unknown action; use /release", http.StatusBadRequest)
return
}
if err := quarantine.ReleaseMessage(msgID); err != nil {
writeJSONError(w, "Failed to release message: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "email_quarantine_release", msgID, "released to the mail queue")
writeOK(w, map[string]interface{}{"message_id": msgID})
case http.MethodDelete:
if action != "" {
writeJSONError(w, "Unknown action", http.StatusBadRequest)
return
}
if err := quarantine.DeleteMessage(msgID); err != nil {
writeJSONError(w, "Failed to delete message: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "email_quarantine_delete", msgID, "deleted permanently")
writeOK(w, map[string]interface{}{"message_id": msgID})
default:
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
}
}
// Seams so the AV status handler can be tested without a live clamd.
var (
resolveClamdSocket = config.ResolveClamdSocket
clamdSocketAvailable = func(path string) bool {
return emailav.NewClamdScanner(path).Available()
}
)
type emailAVStatusResponse struct {
Enabled bool `json:"enabled"`
ClamdAvailable bool `json:"clamd_available"`
ClamdSocket string `json:"clamd_socket"`
YaraXAvailable bool `json:"yarax_available"`
YaraXRuleCount int `json:"yarax_rule_count"`
WatcherMode string `json:"watcher_mode"`
Quarantined int `json:"quarantined"`
}
// apiEmailAVStatus handles GET /api/v1/email/av/status and returns the current
// state of the email AV subsystem (ClamAV, YARA-X, quarantine count, watcher mode).
func (s *Server) apiEmailAVStatus(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
cfg := s.liveCfg()
resp := emailAVStatusResponse{
Enabled: cfg.EmailAV.Enabled,
}
// ClamAV availability. Resolve the socket the same way the daemon's
// scanner does: when the configured path is not answering the daemon
// falls back to a discovered one and mail really is being scanned, so
// probing the configured path here would report the subsystem down over
// a working scanner. Report the socket actually in use, not the stale
// setting, or the page sends the operator after the wrong thing.
clamdSocket, _ := resolveClamdSocket(cfg.EmailAV.ClamdSocket)
if clamdSocket == "" {
// Nothing configured and nothing discovered: name the documented
// default so the card points somewhere rather than showing a blank.
clamdSocket = config.DefaultClamdSocket
}
resp.ClamdSocket = clamdSocket
resp.ClamdAvailable = clamdSocketAvailable(clamdSocket)
// YARA-X availability and rule count. Active() covers both the
// in-process scanner and the out-of-process worker so this card
// reports the real rule count under either backend.
resp.YaraXAvailable = yara.Available()
if b := yara.Active(); b != nil {
resp.YaraXRuleCount = b.RuleCount()
}
// Watcher mode (set by daemon on startup).
resp.WatcherMode = s.emailAVMode()
if resp.WatcherMode == "" {
resp.WatcherMode = "disabled"
}
// Count of currently quarantined messages.
if quarantine := s.emailQuarantineHandle(); quarantine != nil {
msgs, err := quarantine.ListMessages()
if err == nil {
resp.Quarantined = len(msgs)
}
}
writeJSON(w, resp)
}
// eximLookPath and eximRun are injection points so tests can drive the queue
// probes without a live exim.
var (
eximLookPath = exec.LookPath
eximRun = func(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- name is the resolved systemd-run path or the fixed Exim command.
return exec.CommandContext(ctx, name, args...).Output()
}
)
// runExim executes an exim query as a transient unit forked by PID 1. The web
// server runs inside csm.service's ProtectSystem=strict sandbox, where /var/log
// is read-only; exim opens its main log for append even for a read-only query
// and aborts when it cannot, so a direct call returns nothing at all.
func runExim(timeout time.Duration, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
return systemdrun.Run(ctx, eximLookPath, eximRun, systemdrun.Options{
Pipe: true,
RuntimeMax: timeout,
}, "exim", args...)
}
// eximQueueSize returns the current Exim mail queue count. The second result is
// false when the depth could not be read: reporting zero there is
// indistinguishable from a healthy empty queue.
func eximQueueSize() (int, bool) {
out, err := runExim(5*time.Second, "-bpc")
if err != nil {
return 0, false
}
n, err := strconv.Atoi(strings.TrimSpace(string(out)))
if err != nil {
return 0, false
}
return n, true
}
// eximQueueDetails returns the frozen message count and the age of the oldest
// message in the queue. Uses `exim -bp` which lists all queued messages.
func eximQueueDetails() (frozen int, oldestAge string) {
out, err := runExim(10*time.Second, "-bp")
if err != nil {
return 0, ""
}
lines := strings.Split(string(out), "\n")
for _, line := range lines {
if strings.Contains(line, "*** frozen ***") {
frozen++
}
// First field of queue listing lines is the age (e.g., "4d", "15h", "30m")
fields := strings.Fields(line)
if len(fields) >= 3 {
age := fields[0]
// Only consider lines where first field looks like an age
if len(age) >= 2 && (age[len(age)-1] == 'd' || age[len(age)-1] == 'h' || age[len(age)-1] == 'm' || age[len(age)-1] == 's') {
if oldestAge == "" {
oldestAge = age // first entry is the oldest (queue sorted oldest first)
}
}
}
}
return frozen, oldestAge
}
// topMailSenders parses the last N lines of exim_mainlog and returns the
// top K sender domains by outbound message count.
func topMailSenders(tailLines, topK int) []senderEntry {
f, err := os.Open("/var/log/exim_mainlog")
if err != nil {
return nil
}
defer f.Close()
// Read tail of file
info, _ := f.Stat()
var data []byte
if info != nil && info.Size() > 256*1024 {
if _, err := f.Seek(-256*1024, 2); err != nil {
return nil
}
data, _ = io.ReadAll(f)
} else {
data, _ = io.ReadAll(f)
}
lines := strings.Split(string(data), "\n")
// Take last N lines
if len(lines) > tailLines {
lines = lines[len(lines)-tailLines:]
}
counts := make(map[string]int)
for _, line := range lines {
idx := strings.Index(line, " <= ")
if idx < 0 {
continue
}
rest := line[idx+4:]
fields := strings.Fields(rest)
if len(fields) < 1 {
continue
}
sender := fields[0]
atIdx := strings.LastIndex(sender, "@")
if atIdx < 0 {
continue
}
domain := sender[atIdx+1:]
if domain == "" || sender == "<>" || strings.HasPrefix(sender, "cPanel") {
continue
}
counts[domain]++
}
var entries []senderEntry
for domain, count := range counts {
entries = append(entries, senderEntry{Domain: domain, Count: count})
}
sort.Slice(entries, func(i, j int) bool {
return entries[i].Count > entries[j].Count
})
if len(entries) > topK {
entries = entries[:topK]
}
return entries
}
package webui
import (
"net/http"
"net/url"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/store"
)
// emailGroupsScanCap is the hard upper bound on matching findings retained per
// /api/v1/email/groups call. Bounded reads keep the workbench cheap on
// hosts that store thousands of mail-related findings per day.
const emailGroupsScanCap = 5000
// emailGroupsDefaultLimit / Max bound the number of grouped rows returned
// to the operator UI. The plan caps the email first viewport at ~250 nodes
// so 200 is the highest useful ceiling.
const (
emailGroupsDefaultLimit = 50
emailGroupsMaxLimit = 200
)
type emailGroup struct {
Kind string `json:"kind"`
Severity string `json:"severity"`
level alert.Severity
Title string `json:"title"`
Subject string `json:"subject"`
Count int `json:"count"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
Summary string `json:"summary"`
SampleFindings []apiFinding `json:"sample_findings"`
IPs []string `json:"ips,omitempty"`
TopIPs []string `json:"top_ips,omitempty"`
Domains []string `json:"domains,omitempty"`
MessageIDs []string `json:"message_ids,omitempty"`
}
type emailGroupsResponse struct {
Groups []emailGroup `json:"items"`
Total int `json:"total"`
Offset int `json:"offset"`
Limit int `json:"limit"`
From time.Time `json:"from"`
To time.Time `json:"to"`
Scanned int `json:"scanned"`
Truncated bool `json:"truncated"`
}
// emailKindForCheck maps an alert check name to its email-workbench group
// kind. Returns "" when the check is not part of the email surface and
// the finding should be skipped by /api/v1/email/groups.
func emailKindForCheck(check string) string {
switch check {
case "email_compromised_account",
"email_credential_leak",
"email_weak_password",
"mail_account_compromised",
"email_pipe_forwarder",
"email_suspicious_forwarder":
return "compromised_account"
case "email_spam_outbreak",
"email_rate_critical",
"email_rate_warning",
"email_php_relay_abuse",
"email_php_relay_action_failed",
"email_php_relay_rate_limit_hit",
"email_cloud_relay_abuse":
return "spam_outbreak"
case "email_auth_failure_realtime",
"email_suspicious_geo",
"mail_bruteforce",
"mail_bruteforce_suspected",
"mail_subnet_spray",
"mail_account_spray",
"smtp_bruteforce",
"smtp_subnet_spray",
"smtp_account_spray",
"smtp_probe_abuse":
return "auth_failure"
case "email_malware",
"email_phishing_content",
"email_av_degraded",
"email_av_encrypted_archive",
"email_av_timeout",
"email_av_parse_error",
"email_av_quarantine_error":
return "malware"
case "mail_per_account",
"mail_queue",
"mail_queue_unavailable",
"email_defer_fail_governor",
"exim_frozen_realtime":
return "queue_alert"
}
return ""
}
// emailGroupKey is the dedup key used to merge findings into a single
// grouped action row. Different kinds prefer different identity fields:
// auth failures cluster by mailbox/IP, spam/malware/compromised by mailbox
// or domain, and queue alerts by check name.
func emailGroupKey(kind string, f alert.Finding) string {
mailbox := strings.ToLower(strings.TrimSpace(f.Mailbox))
domain := strings.ToLower(strings.TrimSpace(f.Domain))
switch kind {
case "auth_failure":
if mailbox != "" {
return "mailbox:" + mailbox
}
if f.SourceIP != "" {
return "ip:" + f.SourceIP
}
if domain != "" {
return "domain:" + domain
}
return "auth:unknown"
case "queue_alert":
return "queue:" + f.Check
default:
if mailbox != "" {
return kind + ":mailbox:" + mailbox
}
if domain != "" {
return kind + ":domain:" + domain
}
if f.SourceIP != "" {
return kind + ":ip:" + f.SourceIP
}
// Fall back to message text so two distinct payloads with no
// identity fields still produce two groups instead of collapsing.
return kind + ":msg:" + strings.TrimSpace(f.Message)
}
}
// emailGroupTitle renders the human-readable identifier for a grouped row.
// Prefers mailbox > domain > source IP > message text. Queue alerts have
// hard-coded labels because their finding text varies by host.
func emailGroupTitle(kind string, f alert.Finding) string {
if kind == "queue_alert" {
switch f.Check {
case "mail_queue":
return "Mail queue threshold"
case "mail_queue_unavailable":
return "Mail queue unavailable"
case "mail_per_account":
return "Per-account mail volume"
case "exim_frozen_realtime":
return "Frozen mail queue"
}
}
if f.Mailbox != "" {
return f.Mailbox
}
if f.Domain != "" {
return f.Domain
}
if f.SourceIP != "" {
return f.SourceIP
}
return strings.TrimSpace(f.Message)
}
// emailGroupSubject describes the identity dimension behind the group --
// "mailbox", "domain", "ip", or "queue" -- so the UI can pick the right
// detail-panel tabs without re-reading the raw findings.
func emailGroupSubject(kind string, f alert.Finding) string {
if kind == "queue_alert" {
return "queue"
}
if f.Mailbox != "" {
return "mailbox"
}
if f.Domain != "" {
return "domain"
}
if f.SourceIP != "" {
return "ip"
}
return "unknown"
}
// buildEmailGroups walks the supplied findings (already bounded), merges
// matching findings into grouped rows, and returns the result sorted by
// severity (desc) then last-seen (desc). Pure function -- the HTTP
// handler is a thin wrapper so tests can drive grouping directly.
func buildEmailGroups(findings []alert.Finding, from, to time.Time, kindFilter string) []emailGroup {
type aggregator struct {
group *emailGroup
ipCounts map[string]int
domainSet map[string]struct{}
msgIDSet map[string]struct{}
samples []alert.Finding // newest-first
}
groups := make(map[string]*aggregator)
order := make([]string, 0)
for _, f := range findings {
ts := f.Timestamp
if !from.IsZero() && ts.Before(from) {
continue
}
if !to.IsZero() && ts.After(to) {
continue
}
kind := emailKindForCheck(f.Check)
if kind == "" {
continue
}
if kindFilter != "" && kindFilter != kind {
continue
}
key := emailGroupKey(kind, f)
agg, ok := groups[key]
if !ok {
agg = &aggregator{
group: &emailGroup{
Kind: kind,
level: f.Severity,
Title: emailGroupTitle(kind, f),
Subject: emailGroupSubject(kind, f),
FirstSeen: ts.UTC(),
LastSeen: ts.UTC(),
},
ipCounts: make(map[string]int),
domainSet: make(map[string]struct{}),
msgIDSet: make(map[string]struct{}),
}
groups[key] = agg
order = append(order, key)
}
agg.group.Count++
if f.Severity > agg.group.level {
agg.group.level = f.Severity
}
if ts.Before(agg.group.FirstSeen) {
agg.group.FirstSeen = ts.UTC()
}
if ts.After(agg.group.LastSeen) {
agg.group.LastSeen = ts.UTC()
}
if f.SourceIP != "" {
agg.ipCounts[f.SourceIP]++
}
if f.Domain != "" {
agg.domainSet[strings.ToLower(f.Domain)] = struct{}{}
}
for _, id := range f.MsgIDs {
if id != "" {
agg.msgIDSet[id] = struct{}{}
}
}
// Keep up to 3 most recent samples (assumes input is newest-first).
if len(agg.samples) < 3 {
agg.samples = append(agg.samples, f)
}
}
out := make([]emailGroup, 0, len(order))
for _, key := range order {
agg := groups[key]
g := agg.group
// Compose summary text: count + identity + IP/domain hint.
hint := ""
if g.Kind == "auth_failure" && len(agg.ipCounts) > 0 {
hint = " from " + plural(len(agg.ipCounts), "IP")
} else if len(agg.domainSet) > 1 {
hint = " across " + plural(len(agg.domainSet), "domain")
}
g.Summary = plural(g.Count, "event") + hint
g.Severity = g.level.String()
g.SampleFindings = toAPIFindings(agg.samples)
if len(agg.ipCounts) > 0 {
g.IPs = sortedKeys(agg.ipCounts)
g.TopIPs = topKeysByCount(agg.ipCounts, 5)
}
if len(agg.domainSet) > 0 {
g.Domains = sortedSetKeys(agg.domainSet)
}
if len(agg.msgIDSet) > 0 {
g.MessageIDs = sortedSetKeys(agg.msgIDSet)
if len(g.MessageIDs) > 10 {
g.MessageIDs = g.MessageIDs[:10]
}
}
out = append(out, *g)
}
sort.SliceStable(out, func(i, j int) bool {
if out[i].level != out[j].level {
return out[i].level > out[j].level
}
if out[i].Count != out[j].Count {
return out[i].Count > out[j].Count
}
return out[i].LastSeen.After(out[j].LastSeen)
})
return out
}
func plural(n int, label string) string {
if n == 1 {
return "1 " + label
}
return itoa(n) + " " + label + "s"
}
func itoa(n int) string {
// Avoid pulling strconv just for this hot path; keeps the helper inline.
if n == 0 {
return "0"
}
neg := n < 0
if neg {
n = -n
}
var buf [20]byte
i := len(buf)
for n > 0 {
i--
buf[i] = byte('0' + n%10)
n /= 10
}
if neg {
i--
buf[i] = '-'
}
return string(buf[i:])
}
func sortedKeys(m map[string]int) []string {
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
return keys
}
func sortedSetKeys(m map[string]struct{}) []string {
keys := make([]string, 0, len(m))
for k := range m {
keys = append(keys, k)
}
sort.Strings(keys)
return keys
}
// topKeysByCount returns up to k entries from m sorted by descending count
// (ties broken alphabetically) so the UI shows the dominant attackers
// first.
func topKeysByCount(m map[string]int, k int) []string {
type entry struct {
key string
count int
}
entries := make([]entry, 0, len(m))
for key, c := range m {
entries = append(entries, entry{key, c})
}
sort.SliceStable(entries, func(i, j int) bool {
if entries[i].count != entries[j].count {
return entries[i].count > entries[j].count
}
return entries[i].key < entries[j].key
})
if k < len(entries) {
entries = entries[:k]
}
out := make([]string, len(entries))
for i, e := range entries {
out[i] = e.key
}
return out
}
// historyRangeQuery reads the from and to parameters the history endpoints
// share, with the meaning store.ParseHistoryBound gives them: to is
// exclusive. A missing bound takes its default. An unreadable one is a 400,
// written here, and ok is false.
func historyRangeQuery(w http.ResponseWriter, q url.Values, defFrom, defTo time.Time) (from, to time.Time, ok bool) {
from, err := store.ParseHistoryBound(q.Get("from"), false)
if err != nil {
writeJSONError(w, "Invalid from: use YYYY-MM-DD or an RFC 3339 time", http.StatusBadRequest)
return from, to, false
}
to, err = store.ParseHistoryBound(q.Get("to"), true)
if err != nil {
writeJSONError(w, "Invalid to: use YYYY-MM-DD or an RFC 3339 time", http.StatusBadRequest)
return from, to, false
}
if from.IsZero() {
from = defFrom
}
if to.IsZero() {
to = defTo
}
return from, to, true
}
// apiEmailGroups handles GET /api/v1/email/groups. Returns server-side
// grouped action rows for the email workbench. Read-scope tokens may
// call this endpoint -- it does not mutate state.
func (s *Server) apiEmailGroups(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
q := r.URL.Query()
limit := queryInt(r, "limit", emailGroupsDefaultLimit)
if limit <= 0 || limit > emailGroupsMaxLimit {
limit = emailGroupsDefaultLimit
}
now := time.Now()
from, to, ok := historyRangeQuery(w, q, now.Add(-24*time.Hour), now)
if !ok {
return
}
if to.Before(from) {
from, to = to, from
}
kindFilter := q.Get("kind")
writeJSON(w, s.emailMemo("groups?"+q.Encode(), func() any {
return s.buildEmailGroupsResponse(from, to, kindFilter, limit)
}))
}
func (s *Server) buildEmailGroupsResponse(from, to time.Time, kindFilter string, limit int) emailGroupsResponse {
var findings []alert.Finding
if s.store != nil {
// Filter while walking history so unrelated findings, or findings
// newer than the requested range, never use up the scan budget.
findings = s.store.SearchHistorySince(from, emailGroupsScanCap+1, func(f alert.Finding) bool {
if !f.Timestamp.Before(to) {
return false
}
kind := emailKindForCheck(f.Check)
return kind != "" && (kindFilter == "" || kind == kindFilter)
})
}
truncated := false
if len(findings) > emailGroupsScanCap {
findings = findings[:emailGroupsScanCap]
truncated = true
}
scanned := len(findings)
groups := buildEmailGroups(findings, from, to, kindFilter)
total := len(groups)
if len(groups) > limit {
truncated = true
groups = groups[:limit]
}
return emailGroupsResponse{
Groups: groups,
Total: total,
Limit: limit,
From: from.UTC(),
To: to.UTC(),
Scanned: scanned,
Truncated: truncated,
}
}
// emailMemo reuses an email workbench result for the same query while
// history is unchanged; the page polls these every minute.
func (s *Server) emailMemo(key string, compute func() any) any {
if s.store == nil {
return compute()
}
return s.emailMemos.memo(key).get(s.store.HistoryMark(), compute)
}
package webui
import (
"context"
"encoding/json"
"fmt"
"net"
"net/http"
"os/exec"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall"
csmlog "github.com/pidginhost/csm/internal/log"
"github.com/pidginhost/csm/internal/platform"
"github.com/pidginhost/csm/internal/store"
)
// dropAutoBlockThreatRow removes the auto-block threat row for ip after an
// operator unblock. Operator permanent blocks are left in place: a
// firewall-only unblock must not silently clear a deliberate block. Under
// older builds a stale auto-block row could outlive the firewall block and
// ip_reputation would re-flag the IP into a new block loop.
func dropAutoBlockThreatRow(ip string) {
if parsed := net.ParseIP(ip); parsed != nil {
ip = parsed.String()
}
if sdb := store.Global(); sdb != nil {
_, _ = sdb.RemoveTemporaryBlock(ip)
}
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.RemoveTemporary(ip)
}
}
type firewallAllowView struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source"`
// ExpiresAt is left out for a permanent rule.
ExpiresAt time.Time `json:"expires_at,omitzero"`
}
type firewallPortAllowView struct {
IP string `json:"ip"`
Port int `json:"port"`
Proto string `json:"proto"`
Reason string `json:"reason"`
Source string `json:"source"`
}
const cphulkFirewallCheckTimeout = 8 * time.Second
var firewallCheckCommandOutput = func(ctx context.Context, name string, args ...string) ([]byte, error) {
// #nosec G204 -- command names are fixed by trusted call sites; HTTP input
// is parsed as an IP and passed as an execve argument without shell expansion.
return exec.CommandContext(ctx, name, args...).Output()
}
// apiFirewallStatus returns the firewall engine configuration and state summary.
func (s *Server) apiFirewallStatus(w http.ResponseWriter, _ *http.Request) {
cfg := config.EffectiveFirewallConfig(s.liveCfg())
state, err := firewall.LoadState(s.cfg.StatePath)
if err != nil {
writeJSONError(w, "firewall state unavailable (corrupt state file)", http.StatusInternalServerError)
return
}
now := time.Now()
blockedPermanent := 0
blockedTemporary := 0
for _, entry := range state.Blocked {
if entry.ExpiresAt.IsZero() {
blockedPermanent++
continue
}
if now.Before(entry.ExpiresAt) {
blockedTemporary++
}
}
allowPermanent := 0
allowTemporary := 0
for _, entry := range state.Allowed {
if entry.ExpiresAt.IsZero() {
allowPermanent++
continue
}
if now.Before(entry.ExpiresAt) {
allowTemporary++
}
}
result := map[string]interface{}{
"enabled": cfg.Enabled,
"ipv6": cfg.IPv6,
"tcp_in": cfg.TCPIn,
"tcp_out": cfg.TCPOut,
"udp_in": cfg.UDPIn,
"udp_out": cfg.UDPOut,
"restricted_tcp": cfg.RestrictedTCP,
"passive_ftp": [2]int{cfg.PassiveFTPStart, cfg.PassiveFTPEnd},
"conn_rate_limit": cfg.ConnRateLimit,
"conn_limit": cfg.ConnLimit,
"syn_flood_protection": cfg.SYNFloodProtection,
"udp_flood": cfg.UDPFlood,
"smtp_block": cfg.SMTPBlock,
"log_dropped": cfg.LogDropped,
"deny_ip_limit": cfg.DenyIPLimit,
"blocked_count": blockedPermanent + blockedTemporary,
"blocked_net_count": len(state.BlockedNet),
"blocked_permanent": blockedPermanent,
"blocked_temporary": blockedTemporary,
"allowed_count": allowPermanent + allowTemporary,
"allow_permanent": allowPermanent,
"allow_temporary": allowTemporary,
"port_allow_count": len(state.PortAllowed),
"infra_ips": cfg.InfraIPs,
"infra_count": len(cfg.InfraIPs),
"port_flood_rules": len(cfg.PortFlood),
"country_block": cfg.CountryBlock,
"dyndns_hosts": cfg.DynDNSHosts,
}
writeJSON(w, result)
}
// apiFirewallAllowed returns active firewall allow rules and port exceptions.
func (s *Server) apiFirewallAllowed(w http.ResponseWriter, _ *http.Request) {
state, err := firewall.LoadState(s.cfg.StatePath)
if err != nil {
writeJSONError(w, "firewall state unavailable (corrupt state file)", http.StatusInternalServerError)
return
}
now := time.Now()
allowed := make([]firewallAllowView, 0, len(state.Allowed))
for _, entry := range state.Allowed {
if !entry.ExpiresAt.IsZero() && !now.Before(entry.ExpiresAt) {
continue
}
view := firewallAllowView{
IP: entry.IP,
Reason: entry.Reason,
Source: entry.Source,
ExpiresAt: entry.ExpiresAt.UTC(),
}
if view.Source == "" {
view.Source = firewall.InferProvenance("allow", entry.Reason)
}
allowed = append(allowed, view)
}
sort.Slice(allowed, func(i, j int) bool {
return allowed[i].IP < allowed[j].IP
})
portAllowed := make([]firewallPortAllowView, 0, len(state.PortAllowed))
for _, entry := range state.PortAllowed {
portAllowed = append(portAllowed, firewallPortAllowView{
IP: entry.IP,
Port: entry.Port,
Proto: entry.Proto,
Reason: entry.Reason,
Source: entry.Source,
})
if portAllowed[len(portAllowed)-1].Source == "" {
portAllowed[len(portAllowed)-1].Source = firewall.InferProvenance("allow_port", entry.Reason)
}
}
sort.Slice(portAllowed, func(i, j int) bool {
if portAllowed[i].IP != portAllowed[j].IP {
return portAllowed[i].IP < portAllowed[j].IP
}
if portAllowed[i].Port != portAllowed[j].Port {
return portAllowed[i].Port < portAllowed[j].Port
}
return portAllowed[i].Proto < portAllowed[j].Proto
})
writeJSON(w, map[string]interface{}{
"allowed": allowed,
"port_allowed": portAllowed,
})
}
// apiFirewallAllowIP adds a firewall allow rule, temporary when duration > 0.
func (s *Server) apiFirewallAllowIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Duration string `json:"duration"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Audit, incident and threat records key on the canonical spelling.
req.IP = parsedIP.String()
if req.Reason == "" {
req.Reason = "Allowed via CSM Web UI"
}
dur, err := parseDuration(req.Duration)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
if dur > 0 {
allower, ok := s.blocker.(ipTempAllower)
if !ok || allower == nil {
writeJSONError(w, "Firewall allow rules are not available", http.StatusServiceUnavailable)
return
}
if err := allower.TempAllowIP(req.IP, req.Reason, dur); err != nil {
writeJSONError(w, fmt.Sprintf("Allow failed: %v", err), http.StatusInternalServerError)
return
}
s.auditLog(r, "firewall_allow", req.IP, fmt.Sprintf("temporary allow %s: %s", dur, req.Reason))
writeOK(w, map[string]interface{}{"ip": req.IP, "temporary": true})
return
}
allower, ok := s.blocker.(ipAllower)
if !ok || allower == nil {
writeJSONError(w, "Firewall allow rules are not available", http.StatusServiceUnavailable)
return
}
if err := allower.AllowIP(req.IP, req.Reason); err != nil {
writeJSONError(w, fmt.Sprintf("Allow failed: %v", err), http.StatusInternalServerError)
return
}
s.auditLog(r, "firewall_allow", req.IP, "permanent allow: "+req.Reason)
writeOK(w, map[string]interface{}{"ip": req.IP, "temporary": false})
}
// apiFirewallRemoveAllow removes a firewall allow rule.
func (s *Server) apiFirewallRemoveAllow(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, fmt.Sprintf("invalid IP address: %s", req.IP), http.StatusBadRequest)
return
}
// Audit, incident and threat records key on the canonical spelling.
req.IP = parsedIP.String()
allower, ok := s.blocker.(allowRemover)
if !ok || allower == nil {
writeJSONError(w, "Firewall allow rules are not available", http.StatusServiceUnavailable)
return
}
if err := allower.RemoveAllowIP(req.IP); err != nil {
writeJSONError(w, fmt.Sprintf("Remove failed: %v", err), http.StatusInternalServerError)
return
}
s.auditLog(r, "firewall_remove_allow", req.IP, "removed allow rule")
writeOK(w, map[string]interface{}{"ip": req.IP})
}
// apiFirewallAudit returns recent firewall audit log entries.
func (s *Server) apiFirewallAudit(w http.ResponseWriter, r *http.Request) {
limit := queryInt(r, "limit", 100)
// Filters apply to the whole log and the limit to what they matched, so a
// search reaches entries older than the newest page.
entries := firewall.ReadAuditLog(s.cfg.StatePath, 0)
type auditView struct {
Timestamp time.Time `json:"timestamp"`
Action string `json:"action"`
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source"`
// DurationSeconds is the block or allow lifetime; left out when
// the action had none.
DurationSeconds float64 `json:"duration_seconds,omitempty"`
}
search := strings.ToLower(strings.TrimSpace(r.URL.Query().Get("search")))
actionFilter := strings.ToLower(strings.TrimSpace(r.URL.Query().Get("action")))
sourceFilter := strings.ToLower(strings.TrimSpace(r.URL.Query().Get("source")))
var result []auditView
for _, e := range entries {
source := e.Source
if source == "" {
source = firewall.InferProvenance(e.Action, e.Reason)
}
if actionFilter != "" && strings.ToLower(e.Action) != actionFilter {
continue
}
if sourceFilter != "" && strings.ToLower(source) != sourceFilter {
continue
}
if search != "" {
haystack := strings.ToLower(strings.Join([]string{e.Action, e.IP, e.Reason, source}, " "))
if !strings.Contains(haystack, search) {
continue
}
}
view := auditView{
Timestamp: e.Timestamp.UTC(),
Action: e.Action,
IP: e.IP,
Reason: e.Reason,
Source: source,
}
// The log stores the lifetime as Go duration text such as "24h0m0s".
if secs, ok := durationSeconds(e.Duration); ok {
view.DurationSeconds = secs
}
result = append(result, view)
}
if limit == 0 {
writeAll(w, result)
return
}
// The newest entries are last in the log.
total := len(result)
if total > limit {
result = result[total-limit:]
}
writeCapped(w, result, total, limit, nil)
}
// apiFirewallSubnets returns currently blocked subnets.
func (s *Server) apiFirewallSubnets(w http.ResponseWriter, _ *http.Request) {
state, err := firewall.LoadState(s.cfg.StatePath)
if err != nil {
writeJSONError(w, "firewall state unavailable (corrupt state file)", http.StatusInternalServerError)
return
}
type subnetView struct {
CIDR string `json:"cidr"`
Reason string `json:"reason"`
Source string `json:"source"`
BlockedAt time.Time `json:"blocked_at,omitzero"`
// ExpiresAt is left out for a permanent block.
ExpiresAt time.Time `json:"expires_at,omitzero"`
}
var result []subnetView
for _, sn := range state.BlockedNet {
v := subnetView{
CIDR: sn.CIDR,
Reason: sn.Reason,
Source: sn.Source,
BlockedAt: sn.BlockedAt.UTC(),
ExpiresAt: sn.ExpiresAt.UTC(),
}
if v.Source == "" {
v.Source = firewall.InferProvenance("block_subnet", sn.Reason)
}
result = append(result, v)
}
writeAll(w, result)
}
// apiFirewallDenySubnet blocks a subnet via the firewall engine.
func (s *Server) apiFirewallDenySubnet(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
CIDR string `json:"cidr"`
Reason string `json:"reason"`
Duration string `json:"duration"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.CIDR == "" {
writeJSONError(w, "CIDR is required", http.StatusBadRequest)
return
}
if _, err := validateCIDR(req.CIDR); err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
if req.Reason == "" {
req.Reason = "Blocked via CSM Web UI"
}
dur, err := parseDuration(req.Duration)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
sb, ok := s.blocker.(subnetBlocker)
if !ok || sb == nil {
writeJSONError(w, "Firewall engine not available", http.StatusServiceUnavailable)
return
}
if err := sb.BlockSubnet(req.CIDR, req.Reason, dur); err != nil {
writeJSONError(w, fmt.Sprintf("Block failed: %v", err), http.StatusInternalServerError)
return
}
lifetime := "permanent"
if dur > 0 {
lifetime = dur.String()
}
s.auditLog(r, "firewall_deny_subnet", req.CIDR, lifetime+": "+req.Reason)
writeOK(w, map[string]interface{}{"cidr": req.CIDR})
}
// apiFirewallRemoveSubnet removes a subnet block.
func (s *Server) apiFirewallRemoveSubnet(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
CIDR string `json:"cidr"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.CIDR == "" {
writeJSONError(w, "CIDR is required", http.StatusBadRequest)
return
}
if _, _, err := net.ParseCIDR(req.CIDR); err != nil {
writeJSONError(w, "Invalid CIDR notation", http.StatusBadRequest)
return
}
sb, ok := s.blocker.(subnetUnblocker)
if !ok || sb == nil {
writeJSONError(w, "Firewall engine not available", http.StatusServiceUnavailable)
return
}
if err := sb.UnblockSubnet(req.CIDR); err != nil {
writeJSONError(w, fmt.Sprintf("Remove failed: %v", err), http.StatusInternalServerError)
return
}
s.auditLog(r, "firewall_remove_subnet", req.CIDR, "removed subnet block")
writeOK(w, map[string]interface{}{"cidr": req.CIDR})
}
// apiFirewallFlush clears all blocked IPs.
func (s *Server) apiFirewallFlush(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
fb, ok := s.blocker.(blockFlusher)
if !ok || fb == nil {
writeJSONError(w, "Firewall engine not available", http.StatusServiceUnavailable)
return
}
result, err := checks.FlushAutoBlockState(s.cfg.StatePath, fb.FlushBlocked)
if result.SnapshotErr != nil {
csmlog.Warn("web firewall flush could not snapshot persisted blocks", "err", result.SnapshotErr)
}
if err != nil {
if result.Flushed {
writeJSONError(w, fmt.Sprintf("Firewall flushed but auto-block cleanup failed: %v", err), http.StatusInternalServerError)
} else {
writeJSONError(w, fmt.Sprintf("Flush failed: %v", err), http.StatusInternalServerError)
}
return
}
s.auditLog(r, "firewall_flush", "blocked set", fmt.Sprintf("flushed: %v", result.Flushed))
writeOK(w, nil)
}
// apiFirewallFlushCphulk clears cPHulk login history for one IP without touching firewall state.
func (s *Server) apiFirewallFlushCphulk(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Audit, incident and threat records key on the canonical spelling.
req.IP = parsedIP.String()
if err := flushCphulk(req.IP); err != nil {
writeJSONError(w, "Could not clear cPHulk login history: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "cphulk_clear", req.IP, "cleared cPHulk login history")
writeOK(w, map[string]interface{}{"ip": req.IP})
}
// apiFirewallCheck checks if an IP is blocked in CSM or cphulk.
// GET /api/v1/firewall/check?ip=1.2.3.4
// phclient calls this route and reads success, permanent, temporary and
// cphulk (the cpanel-service shape it replaced), so the body keeps
// "success": true next to the fields; it is a deprecated alias:
//
// {"success": true, "ip": "1.2.3.4", "permanent": "reason or null", "temporary": "reason or null", "cphulk": true/false}
//
// Failures are error statuses like every other route; phclient treats any
// non-2xx as a failed call.
func (s *Server) apiFirewallCheck(w http.ResponseWriter, r *http.Request) {
ip := r.URL.Query().Get("ip")
if ip == "" || net.ParseIP(ip) == nil {
writeJSONError(w, "The ip is not valid or it was not set.", http.StatusBadRequest)
return
}
result := map[string]interface{}{
"success": true,
"ip": ip,
"permanent": nil,
"temporary": nil,
"cphulk": false,
}
// Check CSM firewall state
state, err := firewall.LoadState(s.cfg.StatePath)
if err != nil {
// A corrupt state file means we cannot tell whether the IP is
// blocked; report failure rather than a misleading "not blocked".
writeJSONError(w, "Firewall state unavailable", http.StatusInternalServerError)
return
}
now := time.Now()
// Compare as parsed addresses: an IPv6 block saved in one spelling
// must be found when queried in another.
queried := net.ParseIP(ip)
for _, b := range state.Blocked {
if entryIP := net.ParseIP(b.IP); entryIP != nil && entryIP.Equal(queried) {
if b.ExpiresAt.IsZero() {
result["permanent"] = b.Reason
} else if now.Before(b.ExpiresAt) {
// temporary keeps its text form for existing callers;
// expires_at is the instant.
result["temporary"] = fmt.Sprintf("%s (expires in %s)", b.Reason,
time.Until(b.ExpiresAt).Truncate(time.Minute))
result["expires_at"] = b.ExpiresAt.UTC()
}
}
}
// Check blocked subnets
parsedIP := net.ParseIP(ip)
for _, sn := range state.BlockedNet {
_, network, err := net.ParseCIDR(sn.CIDR)
if err == nil && network.Contains(parsedIP) {
result["permanent"] = fmt.Sprintf("Subnet block: %s - %s", sn.CIDR, sn.Reason)
}
}
// Check cphulk (cPanel brute force detector) - read-only check.
if platform.Detect().IsCPanel() && cphulkTempBanContainsIP(r.Context(), ip) {
result["cphulk"] = true
} else {
cphulkOut, cphulkErr := runFirewallCheckCommand(r.Context(), "whmapi1", "read_cphulk_records",
"list_name=black", "--output=json")
if cphulkErr == nil && cphulkBlocksIP(cphulkOut, ip) {
result["cphulk"] = true
}
}
writeJSON(w, result)
}
// cphulkBlocksIP scopes the match to cPHulk record IP fields. A raw token
// search would still match unrelated strings such as operator notes.
func cphulkBlocksIP(jsonOut []byte, ip string) bool {
if ip == "" {
return false
}
var payload struct {
Data struct {
Records []map[string]json.RawMessage `json:"records"`
} `json:"data"`
}
if err := json.Unmarshal(jsonOut, &payload); err != nil {
return false
}
for _, record := range payload.Data.Records {
for key, raw := range record {
if !isCphulkRecordIPField(key) {
continue
}
var value string
if err := json.Unmarshal(raw, &value); err != nil {
continue
}
if value == ip {
return true
}
}
}
return false
}
func isCphulkRecordIPField(key string) bool {
switch strings.ToLower(key) {
case "ip", "ip_address":
return true
default:
return false
}
}
func runFirewallCheckCommand(parent context.Context, name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(parent, cphulkFirewallCheckTimeout)
defer cancel()
return firewallCheckCommandOutput(ctx, name, args...)
}
// cPHulk brute-force temp bans live in the cphulk-TempBan nftables set, not in
// read_cphulk_records, so ask nftables for exact element membership instead of
// scanning a textual set dump where one IP could match unrelated text.
func cphulkTempBanContainsIP(parent context.Context, ip string) bool {
lookupIP, ok := cphulkTempBanLookupIP(ip)
if !ok {
return false
}
_, err := runFirewallCheckCommand(parent, "nft", "get", "element",
"inet", "filter", "cphulk-TempBan", "{", lookupIP, "}")
return err == nil
}
func cphulkTempBanLookupIP(ip string) (string, bool) {
parsed := net.ParseIP(strings.TrimSpace(ip))
if parsed == nil {
return "", false
}
if ip4 := parsed.To4(); ip4 != nil {
return net.IP(ip4).String(), true
}
return parsed.String(), true
}
// apiFirewallUnban unblocks an IP from CSM + cphulk in one call.
// POST /api/v1/firewall/unban body: {"ip": "1.2.3.4"}
func (s *Server) apiFirewallUnban(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "The ip is not valid or it was not set.", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Audit, incident and threat records key on the canonical spelling.
req.IP = parsedIP.String()
// 1. Unblock from CSM firewall (individual IP)
if s.blocker != nil {
_ = s.blocker.UnblockIP(req.IP)
}
dropAutoBlockThreatRow(req.IP)
// 2. Also remove from any covering subnet block. A corrupt state file
// only costs us the subnet sweep; the IP unblock above already ran, so
// skip this step rather than fail the whole unban.
state, stateErr := firewall.LoadState(s.cfg.StatePath)
subnetRemoved := ""
if sb, ok := s.blocker.(subnetUnblocker); ok && stateErr == nil && state != nil {
for _, sn := range state.BlockedNet {
_, network, err := net.ParseCIDR(sn.CIDR)
if err == nil && network.Contains(parsedIP) {
if sb.UnblockSubnet(sn.CIDR) == nil {
subnetRemoved = sn.CIDR
}
break
}
}
}
// 3. Flush from cphulk
_ = flushCphulk(req.IP) // best effort: the unblock is what was asked
// "success" is the deprecated alias phclient reads; see apiFirewallCheck.
result := map[string]interface{}{"ok": true, "success": true, "ip": req.IP}
if subnetRemoved != "" {
result["subnet_removed"] = subnetRemoved
}
details := "unblock, clear auto-block state, flush cPHulk"
if subnetRemoved != "" {
details += ", removed subnet " + subnetRemoved
}
s.auditLog(r, "firewall_unban", req.IP, details)
writeJSON(w, result)
}
package webui
import (
"context"
"encoding/json"
"fmt"
"net/http"
"os"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/firewall/rollback"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/obs"
)
// apiFirewallTentativeApply handles POST /api/v1/settings/firewall/tentative-apply.
// Body shape mirrors the regular settings POST plus an optional
// timeout_min field (1..30, default 5). The handler runs the same
// change validation as the normal save path, snapshots the previous
// csm.yaml bytes into bbolt, writes the new file, and triggers a
// daemon restart. The rollback manager arms an in-process timer; if
// the operator does not POST /confirm before the deadline the daemon
// restores the snapshot and restarts itself.
func (s *Server) apiFirewallTentativeApply(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
// This handler rewrites csm.yaml like the settings and verified-bots
// saves do; it shares their lock so a concurrent save cannot be lost.
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
mgr := rollback.Global()
if mgr == nil {
writeJSONError(w, "rollback manager not available", http.StatusServiceUnavailable)
return
}
if mgr.Status().Pending {
writeJSONError(w, "a firewall rollback is already pending; confirm or revert first", http.StatusConflict)
return
}
section, ok := LookupSettingsSection("firewall")
if !ok {
writeJSONError(w, "firewall section not registered", http.StatusInternalServerError)
return
}
ifMatch := r.Header.Get("If-Match")
if ifMatch == "" {
writeJSONError(w, "If-Match header required", http.StatusBadRequest)
return
}
var body struct {
Changes map[string]json.RawMessage `json:"changes"`
TimeoutMin int `json:"timeout_min"`
}
if err := decodeJSONBodyLimited(w, r, 256*1024, &body); err != nil {
writeJSONError(w, "invalid body: "+err.Error(), http.StatusBadRequest)
return
}
diskBytes, err := os.ReadFile(s.cfg.ConfigFile) // #nosec G304 -- operator-supplied config path
if err != nil {
writeJSONError(w, "read config: "+err.Error(), http.StatusInternalServerError)
return
}
disk, err := config.LoadBytes(diskBytes)
if err != nil {
writeJSONError(w, "parse config: "+err.Error(), http.StatusInternalServerError)
return
}
disk.ConfigFile = s.cfg.ConfigFile
disk.ConfigDir = s.cfg.ConfigDir
if disk.Integrity.ConfigHash != ifMatch {
writeJSONError(w, "config changed on disk, reload", http.StatusPreconditionFailed)
return
}
if rejectIfConfDirChanged(w, s.cfg.ConfigDir, disk) {
return
}
clone := *disk
if disk.Firewall != nil {
fw := *disk.Firewall
clone.Firewall = &fw
}
yamlChanges, errs := buildChangeSet(section, &clone, body.Changes)
if len(errs) > 0 {
writeValidationErrors(w, errs)
return
}
validationResults := append(config.Validate(&clone), config.ValidateDeepSection(&clone, section.ID)...)
fieldErrors, warnings := splitValidationResults(validationResults)
if len(fieldErrors) > 0 {
writeValidationErrors(w, fieldErrors)
return
}
localizeValidationFields(warnings, "firewall")
edited, err := config.YAMLEdit(diskBytes, yamlChanges)
if err != nil {
writeJSONError(w, "yaml edit: "+err.Error(), http.StatusInternalServerError)
return
}
// Stage rollback BEFORE the on-disk write so a crash between the two
// leaves the snapshot recoverable: if there is no new file on disk
// yet, the snapshot is a no-op revert. The reverse order would let
// the daemon come back to the new config with no rollback record and
// no way to undo without operator intervention.
timeout := time.Duration(body.TimeoutMin) * time.Minute
st, err := mgr.Apply(diskBytes, edited, timeout, extractClientIP(r))
if err != nil {
writeJSONError(w, "stage rollback: "+err.Error(), http.StatusInternalServerError)
return
}
if err := integrity.SignAndSavePreserving(s.cfg.ConfigFile, s.cfg.ConfigDir, edited, &clone, disk.Integrity.BinaryHash); err != nil {
// Best-effort cleanup: the snapshot is now misleading because
// the on-disk file never changed. Drop it so the operator does
// not see a phantom pending rollback in the UI.
_ = mgr.AbortApplyIfCurrent(st)
writeJSONError(w, "save: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "settings-tentative-apply", "firewall", auditDetailsFor(section, body.Changes))
// Defer the restart so the response flushes first; otherwise the
// client sees a connection reset and cannot read the rollback ETA
// it needs to drive the countdown banner.
s.scheduleDaemonRestart(250 * time.Millisecond)
if warnings == nil {
warnings = []fieldError{}
}
applied := changedFieldList(body.Changes, section)
// The daemon restarts to apply a firewall change, so every applied field
// waits for it; the shape matches the settings save response.
writeOK(w, map[string]interface{}{
"warnings": warnings,
"rollback": st,
"new_etag": clone.Integrity.ConfigHash,
"applied": applied,
"requires_restart": applied,
"pending_restart": true,
})
}
// changedFieldList returns the dotted YAML paths of the keys in changes
// scoped to the section. Used in the response so the UI knows which
// fields to highlight as pending.
func changedFieldList(changes map[string]json.RawMessage, section SettingsSection) []string {
out := make([]string, 0, len(changes))
for k := range changes {
if k == "" {
out = append(out, section.YAMLPath)
continue
}
out = append(out, section.YAMLPath+"."+k)
}
return out
}
// apiFirewallRollbackStatus returns the pending rollback record, if any.
// Read endpoint, no CSRF; safe to poll for the countdown banner.
func (s *Server) apiFirewallRollbackStatus(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
mgr := rollback.Global()
if mgr == nil {
writeJSON(w, rollback.Status{})
return
}
writeJSON(w, mgr.Status())
}
func rejectConfigWriteDuringRollback(w http.ResponseWriter) bool {
mgr := rollback.Global()
if mgr == nil || !mgr.Status().Pending {
return false
}
writeJSONError(w, "a firewall rollback is pending; confirm or revert it before saving other settings", http.StatusConflict)
return true
}
// apiFirewallRollbackConfirm handles POST .../confirm. Drops the snapshot;
// the new config stays.
func (s *Server) apiFirewallRollbackConfirm(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
mgr := rollback.Global()
if mgr == nil {
writeJSONError(w, "rollback manager not available", http.StatusServiceUnavailable)
return
}
st := mgr.Status()
if !st.Pending {
writeJSONError(w, "no pending rollback", http.StatusConflict)
return
}
if err := mgr.ConfirmIfCurrent(st); err != nil {
writeJSONError(w, "confirm: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "settings-rollback-confirm", "firewall", "")
writeOK(w, nil)
}
// apiFirewallRollbackRevert handles POST .../revert. Restores the
// snapshot to disk and triggers a daemon restart. Returns 200 with the
// pre-revert status; the actual restart happens on a goroutine so the
// response can flush.
func (s *Server) apiFirewallRollbackRevert(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
mgr := rollback.Global()
if mgr == nil {
writeJSONError(w, "rollback manager not available", http.StatusServiceUnavailable)
return
}
st := mgr.Status()
if !st.Pending {
writeJSONError(w, "no pending rollback", http.StatusConflict)
return
}
s.scheduleRollbackRevert(mgr, st, 30*time.Second)
s.auditLog(r, "settings-rollback-revert", "firewall", "")
// The revert and restart run after the response.
writeOKStatus(w, http.StatusAccepted, nil)
}
// scheduleDaemonRestart fires restartDaemon after delay in a supervised
// goroutine. The select on pruneDone lets the goroutine exit cleanly if
// the server begins shutdown during the pre-restart delay, so an
// operator-initiated stop is not chased by a phantom restart.
func (s *Server) scheduleDaemonRestart(delay time.Duration) {
obs.SafeGo("webui-daemon-restart", func() {
select {
case <-s.pruneDone:
return
case <-time.After(delay):
}
if _, err := s.restartDaemon(); err != nil {
fmt.Fprintf(os.Stderr, "webui: daemon restart failed: %v\n", err)
}
})
}
// scheduleRollbackRevert runs the revert in a supervised goroutine with
// a hard timeout. If shutdown already started before the worker runs, it
// does not begin a new revert; once started, the revert owns its restart
// context so the restart it triggers cannot cancel itself via Shutdown.
func (s *Server) scheduleRollbackRevert(mgr *rollback.Manager, expected rollback.Status, timeout time.Duration) {
obs.SafeGo("webui-rollback-revert", func() {
select {
case <-s.pruneDone:
return
default:
}
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
if err := mgr.RevertIfCurrent(ctx, expected); err != nil {
fmt.Fprintf(os.Stderr, "webui: rollback revert failed: %v\n", err)
}
})
}
package webui
import (
"net/http"
"sort"
"github.com/pidginhost/csm/internal/mailfwd/inventory"
"github.com/pidginhost/csm/internal/platform"
)
// forwarderDestination is one resolved target of a forwarder, as served to the
// UI. Provider is the inventory class string (local/yahoo/gmail/outlook/external)
// the table renders as a badge.
type forwarderDestination struct {
Address string `json:"address"`
Domain string `json:"domain"`
Provider string `json:"provider"`
}
// forwarderEntry is a single source address and everything it relays to.
type forwarderEntry struct {
Source string `json:"source"`
Domain string `json:"domain"`
Owner string `json:"owner"`
Destinations []forwarderDestination `json:"destinations"`
Providers []string `json:"providers"` // distinct destination classes, sorted
KeepLocal bool `json:"keep_local"`
ForwardOnly bool `json:"forward_only"`
HasExternal bool `json:"has_external"`
HasFreeProvider bool `json:"has_free_provider"`
}
// forwardersSummary is the page-header rollup: how many forwarders exist and
// how many carry reputation risk (leave the server / target a free provider).
type forwardersSummary struct {
Total int `json:"total"`
External int `json:"external"`
FreeProvider int `json:"free_provider"`
}
type forwardersResponse struct {
Forwarders []forwarderEntry `json:"items"`
Total int `json:"total"`
Summary forwardersSummary `json:"summary"`
}
// selectForwarderSource picks the inventory source for the host. Only cPanel
// enumeration is wired; other platforms get the empty source until their
// adapters land (Phase 3).
func selectForwarderSource() inventory.Source {
if platform.Detect().IsCPanel() {
return inventory.NewCPanelSource()
}
return inventory.EmptySource{}
}
// apiEmailForwarders handles GET /api/v1/email/forwarders and returns the host's
// forwarder inventory: each source, its destinations with provider class, owner,
// and whether it keeps a local copy or forwards only.
func (s *Server) apiEmailForwarders(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
resp := forwardersResponse{Forwarders: []forwarderEntry{}}
if s.forwarderSource == nil {
writeJSON(w, resp)
return
}
fwds, err := s.forwarderSource.Forwarders()
if err != nil {
writeJSONError(w, "Failed to enumerate forwarders", http.StatusInternalServerError)
return
}
for _, f := range fwds {
resp.Forwarders = append(resp.Forwarders, toForwarderEntry(f))
resp.Summary.Total++
resp.Total++
if f.HasExternal() {
resp.Summary.External++
}
if f.HasFreeProvider() {
resp.Summary.FreeProvider++
}
}
writeJSON(w, resp)
}
func toForwarderEntry(f inventory.Forwarder) forwarderEntry {
dests := make([]forwarderDestination, 0, len(f.Destinations))
seen := make(map[string]bool, len(f.Destinations))
providers := make([]string, 0, len(f.Destinations))
for _, d := range f.Destinations {
dests = append(dests, forwarderDestination{
Address: d.Address,
Domain: d.Domain,
Provider: string(d.Provider),
})
if p := string(d.Provider); !seen[p] {
seen[p] = true
providers = append(providers, p)
}
}
sort.Strings(providers)
return forwarderEntry{
Source: f.Source,
Domain: f.Domain,
Owner: f.Owner,
Destinations: dests,
Providers: providers,
KeepLocal: f.KeepLocal,
ForwardOnly: f.ForwardOnly,
HasExternal: f.HasExternal(),
HasFreeProvider: f.HasFreeProvider(),
}
}
package webui
import (
"net"
"net/http"
"github.com/pidginhost/csm/internal/geoip"
)
// SetGeoIPDB sets the GeoIP database for IP lookups.
func (s *Server) SetGeoIPDB(db *geoip.DB) {
s.geoIPDB.Store(db)
}
// apiGeoIPLookup returns geolocation info for an IP.
// GET /api/v1/geoip?ip=1.2.3.4 - fast local lookup (country + ASN)
// GET /api/v1/geoip?ip=1.2.3.4&detail=1 - includes RDAP org/ISP (may take 1-3s)
func (s *Server) apiGeoIPLookup(w http.ResponseWriter, r *http.Request) {
ip := r.URL.Query().Get("ip")
if ip == "" {
writeJSONError(w, "ip parameter required", http.StatusBadRequest)
return
}
if net.ParseIP(ip) == nil {
writeJSONError(w, "invalid IP address", http.StatusBadRequest)
return
}
db := s.geoIPDB.Load()
if db == nil {
writeJSONError(w, "GeoIP databases not loaded", http.StatusServiceUnavailable)
return
}
var info geoip.Info
if r.URL.Query().Get("detail") == "1" {
info = db.LookupWithRDAP(ip)
} else {
info = db.Lookup(ip)
}
writeJSON(w, info)
}
// apiGeoIPBatch returns geolocation info for multiple IPs.
// POST /api/v1/geoip/batch body: {"ips": ["1.2.3.4", "5.6.7.8"]}
func (s *Server) apiGeoIPBatch(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IPs []string `json:"ips"`
}
if err := decodeJSONBodyLimited(w, r, 32*1024, &req); err != nil {
writeJSONError(w, "invalid request body", http.StatusBadRequest)
return
}
if len(req.IPs) > 500 {
writeJSONError(w, "maximum 500 IPs per request", http.StatusBadRequest)
return
}
type geoResult struct {
Country string `json:"country"`
CountryName string `json:"country_name"`
City string `json:"city"`
ASOrg string `json:"as_org"`
Error string `json:"error,omitempty"`
}
db := s.geoIPDB.Load()
results := make(map[string]geoResult, len(req.IPs))
for _, ip := range req.IPs {
if net.ParseIP(ip) == nil {
results[ip] = geoResult{Error: "invalid IP format"}
continue
}
if db == nil {
results[ip] = geoResult{Error: "GeoIP database not loaded"}
continue
}
info := db.Lookup(ip)
results[ip] = geoResult{
Country: info.Country,
CountryName: info.CountryName,
City: info.City,
ASOrg: info.ASOrg,
}
}
writeJSON(w, map[string]interface{}{"results": results})
}
package webui
import (
"bytes"
"fmt"
"html/template"
"net/http"
"os"
"time"
"github.com/pidginhost/csm/internal/checks"
)
func (s *Server) renderTemplate(w http.ResponseWriter, r *http.Request, name string, data interface{}) {
base := s.templates[name]
if base == nil {
fmt.Fprintf(os.Stderr, "[webui] template %s missing\n", name)
http.Error(w, "template not found", http.StatusInternalServerError)
return
}
// The CSRF token belongs to the browser session loading the page, so
// each render binds it on a clone of the parsed template.
tmpl, err := base.Clone()
if err != nil {
fmt.Fprintf(os.Stderr, "[webui] template %s clone error: %v\n", name, err)
http.Error(w, "template render error", http.StatusInternalServerError)
return
}
token := s.csrfTokenFor(r)
tmpl.Funcs(template.FuncMap{"csrfToken": func() string { return token }})
// Render into a buffer first so an execution error can still surface as a
// 500 — html/template streams directly to its writer, and once any byte
// has been flushed the status header is locked in.
var buf bytes.Buffer
if err := tmpl.ExecuteTemplate(&buf, name, data); err != nil {
fmt.Fprintf(os.Stderr, "[webui] template %s error: %v\n", name, err)
http.Error(w, "template render error", http.StatusInternalServerError)
return
}
if _, err := w.Write(buf.Bytes()); err != nil {
fmt.Fprintf(os.Stderr, "[webui] template %s write error: %v\n", name, err)
}
}
type dashboardData struct {
Hostname string
Uptime string
Critical int
High int
Warning int
Total int
SigCount int
FanotifyActive bool
LogWatchers int
LastCriticalAgo string
LastCriticalISO string // RFC3339 of most recent critical, "" if none (so relative time can tick client-side)
RecentFindings []historyEntry
}
type historyEntry struct {
Severity string
SevClass string
Check string
Message string
Details string
Timestamp string
TimestampISO string // RFC3339 for JS comparison
TimeAgo string
HasFix bool
FixDesc string
Key string // canonical dedup key (matches alert.Finding.Key())
}
type quarantineData struct {
Hostname string
Files []quarantineEntry
}
type quarantineEntry struct {
ID string
OriginalPath string
Size int64
QuarantineAt string
Reason string
}
func (s *Server) handleDashboard(w http.ResponseWriter, r *http.Request) {
sum := s.statsSummary24h()
recent := make([]historyEntry, 0, len(sum.recent))
for _, f := range sum.recent {
recent = append(recent, historyEntry{
Severity: severityLabel(f.Severity),
SevClass: severityClass(f.Severity),
Check: f.Check,
Message: f.Message,
Details: f.Details,
Timestamp: f.Timestamp.Format("15:04:05"),
TimestampISO: f.Timestamp.Format(time.RFC3339),
TimeAgo: timeAgo(f.Timestamp),
HasFix: checks.HasFix(f.Check),
FixDesc: checks.FixDescription(f.Check, f.Message, f.FilePath),
Key: f.Key(),
})
}
lastCriticalAgo, lastCriticalISO := "None", ""
if !sum.lastCritical.IsZero() {
lastCriticalAgo = timeAgo(sum.lastCritical)
lastCriticalISO = sum.lastCritical.Format(time.RFC3339)
}
data := dashboardData{
Hostname: s.cfg.Hostname,
Uptime: time.Since(s.startTime).Round(time.Second).String(),
Critical: sum.critical,
High: sum.high,
Warning: sum.warning,
Total: sum.critical + sum.high + sum.warning,
SigCount: s.signatureCount(),
FanotifyActive: s.fanotifyRunning(),
LogWatchers: s.logWatchersRunning(),
LastCriticalAgo: lastCriticalAgo,
LastCriticalISO: lastCriticalISO,
RecentFindings: recent,
}
s.renderTemplate(w, r, "dashboard.html", data)
}
func (s *Server) handleFindings(w http.ResponseWriter, r *http.Request) {
// Findings page is now JS-driven - enriched API provides data
s.renderTemplate(w, r, "findings.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
func (s *Server) handleHistoryRedirect(w http.ResponseWriter, r *http.Request) {
// History is now a tab on the findings page - redirect for backward compat
target := "/findings?tab=history"
if qs := r.URL.RawQuery; qs != "" {
target = "/findings?tab=history&" + qs
}
// #nosec G710 -- target always starts with the fixed same-origin
// /findings path; the incoming query can only add parameters.
http.Redirect(w, r, target, http.StatusFound)
}
// handleBlockedRedirect sends the Firewall page's old address to the page.
func (s *Server) handleBlockedRedirect(w http.ResponseWriter, r *http.Request) {
target := "/firewall"
if qs := r.URL.RawQuery; qs != "" {
target += "?" + qs
}
// #nosec G710 -- target always starts with the fixed same-origin
// /firewall path; the incoming query can only add parameters.
http.Redirect(w, r, target, http.StatusFound)
}
func (s *Server) handleQuarantine(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "quarantine.html", quarantineData{
Hostname: s.cfg.Hostname,
})
}
func (s *Server) handleCleanupHistory(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "cleanup-history.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
func (s *Server) handleFirewall(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "firewall.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
func (s *Server) handleEmail(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "email.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
func (s *Server) handleSettings(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "settings.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
package webui
import (
"net/http"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/store"
)
// apiHardening returns the last stored audit report (GET).
func (s *Server) apiHardening(w http.ResponseWriter, _ *http.Request) {
db := store.Global()
if db == nil {
writeJSON(w, &store.AuditReport{})
return
}
report, err := db.LoadHardeningReport()
if err != nil {
writeJSONError(w, "failed to load report: "+err.Error(), http.StatusInternalServerError)
return
}
writeJSON(w, report)
}
// apiHardeningRun runs the audit, stores the result, and returns it (POST only).
func (s *Server) apiHardeningRun(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if !s.acquireScan() {
writeJSONError(w, "A scan is already in progress. Please wait.", http.StatusConflict)
return
}
defer s.releaseScan()
// The audit can outlast the server's WriteTimeout; extend, never shorten.
rc := http.NewResponseController(w)
_ = rc.SetWriteDeadline(time.Now().Add(longRequestTimeout))
report := checks.RunHardeningAudit(s.liveCfg())
s.auditLog(r, "hardening_run", "server", "hardening audit run")
if db := store.Global(); db != nil {
if err := db.SaveHardeningReport(report); err != nil {
writeJSONError(w, "audit completed but failed to save report: "+err.Error(), http.StatusInternalServerError)
return
}
}
// The fresh report, with the action's ok flag next to its fields.
writeJSON(w, struct {
OK bool `json:"ok"`
*store.AuditReport
}{true, report})
}
// handleHardening renders the hardening audit page.
func (s *Server) handleHardening(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "hardening.html", nil)
}
package webui
import (
"net/http"
"regexp"
"strings"
"github.com/pidginhost/csm/internal/mailfwd/quarantine"
"github.com/pidginhost/csm/internal/platform"
)
// heldIDRe bounds a held-message id to a Maildir filename shape. The store also
// defends (filepath.Base + regular-file check), but validating at the handler
// boundary rejects traversal/control input before it reaches the filesystem.
var heldIDRe = regexp.MustCompile(`^[A-Za-z0-9._-]{1,128}$`)
// heldForwardStore is the held-forward quarantine surface the webui needs.
// *quarantine.Quarantine satisfies it; tests use a fake.
type heldForwardStore interface {
List() ([]quarantine.HeldMessage, error)
Release(id string) error
Delete(id string) error
}
const forwardQuarantineDir = "/var/lib/csm/forward_quarantine/held"
// selectForwardHeld returns the held-forward store for the host. Only
// cPanel/exim writes held copies; other platforms have none.
func selectForwardHeld() heldForwardStore {
if platform.Detect().IsCPanel() {
return quarantine.New(forwardQuarantineDir)
}
return nil
}
// apiEmailHeldList handles GET /api/v1/email/held and returns the forward
// copies the guard has held.
func (s *Server) apiEmailHeldList(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if s.forwardHeld == nil {
writeAll(w, []quarantine.HeldMessage{})
return
}
msgs, err := s.forwardHeld.List()
if err != nil {
writeJSONError(w, "Failed to list held forwards", http.StatusInternalServerError)
return
}
writeAll(w, msgs)
}
// apiEmailHeldAction handles POST /api/v1/email/held/{id}/release (re-inject the
// held copy to its external recipient) and DELETE /api/v1/email/held/{id}
// (discard). Both mutate, so they run under auth + CSRF and are audit-logged.
func (s *Server) apiEmailHeldAction(w http.ResponseWriter, r *http.Request) {
tail := strings.TrimPrefix(r.URL.Path, "/api/v1/email/held/")
if tail == "" {
writeJSONError(w, "Missing held message ID", http.StatusBadRequest)
return
}
parts := strings.SplitN(tail, "/", 2)
id := parts[0]
action := ""
if len(parts) == 2 {
action = parts[1]
}
if !heldIDRe.MatchString(id) || strings.Contains(id, "..") {
writeJSONError(w, "Invalid held message ID", http.StatusBadRequest)
return
}
if s.forwardHeld == nil {
writeJSONError(w, "Forward guard not available on this host", http.StatusServiceUnavailable)
return
}
switch r.Method {
case http.MethodPost:
if action != "release" {
writeJSONError(w, "Unknown action; use /release", http.StatusBadRequest)
return
}
if err := s.forwardHeld.Release(id); err != nil {
writeJSONError(w, "Failed to release held forward: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "email_held_release", id, "re-injected held forward copy to its external recipient")
writeOK(w, map[string]interface{}{"id": id})
case http.MethodDelete:
if action != "" {
writeJSONError(w, "Unknown action", http.StatusBadRequest)
return
}
if err := s.forwardHeld.Delete(id); err != nil {
writeJSONError(w, "Failed to delete held forward: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "email_held_delete", id, "deleted held forward copy")
writeOK(w, map[string]interface{}{"id": id})
default:
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
}
}
package webui
import (
"bytes"
"encoding/json"
"fmt"
"html/template"
"net"
"net/http"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/platform"
)
// jsonForScript marshals v to JSON and returns it as template.JS suitable
// for direct substitution into a <script> block. json.Marshal already
// escapes < > & U+2028 U+2029 to \uXXXX form, so the output cannot break
// out of the surrounding <script> tag or trigger JS line-terminator
// parsing quirks. On marshal failure the fallback is the JS literal
// "null" so the enclosing template still parses.
func jsonForScript(v interface{}) template.JS {
b, err := json.Marshal(v)
if err != nil {
return template.JS("null")
}
// Defense-in-depth: Go's json.Marshal has historically escaped these
// by default, but an explicit pass guarantees the contract even if
// that default ever changes or the input arrived pre-encoded.
b = bytes.ReplaceAll(b, []byte("<"), []byte(`\u003c`))
b = bytes.ReplaceAll(b, []byte(">"), []byte(`\u003e`))
b = bytes.ReplaceAll(b, []byte("&"), []byte(`\u0026`))
b = bytes.ReplaceAll(b, []byte("\u2028"), []byte(`\u2028`))
b = bytes.ReplaceAll(b, []byte("\u2029"), []byte(`\u2029`))
// #nosec G203 -- Output is JSON bytes with HTML/JS-dangerous codepoints
// escaped above; safe to hand to html/template as JS.
return template.JS(b)
}
func decodeJSONBodyLimited(w http.ResponseWriter, r *http.Request, limit int64, dst interface{}) error {
if limit <= 0 {
limit = 64 * 1024
}
r.Body = http.MaxBytesReader(w, r.Body, limit)
dec := json.NewDecoder(r.Body)
dec.DisallowUnknownFields()
if err := dec.Decode(dst); err != nil {
return err
}
if dec.More() {
return fmt.Errorf("request body must contain a single JSON value")
}
return nil
}
// validateEximMessageID rejects message ids that would be unsafe to
// interpolate into filesystem paths. Real Exim ids are of the form
// 6-6-2 on Exim 4.96 and older, and 6-11-4 on Exim 4.97 and newer. The
// validator accepts anything with the same character class so callers can use
// shorter fixtures in tests while still blocking `..`, `/`, `\`, dots, NUL
// bytes, and other shell / path-traversal metacharacters even if the inner
// Quarantine layer's filepath.Base() guard regresses.
func validateEximMessageID(id string) error {
if id == "" {
return fmt.Errorf("message id is required")
}
if len(id) > 32 {
return fmt.Errorf("message id too long (%d chars, max 32)", len(id))
}
for _, c := range id {
if (c < 'a' || c > 'z') && (c < 'A' || c > 'Z') && (c < '0' || c > '9') && c != '-' {
return fmt.Errorf("message id contains invalid character: %q", c)
}
}
return nil
}
// mustBeWithin resolves candidate against root and reports whether the
// result still lives inside root after symlink and `..` resolution.
// Returns the cleaned absolute path on success. Defense-in-depth for any
// handler that takes a user-supplied path fragment, joins it under a
// trusted root, and hands the result to os.Remove / os.RemoveAll / etc.
func mustBeWithin(root, candidate string) (string, error) {
if root == "" {
return "", fmt.Errorf("root is required")
}
absRoot, err := filepath.Abs(filepath.Clean(root))
if err != nil {
return "", fmt.Errorf("resolve root: %w", err)
}
absRoot, err = filepath.EvalSymlinks(absRoot)
if err != nil {
return "", fmt.Errorf("resolve root symlinks: %w", err)
}
abs, err := resolvePathUnderRoot(absRoot, rootRelativeCandidate(candidate))
if err != nil {
return "", err
}
if !isPathWithin(abs, absRoot) {
return "", fmt.Errorf("path %q escapes root %q", candidate, absRoot)
}
return abs, nil
}
func rootRelativeCandidate(candidate string) string {
if volume := filepath.VolumeName(candidate); volume != "" {
candidate = strings.TrimPrefix(candidate, volume)
}
candidate = strings.TrimLeft(candidate, string(filepath.Separator))
if candidate == "" {
return "."
}
return candidate
}
func resolvePathUnderRoot(root, rel string) (string, error) {
parts := pathParts(rel)
current := root
symlinkCount := 0
for i := 0; i < len(parts); i++ {
part := parts[i]
switch part {
case ".", "":
continue
case "..":
current = filepath.Dir(current)
if !isPathWithin(current, root) {
return "", fmt.Errorf("path %q escapes root %q", rel, root)
}
continue
}
next := filepath.Join(current, part)
info, err := os.Lstat(next)
if err != nil {
if os.IsNotExist(err) {
current = next
continue
}
return "", fmt.Errorf("stat candidate: %w", err)
}
if info.Mode()&os.ModeSymlink == 0 {
current = next
continue
}
symlinkCount++
if symlinkCount > 255 {
return "", fmt.Errorf("too many symlinks resolving %q", rel)
}
target, err := os.Readlink(next)
if err != nil {
return "", fmt.Errorf("read symlink: %w", err)
}
if !filepath.IsAbs(target) {
target = filepath.Join(filepath.Dir(next), target)
}
target = filepath.Clean(target)
targetRel, err := filepath.Rel(root, target)
if err != nil {
return "", fmt.Errorf("resolve symlink target: %w", err)
}
nextParts := append([]string{}, pathParts(targetRel)...)
nextParts = append(nextParts, parts[i+1:]...)
parts = nextParts
current = root
i = -1
}
return filepath.Clean(current), nil
}
func pathParts(path string) []string {
if path == "" || path == "." {
return nil
}
raw := strings.Split(path, string(filepath.Separator))
parts := make([]string, 0, len(raw))
for _, part := range raw {
if part != "" && part != "." {
parts = append(parts, part)
}
}
return parts
}
// validateAccountName checks that name is a valid cPanel account name:
// 1-64 characters, alphanumeric and underscore only.
func validateAccountName(name string) error {
if name == "" {
return fmt.Errorf("account name is required")
}
if len(name) > 64 {
return fmt.Errorf("account name too long (%d chars, max 64)", len(name))
}
if (name[0] < 'a' || name[0] > 'z') && (name[0] < 'A' || name[0] > 'Z') {
return fmt.Errorf("account name must start with a letter")
}
for _, c := range name {
if (c < 'a' || c > 'z') && (c < 'A' || c > 'Z') && (c < '0' || c > '9') && c != '_' {
return fmt.Errorf("account name contains invalid character: %c", c)
}
}
return nil
}
// parseAndValidateIP parses an IP string and rejects non-routable addresses
// (loopback, private RFC 1918, link-local, multicast, unspecified, broadcast).
// RFC 5737 documentation ranges (192.0.2.0/24, 198.51.100.0/24, 203.0.113.0/24)
// are intentionally allowed.
func parseAndValidateIP(s string) (net.IP, error) {
s = strings.TrimSpace(s)
if s == "" {
return nil, fmt.Errorf("IP address is required")
}
ip := net.ParseIP(s)
if ip == nil {
return nil, fmt.Errorf("invalid IP address: %s", s)
}
if ip.IsLoopback() {
return nil, fmt.Errorf("loopback address not allowed: %s", s)
}
if ip.IsPrivate() {
return nil, fmt.Errorf("private address not allowed: %s", s)
}
if ip.IsLinkLocalUnicast() || ip.IsLinkLocalMulticast() {
return nil, fmt.Errorf("link-local address not allowed: %s", s)
}
if ip.IsMulticast() {
return nil, fmt.Errorf("multicast address not allowed: %s", s)
}
if ip.IsUnspecified() {
return nil, fmt.Errorf("unspecified address not allowed: %s", s)
}
// Broadcast: 255.255.255.255
if ip.Equal(net.IPv4bcast) {
return nil, fmt.Errorf("broadcast address not allowed: %s", s)
}
return ip, nil
}
// validateCIDR parses a CIDR string and rejects overly broad prefixes
// (/0 through /7; minimum allowed is /8).
func validateCIDR(s string) (*net.IPNet, error) {
if s == "" {
return nil, fmt.Errorf("CIDR is required")
}
_, ipNet, err := net.ParseCIDR(s)
if err != nil {
return nil, fmt.Errorf("invalid CIDR: %w", err)
}
ones, bits := ipNet.Mask.Size()
minPrefix := 8
if bits == 128 { // IPv6
minPrefix = 32
}
if ones < minPrefix {
return nil, fmt.Errorf("CIDR prefix /%d is too broad (minimum /%d)", ones, minPrefix)
}
return ipNet, nil
}
// parseDuration parses a human-friendly duration string from the web UI.
// Supported formats: "24h", "7d", "30d", "0" (permanent), "" (permanent).
// Zero means permanent in the block/allow contexts that consume this.
// Unparseable, negative, or out-of-range input is an error to keep a typo or
// integer overflow from silently becoming a permanent rule.
func parseDuration(s string) (time.Duration, error) {
s = strings.TrimSpace(s)
if s == "" || s == "0" {
return 0, nil
}
if strings.HasPrefix(s, "-") {
return 0, fmt.Errorf("invalid duration %q", s)
}
if strings.HasSuffix(s, "d") {
num := strings.TrimSuffix(s, "d")
if num == "" {
return 0, fmt.Errorf("invalid duration %q", s)
}
for _, c := range num {
if c < '0' || c > '9' {
return 0, fmt.Errorf("invalid duration %q", s)
}
}
days, err := strconv.ParseUint(num, 10, 64)
const maxDuration = time.Duration(1<<63 - 1)
if err != nil || days > uint64(maxDuration/(24*time.Hour)) {
return 0, fmt.Errorf("invalid duration %q", s)
}
return time.Duration(days) * 24 * time.Hour, nil
}
d, err := time.ParseDuration(s)
if err != nil || d == 0 && strings.ContainsAny(s, "123456789") {
return 0, fmt.Errorf("invalid duration %q", s)
}
return d, nil
}
// operatorFacingCheck reports whether findings of a check belong in the
// finding lists operators act on. auto_response and auto_block record what
// CSM already did; check_timeout and health describe CSM itself.
func operatorFacingCheck(check string) bool {
switch check {
case "auto_response", "auto_block", "check_timeout", "health":
return false
}
return true
}
// isPathUnder returns true if the cleaned path is strictly under the base
// directory. It prevents path traversal via ".." and prefix tricks
// (e.g., /home/username is not under /home/user).
func isPathUnder(path, base string) bool {
cleanPath := filepath.Clean(path)
cleanBase := filepath.Clean(base)
// Ensure base ends with separator so "/home/user" doesn't match "/home/username"
prefix := cleanBase + string(filepath.Separator)
return strings.HasPrefix(cleanPath, prefix)
}
func isPathWithin(path, base string) bool {
cleanPath := filepath.Clean(path)
cleanBase := filepath.Clean(base)
return cleanPath == cleanBase || strings.HasPrefix(cleanPath, cleanBase+string(filepath.Separator))
}
// hashLiveStateFile hashes a file for quarantineLiveState; a var so tests
// can count reads.
var hashLiveStateFile = integrity.HashFile
// liveStates remembers how a quarantine archive compared with its live file,
// keyed by both files' change keys, so the list does not hash every pair on
// every request. It is cleared when it grows past liveStatesMax entries.
var (
liveStatesMu sync.Mutex
liveStates = map[string]liveStateResult{}
)
const liveStatesMax = 10000
type liveStateResult struct {
keys string
state string
}
func quarantineLiveState(archivePath, originalPath string) string {
origInfo, err := os.Stat(originalPath)
if os.IsNotExist(err) {
return "original_missing"
}
if err != nil {
return "unknown"
}
if !origInfo.Mode().IsRegular() {
return "original_not_file"
}
archInfo, err := os.Stat(archivePath)
if os.IsNotExist(err) {
return "archive_missing"
}
if err != nil {
return "unknown"
}
if !archInfo.Mode().IsRegular() {
return "archive_not_file"
}
if origInfo.Size() != archInfo.Size() {
return "live_differs"
}
pair := archivePath + "\x00" + originalPath
keys := integrity.FileChangeKey(origInfo) + "|" + integrity.FileChangeKey(archInfo)
liveStatesMu.Lock()
cached, ok := liveStates[pair]
liveStatesMu.Unlock()
if ok && cached.keys == keys {
return cached.state
}
origHash, err := hashLiveStateFile(originalPath)
if err != nil {
return "unknown"
}
archHash, err := hashLiveStateFile(archivePath)
if err != nil {
return "unknown"
}
state := "live_differs"
if origHash == archHash {
state = "restored_identical"
}
liveStatesMu.Lock()
if len(liveStates) >= liveStatesMax {
liveStates = map[string]liveStateResult{}
}
liveStates[pair] = liveStateResult{keys: keys, state: state}
liveStatesMu.Unlock()
return state
}
const preCleanQuarantineIDPrefix = "pre_clean:"
type quarantineEntryRef struct {
ID string
ItemPath string
MetaPath string
}
func quarantineEntryID(metaPath string) string {
id := strings.TrimSuffix(filepath.Base(metaPath), ".meta")
if filepath.Clean(filepath.Dir(metaPath)) == filepath.Join(quarantineDir, "pre_clean") {
return preCleanQuarantineIDPrefix + id
}
return id
}
// reservedQuarantineNames are subtrees of the quarantine root that are not
// entries: the pre-clean backups and the email quarantine. An ID naming one
// of them resolves to the subtree itself and must never be deleted.
var reservedQuarantineNames = map[string]bool{"pre_clean": true, "email": true}
// quarantineEntryDeletable reports whether an ID resolved to a real
// quarantine entry: not a reserved subtree, and carrying the metadata
// sidecar every quarantined item is written with.
func quarantineEntryDeletable(entry quarantineEntryRef) bool {
if reservedQuarantineNames[filepath.Base(entry.ItemPath)] {
return false
}
_, err := os.Stat(entry.MetaPath)
return err == nil
}
func resolveQuarantineEntry(id string) (quarantineEntryRef, error) {
rawID := strings.TrimSpace(id)
if rawID == "" {
return quarantineEntryRef{}, fmt.Errorf("quarantine ID is required")
}
baseDir := quarantineDir
name := rawID
if strings.HasPrefix(rawID, preCleanQuarantineIDPrefix) {
baseDir = filepath.Join(quarantineDir, "pre_clean")
name = strings.TrimPrefix(rawID, preCleanQuarantineIDPrefix)
}
name = filepath.Base(name)
if name == "" || name == "." || name == ".." {
return quarantineEntryRef{}, fmt.Errorf("invalid quarantine ID")
}
itemPath := filepath.Join(baseDir, name)
if !isPathWithin(itemPath, baseDir) {
return quarantineEntryRef{}, fmt.Errorf("invalid quarantine ID")
}
return quarantineEntryRef{
ID: rawID,
ItemPath: itemPath,
MetaPath: itemPath + ".meta",
}, nil
}
// readQuarantineMeta reads and parses a quarantine .meta JSON file.
func readQuarantineMeta(metaPath string) (*checks.QuarantineMeta, error) {
// #nosec G304 -- metaPath is constructed by resolveQuarantineEntry under
// the quarantine base dir with filepath.Base applied to the ID.
data, err := os.ReadFile(metaPath)
if err != nil {
return nil, fmt.Errorf("read quarantine meta: %w", err)
}
var meta checks.QuarantineMeta
if err := json.Unmarshal(data, &meta); err != nil {
return nil, fmt.Errorf("parse quarantine meta: %w", err)
}
return &meta, nil
}
// listMetaFiles returns the full paths of all .meta files in dir (non-recursive).
// Returns nil on any error (e.g., directory does not exist).
func listMetaFiles(dir string) []string {
entries, err := os.ReadDir(dir)
if err != nil {
return nil
}
var metas []string
for _, e := range entries {
if e.IsDir() {
continue
}
if strings.HasSuffix(e.Name(), ".meta") {
metas = append(metas, filepath.Join(dir, e.Name()))
}
}
return metas
}
// Tests redirect platform and scratch roots without changing the config scope.
var quarantineRestoreRoots []string
func quarantineRootsForConfig(cfg *config.Config) ([]string, error) {
roots := quarantineRestoreRoots
if roots == nil {
roots = append(platform.Detect().AccountHomeRoots(), "/tmp", "/dev/shm", "/var/tmp")
}
roots = append([]string(nil), roots...)
// Only platform-owned roots can have aliases (for example /tmp on a
// development host). Snapshot them before adding tenant-owned roots.
for _, root := range roots {
if resolved, err := filepath.EvalSymlinks(root); err == nil && resolved != root {
roots = append(roots, resolved)
}
}
if cfg != nil {
configured, err := platform.ResolveAccountRoots(cfg.AccountRoots)
return append(roots, configured...), err
}
return roots, nil
}
func validateQuarantineRestorePath(path string, roots []string) (string, error) {
cleanPath := filepath.Clean(strings.TrimSpace(path))
if cleanPath == "" {
return "", fmt.Errorf("restore path is required")
}
if !filepath.IsAbs(cleanPath) {
return "", fmt.Errorf("restore path must be absolute")
}
if !pathWithinAny(cleanPath, roots) {
return "", fmt.Errorf("restore path is outside the allowed restore roots: %s", cleanPath)
}
ancestor, err := nearestExistingAncestor(cleanPath)
if err != nil {
return "", err
}
resolvedAncestor, err := filepath.EvalSymlinks(ancestor)
if err != nil {
return "", fmt.Errorf("cannot validate restore path: %w", err)
}
if !pathWithinAny(resolvedAncestor, roots) {
return "", fmt.Errorf("restore path escapes the allowed restore roots: %s", cleanPath)
}
if accountRoot := homeAccountRoot(cleanPath); accountRoot != "" && !isPathWithin(resolvedAncestor, accountRoot) {
return "", fmt.Errorf("restore path escapes the account boundary: %s", cleanPath)
}
return cleanPath, nil
}
func pathWithinAny(path string, bases []string) bool {
for _, base := range bases {
if isPathWithin(path, base) {
return true
}
}
return false
}
func nearestExistingAncestor(path string) (string, error) {
current := filepath.Clean(path)
for {
if _, err := os.Lstat(current); err == nil {
return current, nil
} else if !os.IsNotExist(err) {
return "", fmt.Errorf("cannot stat restore path: %w", err)
}
parent := filepath.Dir(current)
if parent == current {
return "", fmt.Errorf("restore path has no existing ancestor: %s", path)
}
current = parent
}
}
func homeAccountRoot(path string) string {
clean := filepath.Clean(path)
for _, root := range platform.Detect().AccountHomeRoots() {
rest, found := strings.CutPrefix(clean, filepath.Clean(root)+string(filepath.Separator))
if !found {
continue
}
account, tail, inside := strings.Cut(rest, string(filepath.Separator))
if inside && tail != "" {
return filepath.Join(root, account)
}
}
return ""
}
package webui
import (
"errors"
"fmt"
"net/http"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/incident"
)
type timelineEvent struct {
Timestamp time.Time `json:"timestamp"`
Type string `json:"type"` // "finding", "action", "block"
Severity string `json:"severity,omitempty"` // a finding's label; actions have none
Summary string `json:"summary"`
Details string `json:"details,omitempty"`
Source string `json:"source"` // "history", "audit", "firewall"
}
const incidentTimelineEventLimit = 200
// incidentSnapshotScanCap bounds how many incidents the timeline walk
// will inspect from the correlator snapshot per request. Correlators on
// hot hosts can hold thousands of open + recently-closed incidents and
// each one carries a full event timeline; without this cap a single
// /api/v1/incident call can walk hundreds of thousands of timeline
// events before paginating down to incidentTimelineEventLimit. The cap
// is generous (most timelines hit the 200-event ceiling well before the
// 1000-incident ceiling) but bounds worst-case wall time and memory.
const incidentSnapshotScanCap = 1000
func (s *Server) handleIncident(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "incident.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
func (s *Server) apiIncident(w http.ResponseWriter, r *http.Request) {
ip := r.URL.Query().Get("ip")
account := r.URL.Query().Get("account")
if ip == "" && account == "" {
writeJSONError(w, "ip or account parameter is required", http.StatusBadRequest)
return
}
hours := queryInt(r, "hours", 72)
if hours > 720 {
hours = 720
} // max 30 days
cutoff := time.Now().Add(-time.Duration(hours) * time.Hour)
var events []timelineEvent
// Build search terms
var searchTerms []string
if ip != "" {
searchTerms = append(searchTerms, ip)
}
if account != "" {
searchTerms = append(searchTerms, "/home/"+account+"/", account)
}
dedup := make(map[string]struct{})
dedupKey := func(t time.Time, summary string) string {
return t.UTC().Format(time.RFC3339Nano) + "|" + summary
}
matchesHistoryQuery := func(f alert.Finding) bool {
matched := false
for _, term := range searchTerms {
if strings.Contains(f.Message, term) || strings.Contains(f.Details, term) {
matched = true
break
}
}
if !matched {
return false
}
summary := f.Check + ": " + f.Message
key := dedupKey(f.Timestamp, summary)
if _, seen := dedup[key]; seen {
return false
}
dedup[key] = struct{}{}
return true
}
// Search newest-first and stop once the timeline has enough matching
// history rows. Busy hosts can retain large 30-day windows, so response
// size alone is not a safe bound for the read path.
allHistory := s.store.SearchHistorySince(cutoff, incidentTimelineEventLimit+1, matchesHistoryQuery)
historyCapped := len(allHistory) > incidentTimelineEventLimit
if historyCapped {
// The extra row only tells that more exist; forget it so an
// incident event it duplicates is still listed.
dropped := allHistory[incidentTimelineEventLimit]
delete(dedup, dedupKey(dropped.Timestamp, dropped.Check+": "+dropped.Message))
allHistory = allHistory[:incidentTimelineEventLimit]
}
for _, f := range allHistory {
summary := f.Check + ": " + f.Message
events = append(events, timelineEvent{
Timestamp: f.Timestamp.UTC(),
Type: "finding",
Severity: f.Severity.String(),
Summary: summary,
Details: f.Details,
Source: "history",
})
}
// Fold in events from the incident correlator. The finding history
// bucket rotates aggressively on busy hosts so a Critical incident
// from two days ago may have no surviving history row, but the
// incident object still carries the full timeline. Walk every
// incident, match by RemoteIP for IP queries or by Account / Mailbox /
// Domain for account queries, and emit each matching timeline event.
truncated := historyCapped
if s.incidentCorrelator != nil {
snap, totalIncidents := s.incidentCorrelator.SnapshotPageStatuses(nil, 0, incidentSnapshotScanCap)
truncated = truncated || totalIncidents > len(snap)
for _, inc := range snap {
incMatches := incidentMatchesAccount(inc, account)
for _, ev := range inc.Timeline {
if ev.Time.Before(cutoff) {
continue
}
match := false
if ip != "" && ev.RemoteIP == ip {
match = true
}
if !match && incMatches {
match = true
}
if !match {
continue
}
summary := ev.Check
if ev.Message != "" {
if summary != "" {
summary += ": "
}
summary += ev.Message
}
key := dedupKey(ev.Time, summary)
if _, seen := dedup[key]; seen {
continue
}
dedup[key] = struct{}{}
events = append(events, timelineEvent{
Timestamp: ev.Time.UTC(),
Type: "finding",
Severity: inc.Severity.String(),
Summary: summary,
Details: "From incident " + inc.ID + " (" + string(inc.Kind) + ", " + string(inc.Status) + ")",
Source: "incident:" + inc.ID,
})
}
}
}
// Search UI audit log
const auditScanLimit = 500
auditEntries := readUIAuditLog(s.cfg.StatePath, auditScanLimit+1)
if len(auditEntries) > auditScanLimit {
truncated = true
auditEntries = auditEntries[:auditScanLimit]
}
for _, a := range auditEntries {
if a.Timestamp.Before(cutoff) {
continue
}
matched := false
for _, term := range searchTerms {
if strings.Contains(a.Target, term) || strings.Contains(a.Details, term) {
matched = true
break
}
}
if !matched {
continue
}
events = append(events, timelineEvent{
Timestamp: a.Timestamp.UTC(),
Type: "action",
Summary: a.Action + ": " + a.Target,
Details: a.Details,
Source: "audit",
})
}
// Newest first, compared as instants.
sort.SliceStable(events, func(i, j int) bool {
return events[i].Timestamp.After(events[j].Timestamp)
})
total := len(events)
if total > incidentTimelineEventLimit {
events = events[:incidentTimelineEventLimit]
truncated = true
}
if truncated {
w.Header().Set("X-CSM-Truncated", "1")
}
writeItems(w, events, map[string]interface{}{
"total": total,
"offset": 0,
"limit": incidentTimelineEventLimit,
"query_ip": ip,
"query_account": account,
"window_seconds": hours * 3600,
"truncated": truncated,
})
}
// incidentMatchesAccount reports whether an incident's identity fields
// match the account search term. Empty account never matches so an
// IP-only query does not pull in unrelated incidents.
func incidentMatchesAccount(inc incident.Incident, account string) bool {
if account == "" {
return false
}
return inc.Account == account || inc.Mailbox == account || inc.Domain == account
}
// maxIncidentPageSize caps the page size a client may request so a
// misbehaving consumer cannot OOM the daemon by asking for the whole
// world in one round-trip. The web UI's default page size is well
// below this; the ceiling exists for defense in depth.
const maxIncidentPageSize = 500
// defaultIncidentPageSize is the page size when the client passes no
// limit. Tuned to fit comfortably on one screen.
const defaultIncidentPageSize = 50
// apiIncidentList serves GET /api/v1/incidents as one page:
// {"items":[...], "total":N, "offset":N, "limit":N, "status":"..."}.
// limit defaults to defaultIncidentPageSize; total counts every incident
// the status filter matches, so a client pages on through offset.
//
// status accepts the four spec values (open/contained/resolved/dismissed)
// plus the UI-only convenience "active" that means
// open+contained. An empty string means all statuses. Anything else is
// rejected with 400 Bad Request rather than silently widening to all,
// which would hide a typo like ?status=opn.
func (s *Server) apiIncidentList(w http.ResponseWriter, r *http.Request) {
statusParam := r.URL.Query().Get("status")
statuses, err := parseIncidentStatusFilter(statusParam)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
limit := queryInt(r, "limit", defaultIncidentPageSize)
if limit <= 0 {
limit = defaultIncidentPageSize
}
if limit > maxIncidentPageSize {
limit = maxIncidentPageSize
}
offset := queryInt(r, "offset", 0)
if offset < 0 {
offset = 0
}
var items []incident.Incident
total := 0
if s.incidentCorrelator != nil {
items, total = s.incidentPage(statuses, offset, limit)
}
writeItems(w, items, map[string]any{
"total": total,
"offset": offset,
"limit": limit,
"status": statusParam,
"truncated": historyPageTruncated(total, offset, len(items)),
})
}
// parseIncidentStatusFilter validates the status query parameter and
// returns the set of statuses it expands to. Empty input means "all";
// "active" is the UI-only convenience for open+contained.
func parseIncidentStatusFilter(s string) ([]incident.Status, error) {
switch s {
case "":
return nil, nil
case "active":
return []incident.Status{incident.StatusOpen, incident.StatusContained}, nil
case string(incident.StatusOpen),
string(incident.StatusContained),
string(incident.StatusResolved),
string(incident.StatusDismissed):
return []incident.Status{incident.Status(s)}, nil
}
return nil, fmt.Errorf("invalid status %q", s)
}
// incidentPage returns a status-filtered page. The "active" filter
// expands to open+contained and is handled by the correlator in one
// sorted pass so pagination stays stable across statuses.
func (s *Server) incidentPage(statuses []incident.Status, offset, limit int) ([]incident.Incident, int) {
return s.incidentCorrelator.SnapshotPageStatuses(statuses, offset, limit)
}
// apiIncidentShow serves GET /api/v1/incidents/<id>. 404 if not found.
func (s *Server) apiIncidentShow(w http.ResponseWriter, r *http.Request) {
id := strings.TrimPrefix(r.URL.Path, "/api/v1/incidents/")
id = strings.TrimSuffix(id, "/")
if id == "" || s.incidentCorrelator == nil {
writeJSONError(w, "Incident not found", http.StatusNotFound)
return
}
inc, ok := s.incidentCorrelator.Get(id)
if !ok {
writeJSONError(w, "Incident not found", http.StatusNotFound)
return
}
writeJSON(w, inc)
}
// apiIncidentStatus serves POST /api/v1/incidents/<id>/status. Body
// {"status": "resolved", "details": "..."}.
func (s *Server) apiIncidentStatus(w http.ResponseWriter, r *http.Request) {
id := strings.TrimPrefix(r.URL.Path, "/api/v1/incidents/")
id = strings.TrimSuffix(id, "/status")
id = strings.Trim(id, "/")
var body struct {
Status string `json:"status"`
Details string `json:"details"`
}
// Cap the request body like every other mutating handler; a bare
// json.NewDecoder(r.Body) would buffer an unbounded body into memory.
if err := decodeJSONBodyLimited(w, r, 16*1024, &body); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
if s.incidentCorrelator == nil {
writeJSONError(w, "Incidents are not enabled", http.StatusServiceUnavailable)
return
}
if err := s.incidentCorrelator.SetStatus(id, incident.Status(body.Status), body.Details); err != nil {
code := http.StatusBadRequest
if errors.Is(err, incident.ErrIncidentNotFound) {
code = http.StatusNotFound
}
writeJSONError(w, err.Error(), code)
return
}
s.auditLog(r, "incident_status", id, strings.TrimSpace(body.Status+" "+body.Details))
writeJSON(w, map[string]interface{}{"ok": true})
}
// apiIncidentRouter dispatches /api/v1/incidents/<id>[...] sub-paths.
// POST .../status -> apiIncidentStatus; GET .../<id> -> apiIncidentShow.
func (s *Server) apiIncidentRouter(w http.ResponseWriter, r *http.Request) {
if r.Method == http.MethodPost && strings.HasSuffix(r.URL.Path, "/status") {
s.apiIncidentStatus(w, r)
return
}
s.apiIncidentShow(w, r)
}
package webui
import (
"net/http"
"strconv"
"strings"
"github.com/pidginhost/csm/internal/incident"
)
// incidentGroupsDefaultLimit / Max bound how many group rows the UI
// receives. A typical busy production host emits ~12 attacker-IP rows;
// the cap of 200 leaves headroom for the long tail without unbounded
// payload size.
const (
incidentGroupsDefaultLimit = 50
incidentGroupsMaxLimit = 200
)
// apiIncidentGroups handles GET /api/v1/incidents/groups. Buckets the
// in-memory incident snapshot by (kind, source) and returns rolled-up
// group rows. Read-scope eligible; never mutates state.
func (s *Server) apiIncidentGroups(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
q := r.URL.Query()
limit := queryInt(r, "limit", incidentGroupsDefaultLimit)
if limit <= 0 || limit > incidentGroupsMaxLimit {
limit = incidentGroupsDefaultLimit
}
offset := queryInt(r, "offset", 0)
if offset < 0 {
offset = 0
}
filter := incident.GroupFilter{
Kind: incident.Kind(strings.TrimSpace(q.Get("kind"))),
Offset: offset,
MaxGroups: limit,
}
status := strings.ToLower(strings.TrimSpace(q.Get("status")))
switch status {
case "", "active":
// Default surface: open + contained, the UI's primary tab.
filter.StatusSet = []incident.Status{incident.StatusOpen, incident.StatusContained}
case "all":
// No status filter; the operator wants the full picture.
case string(incident.StatusOpen),
string(incident.StatusContained),
string(incident.StatusResolved),
string(incident.StatusDismissed):
filter.StatusSet = []incident.Status{incident.Status(status)}
default:
writeJSONError(w, "unknown status: "+strconv.Quote(q.Get("status")), http.StatusBadRequest)
return
}
var resp incident.GroupsResponse
if s.incidentCorrelator != nil {
resp = incident.BuildGroups(s.incidentCorrelator.Snapshot(), filter)
}
writeItems(w, resp.Groups, map[string]interface{}{
"total": resp.TotalGroups,
"offset": offset,
"limit": limit,
"scanned_incidents": resp.ScannedIncidents,
"truncated": resp.Truncated || historyPageTruncated(resp.TotalGroups, offset, len(resp.Groups)),
"scan_truncated": resp.Truncated,
})
}
package webui
import (
"crypto/subtle"
"fmt"
"net/http"
"os"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/metrics"
)
// handleMetrics serves Prometheus text exposition from the process
// default metrics registry. ROADMAP item 4.
//
// Auth policy (checked in isMetricsAuthenticated):
//
// - If `webui.metrics_token` is set in the live config, a matching
// `Authorization: Bearer` header unlocks the endpoint. The token
// is read from config.Active() per request so a SIGHUP rotation
// (the field is tagged `hotreload:"safe"`) takes effect without
// a restart.
// - As a fallback, a valid UI session cookie or the UI AuthToken
// Bearer is accepted so the dashboard can self-scrape without a
// second credential.
//
// No CSRF required: metrics is read-only and idempotent.
func (s *Server) handleMetrics(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if !s.isMetricsAuthenticated(r) {
w.Header().Set("WWW-Authenticate", `Bearer realm="csm-metrics"`)
http.Error(w, "Unauthorized", http.StatusUnauthorized)
return
}
w.Header().Set("Content-Type", "text/plain; version=0.0.4; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
if r.Method == http.MethodHead {
return
}
if err := metrics.WriteOpenMetrics(w); err != nil {
// The response body is already partially written; there is no
// meaningful HTTP status to flip. Drop a server log line.
fmt.Fprintf(os.Stderr, "webui: metrics WriteOpenMetrics: %v\n", err)
}
}
func (s *Server) isMetricsAuthenticated(r *http.Request) bool {
// Read the metrics token from config.Active() when available so a
// SIGHUP-driven rotation of webui.metrics_token takes effect on
// the next request. MetricsToken is tagged `hotreload:"safe"` for
// this reason, even though its WebUI parent is restart-required
// for the listener/TLS/auth-token fields. Fall back to s.cfg on
// the cold-start window before SetActive is called.
tok := s.cfg.WebUI.MetricsToken
if live := config.Active(); live != nil {
tok = live.WebUI.MetricsToken
}
if tok != "" {
if auth := r.Header.Get("Authorization"); len(auth) > 7 && auth[:7] == "Bearer " {
if subtle.ConstantTimeCompare([]byte(auth[7:]), []byte(tok)) == 1 {
return true
}
}
}
// Fall back to the UI session / AuthToken path so the dashboard
// can scrape itself without a second credential.
return s.isAuthenticated(r)
}
package webui
import (
"net/http"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/store"
)
func (s *Server) handleModSec(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "modsec.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
// modsecBlockView is an aggregated view of blocks per IP+rule. Phase 8.4
// extends the response with first_seen / last_seen_iso (RFC3339), top_uris,
// domain_count, and sample_events so the workbench can drive the detail
// panel without a second round trip. The extra fields are additive; legacy
// field names keep their JSON keys.
type modsecBlockView struct {
IP string `json:"ip"`
Country string `json:"country"`
CountryName string `json:"country_name"`
RuleID string `json:"rule_id"`
Description string `json:"description"`
Domains string `json:"domains"`
DomainList []string `json:"domain_list,omitempty"`
DomainCount int `json:"domain_count"`
Hits int `json:"hits"`
LastSeen time.Time `json:"last_seen,omitzero"`
FirstSeen time.Time `json:"first_seen,omitzero"`
TopURIs []string `json:"top_uris"`
SampleEvents []modsecSampleEvent `json:"sample_events"`
Escalated bool `json:"escalated"`
}
// modsecSampleEvent is a compact per-IP event included in the grouped
// blocks response so the UI can show recent activity without a second
// call to /api/v1/modsec/events.
type modsecSampleEvent struct {
Time time.Time `json:"time"`
RuleID string `json:"rule_id"`
Hostname string `json:"hostname"`
URI string `json:"uri"`
Severity string `json:"severity"`
}
// modsecEventView is a single ModSecurity event. The UI renders Time in the
// operator's time zone and sorts on it.
type modsecEventView struct {
Time time.Time `json:"time"`
IP string `json:"ip"`
Country string `json:"country"`
RuleID string `json:"rule_id"`
Hostname string `json:"hostname"`
URI string `json:"uri"`
Severity string `json:"severity"`
}
// apiModSecStats returns 24h summary stats for ModSecurity blocks.
func (s *Server) apiModSecStats(w http.ResponseWriter, r *http.Request) {
findings := deduplicateModSecFindings(s.modsecFindings(r))
uniqueIPs := make(map[string]bool)
ruleCounts := make(map[string]int)
escalated := 0
for _, f := range findings {
ip := extractModSecIP(f)
if ip != "" {
uniqueIPs[ip] = true
}
rule := extractModSecRule(f)
if rule != "" {
ruleCounts[rule]++
}
if isModSecEscalation(f.Check) {
escalated++
}
}
topRule := "--"
topCount := 0
for rule, count := range ruleCounts {
if count > topCount {
topCount = count
topRule = rule
}
}
writeJSON(w, map[string]interface{}{
"total": len(findings),
"unique_ips": len(uniqueIPs),
"escalated": escalated,
"top_rule": topRule,
})
}
const (
// modsecFindingsScanCap bounds the 24h history read before per-IP
// aggregation starts. The aggregate map has its own cap below, but
// the handler still needs a findings cap for hosts seeing one or two
// hot IP+rule pairs millions of times.
modsecFindingsScanCap = 10000
// modsecBlocksMaxAggregates caps the IP+rule aggregation map so a host
// with millions of unique ModSec rule hits cannot OOM the daemon by
// asking for /api/v1/modsec/blocks. Existing aggregates keep updating
// after the cap is reached; new IP+rule keys past the cap are dropped
// silently with the X-CSM-Truncated response header set so monitoring
// can flag the condition. Default sized for ~50 MB peak: 50000 entries
// times a few hundred bytes per aggregate.
modsecBlocksMaxAggregates = 50000
)
// apiModSecBlocks returns aggregated blocks per IP+rule for the last 24h.
func (s *Server) apiModSecBlocks(w http.ResponseWriter, r *http.Request) {
findings, truncated := s.modsecFindingsWithTruncation(r)
findings = deduplicateModSecFindings(findings)
type blockAgg struct {
ip string
ruleID string
description string
domains map[string]bool
uriCounts map[string]int
hits int
firstSeen time.Time
lastSeen time.Time
escalated bool
samples []modsecSampleEvent // newest-first, capped at 3
}
byBlock := make(map[string]*blockAgg)
// byIP indexes the aggregates per address, so marking escalated
// addresses does not scan every aggregate once per address.
byIP := make(map[string][]*blockAgg)
escalatedIPs := make(map[string]bool)
blockKey := func(ip, rule string) string {
return ip + "\x00" + rule
}
for _, f := range findings {
if isModSecEscalation(f.Check) {
ip := extractModSecIP(f)
if ip != "" {
escalatedIPs[ip] = true
}
continue
}
ip := extractModSecIP(f)
if ip == "" {
continue
}
rule := extractModSecRule(f)
desc := extractModSecDescription(f)
domain := extractModSecHostname(f)
uri := extractModSecURI(f)
key := blockKey(ip, rule)
agg, ok := byBlock[key]
if !ok {
if len(byBlock) >= modsecBlocksMaxAggregates {
truncated = true
continue
}
agg = &blockAgg{
ip: ip,
ruleID: rule,
description: desc,
domains: make(map[string]bool),
uriCounts: make(map[string]int),
firstSeen: f.Timestamp,
}
byBlock[key] = agg
byIP[ip] = append(byIP[ip], agg)
}
agg.hits++
if agg.firstSeen.IsZero() || f.Timestamp.Before(agg.firstSeen) {
agg.firstSeen = f.Timestamp
}
if f.Timestamp.After(agg.lastSeen) {
agg.lastSeen = f.Timestamp
if rule != "" {
agg.ruleID = rule
}
if desc != "" {
agg.description = desc
}
}
if agg.description == "" && desc != "" {
agg.description = desc
}
// Skip server IPs and empty hostnames - only show actual domain names.
// ModSecurity logs the server IP as hostname when the request doesn't
// match a specific vhost (e.g. direct IP access, SNI mismatch).
if domain != "" && !looksLikeIP(domain) {
agg.domains[domain] = true
}
if uri != "" {
agg.uriCounts[uri]++
}
if len(agg.samples) < 3 {
agg.samples = append(agg.samples, modsecSampleEvent{
Time: f.Timestamp.UTC(),
RuleID: rule,
Hostname: domain,
URI: uri,
Severity: f.Severity.String(),
})
}
}
for ip := range escalatedIPs {
for _, agg := range byIP[ip] {
agg.escalated = true
}
if len(byIP[ip]) == 0 {
if len(byBlock) >= modsecBlocksMaxAggregates {
truncated = true
continue
}
byBlock[blockKey(ip, "")] = &blockAgg{
ip: ip,
escalated: true,
domains: make(map[string]bool),
uriCounts: make(map[string]int),
}
}
}
var result []modsecBlockView
for _, agg := range byBlock {
if agg.hits == 0 && !agg.escalated {
continue
}
var domainList []string
for d := range agg.domains {
domainList = append(domainList, d)
}
sort.Strings(domainList)
domains := strings.Join(domainList, ", ")
if len(domains) > 80 {
domains = domains[:77] + "..."
}
topURIs := topKeysByCount(agg.uriCounts, 5)
country, countryName := s.modsecCountryOf(agg.ip)
result = append(result, modsecBlockView{
IP: agg.ip,
Country: country,
CountryName: countryName,
RuleID: agg.ruleID,
Description: agg.description,
Domains: domains,
DomainList: domainList,
DomainCount: len(agg.domains),
Hits: agg.hits,
LastSeen: agg.lastSeen.UTC(),
FirstSeen: agg.firstSeen.UTC(),
TopURIs: topURIs,
SampleEvents: agg.samples,
Escalated: agg.escalated,
})
}
sort.Slice(result, func(i, j int) bool {
if result[i].Hits != result[j].Hits {
return result[i].Hits > result[j].Hits
}
if result[i].Escalated != result[j].Escalated {
return result[i].Escalated
}
if !result[i].LastSeen.Equal(result[j].LastSeen) {
return result[i].LastSeen.After(result[j].LastSeen)
}
if result[i].IP != result[j].IP {
return result[i].IP < result[j].IP
}
return result[i].RuleID < result[j].RuleID
})
if truncated {
w.Header().Set("X-CSM-Truncated", "1")
}
writeItems(w, result, map[string]interface{}{
"total": len(result), "offset": 0, "limit": modsecBlocksMaxAggregates, "truncated": truncated,
})
}
// apiModSecEvents returns the most recent individual ModSecurity events.
func (s *Server) apiModSecEvents(w http.ResponseWriter, r *http.Request) {
limit := 100
if l := r.URL.Query().Get("limit"); l != "" {
if n, err := strconv.Atoi(l); err == nil && n > 0 && n <= 500 {
limit = n
}
}
findings, truncated := s.modsecFindingsWithTruncation(r)
findings = deduplicateModSecFindings(findings)
result := make([]modsecEventView, 0, limit)
total := 0
for _, f := range findings {
if isModSecEscalation(f.Check) {
continue
}
total++
if len(result) >= limit {
continue
}
ip := extractModSecIP(f)
country, _ := s.modsecCountryOf(ip)
result = append(result, modsecEventView{
Time: f.Timestamp.UTC(),
IP: ip,
Country: country,
RuleID: extractModSecRule(f),
Hostname: extractModSecHostname(f),
URI: extractModSecURI(f),
Severity: f.Severity.String(),
})
}
writeCapped(w, result, total, limit, map[string]interface{}{"truncated": truncated})
}
// deduplicateModSecFindings merges Apache + LiteSpeed duplicate events.
// Both log the same block within the same second - keep one with merged fields.
func deduplicateModSecFindings(findings []alert.Finding) []alert.Finding {
type dedupKey struct {
second time.Time
ip string
rule string
}
seen := make(map[dedupKey]int) // key → index in result
var result []alert.Finding
for _, f := range findings {
ip := extractModSecIP(f)
rule := extractModSecRule(f)
ts := f.Timestamp.UTC().Truncate(time.Second)
key := dedupKey{second: ts, ip: ip, rule: rule}
if idx, ok := seen[key]; ok {
// Merge richer details into existing entry
existing := &result[idx]
if extractModSecHostname(f) != "" && extractModSecHostname(*existing) == "" {
existing.Details = f.Details
}
} else {
seen[key] = len(result)
result = append(result, f)
}
}
return result
}
func isModSecEscalation(check string) bool {
return check == "modsec_block_escalation" || check == "modsec_csm_block_escalation"
}
// modsecWindow maps the optional window query param to a lookback duration.
// Missing or unrecognized values fall back to 24h, which also bounds the read
// since history retention and the scan cap assume roughly a day.
func modsecWindow(r *http.Request) time.Duration {
switch r.URL.Query().Get("window") {
case "1h":
return time.Hour
case "6h":
return 6 * time.Hour
default:
return 24 * time.Hour
}
}
// modsecSeverityFilter maps the optional severity query param to a minimum
// severity. ok is false when no (or an unrecognized) filter is requested, in
// which case all severities pass.
func modsecSeverityFilter(r *http.Request) (alert.Severity, bool) {
return parseSeverity(r.URL.Query().Get("severity"))
}
// modsecCountryOf resolves an IP to its ISO country code and full name via the
// loaded GeoIP database. Returns empty strings when the IP is empty or no
// GeoIP database is loaded.
func (s *Server) modsecCountryOf(ip string) (string, string) {
if ip == "" {
return "", ""
}
db := s.geoIPDB.Load()
if db == nil {
return "", ""
}
info := db.Lookup(ip)
return info.Country, info.CountryName
}
// modsecFindings returns modsec findings for the request window and severity
// filter.
func (s *Server) modsecFindings(r *http.Request) []alert.Finding {
findings, _ := s.modsecFindingsWithTruncation(r)
return findings
}
func (s *Server) modsecFindingsWithTruncation(r *http.Request) ([]alert.Finding, bool) {
db := store.Global()
if db == nil {
return nil, false
}
cutoff := time.Now().Add(-modsecWindow(r))
minSev, hasSev := modsecSeverityFilter(r)
findings := db.SearchHistorySince(cutoff, modsecFindingsScanCap+1, func(f alert.Finding) bool {
if !strings.HasPrefix(f.Check, "modsec_") {
return false
}
if hasSev && f.Severity < minSev {
return false
}
return true
})
if len(findings) > modsecFindingsScanCap {
return findings[:modsecFindingsScanCap], true
}
return findings, false
}
// --- Field extraction from finding Details ---
// Details format: "Rule: NNNN\nMessage: ...\nHostname: ...\nURI: ..."
func extractModSecIP(f alert.Finding) string {
// Try from message: "... from IP on ..." or "... from IP ..."
msg := f.Message
if idx := strings.Index(msg, " from "); idx >= 0 {
rest := msg[idx+6:]
if sp := strings.IndexAny(rest, " \n"); sp >= 0 {
rest = rest[:sp]
}
if len(rest) >= 7 && rest[0] >= '0' && rest[0] <= '9' && strings.Count(rest, ".") == 3 {
return rest
}
}
// Fallback: parse [client IP] from raw log line in Details
if ip := extractBetween(f.Details, "[client ", "]"); ip != "" {
// Strip port if present (Apache 2.4: "IP:port")
if strings.Count(ip, ":") == 1 {
if idx := strings.LastIndex(ip, ":"); idx > 0 {
ip = ip[:idx]
}
}
return ip
}
// Fallback: LiteSpeed format - IP in [IP:PORT-CONN#VHOST]
for _, field := range strings.Fields(f.Details) {
if strings.HasPrefix(field, "[") && strings.Contains(field, "#") {
inner := strings.TrimPrefix(field, "[")
if colonIdx := strings.Index(inner, ":"); colonIdx > 0 {
ip := inner[:colonIdx]
if len(ip) >= 7 && ip[0] >= '0' && ip[0] <= '9' {
return ip
}
}
}
}
return ""
}
func extractModSecRule(f alert.Finding) string {
// Try structured format first
if v := extractDetailField(f.Details, "Rule: "); v != "" {
return v
}
// Fallback: parse [id "NNNN"] from raw log line in Details
return extractBetween(f.Details, `[id "`, `"]`)
}
// csmRuleDescriptions provides fallback descriptions for known ModSecurity
// rules.
// LiteSpeed error logs omit the [msg "..."] field, so the log-extracted
// description is often empty. This map ensures the UI always shows a
// meaningful description for rules we define ourselves and common vendor rules.
var csmRuleDescriptions = map[string]string{
"900001": "Blocked LEVIATHAN CGI extension access",
"900002": "Blocked LEVIATHAN directory access",
"900003": "Blocked PHP execution in uploads directory",
"900004": "Blocked PHP execution in languages directory",
"900005": "Blocked direct wp-config.php access",
"900007": "XML-RPC rate limit exceeded",
"900008": "Blocked known webshell filename access",
"900009": "Blocked GSocket User-Agent",
"900100": "WP-Automatic SQLi (CVE-2024-27956)",
"900101": "LayerSlider SQLi (CVE-2024-2879)",
"900102": "Really Simple Security auth bypass (CVE-2024-10924)",
"900103": "LiteSpeed Cache directory traversal (CVE-2024-4345)",
"900104": "Ultimate Member SQLi (CVE-2024-1071)",
"900105": "Backup Migration RCE (CVE-2023-6553)",
"900106": "GiveWP object injection (CVE-2024-5932)",
"900107": "WP File Manager arbitrary upload (CVE-2024-3400)",
"900110": "PHP object injection attempt",
"900111": "Blocked PHP in wp-content/upgrade",
"900112": "WordPress user enumeration blocked",
"900113": "wp-login brute force rate limit",
"900114": "wp-login brute force rate limit",
"900115": "Blocked .env file access",
"900116": "Blocked scanner probe",
"900120": "Blocked wp-coder preview endpoint",
"900121": "Blocked wp-coder attributes endpoint",
"900122": "Blocked wp2shell exploit tool User-Agent",
"900123": "REST batch request counter",
"900124": "REST batch endpoint rate limit",
"900125": "Blocked wp2shell tool fingerprint",
"900126": "Blocked wp2shell tool fingerprint (encoded)",
"900127": "Ultimate Member privilege escalation (CVE-2023-3460)",
// Comodo WAF (CWAF) common rules. Rule IDs in the 21xxxx range are
// from the Comodo vendor ruleset (e.g. /etc/apache2/conf.d/
// modsec_vendor_configs/comodo_litespeed/), NOT from OWASP CRS.
// They were previously mislabeled as "OWASP:" here.
"210710": "Comodo WAF: Request content type is not allowed by policy",
"210381": "Comodo WAF: URL encoding abuse attack attempt",
"214930": "Comodo WAF: Inbound anomaly score threshold exceeded",
"214940": "Comodo WAF: Outbound anomaly score threshold exceeded",
"218420": "Comodo WAF: Request content type restriction",
// OWASP CRS common rules. IDs in the 9xxxxx range are the standard
// OWASP CRS 3.x schema (920xxx protocol, 930xxx LFI, 941xxx XSS,
// 942xxx SQLi).
"920170": "OWASP: Validate GET/HEAD request",
"920420": "OWASP: Request content type is not allowed by policy",
"920600": "OWASP: Illegal Accept header",
"930100": "OWASP: Path traversal attack",
"930110": "OWASP: Path traversal attack",
"930120": "OWASP: OS file access attempt",
"941100": "OWASP: XSS attack detected via libinjection",
"941160": "OWASP: XSS Filter - Category 1",
"942100": "OWASP: SQL injection attack detected via libinjection",
}
func extractModSecDescription(f alert.Finding) string {
if v := extractDetailField(f.Details, "Message: "); v != "" {
return v
}
if v := extractBetween(f.Details, `[msg "`, `"]`); v != "" {
return v
}
// Fallback: use static description for CSM custom rules when the log
// format (e.g. LiteSpeed) doesn't include the [msg "..."] field.
rule := extractModSecRule(f)
if desc, ok := csmRuleDescriptions[rule]; ok {
return desc
}
return ""
}
func extractModSecHostname(f alert.Finding) string {
if v := extractDetailField(f.Details, "Hostname: "); v != "" {
return v
}
return extractBetween(f.Details, `[hostname "`, `"]`)
}
func extractModSecURI(f alert.Finding) string {
if v := extractDetailField(f.Details, "URI: "); v != "" {
return v
}
return extractBetween(f.Details, `[uri "`, `"]`)
}
func extractDetailField(details, prefix string) string {
for _, line := range strings.Split(details, "\n") {
if strings.HasPrefix(line, prefix) {
return strings.TrimPrefix(line, prefix)
}
}
return ""
}
// looksLikeIP returns true if the string looks like an IP address (not a domain).
func looksLikeIP(s string) bool {
if len(s) < 7 {
return false
}
for _, c := range s {
if c != '.' && (c < '0' || c > '9') {
return false
}
}
return strings.Count(s, ".") == 3
}
// extractBetween extracts the value between start and end delimiters.
// Used as fallback for old findings where Details is the raw log line.
func extractBetween(s, start, end string) string {
idx := strings.Index(s, start)
if idx < 0 {
return ""
}
rest := s[idx+len(start):]
endIdx := strings.Index(rest, end)
if endIdx < 0 {
return ""
}
return rest[:endIdx]
}
package webui
import (
"fmt"
"net/http"
"os"
"sort"
"strconv"
"time"
"github.com/pidginhost/csm/internal/modsec"
"github.com/pidginhost/csm/internal/store"
)
func validateModSecDisabledRules(allRules []modsec.Rule, disabled []int) error {
knownIDs := make(map[int]bool)
counterIDs := make(map[int]bool)
for _, rule := range allRules {
knownIDs[rule.ID] = true
if rule.IsCounter {
counterIDs[rule.ID] = true
}
}
for _, id := range disabled {
if !knownIDs[id] {
return fmt.Errorf("rule ID %d is not a known CSM rule", id)
}
if counterIDs[id] {
return fmt.Errorf("rule ID %d is a protected bookkeeping rule", id)
}
}
return nil
}
func (s *Server) handleModSecRules(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "modsec-rules.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
// GET /api/v1/modsec/rules - list all CSM rules with status
func (s *Server) apiModSecRules(w http.ResponseWriter, _ *http.Request) {
cfg := s.liveCfg().ModSec
// Check all three config fields
var missing []string
if cfg.RulesFile == "" {
missing = append(missing, "rules_file")
}
if cfg.OverridesFile == "" {
missing = append(missing, "overrides_file")
}
if cfg.ReloadCommand == "" {
missing = append(missing, "reload_command")
}
if len(missing) > 0 {
writeItems(w, []struct{}{}, map[string]interface{}{
"total": 0,
"configured": false,
"missing": missing,
})
return
}
// Parse rules from config file
allRules, err := modsec.ParseRulesFile(cfg.RulesFile)
if err != nil {
fmt.Fprintf(os.Stderr, "modsec: parse rules failed: %v\n", err)
writeJSONError(w, "Failed to parse rules file", http.StatusInternalServerError)
return
}
// Read disabled IDs from overrides file
disabledIDs, _ := modsec.ReadOverrides(cfg.OverridesFile)
disabledSet := make(map[int]bool)
for _, id := range disabledIDs {
disabledSet[id] = true
}
// Read escalation exclusions and hit counts from store
var noEscalate map[int]bool
var hits map[int]store.RuleHitStats
if db := store.Global(); db != nil {
noEscalate = db.GetModSecNoEscalateRules()
hits = db.GetModSecRuleHits()
}
// Build response - filter out counter rules
type ruleView struct {
ID int `json:"id"`
Description string `json:"description"`
Action string `json:"action"`
StatusCode int `json:"status_code"`
Phase int `json:"phase"`
Enabled bool `json:"enabled"`
Escalate bool `json:"escalate"`
Hits24h int `json:"hits_24h"`
LastHit time.Time `json:"last_hit,omitzero"`
}
var rules []ruleView
for _, r := range allRules {
if r.IsCounter {
continue // hide bookkeeping rules
}
rv := ruleView{
ID: r.ID,
Description: r.Description,
Action: r.Action,
StatusCode: r.StatusCode,
Phase: r.Phase,
Enabled: !disabledSet[r.ID],
Escalate: !noEscalate[r.ID],
}
if h, ok := hits[r.ID]; ok {
rv.Hits24h = h.Hits
rv.LastHit = h.LastHit.UTC()
}
rules = append(rules, rv)
}
active := 0
for _, r := range rules {
if r.Enabled {
active++
}
}
writeItems(w, rules, map[string]interface{}{
"total": len(rules),
"active": active,
"disabled": disabledIDs,
"configured": true,
})
}
// POST /api/v1/modsec/rules/apply - write overrides and reload
func (s *Server) apiModSecRulesApply(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Serialize apply operations - the write+reload+rollback sequence
// must not interleave with concurrent applies.
s.modSecApplyMu.Lock()
defer s.modSecApplyMu.Unlock()
cfg := s.liveCfg().ModSec
if cfg.RulesFile == "" || cfg.OverridesFile == "" || cfg.ReloadCommand == "" {
writeJSONError(w, "ModSecurity not configured", http.StatusBadRequest)
return
}
var req struct {
Disabled []int `json:"disabled"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
// Validate: only allow disabling known CSM rule IDs from the parsed rules file
allRules, err := modsec.ParseRulesFile(cfg.RulesFile)
if err != nil {
fmt.Fprintf(os.Stderr, "modsec: parse rules failed: %v\n", err)
writeJSONError(w, "Failed to parse rules file", http.StatusInternalServerError)
return
}
if err := validateModSecDisabledRules(allRules, req.Disabled); err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Save previous state for rollback
previousContent := modsec.ReadOverridesRaw(cfg.OverridesFile)
// Write new overrides
if writeErr := modsec.WriteOverrides(cfg.OverridesFile, req.Disabled); writeErr != nil {
fmt.Fprintf(os.Stderr, "modsec: write overrides failed: %v\n", writeErr)
writeJSONError(w, "Failed to write overrides", http.StatusInternalServerError)
return
}
// Reload web server
output, reloadErr := modsec.Reload(cfg.ReloadCommand)
if reloadErr != nil {
// Rollback on failure
rollbackErr := modsec.RestoreOverrides(cfg.OverridesFile, previousContent)
outcome := "Web server reload failed, changes rolled back"
if rollbackErr != nil {
outcome = "Web server reload failed; rollback failed, check the overrides before reloading"
fmt.Fprintf(os.Stderr, "modsec: rollback failed: %v\n", rollbackErr)
}
fmt.Fprintf(os.Stderr, "modsec: reload failed: %v\noutput: %s\n", reloadErr, output)
// Truncate output for client - may contain sensitive system paths
clientOutput := output
if len(clientOutput) > 500 {
clientOutput = clientOutput[:500] + "... (truncated)"
}
s.auditLog(r, "modsec_rules_apply_failed", "overrides", outcome)
writeJSONStatus(w, http.StatusInternalServerError, map[string]interface{}{
"error": outcome,
"reload_output": clientOutput,
"rolled_back": rollbackErr == nil,
})
return
}
s.auditLog(r, "modsec_rules_apply", "overrides", fmt.Sprintf("disabled rules: %v", req.Disabled))
writeOK(w, map[string]interface{}{
"disabled_count": len(req.Disabled),
})
}
// GET /api/v1/modsec/rules/escalation - rule IDs excluded from escalation
// POST /api/v1/modsec/rules/escalation - toggle escalation for a single rule
func (s *Server) apiModSecRulesEscalation(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
db := store.Global()
if db == nil {
writeJSONError(w, "Store not available", http.StatusInternalServerError)
return
}
if r.Method == http.MethodGet {
ids := []int{}
for id := range db.GetModSecNoEscalateRules() {
ids = append(ids, id)
}
sort.Ints(ids)
writeAll(w, ids)
return
}
var req struct {
RuleID int `json:"rule_id"`
Escalate bool `json:"escalate"`
}
if err := decodeJSONBodyLimited(w, r, 4*1024, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
// Validate: only accept known CSM rule IDs (900000-900999)
if req.RuleID < 900000 || req.RuleID > 900999 {
writeJSONError(w, "Rule ID must be a CSM custom rule (900000-900999)", http.StatusBadRequest)
return
}
var err error
if req.Escalate {
err = db.RemoveModSecNoEscalateRule(req.RuleID)
} else {
err = db.AddModSecNoEscalateRule(req.RuleID)
}
if err != nil {
fmt.Fprintf(os.Stderr, "modsec: escalation update failed: %v\n", err)
writeJSONError(w, "Failed to update escalation setting", http.StatusInternalServerError)
return
}
setting := "escalation off"
if req.Escalate {
setting = "escalation on"
}
s.auditLog(r, "modsec_rule_escalation", strconv.Itoa(req.RuleID), setting)
writeOK(w, map[string]interface{}{"rule_id": req.RuleID, "escalate": req.Escalate})
}
package webui
import (
"bufio"
"context"
"fmt"
"net/http"
"os"
"os/exec"
"path/filepath"
"sort"
"strconv"
"strings"
"sync"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/mysqlclient"
"github.com/pidginhost/csm/internal/redisinfo"
)
// --- Local response types ---
type perfResponse struct {
Metrics *perfMetrics `json:"metrics"`
Findings []perfFindingView `json:"findings"`
}
type perfMetrics struct {
LoadAvg [3]float64 `json:"load_avg"`
CPUCores int `json:"cpu_cores"`
MemTotalMB uint64 `json:"mem_total_mb"`
MemUsedMB uint64 `json:"mem_used_mb"`
MemAvailMB uint64 `json:"mem_avail_mb"`
SwapTotalMB uint64 `json:"swap_total_mb"`
SwapUsedMB uint64 `json:"swap_used_mb"`
PHPProcs int `json:"php_procs_total"`
TopPHPUsers []userProcs `json:"top_php_users"`
// MySQL telemetry is best-effort. Both fields are nil when csm could
// not read the server's process status or the mysql client failed (no /root/.my.cnf,
// no socket auth, mysqld absent). The webui renders "n/a" in that case
// so operators can tell "MySQL is idle" from "we couldn't ask".
MySQLMemMB *uint64 `json:"mysql_mem_mb"`
MySQLConns *int `json:"mysql_conns"`
RedisMemMB uint64 `json:"redis_mem_mb"`
RedisMaxMB uint64 `json:"redis_maxmem_mb"`
RedisKeys int64 `json:"redis_keys"`
UptimeSeconds int64 `json:"uptime_seconds"`
}
type userProcs struct {
User string `json:"user"`
Count int `json:"count"`
}
type perfFindingView struct {
Severity string `json:"severity"`
level alert.Severity
Check string `json:"check"`
Message string `json:"message"`
Details string `json:"details,omitempty"`
Key string `json:"key"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
}
// --- Cached values ---
var (
perfCoresOnce sync.Once
perfCoresCache int
perfUIDMapOnce sync.Once
perfUIDMapCache map[string]string
)
// cachedCores reads /proc/cpuinfo once and counts "processor\t" lines.
func cachedCores() int {
perfCoresOnce.Do(func() {
f, err := os.Open("/proc/cpuinfo")
if err != nil {
perfCoresCache = 1
return
}
defer func() { _ = f.Close() }()
count := 0
scanner := bufio.NewScanner(f)
for scanner.Scan() {
if strings.HasPrefix(scanner.Text(), "processor\t") {
count++
}
}
if count == 0 {
count = 1
}
perfCoresCache = count
})
return perfCoresCache
}
// cachedUID resolves a UID string to a username via /etc/passwd, cached.
func cachedUID(uid string) string {
perfUIDMapOnce.Do(func() {
perfUIDMapCache = make(map[string]string)
data, err := os.ReadFile("/etc/passwd")
if err != nil {
return
}
for _, line := range strings.Split(string(data), "\n") {
fields := strings.Split(line, ":")
if len(fields) >= 3 {
perfUIDMapCache[fields[2]] = fields[0]
}
}
})
if name, ok := perfUIDMapCache[uid]; ok {
return name
}
return uid
}
// --- Metrics sampler ---
// runCmdQuick runs a command with a 5-second timeout. All call sites pass
// constant binary names (mysql, redis-cli, etc.) and literal argument
// lists — no HTTP-controlled input reaches this function.
func runCmdQuick(name string, args ...string) ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
// #nosec G204 -- see function-level comment: constant names/args only.
out, err := exec.CommandContext(ctx, name, args...).Output()
if ctx.Err() == context.DeadlineExceeded {
return nil, fmt.Errorf("command timed out: %s", name)
}
return out, err
}
// isPHPWorkerCmdline reports whether a /proc cmdline belongs to a PHP process
// that serves requests.
//
// Two forms exist across the supported stacks: LiteSpeed spawns lsphp, and
// cPanel EA4 on Apache runs php-fpm, whose workers retitle themselves
// "php-fpm: pool <name>" and run as the account user. Matching only lsphp
// reported zero PHP activity on every Apache host.
//
// The php-fpm master is excluded on purpose: it runs as root and serves no
// requests, so counting it would attribute per-account load to root.
func isPHPWorkerCmdline(cmdline string) bool {
if strings.HasPrefix(cmdline, "php-fpm: pool ") {
return true
}
if strings.HasPrefix(cmdline, "php-fpm: master") {
return false
}
return strings.Contains(cmdline, "lsphp")
}
// isMySQLServerCmdline reports whether a /proc cmdline is the database server
// itself, as opposed to a client, a wrapper script, or a backup tool.
//
// The pid file used to be read from a hardcoded path, which does not exist on
// the cPanel/MariaDB hosts CSM primarily targets: MariaDB writes
// /var/lib/mysql/<host>.pid instead. Matching the process avoids maintaining a
// list of per-distribution pid paths.
func isMySQLServerCmdline(cmdline string) bool {
fields := strings.Fields(cmdline)
if len(fields) == 0 {
return false
}
base := filepath.Base(fields[0])
// A shell running mysqld_safe has the shell as argv[0]; neither it nor the
// wrapper is the server process.
return base == "mysqld" || base == "mariadbd"
}
// mysqlServerRSSMB returns the resident set size of the running database
// server in MB, or nil when no server process is found.
func mysqlServerRSSMB() *uint64 {
cmdlinePaths, _ := filepath.Glob("/proc/[0-9]*/cmdline")
for _, cmdPath := range cmdlinePaths {
// #nosec G304 -- cmdPath from /proc/*/cmdline glob; kernel pseudo-FS.
data, err := os.ReadFile(cmdPath)
if err != nil {
continue
}
if !isMySQLServerCmdline(strings.ReplaceAll(string(data), "\x00", " ")) {
continue
}
// #nosec G304 -- /proc/<pid>/status; kernel pseudo-FS.
statusData, _ := os.ReadFile(filepath.Join(filepath.Dir(cmdPath), "status"))
for _, line := range strings.Split(string(statusData), "\n") {
if !strings.HasPrefix(line, "VmRSS:") {
continue
}
fields := strings.Fields(line)
if len(fields) >= 2 {
if kb, perr := strconv.ParseUint(fields[1], 10, 64); perr == nil {
mb := kb / 1024
return &mb
}
}
break
}
}
return nil
}
// sampleMetrics gathers live system metrics and returns a populated perfMetrics.
func sampleMetrics() *perfMetrics {
m := &perfMetrics{}
// Load averages
if data, err := os.ReadFile("/proc/loadavg"); err == nil {
fields := strings.Fields(string(data))
if len(fields) >= 3 {
for i := 0; i < 3; i++ {
v, _ := strconv.ParseFloat(fields[i], 64)
m.LoadAvg[i] = v
}
}
}
// CPU cores
m.CPUCores = cachedCores()
// Memory from /proc/meminfo
{
f, err := os.Open("/proc/meminfo")
if err == nil {
var memTotal, memAvail, memFree, memBuffers, memCached, swapTotal, swapFree uint64
scanner := bufio.NewScanner(f)
for scanner.Scan() {
line := scanner.Text()
fields := strings.Fields(line)
if len(fields) < 2 {
continue
}
val, _ := strconv.ParseUint(fields[1], 10, 64)
switch fields[0] {
case "MemTotal:":
memTotal = val
case "MemAvailable:":
memAvail = val
case "MemFree:":
memFree = val
case "Buffers:":
memBuffers = val
case "Cached:":
memCached = val
case "SwapTotal:":
swapTotal = val
case "SwapFree:":
swapFree = val
}
}
_ = f.Close()
m.MemTotalMB = memTotal / 1024
m.MemAvailMB = memAvail / 1024
// Used = Total - Free - Buffers - Cached
used := memTotal
if memFree+memBuffers+memCached <= memTotal {
used = memTotal - memFree - memBuffers - memCached
}
m.MemUsedMB = used / 1024
m.SwapTotalMB = swapTotal / 1024
if swapFree <= swapTotal {
m.SwapUsedMB = (swapTotal - swapFree) / 1024
}
}
}
// PHP processes: scan /proc/*/cmdline for PHP request workers.
{
cmdlinePaths, _ := filepath.Glob("/proc/[0-9]*/cmdline")
userCounts := make(map[string]int)
total := 0
for _, cmdPath := range cmdlinePaths {
// #nosec G304 -- cmdPath from /proc/*/cmdline glob; kernel pseudo-FS.
data, err := os.ReadFile(cmdPath)
if err != nil {
continue
}
cmdStr := strings.ReplaceAll(string(data), "\x00", " ")
if !isPHPWorkerCmdline(cmdStr) {
continue
}
pid := filepath.Base(filepath.Dir(cmdPath))
// #nosec G304 -- /proc/<pid>/status; kernel pseudo-FS, pid from /proc glob.
statusData, _ := os.ReadFile(filepath.Join("/proc", pid, "status"))
uid := ""
for _, line := range strings.Split(string(statusData), "\n") {
if strings.HasPrefix(line, "Uid:\t") {
f := strings.Fields(strings.TrimPrefix(line, "Uid:\t"))
if len(f) > 0 {
uid = f[0]
}
break
}
}
if uid == "" {
uid = "unknown"
}
username := cachedUID(uid)
userCounts[username]++
total++
}
m.PHPProcs = total
// Build sorted top-10 list
type up struct {
user string
count int
}
var ups []up
for u, c := range userCounts {
ups = append(ups, up{u, c})
}
sort.Slice(ups, func(i, j int) bool {
return ups[i].count > ups[j].count
})
if len(ups) > 10 {
ups = ups[:10]
}
m.TopPHPUsers = make([]userProcs, len(ups))
for i, u := range ups {
m.TopPHPUsers[i] = userProcs{User: u.user, Count: u.count}
}
}
// MySQL: PID -> VmRSS, plus Threads_connected. Both fields stay nil
// when the lookup fails so the webui can show "n/a" instead of a
// misleading 0.
{
m.MySQLMemMB = mysqlServerRSSMB()
// Connection count. mysqlclient open returns nil on auth failure,
// missing socket, or absent server -- in every such case we
// leave MySQLConns nil rather than reporting a fake 0.
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
rows, err := mysqlclient.RootQuery(ctx, "SHOW STATUS LIKE 'Threads_connected'")
cancel()
if err == nil && len(rows) > 0 {
fields := strings.Fields(rows[0])
if len(fields) >= 2 {
if n, perr := strconv.Atoi(fields[1]); perr == nil {
m.MySQLConns = &n
}
}
}
}
// Redis: memory + keyspace via in-process client (no redis-cli fork).
{
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if used, max, err := redisinfo.MemoryUsage(ctx); err == nil {
m.RedisMemMB = used / (1024 * 1024)
m.RedisMaxMB = max / (1024 * 1024)
}
if total, err := redisinfo.Keyspace(ctx); err == nil {
m.RedisKeys = total
}
}
// Uptime from /proc/uptime
{
data, err := os.ReadFile("/proc/uptime")
if err == nil {
fields := strings.Fields(string(data))
if len(fields) >= 1 {
secs, _ := strconv.ParseFloat(fields[0], 64)
m.UptimeSeconds = int64(secs)
}
}
}
return m
}
// perfSampleTTL is how long a metrics sample is served before the next
// request takes a new one. Var so tests can force a fresh sample.
var perfSampleTTL = 10 * time.Second
// perfSample is one metrics sample and when it was taken.
type perfSample struct {
metrics *perfMetrics
at time.Time
}
func (s *Server) storePerfSample(m *perfMetrics, at time.Time) {
s.perfSample.Store(&perfSample{metrics: m, at: at})
}
func (s *Server) freshPerfSample() (*perfMetrics, bool) {
p := s.perfSample.Load()
if p == nil || time.Since(p.at) >= perfSampleTTL {
return nil, false
}
return p.metrics, true
}
// currentPerfMetrics returns a sample no older than perfSampleTTL, taking
// one if needed. Requests that arrive while a sample is being taken wait for
// it instead of sampling again.
func (s *Server) currentPerfMetrics() *perfMetrics {
if m, ok := s.freshPerfSample(); ok {
return m
}
s.perfMu.Lock()
defer s.perfMu.Unlock()
if m, ok := s.freshPerfSample(); ok {
return m
}
sample := s.samplePerf
if sample == nil {
sample = sampleMetrics
}
m := sample()
s.storePerfSample(m, time.Now())
return m
}
// apiPerformance returns the latest performance snapshot plus perf_ findings.
func (s *Server) apiPerformance(w http.ResponseWriter, r *http.Request) {
limit := queryInt(r, "limit", 100)
if limit > 500 {
limit = 500
}
metrics := s.currentPerfMetrics()
latest := s.store.LatestFindings()
suppressions := s.store.LoadSuppressions()
var views []perfFindingView
for _, f := range latest {
if !strings.HasPrefix(f.Check, "perf_") {
continue
}
if s.store.IsSuppressed(f, suppressions) {
continue
}
firstSeen := f.Timestamp
lastSeen := f.Timestamp
if entry, ok := s.store.EntryForKey(f.Key()); ok {
firstSeen = entry.FirstSeen
lastSeen = entry.LastSeen
}
key := f.Key()
views = append(views, perfFindingView{
Severity: f.Severity.String(),
level: f.Severity,
Check: f.Check,
Message: f.Message,
Details: f.Details,
Key: key,
FirstSeen: firstSeen.UTC(),
LastSeen: lastSeen.UTC(),
})
}
// Sort by severity descending
sort.Slice(views, func(i, j int) bool {
return views[i].level > views[j].level
})
if len(views) > limit {
views = views[:limit]
}
writeJSON(w, perfResponse{
Metrics: metrics,
Findings: views,
})
}
// apiPerfFixErrorLog truncates an account-owned error_log identified by
// the perf_error_logs finding. Admin scope; CSRF enforced at the route.
func (s *Server) apiPerfFixErrorLog(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Path string `json:"path"`
Key string `json:"key"`
}
if err := decodeJSONBodyLimited(w, r, 1<<14, &req); err != nil {
writeJSONError(w, "invalid request body", http.StatusBadRequest)
return
}
if req.Path == "" {
writeJSONError(w, "path is required", http.StatusBadRequest)
return
}
res := checks.FixErrorLogBloatInRoots(req.Path, s.perfFixAllowedRoots())
if !res.Success {
writeRemediation(w, res)
return
}
s.dismissPerfFinding(req.Key)
s.auditLog(r, "perf_fix_error_log", req.Path, res.Description)
writeRemediation(w, res)
}
// apiPerfFixDisplayErrors disables display_errors in an account-owned
// .user.ini / php.ini / .htaccess identified by the perf_wp_config
// finding's Details field. Admin scope; CSRF enforced at the route.
//
//nolint:dupl // mirrors apiPerfFixErrorLog; separate handlers keep audit actions explicit.
func (s *Server) apiPerfFixDisplayErrors(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Path string `json:"path"`
Key string `json:"key"`
}
if err := decodeJSONBodyLimited(w, r, 1<<14, &req); err != nil {
writeJSONError(w, "invalid request body", http.StatusBadRequest)
return
}
if req.Path == "" {
writeJSONError(w, "path is required", http.StatusBadRequest)
return
}
res := checks.FixDisplayErrorsOnInRoots(req.Path, s.perfFixAllowedRoots())
if !res.Success {
writeRemediation(w, res)
return
}
s.dismissPerfFinding(req.Key)
s.auditLog(r, "perf_fix_display_errors", req.Path, res.Description)
writeRemediation(w, res)
}
// apiPerfFixWPCron disables WP-Cron in an account-owned wp-config.php
// identified by a perf_wp_cron finding and installs a per-user system cron
// that runs wp-cron.php. Admin scope; CSRF enforced at the route.
func (s *Server) apiPerfFixWPCron(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
Path string `json:"path"`
Key string `json:"key"`
}
if err := decodeJSONBodyLimited(w, r, 1<<14, &req); err != nil {
writeJSONError(w, "invalid request body", http.StatusBadRequest)
return
}
if req.Path == "" {
writeJSONError(w, "path is required", http.StatusBadRequest)
return
}
cfg := s.liveCfg()
options := checks.WPCronFixOptions{}
if cfg != nil {
options.IntervalMinutes = cfg.Performance.WPCronFix.IntervalMinutes
options.PHPBin = cfg.Performance.WPCronFix.PHPBin
}
res := checks.FixDisableWPCronInRoots(req.Path, checks.ResolveWPCronRoots(cfg), options)
if !res.Success {
writeRemediation(w, res)
return
}
s.dismissPerfFinding(req.Key)
s.auditLog(r, "perf_fix_wp_cron", req.Path, res.Description)
writeRemediation(w, res)
}
func (s *Server) perfFixAllowedRoots() []string {
cfg := s.liveCfg()
if cfg == nil {
return []string{"/home"}
}
return checks.ResolveWebRoots(cfg)
}
func (s *Server) dismissPerfFinding(key string) {
key = strings.TrimSpace(key)
if key == "" {
return
}
s.store.DismissFinding(key)
s.store.DismissLatestFinding(key)
}
// handlePerformance renders the performance dashboard page.
func (s *Server) handlePerformance(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "performance.html", nil)
}
package webui
import (
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/store"
)
var errNoStore = errors.New("store unavailable")
func nowUnix() int64 { return time.Now().UTC().Unix() }
// operatorKey returns a SHA-256 hex of the auth credential carried by r. It
// is used as the per-operator partition key for the preferences store. The
// store never sees the raw token; only its hash.
//
// Returns "" when no admin credential is present (handlers below run after
// requireAuth, so this normally only happens in tests that bypass middleware).
func (s *Server) operatorKey(r *http.Request) string {
if bearer, ok := s.bearerTokenWithScope(r, "admin"); ok {
return hashOperatorToken(bearer)
}
if cookie, ok := s.cookieTokenWithScope(r, "admin"); ok {
return hashOperatorToken(cookie)
}
return ""
}
func hashOperatorToken(raw string) string {
sum := sha256.Sum256([]byte(raw))
return hex.EncodeToString(sum[:])
}
const (
prefsNamespaceUser = "user"
prefsNamespaceViews = "views"
)
// userPrefsBlob is the JSON document the client posts to /api/v1/prefs/user.
// Fields are validated and clipped to a small enum where appropriate so an
// attacker cannot smuggle arbitrary template data into the layout.
type userPrefsBlob struct {
Density string `json:"density"`
Timezone string `json:"timezone"`
AutoRefresh string `json:"auto_refresh"`
TableColumns map[string][]string `json:"table_columns,omitempty"`
}
func sanitizeUserPrefs(in userPrefsBlob) userPrefsBlob {
out := userPrefsBlob{}
switch in.Density {
case "compact", "comfortable":
out.Density = in.Density
}
switch in.Timezone {
case "server", "local":
out.Timezone = in.Timezone
default:
// Allow IANA-shaped strings (e.g. "Europe/Bucharest"). Reject anything
// containing whitespace or control characters; the value gets reflected
// in JS where Intl.DateTimeFormat will reject malformed zones anyway.
if isIANAish(in.Timezone) {
out.Timezone = in.Timezone
}
}
switch in.AutoRefresh {
case "on", "off":
out.AutoRefresh = in.AutoRefresh
}
if len(in.TableColumns) > 0 {
cleaned := make(map[string][]string, len(in.TableColumns))
for k, vals := range in.TableColumns {
if !isSimpleIdent(k) || len(vals) > 64 {
continue
}
var v []string
for _, name := range vals {
if isSimpleIdent(name) {
v = append(v, name)
}
}
cleaned[k] = v
}
if len(cleaned) > 0 {
out.TableColumns = cleaned
}
}
return out
}
func isIANAish(s string) bool {
if s == "" || len(s) > 64 {
return false
}
for _, r := range s {
switch {
case r >= 'A' && r <= 'Z':
case r >= 'a' && r <= 'z':
case r >= '0' && r <= '9':
case r == '_' || r == '+' || r == '-' || r == '/':
default:
return false
}
}
return true
}
func isSimpleIdent(s string) bool {
if s == "" || len(s) > 64 {
return false
}
for _, r := range s {
switch {
case r >= 'A' && r <= 'Z':
case r >= 'a' && r <= 'z':
case r >= '0' && r <= '9':
case r == '_' || r == '-' || r == '.':
default:
return false
}
}
return true
}
// apiPrefsUser handles GET and PUT for the operator's user-pref blob.
func (s *Server) apiPrefsUser(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
s.handleGetUserPrefs(w, r)
case http.MethodPut:
s.handlePutUserPrefs(w, r)
default:
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
}
}
func (s *Server) handleGetUserPrefs(w http.ResponseWriter, r *http.Request) {
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
sdb := store.Global()
if sdb == nil {
writeJSON(w, userPrefsBlob{})
return
}
raw, err := sdb.GetOperatorPref(opkey, prefsNamespaceUser)
if err != nil {
writeJSONError(w, "Store error", http.StatusInternalServerError)
return
}
if raw == nil {
writeJSON(w, userPrefsBlob{})
return
}
var blob userPrefsBlob
if err := json.Unmarshal(raw, &blob); err != nil {
writeJSON(w, userPrefsBlob{})
return
}
writeJSON(w, sanitizeUserPrefs(blob))
}
func (s *Server) handlePutUserPrefs(w http.ResponseWriter, r *http.Request) {
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
var blob userPrefsBlob
if err := decodeJSONBodyLimited(w, r, store.MaxPrefBlobSize, &blob); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
clean := sanitizeUserPrefs(blob)
raw, err := json.Marshal(clean)
if err != nil {
writeJSONError(w, "Encoding failed", http.StatusInternalServerError)
return
}
sdb := store.Global()
if sdb == nil {
writeJSONError(w, "Store unavailable", http.StatusServiceUnavailable)
return
}
if err := sdb.PutOperatorPref(opkey, prefsNamespaceUser, raw); err != nil {
writeJSONError(w, "Store error", http.StatusInternalServerError)
return
}
// The stored preferences, with the action's ok flag next to them.
stored := map[string]interface{}{
"density": clean.Density,
"timezone": clean.Timezone,
"auto_refresh": clean.AutoRefresh,
}
if len(clean.TableColumns) > 0 {
stored["table_columns"] = clean.TableColumns
}
writeOK(w, stored)
}
// savedView represents one user-named filter combination for a page.
type savedView struct {
Name string `json:"name"`
Page string `json:"page"`
Params map[string]string `json:"params"`
Updated int64 `json:"updated"`
}
// savedViewResponse is a saved view as the API sends it. The store keeps
// updated as Unix seconds; the API sends it as an instant like every time.
type savedViewResponse struct {
Name string `json:"name"`
Page string `json:"page"`
Params map[string]string `json:"params"`
Updated time.Time `json:"updated,omitzero"`
}
const maxSavedViewsPerOperator = 200
// apiPrefsViews handles list (GET), upsert (PUT), and delete (DELETE) of
// saved filter views.
func (s *Server) apiPrefsViews(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
s.handleListSavedViews(w, r)
case http.MethodPut:
s.handlePutSavedView(w, r)
case http.MethodDelete:
s.handleDeleteSavedView(w, r)
default:
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
}
}
func (s *Server) loadSavedViews(opkey string) []savedView {
sdb := store.Global()
if sdb == nil {
return nil
}
raw, err := sdb.GetOperatorPref(opkey, prefsNamespaceViews)
if err != nil || raw == nil {
return nil
}
var views []savedView
if err := json.Unmarshal(raw, &views); err != nil {
return nil
}
return views
}
func (s *Server) saveSavedViews(opkey string, views []savedView) error {
sort.SliceStable(views, func(i, j int) bool {
if views[i].Page != views[j].Page {
return views[i].Page < views[j].Page
}
return views[i].Name < views[j].Name
})
raw, err := json.Marshal(views)
if err != nil {
return err
}
sdb := store.Global()
if sdb == nil {
return errNoStore
}
return sdb.PutOperatorPref(opkey, prefsNamespaceViews, raw)
}
func (s *Server) handleListSavedViews(w http.ResponseWriter, r *http.Request) {
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
page := strings.TrimSpace(r.URL.Query().Get("page"))
views := s.loadSavedViews(opkey)
out := make([]savedViewResponse, 0, len(views))
for _, v := range views {
if page != "" && v.Page != page {
continue
}
view := savedViewResponse{Name: v.Name, Page: v.Page, Params: v.Params}
if v.Updated > 0 {
view.Updated = time.Unix(v.Updated, 0).UTC()
}
out = append(out, view)
}
writeAll(w, out)
}
func (s *Server) handlePutSavedView(w http.ResponseWriter, r *http.Request) {
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
var body struct {
Name string `json:"name"`
Page string `json:"page"`
Params map[string]string `json:"params"`
}
if err := decodeJSONBodyLimited(w, r, store.MaxPrefBlobSize, &body); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
body.Name = strings.TrimSpace(body.Name)
body.Page = strings.TrimSpace(body.Page)
if !isSimpleIdent(body.Page) {
writeJSONError(w, "Invalid page", http.StatusBadRequest)
return
}
if body.Name == "" || len(body.Name) > 80 {
writeJSONError(w, "Name must be 1-80 characters", http.StatusBadRequest)
return
}
if !isPrintableLabel(body.Name) {
writeJSONError(w, "Name contains invalid characters", http.StatusBadRequest)
return
}
if len(body.Params) > 32 {
writeJSONError(w, "Too many params", http.StatusBadRequest)
return
}
cleanParams := make(map[string]string, len(body.Params))
for k, v := range body.Params {
if !isSimpleIdent(k) || len(v) > 256 {
writeJSONError(w, "Invalid param", http.StatusBadRequest)
return
}
cleanParams[k] = v
}
views := s.loadSavedViews(opkey)
now := nowUnix()
updated := false
for i := range views {
if views[i].Page == body.Page && views[i].Name == body.Name {
views[i].Params = cleanParams
views[i].Updated = now
updated = true
break
}
}
if !updated {
if len(views) >= maxSavedViewsPerOperator {
writeJSONError(w, "Saved view limit reached", http.StatusBadRequest)
return
}
views = append(views, savedView{
Name: body.Name,
Page: body.Page,
Params: cleanParams,
Updated: now,
})
}
if err := s.saveSavedViews(opkey, views); err != nil {
writeJSONError(w, "Store error", http.StatusInternalServerError)
return
}
writeOK(w, nil)
}
func (s *Server) handleDeleteSavedView(w http.ResponseWriter, r *http.Request) {
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
var body struct {
Name string `json:"name"`
Page string `json:"page"`
}
if err := decodeJSONBodyLimited(w, r, 8*1024, &body); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
views := s.loadSavedViews(opkey)
out := views[:0]
removed := false
for _, v := range views {
if v.Page == body.Page && v.Name == body.Name {
removed = true
continue
}
out = append(out, v)
}
if !removed {
writeJSONError(w, "View not found", http.StatusNotFound)
return
}
if err := s.saveSavedViews(opkey, out); err != nil {
writeJSONError(w, "Store error", http.StatusInternalServerError)
return
}
writeOK(w, nil)
}
func isPrintableLabel(s string) bool {
for _, r := range s {
if r < 0x20 || r == 0x7f {
return false
}
}
return true
}
package webui
import (
"fmt"
"net/http"
"github.com/pidginhost/csm/internal/mailfwd/intel"
"github.com/pidginhost/csm/internal/platform"
)
// selectQueueReporter picks the queue-composition source for the host. Only
// cPanel/exim is wired; other platforms get the empty reporter.
func selectQueueReporter() intel.QueueReporter {
if platform.Detect().IsCPanel() {
return intel.NewEximQueueSource()
}
return intel.EmptyQueueReporter{}
}
// selectQueueFlusher picks the backscatter-flush executor for the host. Only
// cPanel/exim is wired; other platforms expose the route as unavailable.
func selectQueueFlusher() intel.QueueFlusher {
if platform.Detect().IsCPanel() {
return intel.NewEximQueueFlusher()
}
return nil
}
// apiEmailFlushBackscatter handles POST /api/v1/email/queue/flush-backscatter.
// It removes only frozen null-sender messages -- undeliverable bounce
// backscatter -- from the exim queue. Mutating, so it runs under auth + CSRF
// and is audit-logged.
func (s *Server) apiEmailFlushBackscatter(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if s.queueFlusher == nil {
writeJSONError(w, "Mail queue flush not available on this host", http.StatusServiceUnavailable)
return
}
res, err := s.queueFlusher.FlushBackscatter()
// The attempt must be recorded even when the verification read fails and
// cannot confirm a removal. A target can also leave through concurrent
// delivery, so the audit records observed queue state without attributing
// every disappearance to the removal command.
targeted := res.Targeted
if targeted > 0 {
s.auditLog(r, "email_flush_backscatter", "mail-queue",
fmt.Sprintf("removal requested for %d frozen null-sender message(s); %d confirmed no longer queued afterward",
targeted, res.Removed))
}
if err != nil {
message := fmt.Sprintf("Failed to flush backscatter: %v", err)
if targeted > 0 {
message = fmt.Sprintf("Flush incomplete: %d of %d targeted message(s) confirmed no longer queued: %v",
res.Removed, targeted, err)
}
writeJSONStatus(w, http.StatusInternalServerError, map[string]interface{}{"error": message, "removed": res.Removed})
return
}
writeOK(w, map[string]interface{}{"removed": res.Removed})
}
// apiEmailQueueComposition handles GET /api/v1/email/queue-composition and
// returns the makeup of the exim queue: real mail vs null-sender bounce
// backscatter, frozen count, oldest age, and the most-stuck recipients.
func (s *Server) apiEmailQueueComposition(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
if s.queueReporter == nil {
writeJSON(w, intel.QueueComposition{TopRecipients: []intel.RecipientCount{}})
return
}
comp, err := s.queueReporter.Composition()
if err != nil {
writeJSONError(w, "Failed to read mail queue", http.StatusInternalServerError)
return
}
writeJSON(w, comp)
}
package webui
import (
"sort"
"sync/atomic"
"time"
)
// apiRateLimitMaxIPs caps how many source addresses the unauthenticated
// per-IP rate-limit maps hold. The five-minute prune bounded them over time
// but not between prunes: a scan from many addresses grew them without
// limit in the meantime.
const apiRateLimitMaxIPs = 10000
// rateLimitSweeps counts full passes over a rate-limit map; tests read it.
var rateLimitSweeps atomic.Int64
// boundRateLimitMap keeps m under apiRateLimitMaxIPs before a new address
// is inserted. Caller holds the map's mutex. A full map is swept once down
// to nine tenths of the ceiling: stale entries (no hit after cutoff) go
// first, then the addresses with the oldest last hit. Sweeping one address
// at a time scanned the whole map for every new address, so a flood of new
// sources cost a full scan per request under the lock.
func boundRateLimitMap(m map[string][]time.Time, cutoff time.Time) {
if len(m) < apiRateLimitMaxIPs {
return
}
rateLimitSweeps.Add(1)
type lastHit struct {
ip string
last time.Time
}
live := make([]lastHit, 0, len(m))
for ip, hits := range m {
if len(hits) == 0 || !hits[len(hits)-1].After(cutoff) {
delete(m, ip)
continue
}
live = append(live, lastHit{ip, hits[len(hits)-1]})
}
keep := apiRateLimitMaxIPs - apiRateLimitMaxIPs/10
if len(live) <= keep {
return
}
sort.Slice(live, func(i, j int) bool { return live[i].last.Before(live[j].last) })
for _, h := range live[:len(live)-keep] {
delete(m, h.ip)
}
}
package webui
import (
"net/http"
"sort"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
)
type relayAbuseResponse struct {
Entries []relayAbuseEntry `json:"items"`
Total int `json:"total"`
Offset int `json:"offset"`
Limit int `json:"limit"`
From time.Time `json:"from"`
To time.Time `json:"to"`
Truncated bool `json:"truncated"`
}
type relayAbuseEntry struct {
Path string `json:"path"`
PathLabel string `json:"path_label"`
Severity string `json:"severity"`
SourceIP string `json:"source_ip,omitempty"`
CPUser string `json:"cp_user,omitempty"`
TriggerCount int `json:"trigger_count"`
DetectedAt time.Time `json:"detected_at"`
Sites []relaySiteEntry `json:"sites"`
MsgSample []string `json:"msg_sample,omitempty"`
}
type relaySiteEntry struct {
Site string `json:"site"`
Script string `json:"script"`
Hits int `json:"hits"`
LastSeen time.Time `json:"last_seen"`
SampleSubject string `json:"sample_subject,omitempty"`
}
const relayAbuseDefaultLimit = 20
const relayAbuseMaxLimit = 100
func relayPathLabel(path string) string {
switch path {
case "fanout":
return "Spam outbreak (IP fanout)"
case "volume":
return "High volume script"
case "header":
return "Suspicious headers"
case "volume_account":
return "High volume account"
case "":
return "Unknown path"
default:
return path
}
}
// apiEmailRelayAbuse handles GET /api/v1/email/relay-abuse. Read-only.
// Reads email_php_relay_abuse findings from persisted history (the realtime
// dispatch path does not populate LatestFindings).
func (s *Server) apiEmailRelayAbuse(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
q := r.URL.Query()
limit := queryInt(r, "limit", relayAbuseDefaultLimit)
if limit <= 0 || limit > relayAbuseMaxLimit {
limit = relayAbuseDefaultLimit
}
now := time.Now()
from, to, ok := historyRangeQuery(w, q, now.Add(-24*time.Hour), now)
if !ok {
return
}
if to.Before(from) {
from, to = to, from
}
writeJSON(w, s.emailMemo("relay?"+q.Encode(), func() any {
return s.buildRelayAbuseResponse(from, to, limit)
}))
}
func (s *Server) buildRelayAbuseResponse(from, to time.Time, limit int) relayAbuseResponse {
resp := relayAbuseResponse{
Entries: []relayAbuseEntry{},
Limit: limit,
From: from.UTC(),
To: to.UTC(),
}
if s.store == nil {
return resp
}
// Filter while walking newest-first history and cap the matches, so
// unrelated findings and findings newer than the range never hide a
// match. Truncation covers both this budget and the result limit below.
rows := s.store.SearchHistorySince(from, emailGroupsScanCap+1, func(f alert.Finding) bool {
return f.Check == "email_php_relay_abuse" && f.Timestamp.Before(to)
})
if len(rows) > emailGroupsScanCap {
rows = rows[:emailGroupsScanCap]
resp.Truncated = true
}
resp.Total = len(rows)
sort.SliceStable(rows, func(i, j int) bool {
if !rows[i].Timestamp.Equal(rows[j].Timestamp) {
return rows[i].Timestamp.After(rows[j].Timestamp)
}
if rows[i].Path != rows[j].Path {
return rows[i].Path < rows[j].Path
}
if rows[i].SourceIP != rows[j].SourceIP {
return rows[i].SourceIP < rows[j].SourceIP
}
return rows[i].CPUser < rows[j].CPUser
})
if len(rows) > limit {
resp.Truncated = true
rows = rows[:limit]
}
for _, f := range rows {
resp.Entries = append(resp.Entries, toRelayAbuseEntry(f))
}
return resp
}
func toRelayAbuseEntry(f alert.Finding) relayAbuseEntry {
e := relayAbuseEntry{
Path: f.Path,
PathLabel: relayPathLabel(f.Path),
Severity: f.Severity.String(),
SourceIP: f.SourceIP,
CPUser: f.CPUser,
TriggerCount: relayTriggerCount(f),
DetectedAt: f.Timestamp,
Sites: []relaySiteEntry{},
MsgSample: f.MsgIDs,
}
for _, h := range f.RelayBreakdown {
site, script := splitScriptKey(h.ScriptKey)
e.Sites = append(e.Sites, relaySiteEntry{
Site: site,
Script: script,
Hits: h.Hits,
LastSeen: h.LastSeen,
SampleSubject: h.SampleSubject,
})
}
return e
}
func relayTriggerCount(f alert.Finding) int {
if f.RelayTotal > 0 {
return f.RelayTotal
}
sum := 0
for _, h := range f.RelayBreakdown {
if h.Hits > 0 {
sum += h.Hits
}
}
if sum > 0 {
return sum
}
return len(f.MsgIDs)
}
// splitScriptKey splits a "host:/path" script key into host and path. The key
// is built as host + ":" + path with path starting at "/", so the delimiter is
// the ":/" boundary. Splitting there (not the first colon) keeps host:port and
// IPv6-literal hosts intact. A key without ":/" yields ("", key) so the row
// still renders.
func splitScriptKey(k string) (site, script string) {
if i := strings.Index(k, ":/"); i >= 0 {
return k[:i], k[i+1:]
}
return "", k
}
package webui
import (
"fmt"
"net/http"
"os"
"path/filepath"
"strings"
"github.com/pidginhost/csm/internal/signatures"
"github.com/pidginhost/csm/internal/yara"
)
func (s *Server) handleRules(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "rules.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
// GET /api/v1/rules/status
func (s *Server) apiRulesStatus(w http.ResponseWriter, _ *http.Request) {
cfg := s.liveCfg()
yamlCount := 0
yamlVersion := 0
if scanner := signatures.Global(); scanner != nil {
yamlCount = scanner.RuleCount()
yamlVersion = scanner.Version()
}
yaraCount := 0
if b := yara.Active(); b != nil {
yaraCount = b.RuleCount()
}
result := map[string]interface{}{
"yaml_rules": yamlCount,
"yara_rules": yaraCount,
"yara_available": yara.Available(),
"yaml_version": yamlVersion,
"rules_dir": cfg.Signatures.RulesDir,
"auto_update": cfg.Signatures.UpdateURL != "",
"update_url": cfg.Signatures.UpdateURL,
}
if secs, ok := durationSeconds(cfg.Signatures.UpdateInterval); ok {
result["update_interval_seconds"] = secs
}
writeJSON(w, result)
}
// GET /api/v1/rules/list
func (s *Server) apiRulesList(w http.ResponseWriter, _ *http.Request) {
rulesDir := s.liveCfg().Signatures.RulesDir
type ruleFileInfo struct {
Name string `json:"name"`
Type string `json:"type"` // "yaml" or "yara"
Size int64 `json:"size"`
}
var files []ruleFileInfo
entries, err := os.ReadDir(rulesDir)
if err != nil {
if os.IsNotExist(err) {
writeAll(w, files)
return
}
writeJSONError(w, fmt.Sprintf("reading rules directory: %v", err), http.StatusInternalServerError)
return
}
for _, entry := range entries {
if entry.IsDir() {
continue
}
name := entry.Name()
ext := strings.ToLower(filepath.Ext(name))
var fileType string
switch ext {
case ".yml", ".yaml":
fileType = "yaml"
case ".yar", ".yara":
fileType = "yara"
default:
continue // skip non-rule files
}
info, err := entry.Info()
if err != nil {
continue
}
files = append(files, ruleFileInfo{
Name: name,
Type: fileType,
Size: info.Size(),
})
}
writeAll(w, files)
}
// POST /api/v1/rules/reload
func (s *Server) apiRulesReload(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var yamlErr, yaraErr error
yamlCount := 0
yaraCount := 0
if scanner := signatures.Global(); scanner != nil {
yamlErr = scanner.Reload()
yamlCount = scanner.RuleCount()
// Update the cached sig count shown in dashboard/status
s.SetSigCount(yamlCount)
}
if b := yara.Active(); b != nil {
yaraErr = b.Reload()
yaraCount = b.RuleCount()
}
var errors []string
if yamlErr != nil {
errors = append(errors, fmt.Sprintf("YAML reload: %v", yamlErr))
}
if yaraErr != nil {
errors = append(errors, fmt.Sprintf("YARA reload: %v", yaraErr))
}
s.auditLog(r, "rules_reload", "signatures", fmt.Sprintf("errors: %d", len(errors)))
result := map[string]interface{}{
"yaml_rules": yamlCount,
"yara_rules": yaraCount,
}
if len(errors) > 0 {
result["error"] = strings.Join(errors, "; ")
result["errors"] = errors
writeJSONStatus(w, http.StatusInternalServerError, result)
return
}
writeOK(w, result)
}
package webui
import (
"fmt"
"net/http"
"strings"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/control"
"github.com/pidginhost/csm/internal/store"
)
// ScanJobController is the subset of the daemon's ScanJobManager the WebUI
// needs to enqueue and cancel full-scan jobs. The daemon injects a concrete
// *ScanJobManager via SetScanJobs; nil until wired (POST endpoints → 503).
type ScanJobController interface {
Enqueue(scope, target string, opts checks.AccountScanOptions, quarantine bool) (string, error)
Cancel(id string) error
}
// apiScanJobsList handles GET /api/v1/scan-jobs.
// Returns all scan-job records, newest-first.
func (s *Server) apiScanJobsList(w http.ResponseWriter, _ *http.Request) {
db := store.Global()
if db == nil {
writeJSONError(w, "store unavailable", http.StatusServiceUnavailable)
return
}
jobs, err := db.ListScanJobs()
if err != nil {
writeJSONError(w, "failed to list scan jobs: "+err.Error(), http.StatusInternalServerError)
return
}
writeAll(w, jobs)
}
// apiScanJobsRouter handles /api/v1/scan-jobs/{id} and
// /api/v1/scan-jobs/{id}/findings.
//
// - GET /api/v1/scan-jobs/{id} → job detail
// - GET /api/v1/scan-jobs/{id}/findings → paginated findings
//
// Only GET is accepted; other methods return 405.
func (s *Server) apiScanJobsRouter(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
db := store.Global()
if db == nil {
writeJSONError(w, "store unavailable", http.StatusServiceUnavailable)
return
}
tail := strings.TrimPrefix(r.URL.Path, "/api/v1/scan-jobs/")
if tail == "" {
writeJSONError(w, "scan job id required", http.StatusBadRequest)
return
}
parts := strings.SplitN(tail, "/", 2)
id := parts[0]
sub := ""
if len(parts) == 2 {
sub = parts[1]
}
switch sub {
case "":
s.apiScanJobDetail(w, r, db, id)
case "findings":
s.apiScanJobFindings(w, r, db, id)
default:
writeJSONError(w, "not found", http.StatusNotFound)
}
}
// apiScanJobDetail handles GET /api/v1/scan-jobs/{id}.
func (s *Server) apiScanJobDetail(w http.ResponseWriter, _ *http.Request, db *store.DB, id string) {
rec, ok, err := db.GetScanJob(id)
if err != nil {
writeJSONError(w, "failed to get scan job: "+err.Error(), http.StatusInternalServerError)
return
}
if !ok {
writeJSONError(w, "scan job not found", http.StatusNotFound)
return
}
writeJSON(w, map[string]any{"job": rec})
}
// A full-server scan job can record tens of thousands of findings, so the
// findings endpoint pages them.
const (
scanJobFindingsDefaultLimit = 500
scanJobFindingsMaxLimit = 5000
)
// apiScanJobFindings handles GET /api/v1/scan-jobs/{id}/findings.
// Query params: offset (default 0), limit (default 500, at most 5000).
func (s *Server) apiScanJobFindings(w http.ResponseWriter, r *http.Request, db *store.DB, id string) {
offset := queryInt(r, "offset", 0)
limit := queryInt(r, "limit", scanJobFindingsDefaultLimit)
if limit <= 0 {
limit = scanJobFindingsDefaultLimit
}
if limit > scanJobFindingsMaxLimit {
limit = scanJobFindingsMaxLimit
}
findings, total, err := db.ListScanJobFindings(id, offset, limit)
if err != nil {
writeJSONError(w, "failed to list findings: "+err.Error(), http.StatusInternalServerError)
return
}
writeItems(w, toAPIFindings(findings), map[string]any{
"job_id": id,
"total": total,
"offset": offset,
"limit": limit,
"truncated": offset+len(findings) < total,
})
}
// scanJobEnqueueBody is the JSON request body for POST /api/v1/scan-jobs.
type scanJobEnqueueBody struct {
Scope string `json:"scope"`
Target string `json:"target"`
RespectIgnores bool `json:"respect_ignores"`
Quarantine bool `json:"quarantine"`
}
// apiScanJobsEnqueue handles POST /api/v1/scan-jobs.
// Requires admin auth + CSRF (enforced by the mux registration).
func (s *Server) apiScanJobsEnqueue(w http.ResponseWriter, r *http.Request) {
if s.scanJobs == nil {
writeJSONError(w, "scan job manager not available", http.StatusServiceUnavailable)
return
}
var body scanJobEnqueueBody
if err := decodeJSONBodyLimited(w, r, 4*1024, &body); err != nil {
writeJSONError(w, "invalid request body: "+err.Error(), http.StatusBadRequest)
return
}
opts := checks.FullScanOptions(s.liveCfg(), body.RespectIgnores)
switch body.Scope {
case "account":
if body.Target == "" || !control.ValidScanAccountTarget(body.Target) {
writeJSONError(w, "invalid or missing account target", http.StatusBadRequest)
return
}
id, err := s.scanJobs.Enqueue("account", body.Target, opts, body.Quarantine)
if err != nil {
msg := "enqueue failed: " + err.Error()
status := http.StatusInternalServerError
if strings.Contains(err.Error(), "queue is full") {
status = http.StatusConflict
}
writeJSONError(w, msg, status)
return
}
s.auditLog(r, "scan_job_enqueue", body.Target, fmt.Sprintf("job %s, account scan, quarantine=%v", id, body.Quarantine))
// The scan runs after the response: 202 with the job to poll.
writeOKStatus(w, http.StatusAccepted, map[string]interface{}{"job_id": id, "state": "queued"})
case "all":
if body.Quarantine {
writeJSONError(w, "quarantine is not supported with scope \"all\"", http.StatusBadRequest)
return
}
id, err := s.scanJobs.Enqueue("all", "all", opts, false)
if err != nil {
msg := "enqueue failed: " + err.Error()
status := http.StatusInternalServerError
if strings.Contains(err.Error(), "queue is full") {
status = http.StatusConflict
}
writeJSONError(w, msg, status)
return
}
s.auditLog(r, "scan_job_enqueue", "all", fmt.Sprintf("job %s, full scan", id))
// The scan runs after the response: 202 with the job to poll.
writeOKStatus(w, http.StatusAccepted, map[string]interface{}{"job_id": id, "state": "queued"})
default:
writeJSONError(w, "unsupported scope: must be \"account\" or \"all\"", http.StatusBadRequest)
}
}
// apiScanJobsCancel handles POST /api/v1/scan-jobs/{id}/cancel.
// The path arriving here is everything after /api/v1/scan-jobs/, e.g.
// "sj-abc123/cancel". The tail must end exactly in "/<id>/cancel".
// Requires admin auth + CSRF (enforced by the mux registration).
func (s *Server) apiScanJobsCancel(w http.ResponseWriter, r *http.Request) {
if s.scanJobs == nil {
writeJSONError(w, "scan job manager not available", http.StatusServiceUnavailable)
return
}
// Path: /api/v1/scan-jobs/{id}/cancel
tail := strings.TrimPrefix(r.URL.Path, "/api/v1/scan-jobs/")
if !strings.HasSuffix(tail, "/cancel") {
writeJSONError(w, "not found", http.StatusNotFound)
return
}
id := strings.TrimSuffix(tail, "/cancel")
if id == "" {
writeJSONError(w, "scan job id required", http.StatusNotFound)
return
}
if err := s.scanJobs.Cancel(id); err != nil {
msg := err.Error()
status := http.StatusConflict
if strings.Contains(msg, "not found") {
status = http.StatusNotFound
}
writeJSONError(w, "cancel failed: "+msg, status)
return
}
s.auditLog(r, "scan_job_cancel", id, "cancel requested")
// The job stops after the response: 202 with its state.
writeOKStatus(w, http.StatusAccepted, map[string]interface{}{"job_id": id, "state": "canceling"})
}
package webui
import (
"context"
"crypto/hmac"
"crypto/sha256"
"crypto/subtle"
"crypto/tls"
"encoding/hex"
"fmt"
"html/template"
"io/fs"
"net"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/attackdb"
"github.com/pidginhost/csm/internal/broadcast"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/emailav"
"github.com/pidginhost/csm/internal/geoip"
"github.com/pidginhost/csm/internal/health"
"github.com/pidginhost/csm/internal/incident"
"github.com/pidginhost/csm/internal/mailfwd/intel"
"github.com/pidginhost/csm/internal/mailfwd/inventory"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/session"
"github.com/pidginhost/csm/internal/state"
persiststore "github.com/pidginhost/csm/internal/store"
)
// IPBlocker abstracts the firewall engine for block/unblock operations.
type IPBlocker interface {
BlockIP(ip string, reason string, timeout time.Duration) error
UnblockIP(ip string) error
}
// blockIPForOperator calls BlockIPForce when the blocker supports it (engine
// on live systems), otherwise falls back to BlockIP (test stubs). This ensures
// operator-initiated blocks from the Web UI are never silenced by dry_run.
// Every attempt reports into the shared firewall outcome metric.
func blockIPForOperator(b IPBlocker, ip, reason string, timeout time.Duration) error {
var err error
if fb, ok := b.(forceBlocker); ok {
err = fb.BlockIPForce(ip, reason, timeout)
} else {
err = b.BlockIP(ip, reason, timeout)
}
checks.ObserveOperatorBlock(err, checks.BlockSourceWebUI)
return err
}
// noListDir wraps an http.FileSystem so http.FileServer cannot serve a
// directory index. Opening a directory returns fs.ErrNotExist, which the
// FileServer turns into a 404, while individual files are served normally.
// This keeps the unauthenticated /static/ assets reachable for the login
// page without letting anyone enumerate the shipped file set.
type noListDir struct{ fs http.FileSystem }
func (d noListDir) Open(name string) (http.File, error) {
f, err := d.fs.Open(name)
if err != nil {
return nil, err
}
info, err := f.Stat()
if err != nil {
_ = f.Close()
return nil, err
}
if info.IsDir() {
_ = f.Close()
return nil, fs.ErrNotExist
}
return f, nil
}
// serverWriteTimeout bounds a response. Handlers that run longer (account
// scans, the hardening audit) extend their own deadline to longRequestTimeout.
const (
serverWriteTimeout = 300 * time.Second
longRequestTimeout = 10 * time.Minute
)
// Server is the web UI HTTP server. Serves API always; serves HTML pages
// and static files only if the UI directory exists on disk.
type Server struct {
sessions *session.Manager
sessionNow func() time.Time
cfg *config.Config
store *state.Store
httpSrv *http.Server
templates map[string]*template.Template
hasUI bool // true if UI directory with templates exists
uiDir string // path to UI directory on disk
staticDir string // uiDir/static
assets assetVersions
startTime time.Time
sigCount int // loaded signature rule count
fanotifyActive func() bool // live daemon reader; nil outside the daemon
logWatcherCount func() int // live daemon reader; nil outside the daemon
blocker IPBlocker
geoIPDB atomic.Pointer[geoip.DB]
// emailQuarantine and emailAVWatcherMode are installed by the daemon
// after the listener already serves; request goroutines read them, so
// they are held atomically and read through the accessors below.
emailQuarantine atomic.Pointer[emailav.Quarantine]
emailAVWatcherMode atomic.Value // string
forwarderSource inventory.Source
deferralReporter intel.Reporter
queueReporter intel.QueueReporter
queueFlusher intel.QueueFlusher
forwardHeld heldForwardStore
version string
// Host metrics are sampled when the Performance page asks, at most once
// per perfSampleTTL; perfMu makes concurrent requests share a sample.
perfSample atomic.Pointer[perfSample]
perfMu sync.Mutex
samplePerf func() *perfMetrics
incidentCorrelator *incident.Correlator
// Rate limiting
threatActionMu sync.Mutex // serialize operator block, clear and undo decisions
loginMu sync.Mutex
loginAttempts map[string][]time.Time
apiMu sync.Mutex
apiRequests map[string][]time.Time // per-IP API rate limiting
scanMu sync.Mutex
scanRunning bool // only one scan at a time
auditMu sync.Mutex // serializes UI audit rotation and appends
modSecApplyMu sync.Mutex // serializes modsec rules apply (write+reload+rollback)
sigCountMu sync.RWMutex
settingsSaveHook func()
// verifyFinding is per server so handler tests can inject a verdict without
// replacing process-wide behavior while another server is handling a request.
verifyFinding func(checks.VerifyInput) checks.VerifyResult
// applyFix is per server for the same reason: handler tests observe a
// fix outcome without touching the host.
applyFix func(ctx context.Context, check, message, details string, filePath ...string) checks.RemediationResult
// scanInProgress reports scans the UI did not start; per server so
// handler tests can hold one without running checks.
scanInProgress func() bool
// Results computed from recent history, reused while it is unchanged.
statsMemo historyMemo
timelineMemo historyMemo
emailMemos historyMemos
// accountRoots and scanAccounts are the platform's account inventory;
// per server so handler tests can supply their own tree.
accountRoots func() []string
scanAccounts func(*config.Config) ([]string, error)
provider health.Provider // set by Daemon when it starts the WebUI
mu sync.RWMutex
findingBus *broadcast.Bus // set by Daemon via SetFindingBus
// Graceful shutdown signal for background goroutines and streaming handlers.
// shutdownOnce makes Shutdown idempotent; closing pruneDone twice panics.
pruneDone chan struct{}
shutdownOnce sync.Once
// restartDaemon is called by apiSettingsRestart. Tests override this.
restartDaemon func() (output []byte, err error)
// verifiedBotsReloader pushes a saved verified_bots list into the live
// bot registry + verifier without a restart. Set by the Daemon; nil in tests.
verifiedBotsReloader func() error
// scanJobs is the full-scan job manager injected by the Daemon via
// SetScanJobs. Nil until wired; POST scan-job handlers return 503 when nil.
scanJobs ScanJobController
}
// liveCfg returns the configuration a handler should act on: the last
// reloaded one, or the startup snapshot before any reload. The server keeps
// s.cfg for restart-required wiring (listener, TLS, tokens, state paths);
// thresholds, firewall policy, scan options and alert routing change under a
// SIGHUP and must be read here, or the UI shows and applies stale settings.
func (s *Server) liveCfg() *config.Config {
if live := config.Active(); live != nil {
return live
}
return s.cfg
}
// New creates a new web UI server.
func New(cfg *config.Config, store *state.Store) (*Server, error) {
s := &Server{
cfg: cfg,
store: store,
startTime: time.Now(),
loginAttempts: make(map[string][]time.Time),
apiRequests: make(map[string][]time.Time),
pruneDone: make(chan struct{}),
forwarderSource: selectForwarderSource(),
deferralReporter: selectDeferralReporter(),
queueReporter: selectQueueReporter(),
queueFlusher: selectQueueFlusher(),
forwardHeld: selectForwardHeld(),
verifyFinding: checks.VerifyFindingInput,
applyFix: checks.ApplyFix,
scanInProgress: checks.ScanInProgress,
samplePerf: sampleMetrics,
accountRoots: checks.AccountHomeRoots,
scanAccounts: checks.EnumerateScanAccounts,
}
lifetime, idle, err := cfg.BrowserSessionDurations()
if err != nil {
return nil, err
}
s.sessionNow = time.Now
if db := persiststore.Global(); db != nil {
s.sessions, err = session.New(db, lifetime, idle)
if err != nil {
return nil, fmt.Errorf("initialize browser sessions: %w", err)
}
}
// Check if UI directory exists on disk
s.uiDir = cfg.WebUI.UIDir
if s.uiDir == "" {
s.uiDir = "/opt/csm/ui"
}
funcMap := s.templateFuncs()
// Try to load templates from disk
templateDir := filepath.Join(s.uiDir, "templates")
staticDir := filepath.Join(s.uiDir, "static")
s.staticDir = staticDir
if _, err := os.Stat(templateDir); err == nil {
s.templates = make(map[string]*template.Template)
layoutPath := filepath.Join(templateDir, "layout.html")
for _, page := range []string{"dashboard", "findings", "quarantine", "cleanup-history", "firewall", "modsec", "modsec-rules", "verified-bots", "threat", "rules", "audit", "account", "incident", "email", "performance", "hardening", "settings", "sessions"} {
pagePath := filepath.Join(templateDir, page+".html")
t, err := template.New(page+".html").Funcs(funcMap).ParseFiles(layoutPath, pagePath)
if err != nil {
return nil, fmt.Errorf("parsing template %s from %s: %w", page, templateDir, err)
}
s.templates[page+".html"] = t
}
loginPath := filepath.Join(templateDir, "login.html")
loginTmpl, err := template.New("login.html").Funcs(funcMap).ParseFiles(loginPath)
if err != nil {
return nil, fmt.Errorf("parsing login template: %w", err)
}
s.templates["login.html"] = loginTmpl
s.hasUI = true
fmt.Fprintf(os.Stderr, "WebUI: loaded templates from %s\n", templateDir)
} else {
fmt.Fprintf(os.Stderr, "WebUI: UI directory not found at %s - running in API-only mode\n", s.uiDir)
}
// Set up routes
mux := http.NewServeMux()
// Static files and HTML pages - only if UI directory exists
if s.hasUI {
// Static assets must stay reachable pre-auth (the login page loads its
// own CSS/JS), so they are not behind requireAuth. They must not be
// enumerable, though: noListDir makes directory requests 404 instead of
// returning an index listing of every shipped file.
mux.Handle("/static/", s.staticHandler(staticDir))
mux.HandleFunc("/login", s.handleLogin)
mux.Handle("/", s.requireAuth(http.HandlerFunc(s.handleDashboard)))
mux.Handle("/dashboard", s.requireAuth(http.HandlerFunc(s.handleDashboard)))
mux.Handle("/findings", s.requireAuth(http.HandlerFunc(s.handleFindings)))
mux.Handle("/history", s.requireAuth(http.HandlerFunc(s.handleHistoryRedirect)))
mux.Handle("/quarantine", s.requireAuth(http.HandlerFunc(s.handleQuarantine)))
mux.Handle("/cleanup-history", s.requireAuth(http.HandlerFunc(s.handleCleanupHistory)))
mux.Handle("/blocked", s.requireAuth(http.HandlerFunc(s.handleBlockedRedirect)))
mux.Handle("/firewall", s.requireAuth(http.HandlerFunc(s.handleFirewall)))
mux.Handle("/threat", s.requireAuth(http.HandlerFunc(s.handleThreat)))
mux.Handle("/rules", s.requireAuth(http.HandlerFunc(s.handleRules)))
mux.Handle("/audit", s.requireAuth(http.HandlerFunc(s.handleAudit)))
mux.Handle("/account", s.requireAuth(http.HandlerFunc(s.handleAccount)))
mux.Handle("/incident", s.requireAuth(http.HandlerFunc(s.handleIncident)))
mux.Handle("/email", s.requireAuth(http.HandlerFunc(s.handleEmail)))
mux.Handle("/performance", s.requireAuth(http.HandlerFunc(s.handlePerformance)))
mux.Handle("/hardening", s.requireAuth(http.HandlerFunc(s.handleHardening)))
mux.Handle("/settings", s.requireAuth(http.HandlerFunc(s.handleSettings)))
mux.Handle("GET /sessions", s.requireAuth(http.HandlerFunc(s.handleSessions)))
mux.Handle("POST /sessions/revoke", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.handleSessionRevoke))))
mux.Handle("/modsec", s.requireAuth(http.HandlerFunc(s.handleModSec)))
mux.Handle("/modsec/rules", s.requireAuth(http.HandlerFunc(s.handleModSecRules)))
mux.Handle("/verified-bots", s.requireAuth(http.HandlerFunc(s.handleVerifiedBots)))
}
// Auth-protected API - read (read-scope tokens accepted)
// Any /api/ path no route below matches; see apiNotFound.
mux.Handle("/api/", http.HandlerFunc(s.apiNotFound))
mux.Handle("/api/v1/events", s.requireRead(http.HandlerFunc(s.apiEvents)))
mux.Handle("/api/v1/status", s.requireRead(http.HandlerFunc(s.apiStatus)))
mux.Handle("/api/v1/challenge/stats", s.requireRead(http.HandlerFunc(s.apiChallengeStats)))
mux.Handle("/api/v1/scan-jobs", s.requireRead(http.HandlerFunc(s.apiScanJobsList)))
mux.Handle("/api/v1/scan-jobs/", s.requireRead(http.HandlerFunc(s.apiScanJobsRouter)))
// POST /api/v1/scan-jobs — enqueue a new full-scan job (admin + CSRF).
// Go 1.22 method-qualified patterns take precedence over the unqualified
// prefix above, so GET /api/v1/scan-jobs still goes to the read handler.
mux.Handle("POST /api/v1/scan-jobs", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiScanJobsEnqueue))))
// POST /api/v1/scan-jobs/{id}/cancel — cancel a queued or running job.
mux.Handle("POST /api/v1/scan-jobs/", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiScanJobsCancel))))
mux.Handle("/api/v1/findings", s.requireRead(http.HandlerFunc(s.apiFindings)))
mux.Handle("/api/v1/findings/enriched", s.requireRead(http.HandlerFunc(s.apiFindingsEnriched)))
mux.Handle("/api/v1/history", s.requireRead(http.HandlerFunc(s.apiHistory)))
mux.Handle("/api/v1/stats", s.requireRead(http.HandlerFunc(s.apiStats)))
mux.Handle("/api/v1/stats/trend", s.requireRead(http.HandlerFunc(s.apiStatsTrend)))
mux.Handle("/api/v1/stats/timeline", s.requireRead(http.HandlerFunc(s.apiStatsTimeline)))
mux.Handle("/api/v1/blocked-ips", s.requireRead(http.HandlerFunc(s.apiBlockedIPs)))
mux.Handle("/api/v1/capabilities", s.requireRead(http.HandlerFunc(s.apiCapabilities)))
mux.Handle("/api/v1/health", s.requireRead(http.HandlerFunc(s.apiHealth)))
mux.Handle("/api/v1/components", s.requireRead(http.HandlerFunc(s.apiComponents)))
// Auth-protected API - admin-only reads (data with write-adjacent sensitivity)
mux.Handle("/api/v1/quarantine", s.requireAuth(getOnly(http.HandlerFunc(s.apiQuarantine))))
mux.Handle("/api/v1/modsec/stats", s.requireRead(http.HandlerFunc(s.apiModSecStats)))
mux.Handle("/api/v1/modsec/blocks", s.requireRead(http.HandlerFunc(s.apiModSecBlocks)))
mux.Handle("/api/v1/modsec/events", s.requireRead(http.HandlerFunc(s.apiModSecEvents)))
mux.Handle("/api/v1/modsec/rules", s.requireAuth(getOnly(http.HandlerFunc(s.apiModSecRules))))
mux.Handle("/api/v1/modsec/rules/apply", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiModSecRulesApply))))
mux.Handle("/api/v1/modsec/rules/escalation", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiModSecRulesEscalation))))
mux.Handle("/api/v1/verified-bots", s.requireAuth(http.HandlerFunc(s.apiVerifiedBots)))
mux.Handle("/api/v1/verified-bots/apply", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiVerifiedBotsApply))))
mux.Handle("/api/v1/accounts", s.requireAuth(getOnly(http.HandlerFunc(s.apiAccounts))))
mux.Handle("/api/v1/account", s.requireAuth(getOnly(http.HandlerFunc(s.apiAccountDetail))))
mux.Handle("/api/v1/history/csv", s.requireAuth(getOnly(http.HandlerFunc(s.apiHistoryCSV))))
mux.Handle("/api/v1/export", s.requireAuth(getOnly(http.HandlerFunc(s.apiExport))))
mux.Handle("/api/v1/import", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiImport))))
mux.Handle("/api/v1/incident", s.requireAuth(getOnly(http.HandlerFunc(s.apiIncident))))
// Admin-scope on both routes: ServeMux cannot disambiguate by HTTP method,
// so the POST .../status mutator forces admin; reads under the same prefix
// inherit it (admin is a superset of read). The sub-path also runs CSRF
// because the router can dispatch POST .../status; requireCSRF only acts
// on unsafe methods so GET .../<id> still passes through.
mux.Handle("/api/v1/incidents", s.requireAuth(getOnly(http.HandlerFunc(s.apiIncidentList))))
mux.Handle("/api/v1/incidents/groups", s.requireRead(http.HandlerFunc(s.apiIncidentGroups)))
mux.Handle("/api/v1/incidents/", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiIncidentRouter))))
mux.Handle("/api/v1/email/stats", s.requireAuth(getOnly(http.HandlerFunc(s.apiEmailStats))))
mux.Handle("/api/v1/email/quarantine", s.requireAuth(http.HandlerFunc(s.apiEmailQuarantineList)))
mux.Handle("/api/v1/email/quarantine/", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiEmailQuarantineAction))))
mux.Handle("/api/v1/email/av/status", s.requireAuth(http.HandlerFunc(s.apiEmailAVStatus)))
mux.Handle("/api/v1/email/groups", s.requireRead(http.HandlerFunc(s.apiEmailGroups)))
mux.Handle("/api/v1/email/relay-abuse", s.requireRead(http.HandlerFunc(s.apiEmailRelayAbuse)))
mux.Handle("/api/v1/email/forwarders", s.requireRead(http.HandlerFunc(s.apiEmailForwarders)))
mux.Handle("/api/v1/email/deferrals", s.requireRead(http.HandlerFunc(s.apiEmailDeferrals)))
mux.Handle("/api/v1/email/queue-composition", s.requireRead(http.HandlerFunc(s.apiEmailQueueComposition)))
mux.Handle("/api/v1/email/queue/flush-backscatter", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiEmailFlushBackscatter))))
mux.Handle("/api/v1/email/held", s.requireAuth(http.HandlerFunc(s.apiEmailHeldList)))
mux.Handle("/api/v1/email/held/", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiEmailHeldAction))))
mux.Handle("/api/v1/performance", s.requireAuth(getOnly(http.HandlerFunc(s.apiPerformance))))
mux.Handle("/api/v1/perf/fix-error-log", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiPerfFixErrorLog))))
mux.Handle("/api/v1/perf/fix-display-errors", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiPerfFixDisplayErrors))))
mux.Handle("/api/v1/perf/fix-wp-cron", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiPerfFixWPCron))))
mux.Handle("/api/v1/hardening", s.requireAuth(getOnly(http.HandlerFunc(s.apiHardening))))
mux.Handle("/api/v1/hardening/run", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiHardeningRun))))
// Threat Intelligence API
mux.Handle("/api/v1/threat/stats", s.requireAuth(getOnly(http.HandlerFunc(s.apiThreatStats))))
mux.Handle("/api/v1/threat/top-attackers", s.requireAuth(getOnly(http.HandlerFunc(s.apiThreatTopAttackers))))
mux.Handle("/api/v1/threat/ip", s.requireAuth(getOnly(http.HandlerFunc(s.apiThreatIP))))
mux.Handle("/api/v1/threat/events", s.requireAuth(getOnly(http.HandlerFunc(s.apiThreatEvents))))
mux.Handle("/api/v1/threat/db-stats", s.requireAuth(getOnly(http.HandlerFunc(s.apiThreatDBStats))))
mux.Handle("/api/v1/audit", s.requireAuth(getOnly(http.HandlerFunc(s.apiUIAudit))))
mux.Handle("/api/v1/finding-detail", s.requireAuth(getOnly(http.HandlerFunc(s.apiFindingDetail))))
mux.Handle("/api/v1/threat/whitelist-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatWhitelistIP))))
mux.Handle("/api/v1/threat/whitelist", s.requireAuth(getOnly(http.HandlerFunc(s.apiThreatWhitelist))))
mux.Handle("/api/v1/threat/unwhitelist-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatUnwhitelistIP))))
mux.Handle("/api/v1/threat/block-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatBlockIP))))
mux.Handle("/api/v1/threat/block-ip-permanent", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatBlockIPPermanent))))
mux.Handle("/api/v1/threat/clear-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatClearIP))))
mux.Handle("/api/v1/threat/temp-whitelist-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatTempWhitelistIP))))
mux.Handle("/api/v1/threat/bulk-action", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiThreatBulkAction))))
// Rules API
mux.Handle("/api/v1/rules/status", s.requireAuth(getOnly(http.HandlerFunc(s.apiRulesStatus))))
mux.Handle("/api/v1/rules/list", s.requireAuth(getOnly(http.HandlerFunc(s.apiRulesList))))
mux.Handle("/api/v1/rules/reload", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiRulesReload))))
// Suppressions API
mux.Handle("/api/v1/suppressions", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiSuppressions))))
// Firewall API
mux.Handle("/api/v1/firewall/status", s.requireAuth(getOnly(http.HandlerFunc(s.apiFirewallStatus))))
mux.Handle("/api/v1/firewall/allowed", s.requireAuth(getOnly(http.HandlerFunc(s.apiFirewallAllowed))))
mux.Handle("/api/v1/firewall/audit", s.requireAuth(getOnly(http.HandlerFunc(s.apiFirewallAudit))))
mux.Handle("/api/v1/firewall/subnets", s.requireAuth(getOnly(http.HandlerFunc(s.apiFirewallSubnets))))
mux.Handle("/api/v1/firewall/check", s.requireAuth(getOnly(http.HandlerFunc(s.apiFirewallCheck))))
// Settings API
mux.Handle("/api/v1/settings/restart", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiSettingsRestart))))
mux.Handle("/api/v1/settings/firewall/tentative-apply", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallTentativeApply))))
mux.Handle("/api/v1/settings/firewall/confirm", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallRollbackConfirm))))
mux.Handle("/api/v1/settings/firewall/revert", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallRollbackRevert))))
mux.Handle("/api/v1/settings/firewall/rollback", s.requireAuth(http.HandlerFunc(s.apiFirewallRollbackStatus)))
mux.Handle("/api/v1/settings", s.requireAuth(http.HandlerFunc(s.apiSettingsSections)))
mux.Handle("/api/v1/settings/", s.requireAuth(http.HandlerFunc(s.apiSettings)))
// GeoIP API
mux.Handle("/api/v1/geoip", s.requireAuth(getOnly(http.HandlerFunc(s.apiGeoIPLookup))))
mux.Handle("/api/v1/geoip/batch", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiGeoIPBatch))))
// Auth-protected API - actions (with CSRF validation)
mux.Handle("/api/v1/fix", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFix))))
mux.Handle("/api/v1/verify-finding", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiVerifyFinding))))
mux.Handle("/api/v1/fix-bulk", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiBulkFix))))
mux.Handle("/api/v1/scan-account", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiScanAccount))))
mux.Handle("/api/v1/test-alert", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiTestAlert))))
mux.Handle("/api/v1/block-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiBlockIP))))
mux.Handle("/api/v1/unblock-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiUnblockIP))))
mux.Handle("/api/v1/unblock-bulk", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiUnblockBulk))))
mux.Handle("/api/v1/dismiss", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiDismissFinding))))
mux.Handle("/api/v1/quarantine-preview", s.requireAuth(getOnly(http.HandlerFunc(s.apiQuarantinePreview))))
mux.Handle("/api/v1/quarantine-restore", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiQuarantineRestore))))
mux.Handle("/api/v1/quarantine/bulk-delete", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiQuarantineBulkDelete))))
mux.Handle("/api/v1/db-object-backups", s.requireAuth(getOnly(http.HandlerFunc(s.apiDBObjectBackups))))
mux.Handle("/api/v1/db-object-backup-preview", s.requireAuth(http.HandlerFunc(s.apiDBObjectBackupPreview)))
mux.Handle("/api/v1/db-object-backup-restore", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiDBObjectBackupRestore))))
mux.Handle("/api/v1/firewall/deny-subnet", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallDenySubnet))))
mux.Handle("/api/v1/firewall/allow-ip", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallAllowIP))))
mux.Handle("/api/v1/firewall/remove-allow", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallRemoveAllow))))
mux.Handle("/api/v1/firewall/remove-subnet", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallRemoveSubnet))))
mux.Handle("/api/v1/firewall/cphulk-clear", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallFlushCphulk))))
mux.Handle("/api/v1/firewall/flush", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallFlush))))
mux.Handle("/api/v1/firewall/unban", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiFirewallUnban))))
// Operator preferences (P5.2 saved views, P5.4 user prefs).
mux.Handle("/api/v1/prefs/user", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiPrefsUser))))
mux.Handle("/api/v1/prefs/views", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiPrefsViews))))
// Bulk-action undo (P5.3).
mux.Handle("/api/v1/undo/pending", s.requireAuth(http.HandlerFunc(s.apiUndoPending)))
mux.Handle("/api/v1/undo/run", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiUndoRun))))
// Session management requires admin scope; browser mutations require CSRF.
mux.Handle("/api/v1/sessions", s.requireAuth(http.HandlerFunc(s.apiSessions)))
mux.Handle("DELETE /api/v1/sessions", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiSessions))))
mux.Handle("/api/v1/sessions/", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.apiSessions))))
mux.Handle("/logout", s.requireAuth(s.requireCSRF(http.HandlerFunc(s.handleLogout))))
// /metrics (ROADMAP item 4) has its own auth: the handler accepts
// cfg.WebUI.MetricsToken as a dedicated Bearer token so Prometheus
// scrapers get a credential that does not also unlock the UI, and
// falls back to the existing AuthToken/session path so the UI can
// self-scrape. No CSRF: read-only endpoint.
mux.HandleFunc("/metrics", s.handleMetrics)
s.httpSrv = &http.Server{
Addr: cfg.WebUI.Listen,
Handler: s.securityHeaders(mux),
ReadHeaderTimeout: 10 * time.Second, // slowloris protection
ReadTimeout: 30 * time.Second, // max time to read full request
WriteTimeout: serverWriteTimeout,
IdleTimeout: 120 * time.Second,
MaxHeaderBytes: 1 << 20, // 1MB
TLSConfig: &tls.Config{
MinVersion: tls.VersionTLS12,
// Disable HTTP/2: Go's HTTP/2 implementation applies WriteTimeout
// to the entire connection, not per-stream. Long-running handlers
// (account scans ~5min) cause ERR_HTTP2_PROTOCOL_ERROR in browsers
// when the timeout fires. HTTP/1.1 handles per-request deadlines
// correctly via ResponseController.SetWriteDeadline.
NextProtos: []string{"http/1.1"},
CipherSuites: []uint16{
tls.TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384,
tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256,
tls.TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305,
tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256,
tls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305,
},
},
}
s.restartDaemon = defaultRestartDaemon
return s, nil
}
// pruneLoginAttempts periodically cleans up stale rate-limit entries.
// It returns when s.pruneDone is closed.
func (s *Server) pruneLoginAttempts() {
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for {
select {
case <-s.pruneDone:
return
case <-ticker.C:
s.loginMu.Lock()
cutoff := time.Now().Add(-time.Minute)
for ip, attempts := range s.loginAttempts {
var recent []time.Time
for _, t := range attempts {
if t.After(cutoff) {
recent = append(recent, t)
}
}
if len(recent) == 0 {
delete(s.loginAttempts, ip)
} else {
s.loginAttempts[ip] = recent
}
}
s.loginMu.Unlock()
// Also prune API rate-limit entries
s.apiMu.Lock()
for ip, reqs := range s.apiRequests {
var recent []time.Time
for _, t := range reqs {
if t.After(cutoff) {
recent = append(recent, t)
}
}
if len(recent) == 0 {
delete(s.apiRequests, ip)
} else {
s.apiRequests[ip] = recent
}
}
s.apiMu.Unlock()
}
}
}
// Start starts the HTTPS server. Blocks until shutdown.
func (s *Server) Start() error {
certPath := s.cfg.WebUI.TLSCert
keyPath := s.cfg.WebUI.TLSKey
if certPath == "" && keyPath == "" {
certPath = filepath.Join(s.cfg.StatePath, "webui.crt")
keyPath = filepath.Join(s.cfg.StatePath, "webui.key")
if err := EnsureTLSCert(certPath, keyPath, s.cfg.Hostname); err != nil {
return fmt.Errorf("TLS cert setup: %w", err)
}
} else if err := renewTLSCert(certPath, keyPath); err != nil {
return fmt.Errorf("TLS cert renewal: %w", err)
}
certs, err := newCertReloader(certPath, keyPath)
if err != nil {
return fmt.Errorf("TLS cert load: %w", err)
}
s.httpSrv.TLSConfig.GetCertificate = certs.GetCertificate
// A listener failure ends these workers just as Shutdown does.
defer s.shutdownOnce.Do(func() { close(s.pruneDone) })
obs.Go("webui-prune-logins", s.pruneLoginAttempts)
obs.Go("webui-cert-renewal", func() { s.renewCertLoop(certPath, keyPath) })
fmt.Fprintf(os.Stderr, "WebUI listening on https://%s\n", s.cfg.WebUI.Listen)
return s.httpSrv.ListenAndServeTLS("", "")
}
// renewCertLoop checks the generated certificate daily and renews it before
// it expires; the listener picks up the new files on the next handshake.
func (s *Server) renewCertLoop(certPath, keyPath string) {
ticker := time.NewTicker(24 * time.Hour)
defer ticker.Stop()
for {
select {
case <-s.pruneDone:
return
case <-ticker.C:
if err := renewTLSCert(certPath, keyPath); err != nil {
fmt.Fprintf(os.Stderr, "webui: TLS certificate renewal: %v\n", err)
}
}
}
}
// Shutdown gracefully stops the server. Safe to call more than once;
// the underlying pruneDone close is guarded so duplicate shutdown does
// not panic.
func (s *Server) Shutdown(ctx context.Context) error {
s.shutdownOnce.Do(func() { close(s.pruneDone) })
if s.httpSrv == nil {
return nil
}
return s.httpSrv.Shutdown(ctx)
}
// canonicalAllowedOrigin returns the single CORS origin the web UI
// will accept on /api/ requests. Built from cfg.Hostname plus the
// listen port so a forged HTTP Host header cannot redirect the check.
func (s *Server) canonicalAllowedOrigin() string {
host := canonicalOriginHost(s.cfg.Hostname)
port := webUIListenPort(s.cfg.WebUI.Listen)
if port != "" && port != "443" {
if strings.HasPrefix(host, "[") && strings.HasSuffix(host, "]") {
host = host[1 : len(host)-1]
}
host = net.JoinHostPort(host, port)
}
return "https://" + host
}
func canonicalOriginHost(host string) string {
host = strings.TrimSpace(host)
if strings.HasPrefix(host, "[") {
if end := strings.Index(host, "]"); end > 0 {
if ip := net.ParseIP(host[1:end]); ip != nil {
if ip4 := ip.To4(); ip4 != nil {
return ip4.String()
}
return "[" + ip.String() + "]"
}
}
}
if ip := net.ParseIP(host); ip != nil {
if ip4 := ip.To4(); ip4 != nil {
return ip4.String()
}
return "[" + ip.String() + "]"
}
return strings.ToLower(host)
}
func webUIListenPort(listen string) string {
if _, port, err := net.SplitHostPort(listen); err == nil {
return port
}
if idx := strings.LastIndex(listen, ":"); idx >= 0 {
return listen[idx+1:]
}
return ""
}
// originAllowed reports whether a browser Origin may make credentialed
// requests. A loopback origin (an SSH tunnel on any local port) is trusted
// only when it is the origin the request was sent to: host is the request's
// Host. Another local service in the same browser shares the Web UI's
// cookies, since cookies ignore the port, and must not be trusted.
func (s *Server) originAllowed(origin, host string) bool {
u, err := url.Parse(origin)
if err != nil || !originHeaderURL(u) || !strings.EqualFold(u.Scheme, "https") {
return false
}
if isLoopbackOriginHost(u.Hostname()) {
return host != "" && sameOrigin(origin, "https://"+host)
}
if sameOrigin(origin, s.canonicalAllowedOrigin()) {
return true
}
for _, listed := range s.liveCfg().WebUI.AllowedOrigins {
if sameOrigin(origin, listed) {
return true
}
}
return false
}
func isLoopbackOriginHost(host string) bool {
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
func sameOrigin(got, want string) bool {
gotURL, err := url.Parse(got)
if err != nil || !originHeaderURL(gotURL) {
return false
}
wantURL, err := url.Parse(want)
if err != nil || !originHeaderURL(wantURL) {
return false
}
return strings.EqualFold(gotURL.Scheme, wantURL.Scheme) &&
strings.EqualFold(gotURL.Hostname(), wantURL.Hostname()) &&
originPort(gotURL) == originPort(wantURL)
}
func originHeaderURL(u *url.URL) bool {
return u.Scheme != "" && u.Host != "" && u.User == nil &&
u.Path == "" && u.RawQuery == "" && u.Fragment == ""
}
func originPort(u *url.URL) string {
if port := u.Port(); port != "" {
return port
}
switch strings.ToLower(u.Scheme) {
case "https":
return "443"
case "http":
return "80"
default:
return ""
}
}
// SetSigCount sets the loaded signature count for the status API.
func (s *Server) SetSigCount(count int) {
s.sigCountMu.Lock()
s.sigCount = count
s.sigCountMu.Unlock()
}
func (s *Server) signatureCount() int {
s.sigCountMu.RLock()
defer s.sigCountMu.RUnlock()
return s.sigCount
}
// HasUI returns true if UI templates were loaded from disk.
func (s *Server) HasUI() bool {
return s.hasUI
}
// SetIPBlocker sets the firewall engine for block/unblock operations.
func (s *Server) SetIPBlocker(b IPBlocker) {
s.blocker = b
}
// SetHealthInfo installs the daemon readers the health API and dashboard
// call on each request. Log watchers can start after the web UI, so a
// snapshot taken at startup went stale.
func (s *Server) SetHealthInfo(fanotifyActive func() bool, logWatchers func() int) {
s.fanotifyActive = fanotifyActive
s.logWatcherCount = logWatchers
}
func (s *Server) fanotifyRunning() bool {
return s.fanotifyActive != nil && s.fanotifyActive()
}
func (s *Server) logWatchersRunning() int {
if s.logWatcherCount == nil {
return 0
}
return s.logWatcherCount()
}
// SetEmailQuarantine sets the email quarantine for the email AV API endpoints.
func (s *Server) SetEmailQuarantine(q *emailav.Quarantine) {
s.emailQuarantine.Store(q)
}
// SetEmailAVWatcherMode sets the watcher mode string for the email AV status API.
func (s *Server) SetEmailAVWatcherMode(mode string) {
s.emailAVWatcherMode.Store(mode)
}
// emailQuarantineHandle returns the installed email quarantine, or nil
// before the daemon installs one.
func (s *Server) emailQuarantineHandle() *emailav.Quarantine {
return s.emailQuarantine.Load()
}
// emailAVMode returns the installed AV watcher mode, "" before it is set.
func (s *Server) emailAVMode() string {
if v, ok := s.emailAVWatcherMode.Load().(string); ok {
return v
}
return ""
}
// SetVersion sets the application version for display in the UI.
func (s *Server) SetVersion(v string) {
s.version = v
}
// SetHealthProvider installs the daemon's health provider. The webui
// constructs without one so unit tests can run without a daemon; the
// daemon must call this before any request hits /api/v1/status.
func (s *Server) SetHealthProvider(p health.Provider) {
s.provider = p
}
// SetFindingBus installs the broadcaster the SSE event stream subscribes
// to. The webui constructs without one so unit tests work without a
// daemon; the daemon must call this before any request hits /api/v1/events.
func (s *Server) SetFindingBus(bus *broadcast.Bus) {
s.mu.Lock()
defer s.mu.Unlock()
s.findingBus = bus
}
// SetIncidentCorrelator wires the incident correlator. Called once at
// startup; treated as immutable after first set.
func (s *Server) SetIncidentCorrelator(c *incident.Correlator) {
s.incidentCorrelator = c
}
// SetScanJobs wires the full-scan job manager so POST /api/v1/scan-jobs and
// POST /api/v1/scan-jobs/{id}/cancel can enqueue and cancel jobs through the
// daemon's single-job queue. Must be called before the server starts
// accepting requests; POST handlers return 503 when nil.
func (s *Server) SetScanJobs(c ScanJobController) {
s.scanJobs = c
}
// csmConfig returns the feature-flag map used by the frontend. Templates
// render it through jsonForScript.
func (s *Server) csmConfig() map[string]interface{} {
cfg := s.liveCfg()
return map[string]interface{}{
"version": s.version,
"emailAV": cfg.EmailAV.Enabled,
"firewall": cfg.Firewall != nil && cfg.Firewall.Enabled,
"autoResponse": cfg.AutoResponse.Enabled,
"threatIntel": cfg.Reputation.AbuseIPDBKey != "",
"signatures": cfg.Signatures.RulesDir != "",
"challenge": cfg.Challenge.Difficulty > 0,
"fanotify": s.fanotifyRunning(),
"hostname": s.cfg.Hostname,
// Attack types (brute_force, waf_block, ...) group findings in the
// attack database; their labels come from there.
"attackTypes": attackdb.AttackTypeLabels(),
// #nosec G101 -- Not credentials. This is a lookup from
// finding check name (webshell, email_credential_leak, etc.) to
// the human-readable label rendered in the UI.
"checkNames": map[string]string{
"webshell": "Web Shell",
"cpanel_login": "cPanel Login",
"perf_load": "Load",
"perf_php_processes": "PHP Processes",
"perf_memory": "Memory",
"perf_php_handler": "PHP Handler",
"perf_mysql_config": "MySQL Config",
"perf_redis_config": "Redis Config",
"perf_error_logs": "Error Logs",
"perf_wp_config": "WP Config",
"perf_wp_transients": "WP Transients",
"perf_wp_cron": "WP Cron",
"integrity": "Integrity",
"db_siteurl_hijack": "DB URL Hijack",
"db_siteurl_invalid": "DB Invalid Site Address",
"db_options_injection": "DB Options Injection",
"db_post_injection": "DB Post Injection",
"db_spam_injection": "DB Spam Injection",
"db_rogue_admin": "DB Rogue Admin",
"db_suspicious_admin_email": "DB Suspicious Admin",
"mail_queue": "Mail Queue",
"mail_queue_unavailable": "Mail Queue Unavailable",
"mail_per_account": "Mail Volume",
"email_phishing_content": "Email Phishing",
"email_malware": "Email Malware",
"email_compromised_account": "Compromised Account",
"email_cloud_relay_abuse": "Cloud-Relay Credential Abuse",
"email_spam_outbreak": "Spam Outbreak",
"email_defer_fail_governor": "Defer/Fail Governor",
"email_credential_leak": "Credential Leak",
"email_auth_failure_realtime": "Auth Failure",
"smtp_bruteforce": "SMTP Brute Force",
"smtp_subnet_spray": "SMTP Subnet Spray",
"smtp_account_spray": "SMTP Account Spray",
"mail_bruteforce": "Mail Brute Force",
"mail_bruteforce_suspected": "Mail Auth Advisory",
"mail_subnet_spray": "Mail Subnet Spray",
"mail_auth_backend_degraded": "Mail Auth Backend Degraded",
"admin_panel_bruteforce": "Admin Panel Brute Force",
"mail_account_spray": "Mail Account Spray",
"mail_account_compromised": "Mail Account Compromised",
"http_asn_crawl": "Distributed ASN Crawl",
"http_request_flood": "HTTP Request Flood",
"http_scanner_profile": "HTTP Scanner Profile",
"http_claimed_bot_unverified": "Unverified Claimed Bot",
"http_ua_spoof": "HTTP UA Spoof",
"http_distributed_flood": "Distributed HTTP Flood",
"exim_frozen_realtime": "Frozen Message",
"email_suspicious_geo": "Suspicious Geo Login",
"email_rate_critical": "Email Rate Critical",
"email_rate_warning": "Email Rate Warning",
"email_dkim_failure": "DKIM Failure",
"email_spf_rejection": "SPF Rejection",
"email_pipe_forwarder": "Pipe Forwarder",
"email_suspicious_forwarder": "Suspicious Forwarder",
"cpanel_login_realtime": "cPanel Login",
"cpanel_password_purge_realtime": "Password Purge",
"ssh_login_unknown_ip": "SSH Login",
"pam_login": "PAM Login",
"pam_bruteforce": "PAM Brute Force",
"modsec_block_escalation": "ModSec Escalation",
"modsec_csm_block_escalation": "ModSec Escalation",
"whm_password_change_noninfra": "WHM Password Change",
"password_hijack_confirmed": "Password Hijack",
},
}
}
// --- Template helpers ---
func severityClass(sev alert.Severity) string {
switch sev {
case alert.Critical:
return "critical"
case alert.High:
return "high"
case alert.Warning:
return "warning"
default:
return "secondary"
}
}
func severityLabel(sev alert.Severity) string {
switch sev {
case alert.Critical:
return "CRITICAL"
case alert.High:
return "HIGH"
case alert.Warning:
return "WARNING"
default:
return "INFO"
}
}
// severityRank returns a numeric rank for severity labels (higher = more severe).
func severityRank(label string) int {
switch label {
case "CRITICAL":
return 3
case "HIGH":
return 2
case "WARNING":
return 1
default:
return 0
}
}
func timeAgo(t time.Time) string {
d := time.Since(t)
switch {
case d < time.Minute:
return "just now"
case d < time.Hour:
return fmt.Sprintf("%dm ago", int(d.Minutes()))
case d < 24*time.Hour:
return fmt.Sprintf("%dh ago", int(d.Hours()))
default:
return fmt.Sprintf("%dd ago", int(d.Hours()/24))
}
}
func formatTime(t time.Time) string {
return t.Format("2006-01-02 15:04:05")
}
// isoTime renders an instant for a <time datetime> attribute, which the
// page rewrites in the operator's time zone.
func isoTime(t time.Time) string {
return t.UTC().Format(time.RFC3339)
}
func (s *Server) templateFuncs() template.FuncMap {
return template.FuncMap{
"severityClass": severityClass,
"severityLabel": severityLabel,
"timeAgo": timeAgo,
"formatTime": formatTime,
"isoTime": isoTime,
"asset": s.assetURL,
"csrfToken": func() string { return "" }, // per request, see renderTemplate
"csmConfig": func() template.JS { return jsonForScript(s.csmConfig()) },
"json": jsonForScript,
"multiply": func(a, b int) int { return a * b },
"add": func(a, b int) int { return a + b },
"subtract": func(a, b int) int { return a - b },
"divisibleBy": func(a, b int) bool { return b != 0 && a%b == 0 },
"serverTimeZone": serverTimeZoneName,
"serverUTCOffset": serverUTCOffsetMinutes,
}
}
// --- Security headers middleware ---
func (s *Server) securityHeaders(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("X-Content-Type-Options", "nosniff")
// Turn the legacy XSS auditor off: where browsers still honour it,
// "1; mode=block" can be abused to suppress page scripts.
w.Header().Set("X-XSS-Protection", "0")
w.Header().Set("Referrer-Policy", "strict-origin-when-cross-origin")
w.Header().Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
w.Header().Set("Content-Security-Policy", "default-src 'self'; script-src 'self'; style-src 'self'; connect-src 'self'; img-src 'self' data:; font-src 'self'; object-src 'none'; base-uri 'none'; form-action 'self'; frame-ancestors 'none'")
w.Header().Set("Permissions-Policy", "camera=(), microphone=(), geolocation=(), payment=()")
w.Header().Set("Cache-Control", "no-store")
// CORS/origin validation: reject cross-origin API requests.
// Non-loopback origins must be configured; loopback origins must
// match Host so another local service cannot borrow UI cookies.
// Browser logout and session revocation change server state like API
// writes. Login stays reachable: it already requires the credential.
if strings.HasPrefix(r.URL.Path, "/api/") || r.URL.Path == "/logout" || strings.HasPrefix(r.URL.Path, "/sessions") {
origin := r.Header.Get("Origin")
if origin != "" {
if !s.originAllowed(origin, r.Host) {
writeRequestError(w, r, "Cross-origin request blocked", http.StatusForbidden)
return
}
w.Header().Set("Access-Control-Allow-Origin", origin)
w.Header().Set("Access-Control-Allow-Credentials", "true")
}
// Deny CORS preflight from unknown origins
if r.Method == "OPTIONS" {
w.WriteHeader(http.StatusNoContent)
return
}
}
// API rate limiting: 600 requests per minute per IP. /metrics takes
// a bearer token too, so guessing it shares the same budget.
if strings.HasPrefix(r.URL.Path, "/api/") || r.URL.Path == "/metrics" {
ip := rateLimitKey(r.RemoteAddr)
s.apiMu.Lock()
now := time.Now()
cutoff := now.Add(-time.Minute)
var recent []time.Time
for _, t := range s.apiRequests[ip] {
if t.After(cutoff) {
recent = append(recent, t)
}
}
if len(recent) >= 600 {
s.apiMu.Unlock()
writeRequestError(w, r, "Rate limit exceeded", http.StatusTooManyRequests)
return
}
if _, tracked := s.apiRequests[ip]; !tracked {
boundRateLimitMap(s.apiRequests, cutoff)
}
s.apiRequests[ip] = append(recent, now)
s.apiMu.Unlock()
}
next.ServeHTTP(w, r)
})
}
// getOnly answers every method but GET with a JSON 405, for read routes
// behind requireAuth; requireRead does the same for its routes.
func getOnly(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
w.Header().Set("Allow", http.MethodGet)
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
next.ServeHTTP(w, r)
})
}
// --- CSRF protection ---
// csrfTokenForSession derives the CSRF token of one browser session from an
// active admin secret and the session's cookie secret. Each session gets its
// own token and a new login a new one; knowing a token reveals neither
// secret. Without both secrets there is no token.
func (s *Server) csrfTokenForSession(sessionSecret string) string {
secret := s.csrfSecret()
if secret == "" || sessionSecret == "" {
return ""
}
mac := hmac.New(sha256.New, []byte(secret))
fmt.Fprintf(mac, "csm-csrf-v2:%s", sessionSecret)
return hex.EncodeToString(mac.Sum(nil))[:32]
}
// csrfTokenFor is the CSRF token of the browser session r carries.
func (s *Server) csrfTokenFor(r *http.Request) string {
c, err := r.Cookie("csm_auth")
if err != nil {
return ""
}
return s.csrfTokenForSession(c.Value)
}
func (s *Server) csrfSecret() string {
for _, tok := range s.cfg.WebUI.Tokens {
if tok.Scope == "admin" && tok.Token != "" {
return tok.Token
}
}
if len(s.cfg.WebUI.Tokens) == 0 {
return s.cfg.WebUI.AuthToken
}
return ""
}
// validateCSRF enforces the browser-session CSRF boundary on state-changing
// routes. Bearer-authenticated requests skip the check because cross-origin
// browser requests cannot attach the Authorization header without script access
// to the bearer token.
func (s *Server) validateCSRF(r *http.Request) bool {
if !isUnsafeCSRFMethod(r.Method) {
return true // only validate state-changing methods
}
// Skip CSRF only when the bearer token itself grants admin writes. A
// read-scope bearer presented alongside an admin cookie must not turn
// the cookie-authenticated request into a CSRF-exempt API call.
// CSRF protection is only needed for cookie-based browser sessions.
if s.isAdminBearerAuth(r) {
return true
}
expected := s.csrfTokenFor(r)
// A request without an admin secret or a browser session cannot prove
// it came from a page of that session, so the mutating path stays closed.
if expected == "" {
return false
}
// Check header (API calls from JS use this)
if token := r.Header.Get("X-CSRF-Token"); token != "" {
return subtle.ConstantTimeCompare([]byte(token), []byte(expected)) == 1
}
// Check form field (traditional form posts). Only the body counts: a
// token in the query string would end up in URLs, logs and Referer.
if token := r.PostFormValue("csrf_token"); token != "" {
return subtle.ConstantTimeCompare([]byte(token), []byte(expected)) == 1
}
return false
}
// requireCSRF wraps a handler to validate CSRF on POST, PUT, PATCH, and DELETE
// requests. PUT joined the unsafe set when /api/v1/prefs/user landed; the
// existing list pre-dates that endpoint.
func (s *Server) requireCSRF(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Skip CSRF for admin Bearer token auth. API-to-API callers do not
// need CSRF protection, but read-scope bearer tokens never authorize
// mutating handlers on their own.
if isUnsafeCSRFMethod(r.Method) && !s.isAdminBearerAuth(r) && !s.validateCSRF(r) {
writeRequestError(w, r, "Invalid CSRF token", http.StatusForbidden)
return
}
next.ServeHTTP(w, r)
})
}
func isUnsafeCSRFMethod(method string) bool {
switch method {
case http.MethodPost, http.MethodPut, http.MethodDelete, http.MethodPatch:
return true
default:
return false
}
}
func (s *Server) isBearerAuth(r *http.Request) bool {
_, ok := s.bearerTokenWithScope(r, "read")
return ok
}
func (s *Server) isAdminBearerAuth(r *http.Request) bool {
_, ok := s.bearerTokenWithScope(r, "admin")
return ok
}
// --- Scan rate limiting ---
// acquireScan tries to start a scan. Returns false if a scan is already running.
func (s *Server) acquireScan() bool {
s.scanMu.Lock()
defer s.scanMu.Unlock()
if s.scanRunning {
return false
}
s.scanRunning = true
return true
}
func (s *Server) releaseScan() {
s.scanMu.Lock()
s.scanRunning = false
s.scanMu.Unlock()
}
package webui
import (
"net/http"
"strings"
"time"
"github.com/pidginhost/csm/internal/session"
)
type browserSessionView struct {
ID string `json:"id"`
Name string `json:"name"`
Created time.Time `json:"created"`
LastSeen time.Time `json:"last_seen"`
Expires time.Time `json:"expires"`
RemoteIP string `json:"remote_ip"`
UserAgent string `json:"user_agent"`
Current bool `json:"current"`
}
func (s *Server) browserSessionViews(r *http.Request) ([]browserSessionView, error) {
if s.sessions == nil {
return nil, errNoStore
}
records, err := s.sessions.List(s.sessionNow())
if err != nil {
return nil, err
}
current := ""
if cookie, err := r.Cookie("csm_auth"); err == nil {
current = session.Hash(cookie.Value)
}
views := make([]browserSessionView, 0, len(records))
for _, rec := range records {
views = append(views, browserSessionView{ID: rec.ID, Name: rec.Name, Created: rec.Created,
LastSeen: rec.LastSeen, Expires: rec.Expires, RemoteIP: rec.RemoteIP, UserAgent: rec.UserAgent, Current: rec.Verifier == current})
}
return views, nil
}
func (s *Server) apiSessions(w http.ResponseWriter, r *http.Request) {
id := strings.TrimPrefix(r.URL.Path, "/api/v1/sessions")
if id != "" {
id = strings.TrimPrefix(id, "/")
if !validSessionID(id) {
writeJSONError(w, "Session not found", http.StatusNotFound)
return
}
}
switch r.Method {
case http.MethodGet:
if id != "" {
writeJSONError(w, "Session not found", http.StatusNotFound)
return
}
views, err := s.browserSessionViews(r)
if err != nil {
writeJSONError(w, "Session store unavailable", http.StatusServiceUnavailable)
return
}
writeAll(w, views)
case http.MethodDelete:
if s.sessions == nil {
writeJSONError(w, "Session store unavailable", http.StatusServiceUnavailable)
return
}
actor, via := s.requestActor(r)
if id != "" && !s.sessionExists(id) {
writeJSONError(w, "Session not found", http.StatusNotFound)
return
}
var err error
if id == "" {
err = s.sessions.RevokeAll()
} else {
err = s.sessions.Revoke(id)
}
if err != nil {
writeJSONError(w, "Cannot revoke browser session", http.StatusServiceUnavailable)
return
}
s.auditSessionRevoke(r, actor, via, id)
if id == "" {
clearBrowserCookie(w)
}
writeOK(w, nil)
default:
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
}
}
// sessionExists reports whether id names an active browser session, so
// revoking an unknown one answers 404 instead of a success that did nothing.
func (s *Server) sessionExists(id string) bool {
records, err := s.sessions.List(s.sessionNow())
if err != nil {
return false
}
for _, rec := range records {
if rec.ID == id {
return true
}
}
return false
}
func validSessionID(id string) bool {
if len(id) != 32 {
return false
}
for _, ch := range id {
if (ch < '0' || ch > '9') && (ch < 'a' || ch > 'f') {
return false
}
}
return true
}
func (s *Server) handleSessions(w http.ResponseWriter, r *http.Request) {
views, err := s.browserSessionViews(r)
if err != nil {
http.Error(w, "Session store unavailable", http.StatusServiceUnavailable)
return
}
s.renderTemplate(w, r, "sessions.html", map[string]any{"Sessions": views})
}
func (s *Server) handleSessionRevoke(w http.ResponseWriter, r *http.Request) {
if s.sessions == nil {
http.Error(w, "Session store unavailable", http.StatusServiceUnavailable)
return
}
r.Body = http.MaxBytesReader(w, r.Body, 4096)
if err := r.ParseForm(); err != nil {
http.Error(w, "Invalid request", http.StatusBadRequest)
return
}
id := r.PostForm.Get("id")
actor, via := s.requestActor(r)
var err error
switch {
case id == "all":
err = s.sessions.RevokeAll()
case validSessionID(id):
err = s.sessions.Revoke(id)
default:
http.Error(w, "Invalid session", http.StatusBadRequest)
return
}
if err != nil {
http.Error(w, "Cannot revoke browser session", http.StatusServiceUnavailable)
return
}
target := id
if id == "all" {
target = ""
}
s.auditSessionRevoke(r, actor, via, target)
if id == "all" {
clearBrowserCookie(w)
http.Redirect(w, r, "/login", http.StatusSeeOther)
return
}
http.Redirect(w, r, "/sessions", http.StatusSeeOther)
}
// auditSessionRevoke records a revocation for the actor resolved before it,
// since revoking the caller's own session also ends its attribution. An
// empty id means every browser session.
func (s *Server) auditSessionRevoke(r *http.Request, actor, via, id string) {
if id == "" {
s.auditLogAs(r, actor, via, "session_revoke_all", "browser sessions", "every browser session logged out")
return
}
s.auditLogAs(r, actor, via, "session_revoke", id, "browser session revoked")
}
package webui
import (
"context"
"encoding/json"
"fmt"
"net/http"
"os"
"os/exec"
"reflect"
"sort"
"strconv"
"strings"
"time"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/integrity"
"gopkg.in/yaml.v3"
)
const settingsURLPrefix = "/api/v1/settings/"
type pendingSettingsSection struct {
ID string `json:"id"`
Title string `json:"title"`
}
func pendingRestartSections(live, disk *config.Config) []pendingSettingsSection {
if live == nil || disk == nil {
return nil
}
seen := map[string]struct{}{}
for _, c := range config.Diff(live, disk) {
if c.Tag == config.TagSafe {
continue
}
for _, section := range settingsSections {
if c.Field == section.YAMLPath || strings.HasPrefix(c.Field, section.YAMLPath+".") {
seen[section.ID] = struct{}{}
break
}
}
}
if len(seen) == 0 {
return nil
}
out := make([]pendingSettingsSection, 0, len(seen))
for _, section := range settingsSections {
if _, ok := seen[section.ID]; ok {
out = append(out, pendingSettingsSection{ID: section.ID, Title: section.Title})
}
}
return out
}
func cloneConfigForSettingsApply(src *config.Config) config.Config {
clone := *src
if src.Firewall != nil {
fw := *src.Firewall
clone.Firewall = &fw
}
return clone
}
func copySettingsChangeValues(dst, src *config.Config, section SettingsSection, changes map[string]json.RawMessage) error {
for key := range changes {
path := []string{section.YAMLPath}
if key != "" {
path = append(path, strings.Split(key, ".")...)
}
if err := copyConfigPathValue(dst, src, path); err != nil {
return err
}
}
return nil
}
func copyConfigPathValue(dst, src *config.Config, path []string) error {
dstValue := reflect.ValueOf(dst).Elem()
srcValue := reflect.ValueOf(src).Elem()
for i, key := range path {
if srcValue.Kind() == reflect.Pointer {
if srcValue.IsNil() {
return fmt.Errorf("path %v: source element %d is nil", path, i)
}
if dstValue.IsNil() {
dstValue.Set(reflect.New(dstValue.Type().Elem()))
}
srcValue = srcValue.Elem()
dstValue = dstValue.Elem()
}
field, ok := fieldByYAMLTag(srcValue.Type(), key)
if !ok {
return fmt.Errorf("no yaml field %q under %s", key, strings.Join(path[:i], "."))
}
srcValue = srcValue.FieldByIndex(field.Index)
dstValue = dstValue.FieldByIndex(field.Index)
}
dstValue.Set(srcValue)
return nil
}
func (s *Server) apiSettingsSections(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
writeJSON(w, map[string]interface{}{
"groups": SectionGroupOrder,
"sections": AllSettingsSections(),
})
}
func (s *Server) apiSettings(w http.ResponseWriter, r *http.Request) {
// This prefix serves both GET (read settings) and POST (update); CSRF is
// enforced only on the mutating POST path.
switch r.Method {
case http.MethodGet:
s.apiSettingsGet(w, r)
case http.MethodPost:
s.requireCSRF(http.HandlerFunc(s.apiSettingsPost)).ServeHTTP(w, r)
default:
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
}
}
func (s *Server) apiSettingsGet(w http.ResponseWriter, r *http.Request) {
sectionID := strings.TrimPrefix(r.URL.Path, settingsURLPrefix)
if sectionID == "" || strings.Contains(sectionID, "/") {
writeJSONError(w, "section required", http.StatusBadRequest)
return
}
if sectionID == "restart" {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
section, ok := LookupSettingsSection(sectionID)
if !ok {
writeJSONError(w, "unknown section", http.StatusNotFound)
return
}
resolveFieldOptions(§ion)
diskBytes, err := os.ReadFile(s.cfg.ConfigFile) // #nosec G304 -- operator-configured config path
if err != nil {
writeJSONError(w, "read config: "+err.Error(), http.StatusInternalServerError)
return
}
mergedBytes, err := config.MergeBytesWithDir(diskBytes, s.cfg.ConfigDir)
if err != nil {
writeJSONError(w, "load config: "+err.Error(), http.StatusInternalServerError)
return
}
disk, err := config.LoadBytes(mergedBytes)
if err != nil {
writeJSONError(w, "load config: "+err.Error(), http.StatusInternalServerError)
return
}
disk.ConfigFile = s.cfg.ConfigFile
disk.ConfigDir = s.cfg.ConfigDir
redacted := config.Redact(disk)
values, err := extractSectionValues(mergedBytes, redacted, section)
if err != nil {
writeJSONError(w, "extract: "+err.Error(), http.StatusInternalServerError)
return
}
var pendingFields []string
var pendingSections []pendingSettingsSection
if live := config.Active(); live != nil {
diff := config.Diff(live, disk)
for _, c := range diff {
if c.Tag != config.TagSafe && (c.Field == section.YAMLPath || strings.HasPrefix(c.Field, section.YAMLPath+".")) {
pendingFields = append(pendingFields, c.Field)
}
}
pendingSections = pendingRestartSections(live, disk)
}
etag := integrity.HashConfigStableBytes(diskBytes)
w.Header().Set("ETag", etag)
writeJSON(w, map[string]interface{}{
"section": section,
"values": values,
"etag": etag,
"pending_restart": len(pendingFields) > 0,
"pending_fields": pendingFields,
"pending_sections": pendingSections,
})
}
func extractSectionValues(rawBytes []byte, effectiveCfg *config.Config, section SettingsSection) (map[string]interface{}, error) {
effective, err := extractSectionEffectiveValues(effectiveCfg, section)
if err != nil {
return nil, err
}
raw, err := extractSectionRawValues(rawBytes, section)
if err != nil {
return nil, err
}
values := make(map[string]interface{}, len(effective))
for k, v := range effective {
values[k] = v
}
overlayNullableState(section, values, raw)
return values, nil
}
func extractSectionEffectiveValues(cfg *config.Config, section SettingsSection) (map[string]interface{}, error) {
var wrapper map[string]interface{}
data, err := yaml.Marshal(cfg)
if err != nil {
return nil, err
}
if err := yaml.Unmarshal(data, &wrapper); err != nil {
return nil, err
}
raw, ok := wrapper[section.YAMLPath]
if !ok {
return map[string]interface{}{}, nil
}
if _, isMap := raw.(map[string]interface{}); !isMap {
return map[string]interface{}{section.YAMLPath: raw}, nil
}
return raw.(map[string]interface{}), nil
}
func extractSectionRawValues(rawBytes []byte, section SettingsSection) (map[string]interface{}, error) {
var wrapper map[string]interface{}
if err := yaml.Unmarshal(rawBytes, &wrapper); err != nil {
return nil, err
}
raw, ok := wrapper[section.YAMLPath]
if !ok {
return map[string]interface{}{}, nil
}
if _, isMap := raw.(map[string]interface{}); !isMap {
return map[string]interface{}{section.YAMLPath: raw}, nil
}
return raw.(map[string]interface{}), nil
}
func overlayNullableState(section SettingsSection, values, raw map[string]interface{}) {
for _, field := range section.Fields {
if !field.Nullable {
continue
}
// All v1 nullable fields are direct children of the section.
if strings.Contains(field.YAMLPath, ".") {
continue
}
if v, ok := raw[field.YAMLPath]; ok {
values[field.YAMLPath] = v
continue
}
values[field.YAMLPath] = nil
}
}
func (s *Server) apiSettingsPost(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
sectionID := strings.TrimPrefix(r.URL.Path, settingsURLPrefix)
if sectionID == "" || strings.Contains(sectionID, "/") {
writeJSONError(w, "section required", http.StatusBadRequest)
return
}
section, ok := LookupSettingsSection(sectionID)
if !ok {
writeJSONError(w, "unknown section", http.StatusNotFound)
return
}
ifMatch := r.Header.Get("If-Match")
if ifMatch == "" {
writeJSONError(w, "If-Match header required", http.StatusBadRequest)
return
}
var body struct {
Changes map[string]json.RawMessage `json:"changes"`
}
if err := decodeJSONBodyLimited(w, r, 256*1024, &body); err != nil {
writeJSONError(w, "invalid body: "+err.Error(), http.StatusBadRequest)
return
}
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
if rejectConfigWriteDuringRollback(w) {
return
}
if s.settingsSaveHook != nil {
s.settingsSaveHook()
}
diskBytes, err := os.ReadFile(s.cfg.ConfigFile) // #nosec G304 -- operator-supplied config path
if err != nil {
writeJSONError(w, "read config: "+err.Error(), http.StatusInternalServerError)
return
}
disk, err := config.LoadBytes(diskBytes)
if err != nil {
writeJSONError(w, "parse config: "+err.Error(), http.StatusInternalServerError)
return
}
disk.ConfigFile = s.cfg.ConfigFile
disk.ConfigDir = s.cfg.ConfigDir
if integrity.HashConfigStableBytes(diskBytes) != ifMatch {
writeJSONError(w, "config changed on disk, reload", http.StatusPreconditionFailed)
return
}
if rejectIfConfDirChanged(w, s.cfg.ConfigDir, disk) {
return
}
effectiveDisk, err := config.LoadBytesWithDir(diskBytes, s.cfg.ConfigDir)
if err != nil {
writeJSONError(w, "load merged config: "+err.Error(), http.StatusInternalServerError)
return
}
effectiveDisk.ConfigFile = s.cfg.ConfigFile
clone := cloneConfigForSettingsApply(disk)
yamlChanges, errs := buildChangeSet(section, &clone, body.Changes)
if len(errs) > 0 {
writeValidationErrors(w, errs)
return
}
// Validate the merged result before the edited YAML goes back through
// Load. This looks like a duplicate of the check further down, but Load
// rejects some combinations itself and returns a plain error, which loses
// the field attribution the dashboard needs to mark the offending input.
effectiveCandidate := cloneConfigForSettingsApply(effectiveDisk)
if _, errs := buildChangeSet(section, &effectiveCandidate, body.Changes); len(errs) > 0 {
writeValidationErrors(w, errs)
return
}
fieldErrors, _ := splitValidationResults(config.Validate(&effectiveCandidate))
if len(fieldErrors) > 0 {
writeValidationErrors(w, fieldErrors)
return
}
edited, err := config.YAMLEdit(diskBytes, yamlChanges)
if err != nil {
writeJSONError(w, "yaml edit: "+err.Error(), http.StatusInternalServerError)
return
}
effectiveClone, err := config.LoadBytesWithDir(edited, s.cfg.ConfigDir)
if err != nil {
writeJSONError(w, "load edited merged config: "+err.Error(), http.StatusUnprocessableEntity)
return
}
effectiveClone.ConfigFile = s.cfg.ConfigFile
// Drop-ins can replace the entered credential. Validate the actual
// destination/credential pair before either disk or live config changes.
rebindErrs, err := credentialRebindErrors(section, effectiveDisk, effectiveClone, body.Changes)
if err != nil {
writeJSONError(w, "read merged settings: "+err.Error(), http.StatusInternalServerError)
return
}
if len(rebindErrs) > 0 {
writeValidationErrors(w, rebindErrs)
return
}
validationResults := append(config.Validate(effectiveClone), config.ValidateDeepSection(effectiveClone, section.ID)...)
fieldErrors, warnings := splitValidationResults(validationResults)
if len(fieldErrors) > 0 {
writeValidationErrors(w, fieldErrors)
return
}
if section.ID == "firewall" {
localizeValidationFields(warnings, "firewall")
}
diff := config.Diff(effectiveDisk, effectiveClone)
var restartFields []string
for _, c := range diff {
if c.Tag != config.TagSafe {
restartFields = append(restartFields, c.Field)
}
}
var liveCandidate *config.Config
if len(restartFields) == 0 {
if live := config.Active(); live != nil {
candidate := cloneConfigForSettingsApply(live)
if err := copySettingsChangeValues(&candidate, effectiveClone, section, body.Changes); err != nil {
writeJSONError(w, "prepare live config: "+err.Error(), http.StatusInternalServerError)
return
}
liveCandidate = &candidate
} else {
candidate := cloneConfigForSettingsApply(effectiveClone)
liveCandidate = &candidate
}
}
if err := integrity.SignAndSavePreserving(s.cfg.ConfigFile, s.cfg.ConfigDir, edited, &clone, disk.Integrity.BinaryHash); err != nil {
writeJSONError(w, "save: "+err.Error(), http.StatusInternalServerError)
return
}
newETag := clone.Integrity.ConfigHash
newIntegrity := clone.Integrity
effectiveClone.Integrity = newIntegrity
if liveCandidate != nil {
applySignedIntegrityState(liveCandidate, &clone)
config.SetActive(liveCandidate)
} else if live := config.Active(); live != nil {
livePatched := *live
applySignedIntegrityState(&livePatched, &clone)
config.SetActive(&livePatched)
}
applied := []string{}
for _, c := range diff {
applied = append(applied, c.Field)
}
s.auditLog(r, "settings-save", sectionID, auditDetailsFor(section, body.Changes))
if restartFields == nil {
restartFields = []string{}
}
if warnings == nil {
warnings = []fieldError{}
}
pending := pendingRestartSections(config.Active(), effectiveClone)
if pending == nil {
pending = []pendingSettingsSection{}
}
writeOK(w, map[string]interface{}{
"applied": applied,
"requires_restart": restartFields,
"pending_restart": len(restartFields) > 0,
"pending_sections": pending,
"warnings": warnings,
"new_etag": newETag,
})
}
// applySignedIntegrityState keeps the conf.d digest and the main-config
// exemption policy that selected it together when an API save updates only a
// subset of the running config. Pairing a new hash with an older list makes
// the next periodic Verify report a false mismatch.
func applySignedIntegrityState(dst, signedMain *config.Config) {
dst.Integrity = signedMain.Integrity
dst.ConfD = signedMain.ConfD
}
// rejectIfConfDirChanged refuses to bless a save when the drop-ins on disk no
// longer match what disk (the main config as last signed) recorded. The
// exemption list of that same on-disk config decides which fragments count.
func rejectIfConfDirChanged(w http.ResponseWriter, confDir string, disk *config.Config) bool {
currentHash, err := integrity.HashConfDir(confDir, disk.ConfD.IntegrityExempt)
if err != nil {
writeJSONError(w, "hash conf.d: "+err.Error(), http.StatusInternalServerError)
return true
}
if currentHash != disk.Integrity.ConfdHash {
writeJSONError(w, "conf.d changed on disk, reload", http.StatusPreconditionFailed)
return true
}
return false
}
type fieldError struct {
Field string `json:"field"`
Message string `json:"message"`
}
func splitValidationResults(results []config.ValidationResult) (errs []fieldError, warnings []fieldError) {
for _, v := range results {
if v.Level == "error" {
errs = append(errs, fieldError{Field: v.Field, Message: v.Message})
continue
}
if v.Level == "warn" {
warnings = append(warnings, fieldError{Field: v.Field, Message: v.Message})
}
}
return errs, warnings
}
func localizeValidationFields(results []fieldError, section string) {
prefix := section + "."
for i := range results {
results[i].Field = strings.TrimPrefix(results[i].Field, prefix)
}
}
// writeValidationErrors answers 422 with the one error message every
// failure carries and the per-field problems next to it.
func writeValidationErrors(w http.ResponseWriter, errs []fieldError) {
msg := "Invalid values"
if len(errs) == 1 {
msg = errs[0].Field + ": " + errs[0].Message
}
writeJSONStatus(w, http.StatusUnprocessableEntity, map[string]interface{}{"error": msg, "errors": errs})
}
const fileOnlyFieldMessage = "Change this in csm.yaml. The web UI cannot set commands, file paths, sockets or environment variable names."
// credentialRebindErrors refuses a URL change that would send the stored
// credential of that URL to a new address unless the same save enters the
// credential again. A credential read from an environment variable cannot be
// re-entered here, so its address can only change in csm.yaml.
func credentialRebindErrors(section SettingsSection, current, candidate *config.Config, changes map[string]json.RawMessage) ([]fieldError, error) {
var errs []fieldError
var before, after map[string]interface{}
for _, field := range section.Fields {
if field.CredentialField == "" {
continue
}
if before == nil {
var err error
before, err = extractSectionEffectiveValues(current, section)
if err != nil {
return nil, err
}
after, err = extractSectionEffectiveValues(candidate, section)
if err != nil {
return nil, err
}
}
if settingsStringAt(before, field.YAMLPath) == settingsStringAt(after, field.YAMLPath) {
continue
}
if env := settingsStringAt(after, field.CredentialEnvField); env != "" {
errs = append(errs, fieldError{Field: field.YAMLPath, Message: "The credential for this address is read from environment variable " + env + ". Change the address in csm.yaml."})
continue
}
secret := settingsStringAt(after, field.CredentialField)
if secret == "" {
continue
}
if entered := enteredSecret(changes[field.CredentialField]); entered == "" || entered != secret {
errs = append(errs, fieldError{Field: field.YAMLPath, Message: "Enter the effective credential again when changing this address. If a conf.d drop-in overrides it, change the address and credential in the configuration files."})
}
}
return errs, nil
}
// settingsStringAt returns the string at a dotted path inside a section's
// effective values, or "" when the path is absent or not a string.
func settingsStringAt(values map[string]interface{}, dotted string) string {
var cur interface{} = values
for _, part := range strings.Split(dotted, ".") {
m, ok := cur.(map[string]interface{})
if !ok {
return ""
}
cur = m[part]
}
s, _ := cur.(string)
return s
}
// enteredSecret returns a secret the request actually supplies: absent, empty
// and the redaction placeholder all mean "keep the stored value".
func enteredSecret(raw json.RawMessage) string {
if raw == nil {
return ""
}
var v string
if err := json.Unmarshal(raw, &v); err != nil || v == config.RedactedValue || strings.TrimSpace(v) == "" {
return ""
}
return v
}
func buildChangeSet(section SettingsSection, clone *config.Config, changes map[string]json.RawMessage) ([]config.YAMLChange, []fieldError) {
var out []config.YAMLChange
var errs []fieldError
for key, raw := range changes {
field := lookupSchemaField(section, key)
if field == nil {
errs = append(errs, fieldError{Field: key, Message: "unknown field"})
continue
}
if field.FileOnly {
errs = append(errs, fieldError{Field: key, Message: fileOnlyFieldMessage})
continue
}
if field.Secret {
var sv string
if err := json.Unmarshal(raw, &sv); err == nil && sv == config.RedactedValue {
continue
}
}
if field.Type == "[]enum" {
if field.OptionsSource == "disabled_check_names" || field.OptionsSource == "check_names" {
var err error
raw, err = normaliseDisabledCheckNamesRaw(raw)
if err != nil {
errs = append(errs, fieldError{Field: key, Message: "decode: " + err.Error()})
continue
}
}
if badValues, ok := validateEnumArray(field, raw); !ok {
for _, bv := range badValues {
errs = append(errs, fieldError{Field: key, Message: "unknown value: " + bv})
}
continue
}
}
if field.Type == "enum" {
if badValue, ok := validateEnumScalar(field, raw); !ok {
errs = append(errs, fieldError{Field: key, Message: "unknown value: " + badValue})
continue
}
}
if field.Type == "[]int" {
normalised, badValues, perr := normaliseIntArray(field, raw)
if perr != nil {
errs = append(errs, fieldError{Field: key, Message: perr.Error()})
continue
}
if len(badValues) > 0 {
for _, bv := range badValues {
errs = append(errs, fieldError{Field: key, Message: "invalid value: " + bv})
}
continue
}
raw = normalised
}
// For float fields, coerce JSON string -> JSON number so the downstream
// json.Unmarshal into *float64 (in applyToClone) succeeds.
if field.Type == "float" {
if normalised, ok := coerceFloatRaw(raw); ok {
raw = normalised
} else {
errs = append(errs, fieldError{Field: key, Message: "decode: expected float"})
continue
}
}
fullPath := section.YAMLPath
if key != "" {
fullPath = section.YAMLPath + "." + key
}
decoded, err := decodeJSONForYAML(raw, field)
if err != nil {
errs = append(errs, fieldError{Field: key, Message: "decode: " + err.Error()})
continue
}
out = append(out, config.YAMLChange{Path: strings.Split(fullPath, "."), Value: decoded})
if err := applyToClone(clone, strings.Split(fullPath, "."), raw); err != nil {
errs = append(errs, fieldError{Field: key, Message: err.Error()})
}
}
return out, errs
}
func validateEnumScalar(field *SettingsField, raw json.RawMessage) (bad string, ok bool) {
var value string
if err := json.Unmarshal(raw, &value); err != nil {
return "(not a string)", false
}
resolved := resolvedOptionsForField(field)
if len(resolved) == 0 {
return "", true
}
for _, opt := range resolved {
if value == opt {
return "", true
}
}
return value, false
}
func normaliseDisabledCheckNamesRaw(raw json.RawMessage) (json.RawMessage, error) {
var values []string
if err := json.Unmarshal(raw, &values); err != nil {
return nil, fmt.Errorf("expected string array")
}
seen := make(map[string]struct{}, len(values))
out := make([]string, 0, len(values))
for _, value := range values {
value = config.CanonicalCheckName(strings.TrimSpace(value))
if value == "" {
continue
}
if _, ok := seen[value]; ok {
continue
}
seen[value] = struct{}{}
out = append(out, value)
}
encoded, err := json.Marshal(out)
if err != nil {
return nil, err
}
return encoded, nil
}
// validateEnumArray checks that raw is a JSON array of strings, each of
// which appears in the field's resolved Options. Returns the slice of
// unknown values and ok=false when any value is out-of-set. An empty
// array is allowed (clears the list). Calls resolveFieldOptions-style
// logic by inspecting OptionsSource + Options directly so validation
// doesn't depend on GET ordering.
func validateEnumArray(field *SettingsField, raw json.RawMessage) (bad []string, ok bool) {
var values []string
if err := json.Unmarshal(raw, &values); err != nil {
return []string{"(not a string array)"}, false
}
resolved := resolvedOptionsForField(field)
if len(resolved) == 0 {
return nil, true
}
allowed := make(map[string]struct{}, len(resolved))
for _, v := range resolved {
allowed[v] = struct{}{}
}
seen := map[string]struct{}{}
for _, v := range values {
if _, dup := seen[v]; dup {
continue
}
seen[v] = struct{}{}
if _, okk := allowed[v]; !okk {
bad = append(bad, v)
}
}
return bad, len(bad) == 0
}
// resolvedOptionsForField returns the flat list of allowed values for a
// []enum field, whether it uses static Options or an OptionsSource. Keeps
// POST-side validation independent of the GET-time mutation.
func resolvedOptionsForField(field *SettingsField) []string {
tmp := &SettingsField{Type: field.Type, OptionsSource: field.OptionsSource}
switch field.OptionsSource {
case "check_names":
applyCheckNameOptions(tmp)
case "disabled_check_names":
// Validation accepts a wider set than the UI dropdown (which lists
// public finding names only): the scheduler also honors compatibility
// runner IDs, so rejecting them here would break existing configs.
return disabledCheckValidationOptions()
case "geoip_editions":
applyGeoIPEditionOptions(tmp)
}
if len(tmp.Options) > 0 {
return tmp.Options
}
if len(field.Options) > 0 {
return field.Options
}
return tmp.Options
}
// normaliseIntArray parses raw as a JSON array of integers (or strings that
// parse as integers), enforces field.Min/Max as the per-element bound (default
// 1..65535 for port-list semantics when both are nil), deduplicates, and
// returns a JSON array of distinct ascending integers. Returns the offending
// raw values as badValues when any element is outside the allowed range; the
// parse error path is reserved for malformed JSON.
func normaliseIntArray(field *SettingsField, raw json.RawMessage) (json.RawMessage, []string, error) {
var items []json.RawMessage
if err := json.Unmarshal(raw, &items); err != nil {
return nil, nil, fmt.Errorf("expected array of ints: %s", err.Error())
}
minV := int64(1)
maxV := int64(65535)
if field.Min != nil {
minV = *field.Min
}
if field.Max != nil {
maxV = *field.Max
}
seen := make(map[int64]struct{}, len(items))
out := make([]int64, 0, len(items))
var bad []string
for _, item := range items {
var n int64
if err := json.Unmarshal(item, &n); err != nil {
var s string
if serr := json.Unmarshal(item, &s); serr != nil {
bad = append(bad, string(item))
continue
}
s = strings.TrimSpace(s)
if s == "" {
continue
}
parsed, perr := strconv.ParseInt(s, 10, 64)
if perr != nil {
bad = append(bad, s)
continue
}
n = parsed
}
if n < minV || n > maxV {
bad = append(bad, strconv.FormatInt(n, 10))
continue
}
if _, dup := seen[n]; dup {
continue
}
seen[n] = struct{}{}
out = append(out, n)
}
if len(bad) > 0 {
return nil, bad, nil
}
sort.Slice(out, func(i, j int) bool { return out[i] < out[j] })
encoded, err := json.Marshal(out)
if err != nil {
return nil, nil, err
}
return encoded, nil, nil
}
// coerceFloatRaw returns a JSON-number representation of raw if raw is either
// a JSON number already or a JSON string that parses to float64. The second
// return is false if neither form is valid.
func coerceFloatRaw(raw json.RawMessage) (json.RawMessage, bool) {
var asNum float64
if err := json.Unmarshal(raw, &asNum); err == nil {
b, _ := json.Marshal(asNum)
return b, true
}
var asStr string
if err := json.Unmarshal(raw, &asStr); err == nil {
f, err := strconv.ParseFloat(asStr, 64)
if err != nil {
return nil, false
}
b, _ := json.Marshal(f)
return b, true
}
return nil, false
}
func lookupSchemaField(section SettingsSection, key string) *SettingsField {
for i := range section.Fields {
if section.Fields[i].YAMLPath == key {
return §ion.Fields[i]
}
}
return nil
}
func decodeJSONForYAML(raw json.RawMessage, field *SettingsField) (interface{}, error) {
if string(raw) == "null" {
if !field.Nullable {
return nil, fmt.Errorf("null is only allowed for nullable fields")
}
return nil, nil
}
if field.Type == "float" {
// Accept either a JSON number or a JSON string containing a number.
var asNum float64
if err := json.Unmarshal(raw, &asNum); err == nil {
return asNum, nil
}
var asStr string
if err := json.Unmarshal(raw, &asStr); err == nil {
f, perr := strconv.ParseFloat(asStr, 64)
if perr != nil {
return nil, fmt.Errorf("not a float: %q", asStr)
}
return f, nil
}
return nil, fmt.Errorf("expected float, got %s", string(raw))
}
var v interface{}
if err := json.Unmarshal(raw, &v); err != nil {
return nil, err
}
return v, nil
}
func applyToClone(cfg *config.Config, path []string, raw json.RawMessage) error {
v := reflect.ValueOf(cfg).Elem()
for i, key := range path {
if v.Kind() == reflect.Pointer {
if v.IsNil() {
v.Set(reflect.New(v.Type().Elem()))
}
v = v.Elem()
}
if v.Kind() != reflect.Struct {
return fmt.Errorf("path %v: element %d is not a struct", path, i)
}
field, ok := fieldByYAMLTag(v.Type(), key)
if !ok {
return fmt.Errorf("no yaml field %q under %s", key, strings.Join(path[:i], "."))
}
v = v.FieldByIndex(field.Index)
}
ptr := reflect.New(v.Type())
if err := json.Unmarshal(raw, ptr.Interface()); err != nil {
return fmt.Errorf("unmarshal into %s: %w", v.Type(), err)
}
v.Set(ptr.Elem())
return nil
}
func fieldByYAMLTag(t reflect.Type, yamlName string) (reflect.StructField, bool) {
for i := 0; i < t.NumField(); i++ {
f := t.Field(i)
tag := f.Tag.Get("yaml")
if tag == "" {
continue
}
name := tag
if idx := strings.IndexByte(tag, ','); idx >= 0 {
name = tag[:idx]
}
if name == yamlName {
return f, true
}
}
return reflect.StructField{}, false
}
func auditDetailsFor(section SettingsSection, changes map[string]json.RawMessage) string {
redacted := make(map[string]interface{}, len(changes))
for k, raw := range changes {
if field := lookupSchemaField(section, k); field != nil && field.Secret {
redacted[k] = "***"
continue
}
var v interface{}
_ = json.Unmarshal(raw, &v)
redacted[k] = v
}
b, _ := json.Marshal(redacted)
return string(b)
}
// defaultRestartDaemon is the production implementation. Tests override
// s.restartDaemon with a fake.
func defaultRestartDaemon() ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
// #nosec G204 -- fixed argv, no operator input interpolated.
cmd := exec.CommandContext(ctx, "systemctl", "restart", "csm")
return cmd.CombinedOutput()
}
// settingsRestartDelay is how long the restart is deferred after the 202 is
// written, so the acknowledgement reaches the client before systemctl restart
// SIGTERMs this process. Overridden in tests to keep them fast.
var settingsRestartDelay = 250 * time.Millisecond
// apiSettingsRestart handles POST /api/v1/settings/restart. It schedules the
// restart asynchronously and returns 202 immediately: `systemctl restart csm`
// SIGTERMs this very process, so a synchronous restart cannot send a clean
// response -- the call returns "signal: terminated" and the old handler turned
// that into a spurious 500 even though the restart had succeeded. The frontend
// treats 202 as "restart issued" and polls for the daemon to return; a failed
// restart surfaces as the daemon not coming back, and the error is logged
// server-side by scheduleDaemonRestart.
func (s *Server) apiSettingsRestart(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
s.auditLog(r, "settings-restart", "daemon", "")
writeOKStatus(w, http.StatusAccepted, map[string]interface{}{
"started_at_token": s.daemonStartToken(),
})
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
s.scheduleDaemonRestart(settingsRestartDelay)
}
package webui
import (
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/geoip"
)
// resolveFieldOptions populates Options / OptionGroups for any []enum field
// that declares an OptionsSource. Called once per GET /api/v1/settings/:id
// so the UI always sees a fresh list (for check_names this matters if the
// registry grows over time).
//
// The section's Fields slice header points into the package-level
// settingsSections backing array. Writing into it directly would race
// across concurrent requests. Copy the slice first so mutations are local
// to this request.
func resolveFieldOptions(section *SettingsSection) {
fields := make([]SettingsField, len(section.Fields))
copy(fields, section.Fields)
section.Fields = fields
for i := range section.Fields {
f := §ion.Fields[i]
if f.OptionsSource == "" {
continue
}
switch f.OptionsSource {
case "check_names":
applyCheckNameOptions(f)
case "disabled_check_names":
applyDisabledCheckNameOptions(f)
case "geoip_editions":
applyGeoIPEditionOptions(f)
}
}
}
func applyCheckNameOptions(f *SettingsField) {
infos := checks.PublicCheckInfos()
byCategory := make(map[string][]string)
var order []string
for _, info := range infos {
if _, ok := byCategory[info.Category]; !ok {
order = append(order, info.Category)
}
byCategory[info.Category] = append(byCategory[info.Category], info.Name)
}
groups := make([]OptionGroup, 0, len(order))
flat := make([]string, 0, len(infos))
for _, cat := range order {
groups = append(groups, OptionGroup{Label: cat, Values: byCategory[cat]})
flat = append(flat, byCategory[cat]...)
}
f.Options = flat
f.OptionGroups = groups
}
func applyDisabledCheckNameOptions(f *SettingsField) {
allowed := make(map[string]struct{})
for _, name := range checks.DisabledCheckNames() {
allowed[name] = struct{}{}
}
infos := checks.PublicCheckInfos()
byCategory := make(map[string][]string)
var order []string
for _, info := range infos {
if _, ok := allowed[info.Name]; !ok {
continue
}
if _, ok := byCategory[info.Category]; !ok {
order = append(order, info.Category)
}
byCategory[info.Category] = append(byCategory[info.Category], info.Name)
}
groups := make([]OptionGroup, 0, len(order))
flat := make([]string, 0, len(allowed))
for _, cat := range order {
groups = append(groups, OptionGroup{Label: cat, Values: byCategory[cat]})
flat = append(flat, byCategory[cat]...)
}
f.Options = flat
f.OptionGroups = groups
}
// disabledCheckValidationOptions returns the values a POST to the scheduled-scan
// disabled_checks list may contain: public finding names plus the compatibility
// runner IDs the scheduler still honors. Wider than the UI option set so saving
// an existing config that disables a check by runner ID does not get rejected.
func disabledCheckValidationOptions() []string {
return checks.DisabledCheckConfigNames()
}
func applyGeoIPEditionOptions(f *SettingsField) {
free, commercial := geoip.KnownEditions()
flat := append([]string{}, free...)
flat = append(flat, commercial...)
f.Options = flat
f.OptionGroups = []OptionGroup{
{Label: "GeoLite2 (free)", Values: free},
{Label: "GeoIP2 (paid)", Values: commercial},
}
}
package webui
import "github.com/pidginhost/csm/internal/config"
// OptionGroup is an ordered label + values pair used to render grouped
// multi-select options (e.g. "Authentication & Login" → [cpanel_login, ...]).
type OptionGroup struct {
Label string `json:"label"`
Values []string `json:"values"`
}
// SettingsField describes a single editable leaf within a settings
// section. YAMLPath is the dotted key path relative to the section's
// YAMLPath. For example inside the Alerts section, the field with
// YAMLPath "email.enabled" has full path "alerts.email.enabled".
//
// For Type "[]enum" fields, Options and/or OptionGroups are resolved at
// request time. A field may either declare a static Options list or set
// OptionsSource to have the handler populate Options/OptionGroups from a
// registry ("check_names", "geoip_editions").
type SettingsField struct {
YAMLPath string `json:"yaml_path"`
Type string `json:"type"`
Label string `json:"label"`
Help string `json:"help,omitempty"`
Secret bool `json:"secret,omitempty"`
Nullable bool `json:"nullable,omitempty"`
Min *int64 `json:"min,omitempty"`
Max *int64 `json:"max,omitempty"`
Options []string `json:"options,omitempty"`
OptionGroups []OptionGroup `json:"option_groups,omitempty"`
OptionsSource string `json:"options_source,omitempty"`
Placeholder string `json:"placeholder,omitempty"`
// FieldGroup is the inner subdivider label rendered as a fieldset
// inside a single section (e.g. firewall fields split into Access
// ports / Rate limits / Logging). Empty string means the field
// renders ungrouped under the section's flat grid.
FieldGroup string `json:"field_group,omitempty"`
// FileOnly fields are shown but can only be changed in csm.yaml. They
// name a command, executable, path, socket or environment variable the
// root daemon acts on, so a browser session must not be able to set them.
FileOnly bool `json:"file_only,omitempty"`
// CredentialField and CredentialEnvField name the secret this URL field
// sends to its address. A web UI change of the URL must re-enter that
// secret, so a session cannot point a stored credential at a host of its
// choosing; a secret read from the environment can only move in csm.yaml.
CredentialField string `json:"-"`
CredentialEnvField string `json:"-"`
}
// SettingsSection groups the fields of one top-level Config sub-tree.
// YAMLPath is the root key in csm.yaml (e.g. "auto_response"). ID is
// the URL-path identifier used by the API. Restart is a UI hint based
// on the current hotreload struct tag; final safe-vs-restart authority
// comes from config.Diff at runtime. Icon is a Tabler icon suffix (e.g.
// "bell" for "ti ti-bell"); Group is the nav category the section lives
// in ("Alerting", "Detection", "Integrations", "Ops").
type SettingsSection struct {
ID string `json:"id"`
Title string `json:"title"`
YAMLPath string `json:"yaml_path"`
Restart bool `json:"restart_hint"`
ReloadTag string `json:"reload_tag,omitempty"`
Icon string `json:"icon,omitempty"`
Group string `json:"group,omitempty"`
Fields []SettingsField `json:"fields"`
}
// Section groups for the sidebar. Order here defines order in the UI.
const (
SectionGroupAlerting = "Alerting"
SectionGroupDetection = "Detection"
SectionGroupFirewall = "Firewall"
SectionGroupIntegrations = "Integrations"
SectionGroupOps = "Operations"
)
// SectionGroupOrder is the display order of sidebar group headers.
var SectionGroupOrder = []string{
SectionGroupAlerting,
SectionGroupDetection,
SectionGroupFirewall,
SectionGroupIntegrations,
SectionGroupOps,
}
// Field-group labels used to subdivide large sections. Reuse via
// constants so Phase 6 does not scatter free strings through the
// schema and the static test can lock them in.
const (
FieldGroupAccessPorts = "Access ports"
FieldGroupIPv6 = "IPv6"
FieldGroupRateLimits = "Rate limits"
FieldGroupFloodProtection = "Flood protection"
FieldGroupGeoDynDNS = "Geo and DynDNS"
FieldGroupSMTPControls = "SMTP controls"
FieldGroupLogging = "Logging"
FieldGroupLimits = "Limits"
FieldGroupScanIntervals = "Scan intervals"
FieldGroupWebBruteForce = "Web brute force" // #nosec G101 -- UI label, not a credential.
FieldGroupMailBruteForce = "Mail brute force"
FieldGroupSMTPBruteForce = "SMTP brute force"
FieldGroupAccountSpray = "Account spray"
FieldGroupStateRetention = "State retention"
FieldGroupAbuseReporting = "Abuse reporting"
FieldGroupCentralDB = "Central database"
FieldGroupBlockDigest = "Block digest"
)
func int64p(v int64) *int64 { return &v }
var settingsSections = []SettingsSection{
{
ID: "alerts",
Title: "Alerts",
YAMLPath: "alerts",
Icon: "bell",
Group: SectionGroupAlerting,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "email.enabled", Type: "bool", Label: "Email alerts enabled"},
{YAMLPath: "email.to", Type: "[]string", Label: "Recipients", Help: "One email address per line"},
{YAMLPath: "email.from", Type: "string", Label: "From address"},
{YAMLPath: "email.smtp", Type: "string", Label: "SMTP server", Placeholder: "smtp.example.com:587"},
{YAMLPath: "email.disabled_checks", Type: "[]enum", Label: "Disabled check names", OptionsSource: "check_names", Help: "Findings with these check names never trigger email alerts."},
{YAMLPath: "webhook.enabled", Type: "bool", Label: "Webhook alerts enabled"},
{YAMLPath: "webhook.url", Type: "string", Label: "Webhook URL"},
{YAMLPath: "webhook.type", Type: "enum", Label: "Webhook type", Options: []string{"slack", "discord", "generic", "phpanel"}},
{YAMLPath: "webhook.hmac_secret", Type: "string", Label: "Webhook HMAC secret", Secret: true},
{YAMLPath: "webhook.hmac_secret_env", Type: "string", Label: "Webhook HMAC secret env", FileOnly: true},
{YAMLPath: "webhook.per_finding", Type: "bool", Label: "Per-finding webhook delivery"},
{YAMLPath: "heartbeat.enabled", Type: "bool", Label: "Heartbeat enabled"},
{YAMLPath: "heartbeat.url", Type: "string", Label: "Heartbeat URL"},
{YAMLPath: "max_per_hour", Type: "int", Label: "Max alerts per hour", Min: int64p(0), Max: int64p(10000)},
{YAMLPath: "block_digest.enabled", Type: "bool", Label: "Block digest enabled", FieldGroup: FieldGroupBlockDigest},
{YAMLPath: "block_digest.countries", Type: "[]string", Label: "Block digest countries", Help: "Two-letter country codes, one per line. Empty falls back to trusted countries, then all countries.", FieldGroup: FieldGroupBlockDigest},
{YAMLPath: "block_digest.interval", Type: "string", Label: "Block digest interval", Placeholder: "1h", FieldGroup: FieldGroupBlockDigest},
{YAMLPath: "block_digest.live", Type: "bool", Label: "Send live block digest alerts", FieldGroup: FieldGroupBlockDigest},
{YAMLPath: "block_digest.send_on", Type: "enum", Label: "Block digest send policy", Options: []string{"any", "customer"}, FieldGroup: FieldGroupBlockDigest},
{YAMLPath: "block_digest.channel", Type: "enum", Label: "Block digest channel", Options: []string{"", "email", "webhook"}, FieldGroup: FieldGroupBlockDigest},
{YAMLPath: "block_digest.min_block", Type: "int", Label: "Block digest minimum blocks", Min: int64p(0), Max: int64p(100000), FieldGroup: FieldGroupBlockDigest},
},
},
{
ID: "thresholds",
Title: "Thresholds",
YAMLPath: "thresholds",
Icon: "adjustments",
Group: SectionGroupAlerting,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "mail_queue_warn", Type: "int", Label: "Mail queue warn", Min: int64p(0), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "mail_queue_crit", Type: "int", Label: "Mail queue critical", Min: int64p(0), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "state_expiry_hours", Type: "int", Label: "State expiry (hours)", Min: int64p(1), FieldGroup: FieldGroupStateRetention},
{YAMLPath: "deep_scan_interval_min", Type: "int", Label: "Deep scan interval (min)", Min: int64p(1), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "wp_core_check_interval_min", Type: "int", Label: "WP core check interval (min)", Min: int64p(1), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "webshell_scan_interval_min", Type: "int", Label: "Webshell scan interval (min)", Min: int64p(1), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "filesystem_scan_interval_min", Type: "int", Label: "Filesystem scan interval (min)", Min: int64p(1), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "multi_ip_login_threshold", Type: "int", Label: "Multi-IP login threshold", Min: int64p(1), FieldGroup: FieldGroupAccountSpray},
{YAMLPath: "multi_ip_login_window_min", Type: "int", Label: "Multi-IP login window (min)", Min: int64p(1), FieldGroup: FieldGroupAccountSpray},
{YAMLPath: "cred_stuffing_distinct_accounts", Type: "int", Label: "Credential stuffing distinct accounts", Min: int64p(2), Max: int64p(200), FieldGroup: FieldGroupAccountSpray, Help: "Distinct failed accounts from one source IP inside the auth window before credential_stuffing fires. Default 5."},
{YAMLPath: "pam_bruteforce_threshold", Type: "int", Label: "PAM brute-force threshold", Min: int64p(2), Max: int64p(1000), FieldGroup: FieldGroupAccountSpray, Help: "PAM authentication failures from one source IP inside the window before pam_bruteforce fires and the address is auto-blocked. Default 5."},
{YAMLPath: "pam_bruteforce_window_min", Type: "int", Label: "PAM brute-force window (min)", Min: int64p(1), Max: int64p(1440), FieldGroup: FieldGroupAccountSpray, Help: "Minutes over which PAM failures from one source IP are counted. Default 10."},
{YAMLPath: "plugin_check_interval_min", Type: "int", Label: "Plugin check interval (min)", Min: int64p(1), FieldGroup: FieldGroupScanIntervals},
{YAMLPath: "brute_force_window", Type: "int", Label: "Brute force window", Min: int64p(1), FieldGroup: FieldGroupWebBruteForce},
{YAMLPath: "domlog_max_files", Type: "int", Label: "Domlog max files", Min: int64p(1), Max: int64p(100000), FieldGroup: FieldGroupWebBruteForce},
{YAMLPath: "domlog_tail_lines", Type: "int", Label: "Domlog tail lines", Min: int64p(10), Max: int64p(100000), FieldGroup: FieldGroupWebBruteForce, Help: "Trailing lines tailed from each per-domain access log per WP brute-force cycle. Default 500 covers ~10 minutes of traffic on a busy site."},
{YAMLPath: "domlog_max_age_min", Type: "int", Label: "Domlog max age (min)", Min: int64p(1), Max: int64p(1440), FieldGroup: FieldGroupWebBruteForce, Help: "Skip per-domain access logs untouched for this many minutes. Default 30. Raise on low-traffic hosts where a slow-burn dictionary attack against a quiet domain still needs to fall inside the freshness window."},
{YAMLPath: "mail_log_tail_lines", Type: "int", Label: "Mail log tail lines", Min: int64p(10), Max: int64p(100000), FieldGroup: FieldGroupMailBruteForce, Help: "Trailing lines of /var/log/exim_mainlog read by the per-account mail rate scanner. Default 500. Raise on busy mail hosts where a single account's spam burst spreads across more than 500 lines per cycle."},
{YAMLPath: "syslog_messages_tail_lines", Type: "int", Label: "Syslog messages tail lines", Min: int64p(10), Max: int64p(100000), FieldGroup: FieldGroupScanIntervals, Help: "Trailing lines of /var/log/messages read only by the legacy direct-check FTP path. The daemon's FTP brute-force detector follows the log forward-only and uses ftp_fail_window_min instead. Default 200."},
{YAMLPath: "ftp_fail_window_min", Type: "int", Label: "FTP fail window (min)", Min: int64p(1), Max: int64p(1440), FieldGroup: FieldGroupScanIntervals, Help: "Sliding window in minutes over which the FTP brute-force detector accumulates per-IP pure-ftpd auth failures before raising ftp_bruteforce. Default 30."},
{YAMLPath: "exposed_file_scan_depth", Type: "int", Label: "Exposed-file scan depth", Min: int64p(1), Max: int64p(int64(config.MaxExposedFileScanDepth)), FieldGroup: FieldGroupLimits, Help: "Directory levels below each web document root inspected for downloadable sensitive files. Default 2; maximum 10."},
{YAMLPath: "account_scan_max_files", Type: "int", Label: "Account scan max files", Min: int64p(1), Max: int64p(100000), FieldGroup: FieldGroupLimits, Help: "Per-cycle cap for account and mail-domain scanner paths. Default 10000. Raise on very large multi-tenant hosts."},
{YAMLPath: "crontab_base64_blob_max_bytes", Type: "int", Label: "Crontab base64 blob max bytes", Min: int64p(1024), Max: int64p(1048576), FieldGroup: FieldGroupLimits, Help: "Encoded-byte cap for one crontab base64 candidate before decoded-content matching. Default 16384. Must be a multiple of 4."},
{YAMLPath: "http_flood_threshold", Type: "int", Label: "HTTP flood threshold", Min: int64p(0), FieldGroup: FieldGroupWebBruteForce, Help: "Per-IP requests per window that emits http_request_flood. 0 disables. Sample baseline first."},
{YAMLPath: "http_flood_window_min", Type: "int", Label: "HTTP flood window (min)", Min: int64p(1), FieldGroup: FieldGroupWebBruteForce},
{YAMLPath: "http_ua_spoof_threshold", Type: "int", Label: "UA spoof threshold", Min: int64p(1), FieldGroup: FieldGroupWebBruteForce},
{YAMLPath: "xmlrpc_threshold", Type: "int", Label: "XML-RPC threshold", Min: int64p(0), FieldGroup: FieldGroupWebBruteForce, Help: "Per-IP POST /xmlrpc.php count before access-log xmlrpc_abuse. 0 disables."},
{YAMLPath: "http_distributed_min_ips", Type: "int", Label: "Distributed HTTP min IPs", Min: int64p(0), FieldGroup: FieldGroupWebBruteForce, Help: "Distinct already-abusive source IPs per vhost before the distributed HTTP flood rollup fires. 0 disables."},
{YAMLPath: "http_scanner_min_requests", Type: "int", Label: "URL scanner min requests", Min: int64p(0), FieldGroup: FieldGroupWebBruteForce, Help: "Per-IP request volume before http_scanner_profile evaluates. 0 disables."},
{YAMLPath: "http_scanner_error_pct", Type: "int", Label: "URL scanner error percent", Min: int64p(1), Max: int64p(100), FieldGroup: FieldGroupWebBruteForce, Help: "Minimum percentage of requests that must return a probe-error status. Default 90."},
{YAMLPath: "http_scanner_min_distinct_paths", Type: "int", Label: "URL scanner distinct paths", Min: int64p(1), Max: int64p(int64(config.HTTPScannerMaxDistinctPaths)), FieldGroup: FieldGroupWebBruteForce, Help: "Minimum distinct error paths after query strings are stripped. Default 10."},
{YAMLPath: "http_scanner_status_codes", Type: "[]int", Label: "URL scanner status codes", Min: int64p(100), Max: int64p(599), FieldGroup: FieldGroupWebBruteForce, Help: "HTTP statuses counted as probe errors. Default 404 and 403; add redirects only when legitimate traffic cannot be redirect-heavy."},
{YAMLPath: "smtp_bruteforce_threshold", Type: "int", Label: "SMTP bruteforce threshold", Min: int64p(1), FieldGroup: FieldGroupSMTPBruteForce},
{YAMLPath: "smtp_bruteforce_window_min", Type: "int", Label: "SMTP bruteforce window (min)", Min: int64p(1), FieldGroup: FieldGroupSMTPBruteForce},
{YAMLPath: "smtp_bruteforce_suppress_min", Type: "int", Label: "SMTP bruteforce suppress (min)", Min: int64p(1), FieldGroup: FieldGroupSMTPBruteForce},
{YAMLPath: "smtp_bruteforce_subnet_threshold", Type: "int", Label: "SMTP bruteforce /24 threshold", Min: int64p(1), FieldGroup: FieldGroupSMTPBruteForce},
{YAMLPath: "smtp_account_spray_threshold", Type: "int", Label: "SMTP account spray threshold", Min: int64p(1), FieldGroup: FieldGroupAccountSpray},
{YAMLPath: "smtp_bruteforce_max_tracked", Type: "int", Label: "SMTP bruteforce max tracked", Min: int64p(100), FieldGroup: FieldGroupStateRetention},
{YAMLPath: "smtp_bruteforce_slow_threshold", Type: "int", Label: "SMTP slow bruteforce threshold", Min: int64p(0), Max: int64p(int64(config.SlowBruteMaxThreshold)), FieldGroup: FieldGroupSMTPBruteForce, Help: "Failed auths from one IP across at least three mailboxes inside the slow window; 10 distinct mailboxes trigger on breadth alone. Default 40; 0 disables; nonzero values must be at least 3."},
{YAMLPath: "smtp_bruteforce_slow_window_min", Type: "int", Label: "SMTP slow bruteforce window (min)", Min: int64p(1), Max: int64p(int64(config.SlowBruteMaxWindowMin)), FieldGroup: FieldGroupSMTPBruteForce, Help: "Long-horizon sliding window for paced SMTP auth failures. Default 360; maximum one week."},
{YAMLPath: "mail_bruteforce_threshold", Type: "int", Label: "Mail bruteforce threshold", Min: int64p(1), FieldGroup: FieldGroupMailBruteForce},
{YAMLPath: "mail_bruteforce_window_min", Type: "int", Label: "Mail bruteforce window (min)", Min: int64p(1), FieldGroup: FieldGroupMailBruteForce},
{YAMLPath: "mail_bruteforce_suppress_min", Type: "int", Label: "Mail bruteforce suppress (min)", Min: int64p(1), FieldGroup: FieldGroupMailBruteForce},
{YAMLPath: "mail_bruteforce_subnet_threshold", Type: "int", Label: "Mail bruteforce /24 threshold", Min: int64p(1), FieldGroup: FieldGroupMailBruteForce},
{YAMLPath: "mail_account_spray_threshold", Type: "int", Label: "Mail account spray threshold", Min: int64p(1), FieldGroup: FieldGroupAccountSpray},
{YAMLPath: "mail_bruteforce_max_tracked", Type: "int", Label: "Mail bruteforce max tracked", Min: int64p(100), FieldGroup: FieldGroupStateRetention},
{YAMLPath: "mail_bruteforce_slow_threshold", Type: "int", Label: "Mail slow bruteforce threshold", Min: int64p(0), Max: int64p(int64(config.SlowBruteMaxThreshold)), FieldGroup: FieldGroupMailBruteForce, Help: "Failed auths from one IP across at least three mailboxes inside the slow window; 10 distinct mailboxes trigger on breadth alone. Default 40; 0 disables; nonzero values must be at least 3."},
{YAMLPath: "mail_bruteforce_slow_window_min", Type: "int", Label: "Mail slow bruteforce window (min)", Min: int64p(1), Max: int64p(int64(config.SlowBruteMaxWindowMin)), FieldGroup: FieldGroupMailBruteForce, Help: "Long-horizon sliding window for paced IMAP, POP3, and ManageSieve auth failures. Default 360; maximum one week."},
{YAMLPath: "mail_brute_account_key", Type: "string", Label: "Mail brute account-key extractor", FieldGroup: FieldGroupMailBruteForce, Placeholder: "builtin:dovecot-user", Help: "How to derive the account from a dovecot/postfix log line: builtin:dovecot-user (default), builtin:postfix-sasl, or regex:<pattern> where group 1 is the account."},
},
},
{
ID: "mail_logs",
Title: "Mail logs",
YAMLPath: "mail_logs",
Icon: "mail",
Group: SectionGroupDetection,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "source", Type: "enum", Label: "Log source", Options: []string{"auto", "file", "journal"}, Help: "auto: try platform default file then fall back to journal. file: require log file. journal: read systemd-journald (needs journal build tag)."},
{YAMLPath: "file", Type: "string", Label: "Log file override", Placeholder: "/var/log/maillog", Help: "Override the platform-default file path. Leave blank to keep the default.", FileOnly: true},
{YAMLPath: "units", Type: "[]string", Label: "Journal units", Help: "Systemd units to match when source=journal. One per line (e.g. postfix, dovecot)."},
},
},
{
ID: "suppressions",
Title: "Suppressions",
YAMLPath: "suppressions",
Icon: "volume-off",
Group: SectionGroupAlerting,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "upcp_window_start", Type: "string", Label: "UPCP window start (HH:MM)"},
{YAMLPath: "upcp_window_end", Type: "string", Label: "UPCP window end (HH:MM)"},
{YAMLPath: "known_api_tokens", Type: "[]string", Label: "Known API tokens (hashed)"},
{YAMLPath: "ignore_paths", Type: "[]string", Label: "Ignore paths"},
{YAMLPath: "suppress_webmail_alerts", Type: "bool", Label: "Suppress webmail login alerts"},
{YAMLPath: "suppress_cpanel_login_alerts", Type: "bool", Label: "Suppress cPanel login alerts"},
{YAMLPath: "suppress_blocked_alerts", Type: "bool", Label: "Suppress alerts on auto-blocked IPs"},
{YAMLPath: "trusted_countries", Type: "[]string", Label: "Trusted countries (ISO 3166-1 alpha-2)"},
},
},
{
ID: "auto_response",
Title: "Auto-Response",
YAMLPath: "auto_response",
Icon: "bolt",
Group: SectionGroupDetection,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Auto-response enabled"},
{YAMLPath: "kill_processes", Type: "bool", Label: "Kill malicious processes"},
{YAMLPath: "quarantine_files", Type: "bool", Label: "Quarantine malicious files"},
{YAMLPath: "block_ips", Type: "bool", Label: "Block attacker IPs"},
{YAMLPath: "block_expiry", Type: "string", Label: "Block expiry", Placeholder: "24h", Help: "Positive temporary-block duration. Default 24h."},
{YAMLPath: "max_blocks_per_hour", Type: "int", Label: "Max IP blocks per hour", Min: int64p(0), Help: "0 uses the default 50/hour cap."},
{YAMLPath: "enforce_permissions", Type: "bool", Label: "Auto-chmod 644 world/group-writable PHP"},
{YAMLPath: "fix_wp_cron", Type: "bool", Label: "Auto-disable WP-Cron + install system cron", Help: "On perf_wp_cron findings, edit wp-config.php and add a per-user cron. Tune interval/php under Performance."},
{YAMLPath: "block_cpanel_logins", Type: "bool", Label: "Block on cPanel/webmail login alerts"},
{YAMLPath: "http_scanner_action", Type: "enum", Label: "URL scanner response", Options: []string{"challenge", "block"}, Help: "Response to http_scanner_profile findings. challenge routes the IP to the PoW challenge when the challenge subsystem is enabled, falling back to a block when it is not; block always hard-blocks."},
{YAMLPath: "netblock", Type: "bool", Label: "Auto-block /24 on threshold"},
{YAMLPath: "netblock_threshold", Type: "int", Label: "Netblock threshold", Min: int64p(config.MinBlockEscalationCount), Help: "Blocked addresses in one IPv4 /24 or IPv6 /64 before the subnet itself is blocked. Default 3."},
{YAMLPath: "netblock_window", Type: "string", Label: "Netblock window", Placeholder: "168h", Help: "Positive window in which blocked addresses count toward the netblock threshold, including blocks that already expired. Default 168h."},
{YAMLPath: "permblock", Type: "bool", Label: "Auto-promote to permanent"},
{YAMLPath: "permblock_count", Type: "int", Label: "Temp blocks before permanent", Min: int64p(config.MinBlockEscalationCount), Help: "Temporary blocks inside the window before the address is blocked permanently. Default 4."},
{YAMLPath: "permblock_interval", Type: "string", Label: "Permblock window", Placeholder: "24h", Help: "Positive window for counting temporary blocks. Default 24h."},
{YAMLPath: "clean_database", Type: "bool", Label: "Auto-clean DB injections"},
{YAMLPath: "virtual_patch_exposed_files", Type: "enum", Label: "Virtual-patch web-exposed files", Options: []string{"off", "manual", "auto"}, Help: "Write .htaccess 'Require all denied' rules for confirmed web-exposed files. off: detect only. manual: apply via `csm virtual-patch`. auto: apply each scan except warning-only sample SQL (gated by dry_run)."},
{YAMLPath: "dry_run", Type: "bool", Label: "Auto-response dry run", Help: "Preview automatic IP blocks and web-exposed-file virtual patches without enforcing them."},
{YAMLPath: "verdict_callback.enabled", Type: "bool", Label: "Verdict callback hook"},
{YAMLPath: "verdict_callback.url", Type: "string", Label: "Verdict callback URL"},
{YAMLPath: "verdict_callback.hmac_secret", Type: "string", Label: "Verdict callback HMAC secret", Secret: true},
{YAMLPath: "verdict_callback.hmac_secret_env", Type: "string", Label: "Verdict callback HMAC secret env", FileOnly: true},
{YAMLPath: "verdict_callback.allow_unsigned", Type: "bool", Label: "Allow unsigned verdict callback"},
{YAMLPath: "verdict_callback.require_response_signature", Type: "bool", Label: "Require signed verdict response"},
{YAMLPath: "verdict_callback.timeout_sec", Type: "int", Label: "Verdict callback timeout (sec)", Min: int64p(1), Max: int64p(30)},
},
},
{
ID: "reputation",
Title: "Reputation",
YAMLPath: "reputation",
Icon: "shield-check",
Group: SectionGroupIntegrations,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "abuseipdb_key", Type: "string", Label: "AbuseIPDB API key", Secret: true},
{YAMLPath: "whitelist", Type: "[]string", Label: "Whitelisted IPs", Help: "Never flagged as malicious"},
{YAMLPath: "bot_verify_enabled", Type: "bool", Label: "Verify search-engine bots via rDNS", Nullable: true},
{YAMLPath: "rspamd.enabled", Type: "bool", Label: "Rspamd threat-intel"},
// #nosec G101 -- names of the config fields that hold the credential, not a credential.
{YAMLPath: "rspamd.url", Type: "string", Label: "Rspamd controller URL", CredentialField: "rspamd.token", CredentialEnvField: "rspamd.token_env"},
{YAMLPath: "rspamd.token", Type: "string", Label: "Rspamd controller password", Secret: true},
{YAMLPath: "rspamd.token_env", Type: "string", Label: "Rspamd password env var", FileOnly: true},
{YAMLPath: "upstream.enabled", Type: "bool", Label: "Upstream threat-intel cache"},
// #nosec G101 -- names of the config fields that hold the credential, not a credential.
{YAMLPath: "upstream.url", Type: "string", Label: "Upstream URL", CredentialField: "upstream.token", CredentialEnvField: "upstream.token_env"},
{YAMLPath: "upstream.token", Type: "string", Label: "Upstream bearer token", Secret: true},
{YAMLPath: "upstream.token_env", Type: "string", Label: "Upstream token env var", FileOnly: true},
{YAMLPath: "upstream.cache_ttl_min", Type: "int", Label: "Upstream cache TTL (min)", Min: int64p(1), Max: int64p(1440)},
{YAMLPath: "upstream.timeout_sec", Type: "int", Label: "Upstream request timeout (sec)", Min: int64p(1), Max: int64p(60)},
{YAMLPath: "report.enabled", Type: "bool", Label: "Report confirmed abuse", Help: "Send signed, minimized reports for confirmed-abuse IPs. Targets are configured in YAML.", FieldGroup: FieldGroupAbuseReporting},
{YAMLPath: "report.classes", Type: "[]string", Label: "Reported classes", Help: "bruteforce, php_relay, credential_stuffing, bad_asn_egress", FieldGroup: FieldGroupAbuseReporting},
{YAMLPath: "report.spool_max", Type: "int", Label: "Report spool size", Min: int64p(0), FieldGroup: FieldGroupAbuseReporting},
{YAMLPath: "central.enabled", Type: "bool", Label: "Consume central abuse database", FieldGroup: FieldGroupCentralDB},
{YAMLPath: "central.set_url", Type: "string", Label: "Scored-set URL", FieldGroup: FieldGroupCentralDB},
{YAMLPath: "central.pubkey_env", Type: "string", Label: "Central public-key env var", FieldGroup: FieldGroupCentralDB, FileOnly: true},
{YAMLPath: "central.action", Type: "string", Label: "Action on listed IPs", Help: "off | challenge | block_if_local_corroborated", FieldGroup: FieldGroupCentralDB},
{YAMLPath: "central.block_threshold", Type: "int", Label: "Block score threshold", Min: int64p(0), Max: int64p(100), FieldGroup: FieldGroupCentralDB},
{YAMLPath: "central.refresh_interval", Type: "string", Label: "Refresh interval", Help: "e.g. 6h", FieldGroup: FieldGroupCentralDB},
},
},
{
ID: "email_protection",
Title: "Email Protection",
YAMLPath: "email_protection",
Icon: "mail-shield",
Group: SectionGroupDetection,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "password_check_interval_min", Type: "int", Label: "Password check interval (min)", Min: int64p(1)},
{YAMLPath: "high_volume_senders", Type: "[]string", Label: "High-volume senders"},
{YAMLPath: "rate_warn_threshold", Type: "int", Label: "Rate warn threshold", Min: int64p(1)},
{YAMLPath: "rate_crit_threshold", Type: "int", Label: "Rate critical threshold", Min: int64p(1)},
{YAMLPath: "rate_window_min", Type: "int", Label: "Rate window (min)", Min: int64p(1)},
{YAMLPath: "known_forwarders", Type: "[]string", Label: "Known forwarders"},
},
},
{
ID: "challenge",
Title: "Challenge",
YAMLPath: "challenge",
Icon: "user-question",
Group: SectionGroupDetection,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Challenge pages enabled"},
{YAMLPath: "listen_port", Type: "int", Label: "Listen port", Min: int64p(1), Max: int64p(65535)},
{YAMLPath: "difficulty", Type: "int", Label: "PoW difficulty (0-5)", Min: int64p(0), Max: int64p(5)},
{YAMLPath: "trusted_proxies", Type: "[]string", Label: "Trusted proxy IPs"},
// challenge.secret is auto-generated at daemon startup; intentionally
// omitted so the UI cannot overwrite or leak the HMAC key.
},
},
{
ID: "php_shield",
Title: "PHP Shield",
YAMLPath: "php_shield",
Icon: "brand-php",
Group: SectionGroupDetection,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "PHP Shield enabled"},
},
},
{
ID: "signatures",
Title: "Signatures",
YAMLPath: "signatures",
Icon: "scan",
Group: SectionGroupDetection,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "auto_update", Type: "bool", Label: "Auto-update rules"},
{YAMLPath: "update_interval", Type: "string", Label: "Update interval", Placeholder: "24h"},
{YAMLPath: "yara_forge.enabled", Type: "bool", Label: "YARA-Forge enabled"},
{YAMLPath: "yara_forge.tier", Type: "enum", Label: "YARA-Forge tier", Options: []string{"core", "extended", "full"}},
{YAMLPath: "yara_forge.update_interval", Type: "string", Label: "YARA-Forge interval", Placeholder: "168h"},
{YAMLPath: "yara_forge.download_url", Type: "string", Label: "YARA-Forge signed ZIP URL"},
{YAMLPath: "disabled_rules", Type: "[]string", Label: "Disabled rule names"},
{YAMLPath: "yara_worker_enabled", Type: "bool", Label: "Run YARA-X in supervised worker"},
},
},
{
ID: "email_av",
Title: "Email AV",
YAMLPath: "email_av",
Icon: "virus",
Group: SectionGroupDetection,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Email AV enabled"},
{YAMLPath: "clamd_socket", Type: "string", Label: "clamd socket", FileOnly: true},
{YAMLPath: "scan_timeout", Type: "string", Label: "Scan timeout", Placeholder: "30s"},
{YAMLPath: "max_attachment_size", Type: "int", Label: "Max attachment bytes", Min: int64p(1024)},
{YAMLPath: "max_archive_depth", Type: "int", Label: "Max archive depth", Min: int64p(0)},
{YAMLPath: "max_archive_files", Type: "int", Label: "Max archive files", Min: int64p(1)},
{YAMLPath: "max_extraction_size", Type: "int", Label: "Max extraction bytes", Min: int64p(1024)},
{YAMLPath: "quarantine_infected", Type: "bool", Label: "Quarantine infected"},
{YAMLPath: "scan_concurrency", Type: "int", Label: "Scan concurrency", Min: int64p(1), Max: int64p(64)},
},
},
{
ID: "modsec",
Title: "ModSecurity",
YAMLPath: "modsec",
Icon: "shield-lock",
Group: SectionGroupDetection,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "rules_file", Type: "string", Label: "Rules file path", FileOnly: true},
{YAMLPath: "overrides_file", Type: "string", Label: "Overrides file path", FileOnly: true},
{YAMLPath: "reload_command", Type: "string", Label: "Reload command", FileOnly: true},
},
},
{
ID: "performance",
Title: "Performance",
YAMLPath: "performance",
Icon: "activity",
Group: SectionGroupOps,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Performance checks", Nullable: true, Help: "Leave unset to inherit default (on)"},
{YAMLPath: "load_high_multiplier", Type: "float", Label: "Load high multiplier"},
{YAMLPath: "load_critical_multiplier", Type: "float", Label: "Load critical multiplier"},
{YAMLPath: "php_process_warn_per_user", Type: "int", Label: "PHP process warn per user", Min: int64p(1)},
{YAMLPath: "php_process_critical_total_multiplier", Type: "int", Label: "PHP process crit multiplier", Min: int64p(1)},
{YAMLPath: "error_log_warn_size_mb", Type: "int", Label: "Error log warn size (MB)", Min: int64p(1)},
{YAMLPath: "mysql_join_buffer_max_mb", Type: "int", Label: "MySQL join buffer max (MB)", Min: int64p(1)},
{YAMLPath: "mysql_wait_timeout_max", Type: "int", Label: "MySQL wait timeout max (s)", Min: int64p(1)},
{YAMLPath: "mysql_max_connections_per_user", Type: "int", Label: "MySQL max connections per user", Min: int64p(1)},
{YAMLPath: "redis_bgsave_min_interval", Type: "int", Label: "Redis bgsave min interval (s)", Min: int64p(1)},
{YAMLPath: "redis_large_dataset_gb", Type: "int", Label: "Redis large dataset (GB)", Min: int64p(1)},
{YAMLPath: "wp_memory_limit_max_mb", Type: "int", Label: "WP memory limit max (MB)", Min: int64p(32)},
{YAMLPath: "wp_transient_warn_mb", Type: "int", Label: "WP transient warn (MB)", Min: int64p(1)},
{YAMLPath: "wp_transient_critical_mb", Type: "int", Label: "WP transient critical (MB)", Min: int64p(1)},
{YAMLPath: "wp_cron_fix.interval_minutes", Type: "int", Label: "WP-Cron fix: system cron interval (min)", Min: int64p(1), Max: int64p(60), Help: "How often the installed system cron runs wp-cron.php. Only bounds task latency; WordPress keeps its own event schedule. Default 15."},
{YAMLPath: "wp_cron_fix.php_bin", Type: "string", Label: "WP-Cron fix: PHP binary", Placeholder: "/usr/local/bin/php", Help: "Overrides the cron interpreter for every site. Leave empty to use the unambiguous cPanel MultiPHP version, then auto-detect.", FileOnly: true},
},
},
{
ID: "cloudflare",
Title: "Cloudflare",
YAMLPath: "cloudflare",
Icon: "cloud",
Group: SectionGroupIntegrations,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Cloudflare integration"},
{YAMLPath: "refresh_hours", Type: "int", Label: "Refresh interval (hours)", Min: int64p(1), Max: int64p(168)},
},
},
{
ID: "geoip",
Title: "GeoIP",
YAMLPath: "geoip",
Icon: "world",
Group: SectionGroupIntegrations,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "account_id", Type: "string", Label: "MaxMind account ID"},
{YAMLPath: "license_key", Type: "string", Label: "MaxMind license key", Secret: true},
{YAMLPath: "editions", Type: "[]enum", Label: "Database editions", OptionsSource: "geoip_editions", Help: "Which MaxMind databases to download. GeoLite2-* are free; GeoIP2-* require a paid subscription."},
{YAMLPath: "auto_update", Type: "bool", Label: "Auto-update databases", Nullable: true},
{YAMLPath: "update_interval", Type: "string", Label: "Update interval", Placeholder: "24h"},
},
},
{
ID: "infra_ips",
Title: "Infra IPs",
YAMLPath: "infra_ips",
Icon: "server",
Group: SectionGroupOps,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "", Type: "[]string", Label: "Trusted infra IPs and CIDRs"},
},
},
{
ID: "disabled_checks",
Title: "Disabled checks",
YAMLPath: "disabled_checks",
Icon: "x-circle",
Group: SectionGroupOps,
Restart: false,
Fields: []SettingsField{
{YAMLPath: "", Type: "[]enum", Label: "Skip scheduled check runners", OptionsSource: "disabled_check_names", Help: "Selecting a finding name skips the scheduled check runner or runners that emit it, including sibling findings from the same runner. Realtime findings are not affected. For email-only suppression, use Alerts > Disabled check names instead."},
},
},
{
ID: "sentry",
Title: "Sentry",
YAMLPath: "sentry",
Icon: "bug",
Group: SectionGroupIntegrations,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Sentry enabled"},
{YAMLPath: "dsn", Type: "string", Label: "Sentry DSN", Secret: true},
{YAMLPath: "environment", Type: "string", Label: "Environment", Placeholder: "production"},
{YAMLPath: "sample_rate", Type: "float", Label: "Sample rate (0 to 1.0)"},
{YAMLPath: "debug", Type: "bool", Label: "Debug logging"},
},
},
{
ID: "firewall",
Title: "Firewall",
YAMLPath: "firewall",
Icon: "shield-lock",
Group: SectionGroupFirewall,
Restart: true,
Fields: []SettingsField{
{YAMLPath: "enabled", Type: "bool", Label: "Firewall enabled", Help: "Activates the nftables-based firewall on next daemon restart. Verify port lists below before enabling to avoid lockout.", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "ipv6", Type: "bool", Label: "IPv6 dual-stack", FieldGroup: FieldGroupIPv6},
{YAMLPath: "tcp_in", Type: "[]int", Label: "Inbound TCP ports", Help: "SSH (22) is intentionally not in the default; add it explicitly if sshd listens on 22. WebUI port must be present or you will lose remote access on restart.", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "tcp_out", Type: "[]int", Label: "Outbound TCP ports", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "udp_in", Type: "[]int", Label: "Inbound UDP ports", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "udp_out", Type: "[]int", Label: "Outbound UDP ports", Help: "Includes 6277/24441 by default for SpamAssassin DCC/Pyzor; do not remove unless rspamd-only.", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "tcp6_in", Type: "[]int", Label: "Inbound TCP6 ports", Help: "Empty inherits tcp_in. Only set when IPv6 should differ from IPv4.", FieldGroup: FieldGroupIPv6},
{YAMLPath: "tcp6_out", Type: "[]int", Label: "Outbound TCP6 ports", Help: "Empty inherits tcp_out.", FieldGroup: FieldGroupIPv6},
{YAMLPath: "udp6_in", Type: "[]int", Label: "Inbound UDP6 ports", Help: "Empty inherits udp_in.", FieldGroup: FieldGroupIPv6},
{YAMLPath: "udp6_out", Type: "[]int", Label: "Outbound UDP6 ports", Help: "Empty inherits udp_out.", FieldGroup: FieldGroupIPv6},
{YAMLPath: "restricted_tcp", Type: "[]int", Label: "Restricted TCP (infra-only)", Help: "Reachable only from infra_ips. Manage infra_ips in its own section.", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "required_tcp_out", Type: "[]int", Label: "Required outbound TCP ports", Help: "Checked only, never added to the policy: validation warns when an effective outbound family omits one. Integrations declare the ports they dial in their own conf.d fragment.", FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "passive_ftp_start", Type: "int", Label: "Passive FTP range start", Min: int64p(1024), Max: int64p(65535), FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "passive_ftp_end", Type: "int", Label: "Passive FTP range end", Min: int64p(1024), Max: int64p(65535), FieldGroup: FieldGroupAccessPorts},
{YAMLPath: "drop_nolog", Type: "[]int", Label: "Silent-drop ports", Help: "Dropped without logging to keep scanner noise out of the log.", FieldGroup: FieldGroupLogging},
{YAMLPath: "conn_rate_limit", Type: "int", Label: "Conn rate limit (per IP/min)", Min: int64p(0), Max: int64p(100000), Help: "0 disables. 200 tolerates shared CGNAT egress. IPv6 metered per /64.", FieldGroup: FieldGroupRateLimits},
{YAMLPath: "conn_limit", Type: "int", Label: "Concurrent connections per IPv4", Min: int64p(0), Max: int64p(100000), Help: "0 disables. IPv4 only (no IPv6).", FieldGroup: FieldGroupRateLimits},
{YAMLPath: "syn_flood_protection", Type: "bool", Label: "SYN flood protection", Help: "Dual-stack; IPv6 metered per /64.", FieldGroup: FieldGroupFloodProtection},
{YAMLPath: "udp_flood", Type: "bool", Label: "UDP flood protection", Help: "Dual-stack; IPv6 metered per /64.", FieldGroup: FieldGroupFloodProtection},
{YAMLPath: "udp_flood_rate", Type: "int", Label: "UDP packets/sec", Min: int64p(1), Max: int64p(100000), Help: "Per source; IPv6 metered per /64.", FieldGroup: FieldGroupFloodProtection},
{YAMLPath: "udp_flood_burst", Type: "int", Label: "UDP burst allowance", Min: int64p(1), Max: int64p(1000000), Help: "Per source; IPv6 metered per /64.", FieldGroup: FieldGroupFloodProtection},
{YAMLPath: "deny_ip_limit", Type: "int", Label: "Permanent block cap", Min: int64p(0), Max: int64p(1000000), Help: "0 = unlimited.", FieldGroup: FieldGroupLimits},
{YAMLPath: "deny_temp_ip_limit", Type: "int", Label: "Temporary block cap", Min: int64p(0), Max: int64p(1000000), FieldGroup: FieldGroupLimits},
{YAMLPath: "country_block", Type: "[]string", Label: "Country block (ISO-3166)", Help: "Two-letter codes, one per line.", FieldGroup: FieldGroupGeoDynDNS},
{YAMLPath: "country_db_path", Type: "string", Label: "Country DB path override", Placeholder: "(default: <state_path>/geoip)", FieldGroup: FieldGroupGeoDynDNS, FileOnly: true},
{YAMLPath: "dyndns_hosts", Type: "[]string", Label: "DynDNS hosts", Help: "Resolved every 5 minutes and merged into the trusted set.", FieldGroup: FieldGroupGeoDynDNS},
{YAMLPath: "smtp_block", Type: "bool", Label: "Block outbound SMTP", Help: "When enabled, only smtp_allow_users may originate outbound mail. Verify allow list first.", FieldGroup: FieldGroupSMTPControls},
{YAMLPath: "smtp_allow_users", Type: "[]string", Label: "SMTP allow users", Help: "root is always allowed.", FieldGroup: FieldGroupSMTPControls},
{YAMLPath: "smtp_ports", Type: "[]int", Label: "SMTP ports", FieldGroup: FieldGroupSMTPControls},
{YAMLPath: "log_dropped", Type: "bool", Label: "Log dropped packets", FieldGroup: FieldGroupLogging},
{YAMLPath: "log_rate", Type: "int", Label: "Log entries per minute", Min: int64p(0), Max: int64p(10000), FieldGroup: FieldGroupLogging},
},
},
}
// SettingsSectionIDs returns the ordered list of section IDs.
func SettingsSectionIDs() []string {
out := make([]string, 0, len(settingsSections))
for _, s := range settingsSections {
out = append(out, s.ID)
}
return out
}
// LookupSettingsSection returns the section with the given ID.
func LookupSettingsSection(id string) (SettingsSection, bool) {
for _, s := range settingsSections {
if s.ID == id {
return withReloadPolicy(s), true
}
}
return SettingsSection{}, false
}
// AllSettingsSections returns the list of sections. Intended for
// read-only consumers such as the dashboard navigation.
func AllSettingsSections() []SettingsSection {
out := make([]SettingsSection, len(settingsSections))
for i, s := range settingsSections {
out[i] = withReloadPolicy(s)
}
return out
}
func withReloadPolicy(section SettingsSection) SettingsSection {
for _, policy := range config.HotReloadManifest() {
if policy.Field == section.YAMLPath {
section.Restart = policy.RestartRequired
section.ReloadTag = policy.Tag
break
}
}
return section
}
package webui
import "fmt"
// subnetCoverer is implemented by the firewall engine: it reports the
// blocked subnet, if any, that still covers an address.
type subnetCoverer interface {
BlockedSubnetCovering(ip string) (string, bool)
}
// coveringSubnetWarning explains why an unblocked or allowed address may
// still be dropped: the chain drops @blocked_nets before it accepts
// @allowed_ips, so a covering blocked subnet overrides both actions. Empty
// when nothing covers the address or the blocker cannot tell.
func coveringSubnetWarning(blocker IPBlocker, ip string) string {
cov, ok := blocker.(subnetCoverer)
if !ok {
return ""
}
cidr, covered := cov.BlockedSubnetCovering(ip)
if !covered {
return ""
}
return fmt.Sprintf("still dropped by blocked subnet %s; unblock that subnet to restore access", cidr)
}
package webui
import (
"crypto/rand"
"encoding/hex"
"time"
)
// newSuppressionID returns a fresh opaque rule ID. Shared by the UI add path
// and the import path so an imported rule without an ID can be deleted like
// any other.
func newSuppressionID() string {
b := make([]byte, 8)
if _, err := rand.Read(b); err != nil {
return "sup-" + hex.EncodeToString([]byte(time.Now().UTC().Format("20060102150405.000000000")))
}
return hex.EncodeToString(b)
}
package webui
import (
"net/http"
"path/filepath"
"regexp"
"strings"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/state"
)
// suppressionCheckName matches the check names findings carry. A rule
// matches a finding's check exactly, so a glob or free text would be saved
// as a rule that matches nothing.
var suppressionCheckName = regexp.MustCompile(`^[A-Za-z0-9_][A-Za-z0-9_.:-]{0,127}$`)
// knownCheck reports whether name is a registered check or the check of a
// current finding. Checks from other subsystems may be missing from the
// registry, so an unknown name is a warning, not an error.
func (s *Server) knownCheck(name string) bool {
for _, known := range checks.AllCheckNames() {
if known == name {
return true
}
}
for _, f := range s.store.LatestFindings() {
if f.Check == name {
return true
}
}
return false
}
func (s *Server) apiSuppressions(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
writeAll(w, s.store.LoadSuppressions())
case http.MethodPost:
var req struct {
Check string `json:"check"`
PathPattern string `json:"path_pattern"`
// AllPaths is the explicit opt-in for a rule without a path
// pattern, which hides every finding of the check and stops its
// remediation.
AllPaths bool `json:"all_paths"`
Reason string `json:"reason"`
}
if err := decodeJSONBodyLimited(w, r, 32*1024, &req); err != nil || req.Check == "" {
writeJSONError(w, "check field is required", http.StatusBadRequest)
return
}
if !suppressionCheckName.MatchString(req.Check) {
writeJSONError(w, "check must be a check name such as webshell; patterns and spaces are not allowed", http.StatusBadRequest)
return
}
req.PathPattern = strings.TrimSpace(req.PathPattern)
switch {
case req.PathPattern == "" && !req.AllPaths:
writeJSONError(w, "path_pattern is required; set all_paths to suppress every finding of this check", http.StatusBadRequest)
return
case req.PathPattern != "" && req.AllPaths:
writeJSONError(w, "path_pattern and all_paths are mutually exclusive", http.StatusBadRequest)
return
case req.PathPattern != "":
if _, err := filepath.Match(req.PathPattern, ""); err != nil {
writeJSONError(w, "path_pattern is not a valid glob: "+err.Error(), http.StatusBadRequest)
return
}
}
id := newSuppressionID()
err := s.store.UpdateSuppressions(func(rules []state.SuppressionRule) ([]state.SuppressionRule, error) {
return append(rules, state.SuppressionRule{
ID: id,
Check: req.Check,
PathPattern: req.PathPattern,
Reason: req.Reason,
CreatedAt: time.Now(),
}), nil
})
if err != nil {
writeJSONError(w, "failed to save suppression: "+err.Error(), http.StatusInternalServerError)
return
}
scope := "pattern: " + req.PathPattern
if req.AllPaths {
scope = "all paths"
}
s.auditLog(r, "suppress", req.Check, scope)
resp := map[string]interface{}{"id": id}
if !s.knownCheck(req.Check) {
resp["warning"] = "No known check is named " + req.Check + "; the rule matches nothing until a finding with that check appears."
}
writeOK(w, resp)
case http.MethodDelete:
var req struct {
ID string `json:"id"`
}
if err := decodeJSONBodyLimited(w, r, 16*1024, &req); err != nil || req.ID == "" {
writeJSONError(w, "id is required", http.StatusBadRequest)
return
}
found := false
err := s.store.UpdateSuppressions(func(rules []state.SuppressionRule) ([]state.SuppressionRule, error) {
var filtered []state.SuppressionRule
for _, rule := range rules {
if rule.ID != req.ID {
filtered = append(filtered, rule)
} else {
found = true
}
}
return filtered, nil
})
if err != nil {
writeJSONError(w, "failed to save suppressions: "+err.Error(), http.StatusInternalServerError)
return
}
if !found {
writeJSONError(w, "Suppression rule not found", http.StatusNotFound)
return
}
s.auditLog(r, "unsuppress", req.ID, "removed suppression rule")
writeOK(w, map[string]interface{}{"id": req.ID})
default:
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
}
}
package webui
import (
"errors"
"fmt"
"net"
"net/http"
"strings"
"time"
"github.com/pidginhost/csm/internal/alert"
"github.com/pidginhost/csm/internal/attackdb"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/threat"
)
func (s *Server) handleThreat(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "threat.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
// GET /api/v1/threat/stats
func (s *Server) apiThreatStats(w http.ResponseWriter, r *http.Request) {
adb := attackdb.Global()
if adb == nil {
writeJSONError(w, "attack database not initialized", http.StatusServiceUnavailable)
return
}
writeJSON(w, adb.Stats())
}
// GET /api/v1/threat/top-attackers?limit=25
func (s *Server) apiThreatTopAttackers(w http.ResponseWriter, r *http.Request) {
limit := queryInt(r, "limit", 25)
if limit > 200 {
limit = 200
}
adb := attackdb.Global()
if adb == nil {
writeItems(w, []struct{}{}, map[string]interface{}{"offset": 0, "limit": limit, "truncated": false})
return
}
// One record past the limit tells whether the list was cut.
recs := adb.TopAttackers(limit + 1)
truncated := len(recs) > limit
if truncated {
recs = recs[:limit]
}
// Enrich with unified intelligence
ips := make([]string, len(recs))
for i, rec := range recs {
ips[i] = rec.IP
}
intels := threat.LookupBatch(ips, s.cfg.StatePath)
type enriched struct {
*attackdb.IPRecord
UnifiedScore int `json:"unified_score"`
Verdict string `json:"verdict"`
AbuseScore int `json:"abuse_score"`
InThreatDB bool `json:"in_threat_db"`
Blocked bool `json:"currently_blocked"`
Country string `json:"country,omitempty"`
ASOrg string `json:"as_org,omitempty"`
}
results := make([]enriched, len(recs))
for i, rec := range recs {
results[i] = enriched{
IPRecord: rec,
UnifiedScore: intels[i].UnifiedScore,
Verdict: intels[i].Verdict,
AbuseScore: intels[i].AbuseScore,
InThreatDB: intels[i].InThreatDB,
Blocked: intels[i].CurrentlyBlocked,
}
if gdb := s.geoIPDB.Load(); gdb != nil {
geo := gdb.Lookup(rec.IP)
results[i].Country = geo.Country
results[i].ASOrg = geo.ASOrg
}
}
writeItems(w, results, map[string]interface{}{"offset": 0, "limit": limit, "truncated": truncated})
}
// GET /api/v1/threat/ip?ip=1.2.3.4
func (s *Server) apiThreatIP(w http.ResponseWriter, r *http.Request) {
ip := r.URL.Query().Get("ip")
if ip == "" || net.ParseIP(ip) == nil {
writeJSONError(w, "invalid or missing ip parameter", http.StatusBadRequest)
return
}
intel := threat.Lookup(ip, s.cfg.StatePath)
// Enrich with GeoIP data if available
if gdb := s.geoIPDB.Load(); gdb != nil {
geo := gdb.Lookup(ip)
intel.Country = geo.Country
intel.CountryName = geo.CountryName
intel.City = geo.City
intel.ASN = geo.ASN
intel.ASOrg = geo.ASOrg
intel.Network = geo.Network
}
writeJSON(w, intel)
}
// threatEventView is an attack event as the API sends it: sev is a label.
type threatEventView struct {
attackdb.Event
Severity string `json:"sev"`
}
// GET /api/v1/threat/events?ip=1.2.3.4&limit=50
func (s *Server) apiThreatEvents(w http.ResponseWriter, r *http.Request) {
ip := r.URL.Query().Get("ip")
if ip == "" || net.ParseIP(ip) == nil {
writeJSONError(w, "invalid or missing ip parameter", http.StatusBadRequest)
return
}
limit := queryInt(r, "limit", 50)
if limit > 500 {
limit = 500
}
adb := attackdb.Global()
if adb == nil {
writeItems(w, []threatEventView{}, map[string]interface{}{"offset": 0, "limit": limit, "truncated": false})
return
}
// One event past the limit tells whether older events were left out.
events := adb.QueryEvents(ip, limit+1)
truncated := len(events) > limit
if truncated {
events = events[:limit]
}
views := make([]threatEventView, len(events))
for i, ev := range events {
views[i] = threatEventView{Event: ev, Severity: alert.Severity(ev.Severity).String()}
}
writeItems(w, views, map[string]interface{}{"offset": 0, "limit": limit, "truncated": truncated})
}
// GET /api/v1/threat/db-stats
func (s *Server) apiThreatDBStats(w http.ResponseWriter, r *http.Request) {
result := make(map[string]interface{})
if tdb := checks.GetThreatDB(); tdb != nil {
result["threat_db"] = tdb.Stats()
}
if adb := attackdb.Global(); adb != nil {
result["attack_db"] = map[string]interface{}{
"total_ips": adb.TotalIPs(),
"top_line": adb.FormatTopLine(),
}
}
writeJSON(w, result)
}
// POST /api/v1/threat/whitelist-ip - mark an IP as a known customer
// Unblocks, removes from threat DB + attack DB, adds to whitelist.
func (s *Server) apiThreatWhitelistIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != "POST" {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Every downstream key (firewall, threat DB, attack DB, audit) uses the
// canonical form; a padded or upper-case spelling validated but then
// matched nothing and still answered 200.
req.IP = parsedIP.String()
actions, releaseErr := s.releaseIP(req.IP, ipRelease{allow: releaseAllow, reason: "CSM whitelist: customer IP"})
warning := ""
if s.blocker != nil {
warning = coveringSubnetWarning(s.blocker, req.IP)
}
detail := "permanent whitelist"
if warning != "" {
detail += "; " + warning
}
s.auditLog(r, "whitelist_ip", req.IP, detail)
if releaseErr != nil {
writeJSONError(w, "IP action incomplete: "+releaseErr.Error(), http.StatusInternalServerError)
return
}
resp := map[string]interface{}{
"ip": req.IP,
"actions": actions,
}
if warning != "" {
resp["warning"] = warning
}
writeOK(w, resp)
}
// GET /api/v1/threat/whitelist - list all whitelisted IPs
func (s *Server) apiThreatWhitelist(w http.ResponseWriter, r *http.Request) {
tdb := checks.GetThreatDB()
if tdb == nil {
writeAll(w, []checks.WhitelistIP{})
return
}
writeAll(w, tdb.WhitelistedIPs())
}
// POST /api/v1/threat/unwhitelist-ip - remove an IP from the whitelist
func (s *Server) apiThreatUnwhitelistIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != "POST" {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Every downstream key (firewall, threat DB, attack DB, audit) uses the
// canonical form; a padded or upper-case spelling validated but then
// matched nothing and still answered 200.
req.IP = parsedIP.String()
if tdb := checks.GetThreatDB(); tdb != nil {
if tdb.IsConfigWhitelisted(req.IP) {
writeJSONError(w, "IP is managed by reputation.whitelist; edit the config and reload", http.StatusConflict)
return
}
tdb.RemoveWhitelist(req.IP)
}
// Also remove from firewall allow list
if s.blocker != nil {
if remover, ok := s.blocker.(allowRemover); ok {
if err := remover.RemoveAllowIP(req.IP); err != nil {
s.auditLog(r, "unwhitelist_ip_failed", req.IP, err.Error())
writeJSONError(w, "Remove firewall allow rule: "+err.Error(), http.StatusInternalServerError)
return
}
}
}
s.auditLog(r, "unwhitelist_ip", req.IP, "removed from runtime whitelist and firewall allow list")
writeOK(w, map[string]interface{}{"ip": req.IP})
}
// manualBlockTTL is the lifetime of the Web UI "Block (24h)" action. The
// threat evidence it records carries the same expiry, so the address stops
// counting as malicious when the firewall block lapses.
const manualBlockTTL = 24 * time.Hour
const (
manualBlockReason = "Manually blocked via CSM Web UI"
manualPermBlockReason = "Permanently blocked via CSM Web UI"
bulkBlockReason = "Bulk blocked via CSM Web UI"
bulkPermanentBlockReason = "Bulk permanently blocked via CSM Web UI"
)
// POST /api/v1/threat/block-ip - manually block an IP for 24 hours.
func (s *Server) apiThreatBlockIP(w http.ResponseWriter, r *http.Request) {
s.operatorBlockIP(w, r, false)
}
// POST /api/v1/threat/block-ip-permanent - block an IP with no expiry.
// Permanence is chosen by the authenticated operator action alone: no
// request field can turn the 24h block into a permanent one.
func (s *Server) apiThreatBlockIPPermanent(w http.ResponseWriter, r *http.Request) {
s.operatorBlockIP(w, r, true)
}
func (s *Server) operatorBlockIP(w http.ResponseWriter, r *http.Request, permanent bool) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != "POST" {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Every downstream key (firewall, threat DB, attack DB, audit) uses the
// canonical form; a padded or upper-case spelling validated but then
// matched nothing and still answered 200.
req.IP = parsedIP.String()
reason := manualBlockReason
ttl := manualBlockTTL
if permanent {
reason = manualPermBlockReason
ttl = 0
}
var actions []string
// 1. Block in firewall. A zero timeout is a permanent firewall block.
if s.blocker == nil {
writeJSONError(w, "firewall engine not available", http.StatusServiceUnavailable)
return
}
// Operator-initiated: bypass auto_response.dry_run gate.
if err := s.blockIPPreservingLifetime(req.IP, reason, ttl); err != nil {
status := http.StatusInternalServerError
if errors.Is(err, firewall.ErrPermanentBlock) || errors.Is(err, firewall.ErrLongerBlock) {
status = http.StatusConflict
}
writeJSONError(w, fmt.Sprintf("block failed: %v", err), status)
return
}
invalidateIPUndo(req.IP)
if permanent {
actions = append(actions, "blocked in firewall permanently")
} else {
actions = append(actions, "blocked in firewall for 24h")
}
// 2. Record threat evidence with the same lifetime as the block.
if tdb := checks.GetThreatDB(); tdb != nil {
if permanent {
tdb.AddPermanent(req.IP, reason)
actions = append(actions, "added to threat DB permanently")
} else {
tdb.AddOperatorTemporary(req.IP, reason, ttl)
actions = append(actions, "added to threat DB for 24h")
}
}
// 3. Record in attack DB
if adb := attackdb.Global(); adb != nil {
adb.MarkBlocked(req.IP)
actions = append(actions, "recorded in attack DB")
}
if permanent {
s.auditLog(r, "block_ip_permanent", req.IP, "manual permanent block")
} else {
s.auditLog(r, "block_ip", req.IP, "manual block 24h")
}
writeOK(w, map[string]interface{}{
"ip": req.IP,
"permanent": permanent,
"actions": actions,
})
}
// POST /api/v1/threat/clear-ip - unblock + clear from all DBs without whitelisting.
// For dynamic IP customers: one-time cleanup, IP can be re-blocked later.
// ipRelease says how releaseIP lets an address go. The zero value unblocks
// and forgets it without allowing it (Unblock & Clear).
type ipRelease struct {
allow releaseAllowMode
ttl time.Duration // for releaseTempAllow
reason string // firewall allow rule comment
}
type releaseAllowMode int
const (
releaseNoAllow releaseAllowMode = iota
releaseAllow
releaseTempAllow
)
// releaseIP is the shared tail of whitelist, temporary whitelist, clear and
// bulk whitelist: unblock the address in the firewall, allow it as rel says,
// drop it from the threat database's permanent list (whitelisting it there
// too when allowed), and forget it in the attack database, all inside the
// netblock-history lock so the history cleanup is one operator action with
// the firewall change. cPHulk's login history is flushed after. It returns
// the steps that took effect and any failed required steps. Other changes
// may have happened on error, so callers still audit.
func (s *Server) releaseIP(ip string, rel ipRelease) ([]string, error) {
var actions []string
var failures []error
hours := int(rel.ttl / time.Hour)
err := checks.ForgetNetblockHistory(s.cfg.StatePath, ip, func() {
if s.blocker != nil {
if err := s.blocker.UnblockIP(ip); err == nil {
actions = append(actions, "unblocked from firewall")
} else {
failures = append(failures, fmt.Errorf("unblock firewall: %w", err))
}
switch rel.allow {
case releaseAllow:
if allower, ok := s.blocker.(ipAllower); ok {
if err := allower.AllowIP(ip, rel.reason); err == nil {
actions = append(actions, "added to firewall allow list")
} else {
failures = append(failures, fmt.Errorf("allow in firewall: %w", err))
}
}
case releaseTempAllow:
if allower, ok := s.blocker.(ipTempAllower); ok {
if err := allower.TempAllowIP(ip, rel.reason, rel.ttl); err == nil {
actions = append(actions, fmt.Sprintf("temp allowed in firewall for %dh", hours))
} else {
failures = append(failures, fmt.Errorf("temporarily allow in firewall: %w", err))
}
}
}
}
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.RemovePermanent(ip)
switch rel.allow {
case releaseAllow:
tdb.AddWhitelist(ip)
actions = append(actions, "removed from threat DB, added to whitelist")
case releaseTempAllow:
tdb.TempWhitelist(ip, rel.ttl)
actions = append(actions, fmt.Sprintf("temp whitelisted for %dh", hours))
default:
actions = append(actions, "removed from threat DB")
}
}
if adb := attackdb.Global(); adb != nil {
adb.RemoveIP(ip)
actions = append(actions, "removed from attack DB")
}
})
if err == nil {
actions = append(actions, "removed from subnet block history")
} else {
failures = append(failures, fmt.Errorf("subnet history cleanup failed: %w", err))
}
if flushCphulk(ip) == nil {
actions = append(actions, "flushed cPanel login history")
}
return actions, errors.Join(failures...)
}
func (s *Server) apiThreatClearIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != "POST" {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Every downstream key (firewall, threat DB, attack DB, audit) uses the
// canonical form; a padded or upper-case spelling validated but then
// matched nothing and still answered 200.
req.IP = parsedIP.String()
actions, releaseErr := s.releaseIP(req.IP, ipRelease{})
s.auditLog(r, "clear_ip", req.IP, "unblock & clear")
if releaseErr != nil {
writeJSONError(w, "IP action incomplete: "+releaseErr.Error(), http.StatusInternalServerError)
return
}
writeOK(w, map[string]interface{}{
"ip": req.IP,
"actions": actions,
})
}
// POST /api/v1/threat/temp-whitelist-ip - whitelist for a specified duration.
func (s *Server) apiThreatTempWhitelistIP(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != "POST" {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IP string `json:"ip"`
Hours int `json:"hours"` // default 24
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil || req.IP == "" {
writeJSONError(w, "IP is required", http.StatusBadRequest)
return
}
parsedIP, err := parseAndValidateIP(req.IP)
if err != nil {
writeJSONError(w, err.Error(), http.StatusBadRequest)
return
}
// Every downstream key (firewall, threat DB, attack DB, audit) uses the
// canonical form; a padded or upper-case spelling validated but then
// matched nothing and still answered 200.
req.IP = parsedIP.String()
if req.Hours <= 0 {
req.Hours = 24
}
if req.Hours > 168 { // max 7 days
req.Hours = 168
}
ttl := time.Duration(req.Hours) * time.Hour
actions, releaseErr := s.releaseIP(req.IP, ipRelease{allow: releaseTempAllow, ttl: ttl, reason: "CSM temp whitelist"})
s.auditLog(r, "temp_whitelist_ip", req.IP, fmt.Sprintf("%dh temp whitelist", req.Hours))
if releaseErr != nil {
writeJSONError(w, "IP action incomplete: "+releaseErr.Error(), http.StatusInternalServerError)
return
}
writeOK(w, map[string]interface{}{
"ip": req.IP,
"duration_seconds": ttl.Seconds(),
"actions": actions,
})
}
// threatBulkActionMax bounds the addresses one bulk threat action changes. Each
// request returns one undo token, so the UI refuses larger selections instead
// of splitting them.
const threatBulkActionMax = 100
// POST /api/v1/threat/bulk-action - block or whitelist multiple IPs at once.
func (s *Server) apiThreatBulkAction(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
var req struct {
IPs []string `json:"ips"`
Action string `json:"action"`
}
if err := decodeJSONBodyLimited(w, r, 64*1024, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
if len(req.IPs) == 0 || len(req.IPs) > threatBulkActionMax {
writeJSONError(w, fmt.Sprintf("IPs must be 1-%d items", threatBulkActionMax), http.StatusBadRequest)
return
}
blockAction := req.Action == "block" || req.Action == "block_permanent"
if !blockAction && req.Action != "whitelist" {
writeJSONError(w, "Action must be 'block', 'block_permanent' or 'whitelist'", http.StatusBadRequest)
return
}
if blockAction && s.blocker == nil {
writeJSONError(w, "firewall engine not available", http.StatusServiceUnavailable)
return
}
// Permanence follows the operator action, never a per-IP request field.
permanent := req.Action == "block_permanent"
blockReason := bulkBlockReason
blockTTL := manualBlockTTL
if permanent {
blockReason = bulkPermanentBlockReason
blockTTL = 0
}
priorBlocks := make(map[string]firewall.BlockedEntry)
expectedBlocks := make(map[string]firewall.BlockedEntry)
seen := make(map[string]bool, len(req.IPs))
count := 0
succeeded := make([]string, 0, len(req.IPs))
var removedThreats []undoThreatRow
var warnings []string
for _, ipStr := range req.IPs {
parsedIP, err := parseAndValidateIP(ipStr)
if err != nil {
warnings = append(warnings, ipStr+": "+err.Error())
continue
}
ipStr = parsedIP.String()
if seen[ipStr] {
continue
}
seen[ipStr] = true
switch {
case blockAction:
// Mirror the single-IP block flow.
// Operator-initiated bulk block: bypass auto_response.dry_run gate.
before, after, err := s.blockIPForUndo(ipStr, blockReason, blockTTL)
if err != nil {
warnings = append(warnings, ipStr+": "+err.Error())
continue
}
if before != nil {
priorBlocks[ipStr] = *before
}
if after != nil {
expectedBlocks[ipStr] = *after
}
invalidateIPUndo(ipStr)
// Capture whatever evidence is already on file so the undo
// restores it instead of dropping an older permanent row this
// block did not create.
if row, ok := captureUndoThreatRow(ipStr, false); ok {
removedThreats = append(removedThreats, row)
}
if tdb := checks.GetThreatDB(); tdb != nil {
if permanent {
tdb.AddPermanent(ipStr, blockReason)
} else {
tdb.AddOperatorTemporary(ipStr, blockReason, blockTTL)
}
}
if adb := attackdb.Global(); adb != nil {
adb.MarkBlocked(ipStr)
}
succeeded = append(succeeded, ipStr)
count++
default:
// Capture the live threat row before dropping it so an undo can
// restore it exactly, preserving source/expiry, instead of
// leaving a whitelisted attacker with no threat record.
row, hadThreat := captureUndoThreatRow(ipStr, false)
if _, err := s.releaseIP(ipStr, ipRelease{allow: releaseAllow, reason: "CSM bulk whitelist"}); err != nil {
warnings = append(warnings, ipStr+": "+err.Error())
continue
}
if hadThreat {
removedThreats = append(removedThreats, row)
}
if s.blocker != nil {
if warning := coveringSubnetWarning(s.blocker, ipStr); warning != "" {
warnings = append(warnings, ipStr+": "+warning)
}
}
succeeded = append(succeeded, ipStr)
count++
}
}
auditDetail := ""
if blockAction {
auditDetail = "24h block"
if permanent {
auditDetail = "permanent block"
}
auditDetail += ": "
}
auditDetail += strings.Join(succeeded, ", ")
s.auditLog(r, "threat_bulk_"+req.Action, fmt.Sprintf("%d IPs", count), auditDetail)
var undoToken string
if count > 0 {
inverse := undoInverseThreatBlock
summary := fmt.Sprintf("Blocked %d IPs", count)
action := "threat_bulk_block"
switch {
case permanent:
summary = fmt.Sprintf("Permanently blocked %d IPs", count)
action = "threat_bulk_block_permanent"
case !blockAction:
inverse = undoInverseThreatWhitelist
summary = fmt.Sprintf("Whitelisted %d IPs", count)
action = "threat_bulk_whitelist"
}
undoToken = s.recordUndoEntry(r, action, inverse, summary,
undoPayloadIPs{IPs: succeeded, RestoreThreats: removedThreats, BlockSnapshot: blockAction, RestoreBlocks: priorBlocks, ExpectedBlocks: expectedBlocks})
}
if warnings == nil {
warnings = []string{}
}
fields := map[string]interface{}{
"count": count,
"permanent": permanent,
"warnings": warnings,
}
if count == 0 {
// No address completed: an error, with each address's reason.
msg := "No address action completed"
if len(warnings) > 0 {
msg += ": " + strings.Join(warnings, "; ")
}
fields["error"] = msg
writeJSONStatus(w, http.StatusUnprocessableEntity, fields)
return
}
if undoToken != "" {
fields["undo_token"] = undoToken
}
writeOK(w, fields)
}
// writeJSON is defined in api.go
package webui
import (
"bytes"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"fmt"
"math/big"
"net"
"os"
"path/filepath"
"sync"
"time"
"github.com/pidginhost/csm/internal/integrity"
)
// selfSignedValidity is how long a generated certificate is valid; a var so
// tests can generate one close to expiry.
var selfSignedValidity = 365 * 24 * time.Hour
// Inject write failures without relying on permissions that root bypasses.
var writeTLSFile = writeFileReplace
// renewBefore is how close to expiry a generated certificate is replaced.
const renewBefore = 30 * 24 * time.Hour
// selfSignedOrganization marks the certificates EnsureTLSCert generates.
const selfSignedOrganization = "CSM Security Monitor"
// EnsureTLSCert generates a self-signed ECDSA P-256 certificate if the cert
// or key file doesn't exist, and renews one it generated earlier when it
// expires within renewBefore. Any other certificate, such as one the
// operator installed, is left alone. Includes localhost and the server
// hostname in the certificate SANs.
func EnsureTLSCert(certPath, keyPath string, extraNames ...string) error {
if fileExists(certPath) && fileExists(keyPath) {
return renewTLSCert(certPath, keyPath)
}
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
return fmt.Errorf("generating key: %w", err)
}
serial, _ := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
// Use first extra name (hostname) as CN, fall back to localhost
cn := "localhost"
if len(extraNames) > 0 && extraNames[0] != "" {
cn = extraNames[0]
}
template := &x509.Certificate{
SerialNumber: serial,
Subject: pkix.Name{
Organization: []string{selfSignedOrganization},
CommonName: cn,
},
NotBefore: time.Now(),
NotAfter: time.Now().Add(selfSignedValidity),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
DNSNames: buildDNSNames(extraNames),
IPAddresses: buildIPList(extraNames),
}
certDER, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
if err != nil {
return fmt.Errorf("creating certificate: %w", err)
}
keyDER, err := x509.MarshalECPrivateKey(key)
if err != nil {
return fmt.Errorf("marshaling key: %w", err)
}
// This path creates the initial pair. Renewal retains its key so a
// failed certificate replacement cannot break the pair on disk.
if err := writeTLSFile(keyPath, pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})); err != nil {
return fmt.Errorf("writing key: %w", err)
}
if err := writeTLSFile(certPath, pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})); err != nil {
return fmt.Errorf("writing cert: %w", err)
}
return nil
}
// renewTLSCert never generates missing files: an operator may temporarily
// remove either file while replacing their certificate.
func renewTLSCert(certPath, keyPath string) error {
// Read and validate the same certificate that will be renewed, even
// when an operator replaces the files during this check.
// #nosec G304 -- paths are operator-configured TLS files.
certPEM, err := os.ReadFile(certPath)
if err != nil {
return err
}
block, prefix, suffix := leafCertificateBlock(certPEM)
if block == nil {
return nil
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
return fmt.Errorf("parsing certificate for renewal: %w", err)
}
if !ownCertExpiring(cert) {
return nil
}
// #nosec G304 -- paths are operator-configured TLS files.
keyPEM, err := os.ReadFile(keyPath)
if err != nil {
return err
}
pair, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
return fmt.Errorf("loading certificate for renewal: %w", err)
}
serial, err := rand.Int(rand.Reader, new(big.Int).Lsh(big.NewInt(1), 128))
if err != nil {
return err
}
cert.SerialNumber = serial
cert.NotBefore = time.Now()
cert.NotAfter = cert.NotBefore.Add(selfSignedValidity)
der, err := x509.CreateCertificate(rand.Reader, cert, cert, cert.PublicKey, pair.PrivateKey)
if err != nil {
return fmt.Errorf("renewing certificate: %w", err)
}
renewed := append(bytes.Clone(prefix), pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})...)
renewed = append(renewed, suffix...)
return writeTLSFile(certPath, renewed)
}
// leafCertificateBlock returns the certificate the TLS stack serves: the
// first CERTIFICATE block. A combined file, such as cPanel's service
// certificate, carries the private key ahead of it. The surrounding bytes
// must survive renewal because they may contain the key and chain.
func leafCertificateBlock(data []byte) (block *pem.Block, prefix, suffix []byte) {
rest := data
for {
block, rest = pem.Decode(rest)
if block == nil {
return nil, nil, nil
}
if block.Type == "CERTIFICATE" {
// Decode skips leading text and malformed blocks. Its accepted
// BEGIN marker is the last one before the consumed block's end.
end := len(data) - len(rest)
start := bytes.LastIndex(data[:end], []byte("-----BEGIN CERTIFICATE-----"))
return block, data[:start], rest
}
}
}
// ownCertExpiring recognizes CSM's self-signed certificates near expiry.
// A matching issuer name alone does not prove a self-signature.
func ownCertExpiring(cert *x509.Certificate) bool {
if !bytes.Equal(cert.RawIssuer, cert.RawSubject) {
return false
}
if len(cert.Subject.Organization) != 1 || cert.Subject.Organization[0] != selfSignedOrganization {
return false
}
if err := cert.CheckSignature(cert.SignatureAlgorithm, cert.RawTBSCertificate, cert.Signature); err != nil {
return false
}
return time.Until(cert.NotAfter) < renewBefore
}
// writeFileReplace writes data to a private temporary file beside path and
// renames it over path.
func writeFileReplace(path string, data []byte) error {
tmp, err := os.CreateTemp(filepath.Dir(path), "."+filepath.Base(path)+".*")
if err != nil {
return err
}
defer func() { _ = os.Remove(tmp.Name()) }()
if _, err := tmp.Write(data); err != nil {
_ = tmp.Close()
return err
}
if err := tmp.Close(); err != nil {
return err
}
return os.Rename(tmp.Name(), path)
}
// certReloader serves the certificate on disk and loads it again when the
// files change, so a renewed or replaced certificate needs no restart.
type certReloader struct {
certPath, keyPath string
mu sync.Mutex
cert *tls.Certificate
stamp string
}
func newCertReloader(certPath, keyPath string) (*certReloader, error) {
r := &certReloader{certPath: certPath, keyPath: keyPath}
if _, err := r.GetCertificate(nil); err != nil {
return nil, err
}
return r, nil
}
// GetCertificate is a tls.Config.GetCertificate. A pair that does not load
// (mid-replacement) keeps the previous certificate in service.
func (r *certReloader) GetCertificate(*tls.ClientHelloInfo) (*tls.Certificate, error) {
stamp := fileStamp(r.certPath) + "|" + fileStamp(r.keyPath)
r.mu.Lock()
defer r.mu.Unlock()
if r.cert != nil && stamp == r.stamp {
return r.cert, nil
}
cert, err := tls.LoadX509KeyPair(r.certPath, r.keyPath)
if err != nil {
if r.cert != nil {
return r.cert, nil
}
return nil, err
}
r.cert, r.stamp = &cert, stamp
return r.cert, nil
}
func fileStamp(path string) string {
info, err := os.Stat(path)
if err != nil {
return ""
}
return integrity.FileChangeKey(info)
}
func buildDNSNames(extra []string) []string {
names := []string{"localhost"}
for _, n := range extra {
if net.ParseIP(n) == nil { // not an IP - it's a hostname
names = append(names, n)
}
}
return names
}
func buildIPList(extra []string) []net.IP {
ips := []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("::1")}
for _, n := range extra {
if ip := net.ParseIP(n); ip != nil {
ips = append(ips, ip)
}
}
return ips
}
func fileExists(path string) bool {
_, err := os.Stat(path)
return err == nil
}
package webui
import (
"bytes"
"compress/gzip"
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"strings"
"time"
"github.com/pidginhost/csm/internal/checks"
"github.com/pidginhost/csm/internal/firewall"
"github.com/pidginhost/csm/internal/state"
"github.com/pidginhost/csm/internal/store"
)
// Recognised inverse-action keys. Each handler that records an undo entry
// sets one of these on the entry; apiUndoRun dispatches based on the value.
const (
undoInverseThreatBlock = "threat_bulk_unblock"
undoInverseThreatUnblock = "threat_bulk_block"
undoInverseThreatWhitelist = "threat_bulk_unwhitelist"
undoInverseThreatUnwhitelist = "threat_bulk_whitelist"
undoInverseFirewallUnblock = "firewall_bulk_reblock"
undoInverseFindingUndismiss = "finding_undismiss"
)
// maxUndoPayloadSize bounds decompression of persisted data while leaving
// room for the threat-row snapshot behind the 500-item firewall action.
const maxUndoPayloadSize = 4 * 1024 * 1024
// undoPayloadIPs is the payload schema for every undo entry we currently
// generate: a list of IPs plus an optional reason and timeout. Future undo
// kinds can add their own payload structs alongside this one.
type undoPayloadIPs struct {
BlockSnapshot bool `json:"block_snapshot,omitempty"`
RestoreBlocks map[string]firewall.BlockedEntry `json:"restore_blocks,omitempty"`
ExpectedBlocks map[string]firewall.BlockedEntry `json:"expected_blocks,omitempty"`
IPs []string `json:"ips"`
Reason string `json:"reason,omitempty"`
Timeout string `json:"timeout,omitempty"` // ParseDuration-compatible
// RestoreThreats carries the threat-DB rows a bulk action removed so the
// matching undo can put them back exactly.
RestoreThreats []undoThreatRow `json:"restore_threats,omitempty"`
// Dismissals carries what a finding dismissal changed so its undo can
// list the finding again and re-arm its alerts.
Dismissals []state.DismissUndo `json:"dismissals,omitempty"`
}
// undoThreatRow captures a removed threat-DB row's identity so undo can
// restore it with the same source and expiry, instead of resurrecting an
// auto-block row as a never-expiring operator block (or vice versa).
type undoThreatRow struct {
IP string `json:"ip"`
Reason string `json:"reason"`
Source string `json:"source,omitempty"`
ExpiresAt time.Time `json:"expires_at,omitzero"`
}
// recordUndoEntry persists an undo entry for the operator who issued r and
// returns the new entry's ID. The ID lets the calling handler surface the
// undo token to the client in the same response. Any store error is logged
// and swallowed so a bulk action never fails just because the undo queue
// could not be encoded or written.
func (s *Server) recordUndoEntry(r *http.Request, action, inverse, summary string, payload undoPayloadIPs) string {
if r == nil {
return ""
}
opkey := s.operatorKey(r)
if opkey == "" {
return ""
}
sdb := store.Global()
if sdb == nil {
return ""
}
raw, err := encodeUndoPayload(payload)
if err != nil {
log.Printf("webui: encode undo entry: %v", err)
return ""
}
entry, err := sdb.AppendUndoEntry(opkey, store.UndoEntry{
Targets: undoTargets(payload),
Action: action,
Inverse: inverse,
Payload: raw,
Summary: summary,
})
if err != nil {
log.Printf("webui: record undo entry: %v", err)
return ""
}
return entry.ID
}
// undoTargets names what an undo entry acts on: its IPs, or the finding keys
// of a dismissal.
func undoTargets(payload undoPayloadIPs) []string {
if len(payload.IPs) > 0 {
return payload.IPs
}
keys := make([]string, 0, len(payload.Dismissals))
for _, d := range payload.Dismissals {
keys = append(keys, d.Key)
}
return keys
}
func encodeUndoPayload(payload undoPayloadIPs) ([]byte, error) {
raw, err := json.Marshal(payload)
if err != nil {
return nil, err
}
if len(raw) > maxUndoPayloadSize {
return nil, fmt.Errorf("undo payload exceeds %d bytes", maxUndoPayloadSize)
}
var compressed bytes.Buffer
zw, err := gzip.NewWriterLevel(&compressed, gzip.BestSpeed)
if err != nil {
return nil, err
}
if _, err := zw.Write(raw); err != nil {
_ = zw.Close()
return nil, err
}
if err := zw.Close(); err != nil {
return nil, err
}
if compressed.Len() < len(raw) {
return compressed.Bytes(), nil
}
return raw, nil
}
func decodeUndoPayload(raw []byte, payload *undoPayloadIPs) error {
if len(raw) < 2 || raw[0] != 0x1f || raw[1] != 0x8b {
return json.Unmarshal(raw, payload)
}
zr, err := gzip.NewReader(bytes.NewReader(raw))
if err != nil {
return err
}
decoded, readErr := io.ReadAll(io.LimitReader(zr, maxUndoPayloadSize+1))
closeErr := zr.Close()
if readErr != nil {
return readErr
}
if closeErr != nil {
return closeErr
}
if len(decoded) > maxUndoPayloadSize {
return fmt.Errorf("undo payload exceeds %d bytes", maxUndoPayloadSize)
}
return json.Unmarshal(decoded, payload)
}
// undoPendingView is the JSON shape returned to the client. The payload is
// stripped because the client only needs identity + summary to render the
// banner; the server keeps the payload for the actual undo run.
type undoPendingView struct {
ID string `json:"id"`
Action string `json:"action"`
Inverse string `json:"inverse"`
Summary string `json:"summary"`
RecordedAt time.Time `json:"recorded_at"`
ExpiresAt time.Time `json:"expires_at"`
}
// apiUndoPending returns the latest non-expired undo entry for the operator,
// or an empty object when no entry is queued.
func (s *Server) apiUndoPending(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
sdb := store.Global()
if sdb == nil {
writeJSON(w, map[string]interface{}{})
return
}
entry, ok, err := sdb.LatestUndoEntry(opkey)
if err != nil {
writeJSONError(w, "Store error", http.StatusInternalServerError)
return
}
if !ok {
writeJSON(w, map[string]interface{}{})
return
}
writeJSON(w, undoPendingView{
ID: entry.ID,
Action: entry.Action,
Inverse: entry.Inverse,
Summary: entry.Summary,
RecordedAt: entry.RecordedAt,
ExpiresAt: entry.RecordedAt.Add(store.UndoTTL),
})
}
type undoRunRequest struct {
ID string `json:"id"`
}
type undoRunResponse struct {
OK bool `json:"ok"`
Action string `json:"action"`
Inverse string `json:"inverse"`
Count int `json:"count"`
}
// apiUndoRun consumes the named undo entry (or the most recent one when id
// is empty) and dispatches its inverse. Each successful undo also writes a
// "undo_<original>" audit entry so the trail records the reversal.
func (s *Server) apiUndoRun(w http.ResponseWriter, r *http.Request) {
s.threatActionMu.Lock()
defer s.threatActionMu.Unlock()
if r.Method != http.MethodPost {
writeJSONError(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
opkey := s.operatorKey(r)
if opkey == "" {
writeJSONError(w, "Unauthenticated", http.StatusUnauthorized)
return
}
var req undoRunRequest
if err := decodeJSONBodyLimited(w, r, 4*1024, &req); err != nil {
writeJSONError(w, "Invalid request body", http.StatusBadRequest)
return
}
sdb := store.Global()
if sdb == nil {
writeJSONError(w, "Store unavailable", http.StatusServiceUnavailable)
return
}
var (
entry store.UndoEntry
ok bool
err error
)
if req.ID == "" {
entry, ok, err = sdb.LatestUndoEntry(opkey)
if err == nil && ok {
entry, ok, err = sdb.ConsumeUndoEntry(opkey, entry.ID)
}
} else {
entry, ok, err = sdb.ConsumeUndoEntry(opkey, req.ID)
}
if err != nil {
writeJSONError(w, "Store error", http.StatusInternalServerError)
return
}
if !ok {
writeJSONError(w, "Undo window expired", http.StatusGone)
return
}
targets := strings.Join(entry.Targets, ", ")
resp, runErr := s.runUndoEntry(r, entry)
if runErr != nil {
// The entry is consumed and the inverse may have run part way, so the
// attempt is recorded even though it failed.
s.auditLog(r, "undo_"+entry.Action+"_failed", targets, entry.Summary+": "+runErr.Error())
writeJSONError(w, runErr.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "undo_"+entry.Action, fmt.Sprintf("%d items", resp.Count), entry.Summary+": "+targets)
writeJSON(w, resp)
}
func (s *Server) runUndoEntry(r *http.Request, entry store.UndoEntry) (undoRunResponse, error) {
var payload undoPayloadIPs
if len(entry.Payload) > 0 {
if err := decodeUndoPayload(entry.Payload, &payload); err != nil {
return undoRunResponse{}, fmt.Errorf("decode payload: %w", err)
}
}
resp := undoRunResponse{
OK: true,
Action: entry.Action,
Inverse: entry.Inverse,
}
switch entry.Inverse {
case undoInverseThreatBlock:
// Original action blocked IPs; inverse unblocks them and puts back
// the evidence that was on file before the block.
count, err := s.undoBulkBlock(payload)
if err != nil {
return undoRunResponse{}, err
}
resp.Count = count
case undoInverseThreatUnblock:
// Original unblocked IPs; inverse re-blocks them with the saved reason.
// The payload is self-written state, so a bad timeout falls back to
// the default instead of failing the undo.
timeout, terr := parseDuration(payload.Timeout)
if terr != nil || timeout == 0 {
timeout = 24 * time.Hour
}
reason := payload.Reason
if reason == "" {
reason = "Undo: re-block via CSM Web UI"
}
count, err := s.undoBulkReblock(payload.IPs, reason, timeout)
if err != nil {
return undoRunResponse{}, err
}
resp.Count = count
case undoInverseThreatWhitelist:
resp.Count = s.undoBulkWhitelist(payload)
case undoInverseThreatUnwhitelist:
resp.Count = s.undoBulkUnwhitelist(payload.IPs)
case undoInverseFirewallUnblock:
if payload.BlockSnapshot {
count, err := s.undoSnapshotBlocks(payload, false)
if err != nil {
return undoRunResponse{}, err
}
resp.Count = count
break
}
reason := payload.Reason
if reason == "" {
reason = "Undo: re-block via CSM Web UI"
}
timeout, terr := parseDuration(payload.Timeout)
if terr != nil || timeout == 0 {
timeout = 24 * time.Hour
}
count, err := s.undoBulkReblock(payload.IPs, reason, timeout)
if err != nil {
return undoRunResponse{}, err
}
restoreUndoThreatRows(payload.RestoreThreats)
resp.Count = count
case undoInverseFindingUndismiss:
for _, d := range payload.Dismissals {
if s.store.UndoDismiss(d) {
resp.Count++
}
}
default:
return undoRunResponse{}, fmt.Errorf("unknown inverse action %q", entry.Inverse)
}
return resp, nil
}
// captureUndoThreatRow snapshots the live threat row for ip. With
// temporaryOnly set it captures only rows tied to a firewall block, which is
// exactly the set a firewall unblock drops.
func captureUndoThreatRow(ip string, temporaryOnly bool) (undoThreatRow, bool) {
sdb := store.Global()
if sdb == nil {
return undoThreatRow{}, false
}
entry, ok := sdb.GetPermanentBlock(ip)
if !ok || entry.Expired(time.Now()) {
return undoThreatRow{}, false
}
if temporaryOnly && !entry.TiedToFirewallBlock() {
return undoThreatRow{}, false
}
return undoThreatRow{
IP: entry.IP,
Reason: entry.Reason,
Source: entry.Source,
ExpiresAt: entry.ExpiresAt,
}, true
}
func restoreUndoThreatRows(rows []undoThreatRow) {
tdb := checks.GetThreatDB()
if tdb == nil {
return
}
now := time.Now()
for _, row := range rows {
if _, err := parseAndValidateIP(row.IP); err != nil {
continue
}
if shouldRestoreUndoThreatAsPermanent(row, now) {
tdb.AddPermanent(row.IP, row.Reason)
continue
}
if row.ExpiresAt.IsZero() {
continue
}
ttl := time.Until(row.ExpiresAt)
if ttl <= 0 {
continue // already lapsed; nothing worth restoring
}
// A timed operator block keeps its operator source: restoring it as
// auto-block evidence would let the auto-block cleanup paths delete
// an operator's deliberate block.
if row.Source == store.ThreatSourceOperator {
tdb.AddOperatorTemporary(row.IP, row.Reason, ttl)
continue
}
tdb.AddTemporary(row.IP, row.Reason, ttl)
}
}
func shouldRestoreUndoThreatAsPermanent(row undoThreatRow, now time.Time) bool {
if row.Source == store.ThreatSourceOperator {
return row.ExpiresAt.IsZero()
}
if row.Source != "" || !row.ExpiresAt.IsZero() {
return false
}
legacy := store.PermanentBlockEntry{Reason: row.Reason}
return !legacy.Expired(now)
}
func (s *Server) undoBulkBlock(payload undoPayloadIPs) (int, error) {
if payload.BlockSnapshot {
return s.undoSnapshotBlocks(payload, true)
}
if s.blocker == nil {
return 0, fmt.Errorf("firewall engine not available")
}
count := 0
for _, ip := range payload.IPs {
if _, err := parseAndValidateIP(ip); err != nil {
continue
}
if err := s.blocker.UnblockIP(ip); err != nil {
continue
}
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.RemovePermanent(ip)
}
restoreUndoThreatRows(threatRowsForIP(payload.RestoreThreats, ip))
_ = flushCphulk(ip) // best effort
count++
}
return count, nil
}
func (s *Server) undoBulkReblock(ips []string, reason string, timeout time.Duration) (int, error) {
if s.blocker == nil {
return 0, fmt.Errorf("firewall engine not available")
}
count := 0
for _, ip := range ips {
if _, err := parseAndValidateIP(ip); err != nil {
continue
}
if err := blockIPForOperator(s.blocker, ip, reason, timeout); err != nil {
continue
}
count++
}
return count, nil
}
func (s *Server) undoBulkWhitelist(payload undoPayloadIPs) int {
count := 0
for _, ip := range payload.IPs {
if _, err := parseAndValidateIP(ip); err != nil {
continue
}
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.RemoveWhitelist(ip)
}
// The bulk whitelist added a firewall allow rule; leaving it lets a
// mis-whitelisted attacker bypass every future block indefinitely.
if s.blocker != nil {
if remover, ok := s.blocker.(allowRemover); ok {
_ = remover.RemoveAllowIP(ip)
}
}
count++
}
// Restore the threat rows the whitelist removed. RemoveWhitelist ran for
// every IP above, so the whitelist no longer suppresses these adds.
restoreUndoThreatRows(payload.RestoreThreats)
return count
}
func (s *Server) undoBulkUnwhitelist(ips []string) int {
count := 0
for _, ip := range ips {
if _, err := parseAndValidateIP(ip); err != nil {
continue
}
if tdb := checks.GetThreatDB(); tdb != nil {
tdb.AddWhitelist(ip)
}
count++
}
return count
}
package webui
import (
"net/http"
"os"
"strings"
"github.com/pidginhost/csm/internal/config"
"github.com/pidginhost/csm/internal/integrity"
"github.com/pidginhost/csm/internal/threatintel"
)
// SetVerifiedBotsReloader registers the callback the Daemon uses to push a
// saved verified_bots list into the live registry + rDNS verifier, so edits
// take effect without a restart. nil-safe: tests leave it unset.
func (s *Server) SetVerifiedBotsReloader(fn func() error) {
s.verifiedBotsReloader = fn
}
func (s *Server) handleVerifiedBots(w http.ResponseWriter, r *http.Request) {
s.renderTemplate(w, r, "verified-bots.html", map[string]string{
"Hostname": s.cfg.Hostname,
})
}
// apiVerifiedBots (GET /api/v1/verified-bots) returns the configured list plus
// the config etag for optimistic locking on save.
func (s *Server) apiVerifiedBots(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
diskBytes, err := os.ReadFile(s.cfg.ConfigFile) // #nosec G304 -- operator-supplied config path
if err != nil {
writeJSONError(w, "read config: "+err.Error(), http.StatusInternalServerError)
return
}
disk, err := config.LoadBytes(diskBytes)
if err != nil {
writeJSONError(w, "parse config: "+err.Error(), http.StatusInternalServerError)
return
}
bots := disk.Reputation.VerifiedBots
writeItems(w, bots, map[string]interface{}{
"total": len(bots),
"etag": disk.Integrity.ConfigHash,
"bot_ranges": botRangesSummary(disk),
})
}
// botRangesSummary is the read-only view of the built-in AI-crawler ranges the
// Verified Bots page shows: the configured auto-update posture plus the live
// per-bot prefix counts and last-refresh time from the active overlay.
func botRangesSummary(disk *config.Config) map[string]interface{} {
summary := map[string]interface{}{
"auto_update": disk.BotRangesAutoUpdate(),
"prefixes": threatintel.AICrawlerRangePrefixCounts(),
}
if secs, ok := durationSeconds(disk.Reputation.BotRanges.UpdateInterval); ok {
summary["update_interval_seconds"] = secs
}
if ts := threatintel.LastFetchedRangesRefresh(); !ts.IsZero() {
summary["last_refresh"] = ts
}
return summary
}
// apiVerifiedBotsApply (POST /api/v1/verified-bots/apply) validates and
// persists the whole verified_bots list to csm.yaml, then applies it live.
// The list is validated exactly as a config load would, so the same
// shared-hosting/over-broad-range guards apply here as on disk.
func (s *Server) apiVerifiedBotsApply(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
writeJSONError(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
// Serialize read-validate-write-reload so concurrent saves cannot race.
configMu := integrity.ConfigWriteMutex()
configMu.Lock()
defer configMu.Unlock()
if rejectConfigWriteDuringRollback(w) {
return
}
ifMatch := r.Header.Get("If-Match")
if ifMatch == "" {
writeJSONError(w, "If-Match header required", http.StatusBadRequest)
return
}
var body struct {
Bots *[]config.VerifiedBot `json:"bots"`
}
if err := decodeJSONBodyLimited(w, r, 256*1024, &body); err != nil {
writeJSONError(w, "invalid body: "+err.Error(), http.StatusBadRequest)
return
}
if body.Bots == nil {
writeJSONError(w, "bots is required", http.StatusBadRequest)
return
}
bots := normalizeVerifiedBotsForSave(*body.Bots)
diskBytes, err := os.ReadFile(s.cfg.ConfigFile) // #nosec G304 -- operator-supplied config path
if err != nil {
writeJSONError(w, "read config: "+err.Error(), http.StatusInternalServerError)
return
}
disk, err := config.LoadBytes(diskBytes)
if err != nil {
writeJSONError(w, "parse config: "+err.Error(), http.StatusInternalServerError)
return
}
disk.ConfigFile = s.cfg.ConfigFile
disk.ConfigDir = s.cfg.ConfigDir
if disk.Integrity.ConfigHash != ifMatch {
writeJSONError(w, "config changed on disk, reload", http.StatusPreconditionFailed)
return
}
if rejectIfConfDirChanged(w, s.cfg.ConfigDir, disk) {
return
}
clone := cloneConfigForSettingsApply(disk)
clone.Reputation.VerifiedBots = bots
var verr []fieldError
for _, v := range config.Validate(&clone) {
if v.Level == "error" && strings.HasPrefix(v.Field, "reputation.verified_bots") {
verr = append(verr, fieldError{Field: v.Field, Message: v.Message})
}
}
if len(verr) > 0 {
writeValidationErrors(w, verr)
return
}
// YAMLEdit block-renders []interface{} via yaml.Marshal; a typed
// []config.VerifiedBot is not recognized, so wrap each entry.
botsVal := make([]interface{}, len(bots))
for i, b := range bots {
botsVal[i] = b
}
edited, err := config.YAMLEdit(diskBytes, []config.YAMLChange{
{Path: []string{"reputation", "verified_bots"}, Value: botsVal},
})
if err != nil {
writeJSONError(w, "yaml edit: "+err.Error(), http.StatusInternalServerError)
return
}
if err := integrity.SignAndSavePreserving(s.cfg.ConfigFile, s.cfg.ConfigDir, edited, &clone, disk.Integrity.BinaryHash); err != nil {
writeJSONError(w, "save: "+err.Error(), http.StatusInternalServerError)
return
}
s.auditLog(r, "verified_bots_save", "reputation.verified_bots", "csm.yaml rewritten")
newIntegrity := clone.Integrity
// reputation is a hot-reload-safe section: apply to the live config now.
if live := config.Active(); live != nil {
liveClone := *live
liveClone.Reputation.VerifiedBots = bots
applySignedIntegrityState(&liveClone, &clone)
config.SetActive(&liveClone)
} else {
config.SetActive(&clone)
}
// Push the new list into the running registry + verifier (no restart).
if s.verifiedBotsReloader != nil {
_ = s.verifiedBotsReloader()
}
writeJSON(w, map[string]interface{}{
"ok": true,
"count": len(bots),
"new_etag": newIntegrity.ConfigHash,
})
}
func normalizeVerifiedBotsForSave(in []config.VerifiedBot) []config.VerifiedBot {
out := make([]config.VerifiedBot, len(in))
copy(out, in)
for i := range out {
out[i].UASubstrings = nilIfEmptyStrings(out[i].UASubstrings)
out[i].RDNSSuffixes = nilIfEmptyStrings(out[i].RDNSSuffixes)
out[i].IPRanges = nilIfEmptyStrings(out[i].IPRanges)
}
return out
}
func nilIfEmptyStrings(v []string) []string {
if len(v) == 0 {
return nil
}
return v
}
package wpcheck
import (
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"sync"
"time"
"golang.org/x/sys/unix"
)
type Cache struct {
mu sync.RWMutex
statePath string
checksums map[string]map[string]string // core: "<version>:<locale>" -> relPath -> MD5
pluginChecksums map[string]map[string]string // plugins: "<slug>:<version>" -> relPath -> SHA256
roots map[string]rootEntry
fetching map[string]bool
// Retain admission times after completion so fast responses and repeated
// misses cannot turn tenant-written release names into unlimited fetches.
coreFetchAfter map[string]time.Time
// pluginNotFoundUntil records slug+version pairs that wordpress.org
// returned 404 for, paired with the absolute time at which the
// suppression expires. Plugins hosted outside wp.org (paid forks,
// custom internal plugins) would otherwise re-arm the 4-attempt
// retry cycle on every cache miss. The TTL ensures wp.org adding
// a plugin later still gets picked up.
pluginNotFoundUntil map[string]time.Time
// stopCh, when non-nil and closed, signals pending retry timers to
// drop their scheduled fetch instead of firing. Wired by the daemon
// to the FileMonitor stopCh so checksum-retry chains do not survive
// daemon shutdown.
stopMu sync.RWMutex
stopCh <-chan struct{}
}
type rootEntry struct {
version string
locale string
}
func NewCache(statePath string) *Cache {
c := &Cache{
statePath: statePath,
checksums: make(map[string]map[string]string),
pluginChecksums: make(map[string]map[string]string),
roots: make(map[string]rootEntry),
fetching: make(map[string]bool),
coreFetchAfter: make(map[string]time.Time),
pluginNotFoundUntil: make(map[string]time.Time),
}
c.loadFromDisk()
return c
}
func cacheKey(version, locale string) string {
return version + ":" + locale
}
func diskFilename(version, locale string) string {
return version + "_" + locale + ".json"
}
func (c *Cache) loadFromDisk() {
dir := filepath.Join(c.statePath, "wp-checksums")
entries, err := os.ReadDir(dir)
if err != nil {
return
}
for _, entry := range entries {
name := entry.Name()
if entry.IsDir() || !strings.HasSuffix(name, ".json") {
continue
}
// #nosec G304 -- dir is {statePath}/wp-checksums; name comes from
// our own os.ReadDir of that same dir.
data, err := os.ReadFile(filepath.Join(dir, name))
if err != nil {
continue
}
checksums, err := ParseChecksumResponse(data)
if err != nil {
continue
}
base := strings.TrimSuffix(name, ".json")
parts := strings.SplitN(base, "_", 2)
if len(parts) != 2 {
continue
}
c.checksums[cacheKey(parts[0], parts[1])] = checksums
}
}
// PersistChecksums writes checksum data to disk atomically (tmpfile + rename)
// and populates the in-memory cache. The file is written to {statePath}/wp-checksums/.
func (c *Cache) PersistChecksums(version, locale string, rawJSON []byte, checksums map[string]string) error {
dir := filepath.Join(c.statePath, "wp-checksums")
if err := os.MkdirAll(dir, 0700); err != nil {
return fmt.Errorf("creating wp-checksums dir: %w", err)
}
filename := diskFilename(version, locale)
tmpPath := filepath.Join(dir, filename+".tmp")
finalPath := filepath.Join(dir, filename)
if err := os.WriteFile(tmpPath, rawJSON, 0600); err != nil {
return fmt.Errorf("writing temp file: %w", err)
}
if err := os.Rename(tmpPath, finalPath); err != nil {
os.Remove(tmpPath)
return fmt.Errorf("renaming to final: %w", err)
}
c.mu.Lock()
c.checksums[cacheKey(version, locale)] = checksums
c.mu.Unlock()
return nil
}
func (c *Cache) lookupChecksum(version, locale, relativePath string) (string, bool) {
c.mu.RLock()
defer c.mu.RUnlock()
versionMap, ok := c.checksums[cacheKey(version, locale)]
if !ok {
return "", false
}
md5hex, ok := versionMap[relativePath]
return md5hex, ok
}
func (c *Cache) hasChecksums(version, locale string) bool {
c.mu.RLock()
ok := c.checksums[cacheKey(version, locale)] != nil
c.mu.RUnlock()
return ok
}
func (c *Cache) getRoot(root string) (version, locale string, ok bool) {
c.mu.RLock()
entry, ok := c.roots[root]
c.mu.RUnlock()
if !ok {
return "", "", false
}
return entry.version, entry.locale, true
}
func (c *Cache) setRoot(root, version, locale string) {
c.mu.Lock()
c.roots[root] = rootEntry{version: version, locale: locale}
c.mu.Unlock()
}
func (c *Cache) invalidateRoot(root string) {
c.mu.Lock()
delete(c.roots, root)
c.mu.Unlock()
}
const (
coreFetchMaxPending = 8
coreFetchHistoryMax = 64
coreFetchCooldown = time.Hour
)
func (c *Cache) startBackgroundFetch(version, locale string) {
if c.isStopped() {
return
}
key := cacheKey(version, locale)
c.mu.Lock()
if c.fetching[key] || c.checksums[key] != nil || len(c.fetching) >= coreFetchMaxPending {
c.mu.Unlock()
return
}
now := time.Now()
for oldKey, until := range c.coreFetchAfter {
if !now.Before(until) && !c.fetching[oldKey] {
delete(c.coreFetchAfter, oldKey)
}
}
if now.Before(c.coreFetchAfter[key]) || len(c.coreFetchAfter) >= coreFetchHistoryMax {
c.mu.Unlock()
return
}
c.coreFetchAfter[key] = now.Add(coreFetchCooldown)
c.fetching[key] = true
c.mu.Unlock()
go c.fetchWithRetry(version, locale, 0)
}
func (c *Cache) fetchWithRetry(version, locale string, attempt int) {
backoffs := []time.Duration{1 * time.Minute, 5 * time.Minute, 15 * time.Minute, 1 * time.Hour}
key := cacheKey(version, locale)
if c.isStopped() {
c.clearFetching(key)
return
}
rawJSON, checksums, err := FetchChecksums(version, locale)
if err != nil {
if attempt >= len(backoffs) {
c.mu.Lock()
c.coreFetchAfter[key] = time.Now().Add(coreFetchCooldown)
delete(c.fetching, key)
c.mu.Unlock()
fmt.Fprintf(os.Stderr, "wpcheck: core fetch abandoned for WP %s (%s) after %d attempts: %v\n",
version, locale, attempt+1, err)
return
}
delay := backoffs[attempt]
fmt.Fprintf(os.Stderr, "wpcheck: fetch failed for WP %s (%s), retry in %v: %v\n",
version, locale, delay, err)
c.scheduleRetry(delay, func() {
c.fetchWithRetry(version, locale, attempt+1)
}, func() {
c.clearFetching(key)
})
return
}
if err := c.PersistChecksums(version, locale, rawJSON, checksums); err != nil {
fmt.Fprintf(os.Stderr, "wpcheck: persist failed for WP %s (%s): %v\n", version, locale, err)
}
c.clearFetching(key)
fmt.Fprintf(os.Stderr, "wpcheck: cached %d checksums for WP %s (%s)\n", len(checksums), version, locale)
}
const maxFileSize = 2 << 20
// readCompleteFileForHash returns a stable-size snapshot of a regular file.
// Hash verification must never accept only a prefix: if a known-good file is
// exactly maxFileSize bytes, an attacker could otherwise append a payload that
// the one-shot bounded Pread silently ignores.
func readCompleteFileForHash(fd int) []byte {
var before unix.Stat_t
if err := unix.Fstat(fd, &before); err != nil || before.Size <= 0 || before.Size > maxFileSize {
return nil
}
if before.Mode&unix.S_IFMT != unix.S_IFREG {
return nil
}
data := make([]byte, int(before.Size))
offset := 0
interrupts := 0
for offset < len(data) {
n, err := unix.Pread(fd, data[offset:], int64(offset))
if n > 0 {
offset += n
interrupts = 0
}
if err != nil && !errors.Is(err, unix.EINTR) {
return nil
}
if n == 0 {
if !errors.Is(err, unix.EINTR) {
return nil
}
interrupts++
if interrupts > 100 {
return nil
}
}
}
var after unix.Stat_t
if err := unix.Fstat(fd, &after); err != nil || before.Dev != after.Dev ||
before.Ino != after.Ino || before.Size != after.Size {
return nil
}
return data
}
package wpcheck
import (
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
)
// wpRootLevelFiles are filenames that only exist at the WP installation root.
// index.php is excluded - it requires a secondary check (version.php must exist).
var wpRootLevelFiles = map[string]bool{
"wp-activate.php": true,
"wp-blog-header.php": true,
"wp-comments-post.php": true,
"wp-config-sample.php": true,
"wp-cron.php": true,
"wp-links-opml.php": true,
"wp-load.php": true,
"wp-login.php": true,
"wp-mail.php": true,
"wp-settings.php": true,
"wp-signup.php": true,
"wp-trackback.php": true,
"xmlrpc.php": true,
}
// DetectWPRoot returns the WordPress installation root directory for a file path,
// or empty string if the path is not inside a WP core location.
//
// Detection methods:
// - Path contains /wp-includes/ or /wp-admin/ → root is everything before that segment
// - Filename is a known root-level WP file → root is the parent directory
// - Filename is index.php and version.php exists in wp-includes/ → root is the parent directory
func DetectWPRoot(path string) string {
// Check for /wp-includes/ or /wp-admin/ in path
for _, marker := range []string{"/wp-includes/", "/wp-admin/"} {
if idx := strings.Index(path, marker); idx >= 0 {
return path[:idx]
}
}
// Check for direct wp-includes or wp-admin (file directly inside)
dir := filepath.Dir(path)
base := filepath.Base(dir)
if base == "wp-includes" || base == "wp-admin" {
return filepath.Dir(dir)
}
// Check for known root-level WP files
name := filepath.Base(path)
if wpRootLevelFiles[name] {
return dir
}
// Special case: index.php requires version.php to confirm WP root
if name == "index.php" {
versionPath := filepath.Join(dir, "wp-includes", "version.php")
if _, err := os.Stat(versionPath); err == nil {
return dir
}
}
return ""
}
// RelativePath computes the path of a file relative to the WP root.
// Returns empty string if the file is not under root.
func RelativePath(root, path string) string {
rel, err := filepath.Rel(root, path)
if err != nil || strings.HasPrefix(rel, "..") {
return ""
}
return rel
}
var (
reVersion = regexp.MustCompile(`\$wp_version\s*=\s*'([^']+)'`)
reLocale = regexp.MustCompile(`\$wp_local_package\s*=\s*'([^']+)'`)
// version.php is tenant-writable and its strings name a root-written
// cache file and a query string, so both are held to the shapes
// WordPress actually ships: "6.5", "6.5.2", "6.7-RC1", "6.8-alpha-59245"
// and locales such as "en_US", "de_DE_formal", "ary".
reValidVersion = regexp.MustCompile(`^[0-9]+(?:\.[0-9]+){0,3}(?:-[A-Za-z0-9]+(?:[.-][A-Za-z0-9]+)*)?$`)
reValidLocale = regexp.MustCompile(`^[a-z]{2,3}(?:_[A-Za-z0-9]{2,10}){0,2}$`)
)
// Release identifiers are retained per candidate and used in request URLs
// and cache filenames. Bound their bytes as well as the number of lookups.
const maxCoreVersionLength = 64
// ParseVersionContent extracts the WP version and locale from version.php content.
// Locale defaults to "en_US" if $wp_local_package is not present. A version or
// locale outside the shapes WordPress ships is an error: the install is then
// treated as unverifiable rather than letting tenant-chosen text reach the
// checksum cache path or the API query.
func ParseVersionContent(data []byte) (version, locale string, err error) {
m := reVersion.FindSubmatch(data)
if m == nil {
return "", "", errors.New("wp_version not found in version.php")
}
if len(m[1]) > maxCoreVersionLength {
return "", "", errors.New("wp_version in version.php is too long")
}
version = string(m[1])
if !reValidVersion.MatchString(version) {
return "", "", fmt.Errorf("wp_version %q in version.php is not a WordPress version string", version)
}
locale = "en_US"
if lm := reLocale.FindSubmatch(data); lm != nil {
locale = string(lm[1])
}
if !reValidLocale.MatchString(locale) {
return "", "", fmt.Errorf("wp_local_package %q in version.php is not a WordPress locale", locale)
}
return version, locale, nil
}
// ReadVersionFile reads and parses {root}/wp-includes/version.php.
func ReadVersionFile(root string) (version, locale string, err error) {
// #nosec G304 -- path derived from a WordPress install root discovered
// by the scanner under configured /home/*/public_html paths.
data, err := os.ReadFile(filepath.Join(root, "wp-includes", "version.php"))
if err != nil {
return "", "", err
}
return ParseVersionContent(data)
}
package wpcheck
import (
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"time"
)
var httpClient = &http.Client{Timeout: 10 * time.Second}
// checksumAPIURL returns the WordPress.org checksum API URL for a version and
// locale. Both values are validated at parse time; escaping them here keeps
// the query shape independent of that validation.
func checksumAPIURL(version, locale string) string {
return fmt.Sprintf("https://api.wordpress.org/core/checksums/1.0/?version=%s&locale=%s", url.QueryEscape(version), url.QueryEscape(locale))
}
// checksumResponse is the JSON structure returned by the WP checksum API.
type checksumResponse struct {
Checksums map[string]string `json:"checksums"`
}
// ParseChecksumResponse parses the JSON response from the WP checksum API.
// Returns a map of relative_path -> md5_hex.
func ParseChecksumResponse(data []byte) (map[string]string, error) {
var resp checksumResponse
if err := json.Unmarshal(data, &resp); err != nil {
return nil, fmt.Errorf("invalid JSON: %w", err)
}
if len(resp.Checksums) == 0 {
return nil, errors.New("empty or missing checksums in response")
}
return resp.Checksums, nil
}
// FetchChecksums fetches official checksums from api.wordpress.org for a given
// version and locale. Returns the raw response body and parsed checksums.
func FetchChecksums(version, locale string) (rawJSON []byte, checksums map[string]string, err error) {
url := checksumAPIURL(version, locale)
resp, err := httpClient.Get(url)
if err != nil {
return nil, nil, fmt.Errorf("HTTP request failed: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, nil, fmt.Errorf("HTTP %d from %s", resp.StatusCode, url)
}
body, err := io.ReadAll(io.LimitReader(resp.Body, 2<<20)) // 2MB limit
if err != nil {
return nil, nil, fmt.Errorf("reading response: %w", err)
}
checksums, err = ParseChecksumResponse(body)
if err != nil {
return nil, nil, err
}
return body, checksums, nil
}
package wpcheck
import (
"crypto/subtle"
"encoding/hex"
)
const maxHexDigestLength = 64
func constantTimeHexDigestEqual(actualDigest []byte, expectedHex string) bool {
actualHexLen := len(actualDigest) * 2
if actualHexLen == 0 || actualHexLen > maxHexDigestLength {
return false
}
var actual [maxHexDigestLength]byte
var expected [maxHexDigestLength]byte
hex.Encode(actual[:actualHexLen], actualDigest)
if len(expectedHex) >= actualHexLen {
copy(expected[:actualHexLen], expectedHex[:actualHexLen])
} else {
copy(expected[:actualHexLen], expectedHex)
}
lengthEqual := 0
if len(expectedHex) == actualHexLen {
lengthEqual = 1
}
digestEqual := subtle.ConstantTimeCompare(actual[:actualHexLen], expected[:actualHexLen])
return lengthEqual&digestEqual == 1
}
package wpcheck
import (
"archive/zip"
"bytes"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"io"
"net/http"
"os"
"path/filepath"
"regexp"
"strings"
"time"
)
// Plugin verification mirrors the core-file verification path: when an
// installed or update-staged plugin file matches the hash we computed from the
// plugin's official wordpress.org ZIP, signature/YARA rule matches on it are
// false positives and should not fire.
const (
pluginsSegment = "/wp-content/plugins/"
upgradeSegment = "/wp-content/upgrade/"
)
// DetectPluginRoot returns the plugin root directory and slug for a path that
// sits inside a plugin, either installed under /wp-content/plugins/<slug>/ or
// staged by an in-progress update under /wp-content/upgrade/<package>/<slug>/.
// Returns empty strings if the path is not inside a plugin.
func DetectPluginRoot(path string) (root, slug string) {
if root, slug := detectInstalledPluginRoot(path); root != "" {
return root, slug
}
return detectStagedPluginRoot(path)
}
func detectInstalledPluginRoot(path string) (root, slug string) {
idx := strings.Index(path, pluginsSegment)
if idx < 0 {
return "", ""
}
tail := path[idx+len(pluginsSegment):]
slashIdx := strings.Index(tail, "/")
if slashIdx <= 0 {
return "", ""
}
slug = tail[:slashIdx]
if !safePluginPathComponent(slug) {
return "", ""
}
root = path[:idx+len(pluginsSegment)] + slug
return root, slug
}
// detectStagedPluginRoot resolves the layout WordPress unpacks an update into:
// wp-content/upgrade/<package>/<slug>/<rest>, moved into wp-content/plugins/
// only once the install succeeds. The staged tree carries the same files and
// plugin header, so the per-file hash comparison against the official ZIP also
// works before the move. Without this, a routine plugin update leaves every
// one of its files unverifiable while it is staged.
func detectStagedPluginRoot(path string) (root, slug string) {
idx := strings.Index(path, upgradeSegment)
if idx < 0 {
return "", ""
}
pkg, tail, ok := strings.Cut(path[idx+len(upgradeSegment):], "/")
if !ok || !safePluginPathComponent(pkg) {
return "", ""
}
slug, tail, ok = strings.Cut(tail, "/")
// A package directory with no file below <slug>/ is not a staged plugin.
if !ok || tail == "" || !safePluginPathComponent(slug) {
return "", ""
}
return path[:idx+len(upgradeSegment)] + pkg + "/" + slug, slug
}
// safePluginPathComponent rejects the components that would let a crafted path
// resolve a root outside the directory it appears to name.
func safePluginPathComponent(name string) bool {
return name != "" && name != "." && name != ".."
}
const (
pluginHeaderReadLimit = 8192
maxPluginRootEntries = 256
)
var (
rePluginNameHeader = regexp.MustCompile(`(?im)^[ \t/*#@]*Plugin Name:[ \t]*[^ \t\r\n]`)
rePluginVersionHeader = regexp.MustCompile(`(?im)^[ \t/*#@]*Version:[ \t]*([^\s]+)`)
)
// ReadPluginVersion extracts the Version: header from the plugin's main
// file. Most plugins use <slug>.php, but WordPress permits any root-level PHP
// filename. The fallback directory scan is bounded and requires Plugin Name:
// as well as Version: so theme and core staging trees fail closed.
func ReadPluginVersion(pluginRoot, slug string) (string, error) {
if !safePluginPathComponent(slug) {
return "", errors.New("invalid plugin slug")
}
preferredName := slug + ".php"
preferredPath := filepath.Join(pluginRoot, preferredName)
version, found, err := readPluginVersionHeader(preferredPath)
if err == nil && found {
return version, nil
}
if err != nil && !errors.Is(err, os.ErrNotExist) {
return "", err
}
// #nosec G304 -- pluginRoot is derived from a path the scanner received
// from fanotify under a recognized plugin or update-staging layout.
dir, err := os.Open(pluginRoot)
if err != nil {
return "", err
}
defer func() { _ = dir.Close() }()
entries, err := dir.ReadDir(maxPluginRootEntries + 1)
if err != nil && !errors.Is(err, io.EOF) {
return "", fmt.Errorf("reading plugin root: %w", err)
}
if len(entries) > maxPluginRootEntries {
return "", fmt.Errorf("plugin root has more than %d entries", maxPluginRootEntries)
}
versions := make(map[string]struct{})
for _, entry := range entries {
if entry.Name() == preferredName || filepath.Ext(entry.Name()) != ".php" {
continue
}
info, infoErr := entry.Info()
if infoErr != nil {
return "", fmt.Errorf("stat plugin entry %s: %w", entry.Name(), infoErr)
}
if !info.Mode().IsRegular() {
continue
}
candidatePath := filepath.Join(pluginRoot, entry.Name())
candidateVersion, candidateFound, readErr := readPluginVersionHeader(candidatePath)
if readErr != nil {
return "", readErr
}
if candidateFound {
versions[candidateVersion] = struct{}{}
}
}
if len(versions) == 0 {
return "", errors.New("plugin header not found in root PHP files")
}
if len(versions) > 1 {
return "", errors.New("ambiguous plugin versions in root PHP files")
}
for candidateVersion := range versions {
return candidateVersion, nil
}
return "", errors.New("plugin version unavailable")
}
func readPluginVersionHeader(path string) (version string, found bool, err error) {
// #nosec G304 -- callers construct path from a scanner-derived root and a
// safe immediate child. Reads are bounded to WordPress's 8 KiB header limit.
f, err := os.Open(path)
if err != nil {
return "", false, err
}
defer func() { _ = f.Close() }()
buf, err := io.ReadAll(io.LimitReader(f, pluginHeaderReadLimit))
if err != nil {
return "", false, fmt.Errorf("reading plugin header: %w", err)
}
if !rePluginNameHeader.Match(buf) {
return "", false, nil
}
m := rePluginVersionHeader.FindSubmatch(buf)
if m == nil {
return "", false, nil
}
return string(m[1]), true, nil
}
// pluginZipURL returns the canonical wordpress.org download URL for a given
// plugin slug and version.
func pluginZipURL(slug, version string) string {
return fmt.Sprintf("https://downloads.wordpress.org/plugin/%s.%s.zip", slug, version)
}
// FetchPluginChecksums downloads the plugin ZIP from wordpress.org,
// extracts each file, and returns a map of relative path -> SHA256 hex.
// The returned paths are relative to the plugin root (the leading
// "<slug>/" prefix from the ZIP entries is stripped).
func FetchPluginChecksums(slug, version string) (map[string]string, error) {
return fetchPluginChecksumsFromURL(pluginZipURL(slug, version), slug)
}
const maxPluginZipBytes = 100 << 20 // 100 MB ceiling
// ErrPluginNotInWPOrg is returned when wordpress.org responds with HTTP 404
// for a plugin slug+version. Plugins that are not in the wp.org repository
// (paid forks, custom internal plugins, slugs that simply do not exist) need
// to be distinguished from transient errors so the cache can suppress
// further fetch attempts for a TTL.
//
// 5xx responses, network errors, and malformed responses are NOT this
// error - those keep their normal retry behaviour because the plugin may
// still exist in the catalogue and wp.org may simply be having an outage.
var ErrPluginNotInWPOrg = errors.New("plugin not in wordpress.org repository")
func fetchPluginChecksumsFromURL(url, slug string) (map[string]string, error) {
resp, err := httpClient.Get(url) //nolint:gosec,bodyclose // httpClient has a timeout; body is closed below.
if err != nil {
return nil, fmt.Errorf("plugin zip GET failed: %w", err)
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode == http.StatusNotFound {
return nil, fmt.Errorf("plugin zip HTTP 404 from %s: %w", url, ErrPluginNotInWPOrg)
}
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("plugin zip HTTP %d from %s", resp.StatusCode, url)
}
body, err := io.ReadAll(io.LimitReader(resp.Body, maxPluginZipBytes))
if err != nil {
return nil, fmt.Errorf("reading plugin zip: %w", err)
}
zr, err := zip.NewReader(bytes.NewReader(body), int64(len(body)))
if err != nil {
return nil, fmt.Errorf("opening plugin zip: %w", err)
}
out := make(map[string]string, len(zr.File))
prefix := slug + "/"
for _, zf := range zr.File {
if zf.FileInfo().IsDir() {
continue
}
name := zf.Name
if !strings.HasPrefix(name, prefix) {
// Malformed ZIP (e.g. nested into a differently-named folder).
// Skip; callers detect partial results by checking cache emptiness.
continue
}
rel := filepath.Clean(strings.TrimPrefix(name, prefix))
// Reject path-traversal and absolute paths: a crafted ZIP entry
// named "<slug>/../../etc/passwd" would otherwise land in the
// checksum map. Defense-in-depth against a compromised CDN.
if rel == "." || strings.HasPrefix(rel, "..") || strings.HasPrefix(rel, "/") {
continue
}
rc, err := zf.Open()
if err != nil {
return nil, fmt.Errorf("opening zip entry %s: %w", name, err)
}
// Cap decompressed size per entry. Without this, a zip-bomb whose
// compressed body fits under maxPluginZipBytes can still exhaust
// memory during io.Copy. +1 lets us detect overflow.
limited := io.LimitReader(rc, maxPluginZipBytes+1)
h := sha256.New()
nCopied, err := io.Copy(h, limited)
_ = rc.Close()
if err != nil {
return nil, fmt.Errorf("hashing zip entry %s: %w", name, err)
}
if nCopied > maxPluginZipBytes {
return nil, fmt.Errorf("zip entry %s exceeds per-entry size cap", name)
}
out[rel] = hex.EncodeToString(h.Sum(nil))
}
if len(out) == 0 {
return nil, errors.New("plugin zip yielded no checksums")
}
return out, nil
}
// --- Cache plugin support ------------------------------------------------
func pluginKey(slug, version string) string {
return slug + ":" + version
}
func (c *Cache) setPluginChecksums(slug, version string, checksums map[string]string) {
c.mu.Lock()
defer c.mu.Unlock()
if c.pluginChecksums == nil {
c.pluginChecksums = make(map[string]map[string]string)
}
c.pluginChecksums[pluginKey(slug, version)] = checksums
}
func (c *Cache) lookupPluginChecksum(slug, version, relPath string) (string, bool) {
c.mu.RLock()
defer c.mu.RUnlock()
m, ok := c.pluginChecksums[pluginKey(slug, version)]
if !ok {
return "", false
}
h, ok := m[relPath]
return h, ok
}
func (c *Cache) hasPluginChecksums(slug, version string) bool {
c.mu.RLock()
_, ok := c.pluginChecksums[pluginKey(slug, version)]
c.mu.RUnlock()
return ok
}
// pluginNotFoundTTL bounds how long a wp.org 404 outcome suppresses
// re-fetches for the same slug+version. After expiry the next cache miss
// retries normally, so a plugin that wp.org publishes later will be
// picked up. 72 hours strikes a balance between not flooding wp.org with
// requests for non-existent plugins and propagating corrections in
// reasonable time.
const pluginNotFoundTTL = 72 * time.Hour
// markPluginNotFound records a wp.org 404 outcome so subsequent fetches
// short-circuit until ttl elapses. Caller passes ttl explicitly so tests
// can shorten or invert it; production code should use pluginNotFoundTTL.
func (c *Cache) markPluginNotFound(slug, version string, ttl time.Duration) {
key := pluginKey(slug, version)
c.mu.Lock()
if c.pluginNotFoundUntil == nil {
c.pluginNotFoundUntil = make(map[string]time.Time)
}
c.pluginNotFoundUntil[key] = time.Now().Add(ttl)
c.mu.Unlock()
}
// isPluginNotFound reports whether an unexpired wp.org 404 marker exists
// for slug+version. Markers are scoped to slug+version so a fork of a
// plugin under a new version number that DOES exist on wp.org is still
// fetched.
func (c *Cache) isPluginNotFound(slug, version string) bool {
key := pluginKey(slug, version)
c.mu.RLock()
until, ok := c.pluginNotFoundUntil[key]
c.mu.RUnlock()
if !ok {
return false
}
return time.Now().Before(until)
}
func (c *Cache) startBackgroundPluginFetch(slug, version string) {
if c.isStopped() {
return
}
// wp.org has already told us this slug+version does not exist;
// suppress the fetch entirely until the marker expires. Without this
// gate every cache miss for a non-wp.org plugin would re-arm the
// 4-attempt retry cycle.
if c.isPluginNotFound(slug, version) {
return
}
key := pluginKey(slug, version)
c.mu.Lock()
if c.fetching == nil {
c.fetching = make(map[string]bool)
}
if c.fetching[key] {
c.mu.Unlock()
return
}
c.fetching[key] = true
c.mu.Unlock()
go c.fetchPluginWithRetry(slug, version, 0)
}
// fetchPluginWithRetry mirrors the core-checksum fetchWithRetry: the
// fetching flag stays set across retries so cache-miss events for the
// same slug/version do not spawn new goroutines. On exhaustion the flag
// is cleared so a future event can retry fresh.
//
// Special case: an HTTP 404 from wordpress.org is treated as a definitive
// "this plugin is not in the wp.org repository" signal. We mark the
// slug+version not-found for pluginNotFoundTTL and skip the retry cycle
// entirely. Network errors and 5xx responses keep their normal retry
// behaviour - those are transient.
func (c *Cache) fetchPluginWithRetry(slug, version string, attempt int) {
backoffs := []time.Duration{1 * time.Minute, 5 * time.Minute, 15 * time.Minute, 1 * time.Hour}
key := pluginKey(slug, version)
if c.isStopped() {
c.clearFetching(key)
return
}
checksums, err := FetchPluginChecksums(slug, version)
if err == nil {
c.setPluginChecksums(slug, version, checksums)
c.clearFetching(key)
fmt.Fprintf(os.Stderr, "wpcheck: cached %d checksums for plugin %s %s\n", len(checksums), slug, version)
return
}
if errors.Is(err, ErrPluginNotInWPOrg) {
c.markPluginNotFound(slug, version, pluginNotFoundTTL)
c.clearFetching(key)
fmt.Fprintf(os.Stderr, "wpcheck: plugin %s %s not in wp.org repository, suppressing retries for %s\n",
slug, version, pluginNotFoundTTL)
return
}
if attempt >= len(backoffs) {
c.clearFetching(key)
fmt.Fprintf(os.Stderr, "wpcheck: plugin fetch abandoned for %s %s after %d attempts: %v\n",
slug, version, attempt+1, err)
return
}
delay := backoffs[attempt]
fmt.Fprintf(os.Stderr, "wpcheck: plugin fetch failed for %s %s, retry in %v: %v\n",
slug, version, delay, err)
c.scheduleRetry(delay, func() {
c.fetchPluginWithRetry(slug, version, attempt+1)
}, func() {
c.clearFetching(key)
})
}
package wpcheck
import (
"time"
)
// SetStopCh wires a daemon-level cancellation channel into the cache.
// Closing the channel signals every pending checksum-retry timer to drop
// the scheduled fetch instead of firing. Safe to call once after NewCache,
// before any fetches start.
func (c *Cache) SetStopCh(stop <-chan struct{}) {
c.stopMu.Lock()
c.stopCh = stop
c.stopMu.Unlock()
}
func (c *Cache) currentStopCh() <-chan struct{} {
c.stopMu.RLock()
defer c.stopMu.RUnlock()
return c.stopCh
}
func (c *Cache) isStopped() bool {
return stopClosed(c.currentStopCh())
}
func stopClosed(stop <-chan struct{}) bool {
if stop == nil {
return false
}
select {
case <-stop:
return true
default:
return false
}
}
func (c *Cache) clearFetching(key string) {
c.mu.Lock()
delete(c.fetching, key)
c.mu.Unlock()
}
// scheduleRetry runs fn after delay unless the cache's stop channel
// closes first. Always returns immediately. Cancellation is necessary
// because the longest checksum-retry backoff (1 hour) survives daemon
// shutdown otherwise, causing wp.org fetches against torn-down state.
func (c *Cache) scheduleRetry(delay time.Duration, fn func(), onCancel func()) {
stop := c.currentStopCh()
go func() {
if stopClosed(stop) {
runCancel(onCancel)
return
}
timer := time.NewTimer(delay)
defer timer.Stop()
select {
case <-timer.C:
if stopClosed(stop) {
runCancel(onCancel)
return
}
fn()
case <-stop:
runCancel(onCancel)
}
}()
}
func runCancel(onCancel func()) {
if onCancel != nil {
onCancel()
}
}
package wpcheck
import (
"crypto/md5" // #nosec G501 -- wordpress.org publishes MD5 digests for core files
"crypto/sha256"
"encoding/hex"
"io"
"os"
"path/filepath"
"regexp"
"strconv"
"strings"
)
// PackageKind names the WordPress package family a path belongs to.
type PackageKind int
const (
KindNone PackageKind = iota
KindCore
KindPlugin
KindTheme
)
// Verdict is the outcome of locating a file inside a WordPress package and,
// when wordpress.org publishes checksums for that package version, comparing
// the file against them. The first four values describe the package; the
// last three are comparison results.
type Verdict int
const (
// VerdictUnknown: the path is not inside a recognised package.
VerdictUnknown Verdict = iota
// VerdictNoVersion: the package is recognised but its version header is
// not readable. During an unpack the header file may simply not exist
// yet, so a caller should describe the path again later.
VerdictNoVersion
// VerdictPending: the checksums for this package version are being
// fetched. A digest taken now can be verified once they arrive.
VerdictPending
// VerdictUnavailable: wordpress.org has no checksums for this package
// version. Themes, premium plugins and private forks land here.
VerdictUnavailable
// VerdictReady: the checksums are cached and Verify can compare a digest.
VerdictReady
VerdictVerified
// VerdictMismatch: the digest differs from the official file, or the
// official package has no file at this relative path.
VerdictMismatch
// VerdictUnverifiable: checksums are cached but the file could not be
// hashed (over the size cap, or not a regular file).
VerdictUnverifiable
)
func (v Verdict) String() string {
switch v {
case VerdictUnknown:
return "unknown"
case VerdictNoVersion:
return "no-version"
case VerdictPending:
return "pending"
case VerdictUnavailable:
return "unavailable"
case VerdictReady:
return "ready"
case VerdictVerified:
return "verified"
case VerdictMismatch:
return "mismatch"
case VerdictUnverifiable:
return "unverifiable"
}
return "verdict(" + strconv.Itoa(int(v)) + ")"
}
// Verification identifies one file inside a WordPress package. Root is the
// package root on disk, Rel the path below it as the checksum manifest keys
// it. Digest holds the hex digest of the file content in the algorithm
// wordpress.org publishes for the kind: SHA-256 for plugins, MD5 for core.
type Verification struct {
Verdict Verdict
Kind PackageKind
Root string
Slug string
Version string
Locale string
Rel string
Digest string
// Staged pins a retained digest to this release instead of retrying it
// against a later installed-core header. RootInfo identifies the staged
// directory observed while reading the header, before content analysis.
Staged bool
RootInfo os.FileInfo
}
var reThemeVersionHeader = regexp.MustCompile(`(?im)^[ \t/*#@]*Version:[ \t]*([^\s]+)`)
var reThemeNameHeader = regexp.MustCompile(`(?im)^[ \t/*#@]*Theme Name:[ \t]*[^ \t\r\n]`)
// Describe locates path inside a WordPress core, plugin or theme tree and
// reports whether official checksums for that package version are cached,
// being fetched, or do not exist. It reads package headers only, never the
// file at path, and starts a background fetch on a cache miss.
func (c *Cache) Describe(path string) Verification {
// The outer package owns every manifest path, including bundled plugins
// in a core ZIP and plugin files with core-like names or subdirectories.
stagedRoot, stagedSlug := detectStagedPluginRoot(path)
installedRoot, installedSlug := detectInstalledPluginRoot(path)
if installedRoot != "" && (stagedRoot == "" || len(installedRoot) < len(stagedRoot)) {
return c.describePluginRoot(installedRoot, installedSlug, path)
}
if root, slug := stagedRoot, stagedSlug; root != "" {
return c.describeStagedRoot(root, slug, path)
}
core := c.describeCore(path)
if core.Verdict != VerdictUnknown && core.Verdict != VerdictNoVersion {
return core
}
if theme := describeTheme(path); theme.Verdict != VerdictUnknown {
return theme
}
return core
}
func (c *Cache) describeStagedRoot(root, slug, path string) Verification {
before, beforeErr := os.Lstat(root)
var v Verification
if slug == "wordpress" {
v = c.describeCoreRoot(root, path, true)
} else {
v = c.describePluginRoot(root, slug, path)
}
v.Staged = true
after, afterErr := os.Lstat(root)
if beforeErr == nil && afterErr == nil && before.IsDir() && os.SameFile(before, after) {
v.RootInfo = before
} else {
// A header from a replacement tree cannot identify this event. Keep
// its kind and relative path so its digest can still be retained.
v.Version, v.Locale = "", ""
v.Verdict = VerdictNoVersion
}
return v
}
func (c *Cache) describeCore(path string) Verification {
root := DetectWPRoot(path)
return c.describeCoreRoot(root, path, false)
}
func (c *Cache) describeCoreRoot(root, path string, staged bool) Verification {
if root == "" {
return Verification{Verdict: VerdictUnknown}
}
rel := RelativePath(root, path)
if rel == "" {
return Verification{Verdict: VerdictUnknown}
}
v := Verification{Kind: KindCore, Root: root, Rel: rel}
if rel == filepath.Join("wp-includes", "version.php") {
c.invalidateRoot(root)
}
// A staging pathname can be reused while its next header is still
// missing. Only installed roots may reuse a cached version.
var version, locale string
var ok bool
if !staged {
version, locale, ok = c.getRoot(root)
}
if !ok {
var err error
version, locale, err = ReadVersionFile(root)
if err != nil {
v.Verdict = VerdictNoVersion
return v
}
if !staged {
c.setRoot(root, version, locale)
}
}
v.Version, v.Locale = version, locale
if c.hasChecksums(version, locale) {
v.Verdict = VerdictReady
return v
}
c.startBackgroundFetch(version, locale)
v.Verdict = VerdictPending
return v
}
func (c *Cache) describePluginRoot(root, slug, path string) Verification {
rel := RelativePath(root, path)
if rel == "" {
return Verification{Verdict: VerdictUnknown}
}
v := Verification{Kind: KindPlugin, Root: root, Slug: slug, Rel: rel}
version, err := ReadPluginVersion(root, slug)
if err != nil || version == "" {
// The same staged layout carries themes. A style.css with a theme
// header settles the kind, and a theme has no checksum source.
if theme := describeThemeRoot(root, slug, rel); theme.Verdict != VerdictUnknown {
return theme
}
v.Verdict = VerdictNoVersion
return v
}
v.Version = version
switch {
case c.hasPluginChecksums(slug, version):
v.Verdict = VerdictReady
case c.isPluginNotFound(slug, version):
v.Verdict = VerdictUnavailable
default:
c.startBackgroundPluginFetch(slug, version)
v.Verdict = VerdictPending
}
return v
}
// describeTheme handles the installed theme layout. Staged themes share the
// staged plugin layout and are recognised from describePluginRoot.
func describeTheme(path string) Verification {
const themesSegment = "/wp-content/themes/"
idx := strings.Index(path, themesSegment)
if idx < 0 {
return Verification{Verdict: VerdictUnknown}
}
slug, rel, ok := strings.Cut(path[idx+len(themesSegment):], "/")
if !ok || rel == "" || !safePluginPathComponent(slug) {
return Verification{Verdict: VerdictUnknown}
}
return describeThemeRoot(path[:idx+len(themesSegment)]+slug, slug, rel)
}
func describeThemeRoot(root, slug, rel string) Verification {
// #nosec G304 -- root is derived from a scanner-received path under a
// recognised theme or update-staging layout; the read is header-bounded.
f, err := os.Open(filepath.Join(root, "style.css"))
if err != nil {
return Verification{Verdict: VerdictUnknown}
}
defer func() { _ = f.Close() }()
buf, err := io.ReadAll(io.LimitReader(f, pluginHeaderReadLimit))
if err != nil || !reThemeNameHeader.Match(buf) {
return Verification{Verdict: VerdictUnknown}
}
v := Verification{Verdict: VerdictUnavailable, Kind: KindTheme, Root: root, Slug: slug, Rel: rel}
if m := reThemeVersionHeader.FindSubmatch(buf); m != nil {
v.Version = string(m[1])
}
return v
}
// Digest hashes the complete content behind fd with the algorithm
// wordpress.org publishes for kind. It returns "" when the file cannot be
// hashed whole, so a partial read can never verify as stock.
func (c *Cache) Digest(kind PackageKind, fd int) string {
switch kind {
case KindPlugin, KindCore:
default:
return ""
}
data := readCompleteFileForHash(fd)
if data == nil {
return ""
}
if kind == KindPlugin {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
// #nosec G401 -- MD5 is required here: wordpress.org ships MD5 digests
// as the canonical integrity reference for core files. We compare
// against their published values, not derive authority from the hash.
sum := md5.Sum(data)
return hex.EncodeToString(sum[:])
}
// Verify compares v.Digest against the official checksum for v. A
// description taken while the fetch was still running resolves here once
// the checksums are cached, so callers can hash at event time and compare
// later without holding the file open.
func (c *Cache) Verify(v Verification) Verdict {
switch v.Kind {
case KindPlugin:
return c.verifyPlugin(v)
case KindCore:
return c.verifyCore(v)
}
return v.Verdict
}
func (c *Cache) verifyPlugin(v Verification) Verdict {
if v.Version == "" {
return VerdictNoVersion
}
if !c.hasPluginChecksums(v.Slug, v.Version) {
if c.isPluginNotFound(v.Slug, v.Version) {
return VerdictUnavailable
}
c.startBackgroundPluginFetch(v.Slug, v.Version)
return VerdictPending
}
if v.Digest == "" {
return VerdictUnverifiable
}
expected, ok := c.lookupPluginChecksum(v.Slug, v.Version, v.Rel)
if !ok {
return VerdictMismatch
}
if hexDigestEqual(v.Digest, expected) {
return VerdictVerified
}
return VerdictMismatch
}
func (c *Cache) verifyCore(v Verification) Verdict {
if v.Version == "" {
return VerdictNoVersion
}
if !c.hasChecksums(v.Version, v.Locale) {
c.startBackgroundFetch(v.Version, v.Locale)
return VerdictPending
}
if v.Digest == "" {
return VerdictUnverifiable
}
if expected, ok := c.lookupChecksum(v.Version, v.Locale, v.Rel); ok && hexDigestEqual(v.Digest, expected) {
return VerdictVerified
}
if v.Staged {
return VerdictMismatch
}
// A core update rewrites version.php part-way through. A file that fails
// against the version cached for this root may be stock for the version
// the root now declares.
c.invalidateRoot(v.Root)
version, locale, err := ReadVersionFile(v.Root)
if err != nil || (version == v.Version && locale == v.Locale) {
return VerdictMismatch
}
c.setRoot(v.Root, version, locale)
if !c.hasChecksums(version, locale) {
c.startBackgroundFetch(version, locale)
return VerdictPending
}
if expected, ok := c.lookupChecksum(version, locale, v.Rel); ok && hexDigestEqual(v.Digest, expected) {
return VerdictVerified
}
return VerdictMismatch
}
// VerifyFile describes path, hashes fd when the package kind has a checksum
// source, and resolves the comparison when the checksums are already cached.
func (c *Cache) VerifyFile(fd int, path string) Verification {
v := c.Describe(path)
switch v.Verdict {
case VerdictNoVersion, VerdictPending, VerdictReady:
v.Digest = c.Digest(v.Kind, fd)
}
if v.Verdict == VerdictReady {
v.Verdict = c.Verify(v)
}
return v
}
func hexDigestEqual(digest, expected string) bool {
raw, err := hex.DecodeString(digest)
if err != nil {
return false
}
return constantTimeHexDigestEqual(raw, expected)
}
// resolve hashes fd when the checksums for v are cached and returns the
// comparison, or v's own verdict when there is nothing to compare against.
func (c *Cache) resolve(fd int, v Verification) Verdict {
if v.Verdict != VerdictReady {
return v.Verdict
}
v.Digest = c.Digest(v.Kind, fd)
return c.Verify(v)
}
// IsVerifiedCoreFile reports whether path is an unmodified file of the
// WordPress core version its install declares.
func (c *Cache) IsVerifiedCoreFile(fd int, path string) bool {
v := c.Describe(path)
return v.Kind == KindCore && c.resolve(fd, v) == VerdictVerified
}
// IsVerifiedPluginFile reports whether path is an unmodified file of the
// wordpress.org release its plugin header declares, installed or staged.
func (c *Cache) IsVerifiedPluginFile(fd int, path string) bool {
v := c.Describe(path)
return v.Kind == KindPlugin && c.resolve(fd, v) == VerdictVerified
}
package yara
import (
"encoding/hex"
"errors"
"fmt"
"sync/atomic"
"github.com/pidginhost/csm/internal/contenttype"
"github.com/pidginhost/csm/internal/yaraipc"
)
// Backend is the consumable scanning surface shared by the in-process
// *Scanner and out-of-process process supervisor. Callers should depend
// on this interface (via Active()) so they keep working when the daemon
// switches backends at startup. String-valued rule metadata travels on
// Match.Meta, so adapters that historically reached for the compiled
// *yara_x.Rules object (e.g. emailav) now work uniformly under both
// backends -- see internal/emailav/yarax.go.
type Backend interface {
ScanFile(path string, maxBytes int) []Match
ScanBytes(data []byte) []Match
RuleCount() int
Reload() error
}
// CheckedScanner is an optional capability on a Backend. ScanBytesChecked
// reports a scan failure (worker down, payload too large for the IPC frame, a
// transport error) distinctly from "no matches". Plain ScanBytes cannot: it
// returns nil for both a clean file and a failed scan, so a caller with a
// fail-closed policy -- email AV, finding re-check -- would auto-clear a file
// it never actually scanned. Both production backends implement it; the
// in-process stub returns a nil error because it cannot fail this way.
type CheckedScanner interface {
ScanBytesChecked(data []byte) ([]Match, error)
}
// FileScanResult binds path-scan matches to the exact bytes the backend
// scanned. ContentSHA256 is lowercase hex so it can be copied directly into a
// finding and transported over the worker IPC boundary.
type FileScanResult struct {
Matches []Match
ContentSHA256 string
}
// CheckedFileScanner is the path-based equivalent of CheckedScanner. It is
// used when a payload is too large for inline worker IPC and the worker must
// reopen the file without turning a worker or read failure into a clean scan.
type CheckedFileScanner interface {
ScanFileChecked(path string, maxBytes int) (FileScanResult, error)
}
// ScanBytesChecked scans data, read from the file called name, via b,
// surfacing a scan error when b supports the CheckedScanner capability.
// Backends without it fall back to the error-free ScanBytes. Callers that must
// fail closed on an unscannable payload should use this instead of
// Backend.ScanBytes.
//
// The archive policy lives here, at the backend-agnostic boundary, because it
// needs the name: raw archive bytes are not scannable and their stored entries
// trip rules with spurious tokens, so a file that is an archive by name and by
// magic is left to the extraction-time scan. Magic alone is not enough -- PHP
// executes past any leading bytes -- so an executable name is always scanned.
func ScanBytesChecked(b Backend, name string, data []byte) ([]Match, error) {
if b == nil {
return nil, errors.New("yara: backend unavailable")
}
if contenttype.IsArchiveFile(name, data) {
return nil, nil
}
if cs, ok := b.(CheckedScanner); ok {
return cs.ScanBytesChecked(data)
}
return b.ScanBytes(data), nil
}
// ScanFileChecked scans a path via b. Backends without the checked capability
// return an error because a caller cannot safely distinguish failure from a
// clean result through Backend.ScanFile.
func ScanFileChecked(b Backend, path string, maxBytes int) (FileScanResult, error) {
if b == nil {
return FileScanResult{}, errors.New("yara: backend unavailable")
}
if cs, ok := b.(CheckedFileScanner); ok {
result, err := cs.ScanFileChecked(path, maxBytes)
if err != nil {
return FileScanResult{}, err
}
if len(result.ContentSHA256) != 64 {
return FileScanResult{}, errors.New("yara: checked file scan returned an invalid content hash")
}
if _, err := hex.DecodeString(result.ContentSHA256); err != nil {
return FileScanResult{}, fmt.Errorf("yara: checked file scan returned an invalid content hash: %w", err)
}
return result, nil
}
return FileScanResult{}, errors.New("yara: backend does not support checked file scans")
}
var activeBackend atomic.Pointer[backendHolder]
type backendHolder struct{ b Backend }
// Active returns the configured scanning backend. When SetActive has
// not been called, it falls back to the in-process singleton Global().
// Returns a nil interface if neither is available (e.g. a !yara build
// with no supervisor wired up); callers must nil-check.
func Active() Backend {
if h := activeBackend.Load(); h != nil && h.b != nil {
return h.b
}
if g := Global(); g != nil {
return g
}
return nil
}
// SetActive installs a scanning backend. Calling with nil clears the
// override and restores the Global() fallback. Safe to call at any
// time; reads in-flight see the prior backend finish and the next read
// sees the new one.
func SetActive(b Backend) {
if b == nil {
activeBackend.Store(nil)
return
}
activeBackend.Store(&backendHolder{b: b})
}
// ScanContentOrPathChecked scans data inline and, when the payload cannot fit
// one IPC frame, retries by asking the backend to read path itself. The frame
// ceiling is a property of the transport, not of the file, so reporting it as a
// scan failure leaves a file unscanned that the worker could have read directly
// -- and lets an attacker put a payload out of reach by padding it past the
// inline ceiling.
//
// The returned digest is the SHA-256 of the bytes the backend actually scanned,
// and is empty when the inline scan succeeded: the caller already holds the
// digest of the snapshot it passed in. Only the oversize case retries; any
// other scan error surfaces unchanged so fail-closed callers stay fail-closed.
func ScanContentOrPathChecked(b Backend, path string, data []byte, maxBytes int) ([]Match, string, error) {
matches, err := ScanBytesChecked(b, path, data)
if !errors.Is(err, yaraipc.ErrPayloadTooLarge) {
return matches, "", err
}
res, pathErr := ScanFileChecked(b, path, maxBytes)
if pathErr != nil {
return nil, "", pathErr
}
return res.Matches, res.ContentSHA256, nil
}
package yara
import (
"fmt"
"os"
"sync"
)
var (
globalScanner *Scanner
globalOnce sync.Once
)
// Init initializes the global YARA-X scanner.
// Returns nil scanner if YARA-X is not compiled in or no rules found.
func Init(rulesDir string, disabled ...string) *Scanner {
if !Available() {
return nil
}
globalOnce.Do(func() {
s, err := NewScanner(rulesDir, disabled...)
if err != nil {
fmt.Fprintf(os.Stderr, "yara: init error: %v\n", err)
return
}
globalScanner = s
})
return globalScanner
}
// Global returns the global YARA-X scanner, or nil if not initialized.
func Global() *Scanner {
return globalScanner
}
package yara
import (
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"syscall"
)
// validateRulesDir refuses to compile YARA rules from a directory or
// file whose ownership or permissions would let a non-root non-self
// account drop a rule that disables detection. Mirrors the same trust
// rules CSM applies to /etc/csm/conf.d: only root or the running
// process may own the dir or any rule file, and group/world write bits
// are refused. A missing directory is a no-op so an operator who has
// not installed YARA rules yet does not see startup failures.
func validateRulesDir(dir string) error {
if dir == "" {
return nil
}
info, err := os.Stat(dir)
if err != nil {
if errors.Is(err, os.ErrNotExist) {
return nil
}
return fmt.Errorf("rules dir stat: %w", err)
}
if !info.IsDir() {
return fmt.Errorf("rules dir is not a directory: %s", dir)
}
if trustErr := checkYaraEntryTrust(dir, info); trustErr != nil {
return trustErr
}
entries, err := os.ReadDir(dir)
if err != nil {
return fmt.Errorf("rules dir read: %w", err)
}
for _, entry := range entries {
if entry.IsDir() {
continue
}
ext := strings.ToLower(filepath.Ext(entry.Name()))
if ext != ".yar" && ext != ".yara" {
continue
}
path := filepath.Join(dir, entry.Name())
fileInfo, err := os.Lstat(path)
if err != nil {
return fmt.Errorf("rule file stat %s: %w", path, err)
}
if fileInfo.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("rule file is a symlink: %s", path)
}
if !fileInfo.Mode().IsRegular() {
return fmt.Errorf("rule file is not a regular file: %s", path)
}
if trustErr := checkYaraEntryTrust(path, fileInfo); trustErr != nil {
return trustErr
}
}
return nil
}
// checkYaraEntryTrust enforces the perm + ownership trust rules on a
// single path. Used for both the rules dir itself and each rule file.
func checkYaraEntryTrust(path string, info os.FileInfo) error {
if mode := info.Mode().Perm(); mode&0022 != 0 {
return fmt.Errorf("rules path %s has unsafe mode %04o (group or world writable); set 0750 or stricter", path, mode)
}
sys, ok := info.Sys().(*syscall.Stat_t)
if !ok {
return nil
}
selfUID := uint32(os.Geteuid()) // #nosec G115 -- Linux uid_t is uint32; os.Geteuid returns a non-negative process uid.
if sys.Uid != 0 && sys.Uid != selfUID {
return fmt.Errorf("rules path %s owner uid=%d is neither root nor process uid=%d; refusing to load untrusted rules", path, sys.Uid, selfUID)
}
return nil
}
//go:build !yara
package yara
import "errors"
// Scanner is a no-op stub when YARA-X is not compiled in.
type Scanner struct{}
// Match represents a YARA rule that matched. Meta carries string-valued
// rule metadata pulled from `rule.Metadata()` (see yarax build); under
// this stub no rule ever matches, so Meta is never populated.
type Match struct {
RuleName string
Meta map[string]string
}
// NewScanner returns nil when YARA-X is not available.
func NewScanner(_ string, _ ...string) (*Scanner, error) {
return nil, nil
}
// DisabledRules returns nothing without YARA-X; no rules are compiled.
func (s *Scanner) DisabledRules() []string { return nil }
// DisabledRuleCount returns zero without YARA-X.
func (s *Scanner) DisabledRuleCount() int { return 0 }
// Reload is a no-op without YARA-X.
func (s *Scanner) Reload() error { return nil }
// ScanBytes returns nil without YARA-X.
func (s *Scanner) ScanBytes(_ []byte) []Match { return nil }
// ScanBytesChecked returns nil without YARA-X; there is no engine to fail.
func (s *Scanner) ScanBytesChecked(_ []byte) ([]Match, error) { return nil, nil }
// ScanFile returns nil without YARA-X.
func (s *Scanner) ScanFile(_ string, _ int) []Match { return nil }
// ScanFileChecked reports that no YARA engine is available. Checked callers
// must not treat an unscanned file as clean.
func (s *Scanner) ScanFileChecked(_ string, _ int) (FileScanResult, error) {
return FileScanResult{}, errors.New("yara: scanner unavailable in this build")
}
// RuleCount returns 0 without YARA-X.
func (s *Scanner) RuleCount() int { return 0 }
// GlobalRules returns nil without YARA-X (no compiled rules available).
func (s *Scanner) GlobalRules() interface{} { return nil }
// Available returns false (YARA-X is not compiled in).
func Available() bool {
return false
}
// TestCompile is a no-op when YARA-X is not compiled in.
func TestCompile(source string) error {
return nil
}
package yara
import "strings"
// YARA Forge aggregates community rulesets of mixed quality. Most are useful,
// but a rule whose condition ordinary code satisfies fires on every account on
// a shared host at once, and that volume hides real detections. Those rules are
// stripped both when a tier is downloaded and when rules are compiled, so a
// tier already on disk stops firing at the next reload instead of at the next
// weekly update.
//
// A rule belongs here only when its condition -- not merely its subject -- is
// unsound, and only when CSM detects the same technique itself. Suppressing a
// rule is never a substitute for detection: pair every entry with a rule in
// configs/malware.yar.
var suppressedRuleNames = []string{
// Detects HTML smuggling by requiring exactly one "payload marker", but the
// marker set mixes base64 file headers with the generic `.charCodeAt(i)^`
// XOR idiom. Any minified bundle that offers a client-side download and
// hashes a string satisfies it with nothing smuggled, which reported stock
// analytics plugins and cached page assets as malware delivery.
// Replaced by html_smuggling_payload, which requires the encoded header.
"ELCEEF_HTML_Smuggling_A",
// Detects a Windows IIS native-module family, but its condition is
// `native_module_private and 2 of ($i*) or 5 of them`. The trailing arm
// drops the native-module guard, so five generic strings decide the
// match -- and four of them are crawler user agents (Baiduspider,
// 360Spider, Sogou, YisouSpider) that any PHP bot-filter table lists
// alongside HTTP_X_FORWARDED_FOR. Stock analytics plugins matched.
// Replaced by seo_cloak_group14_ioc, which binds the campaign markers to
// native-module or campaign-infrastructure context and drops the user-agent
// arm.
"ESET_IIS_Group14",
}
// SuppressedRuleNames returns the built-in rule suppressions.
func SuppressedRuleNames() []string {
out := make([]string, len(suppressedRuleNames))
copy(out, suppressedRuleNames)
return out
}
// StripRules removes the named rules, bodies included, from YARA source text.
// Unknown names are ignored so a suppression list may name rules that a given
// tier does not ship.
func StripRules(content []byte, names []string) []byte {
filtered, _ := stripRules(content, names)
return filtered
}
// stripRules also counts removed declarations so a deliberately empty set
// remains distinguishable from a missing or empty rule file.
func stripRules(content []byte, names []string) ([]byte, int) {
if len(names) == 0 {
return content, 0
}
drop := make(map[string]bool, len(names))
for _, name := range names {
if name = strings.ToLower(strings.TrimSpace(name)); name != "" {
drop[name] = true
}
}
if len(drop) == 0 {
return content, 0
}
var result []byte
kept := 0
removed := 0
for _, rule := range sourceRules(content) {
if drop[strings.ToLower(rule.name)] {
removed++
result = append(result, content[kept:rule.start]...)
kept = rule.end
}
}
if kept == 0 {
return content, 0
}
return append(result, content[kept:]...), removed
}
type sourceRule struct {
name string
start, end int
}
// RuleNames returns declarations outside comments and literals. Validation
// uses the same boundaries as stripping, including modifiers and compact files.
func RuleNames(content []byte) []string {
var names []string
for _, rule := range sourceRules(content) {
names = append(names, rule.name)
}
return names
}
// sourceRules locates complete rule bodies without interpreting conditions.
// Counting source braces or removing whole lines can silently erase a neighbor:
// braces also occur in strings, comments, regexes and hex patterns.
func sourceRules(content []byte) []sourceRule {
var rules []sourceRule
start := -1
for pos := 0; pos < len(content); {
token, from, end := yaraToken(content, pos)
pos = end
if token == "private" || token == "global" {
if start < 0 {
start = from
}
continue
}
if token != "rule" {
start = -1
continue
}
if start < 0 {
start = from
}
name, _, next := yaraToken(content, pos)
pos = next
if !yaraIdentifier(name) {
break
}
token, _, pos = yaraToken(content, pos)
if token == ":" {
token, _, pos = yaraToken(content, pos)
if !yaraIdentifier(token) {
break
}
for yaraIdentifier(token) {
token, _, pos = yaraToken(content, pos)
}
}
if token != "{" {
break
}
depth := 1
for pos < len(content) && depth > 0 {
token, _, pos = yaraToken(content, pos)
switch token {
case "{":
depth++
case "}":
depth--
}
}
// Leave malformed source intact so compilation still reports it.
if depth != 0 {
break
}
rules = append(rules, sourceRule{name: name, start: start, end: pos})
start = -1
}
return rules
}
// yaraToken skips trivia and consumes quoted/regex literals as single tokens.
// YARA uses backslash for division, so slash always starts a regex or comment.
func yaraToken(src []byte, pos int) (token string, start, end int) {
for pos < len(src) {
switch src[pos] {
case ' ', '\t', '\r', '\n', '\f', '\v':
pos++
continue
case '/':
if pos+1 < len(src) && src[pos+1] == '/' {
pos += 2
for pos < len(src) && src[pos] != '\n' && src[pos] != '\r' {
pos++
}
continue
}
if pos+1 < len(src) && src[pos+1] == '*' {
pos += 2
for pos+1 < len(src) && (src[pos] != '*' || src[pos+1] != '/') {
pos++
}
pos = min(pos+2, len(src))
continue
}
}
break
}
start = pos
if pos == len(src) {
return "", pos, pos
}
ch := src[pos]
pos++
if ch == '"' || ch == '/' {
inClass := false
for pos < len(src) {
c := src[pos]
pos++
if c == '\\' {
pos = min(pos+1, len(src))
continue
}
if ch == '/' {
switch c {
case '[':
inClass = true
case ']':
inClass = false
}
}
if c == ch && !inClass {
break
}
}
return string(src[start:pos]), start, pos
}
if yaraIdent(ch) {
for pos < len(src) && yaraIdent(src[pos]) {
pos++
}
}
return string(src[start:pos]), start, pos
}
func yaraIdent(ch byte) bool {
return ch >= 'a' && ch <= 'z' || ch >= 'A' && ch <= 'Z' || ch >= '0' && ch <= '9' || ch == '_'
}
func yaraIdentifier(token string) bool {
// A following declaration must not become part of a malformed tag list.
if token == "rule" || token == "private" || token == "global" {
return false
}
if token == "" || token[0] >= '0' && token[0] <= '9' {
return false
}
for i := range len(token) {
if !yaraIdent(token[i]) {
return false
}
}
return true
}
// RuleNameFromLine returns the rule name declared on a source line, or "" when
// the line does not open a rule.
func RuleNameFromLine(line string) string {
s := strings.TrimPrefix(line, "private ")
if !strings.HasPrefix(s, "rule ") {
return ""
}
s = s[5:]
for i, ch := range s {
if ch == ' ' || ch == '\t' || ch == ':' || ch == '{' {
return s[:i]
}
}
return s
}
package yaraipc
import (
"errors"
"fmt"
"io"
"net"
"sync"
"syscall"
"time"
)
// ErrWorkerClosed means the worker hung up mid-request.
var ErrWorkerClosed = errors.New("yaraipc: worker connection closed")
// Dialer returns a fresh net.Conn to the worker. Decoupled from
// net.Dial so tests can substitute net.Pipe.
type Dialer func() (net.Conn, error)
// Client is a persistent-connection client. One in-flight request at a
// time: scanner callers do not need concurrency on a single socket and
// serialising simplifies failure semantics.
type Client struct {
mu sync.Mutex
conn net.Conn
dialer Dialer
timeout time.Duration
}
// NewClient constructs a Client that dials socketPath on demand.
func NewClient(socketPath string, timeout time.Duration) *Client {
return NewClientWithDialer(func() (net.Conn, error) {
return net.DialTimeout("unix", socketPath, timeout)
}, timeout)
}
// NewClientWithDialer is the test-friendly constructor.
func NewClientWithDialer(d Dialer, timeout time.Duration) *Client {
return &Client{dialer: d, timeout: timeout}
}
// Close drops the underlying connection if any. The Client stays usable;
// the next call dials again.
func (c *Client) Close() error {
c.mu.Lock()
defer c.mu.Unlock()
return c.dropLocked()
}
func (c *Client) ensureConnLocked() (net.Conn, error) {
if c.conn != nil {
return c.conn, nil
}
conn, err := c.dialer()
if err != nil {
return nil, fmt.Errorf("yaraipc: dial: %w", err)
}
c.conn = conn
return conn, nil
}
func (c *Client) dropLocked() error {
if c.conn == nil {
return nil
}
err := c.conn.Close()
c.conn = nil
return err
}
// roundTrip sends req and returns the response frame. A closed socket is
// retried once because the server may intentionally close idle
// connections after their message budget is spent.
func (c *Client) roundTrip(req Frame) (Frame, error) {
c.mu.Lock()
defer c.mu.Unlock()
resp, err := c.roundTripLocked(req)
if err == nil || !isRetryableClosedConn(err) {
return resp, err
}
return c.roundTripLocked(req)
}
func (c *Client) roundTripLocked(req Frame) (Frame, error) {
conn, err := c.ensureConnLocked()
if err != nil {
return Frame{}, err
}
if c.timeout > 0 {
_ = conn.SetDeadline(time.Now().Add(c.timeout))
}
if werr := WriteFrame(conn, req); werr != nil {
_ = c.dropLocked()
return Frame{}, fmt.Errorf("yaraipc: write: %w", werr)
}
resp, rerr := ReadFrame(conn)
if rerr != nil {
_ = c.dropLocked()
if errors.Is(rerr, io.EOF) {
return Frame{}, ErrWorkerClosed
}
return Frame{}, fmt.Errorf("yaraipc: read: %w", rerr)
}
if resp.Error != "" {
return Frame{}, fmt.Errorf("yaraipc: worker: %s", resp.Error)
}
return resp, nil
}
func isRetryableClosedConn(err error) bool {
return errors.Is(err, ErrWorkerClosed) ||
errors.Is(err, io.EOF) ||
errors.Is(err, io.ErrUnexpectedEOF) ||
errors.Is(err, net.ErrClosed) ||
errors.Is(err, syscall.EPIPE) ||
errors.Is(err, syscall.ECONNRESET) ||
errors.Is(err, syscall.ECONNABORTED)
}
// ScanFile is the daemon-side shim for OpScanFile.
func (c *Client) ScanFile(args ScanFileArgs) (ScanResult, error) {
req, err := EncodePayload(OpScanFile, args)
if err != nil {
return ScanResult{}, err
}
resp, err := c.roundTrip(req)
if err != nil {
return ScanResult{}, err
}
var res ScanResult
if len(resp.Payload) > 0 {
if err := DecodePayload(resp, &res); err != nil {
return ScanResult{}, err
}
}
return res, nil
}
// ScanBytes is the daemon-side shim for OpScanBytes.
func (c *Client) ScanBytes(args ScanBytesArgs) (ScanResult, error) {
// Reject up front so an oversize buffer becomes a legible, typed error
// rather than a generic WriteFrame failure after a wasteful multi-MiB
// marshal. Callers must treat this as "could not scan", never as clean.
if len(args.Data) > MaxScanBytes {
return ScanResult{}, fmt.Errorf("%w (%d > %d bytes)", ErrPayloadTooLarge, len(args.Data), MaxScanBytes)
}
req, err := EncodePayload(OpScanBytes, args)
if err != nil {
return ScanResult{}, err
}
resp, err := c.roundTrip(req)
if err != nil {
return ScanResult{}, err
}
var res ScanResult
if len(resp.Payload) > 0 {
if err := DecodePayload(resp, &res); err != nil {
return ScanResult{}, err
}
}
return res, nil
}
// Reload is the daemon-side shim for OpReload.
func (c *Client) Reload(args ReloadArgs) (ReloadResult, error) {
req, err := EncodePayload(OpReload, args)
if err != nil {
return ReloadResult{}, err
}
resp, err := c.roundTrip(req)
if err != nil {
return ReloadResult{}, err
}
var res ReloadResult
if len(resp.Payload) > 0 {
if err := DecodePayload(resp, &res); err != nil {
return ReloadResult{}, err
}
}
return res, nil
}
// Ping is the daemon-side shim for OpPing.
func (c *Client) Ping() (PingResult, error) {
req, err := EncodePayload(OpPing, nil)
if err != nil {
return PingResult{}, err
}
resp, err := c.roundTrip(req)
if err != nil {
return PingResult{}, err
}
var res PingResult
if len(resp.Payload) > 0 {
if err := DecodePayload(resp, &res); err != nil {
return PingResult{}, err
}
}
return res, nil
}
// Package yaraipc defines the wire protocol spoken between the CSM daemon
// and the supervised `csm yara-worker` child process. The worker exists to
// isolate the YARA-X cgo surface; a crash in the worker must not take the
// daemon down. See ROADMAP.md item 2 for the decision record.
//
// The protocol is length-prefixed JSON frames on a Unix-domain socket.
// Connections are persistent: the daemon opens one, streams scan and
// reload requests, and reconnects if the worker dies.
package yaraipc
import (
"encoding/binary"
"encoding/json"
"errors"
"fmt"
"io"
)
// MaxFrameBytes caps a single request or response payload. Sized to cover
// a file sent inline via OpScanBytes up to the scanner's usual 8 MiB read
// ceiling, with headroom for JSON base64 expansion.
const MaxFrameBytes = 16 << 20
// MaxScanBytes is the largest raw payload OpScanBytes can carry in one frame.
// A []byte is base64-encoded in JSON (4/3 expansion) plus a small envelope, so
// the raw ceiling sits well below MaxFrameBytes. Callers reject an oversize
// buffer up front with ErrPayloadTooLarge instead of marshalling a multi-MiB
// frame only to have WriteFrame fail -- and, crucially, an oversize payload
// must surface as an error, never be silently treated as a clean scan.
const MaxScanBytes = (MaxFrameBytes - 256) * 3 / 4
// ErrPayloadTooLarge is returned by the client when an inline OpScanBytes
// payload would not fit in a single protocol frame.
var ErrPayloadTooLarge = errors.New("yaraipc: scan payload exceeds max inline size")
// Op selects the handler on the worker side. Strings (not iota ints) so
// adding a new op is additive and mismatched client/worker versions fail
// with a recognisable "unknown op" error instead of silently dispatching
// the wrong handler.
const (
OpScanFile = "scan_file"
OpScanBytes = "scan_bytes"
OpReload = "reload"
OpPing = "ping"
)
// Frame is the envelope. Request frames carry an Op and typed args in
// Payload; response frames carry the typed result (or an Error) in
// Payload and leave Op empty.
type Frame struct {
Op string `json:"op,omitempty"`
Payload json.RawMessage `json:"payload,omitempty"`
Error string `json:"error,omitempty"`
}
// ScanFileArgs asks the worker to read and scan a file by path. MaxBytes
// bounds the read so the daemon cannot make the worker allocate more than
// it agreed to.
type ScanFileArgs struct {
Path string `json:"path"`
MaxBytes int `json:"max_bytes"`
}
// ScanBytesArgs carries file content inline. Used when the caller already
// has the bytes (fanotify buffered reads) and avoids a second file open in
// the worker.
type ScanBytesArgs struct {
Data []byte `json:"data"`
}
// ReloadArgs triggers a rule recompile. RulesDir is optional; if empty
// the worker reuses the directory it was started with.
type ReloadArgs struct {
RulesDir string `json:"rules_dir,omitempty"`
}
// Match mirrors yara.Match but is part of this package's public wire
// contract so the daemon does not need to import the yara package just to
// speak to its worker.
//
// Meta carries string-valued rule metadata (identifier -> value) pulled
// from yara_x `rule.Metadata()` inside the worker, where the compiled
// rules live. Non-string metadata (int / float / bool / bytes) is
// dropped: wiring only string values is a deliberate policy, not a
// fidelity claim. Consumers that need a specific key document their own
// default; e.g. emailav maps a missing "severity" entry to "high".
// Omitted from the wire when empty so the per-scan payload cost is zero
// for the common clean-file case.
type Match struct {
RuleName string `json:"rule"`
Meta map[string]string `json:"meta,omitempty"`
}
// ScanResult is returned for OpScanFile and OpScanBytes. ContentSHA256 is set
// for path scans so callers bind findings to the bytes scanned by the worker.
type ScanResult struct {
Matches []Match `json:"matches,omitempty"`
ContentSHA256 string `json:"content_sha256,omitempty"`
}
// ReloadResult is returned for OpReload. CompileError is non-empty when the
// worker is up but its rules failed to compile, so the daemon can tell a
// successful reload (RuleCount>0) from a still-broken rule set instead of
// treating a silent no-op as success.
type ReloadResult struct {
RuleCount int `json:"rule_count"`
CompileError string `json:"compile_error,omitempty"`
}
// PingResult is returned for OpPing. Used by the supervisor's liveness
// check and as the first frame after a reconnect to confirm the worker is
// past its rule-compile step before real scan traffic begins.
//
// CompileError is non-empty when the worker process is alive but its rule
// compile failed at startup and has not yet been recovered by a reload. It
// disambiguates "0 rules because the compile broke" (needs an operator alert)
// from "0 rules because there is no engine / empty rules dir".
type PingResult struct {
Alive bool `json:"alive"`
RuleCount int `json:"rule_count"`
CompileError string `json:"compile_error,omitempty"`
}
// WriteFrame writes a 4-byte big-endian length prefix followed by the
// JSON-encoded frame. The caller owns any deadline on the underlying
// writer.
func WriteFrame(w io.Writer, f Frame) error {
body, err := json.Marshal(f)
if err != nil {
return fmt.Errorf("marshal frame: %w", err)
}
if len(body) > MaxFrameBytes {
return fmt.Errorf("frame body %d bytes exceeds cap %d", len(body), MaxFrameBytes)
}
var hdr [4]byte
// #nosec G115 -- len(body) is bounded above by MaxFrameBytes (16 MiB), which fits in uint32.
binary.BigEndian.PutUint32(hdr[:], uint32(len(body)))
if _, err := w.Write(hdr[:]); err != nil {
return err
}
if _, err := w.Write(body); err != nil {
return err
}
return nil
}
// ReadFrame reads one length-prefixed JSON frame from r. Frames larger
// than MaxFrameBytes are rejected before the body is read so a hostile or
// corrupt peer cannot make us allocate an unbounded buffer.
func ReadFrame(r io.Reader) (Frame, error) {
var hdr [4]byte
if _, err := io.ReadFull(r, hdr[:]); err != nil {
return Frame{}, err
}
n := binary.BigEndian.Uint32(hdr[:])
if n == 0 {
return Frame{}, errors.New("yaraipc: zero-length frame")
}
if n > MaxFrameBytes {
return Frame{}, fmt.Errorf("yaraipc: frame length %d exceeds cap %d", n, MaxFrameBytes)
}
body := make([]byte, n)
if _, err := io.ReadFull(r, body); err != nil {
return Frame{}, err
}
var f Frame
if err := json.Unmarshal(body, &f); err != nil {
return Frame{}, fmt.Errorf("yaraipc: unmarshal frame: %w", err)
}
return f, nil
}
// DecodePayload unmarshals f.Payload into out. Kept as a helper because
// every handler does it and a typo in the json tag on one side is the
// kind of bug that hides until production.
func DecodePayload(f Frame, out any) error {
if len(f.Payload) == 0 {
return errors.New("yaraipc: empty payload")
}
return json.Unmarshal(f.Payload, out)
}
// EncodePayload marshals v and returns a Frame with Op set. Convenience
// for the client side.
func EncodePayload(op string, v any) (Frame, error) {
if v == nil {
return Frame{Op: op}, nil
}
raw, err := json.Marshal(v)
if err != nil {
return Frame{}, fmt.Errorf("yaraipc: marshal payload: %w", err)
}
return Frame{Op: op, Payload: raw}, nil
}
package yaraipc
import (
"context"
"errors"
"fmt"
"io"
"net"
"sync"
"github.com/pidginhost/csm/internal/obs"
)
// Handler is the worker-side interface. The production implementation
// wraps internal/yara; tests supply fakes to drive the Serve loop.
type Handler interface {
ScanFile(ScanFileArgs) (ScanResult, error)
ScanBytes(ScanBytesArgs) (ScanResult, error)
Reload(ReloadArgs) (ReloadResult, error)
Ping() (PingResult, error)
}
// ServeOptions tunes Serve's behaviour. ErrorLog is called for per-frame
// decode or transport errors. Nil is fine; these errors are not fatal to
// the worker process.
type ServeOptions struct {
ErrorLog func(error)
// MaxMessagesPerConn caps how many frames a single connection may
// serve before the server force-closes it. Zero applies the
// default; a negative value disables the cap. Prevents a malicious
// or runaway peer from holding one socket and looping millions of
// frames through a single goroutine.
MaxMessagesPerConn int
}
// defaultMaxMessagesPerConn is the per-connection frame budget when
// the operator does not configure one. Production daemon-to-worker
// traffic burns a handful of frames per scan; 1e6 is several orders of
// magnitude above the legitimate ceiling and still catches a wedged
// peer in finite time.
const defaultMaxMessagesPerConn = 1_000_000
// Serve accepts connections on ln and dispatches frames to h until ctx
// is cancelled or ln returns an error. Per-connection goroutines are
// spawned; each handles its connection serially (single in-flight
// request), which matches the daemon-side client and keeps failure
// semantics simple.
//
// On ctx cancellation Serve closes the listener and any active
// connections. Without this, clients holding a cached conn would keep
// talking to a zombie handler instead of seeing EOF and reconnecting
// to whatever replaces the worker. In production that "zombie" is a
// crashed process (the kernel closes its sockets), but in-process
// tests and a graceful SIGTERM shutdown both need the explicit close.
func Serve(ctx context.Context, ln net.Listener, h Handler, opts ServeOptions) error {
var (
mu sync.Mutex
active = map[net.Conn]struct{}{}
closed bool
servers sync.WaitGroup
)
closeAll := func() {
mu.Lock()
if closed {
mu.Unlock()
return
}
closed = true
conns := make([]net.Conn, 0, len(active))
for c := range active {
conns = append(conns, c)
}
mu.Unlock()
for _, c := range conns {
_ = c.Close()
}
}
obs.SafeGo("yaraipc-ctx", func() {
<-ctx.Done()
_ = ln.Close()
closeAll()
})
for {
conn, err := ln.Accept()
if err != nil {
if ctx.Err() != nil {
servers.Wait()
// Accept error after ctx cancellation is the expected
// shutdown path (ln.Close from the watcher goroutine);
// swallowing it here keeps Serve callers from needing to
// distinguish "clean stop" from "real failure".
return nil //nolint:nilerr
}
closeAll()
servers.Wait()
return fmt.Errorf("yaraipc: accept: %w", err)
}
mu.Lock()
if closed {
mu.Unlock()
_ = conn.Close()
continue
}
active[conn] = struct{}{}
mu.Unlock()
servers.Add(1)
c := conn
obs.SafeGo("yaraipc-conn", func() {
defer servers.Done()
serveConn(c, h, opts)
mu.Lock()
delete(active, c)
mu.Unlock()
})
}
}
func serveConn(conn net.Conn, h Handler, opts ServeOptions) {
defer func() { _ = conn.Close() }()
budget := opts.MaxMessagesPerConn
if budget == 0 {
budget = defaultMaxMessagesPerConn
}
served := 0
for {
if budget > 0 && served >= budget {
if opts.ErrorLog != nil {
opts.ErrorLog(fmt.Errorf("max messages per connection (%d) reached; closing", budget))
}
return
}
req, err := ReadFrame(conn)
if err != nil {
if errors.Is(err, io.EOF) {
return
}
if opts.ErrorLog != nil {
opts.ErrorLog(fmt.Errorf("read: %w", err))
}
return
}
resp := safeDispatch(req, h, opts)
if err := WriteFrame(conn, resp); err != nil {
if opts.ErrorLog != nil {
opts.ErrorLog(fmt.Errorf("write: %w", err))
}
return
}
served++
}
}
// safeDispatch runs dispatch and converts a handler panic into an error
// response frame. Without it a panic (malformed rules, a pathological scan
// input, a bad arg) would unwind serveConn with no frame written: the client
// blocks to its deadline, retries, and panics again, while the worker process
// stays up so the supervisor never restarts it -- that scan never completes.
func safeDispatch(req Frame, h Handler, opts ServeOptions) (resp Frame) {
defer func() {
if r := recover(); r != nil {
if opts.ErrorLog != nil {
opts.ErrorLog(fmt.Errorf("handler panic on op %q: %v", req.Op, r))
}
resp = Frame{Error: fmt.Sprintf("handler panic: %v", r)}
}
}()
return dispatch(req, h)
}
func dispatch(req Frame, h Handler) Frame {
switch req.Op {
case OpScanFile:
var args ScanFileArgs
if err := DecodePayload(req, &args); err != nil {
return Frame{Error: fmt.Sprintf("decode scan_file: %v", err)}
}
res, err := h.ScanFile(args)
return responseFrame(res, err)
case OpScanBytes:
var args ScanBytesArgs
if err := DecodePayload(req, &args); err != nil {
return Frame{Error: fmt.Sprintf("decode scan_bytes: %v", err)}
}
res, err := h.ScanBytes(args)
return responseFrame(res, err)
case OpReload:
// Reload payload is optional; an empty frame means reuse the
// worker's startup RulesDir.
var args ReloadArgs
if len(req.Payload) > 0 {
if err := DecodePayload(req, &args); err != nil {
return Frame{Error: fmt.Sprintf("decode reload: %v", err)}
}
}
res, err := h.Reload(args)
return responseFrame(res, err)
case OpPing:
res, err := h.Ping()
return responseFrame(res, err)
default:
return Frame{Error: fmt.Sprintf("yaraipc: unknown op %q", req.Op)}
}
}
func responseFrame(result any, err error) Frame {
if err != nil {
return Frame{Error: err.Error()}
}
f, encErr := EncodePayload("", result)
if encErr != nil {
return Frame{Error: encErr.Error()}
}
return f
}
// Package yaraworker implements the `csm yara-worker` subcommand: a
// child process that exists only to host the YARA-X cgo surface and
// reply to scan requests over a Unix socket. See ROADMAP.md item 2.
//
// The handler here adapts a yara.Scanner (real in `-tags yara` builds,
// no-op in plain builds) to the yaraipc.Handler wire contract. The
// package is deliberately thin: IPC lives in internal/yaraipc, rule
// compilation lives in internal/yara, and supervision lives in
// internal/daemon.
package yaraworker
import (
"fmt"
"sync"
"github.com/pidginhost/csm/internal/yara"
"github.com/pidginhost/csm/internal/yaraipc"
)
// Scanner is the subset of *yara.Scanner that the handler uses. An
// interface (rather than the concrete type) so tests can inject a fake
// without pulling in the cgo build tag.
type Scanner interface {
ScanFile(path string, maxBytes int) []yara.Match
ScanFileChecked(path string, maxBytes int) (yara.FileScanResult, error)
ScanBytes(data []byte) []yara.Match
ScanBytesChecked(data []byte) ([]yara.Match, error)
Reload() error
RuleCount() int
}
// NewHandler returns a yaraipc.Handler backed by s. A nil scanner is
// permitted: the handler reports Alive=true with zero matches, which is
// the expected behaviour on builds compiled without the yara tag and on
// hosts where no rules directory has been provisioned yet. This constructor
// has no recovery factory, so a nil scanner stays nil (used by tests and the
// no-engine path).
func NewHandler(s Scanner) yaraipc.Handler {
return &handler{scanner: s}
}
// newRecoverableHandler adds a rebuild factory and a startup compile error, so
// a worker that came up with a failed rule compile (scanner == nil,
// compileErr != "") can recover on a later Reload instead of no-op'ing
// forever. rebuild returns a fresh scanner from the current rules on disk, an
// error if they still do not compile, or (nil, nil) when there is no engine to
// build (plain build).
func newRecoverableHandler(s Scanner, rebuild func() (Scanner, error), compileErr string) yaraipc.Handler {
return &handler{scanner: s, rebuild: rebuild, compileErr: compileErr}
}
type handler struct {
// mu guards scanner and compileErr: the wire contract allows more than
// one connection, and a Reload that swaps the scanner in must not race a
// concurrent Ping/Scan reading it.
mu sync.Mutex
scanner Scanner
rebuild func() (Scanner, error)
compileErr string
}
func (h *handler) currentState() (Scanner, string) {
h.mu.Lock()
defer h.mu.Unlock()
return h.scanner, h.compileErr
}
func (h *handler) ScanFile(a yaraipc.ScanFileArgs) (yaraipc.ScanResult, error) {
sc, compileErr := h.currentState()
if sc == nil {
if compileErr != "" {
return yaraipc.ScanResult{}, fmt.Errorf("yara scanner unavailable: %s", compileErr)
}
return yaraipc.ScanResult{}, nil
}
result, err := sc.ScanFileChecked(a.Path, a.MaxBytes)
if err != nil {
return yaraipc.ScanResult{}, err
}
return yaraipc.ScanResult{
Matches: convertMatches(result.Matches),
ContentSHA256: result.ContentSHA256,
}, nil
}
func (h *handler) ScanBytes(a yaraipc.ScanBytesArgs) (yaraipc.ScanResult, error) {
sc, compileErr := h.currentState()
if sc == nil {
if compileErr != "" {
return yaraipc.ScanResult{}, fmt.Errorf("yara scanner unavailable: %s", compileErr)
}
return yaraipc.ScanResult{}, nil
}
matches, err := sc.ScanBytesChecked(a.Data)
if err != nil {
return yaraipc.ScanResult{}, err
}
return yaraipc.ScanResult{Matches: convertMatches(matches)}, nil
}
func (h *handler) Reload(_ yaraipc.ReloadArgs) (yaraipc.ReloadResult, error) {
h.mu.Lock()
defer h.mu.Unlock()
if h.scanner == nil {
// No live scanner. Either there is no engine (no rebuild factory ->
// no-op, matching the pre-recovery behaviour) or the startup compile
// failed and we now retry it.
if h.rebuild == nil {
return yaraipc.ReloadResult{}, nil
}
newSc, err := h.rebuild()
if err != nil {
h.compileErr = err.Error()
return yaraipc.ReloadResult{CompileError: h.compileErr}, err
}
if newSc == nil {
// No engine to build (plain build): stay a no-op, not an error.
h.compileErr = ""
return yaraipc.ReloadResult{}, nil
}
h.scanner = newSc
h.compileErr = ""
return yaraipc.ReloadResult{RuleCount: newSc.RuleCount()}, nil
}
if err := h.scanner.Reload(); err != nil {
return yaraipc.ReloadResult{}, err
}
return yaraipc.ReloadResult{RuleCount: h.scanner.RuleCount()}, nil
}
func (h *handler) Ping() (yaraipc.PingResult, error) {
h.mu.Lock()
defer h.mu.Unlock()
if h.scanner == nil {
return yaraipc.PingResult{Alive: true, CompileError: h.compileErr}, nil
}
return yaraipc.PingResult{Alive: true, RuleCount: h.scanner.RuleCount()}, nil
}
func convertMatches(in []yara.Match) []yaraipc.Match {
if len(in) == 0 {
return nil
}
out := make([]yaraipc.Match, len(in))
for i := range in {
out[i] = yaraipc.Match{
RuleName: in[i].RuleName,
Meta: in[i].Meta,
}
}
return out
}
package yaraworker
import (
"context"
"fmt"
"net"
"os"
"path/filepath"
"syscall"
"github.com/pidginhost/csm/internal/yara"
"github.com/pidginhost/csm/internal/yaraipc"
)
// Config is what the `csm yara-worker` subcommand receives from its
// parent process. SocketPath and RulesDir are mandatory; ErrorLog is
// optional.
type Config struct {
SocketPath string
RulesDir string
// DisabledRules are the operator's signatures.disabled_rules names.
// The worker compiles the same rules the daemon would, so it has to
// honour the same list or a rule switched off in config keeps firing
// through the worker.
DisabledRules []string
ErrorLog func(error)
}
// Run is the entrypoint the `csm yara-worker` subcommand calls. It
// binds the Unix socket, compiles rules, and serves until ctx is
// cancelled or an unrecoverable accept error occurs.
//
// Rule-compile failures at startup are not fatal: the worker still
// serves and returns zero matches, so the supervisor can observe the
// condition via Ping and the next OpReload can recover. A fatal failure
// here (bad socket path, permission denied on bind, stale socket that
// cannot be removed) is returned so systemd sees a non-zero exit and
// the supervisor escalates through its backoff.
// listenPrivateUnix binds path under a 0077 umask so the socket is 0600 from
// the moment it exists. Creating it with the process umask and chmodding
// afterwards left a window in which any local user could connect, and a
// connection made in that window survives the chmod because permissions are
// checked at connect time. The chmod stays as a belt for filesystems that
// ignore the umask on socket creation.
func listenPrivateUnix(path string, listen func(string) (net.Listener, error)) (net.Listener, error) {
old := syscall.Umask(0o077)
ln, err := listen(path)
syscall.Umask(old)
if err != nil {
return nil, err
}
if err := os.Chmod(path, 0o600); err != nil {
_ = ln.Close()
return nil, fmt.Errorf("chmod socket: %w", err)
}
return ln, nil
}
func Run(ctx context.Context, cfg Config) error {
if cfg.SocketPath == "" {
return fmt.Errorf("yaraworker: socket path is empty")
}
if err := os.MkdirAll(filepath.Dir(cfg.SocketPath), 0o700); err != nil {
return fmt.Errorf("yaraworker: mkdir socket dir: %w", err)
}
// A stale socket file from a previous worker crash blocks bind.
// The supervisor only starts one worker at a time, so there is no
// concurrent binder to race with.
if err := os.Remove(cfg.SocketPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("yaraworker: removing stale socket: %w", err)
}
ln, err := listenPrivateUnix(cfg.SocketPath, func(p string) (net.Listener, error) {
return net.Listen("unix", p)
})
if err != nil {
return fmt.Errorf("yaraworker: listen: %w", err)
}
scanner, compileErr := yara.NewScanner(cfg.RulesDir, cfg.DisabledRules...)
compileErrStr := ""
if compileErr != nil {
compileErrStr = compileErr.Error()
if cfg.ErrorLog != nil {
cfg.ErrorLog(fmt.Errorf("yaraworker: scanner init: %w", compileErr))
}
}
// A failed startup compile is not permanent: give the handler a factory so
// a later OpReload (forge update, SIGHUP) can rebuild from the fixed rules
// on disk instead of the worker staying silently dead until it crashes.
rulesDir := cfg.RulesDir
disabledRules := cfg.DisabledRules
rebuild := func() (Scanner, error) {
s, err := yara.NewScanner(rulesDir, disabledRules...)
if err != nil {
return nil, err
}
if s == nil {
// !yara build: no engine to build.
return nil, nil
}
return s, nil
}
h := newRecoverableHandler(asScanner(scanner), rebuild, compileErrStr)
return yaraipc.Serve(ctx, ln, h, yaraipc.ServeOptions{ErrorLog: cfg.ErrorLog})
}
// asScanner converts a possibly-nil *yara.Scanner into the Scanner interface
// without producing a typed-nil-inside-interface. `*yara.Scanner` is nil in
// two cases: (1) !yara builds always return (nil, nil), and (2) yara-build
// NewScanner returned (nil, err). In both we want an untyped-nil interface so
// the handler's nil fast path fires instead of a method call on a nil pointer.
func asScanner(s *yara.Scanner) Scanner {
if s == nil {
return nil
}
return s
}
package yaraworker
import (
"context"
"encoding/json"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"sync"
"sync/atomic"
"syscall"
"time"
"github.com/pidginhost/csm/internal/obs"
"github.com/pidginhost/csm/internal/yara"
"github.com/pidginhost/csm/internal/yaraipc"
)
// SupervisorConfig parameterises the daemon-side lifecycle manager for
// the `csm yara-worker` child process.
//
// Restart backoff policy:
//
// - The first crash triggers a restart after MinRestartInterval.
// - Each consecutive crash doubles the delay up to MaxRestartInterval.
// - A worker that stays up for StableDuration resets the backoff to
// MinRestartInterval so a single bad rule deploy is not punished
// forever.
type SupervisorConfig struct {
BinaryPath string
SocketPath string
RulesDir string
ConfigFile string
ConfigDir string
// Carry the effective list across crashes, even if config on disk changed
// while the daemon is waiting for a restart to apply those changes.
DisabledRules []string
StartTimeout time.Duration
MinRestartInterval time.Duration
MaxRestartInterval time.Duration
StableDuration time.Duration
// ExtraArgs appended to `csm yara-worker`. Tests use this to flag
// the helper process into mock-worker mode.
ExtraArgs []string
// Env override; nil means inherit os.Environ.
Env []string
// OnRestart is called after each unplanned worker exit. Exit code
// is the process exit status (or -1 if the process was killed by a
// signal whose number is signal). Daemons wire this to a finding
// emitter.
OnRestart func(exitCode int, signal syscall.Signal, runDuration time.Duration)
// OnStable is called each time a worker has stayed up for StableDuration
// after becoming ready. A restart that passes its readiness probe and
// dies again soon after never reports stable.
// Like OnRestart, it must return promptly and must not call Stop, which
// waits for callbacks to finish.
OnStable func()
// Logf is an optional structured-log hook. Supervisor internals log
// restarts + transient errors here. Nil is fine.
Logf func(format string, args ...any)
// ClientTimeout is the per-call read/write deadline the supervisor
// imposes on the worker. Scan calls inherit this.
ClientTimeout time.Duration
}
// Supervisor manages the `csm yara-worker` child process and exposes a
// Scanner-shaped surface to the rest of the daemon. One supervisor per
// daemon.
type Supervisor struct {
cfg SupervisorConfig
mu sync.Mutex
cmd *exec.Cmd
client *yaraipc.Client
started time.Time
stopped bool
running atomic.Bool
// callbackMu serializes OnRestart and OnStable. waitForChild clears cmd
// before OnRestart runs, so a stable check that loses the race to a crash
// sees a different child and cannot report the dead worker as stable.
callbackMu sync.Mutex
ctx context.Context
cancel context.CancelFunc
done chan struct{}
// Restart counters for observability / tests. Under mu.
restartCount int
lastExitCode int
lastExitSignal syscall.Signal
// Scan-failure log suppression. A scanner that cannot load its rules
// fails identically for every buffer, so the same message would
// otherwise be written once per scanned file.
scanErrMu sync.Mutex
scanErrSeen map[string]*scanErrRecord
scanErrOverflow scanErrRecord
}
// scanErrRecord tracks one distinct failure message: when it was last
// written, and how many identical failures happened since.
type scanErrRecord struct {
at time.Time
suppressed int
}
// scanErrMaxTracked bounds the distinct-message table so a failure carrying
// unique text (a path, an offset) cannot grow it without limit.
const scanErrMaxTracked = 32
// scanErrLogWindow bounds how often one recurring scan failure is written.
// Long enough that a broken rules directory cannot flood the journal, short
// enough that a persistent fault keeps reappearing.
const scanErrLogWindow = time.Minute
// NewSupervisor validates cfg and returns an unstarted supervisor.
// Defaults: StartTimeout 10s, MinRestartInterval 1s, MaxRestartInterval
// 60s, StableDuration 30s, ClientTimeout 30s.
func NewSupervisor(cfg SupervisorConfig) (*Supervisor, error) {
if cfg.BinaryPath == "" {
return nil, errors.New("yaraworker: BinaryPath is required")
}
if cfg.SocketPath == "" {
return nil, errors.New("yaraworker: SocketPath is required")
}
if cfg.StartTimeout == 0 {
cfg.StartTimeout = 10 * time.Second
}
if cfg.MinRestartInterval == 0 {
cfg.MinRestartInterval = time.Second
}
if cfg.MaxRestartInterval == 0 {
cfg.MaxRestartInterval = 60 * time.Second
}
if cfg.StableDuration == 0 {
cfg.StableDuration = 30 * time.Second
}
if cfg.ClientTimeout == 0 {
cfg.ClientTimeout = 30 * time.Second
}
cfg.DisabledRules = append([]string(nil), cfg.DisabledRules...)
return &Supervisor{cfg: cfg}, nil
}
// Start launches the worker and blocks until the first Ping succeeds or
// StartTimeout elapses. Subsequent calls return an error.
func (s *Supervisor) Start(ctx context.Context) error {
s.mu.Lock()
if s.running.Load() {
s.mu.Unlock()
return errors.New("yaraworker: supervisor already started")
}
if s.stopped {
s.mu.Unlock()
return errors.New("yaraworker: supervisor already stopped")
}
s.ctx, s.cancel = context.WithCancel(ctx)
s.done = make(chan struct{})
s.mu.Unlock()
if err := s.spawnAndWaitReady(); err != nil {
s.cancel()
close(s.done)
return err
}
s.mu.Lock()
stopped := s.stopped
if !stopped {
s.running.Store(true)
}
s.mu.Unlock()
obs.Go("yara-supervisor", s.supervise)
if stopped {
return errors.New("yaraworker: supervisor already stopped")
}
return nil
}
// Stop signals the worker to exit, waits for it, and prevents further
// restarts. Safe to call multiple times; subsequent calls are no-ops.
func (s *Supervisor) Stop() error {
s.mu.Lock()
if s.stopped {
s.mu.Unlock()
return nil
}
s.stopped = true
// Clear running so post-Stop ScanFile/ScanBytes/Reload short-circuit to
// the degraded path instead of redialing the now-closed worker socket on
// every call and logging dial errors.
s.running.Store(false)
cancel := s.cancel
done := s.done
s.mu.Unlock()
if cancel != nil {
cancel()
}
s.mu.Lock()
if s.cmd != nil && s.cmd.Process != nil {
_ = s.cmd.Process.Signal(syscall.SIGTERM)
}
if s.client != nil {
_ = s.client.Close()
}
s.mu.Unlock()
if done != nil {
<-done
}
// A timer may have passed its stopped/context check before shutdown.
// Join that callback before allowing its owner to tear down health state.
s.callbackMu.Lock()
defer s.callbackMu.Unlock()
return nil
}
// ScanFile is the compatibility entrypoint. Callers that must distinguish a
// worker failure from a clean path scan should use ScanFileChecked.
func (s *Supervisor) ScanFile(path string, maxBytes int) []yara.Match {
result, _ := s.ScanFileChecked(path, maxBytes)
return result.Matches
}
// ScanFileChecked returns a non-nil error when the worker is down or the path
// scan could not be delivered or completed.
func (s *Supervisor) ScanFileChecked(path string, maxBytes int) (yara.FileScanResult, error) {
if !s.running.Load() {
return yara.FileScanResult{}, errors.New("yaraworker: supervisor not running")
}
s.mu.Lock()
client := s.client
s.mu.Unlock()
if client == nil {
return yara.FileScanResult{}, errors.New("yaraworker: no client")
}
res, err := client.ScanFile(yaraipc.ScanFileArgs{Path: path, MaxBytes: maxBytes})
if err != nil {
s.logf("scan_file: %v", err)
return yara.FileScanResult{}, fmt.Errorf("yaraworker scan_file: %w", err)
}
return yara.FileScanResult{
Matches: toYaraMatches(res.Matches),
ContentSHA256: res.ContentSHA256,
}, nil
}
// ScanBytes is the daemon-facing entrypoint for already-in-memory data. A
// worker error surfaces as zero matches; callers that must fail closed on an
// unscannable payload should use ScanBytesChecked instead.
func (s *Supervisor) ScanBytes(data []byte) []yara.Match {
m, _ := s.ScanBytesChecked(data)
return m
}
// ScanBytesChecked is the fail-closed entrypoint: unlike ScanBytes it returns
// a non-nil error when the worker is down or the request could not be
// delivered (e.g. the payload exceeds the IPC frame budget, or the worker
// crashed mid-request), so a caller does not mistake an unscanned payload for
// a clean one.
func (s *Supervisor) ScanBytesChecked(data []byte) ([]yara.Match, error) {
if !s.running.Load() {
return nil, errors.New("yaraworker: supervisor not running")
}
s.mu.Lock()
client := s.client
s.mu.Unlock()
if client == nil {
return nil, errors.New("yaraworker: no client")
}
res, err := client.ScanBytes(yaraipc.ScanBytesArgs{Data: data})
if err != nil {
s.logScanErr(err)
return nil, fmt.Errorf("yaraworker scan_bytes: %w", err)
}
return toYaraMatches(res.Matches), nil
}
// logScanErr writes a scan failure at most once per scanErrLogWindow per
// distinct message, reporting how many identical failures were folded into
// the gap so the volume stays visible. Tracking is per message rather than
// per most-recent, so two failures alternating cannot defeat suppression.
func (s *Supervisor) logScanErr(err error) {
msg := err.Error()
now := time.Now()
s.scanErrMu.Lock()
if s.scanErrSeen == nil {
s.scanErrSeen = make(map[string]*scanErrRecord)
}
rec, ok := s.scanErrSeen[msg]
overflow := !ok && len(s.scanErrSeen) >= scanErrMaxTracked
if overflow {
// Preserve known recurring failures. Resetting the table lets a stream
// of unique offsets or paths disable throttling for every message.
rec = &s.scanErrOverflow
ok = !rec.at.IsZero()
}
if ok && now.Sub(rec.at) < scanErrLogWindow {
rec.suppressed++
s.scanErrMu.Unlock()
return
}
if !ok {
if !overflow {
rec = &scanErrRecord{}
s.scanErrSeen[msg] = rec
}
}
suppressed := rec.suppressed
rec.at = now
rec.suppressed = 0
s.scanErrMu.Unlock()
if overflow {
s.logf("scan_bytes: %v (%d additional scan failures suppressed)", err, suppressed)
return
}
if suppressed > 0 {
s.logf("scan_bytes: %v (%d identical failures suppressed)", err, suppressed)
return
}
s.logf("scan_bytes: %v", err)
}
// Reload asks the worker to recompile its rules directory.
func (s *Supervisor) Reload() error {
if !s.running.Load() {
return errors.New("yaraworker: supervisor not running")
}
s.mu.Lock()
client := s.client
s.mu.Unlock()
if client == nil {
return errors.New("yaraworker: no client")
}
res, err := client.Reload(yaraipc.ReloadArgs{})
if err != nil {
return err
}
if res.CompileError != "" {
return fmt.Errorf("yaraworker reload compile error: %s", res.CompileError)
}
return nil
}
// CompileError returns the worker's current rule-compile error, or "" when
// rules compiled cleanly (or the worker is unreachable). A non-empty value
// means the worker process is alive but has no usable rules until a reload
// fixes them -- the daemon surfaces this as a finding so a silent-dead YARA
// backend is visible instead of masquerading as "0 rules, all fine".
func (s *Supervisor) CompileError() string {
if !s.running.Load() {
return ""
}
s.mu.Lock()
client := s.client
s.mu.Unlock()
if client == nil {
return ""
}
res, err := client.Ping()
if err != nil {
return ""
}
return res.CompileError
}
// RuleCount queries the worker. Zero on any error, matching the scanner
// semantics the daemon already expects.
func (s *Supervisor) RuleCount() int {
if !s.running.Load() {
return 0
}
s.mu.Lock()
client := s.client
s.mu.Unlock()
if client == nil {
return 0
}
res, err := client.Ping()
if err != nil {
return 0
}
return res.RuleCount
}
// RestartCount is exposed for metrics + tests.
func (s *Supervisor) RestartCount() int {
s.mu.Lock()
defer s.mu.Unlock()
return s.restartCount
}
// ChildPID returns the current worker's pid, or 0 when no worker is
// running. For operator-facing log lines.
func (s *Supervisor) ChildPID() int {
s.mu.Lock()
defer s.mu.Unlock()
if s.cmd == nil || s.cmd.Process == nil {
return 0
}
return s.cmd.Process.Pid
}
// RestartWorker signals the current worker to exit so the supervise
// loop respawns it against whatever state is now on disk (new rules
// directory, updated binary, etc.). Callers should prefer Reload for
// normal rule updates; RestartWorker is the escalation path for the
// rare case where an in-process recompile cannot or must not be
// trusted. The call returns immediately; the restart is asynchronous
// and observable via the OnRestart callback.
//
// No-op when the supervisor is stopped or has no running child.
func (s *Supervisor) RestartWorker() error {
s.mu.Lock()
if s.stopped {
s.mu.Unlock()
return errors.New("yaraworker: supervisor is stopped")
}
cmd := s.cmd
s.mu.Unlock()
if cmd == nil || cmd.Process == nil {
return errors.New("yaraworker: no running worker")
}
return cmd.Process.Signal(syscall.SIGTERM)
}
// supervise watches the current child and restarts it on exit until
// ctx is cancelled.
func (s *Supervisor) supervise() {
defer close(s.done)
backoff := s.cfg.MinRestartInterval
for {
exitCode, sig := s.waitForChild()
if s.ctx.Err() != nil {
return
}
runDuration := time.Since(s.started)
s.mu.Lock()
s.restartCount++
s.lastExitCode = exitCode
s.lastExitSignal = sig
s.mu.Unlock()
if s.cfg.OnRestart != nil {
s.callbackMu.Lock()
s.cfg.OnRestart(exitCode, sig, runDuration)
s.callbackMu.Unlock()
}
// A stable exit already reset the delay. Short-lived workers and
// failed spawn retries still advance the backoff.
exitBackoffRecorded := runDuration >= s.cfg.StableDuration
if exitBackoffRecorded {
backoff = s.cfg.MinRestartInterval
}
for {
s.logf("worker exited code=%d signal=%v ran=%s, restarting in %s",
exitCode, sig, runDuration.Round(time.Millisecond), backoff)
select {
case <-time.After(backoff):
case <-s.ctx.Done():
return
}
err := s.spawnAndWaitReady()
backoff, exitBackoffRecorded = restartBackoffAfterAttempt(
backoff,
exitBackoffRecorded,
err != nil,
s.cfg.MaxRestartInterval,
)
if err == nil {
break
}
if s.ctx.Err() != nil {
return
}
s.logf("restart failed: %v", err)
}
}
}
func restartBackoffAfterAttempt(
current time.Duration,
exitBackoffRecorded bool,
spawnFailed bool,
max time.Duration,
) (time.Duration, bool) {
if exitBackoffRecorded && !spawnFailed {
return current, exitBackoffRecorded
}
next := current * 2
if next > max {
next = max
}
return next, true
}
// waitForChild blocks until the current worker exits, then returns its
// exit code and signal. -1/0 for either field means "unknown" or "not
// applicable".
func (s *Supervisor) waitForChild() (int, syscall.Signal) {
s.mu.Lock()
cmd := s.cmd
s.mu.Unlock()
if cmd == nil {
return -1, 0
}
err := cmd.Wait()
s.mu.Lock()
if s.cmd == cmd {
s.cmd = nil
}
if s.client != nil {
_ = s.client.Close()
s.client = nil
}
s.mu.Unlock()
if err == nil {
return 0, 0
}
var ee *exec.ExitError
if errors.As(err, &ee) {
if status, ok := ee.Sys().(syscall.WaitStatus); ok {
if status.Signaled() {
return -1, status.Signal()
}
return status.ExitStatus(), 0
}
return ee.ExitCode(), 0
}
return -1, 0
}
func (s *Supervisor) spawnAndWaitReady() error {
s.mu.Lock()
if s.client != nil {
_ = s.client.Close()
s.client = nil
}
s.mu.Unlock()
if err := s.ctx.Err(); err != nil {
return err
}
// Unlink stale socket here too, even though the worker also does
// it: the worker may fail before reaching its own unlink, leaving
// a stale file that blocks dial attempts during a failed start.
if err := os.Remove(s.cfg.SocketPath); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("yaraworker: removing stale socket: %w", err)
}
args := []string{"yara-worker",
"--socket", s.cfg.SocketPath,
"--rules-dir", s.cfg.RulesDir,
}
if s.cfg.ConfigFile != "" {
args = append(args, "--config", s.cfg.ConfigFile)
}
// Preserve the daemon's selection even when it is empty or missing.
// Omitting it would re-enable the worker's environment/default lookup.
args = append(args, "--inherited-config-dir", s.cfg.ConfigDir)
disabled, _ := json.Marshal(s.cfg.DisabledRules) // []string cannot fail to encode.
args = append(args, "--disabled-rules", string(disabled))
args = append(args, s.cfg.ExtraArgs...)
// #nosec G204 -- BinaryPath is supervisor-operator-configured (see
// cmd/csm/main.go binaryPath), not attacker-controlled.
cmd := exec.Command(s.cfg.BinaryPath, args...)
if s.cfg.Env != nil {
cmd.Env = s.cfg.Env
}
cmd.Stdout = os.Stderr
cmd.Stderr = os.Stderr
if err := cmd.Start(); err != nil {
return fmt.Errorf("yaraworker: start worker: %w", err)
}
s.mu.Lock()
s.cmd = cmd
s.started = time.Now()
client := yaraipc.NewClient(s.cfg.SocketPath, s.cfg.ClientTimeout)
s.client = client
s.mu.Unlock()
if err := s.waitForReady(client); err != nil {
_ = cmd.Process.Kill()
_ = cmd.Wait()
s.mu.Lock()
s.cmd = nil
_ = s.client.Close()
s.client = nil
s.mu.Unlock()
return err
}
s.reportStableAfter(cmd)
return nil
}
// reportStableAfter calls OnStable once cmd has stayed the current worker for
// StableDuration without the supervisor stopping.
func (s *Supervisor) reportStableAfter(cmd *exec.Cmd) {
if s.cfg.OnStable == nil {
return
}
ctx := s.ctx
time.AfterFunc(s.cfg.StableDuration, func() {
s.callbackMu.Lock()
defer s.callbackMu.Unlock()
s.mu.Lock()
current := s.cmd == cmd && !s.stopped
s.mu.Unlock()
if current && ctx.Err() == nil {
s.cfg.OnStable()
}
})
}
func (s *Supervisor) waitForReady(client *yaraipc.Client) error {
deadline := time.Now().Add(s.cfg.StartTimeout)
for time.Now().Before(deadline) {
if err := s.ctx.Err(); err != nil {
return err
}
if info, err := os.Stat(s.cfg.SocketPath); err == nil && info.Mode()&os.ModeSocket != 0 {
if _, err := client.Ping(); err == nil {
return nil
}
}
select {
case <-time.After(25 * time.Millisecond):
case <-s.ctx.Done():
return s.ctx.Err()
}
}
return fmt.Errorf("yaraworker: worker did not become ready within %s", s.cfg.StartTimeout)
}
func (s *Supervisor) logf(format string, args ...any) {
if s.cfg.Logf != nil {
s.cfg.Logf(format, args...)
}
}
func toYaraMatches(in []yaraipc.Match) []yara.Match {
if len(in) == 0 {
return nil
}
out := make([]yara.Match, len(in))
for i := range in {
out[i] = yara.Match{
RuleName: in[i].RuleName,
Meta: in[i].Meta,
}
}
return out
}
// defaultSocketPath mirrors the roadmap-agreed location. Exposed for
// the daemon to reach when wiring up config defaults.
func DefaultSocketPath() string {
return filepath.Join("/var", "run", "csm", "yara-worker.sock")
}