atlas-iac/testing/tests/test_hermes_node_io_coverage.py
jenkins 0a4ff7f1f6 hermes: read node access keys through CSI symlinks
The Vault CSI secrets-store projects /vault/secrets/node-ssh-public-key
as a symlink chain (file -> ..data/file -> ..<timestamp>/file), which
the O_NOFOLLOW hardened reader refuses with ELOOP, crash-looping the
hermes-node-ssh-access key reconciler on most nodes.

Resolve the projection safely instead of relaxing the reader: realpath
the secret strictly (loops and dangling links become a clean refusal),
refuse any resolution that escapes the read-only mount, then apply the
unchanged O_NOFOLLOW regular-file read to the resolved target. Host
account databases and authorized_keys keep the strict no-symlink read.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-18 16:28:25 -03:00

250 lines
8.6 KiB
Python

"""Behavioral branch coverage for crash-safe node account file I/O."""
from __future__ import annotations
import os
import stat
import sys
from pathlib import Path
import pytest
from testing.tests.test_hermes_node_account_support import _load
def _io_module():
hardening = _load()
return sys.modules[hardening.atomic_write.__module__]
def _snapshot(module, path: Path, *, value=None, xattrs=()):
metadata = path.stat()
content = path.read_bytes() if value is None else value
return module.FileSnapshot(
value=content,
device=metadata.st_dev,
inode=metadata.st_ino,
mode=stat.S_IMODE(metadata.st_mode),
uid=metadata.st_uid,
gid=metadata.st_gid,
size=len(content),
mtime_ns=metadata.st_mtime_ns,
ctime_ns=metadata.st_ctime_ns,
xattrs=tuple(xattrs),
)
def test_xattr_reader_uses_compatibility_fallback(monkeypatch):
module = _io_module()
calls = []
def listxattr(target, **kwargs):
calls.append((target, kwargs))
if kwargs:
raise TypeError("fd API")
return ["user.safe"]
def getxattr(target, name, **kwargs):
if kwargs:
raise TypeError("fd API")
return b"value"
monkeypatch.setattr(module.os, "listxattr", listxattr)
monkeypatch.setattr(module.os, "getxattr", getxattr)
assert module._xattrs(7) == (("user.safe", b"value"),)
assert len(calls) == 2
@pytest.mark.parametrize("failure", ["count", "name", "value", "total"])
def test_xattr_reader_enforces_all_bounds(monkeypatch, failure):
module = _io_module()
if failure == "count":
names = [f"user.{index}" for index in range(module.MAX_XATTRS + 1)]
elif failure == "name":
names = ["x" * (module.MAX_XATTR_NAME + 1)]
elif failure == "total":
names = ["user.one", "user.two", "user.three", "user.four", "user.five"]
else:
names = ["user.large"]
monkeypatch.setattr(module.os, "listxattr", lambda *_a, **_k: names)
size = module.MAX_XATTR_VALUE + 1 if failure == "value" else module.MAX_XATTR_VALUE
monkeypatch.setattr(module.os, "getxattr", lambda *_a, **_k: b"x" * size)
with pytest.raises(module.HardeningError):
module._xattrs(7)
def test_restore_xattrs_removes_unexpected_and_sets_expected(monkeypatch):
module = _io_module()
monkeypatch.setattr(
module, "_xattrs", lambda _fd: (("user.old", b"old"), ("user.keep", b"before"))
)
removed = []
written = []
monkeypatch.setattr(
module.os, "removexattr", lambda fd, name: removed.append((fd, name))
)
monkeypatch.setattr(
module.os, "setxattr", lambda fd, name, value: written.append((fd, name, value))
)
module._restore_xattrs(4, (("user.keep", b"after"),))
assert removed == [(4, "user.old")]
assert written == [(4, "user.keep", b"after")]
def test_read_regular_rejects_directory_oversize_and_changed_read(
tmp_path, monkeypatch
):
module = _io_module()
with pytest.raises(module.HardeningError, match="unsafe regular"):
module.read_regular(tmp_path, 10)
source = tmp_path / "file"
source.write_bytes(b"12345")
with pytest.raises(module.HardeningError, match="unsafe regular"):
module.read_regular(source, 4)
monkeypatch.setattr(module.os, "read", lambda _fd, _size: b"1")
with pytest.raises(module.HardeningError, match="short file read"):
module.read_regular(source, 10)
def test_read_mounted_secret_follows_csi_projection_inside_mount(tmp_path):
module = _io_module()
mount = tmp_path / "vault-secrets"
data = mount / "..2026_08_18.1"
data.mkdir(parents=True)
(data / "node-ssh-public-key").write_bytes(b"ssh-ed25519 synthetic\n")
(mount / "..data").symlink_to(data.name)
(mount / "node-ssh-public-key").symlink_to(Path("..data") / "node-ssh-public-key")
snapshot = module.read_mounted_secret(mount / "node-ssh-public-key", mount, 64)
assert snapshot.value == b"ssh-ed25519 synthetic\n"
def test_read_mounted_secret_reads_plain_regular_file_unchanged(tmp_path):
module = _io_module()
mount = tmp_path / "vault-secrets"
mount.mkdir()
secret = mount / "node-ssh-public-key"
secret.write_bytes(b"ssh-ed25519 synthetic\n")
snapshot = module.read_mounted_secret(secret, mount, 64)
assert snapshot == module.read_regular(secret, 64)
def test_read_mounted_secret_refuses_escape_from_mount(tmp_path):
module = _io_module()
mount = tmp_path / "vault-secrets"
mount.mkdir()
outside = tmp_path / "outside"
outside.write_bytes(b"planted\n")
(mount / "node-ssh-public-key").symlink_to(outside)
with pytest.raises(module.HardeningError, match="escapes its mount"):
module.read_mounted_secret(mount / "node-ssh-public-key", mount, 64)
def test_read_mounted_secret_refuses_symlink_loop(tmp_path):
module = _io_module()
mount = tmp_path / "vault-secrets"
mount.mkdir()
(mount / "node-ssh-public-key").symlink_to("node-ssh-public-key")
with pytest.raises(module.HardeningError, match="unresolvable mounted secret"):
module.read_mounted_secret(mount / "node-ssh-public-key", mount, 64)
def test_read_mounted_secret_reports_missing_secret_as_not_found(tmp_path):
module = _io_module()
mount = tmp_path / "vault-secrets"
mount.mkdir()
with pytest.raises(FileNotFoundError):
module.read_mounted_secret(mount / "node-ssh-public-key", mount, 64)
def test_read_mounted_secret_refuses_non_regular_resolved_target(tmp_path):
module = _io_module()
mount = tmp_path / "vault-secrets"
data = mount / "..data"
data.mkdir(parents=True)
(mount / "node-ssh-public-key").symlink_to("..data")
with pytest.raises(module.HardeningError, match="unsafe regular"):
module.read_mounted_secret(mount / "node-ssh-public-key", mount, 64)
def test_assert_unchanged_rejects_changed_snapshot(tmp_path):
module = _io_module()
source = tmp_path / "passwd"
source.write_bytes(b"old")
snapshot = module.read_regular(source, 16)
source.write_bytes(b"new")
with pytest.raises(module.HardeningError, match="concurrent"):
module.assert_unchanged(source, snapshot, 16)
def test_atomic_write_short_write_unlinks_temporary(tmp_path, monkeypatch):
module = _io_module()
target = tmp_path / "passwd"
target.write_bytes(b"old")
snapshot = module.read_regular(target, 16)
monkeypatch.setattr(module.os, "write", lambda _fd, _value: 0)
with pytest.raises(module.HardeningError, match="short atomic write"):
module.atomic_write(target, b"new", snapshot)
assert target.read_bytes() == b"old"
assert not list(tmp_path.glob(".passwd.hermes-*"))
def test_backup_existing_is_validated_and_short_write_is_removed(tmp_path, monkeypatch):
module = _io_module()
source = tmp_path / "passwd"
source.write_bytes(b"old")
snapshot = module.read_regular(source, 16)
backup = tmp_path / "passwd.hermes-boundary-backup"
backup.write_bytes(b"preserved")
assert module.backup_once(source, snapshot, 16) == backup
backup.unlink()
monkeypatch.setattr(module.os, "write", lambda _fd, _value: 0)
with pytest.raises(module.HardeningError, match="short backup write"):
module.backup_once(source, snapshot, 16)
assert not backup.exists()
def test_account_lock_rejects_unsafe_metadata(tmp_path):
module = _io_module()
lock = tmp_path / ".pwd.lock"
lock.write_text("", encoding="utf-8")
lock.chmod(0o644)
with (
pytest.raises(module.HardeningError, match="unsafe standard account lock"),
module.account_lock(tmp_path, expected_uid=os.getuid()),
):
pass
def test_account_lock_retries_then_succeeds(tmp_path, monkeypatch):
module = _io_module()
attempts = []
def lockf(_descriptor, operation):
attempts.append(operation)
if len(attempts) == 1:
raise BlockingIOError()
monkeypatch.setattr(module.fcntl, "lockf", lockf)
monkeypatch.setattr(module.time, "sleep", lambda _seconds: None)
with module.account_lock(tmp_path, expected_uid=os.getuid()):
assert attempts
assert len(attempts) == 3
def test_account_lock_reports_timeout(tmp_path, monkeypatch):
module = _io_module()
def lockf(_descriptor, operation):
if operation != module.fcntl.LOCK_UN:
raise BlockingIOError()
monkeypatch.setattr(module.fcntl, "lockf", lockf)
ticks = iter((0.0, 16.0))
monkeypatch.setattr(module.time, "monotonic", lambda: next(ticks))
with (
pytest.raises(module.HardeningError, match="lock is busy"),
module.account_lock(tmp_path, expected_uid=os.getuid()),
):
pass