add allow commands list

This commit is contained in:
ars
2026-05-24 19:53:44 +03:00
parent f9aed123fd
commit 394e4e924f
6 changed files with 183 additions and 1 deletions
+15
View File
@@ -29,3 +29,18 @@ client -> server: hello!, i'm $(hostname), uuid= , ts=, stdout= , stderr= , ...
+ возможно на клиенте стоит ограничить набор допустимых команд + возможно на клиенте стоит ограничить набор допустимых команд
+ стоит явно задуматься о шифровании/идентификации сервера + стоит явно задуматься о шифровании/идентификации сервера
## Разрешённые команды
Клиент исполняет только локально разрешённые command presets. Базовая
настройка задаётся переменной окружения `NEXUS_SYNC_ALLOWED_COMMANDS`.
Примеры:
```bash
NEXUS_SYNC_ALLOWED_COMMANDS=hostname,network_interfaces
NEXUS_SYNC_ALLOWED_COMMANDS=full_access
```
`full_access` означает доступ ко всем локально зарегистрированным presets. Это
не разрешает выполнение произвольных shell-строк от сервера.
+17
View File
@@ -0,0 +1,17 @@
from nexus_sync.client.config import (
COMMAND_ACCESS_ENV,
FULL_ACCESS_VALUE,
load_command_access_policy,
)
from nexus_sync.client.execute import (
CommandAccessPolicy,
execute_command,
)
__all__ = [
"COMMAND_ACCESS_ENV",
"CommandAccessPolicy",
"FULL_ACCESS_VALUE",
"execute_command",
"load_command_access_policy",
]
+25
View File
@@ -0,0 +1,25 @@
import os
from collections.abc import Mapping
from nexus_sync.client.execute import CommandAccessPolicy
COMMAND_ACCESS_ENV = "NEXUS_SYNC_ALLOWED_COMMANDS"
FULL_ACCESS_VALUE = "full_access"
def load_command_access_policy(
env: Mapping[str, str] = os.environ,
*,
default: CommandAccessPolicy | None = None,
) -> CommandAccessPolicy:
raw_value = env.get(COMMAND_ACCESS_ENV)
if raw_value is None or not raw_value.strip():
return default or CommandAccessPolicy.deny_all()
command_names = [item.strip() for item in raw_value.split(",") if item.strip()]
if len(command_names) == 1 and command_names[0].lower() == FULL_ACCESS_VALUE:
return CommandAccessPolicy.allow_all()
if any(command_name.lower() == FULL_ACCESS_VALUE for command_name in command_names):
raise ValueError(f"{FULL_ACCESS_VALUE} cannot be mixed with explicit command names")
return CommandAccessPolicy.allow(command_names)
+27
View File
@@ -1,5 +1,7 @@
import platform import platform
import subprocess import subprocess
from collections.abc import Iterable
from dataclasses import dataclass, field
from datetime import UTC, datetime from datetime import UTC, datetime
from typing import Any, Callable, Mapping, Sequence from typing import Any, Callable, Mapping, Sequence
@@ -10,6 +12,27 @@ PresetBuilder = Callable[[Mapping[str, Any]], Sequence[str]]
DEFAULT_OUTPUT_LIMIT_BYTES = 64 * 1024 DEFAULT_OUTPUT_LIMIT_BYTES = 64 * 1024
@dataclass(frozen=True)
class CommandAccessPolicy:
allowed_commands: frozenset[str] = field(default_factory=frozenset)
full_access: bool = False
@classmethod
def allow(cls, command_names: Iterable[str]) -> "CommandAccessPolicy":
return cls(allowed_commands=frozenset(command_names))
@classmethod
def allow_all(cls) -> "CommandAccessPolicy":
return cls(full_access=True)
@classmethod
def deny_all(cls) -> "CommandAccessPolicy":
return cls()
def allows(self, command_name: str) -> bool:
return self.full_access or command_name in self.allowed_commands
def _reject(command: Command, message: str) -> CommandResult: def _reject(command: Command, message: str) -> CommandResult:
now = datetime.now(UTC) now = datetime.now(UTC)
return CommandResult( return CommandResult(
@@ -45,6 +68,7 @@ DEFAULT_PRESETS: dict[str, PresetBuilder] = {
"hostname": _hostname, "hostname": _hostname,
"network_interfaces": _network_interfaces, "network_interfaces": _network_interfaces,
} }
DEFAULT_COMMAND_ACCESS_POLICY = CommandAccessPolicy.allow_all()
def execute_command( def execute_command(
@@ -52,6 +76,7 @@ def execute_command(
*, *,
stdin: str | None = None, stdin: str | None = None,
presets: Mapping[str, PresetBuilder] = DEFAULT_PRESETS, presets: Mapping[str, PresetBuilder] = DEFAULT_PRESETS,
access_policy: CommandAccessPolicy = DEFAULT_COMMAND_ACCESS_POLICY,
output_limit_bytes: int = DEFAULT_OUTPUT_LIMIT_BYTES, output_limit_bytes: int = DEFAULT_OUTPUT_LIMIT_BYTES,
) -> CommandResult: ) -> CommandResult:
if command.kind != CommandKind.EXEC: if command.kind != CommandKind.EXEC:
@@ -60,6 +85,8 @@ def execute_command(
builder = presets.get(command.name) builder = presets.get(command.name)
if builder is None: if builder is None:
return _reject(command, f"unknown command preset: {command.name}") return _reject(command, f"unknown command preset: {command.name}")
if not access_policy.allows(command.name):
return _reject(command, f"command preset is not allowed: {command.name}")
try: try:
argv = list(builder(command.args)) argv = list(builder(command.args))
+42
View File
@@ -0,0 +1,42 @@
import pytest
from nexus_sync.client.config import COMMAND_ACCESS_ENV, load_command_access_policy
from nexus_sync.client.execute import CommandAccessPolicy
def test_load_command_access_policy_defaults_to_deny_all() -> None:
policy = load_command_access_policy({})
assert not policy.full_access
assert not policy.allows("hostname")
def test_load_command_access_policy_uses_default_when_env_is_missing() -> None:
default = CommandAccessPolicy.allow(["hostname"])
policy = load_command_access_policy({}, default=default)
assert policy == default
def test_load_command_access_policy_supports_command_allowlist() -> None:
policy = load_command_access_policy(
{COMMAND_ACCESS_ENV: "hostname, network_interfaces"},
)
assert policy.allows("hostname")
assert policy.allows("network_interfaces")
assert not policy.allows("unknown")
def test_load_command_access_policy_supports_full_access() -> None:
policy = load_command_access_policy({COMMAND_ACCESS_ENV: "full_access"})
assert policy.full_access
assert policy.allows("hostname")
assert policy.allows("future_registered_preset")
def test_load_command_access_policy_rejects_mixed_full_access() -> None:
with pytest.raises(ValueError, match="cannot be mixed"):
load_command_access_policy({COMMAND_ACCESS_ENV: "hostname,full_access"})
+57 -1
View File
@@ -1,6 +1,6 @@
import subprocess import subprocess
from nexus_sync.client.execute import execute_command from nexus_sync.client.execute import CommandAccessPolicy, execute_command
from nexus_sync.common import Command, CommandKind, CommandResultStatus from nexus_sync.common import Command, CommandKind, CommandResultStatus
@@ -51,6 +51,62 @@ def test_execute_command_rejects_unknown_preset() -> None:
assert "unknown command preset" in result.stderr assert "unknown command preset" in result.stderr
def test_execute_command_allows_selected_preset(monkeypatch) -> None:
calls = []
def fake_run(argv, **kwargs):
calls.append(argv)
return subprocess.CompletedProcess(argv, 0, stdout="host\n", stderr="")
monkeypatch.setattr(subprocess, "run", fake_run)
result = execute_command(
_command(),
access_policy=CommandAccessPolicy.allow(["hostname"]),
)
assert result.status == CommandResultStatus.SUCCEEDED
assert calls == [["hostname"]]
def test_execute_command_rejects_disallowed_preset(monkeypatch) -> None:
calls = []
def fake_run(argv, **kwargs):
calls.append(argv)
return subprocess.CompletedProcess(argv, 0, stdout="host\n", stderr="")
monkeypatch.setattr(subprocess, "run", fake_run)
result = execute_command(
_command(name="hostname"),
access_policy=CommandAccessPolicy.allow(["network_interfaces"]),
)
assert result.status == CommandResultStatus.REJECTED
assert result.return_code is None
assert "not allowed" in result.stderr
assert calls == []
def test_execute_command_full_access_allows_registered_presets(monkeypatch) -> None:
calls = []
def fake_run(argv, **kwargs):
calls.append(argv)
return subprocess.CompletedProcess(argv, 0, stdout="host\n", stderr="")
monkeypatch.setattr(subprocess, "run", fake_run)
result = execute_command(
_command(name="hostname"),
access_policy=CommandAccessPolicy.allow_all(),
)
assert result.status == CommandResultStatus.SUCCEEDED
assert calls == [["hostname"]]
def test_execute_command_maps_non_zero_exit_to_failed(monkeypatch) -> None: def test_execute_command_maps_non_zero_exit_to_failed(monkeypatch) -> None:
def fake_run(argv, **kwargs): def fake_run(argv, **kwargs):
return subprocess.CompletedProcess(argv, 2, stdout="", stderr="failed\n") return subprocess.CompletedProcess(argv, 2, stdout="", stderr="failed\n")