Files
dockwatch/internal/hostsecurity/security.go
T
groot 6eb4e093ec
release-tag / release-image (push) Successful in 2m45s
Update
2026-09-01 13:59:25 +02:00

1711 lines
55 KiB
Go

package hostsecurity
import (
"bytes"
"context"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"os"
"os/exec"
"path/filepath"
"regexp"
"sort"
"strconv"
"strings"
"sync"
"syscall"
"time"
)
// Config controls access to the host security subsystem. The three capability
// flags are deliberately independent: security can be inspected without
// allowing configuration changes, and package installation requires a second
// explicit opt-in.
type Config struct {
Enabled bool
AllowChanges bool
AllowPackageManagement bool
HostRoot string
DataDir string
HostPID int
}
type Service struct {
cfg Config
mu sync.Mutex
rollbackMu sync.Mutex
rollbacks map[string]*time.Timer
}
type Capabilities struct {
Enabled bool `json:"enabled"`
HostRootAvailable bool `json:"host_root_available"`
HostRootWritable bool `json:"host_root_writable"`
ExecutorAvailable bool `json:"executor_available"`
TargetVerified bool `json:"target_verified"`
AllowChanges bool `json:"allow_changes"`
AllowPackageManagement bool `json:"allow_package_management"`
AuditOnly bool `json:"audit_only"`
Reason string `json:"reason,omitempty"`
}
type OSInfo struct {
ID string `json:"id"`
Name string `json:"name"`
VersionID string `json:"version_id"`
Pretty string `json:"pretty_name"`
}
type ComponentStatus struct {
Name string `json:"name"`
Installed bool `json:"installed"`
Active bool `json:"active"`
Enabled bool `json:"enabled"`
Version string `json:"version,omitempty"`
Service string `json:"service,omitempty"`
ConfigPath string `json:"config_path,omitempty"`
Drift bool `json:"drift"`
Detail string `json:"detail,omitempty"`
Error string `json:"error,omitempty"`
}
type Finding struct {
Severity string `json:"severity"`
Title string `json:"title"`
Detail string `json:"detail"`
Action string `json:"action,omitempty"`
}
type Status struct {
Capabilities Capabilities `json:"capabilities"`
OS OSInfo `json:"os"`
PackageManager string `json:"package_manager,omitempty"`
InitSystem string `json:"init_system,omitempty"`
Firewall ComponentStatus `json:"firewall"`
FirewallBackend FirewallBackendInfo `json:"firewall_backend"`
Fail2Ban ComponentStatus `json:"fail2ban"`
Auditd ComponentStatus `json:"auditd"`
Conflicts []string `json:"firewall_conflicts"`
Findings []Finding `json:"findings"`
Score int `json:"score"`
Managed map[string]ManagedConfig `json:"managed"`
}
type ManagedConfig struct {
Configured bool `json:"configured"`
Drift bool `json:"drift"`
Path string `json:"path,omitempty"`
Hash string `json:"hash,omitempty"`
}
type FirewallRule struct {
Action string `json:"action"`
Protocol string `json:"protocol"`
Port string `json:"port"`
Source string `json:"source,omitempty"`
Comment string `json:"comment,omitempty"`
}
type FirewallPolicy struct {
Provider string `json:"provider,omitempty"`
ResolvedProvider string `json:"resolved_provider,omitempty"`
Enabled bool `json:"enabled"`
ManageDefault bool `json:"manage_default"`
DefaultInbound string `json:"default_inbound"`
Zone string `json:"zone,omitempty"`
AllowICMP bool `json:"allow_icmp"`
TrustedCIDRs []string `json:"trusted_cidrs"`
Rules []FirewallRule `json:"rules"`
}
type FirewallPreview struct {
Policy FirewallPolicy `json:"policy"`
Backend FirewallBackendInfo `json:"backend"`
Runtime FirewallRuntimeView `json:"runtime"`
Rendered string `json:"rendered"`
Warnings []string `json:"warnings"`
Conflict []string `json:"conflicts"`
CanApply bool `json:"can_apply"`
Rollback bool `json:"rollback_supported"`
ManagedTable string `json:"managed_table,omitempty"`
Persistence string `json:"persistence"`
}
type FirewallApplyResult struct {
OK bool `json:"ok"`
ChangeID string `json:"change_id,omitempty"`
ExpiresAt int64 `json:"expires_at,omitempty"`
Preview FirewallPreview `json:"preview"`
Message string `json:"message"`
}
type Fail2BanJail struct {
Name string `json:"name"`
Enabled bool `json:"enabled"`
Port string `json:"port,omitempty"`
Filter string `json:"filter,omitempty"`
Backend string `json:"backend,omitempty"`
LogPath string `json:"logpath,omitempty"`
MaxRetry int `json:"maxretry,omitempty"`
}
type Fail2BanPolicy struct {
Bantime string `json:"bantime"`
Findtime string `json:"findtime"`
MaxRetry int `json:"maxretry"`
Backend string `json:"backend"`
IgnoreIP []string `json:"ignore_ip"`
Jails []Fail2BanJail `json:"jails"`
}
type AuditWatch struct {
Path string `json:"path"`
Permissions string `json:"permissions"`
Key string `json:"key"`
}
type AuditdPolicy struct {
IdentityFiles bool `json:"identity_files"`
Sudoers bool `json:"sudoers"`
SSH bool `json:"ssh"`
Docker bool `json:"docker"`
Systemd bool `json:"systemd"`
KernelModules bool `json:"kernel_modules"`
Custom []AuditWatch `json:"custom"`
}
type PolicyResult struct {
OK bool `json:"ok"`
Message string `json:"message"`
Output string `json:"output,omitempty"`
Warnings []string `json:"warnings,omitempty"`
}
type InstallInput struct {
Enable bool `json:"enable"`
Provider string `json:"provider,omitempty"`
}
type pendingFirewall struct {
ID string `json:"id"`
ExpiresAt int64 `json:"expires_at"`
Provider string `json:"provider,omitempty"`
Previous *FirewallPolicy `json:"previous,omitempty"`
Snapshot FirewallRuntimeSnapshot `json:"snapshot,omitempty"`
}
const (
firewallHostPath = "/etc/dockwatch/firewall.nft"
firewallUnitPath = "/etc/systemd/system/dockwatch-firewall.service"
firewallOpenRC = "/etc/local.d/dockwatch-firewall.start"
fail2banHostPath = "/etc/fail2ban/jail.d/dockwatch.local"
auditdHostPath = "/etc/audit/rules.d/90-dockwatch.rules"
)
var (
safeNameRE = regexp.MustCompile(`^[A-Za-z0-9_.-]{1,64}$`)
safePortRE = regexp.MustCompile(`^[0-9]{1,5}(-[0-9]{1,5})?$`)
safePortListRE = regexp.MustCompile(`^[A-Za-z0-9_,:-]{1,128}$`)
durationRE = regexp.MustCompile(`^[0-9]{1,9}([smhdwy])?$`)
auditKeyRE = regexp.MustCompile(`^[A-Za-z0-9_.-]{1,64}$`)
)
func New(cfg Config) *Service {
if cfg.HostPID <= 0 {
cfg.HostPID = 1
}
s := &Service{cfg: cfg, rollbacks: map[string]*time.Timer{}}
s.recoverPendingFirewall()
return s
}
func (s *Service) Status(ctx context.Context) Status {
caps := s.capabilities(ctx)
osInfo := s.osInfo()
pm := s.packageManager()
initSystem := s.initSystem()
fwPolicy := s.FirewallPolicy()
fwBackend := s.firewallBackend(ctx, fwPolicy.Provider)
st := Status{
Capabilities: caps,
OS: osInfo,
PackageManager: pm,
InitSystem: initSystem,
Firewall: s.firewallComponentStatus(ctx, initSystem, fwBackend),
FirewallBackend: fwBackend,
Fail2Ban: s.componentStatus(ctx, "fail2ban", initSystem),
Auditd: s.componentStatus(ctx, "auditd", initSystem),
Managed: map[string]ManagedConfig{},
}
st.Conflicts = append([]string(nil), fwBackend.Conflicts...)
if _, ok := s.loadFirewallPolicy(); ok {
st.Managed["firewall"] = s.firewallManagedConfig(ctx, fwPolicy, fwBackend)
st.Firewall.Drift = st.Managed["firewall"].Drift
}
if p, ok := s.loadFail2BanPolicy(); ok {
rendered, _ := renderFail2Ban(p)
st.Managed["fail2ban"] = s.managedConfig(fail2banHostPath, rendered)
st.Fail2Ban.Drift = st.Managed["fail2ban"].Drift
}
if p, ok := s.loadAuditdPolicy(); ok {
rendered, _ := s.renderAuditdForHost(p)
st.Managed["auditd"] = s.managedConfig(auditdHostPath, rendered)
st.Auditd.Drift = st.Managed["auditd"].Drift
}
st.Findings, st.Score = findings(st)
return st
}
func findings(st Status) ([]Finding, int) {
fs := []Finding{}
score := 100
if !st.Capabilities.Enabled {
return []Finding{{Severity: "info", Title: "Host Security layer disabled", Detail: "Enable HOST_SECURITY_ENABLED to inspect the host security posture."}}, 0
}
if !st.Capabilities.HostRootAvailable {
fs = append(fs, Finding{Severity: "high", Title: "Host root unavailable", Detail: "HOST_ROOT is not mounted; file-based host security checks cannot run."})
score -= 25
}
if !st.Firewall.Installed {
provider := st.FirewallBackend.Selected
if provider == "" {
provider = "firewall"
}
fs = append(fs, Finding{Severity: "high", Title: "Firewall provider not installed", Detail: "Selected provider " + provider + " is not installed on this host.", Action: "Install " + provider})
score -= 25
} else if !st.Firewall.Active {
fs = append(fs, Finding{Severity: "medium", Title: "Firewall frontend inactive", Detail: "Selected provider " + st.FirewallBackend.Selected + " is installed but not currently active."})
score -= 15
}
if len(st.Conflicts) > 0 {
fs = append(fs, Finding{Severity: "high", Title: "Competing firewall frontends detected", Detail: "Dockwatch detected conflicting active frontends: " + strings.Join(st.Conflicts, ", ") + ". Resolve this before applying changes."})
score -= 10
}
if !st.Fail2Ban.Installed {
fs = append(fs, Finding{Severity: "medium", Title: "Fail2Ban not installed", Detail: "Brute-force protection is not available through Dockwatch.", Action: "Install Fail2Ban"})
score -= 15
} else if !st.Fail2Ban.Active {
fs = append(fs, Finding{Severity: "medium", Title: "Fail2Ban inactive", Detail: "Fail2Ban is installed but its service is not active."})
score -= 10
}
if !st.Auditd.Installed {
fs = append(fs, Finding{Severity: "medium", Title: "Linux audit tooling not installed", Detail: "auditd/auditctl is not available.", Action: "Install auditd"})
score -= 15
} else if !st.Auditd.Active {
fs = append(fs, Finding{Severity: "medium", Title: "auditd inactive", Detail: "Linux audit tooling is installed but auditd is not active."})
score -= 10
}
for name, m := range st.Managed {
if m.Drift {
fs = append(fs, Finding{Severity: "medium", Title: "Managed configuration drift", Detail: fmt.Sprintf("The host copy for %s differs from Dockwatch's desired policy.", name)})
score -= 5
}
}
if score < 0 {
score = 0
}
if len(fs) == 0 {
fs = append(fs, Finding{Severity: "ok", Title: "No managed security findings", Detail: "All enabled Dockwatch-managed host security components are healthy."})
}
return fs, score
}
func (s *Service) capabilities(ctx context.Context) Capabilities {
c := Capabilities{Enabled: s.cfg.Enabled, AllowChanges: s.cfg.AllowChanges, AllowPackageManagement: s.cfg.AllowPackageManagement}
if !s.cfg.Enabled {
c.Reason = "HOST_SECURITY_ENABLED=false"
return c
}
if s.cfg.HostRoot == "" {
c.Reason = "HOST_ROOT is not configured"
return c
}
if _, err := os.Stat(s.cfg.HostRoot); err != nil {
c.Reason = "HOST_ROOT is unavailable: " + err.Error()
return c
}
c.HostRootAvailable = true
c.HostRootWritable = hostRootWritable(s.cfg.HostRoot)
if _, err := exec.LookPath("nsenter"); err != nil {
c.Reason = "nsenter is not installed in the Dockwatch container"
c.AuditOnly = true
return c
}
if err := s.verifyHostTarget(); err != nil {
c.Reason = err.Error()
c.AuditOnly = true
return c
}
c.TargetVerified = true
x, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
if _, err := s.hostCommand(x, nil, "true"); err != nil {
c.Reason = "host namespace executor unavailable: " + err.Error()
c.AuditOnly = true
return c
}
c.ExecutorAvailable = true
c.AuditOnly = !s.cfg.AllowChanges
if s.cfg.AllowChanges && !c.HostRootWritable {
c.Reason = "HOST_ROOT is read-only; host configuration files cannot be changed"
}
return c
}
func hostRootWritable(path string) bool {
var st syscall.Statfs_t
if err := syscall.Statfs(path, &st); err != nil {
return false
}
const stReadonly = 1
return st.Flags&stReadonly == 0
}
func (s *Service) verifyHostTarget() error {
if s.cfg.HostRoot == "" {
return errors.New("HOST_ROOT is not configured")
}
if os.Getpid() == s.cfg.HostPID {
return errors.New("Dockwatch appears to be PID 1; enable pid: host for host command execution")
}
hostInfo, err := os.Stat(s.cfg.HostRoot)
if err != nil {
return fmt.Errorf("stat HOST_ROOT: %w", err)
}
pidRoot := fmt.Sprintf("/proc/%d/root", s.cfg.HostPID)
pidInfo, err := os.Stat(pidRoot)
if err != nil {
return fmt.Errorf("stat host PID root: %w", err)
}
if !os.SameFile(hostInfo, pidInfo) {
return errors.New("HOST_ROOT does not match the configured host PID root; mount /:/host and use pid: host")
}
return nil
}
func (s *Service) requireChanges(ctx context.Context) error {
if !s.cfg.Enabled {
return errors.New("host security layer is disabled")
}
if !s.cfg.AllowChanges {
return errors.New("host security changes are disabled; set ALLOW_HOST_SECURITY_CHANGES=true")
}
caps := s.capabilities(ctx)
if !caps.ExecutorAvailable || !caps.TargetVerified {
return errors.New(caps.Reason)
}
if !caps.HostRootWritable {
return errors.New("HOST_ROOT must be mounted read-write for host security changes")
}
return nil
}
func (s *Service) requirePackages(ctx context.Context) error {
if err := s.requireChanges(ctx); err != nil {
return err
}
if !s.cfg.AllowPackageManagement {
return errors.New("host package management is disabled; set ALLOW_HOST_PACKAGE_MANAGEMENT=true")
}
if s.packageManager() == "" {
return errors.New("no supported host package manager detected")
}
return nil
}
func (s *Service) hostCommand(ctx context.Context, stdin []byte, args ...string) (string, error) {
if len(args) == 0 {
return "", errors.New("host command is empty")
}
base := []string{"-t", strconv.Itoa(s.cfg.HostPID), "-m", "-u", "-i", "-n", "-p", "-r", "--"}
base = append(base, args...)
cmd := exec.CommandContext(ctx, "nsenter", base...)
if stdin != nil {
cmd.Stdin = bytes.NewReader(stdin)
}
var b bytes.Buffer
cmd.Stdout = &b
cmd.Stderr = &b
if err := cmd.Run(); err != nil {
msg := strings.TrimSpace(b.String())
if msg == "" {
msg = err.Error()
}
return msg, fmt.Errorf("host command failed: %s", msg)
}
return strings.TrimSpace(b.String()), nil
}
func (s *Service) osInfo() OSInfo {
b, err := os.ReadFile(s.hostPath("/etc/os-release"))
if err != nil {
return OSInfo{}
}
m := map[string]string{}
for _, line := range strings.Split(string(b), "\n") {
k, v, ok := strings.Cut(line, "=")
if !ok {
continue
}
v = strings.Trim(strings.TrimSpace(v), `"`)
m[k] = v
}
return OSInfo{ID: m["ID"], Name: m["NAME"], VersionID: m["VERSION_ID"], Pretty: m["PRETTY_NAME"]}
}
func (s *Service) packageManager() string {
checks := []struct{ name, path string }{
{"apt", "/usr/bin/apt-get"}, {"dnf", "/usr/bin/dnf"}, {"yum", "/usr/bin/yum"},
{"zypper", "/usr/bin/zypper"}, {"apk", "/sbin/apk"}, {"pacman", "/usr/bin/pacman"},
}
for _, c := range checks {
if fileExists(s.hostPath(c.path)) {
return c.name
}
}
return ""
}
func (s *Service) initSystem() string {
if dirExists(s.hostPath("/run/systemd/system")) && (fileExists(s.hostPath("/bin/systemctl")) || fileExists(s.hostPath("/usr/bin/systemctl"))) {
return "systemd"
}
if fileExists(s.hostPath("/sbin/openrc")) || fileExists(s.hostPath("/sbin/rc-service")) {
return "openrc"
}
return "unknown"
}
func (s *Service) componentStatus(ctx context.Context, component, initSystem string) ComponentStatus {
var bin, service, cfg string
switch component {
case "firewall":
bin, service, cfg = "/usr/sbin/nft", "dockwatch-firewall", firewallHostPath
if !fileExists(s.hostPath(bin)) {
bin = "/sbin/nft"
}
case "fail2ban":
bin, service, cfg = "/usr/bin/fail2ban-client", "fail2ban", fail2banHostPath
case "auditd":
bin, service, cfg = "/sbin/auditctl", "auditd", auditdHostPath
if !fileExists(s.hostPath(bin)) {
bin = "/usr/sbin/auditctl"
}
}
st := ComponentStatus{Name: component, Service: service, ConfigPath: cfg, Installed: fileExists(s.hostPath(bin))}
if !st.Installed {
return st
}
caps := s.capabilitiesNoCommand()
if !caps.TargetVerified {
st.Detail = "installed; runtime service status unavailable without verified host namespace access"
return st
}
x, cancel := context.WithTimeout(ctx, 4*time.Second)
defer cancel()
var versionArgs []string
switch component {
case "firewall":
versionArgs = []string{"nft", "--version"}
case "fail2ban":
versionArgs = []string{"fail2ban-client", "--version"}
case "auditd":
versionArgs = []string{"auditctl", "-v"}
}
if out, err := s.hostCommand(x, nil, versionArgs...); err == nil {
st.Version = firstLine(out)
}
if component == "firewall" {
if _, err := s.hostCommand(x, nil, "nft", "list", "table", "inet", "dockwatch"); err == nil {
st.Active = true
}
st.Enabled = s.persistenceEnabled(ctx, initSystem, "dockwatch-firewall")
return st
}
st.Active, st.Enabled = s.serviceState(ctx, initSystem, service)
if component == "fail2ban" && st.Active {
if out, err := s.hostCommand(x, nil, "fail2ban-client", "status"); err == nil {
st.Detail = limit(out, 2048)
}
}
if component == "auditd" {
if out, err := s.hostCommand(x, nil, "auditctl", "-s"); err == nil {
st.Detail = limit(out, 2048)
}
}
return st
}
func (s *Service) capabilitiesNoCommand() Capabilities {
c := Capabilities{Enabled: s.cfg.Enabled}
if !s.cfg.Enabled || s.cfg.HostRoot == "" {
return c
}
if _, err := os.Stat(s.cfg.HostRoot); err != nil {
return c
}
c.HostRootAvailable = true
if _, err := exec.LookPath("nsenter"); err != nil {
return c
}
if err := s.verifyHostTarget(); err == nil {
c.TargetVerified = true
c.ExecutorAvailable = true
}
return c
}
func (s *Service) serviceState(ctx context.Context, initSystem, service string) (active, enabled bool) {
x, cancel := context.WithTimeout(ctx, 3*time.Second)
defer cancel()
switch initSystem {
case "systemd":
if _, err := s.hostCommand(x, nil, "systemctl", "is-active", "--quiet", service); err == nil {
active = true
}
if _, err := s.hostCommand(x, nil, "systemctl", "is-enabled", "--quiet", service); err == nil {
enabled = true
}
case "openrc":
if out, err := s.hostCommand(x, nil, "rc-service", service, "status"); err == nil && strings.Contains(strings.ToLower(out), "started") {
active = true
}
if out, err := s.hostCommand(x, nil, "rc-update", "show", "default"); err == nil && strings.Contains(out, service) {
enabled = true
}
}
return
}
func (s *Service) persistenceEnabled(ctx context.Context, initSystem, service string) bool {
_, enabled := s.serviceState(ctx, initSystem, service)
return enabled
}
func (s *Service) firewallConflicts(ctx context.Context) []string {
return append([]string(nil), s.firewallBackend(ctx, FirewallProviderAuto).Conflicts...)
}
func (s *Service) FirewallPolicy() FirewallPolicy {
if p, ok := s.loadFirewallPolicy(); ok {
if normalizeFirewallProvider(p.Provider) == "" {
p.Provider = FirewallProviderAuto
}
if p.DefaultInbound == "" {
p.DefaultInbound = "accept"
}
if p.TrustedCIDRs == nil {
p.TrustedCIDRs = []string{}
}
if p.Rules == nil {
p.Rules = []FirewallRule{}
}
return p
}
return FirewallPolicy{Provider: FirewallProviderAuto, Enabled: false, ManageDefault: false, DefaultInbound: "accept", AllowICMP: true, TrustedCIDRs: []string{}, Rules: []FirewallRule{{Action: "accept", Protocol: "tcp", Port: "22", Comment: "SSH"}}}
}
func (s *Service) PreviewFirewall(ctx context.Context, p FirewallPolicy) (FirewallPreview, error) {
provider := normalizeFirewallProvider(p.Provider)
if provider == "" {
return FirewallPreview{}, errors.New("provider must be auto, ufw, firewalld or nftables")
}
p.Provider = provider
if p.DefaultInbound == "" {
p.DefaultInbound = "accept"
}
if p.TrustedCIDRs == nil {
p.TrustedCIDRs = []string{}
}
if p.Rules == nil {
p.Rules = []FirewallRule{}
}
backend := s.firewallBackend(ctx, provider)
p.ResolvedProvider = backend.Selected
runtime := s.firewallRuntime(ctx, provider)
if backend.Selected == FirewallProviderFirewalld && p.Zone == "" {
p.Zone = runtime.Zone
if p.Zone == "" {
p.Zone = backend.DefaultZone
}
}
if backend.Selected == FirewallProviderFirewalld && p.Zone != "" && (!safeNameRE.MatchString(p.Zone) || strings.HasPrefix(p.Zone, "-")) {
return FirewallPreview{}, errors.New("firewalld zone name is invalid")
}
rendered, warnings, err := renderFirewallPlan(p, backend, runtime)
if err != nil {
return FirewallPreview{}, err
}
if p.DefaultInbound == "drop" && (backend.Selected == FirewallProviderNftables || p.ManageDefault) {
warnings = append(warnings, "Default inbound DROP can lock out SSH or Dockwatch. Apply uses a timed rollback until explicitly committed.")
}
if len(backend.Conflicts) > 0 {
warnings = append(warnings, "Competing active firewall frontends detected: "+strings.Join(backend.Conflicts, ", ")+". Apply is blocked until only the selected frontend owns host filtering.")
}
if backend.Selected == FirewallProviderFirewalld && p.ManageDefault && !runtime.Active {
warnings = append(warnings, "firewalld is inactive. Start it before changing the zone target so Dockwatch can snapshot the current runtime/permanent target for rollback.")
}
if oneOf(backend.Selected, FirewallProviderUFW, FirewallProviderFirewalld) && p.Enabled && !runtime.Active {
warnings = append(warnings, "The selected firewall frontend is inactive. Applying an enabled Dockwatch policy will activate it; rollback restores the previous active state.")
}
if !containsString(backend.Available, backend.Selected) {
warnings = append(warnings, backend.Selected+" is not installed on this host.")
}
caps := s.capabilities(ctx)
managedTable := ""
if backend.Selected == FirewallProviderNftables {
managedTable = "inet dockwatch"
}
canApply := caps.AllowChanges && caps.ExecutorAvailable && caps.HostRootWritable && len(backend.Conflicts) == 0 && containsString(backend.Available, backend.Selected)
if backend.Selected == FirewallProviderFirewalld && p.ManageDefault && !runtime.Active {
canApply = false
}
return FirewallPreview{Policy: p, Backend: backend, Runtime: runtime, Rendered: rendered, Warnings: warnings, Conflict: backend.Conflicts, CanApply: canApply, Rollback: true, ManagedTable: managedTable, Persistence: s.initSystem()}, nil
}
func renderFirewall(p FirewallPolicy) (string, error) {
if p.DefaultInbound == "" {
p.DefaultInbound = "accept"
}
if p.DefaultInbound != "accept" && p.DefaultInbound != "drop" {
return "", errors.New("default_inbound must be accept or drop")
}
for _, cidr := range p.TrustedCIDRs {
if _, _, err := net.ParseCIDR(strings.TrimSpace(cidr)); err != nil {
return "", fmt.Errorf("invalid trusted CIDR %q", cidr)
}
}
for i, r := range p.Rules {
if !oneOf(r.Action, "accept", "drop", "reject", "limit") {
return "", fmt.Errorf("rule %d action must be accept, drop, reject or limit", i+1)
}
if r.Action == "limit" && r.Protocol != "tcp" {
return "", fmt.Errorf("rule %d limit action is supported only for tcp", i+1)
}
if r.Protocol != "tcp" && r.Protocol != "udp" {
return "", fmt.Errorf("rule %d protocol must be tcp or udp", i+1)
}
if !safePortRE.MatchString(r.Port) || !validPortRange(r.Port) {
return "", fmt.Errorf("rule %d has invalid port/range", i+1)
}
if r.Source != "" {
if _, _, err := net.ParseCIDR(r.Source); err != nil {
return "", fmt.Errorf("rule %d has invalid source CIDR", i+1)
}
}
if strings.ContainsAny(r.Comment, "\n\r") || len(r.Comment) > 100 {
return "", fmt.Errorf("rule %d comment is invalid", i+1)
}
}
var b strings.Builder
b.WriteString("# Managed by Dockwatch. Manual edits are detected as drift.\n")
b.WriteString("table inet dockwatch {\n")
b.WriteString(" chain input {\n")
b.WriteString(" type filter hook input priority 10; policy " + p.DefaultInbound + ";\n")
b.WriteString(" ct state established,related accept comment \"dockwatch established\"\n")
b.WriteString(" iifname \"lo\" accept comment \"dockwatch loopback\"\n")
b.WriteString(" ct state invalid drop comment \"dockwatch invalid\"\n")
if p.AllowICMP {
b.WriteString(" ip protocol icmp accept comment \"dockwatch icmp4\"\n")
b.WriteString(" ip6 nexthdr ipv6-icmp accept comment \"dockwatch icmp6\"\n")
}
for _, cidr := range p.TrustedCIDRs {
family := "ip"
if strings.Contains(cidr, ":") {
family = "ip6"
}
b.WriteString(fmt.Sprintf(" %s saddr %s accept comment \"dockwatch trusted\"\n", family, cidr))
}
for _, r := range p.Rules {
var src string
if r.Source != "" {
family := "ip"
if strings.Contains(r.Source, ":") {
family = "ip6"
}
src = family + " saddr " + r.Source + " "
}
comment := "dockwatch rule"
if strings.TrimSpace(r.Comment) != "" {
comment = "dockwatch " + strings.TrimSpace(r.Comment)
}
comment = strings.ReplaceAll(comment, `"`, `'`)
action := r.Action
if action == "limit" {
b.WriteString(fmt.Sprintf(" %s%s dport %s ct state new limit rate 6/minute accept comment %s\n", src, r.Protocol, r.Port, strconv.Quote(comment)))
continue
}
b.WriteString(fmt.Sprintf(" %s%s dport %s %s comment %s\n", src, r.Protocol, r.Port, action, strconv.Quote(comment)))
}
b.WriteString(" }\n}\n")
return b.String(), nil
}
func validPortRange(v string) bool {
parts := strings.Split(v, "-")
for _, p := range parts {
n, err := strconv.Atoi(p)
if err != nil || n < 1 || n > 65535 {
return false
}
}
if len(parts) == 2 {
a, _ := strconv.Atoi(parts[0])
b, _ := strconv.Atoi(parts[1])
return a <= b
}
return len(parts) == 1
}
func (s *Service) ApplyFirewall(ctx context.Context, p FirewallPolicy, rollbackSeconds int) (FirewallApplyResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
if err := s.requireChanges(ctx); err != nil {
return FirewallApplyResult{}, err
}
preview, err := s.PreviewFirewall(ctx, p)
if err != nil {
return FirewallApplyResult{}, err
}
p = preview.Policy
if len(preview.Backend.Conflicts) > 0 {
return FirewallApplyResult{}, fmt.Errorf("competing firewall frontend(s) active: %s", strings.Join(preview.Backend.Conflicts, ", "))
}
if !containsString(preview.Backend.Available, preview.Backend.Selected) {
return FirewallApplyResult{}, fmt.Errorf("selected firewall provider %s is not installed", preview.Backend.Selected)
}
if preview.Backend.Selected == FirewallProviderFirewalld && p.ManageDefault && !preview.Runtime.Active {
return FirewallApplyResult{}, errors.New("firewalld must be active before Dockwatch can change and safely roll back the zone target")
}
if rollbackSeconds == 0 {
rollbackSeconds = 90
}
if rollbackSeconds < 30 || rollbackSeconds > 600 {
return FirewallApplyResult{}, errors.New("rollback_seconds must be between 30 and 600")
}
previous, hadPrevious := s.loadFirewallPolicy()
var previousPtr *FirewallPolicy
if hadPrevious {
previousPtr = &previous
}
snapshot := s.snapshotFirewallRuntime(ctx, preview.Backend.Selected, p.Zone)
rollbackState := pendingFirewall{Provider: preview.Backend.Selected, Previous: previousPtr, Snapshot: snapshot}
if hadPrevious {
oldProvider := s.providerForPolicy(ctx, previous)
if oldProvider != "" && oldProvider != preview.Backend.Selected {
if err := s.removeCurrentManagedFirewall(ctx, previous); err != nil {
return FirewallApplyResult{}, fmt.Errorf("remove previous %s managed firewall before provider switch: %w", oldProvider, err)
}
}
}
if err := s.applyFirewallProvider(ctx, p, previousPtr, preview.Backend); err != nil {
s.cleanupAttemptedFirewall(ctx, p, preview.Backend)
_ = s.rollbackFirewallLocked(ctx, rollbackState)
return FirewallApplyResult{}, err
}
if err := s.savePolicy("firewall.json", p); err != nil {
s.cleanupAttemptedFirewall(ctx, p, preview.Backend)
_ = s.rollbackFirewallLocked(ctx, rollbackState)
return FirewallApplyResult{}, err
}
id := randomID()
pending := pendingFirewall{ID: id, ExpiresAt: time.Now().Add(time.Duration(rollbackSeconds) * time.Second).Unix(), Provider: preview.Backend.Selected, Previous: previousPtr, Snapshot: snapshot}
if err := s.savePendingFirewall(pending); err != nil {
_ = s.rollbackFirewallLocked(ctx, rollbackState)
return FirewallApplyResult{}, err
}
s.scheduleFirewallRollback(pending)
return FirewallApplyResult{OK: true, ChangeID: id, ExpiresAt: pending.ExpiresAt, Preview: preview, Message: "Firewall policy applied through " + preview.Backend.Selected + ". Commit the change before the rollback timer expires."}, nil
}
func (s *Service) CommitFirewall(id string) error {
s.rollbackMu.Lock()
defer s.rollbackMu.Unlock()
p, err := s.readPendingFirewall()
if err != nil {
return err
}
if p.ID == "" || p.ID != id {
return errors.New("no matching pending firewall change")
}
if t := s.rollbacks[id]; t != nil {
t.Stop()
delete(s.rollbacks, id)
}
_ = os.Remove(s.securityPath("pending-firewall.json"))
return nil
}
func (s *Service) RollbackFirewall(ctx context.Context, id string) error {
s.mu.Lock()
defer s.mu.Unlock()
p, err := s.readPendingFirewall()
if err != nil {
return err
}
if p.ID == "" || (id != "" && p.ID != id) {
return errors.New("no matching pending firewall change")
}
if err := s.rollbackFirewallLocked(ctx, p); err != nil {
return err
}
s.rollbackMu.Lock()
if t := s.rollbacks[p.ID]; t != nil {
t.Stop()
delete(s.rollbacks, p.ID)
}
s.rollbackMu.Unlock()
_ = os.Remove(s.securityPath("pending-firewall.json"))
return nil
}
func (s *Service) validateFirewallRuntime(ctx context.Context, rendered string) error {
x, cancel := context.WithTimeout(ctx, 8*time.Second)
defer cancel()
if _, err := s.hostCommand(x, []byte(rendered), "nft", "-c", "-f", "-"); err != nil {
return fmt.Errorf("nftables validation failed: %w", err)
}
return nil
}
func (s *Service) applyFirewallRuntime(ctx context.Context, p FirewallPolicy, rendered string) error {
x, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if _, err := s.hostCommand(x, nil, "nft", "list", "table", "inet", "dockwatch"); err == nil {
_, _ = s.hostCommand(x, nil, "nft", "delete", "table", "inet", "dockwatch")
}
if !p.Enabled {
return nil
}
if _, err := s.hostCommand(x, []byte(rendered), "nft", "-f", "-"); err != nil {
return fmt.Errorf("apply nftables policy: %w", err)
}
return nil
}
func (s *Service) configureFirewallPersistence(ctx context.Context, enabled bool, rendered string) error {
if enabled {
if err := s.writeHostManaged(firewallHostPath, []byte(rendered), 0o600); err != nil {
return err
}
return s.ensureFirewallPersistence(ctx)
}
initSystem := s.initSystem()
x, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
switch initSystem {
case "systemd":
_, _ = s.hostCommand(x, nil, "systemctl", "disable", "--now", "dockwatch-firewall.service")
_ = s.removeHostManaged(firewallUnitPath)
_, _ = s.hostCommand(x, nil, "systemctl", "daemon-reload")
case "openrc":
_ = s.removeHostManaged(firewallOpenRC)
}
_ = s.removeHostManaged(firewallHostPath)
return nil
}
func (s *Service) ensureFirewallPersistence(ctx context.Context) error {
initSystem := s.initSystem()
switch initSystem {
case "systemd":
unit := `[Unit]
Description=Dockwatch managed host firewall
After=local-fs.target
Before=network-pre.target docker.service
Wants=network-pre.target
[Service]
Type=oneshot
ExecStart=/usr/sbin/nft -f /etc/dockwatch/firewall.nft
ExecReload=/usr/sbin/nft -f /etc/dockwatch/firewall.nft
ExecStop=-/usr/sbin/nft delete table inet dockwatch
RemainAfterExit=yes
[Install]
WantedBy=multi-user.target
`
// /usr/sbin/nft is the common host path; systems with /sbin/nft usually
// provide /usr/sbin via merged-/usr. Refuse if neither is present.
if !fileExists(s.hostPath("/usr/sbin/nft")) && fileExists(s.hostPath("/sbin/nft")) {
unit = strings.ReplaceAll(unit, "/usr/sbin/nft", "/sbin/nft")
}
if err := s.writeHostManaged(firewallUnitPath, []byte(unit), 0o644); err != nil {
return err
}
x, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
if _, err := s.hostCommand(x, nil, "systemctl", "daemon-reload"); err != nil {
return err
}
if _, err := s.hostCommand(x, nil, "systemctl", "enable", "dockwatch-firewall.service"); err != nil {
return err
}
case "openrc":
script := "#!/bin/sh\nexec nft -f /etc/dockwatch/firewall.nft\n"
if err := s.writeHostManaged(firewallOpenRC, []byte(script), 0o755); err != nil {
return err
}
x, cancel := context.WithTimeout(ctx, 10*time.Second)
defer cancel()
_, _ = s.hostCommand(x, nil, "rc-update", "add", "local", "default")
default:
return errors.New("unsupported init system for persistent firewall; systemd or OpenRC required")
}
return nil
}
func (s *Service) scheduleFirewallRollback(p pendingFirewall) {
d := time.Until(time.Unix(p.ExpiresAt, 0))
if d < 0 {
d = 0
}
s.rollbackMu.Lock()
if old := s.rollbacks[p.ID]; old != nil {
old.Stop()
}
s.rollbacks[p.ID] = time.AfterFunc(d, func() {
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second)
defer cancel()
_ = s.RollbackFirewall(ctx, p.ID)
})
s.rollbackMu.Unlock()
}
func (s *Service) rollbackFirewallLocked(ctx context.Context, p pendingFirewall) error {
current, hasCurrent := s.loadFirewallPolicy()
if p.Previous == nil {
if hasCurrent {
if err := s.removeCurrentManagedFirewall(ctx, current); err != nil {
return err
}
}
if err := s.restoreFirewallSnapshot(ctx, p.Snapshot); err != nil {
return err
}
_ = os.Remove(s.securityPath("firewall.json"))
return nil
}
provider := s.providerForPolicy(ctx, *p.Previous)
backend := s.firewallBackend(ctx, provider)
var currentPtr *FirewallPolicy
restoredSnapshot := false
if hasCurrent {
currentPtr = &current
currentProvider := s.providerForPolicy(ctx, current)
if currentProvider != "" && currentProvider != provider {
if err := s.removeCurrentManagedFirewall(ctx, current); err != nil {
return err
}
if err := s.restoreFirewallSnapshot(ctx, p.Snapshot); err != nil {
return err
}
restoredSnapshot = true
currentPtr = nil
}
}
if !restoredSnapshot && !p.Previous.ManageDefault && oneOf(p.Snapshot.Provider, FirewallProviderUFW, FirewallProviderFirewalld) {
if err := s.restoreFirewallSnapshot(ctx, p.Snapshot); err != nil {
return err
}
}
if err := s.applyFirewallProvider(ctx, *p.Previous, currentPtr, backend); err != nil {
return err
}
return s.savePolicy("firewall.json", *p.Previous)
}
func (s *Service) recoverPendingFirewall() {
p, err := s.readPendingFirewall()
if err != nil || p.ID == "" {
return
}
s.scheduleFirewallRollback(p)
}
func (s *Service) Fail2BanPolicy() Fail2BanPolicy {
if p, ok := s.loadFail2BanPolicy(); ok {
return p
}
backend := "auto"
jailBackend := "auto"
if s.initSystem() == "systemd" {
jailBackend = "systemd"
}
return Fail2BanPolicy{Bantime: "1h", Findtime: "10m", MaxRetry: 5, Backend: backend, IgnoreIP: []string{"127.0.0.1/8", "::1"}, Jails: []Fail2BanJail{{Name: "sshd", Enabled: true, Port: "ssh", Filter: "sshd", Backend: jailBackend}}}
}
func renderFail2Ban(p Fail2BanPolicy) (string, error) {
if !durationRE.MatchString(p.Bantime) || !durationRE.MatchString(p.Findtime) {
return "", errors.New("bantime/findtime must be a number optionally followed by s,m,h,d,w,y")
}
if p.MaxRetry < 1 || p.MaxRetry > 1000 {
return "", errors.New("maxretry must be between 1 and 1000")
}
if p.Backend == "" {
p.Backend = "auto"
}
if !oneOf(p.Backend, "auto", "systemd", "polling", "pyinotify") {
return "", errors.New("unsupported Fail2Ban backend")
}
for _, ip := range p.IgnoreIP {
if !validIPOrCIDR(ip) {
return "", fmt.Errorf("invalid ignore_ip %q", ip)
}
}
var b strings.Builder
b.WriteString("# Managed by Dockwatch. Put manual overrides in a different jail.d file.\n[DEFAULT]\n")
b.WriteString("bantime = " + p.Bantime + "\nfindtime = " + p.Findtime + "\nmaxretry = " + strconv.Itoa(p.MaxRetry) + "\nbackend = " + p.Backend + "\n")
if len(p.IgnoreIP) > 0 {
b.WriteString("ignoreip = " + strings.Join(p.IgnoreIP, " ") + "\n")
}
for i, j := range p.Jails {
if !safeNameRE.MatchString(j.Name) {
return "", fmt.Errorf("jail %d has invalid name", i+1)
}
if j.Port != "" && !safePortListRE.MatchString(j.Port) {
return "", fmt.Errorf("jail %s has invalid port", j.Name)
}
if j.Filter != "" && !safeNameRE.MatchString(j.Filter) {
return "", fmt.Errorf("jail %s has invalid filter", j.Name)
}
if j.Backend != "" && !oneOf(j.Backend, "auto", "systemd", "polling", "pyinotify") {
return "", fmt.Errorf("jail %s has unsupported backend", j.Name)
}
if j.LogPath != "" && (!strings.HasPrefix(j.LogPath, "/") || strings.ContainsAny(j.LogPath, "\r\n;`$")) {
return "", fmt.Errorf("jail %s has invalid logpath", j.Name)
}
if j.MaxRetry < 0 || j.MaxRetry > 1000 {
return "", fmt.Errorf("jail %s maxretry is invalid", j.Name)
}
b.WriteString("\n[" + j.Name + "]\nenabled = " + strconv.FormatBool(j.Enabled) + "\n")
if j.Port != "" {
b.WriteString("port = " + j.Port + "\n")
}
if j.Filter != "" {
b.WriteString("filter = " + j.Filter + "\n")
}
if j.Backend != "" {
b.WriteString("backend = " + j.Backend + "\n")
}
if j.LogPath != "" {
b.WriteString("logpath = " + j.LogPath + "\n")
}
if j.MaxRetry > 0 {
b.WriteString("maxretry = " + strconv.Itoa(j.MaxRetry) + "\n")
}
}
return b.String(), nil
}
func (s *Service) ApplyFail2Ban(ctx context.Context, p Fail2BanPolicy) (PolicyResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
if err := s.requireChanges(ctx); err != nil {
return PolicyResult{}, err
}
if !fileExists(s.hostPath("/usr/bin/fail2ban-client")) {
return PolicyResult{}, errors.New("Fail2Ban is not installed")
}
rendered, err := renderFail2Ban(p)
if err != nil {
return PolicyResult{}, err
}
previous, hadPrevious := s.readHostSnapshot(fail2banHostPath)
if err = s.writeHostManaged(fail2banHostPath, []byte(rendered), 0o640); err != nil {
return PolicyResult{}, err
}
x, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
out, err := s.hostCommand(x, nil, "fail2ban-client", "-t")
if err != nil {
s.restoreHostSnapshot(fail2banHostPath, previous, hadPrevious, 0o640)
return PolicyResult{}, fmt.Errorf("Fail2Ban validation failed and configuration was restored: %w", err)
}
if err = s.savePolicy("fail2ban.json", p); err != nil {
return PolicyResult{}, err
}
if active, _ := s.serviceState(ctx, s.initSystem(), "fail2ban"); active {
if v, e := s.hostCommand(x, nil, "fail2ban-client", "reload"); e == nil && v != "" {
out += "\n" + v
}
}
return PolicyResult{OK: true, Message: "Fail2Ban policy validated and applied.", Output: limit(out, 4096)}, nil
}
func (s *Service) AuditdPolicy() AuditdPolicy {
if p, ok := s.loadAuditdPolicy(); ok {
return p
}
return AuditdPolicy{IdentityFiles: true, Sudoers: true, SSH: true, Docker: true, Systemd: true, KernelModules: true, Custom: []AuditWatch{}}
}
func renderAuditd(p AuditdPolicy) (string, error) {
var rules []string
add := func(path, perms, key string) {
rules = append(rules, fmt.Sprintf("-w %s -p %s -k %s", path, perms, key))
}
if p.IdentityFiles {
for _, x := range []string{"/etc/passwd", "/etc/group", "/etc/shadow", "/etc/gshadow"} {
add(x, "wa", "identity")
}
}
if p.Sudoers {
add("/etc/sudoers", "wa", "scope")
add("/etc/sudoers.d", "wa", "scope")
}
if p.SSH {
add("/etc/ssh/sshd_config", "wa", "sshd")
add("/etc/ssh/sshd_config.d", "wa", "sshd")
}
if p.Docker {
add("/var/run/docker.sock", "rwxa", "docker")
add("/etc/docker", "wa", "docker")
}
if p.Systemd {
add("/etc/systemd/system", "wa", "systemd")
add("/usr/lib/systemd/system", "wa", "systemd")
}
if p.KernelModules {
add("/etc/modules-load.d", "wa", "modules")
add("/etc/modprobe.d", "wa", "modules")
}
for i, w := range p.Custom {
clean := filepath.Clean(w.Path)
if !strings.HasPrefix(clean, "/") || clean == "/" || strings.ContainsAny(clean, "\r\n") {
return "", fmt.Errorf("custom watch %d path is invalid", i+1)
}
if !validAuditPerms(w.Permissions) {
return "", fmt.Errorf("custom watch %d permissions must contain only r,w,x,a", i+1)
}
if !auditKeyRE.MatchString(w.Key) {
return "", fmt.Errorf("custom watch %d key is invalid", i+1)
}
add(clean, w.Permissions, w.Key)
}
sort.Strings(rules)
return "# Managed by Dockwatch. Rules are intentionally limited to file watches.\n" + strings.Join(rules, "\n") + "\n", nil
}
func (s *Service) renderAuditdForHost(p AuditdPolicy) (string, []string) {
rendered, err := renderAuditd(p)
if err != nil {
return "", []string{err.Error()}
}
var kept []string
var warnings []string
for _, line := range strings.Split(rendered, "\n") {
if !strings.HasPrefix(line, "-w ") {
if line != "" {
kept = append(kept, line)
}
continue
}
parts := strings.Fields(line)
if len(parts) < 2 {
continue
}
path := parts[1]
if _, err := os.Stat(s.hostPath(path)); err != nil {
warnings = append(warnings, "Skipped missing audit watch path: "+path)
continue
}
kept = append(kept, line)
}
return strings.Join(kept, "\n") + "\n", warnings
}
func (s *Service) ApplyAuditd(ctx context.Context, p AuditdPolicy) (PolicyResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
if err := s.requireChanges(ctx); err != nil {
return PolicyResult{}, err
}
if !fileExists(s.hostPath("/sbin/auditctl")) && !fileExists(s.hostPath("/usr/sbin/auditctl")) {
return PolicyResult{}, errors.New("Linux audit tooling is not installed")
}
if _, err := renderAuditd(p); err != nil {
return PolicyResult{}, err
}
rendered, warnings := s.renderAuditdForHost(p)
previous, hadPrevious := s.readHostSnapshot(auditdHostPath)
if err := s.writeHostManaged(auditdHostPath, []byte(rendered), 0o640); err != nil {
return PolicyResult{}, err
}
x, cancel := context.WithTimeout(ctx, 15*time.Second)
defer cancel()
out, err := s.hostCommand(x, nil, "augenrules", "--check")
if err != nil {
s.restoreHostSnapshot(auditdHostPath, previous, hadPrevious, 0o640)
return PolicyResult{}, fmt.Errorf("audit rule validation failed and configuration was restored: %w", err)
}
load, err := s.hostCommand(x, nil, "augenrules", "--load")
if err != nil {
s.restoreHostSnapshot(auditdHostPath, previous, hadPrevious, 0o640)
return PolicyResult{}, fmt.Errorf("audit rule load failed and configuration was restored: %w", err)
}
if err = s.savePolicy("auditd.json", p); err != nil {
return PolicyResult{}, err
}
if load != "" {
out += "\n" + load
}
return PolicyResult{OK: true, Message: "auditd rules validated and loaded.", Output: limit(out, 4096), Warnings: warnings}, nil
}
func (s *Service) Install(ctx context.Context, component string, in InstallInput) (PolicyResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
if err := s.requirePackages(ctx); err != nil {
return PolicyResult{}, err
}
var pkg, service string
var err error
if component == "firewall" {
pkg, service, err = s.firewallPackage(in.Provider)
} else {
pkg, service, err = s.packageFor(component)
}
if err != nil {
return PolicyResult{}, err
}
pm := s.packageManager()
x, cancel := context.WithTimeout(ctx, 5*time.Minute)
defer cancel()
var out string
switch pm {
case "apt":
if v, e := s.hostCommand(x, nil, "apt-get", "update"); e != nil {
return PolicyResult{}, e
} else {
out = v
}
v, e := s.hostCommand(x, nil, "env", "DEBIAN_FRONTEND=noninteractive", "apt-get", "install", "-y", pkg)
if e != nil {
return PolicyResult{}, e
}
out += "\n" + v
case "dnf", "yum":
v, e := s.hostCommand(x, nil, pm, "install", "-y", pkg)
if e != nil {
return PolicyResult{}, e
}
out = v
case "zypper":
v, e := s.hostCommand(x, nil, "zypper", "--non-interactive", "install", pkg)
if e != nil {
return PolicyResult{}, e
}
out = v
case "apk":
v, e := s.hostCommand(x, nil, "apk", "add", "--no-cache", pkg)
if e != nil {
return PolicyResult{}, e
}
out = v
case "pacman":
v, e := s.hostCommand(x, nil, "pacman", "-Sy", "--noconfirm", pkg)
if e != nil {
return PolicyResult{}, e
}
out = v
default:
return PolicyResult{}, errors.New("unsupported package manager")
}
if component == "firewall" {
// Persist an explicit provider selection even when installation does not
// activate the frontend. This keeps Auto from immediately switching back
// to another already-installed firewall after package installation.
p := s.FirewallPolicy()
if v := normalizeFirewallProvider(in.Provider); v != "" && v != FirewallProviderAuto {
p.Provider = v
p.ResolvedProvider = v
_ = s.savePolicy("firewall.json", p)
}
}
if in.Enable {
if component == "firewall" {
_, _ = s.firewallServiceAction(ctx, "enable")
} else if service != "" {
_, _ = s.serviceActionUnlocked(ctx, component, "enable")
}
}
return PolicyResult{OK: true, Message: component + " installed/updated.", Output: limit(out, 12000)}, nil
}
func (s *Service) ComponentAction(ctx context.Context, component, action string) (PolicyResult, error) {
s.mu.Lock()
defer s.mu.Unlock()
if err := s.requireChanges(ctx); err != nil {
return PolicyResult{}, err
}
out, err := s.serviceActionUnlocked(ctx, component, action)
if err != nil {
return PolicyResult{}, err
}
return PolicyResult{OK: true, Message: component + " " + action + " completed.", Output: limit(out, 4096)}, nil
}
func (s *Service) serviceActionUnlocked(ctx context.Context, component, action string) (string, error) {
if !oneOf(component, "firewall", "fail2ban", "auditd") {
return "", errors.New("unsupported security component")
}
if !oneOf(action, "enable", "disable", "restart", "reload", "start", "stop") {
return "", errors.New("unsupported service action")
}
if component == "firewall" {
return s.firewallServiceAction(ctx, action)
}
service := map[string]string{"fail2ban": "fail2ban", "auditd": "auditd"}[component]
if component == "auditd" && action == "reload" {
x, c := context.WithTimeout(ctx, 15*time.Second)
defer c()
return s.hostCommand(x, nil, "augenrules", "--load")
}
initSystem := s.initSystem()
x, c := context.WithTimeout(ctx, 30*time.Second)
defer c()
switch initSystem {
case "systemd":
args := []string{"systemctl"}
switch action {
case "enable":
args = append(args, "enable", "--now", service)
case "disable":
args = append(args, "disable", "--now", service)
default:
args = append(args, action, service)
}
return s.hostCommand(x, nil, args...)
case "openrc":
switch action {
case "enable":
_, err := s.hostCommand(x, nil, "rc-update", "add", service, "default")
if err != nil {
return "", err
}
return s.hostCommand(x, nil, "rc-service", service, "start")
case "disable":
_, _ = s.hostCommand(x, nil, "rc-service", service, "stop")
return s.hostCommand(x, nil, "rc-update", "del", service, "default")
default:
return s.hostCommand(x, nil, "rc-service", service, action)
}
default:
return "", errors.New("unsupported init system")
}
}
func (s *Service) packageFor(component string) (pkg, service string, err error) {
osid := s.osInfo().ID
switch component {
case "firewall":
return s.firewallPackage(FirewallProviderAuto)
case "fail2ban":
return "fail2ban", "fail2ban", nil
case "auditd":
if oneOf(osid, "rhel", "fedora", "rocky", "almalinux", "centos", "arch", "alpine", "opensuse-leap", "opensuse-tumbleweed", "sles") {
return "audit", "auditd", nil
}
return "auditd", "auditd", nil
default:
return "", "", errors.New("unsupported security component")
}
}
func (s *Service) readHostSnapshot(hostPath string) ([]byte, bool) {
p, err := s.safeHostPath(hostPath)
if err != nil {
return nil, false
}
b, err := os.ReadFile(p)
return b, err == nil
}
func (s *Service) restoreHostSnapshot(hostPath string, data []byte, existed bool, mode os.FileMode) {
if existed {
_ = s.writeHostManaged(hostPath, data, mode)
return
}
_ = s.removeHostManaged(hostPath)
}
func (s *Service) managedConfig(hostPath, rendered string) ManagedConfig {
m := ManagedConfig{Configured: true, Path: hostPath, Hash: hashString(rendered)}
b, err := os.ReadFile(s.hostPath(hostPath))
if err != nil {
m.Drift = true
return m
}
m.Drift = hashBytes(b) != m.Hash
return m
}
func (s *Service) writeHostManaged(hostPath string, data []byte, mode os.FileMode) error {
if !s.cfg.AllowChanges {
return errors.New("host security changes are disabled")
}
p, err := s.safeHostPath(hostPath)
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(p), 0o755); err != nil {
return err
}
if err := rejectSymlinkParents(s.cfg.HostRoot, p); err != nil {
return err
}
if old, err := os.ReadFile(p); err == nil {
if err = s.backup(hostPath, old); err != nil {
return err
}
}
tmp, err := os.CreateTemp(filepath.Dir(p), ".dockwatch-*")
if err != nil {
return err
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
if err = tmp.Chmod(mode); err == nil {
_, err = tmp.Write(data)
}
if syncErr := tmp.Sync(); err == nil {
err = syncErr
}
if closeErr := tmp.Close(); err == nil {
err = closeErr
}
if err != nil {
return err
}
return os.Rename(tmpName, p)
}
func (s *Service) removeHostManaged(hostPath string) error {
p, err := s.safeHostPath(hostPath)
if err != nil {
return err
}
if old, err := os.ReadFile(p); err == nil {
_ = s.backup(hostPath, old)
}
if err := os.Remove(p); err != nil && !os.IsNotExist(err) {
return err
}
return nil
}
func (s *Service) restoreManagedBackup(hostPath string) error {
entries, err := os.ReadDir(s.backupDir())
if err != nil {
return err
}
prefix := strings.ReplaceAll(strings.Trim(hostPath, "/"), "/", "_") + "-"
var names []string
for _, e := range entries {
if strings.HasPrefix(e.Name(), prefix) {
names = append(names, e.Name())
}
}
if len(names) == 0 {
return errors.New("no managed backup available")
}
sort.Strings(names)
b, err := os.ReadFile(filepath.Join(s.backupDir(), names[len(names)-1]))
if err != nil {
return err
}
return s.writeHostManaged(hostPath, b, 0o640)
}
func (s *Service) backup(hostPath string, data []byte) error {
if err := os.MkdirAll(s.backupDir(), 0o700); err != nil {
return err
}
name := strings.ReplaceAll(strings.Trim(hostPath, "/"), "/", "_") + "-" + time.Now().UTC().Format("20060102T150405.000000000Z")
return os.WriteFile(filepath.Join(s.backupDir(), name), data, 0o600)
}
func (s *Service) savePolicy(name string, v any) error {
if err := os.MkdirAll(s.securityDir(), 0o700); err != nil {
return err
}
b, err := json.MarshalIndent(v, "", " ")
if err != nil {
return err
}
tmp, err := os.CreateTemp(s.securityDir(), ".policy-*")
if err != nil {
return err
}
tmpName := tmp.Name()
defer os.Remove(tmpName)
if err = tmp.Chmod(0o600); err == nil {
_, err = tmp.Write(append(b, '\n'))
}
if closeErr := tmp.Close(); err == nil {
err = closeErr
}
if err != nil {
return err
}
return os.Rename(tmpName, s.securityPath(name))
}
func (s *Service) loadPolicy(name string, v any) bool {
b, err := os.ReadFile(s.securityPath(name))
if err != nil {
return false
}
return json.Unmarshal(b, v) == nil
}
func (s *Service) loadFirewallPolicy() (FirewallPolicy, bool) {
var p FirewallPolicy
ok := s.loadPolicy("firewall.json", &p)
return p, ok
}
func (s *Service) loadFail2BanPolicy() (Fail2BanPolicy, bool) {
var p Fail2BanPolicy
ok := s.loadPolicy("fail2ban.json", &p)
return p, ok
}
func (s *Service) loadAuditdPolicy() (AuditdPolicy, bool) {
var p AuditdPolicy
ok := s.loadPolicy("auditd.json", &p)
return p, ok
}
func (s *Service) savePendingFirewall(p pendingFirewall) error {
return s.savePolicy("pending-firewall.json", p)
}
func (s *Service) readPendingFirewall() (pendingFirewall, error) {
var p pendingFirewall
if !s.loadPolicy("pending-firewall.json", &p) {
return p, errors.New("no pending firewall change")
}
return p, nil
}
func (s *Service) securityDir() string { return filepath.Join(s.cfg.DataDir, "host-security") }
func (s *Service) securityPath(name string) string { return filepath.Join(s.securityDir(), name) }
func (s *Service) backupDir() string { return filepath.Join(s.securityDir(), "backups") }
func (s *Service) hostPath(abs string) string {
if s.cfg.HostRoot == "" {
return ""
}
return filepath.Join(s.cfg.HostRoot, strings.TrimPrefix(filepath.Clean(abs), string(filepath.Separator)))
}
func (s *Service) safeHostPath(abs string) (string, error) {
if !strings.HasPrefix(abs, "/") {
return "", errors.New("host path must be absolute")
}
clean := filepath.Clean(abs)
if clean == "/" {
return "", errors.New("host root cannot be a managed file")
}
p := s.hostPath(clean)
root := filepath.Clean(s.cfg.HostRoot)
if p == root || !strings.HasPrefix(p, root+string(filepath.Separator)) {
return "", errors.New("host path escapes HOST_ROOT")
}
return p, nil
}
func rejectSymlinkParents(root, path string) error {
root = filepath.Clean(root)
dir := filepath.Dir(path)
rel, err := filepath.Rel(root, dir)
if err != nil || strings.HasPrefix(rel, "..") {
return errors.New("managed host path escapes HOST_ROOT")
}
cur := root
for _, part := range strings.Split(rel, string(filepath.Separator)) {
if part == "" || part == "." {
continue
}
cur = filepath.Join(cur, part)
fi, err := os.Lstat(cur)
if os.IsNotExist(err) {
continue
}
if err != nil {
return err
}
if fi.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("refusing symlinked host configuration parent %s", cur)
}
}
return nil
}
func fileExists(p string) bool {
if p == "" {
return false
}
fi, err := os.Stat(p)
return err == nil && !fi.IsDir()
}
func dirExists(p string) bool {
if p == "" {
return false
}
fi, err := os.Stat(p)
return err == nil && fi.IsDir()
}
func oneOf(v string, xs ...string) bool {
for _, x := range xs {
if v == x {
return true
}
}
return false
}
func firstLine(v string) string {
if i := strings.IndexByte(v, '\n'); i >= 0 {
return v[:i]
}
return v
}
func limit(v string, n int) string {
if len(v) <= n {
return v
}
return v[:n] + "…"
}
func hashString(v string) string { return hashBytes([]byte(v)) }
func hashBytes(v []byte) string { x := sha256.Sum256(v); return hex.EncodeToString(x[:]) }
func randomID() string {
b := make([]byte, 16)
if _, err := io.ReadFull(rand.Reader, b); err != nil {
return strconv.FormatInt(time.Now().UnixNano(), 36)
}
return hex.EncodeToString(b)
}
func validIPOrCIDR(v string) bool {
v = strings.TrimSpace(v)
if net.ParseIP(v) != nil {
return true
}
_, _, err := net.ParseCIDR(v)
return err == nil
}
func validAuditPerms(v string) bool {
if v == "" || len(v) > 4 {
return false
}
seen := map[rune]bool{}
for _, r := range v {
if !strings.ContainsRune("rwxa", r) || seen[r] {
return false
}
seen[r] = true
}
return true
}