Files
truf-server/deploy/fail2ban/truf_caddy_admin_denylist.py
2026-09-30 20:30:56 +03:00

439 lines
16 KiB
Python

#!/usr/bin/env python3
"""Maintain the Caddy admin-only IP denylist with durable expiry state."""
import argparse
from contextlib import contextmanager
import hashlib
import ipaddress
import json
import os
from pathlib import Path
import re
import stat
import subprocess
import sys
import tempfile
import time
BAN_SECONDS = 24 * 60 * 60
MAX_BANS = 4096
MAX_FILE_BYTES = 512 * 1024
STATE_VERSION = 1
EMPTY_SNIPPET = "# Managed by truf-caddy-admin-denylist. Admin-route import only.\n"
SHA256_RE = re.compile(r"^[0-9a-f]{64}$")
PROFILE_PATH = Path("/etc/truf/deployment-profile")
STANDALONE_PROFILE = "standalone-edge-v1"
SHARED_HOST_PROFILE = "shared-host-edge-v1"
class UpdateError(RuntimeError):
pass
class CommandFailure(UpdateError):
pass
class RollbackFailure(UpdateError):
pass
def canonical_ip(value):
text = str(value or "")
if not text or len(text) > 64 or "%" in text or any(char.isspace() for char in text):
raise UpdateError("invalid IP address")
try:
address = ipaddress.ip_address(text)
except ValueError as exc:
raise UpdateError("invalid IP address") from exc
if address.is_unspecified or address.is_multicast:
raise UpdateError("unsupported IP address")
return address.compressed.lower()
def render_snippet(bans, matcher="remote_ip"):
if matcher not in {"remote_ip", "client_ip"}:
raise UpdateError("unsupported denylist matcher")
addresses = sorted(
(ipaddress.ip_address(address) for address in bans),
key=lambda address: (address.version, int(address)),
)
if not addresses:
return EMPTY_SNIPPET.encode("ascii")
lines = [EMPTY_SNIPPET.rstrip("\n")]
for offset in range(0, len(addresses), 64):
name = f"truf_admin_denied_{offset // 64:04d}"
values = " ".join(address.compressed.lower() for address in addresses[offset:offset + 64])
lines.append(f"@{name} {matcher} {values}")
lines.append(f'respond @{name} "" 403')
return ("\n".join(lines) + "\n").encode("ascii")
def _digest(content):
return hashlib.sha256(content).hexdigest()
def _check_parent(path):
parent = path.parent
details = parent.lstat()
if not stat.S_ISDIR(details.st_mode) or stat.S_ISLNK(details.st_mode):
raise UpdateError("managed parent must be a real directory")
if os.name == "posix" and stat.S_IMODE(details.st_mode) & 0o002:
raise UpdateError("managed parent must not be world-writable")
def _read_optional(path):
_check_parent(path)
try:
details = path.lstat()
except FileNotFoundError:
return None
if not stat.S_ISREG(details.st_mode) or stat.S_ISLNK(details.st_mode):
raise UpdateError("managed path must be a regular file")
flags = os.O_RDONLY | getattr(os, "O_BINARY", 0) | getattr(os, "O_NOFOLLOW", 0)
descriptor = os.open(path, flags)
try:
current = os.fstat(descriptor)
if not stat.S_ISREG(current.st_mode) or current.st_size > MAX_FILE_BYTES:
raise UpdateError("managed file is invalid or too large")
chunks = []
remaining = MAX_FILE_BYTES + 1
while remaining:
chunk = os.read(descriptor, min(65536, remaining))
if not chunk:
break
chunks.append(chunk)
remaining -= len(chunk)
content = b"".join(chunks)
if len(content) > MAX_FILE_BYTES:
raise UpdateError("managed file is too large")
return content
finally:
os.close(descriptor)
def _sync_parent(parent):
if os.name != "posix":
return
descriptor = os.open(parent, os.O_RDONLY | getattr(os, "O_DIRECTORY", 0))
try:
os.fsync(descriptor)
finally:
os.close(descriptor)
def _atomic_write(path, content, mode):
_check_parent(path)
if len(content) > MAX_FILE_BYTES:
raise UpdateError("managed content is too large")
try:
existing = path.lstat()
except FileNotFoundError:
existing = None
if existing is not None and (not stat.S_ISREG(existing.st_mode) or stat.S_ISLNK(existing.st_mode)):
raise UpdateError("managed path must be a regular file")
descriptor, temporary = tempfile.mkstemp(prefix=".truf-denylist-", dir=path.parent)
temporary_path = Path(temporary)
try:
if hasattr(os, "fchmod"):
os.fchmod(descriptor, mode)
else:
os.chmod(temporary_path, mode)
with os.fdopen(descriptor, "wb", closefd=True) as handle:
descriptor = -1
handle.write(content)
handle.flush()
os.fsync(handle.fileno())
os.replace(temporary_path, path)
_sync_parent(path.parent)
finally:
if descriptor >= 0:
os.close(descriptor)
try:
temporary_path.unlink()
except FileNotFoundError:
pass
def _restore(path, content, mode):
if content is not None:
_atomic_write(path, content, mode)
return
try:
details = path.lstat()
except FileNotFoundError:
return
if not stat.S_ISREG(details.st_mode) or stat.S_ISLNK(details.st_mode):
raise UpdateError("managed path changed during rollback")
path.unlink()
_sync_parent(path.parent)
@contextmanager
def _exclusive_lock(path):
_check_parent(path)
flags = os.O_RDWR | os.O_CREAT | getattr(os, "O_BINARY", 0) | getattr(os, "O_NOFOLLOW", 0)
descriptor = os.open(path, flags, 0o600)
try:
details = os.fstat(descriptor)
if not stat.S_ISREG(details.st_mode):
raise UpdateError("lock path must be a regular file")
if os.name == "posix":
import fcntl
fcntl.flock(descriptor, fcntl.LOCK_EX)
yield
finally:
os.close(descriptor)
def _load_state(content):
if content is None:
return {"version": STATE_VERSION, "bans": {}, "applied_sha256": ""}
try:
value = json.loads(content.decode("ascii"))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise UpdateError("denylist state is not valid JSON") from exc
if not isinstance(value, dict) or set(value) != {"version", "bans", "applied_sha256"}:
raise UpdateError("denylist state has an invalid schema")
if value["version"] != STATE_VERSION or not isinstance(value["bans"], dict):
raise UpdateError("denylist state has an unsupported version")
if len(value["bans"]) > MAX_BANS:
raise UpdateError("denylist state exceeds its entry bound")
applied = value["applied_sha256"]
if not isinstance(applied, str) or (applied and not SHA256_RE.fullmatch(applied)):
raise UpdateError("denylist state has an invalid applied digest")
bans = {}
for address, expires_at in value["bans"].items():
canonical = canonical_ip(address)
if canonical != address or isinstance(expires_at, bool) or not isinstance(expires_at, int):
raise UpdateError("denylist state has a noncanonical entry")
if expires_at <= 0 or expires_at > 253402300799:
raise UpdateError("denylist state has an invalid expiry")
bans[canonical] = expires_at
return {"version": STATE_VERSION, "bans": bans, "applied_sha256": applied}
def _encode_state(state):
return (json.dumps(state, sort_keys=True, separators=(",", ":")) + "\n").encode("ascii")
def _subprocess_runner(command):
environment = {
"HOME": "/root",
"LANG": "C.UTF-8",
"LC_ALL": "C.UTF-8",
"PATH": "/usr/local/sbin:/usr/local/bin:/usr/sbin:/usr/bin:/sbin:/bin",
}
try:
completed = subprocess.run(
command,
stdin=subprocess.DEVNULL,
stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL,
env=environment,
timeout=45,
check=False,
)
except (OSError, subprocess.SubprocessError):
return False
return completed.returncode == 0
class DenylistUpdater:
def __init__(
self,
state_path,
snippet_path,
project_directory="/opt/truf",
env_file="/etc/truf-edge/edge.env",
profile=None,
runner=None,
clock=None,
):
self.state_path = Path(state_path)
self.snippet_path = Path(snippet_path)
self.lock_path = self.state_path.with_suffix(self.state_path.suffix + ".lock")
if profile is None:
try:
profile = PROFILE_PATH.read_text(encoding="ascii").strip()
except FileNotFoundError:
profile = STANDALONE_PROFILE
except (OSError, UnicodeError):
raise UpdateError("deployment profile is unreadable") from None
if profile not in {STANDALONE_PROFILE, SHARED_HOST_PROFILE}:
raise UpdateError("deployment profile is unsupported")
compose_file = (
"compose.shared-host.yaml"
if profile == SHARED_HOST_PROFILE else "compose.edge.yaml"
)
caddyfile = (
"/etc/caddy/Caddyfile.shared-host"
if profile == SHARED_HOST_PROFILE else "/etc/caddy/Caddyfile"
)
self.matcher = "client_ip" if profile == SHARED_HOST_PROFILE else "remote_ip"
compose = (
"docker", "compose", "--ansi", "never", "--env-file", str(env_file),
"--project-directory", str(project_directory),
"--file", str(Path(project_directory) / "compose.yaml"),
"--file", str(Path(project_directory) / compose_file),
)
self.validate_command = compose + (
"exec", "-T", "edge", "caddy", "validate",
"--config", caddyfile, "--adapter", "caddyfile",
)
self.reload_command = compose + (
"exec", "-T", "edge", "caddy", "reload",
"--config", caddyfile, "--adapter", "caddyfile",
"--address", "unix//run/caddy-admin.sock",
)
self.runner = runner or _subprocess_runner
self.clock = clock or time.time
def _run(self, command, phase):
try:
succeeded = self.runner(command)
except Exception as exc:
raise CommandFailure(f"{phase} command failed") from exc
if not succeeded:
raise CommandFailure(f"{phase} command failed")
def update(self, operation, address=None):
if operation not in {"ban", "unban", "expire", "status"}:
raise UpdateError("unsupported operation")
canonical = canonical_ip(address) if operation in {"ban", "unban"} else None
now = int(self.clock())
if now <= 0:
raise UpdateError("system clock is invalid")
with _exclusive_lock(self.lock_path):
old_state_content = _read_optional(self.state_path)
old_snippet_content = _read_optional(self.snippet_path)
state = _load_state(old_state_content)
bans = {
ip: expires_at for ip, expires_at in state["bans"].items()
if expires_at > now
}
expired = len(state["bans"]) - len(bans)
if operation == "ban":
if canonical not in bans and len(bans) >= MAX_BANS:
raise UpdateError("denylist entry bound reached")
bans[canonical] = max(bans.get(canonical, 0), now + BAN_SECONDS)
elif operation == "unban":
bans.pop(canonical, None)
desired_snippet = render_snippet(bans, self.matcher)
desired_digest = _digest(desired_snippet)
pending_state = {
"version": STATE_VERSION,
"bans": bans,
"applied_sha256": state["applied_sha256"],
}
pending_content = _encode_state(pending_state)
needs_reload = (
old_snippet_content != desired_snippet
or state["applied_sha256"] != desired_digest
)
needs_state_write = old_state_content != pending_content
if needs_reload:
reload_attempted = False
try:
_atomic_write(self.state_path, pending_content, 0o600)
_atomic_write(self.snippet_path, desired_snippet, 0o640)
self._run(self.validate_command, "validation")
reload_attempted = True
self._run(self.reload_command, "reload")
pending_state["applied_sha256"] = desired_digest
_atomic_write(self.state_path, _encode_state(pending_state), 0o600)
except Exception as original:
try:
_restore(self.state_path, old_state_content, 0o600)
_restore(self.snippet_path, old_snippet_content, 0o640)
if reload_attempted:
self._run(self.validate_command, "rollback validation")
self._run(self.reload_command, "rollback reload")
except Exception as rollback:
raise RollbackFailure("denylist rollback failed") from rollback
if isinstance(original, UpdateError):
raise
raise UpdateError("denylist update failed") from original
elif needs_state_write:
pending_state["applied_sha256"] = desired_digest
_atomic_write(self.state_path, _encode_state(pending_state), 0o600)
return {
"operation": operation,
"ip": canonical,
"expired": expired,
"bans": dict(bans),
}
def _emit(result):
bans = result["bans"]
if result["operation"] == "status":
addresses = sorted(
bans, key=lambda value: (ipaddress.ip_address(value).version, int(ipaddress.ip_address(value)))
)
payload = {
"active": len(addresses),
"bans": [
{"ip": address, "expires_at": bans[address]}
for address in addresses[:256]
],
"event": "admin_denylist_status",
"truncated": len(addresses) > 256,
}
else:
payload = {
"active": len(bans),
"event": "admin_denylist_" + result["operation"],
"expired": result["expired"],
}
if result["ip"] is not None:
payload["ip"] = result["ip"]
print(json.dumps(payload, sort_keys=True, separators=(",", ":")), flush=True)
def parse_args(argv=None):
parser = argparse.ArgumentParser(allow_abbrev=False)
parser.add_argument("--state-path", default="/var/lib/truf-edge/admin-denylist.json")
parser.add_argument("--snippet-path", default="/etc/truf-edge/denylist/admin-denylist.caddy")
parser.add_argument("--project-directory", default="/opt/truf")
parser.add_argument("--env-file", default="/etc/truf-edge/edge.env")
commands = parser.add_subparsers(dest="operation", required=True)
for name in ("ban", "unban"):
command = commands.add_parser(name, allow_abbrev=False)
command.add_argument("ip")
commands.add_parser("expire", allow_abbrev=False)
commands.add_parser("status", allow_abbrev=False)
return parser.parse_args(argv)
def main(argv=None):
args = parse_args(argv)
updater = DenylistUpdater(
args.state_path,
args.snippet_path,
project_directory=args.project_directory,
env_file=args.env_file,
)
try:
result = updater.update(args.operation, getattr(args, "ip", None))
except Exception as exc:
payload = {
"event": "admin_denylist_error",
"operation": args.operation,
"reason": type(exc).__name__,
}
print(json.dumps(payload, sort_keys=True, separators=(",", ":")), file=sys.stderr)
return 1
_emit(result)
return 0
if __name__ == "__main__":
raise SystemExit(main())