diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml index dd86ba23..c98b7c95 100644 --- a/.github/workflows/publish.yml +++ b/.github/workflows/publish.yml @@ -68,6 +68,7 @@ jobs: uses: softprops/action-gh-release@v2 with: tag_name: ${{ env.VERSION }} + target_commitish: ${{ github.sha }} name: Release ${{ env.VERSION }} draft: false prerelease: ${{ steps.release_info.outputs.prerelease }} diff --git a/README.md b/README.md index 5a2ef37f..9b629489 100644 --- a/README.md +++ b/README.md @@ -20,7 +20,8 @@ SKALE Node CLI, part of the SKALE suite of validator tools, is the command line 4. [sChain commands (Standard)](#schain-commands-standard) 5. [Health commands (Standard)](#health-commands-standard) 6. [SSL commands (Standard)](#ssl-commands-standard) - 7. [Logs commands (Standard)](#logs-commands-standard) + 7. [SGX commands (Standard)](#sgx-commands-standard) + 8. [Logs commands (Standard)](#logs-commands-standard) 3. [Passive Node Usage (`skale` - Passive Build)](#passive-node-usage-skale---passive-build) 1. [Top level commands (Passive)](#top-level-commands-passive) 2. [Passive node commands](#passive-node-commands) @@ -32,8 +33,9 @@ SKALE Node CLI, part of the SKALE suite of validator tools, is the command line 5. [Fair Wallet commands](#fair-wallet-commands) 6. [Fair Logs commands](#fair-logs-commands) 7. [Fair SSL commands](#fair-ssl-commands) - 8. [Fair Staking commands](#fair-staking-commands) - 9. [Passive Fair Node commands](#passive-fair-node-commands) + 8. [Fair SGX commands](#fair-sgx-commands) + 9. [Fair Staking commands](#fair-staking-commands) + 10. [Passive Fair Node commands](#passive-fair-node-commands) 5. [Exit codes](#exit-codes) 6. [Development](#development) @@ -119,6 +121,29 @@ Options: > Prefix: `skale node` +#### Configure firewall + +Firewall setup does not automatically open monitoring ports 9100 and 8080. +Reconfiguration removes their legacy allow rules from the managed base chain +and saves the updated rules for reboot. `MONITORING_CONTAINERS` controls the +containers only; the firewall's `--monitoring` option has been removed. +Explicit rules in `/etc/nft.conf.d/skale/user.conf` remain under operator control. + +SSH allow rules use all listening ports reported by `sshd -T`, including ports +configured through included files and `ListenAddress`. If detection fails, the +command stops before enabling the default-drop policy. + +For a port configured through sshd command-line options or socket activation, +set the `SSH_PORT` environment variable explicitly when running commands that +configure the firewall (including node init and update): + +```shell +sudo SSH_PORT=2222 skale node configure-firewall +``` + +The override replaces automatic detection and accepts one port from 1 to 65535. +It must be passed in the command environment, not only in the node settings file. + #### Node information Get base info about the standard SKALE node. @@ -476,6 +501,48 @@ Options: * `--port/-p` - Port to start healthcheck server (default: `4536`). * `--no-client` - Skip client connection (only make sure server started without errors). +### SGX commands (Standard) + +> Prefix: `skale sgx` + +Manage the client certificate that node services use to authenticate to the SGX wallet. +The files live in `~/.skale/node_data/sgx_certs` and are read by the SKALE containers. +These commands work directly with those files and the SGX server; they do not go through +the node API. + +#### SGX certificate status + +Show the certificate files, the certificate details and its expiry. + +```shell +skale sgx status [--json] [--check] +``` + +Options: + +* `--json` - Show data in JSON format. +* `--check` - Also verify that the SGX server accepts the certificate. + +#### Renew SGX certificate + +Issue a new client certificate from the SGX server and install it. The current +certificate stays in place until the new one is signed and verified against the server. +The previous files are copied to `~/.skale/node_data/sgx_certs_backup/`. +If the SGX server requires manual approval of signing requests, the command prints the +request hash and waits until it is approved. Node services pick up the new certificate +on their next SGX request; no restart is needed. `skale health sgx` confirms afterwards +that node services reach the SGX server. + +```shell +skale sgx renew [--yes] [--timeout ] [--skip-verify] +``` + +Options: + +* `--yes` - Do not ask for confirmation. +* `--timeout` - Seconds to wait for the SGX server to sign the request (default: `600`). +* `--skip-verify` - Install the certificate without testing it against the SGX server first. + ### Logs commands (Standard) > Prefix: `skale logs` @@ -1096,6 +1163,25 @@ Options: * `--no-client` - Skip client connection for openssl check. * `--no-wss` - Skip WSS server starting for skaled check. +### Fair SGX commands + +> Prefix: `fair sgx` + +Manage the client certificate that node services use to authenticate to the SGX wallet. +See [SGX commands (Standard)](#sgx-commands-standard) for details; the behaviour is the same. + +#### Fair SGX Status + +```shell +fair sgx status [--json] [--check] +``` + +#### Fair SGX Renew + +```shell +fair sgx renew [--yes] [--timeout ] [--skip-verify] +``` + ### Fair Staking commands > Prefix: `fair staking` diff --git a/node_cli/cli/node.py b/node_cli/cli/node.py index b34fd155..20861a55 100644 --- a/node_cli/cli/node.py +++ b/node_cli/cli/node.py @@ -239,7 +239,6 @@ def check(network): @node.command(help='Reconfigure nftables rules') -@click.option('--monitoring', is_flag=True) @click.option( '--yes', is_flag=True, @@ -247,8 +246,8 @@ def check(network): expose_value=False, prompt='Are you sure you want to reconfigure firewall rules?', ) -def configure_firewall(monitoring): - configure_firewall_rules(enable_monitoring=monitoring) +def configure_firewall(): + configure_firewall_rules() @node.command(help='Show node version information') diff --git a/node_cli/cli/sgx.py b/node_cli/cli/sgx.py new file mode 100644 index 00000000..e1635617 --- /dev/null +++ b/node_cli/cli/sgx.py @@ -0,0 +1,191 @@ +# -*- coding: utf-8 -*- +# +# This file is part of node-cli +# +# Copyright (C) 2026 SKALE Labs +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +import json + +import click +from terminaltables import SingleTable + +from node_cli.configs.sgx import SGX_SIGN_TIMEOUT +from node_cli.core.sgx import ( + SgxCertificateError, + check_certificate, + get_certificate_status, + get_server_options, + renew_certificate, +) +from node_cli.utils.decorators import check_inited, check_user +from node_cli.utils.exit_codes import CLIExitCodes +from node_cli.utils.helper import abort_if_false, error_exit +from node_cli.utils.settings import get_sgx_url +from node_cli.utils.texts import safe_load_texts + +G_TEXTS = safe_load_texts() +TEXTS = G_TEXTS['sgx'] + + +@click.group() +def sgx_cli(): + pass + + +@sgx_cli.group('sgx', help=TEXTS['help']) +def sgx(): + pass + + +@sgx.command('options', help=TEXTS['options']['help']) +@click.option('--json', 'json_format', is_flag=True, help=G_TEXTS['common']['json']) +@check_inited +@check_user +def options(json_format: bool) -> None: + _configured_sgx_url() + status, payload = get_server_options() + if status != 'ok': + error_exit(payload, exit_code=CLIExitCodes.BAD_API_RESPONSE) + if json_format: + print(json.dumps(payload)) + else: + rows = [['SGX option', 'Value']] + for group, values in payload.items(): + entries = ( + [(f'{group}.{key}', value) for key, value in values.items()] + if isinstance(values, dict) + else [(group, values)] + ) + rows.extend( + [key, value if isinstance(value, str) else json.dumps(value)] + for key, value in entries + ) + print(SingleTable(rows).table) + + +@sgx.command('cert-status', help=TEXTS['status']['help']) +@click.option('--json', 'json_format', is_flag=True, help=G_TEXTS['common']['json']) +@click.option('--check', is_flag=True, help=TEXTS['status']['check']) +def cert_status(json_format: bool, check: bool) -> None: + try: + info = get_certificate_status() + except SgxCertificateError as err: + error_exit(str(err), exit_code=CLIExitCodes.OPERATION_EXECUTION_ERROR) + check_error = None + if check: + try: + info['server_version'] = check_certificate(_configured_sgx_url()) + except SgxCertificateError as err: + check_error = str(err) + if json_format: + if check_error: + info['check_error'] = check_error + print(json.dumps(info)) + else: + print_certificate_status(info) + if check_error: + error_exit(check_error, exit_code=CLIExitCodes.OPERATION_EXECUTION_ERROR) + + +@sgx.command('renew', help=TEXTS['renew']['help']) +@click.option( + '--yes', + is_flag=True, + callback=abort_if_false, + expose_value=False, + prompt=TEXTS['renew']['prompt'], +) +@click.option( + '--timeout', + type=int, + default=SGX_SIGN_TIMEOUT, + show_default=True, + help=TEXTS['renew']['timeout'], +) +@click.option('--skip-verify', is_flag=True, help=TEXTS['renew']['skip_verify']) +@check_inited +@check_user +def renew(timeout: int, skip_verify: bool) -> None: + sgx_url = _configured_sgx_url() + try: + result = renew_certificate(sgx_url, timeout=timeout, verify=not skip_verify, log=print) + except SgxCertificateError as err: + error_exit(str(err), exit_code=CLIExitCodes.OPERATION_EXECUTION_ERROR) + print_certificate_status(result) + if result['backup']: + print(TEXTS['renew']['backup'].format(path=result['backup'])) + print(TEXTS['renew']['done']) + + +def _configured_sgx_url() -> str: + try: + sgx_url = get_sgx_url() + except Exception as err: # settings files are missing or invalid + error_exit(f'Cannot read node settings: {err}', exit_code=CLIExitCodes.NODE_STATE_ERROR) + if not sgx_url: + error_exit(TEXTS['no_sgx'], exit_code=CLIExitCodes.NODE_STATE_ERROR) + return sgx_url + + +def print_certificate_status(info: dict) -> None: + present = info['present'] + rows = [ + ['SGX client certificate', ''], + ['Directory', info['directory']], + ['Private key', _presence(present['key'])], + ['Signing request', _presence(present['csr'])], + ['Certificate', _presence(present['crt'])], + ] + if 'subject' in info: + rows.extend( + [ + ['Subject CN', info['subject']], + ['Issuer CN', info['issuer']], + ['Valid from', info['not_valid_before']], + ['Valid until', info['not_valid_after']], + ['Days left', str(info['days_left'])], + ['SHA-256', info['fingerprint_sha256']], + ['Key matches', _yes_no(info['key_matches'])], + ] + ) + if info.get('server_version'): + rows.append(['SGX server', f'accepted the certificate, version {info["server_version"]}']) + print(SingleTable(rows).table) + for notice in _notices(info): + print(notice) + + +def _notices(info: dict) -> list[str]: + notices = [] + if not info['complete']: + notices.append(TEXTS['status']['missing']) + elif info.get('expired'): + notices.append(TEXTS['status']['expired']) + elif info.get('expires_soon'): + notices.append(TEXTS['status']['expires_soon'].format(days=info['days_left'])) + if info.get('key_matches') is False: + notices.append(TEXTS['status']['key_mismatch']) + return notices + + +def _presence(present: bool) -> str: + return 'present' if present else 'missing' + + +def _yes_no(value: bool | None) -> str: + if value is None: + return 'unknown' + return 'yes' if value else 'no' diff --git a/node_cli/configs/__init__.py b/node_cli/configs/__init__.py index 7ed6afee..d65e4613 100644 --- a/node_cli/configs/__init__.py +++ b/node_cli/configs/__init__.py @@ -40,6 +40,8 @@ SKALE_DIR = os.path.join(G_CONF_HOME, '.skale') SKALE_TMP_DIR = os.path.join(SKALE_DIR, '.tmp') +AUTH_DIR = Path(SKALE_DIR) / 'auth' +ADMIN_API_TOKEN_PATH = AUTH_DIR / 'admin-api.token' NODE_DATA_PATH = os.path.join(SKALE_DIR, 'node_data') SCHAIN_NODE_DATA_PATH = os.path.join(NODE_DATA_PATH, 'schains') diff --git a/node_cli/configs/routes.py b/node_cli/configs/routes.py index d26dccb9..3cda6bee 100644 --- a/node_cli/configs/routes.py +++ b/node_cli/configs/routes.py @@ -37,7 +37,7 @@ 'update-safe', ], 'health': ['containers', 'schains'], - 'info': ['sgx'], + 'info': ['sgx', 'sgx-options'], 'schains': ['config', 'list', 'dkg-statuses', 'firewall-rules', 'repair', 'get'], 'ssl': ['status', 'upload'], 'wallet': ['info', 'send-eth'], diff --git a/node_cli/configs/sgx.py b/node_cli/configs/sgx.py new file mode 100644 index 00000000..d37c3a78 --- /dev/null +++ b/node_cli/configs/sgx.py @@ -0,0 +1,54 @@ +# -*- coding: utf-8 -*- +# +# This file is part of node-cli +# +# Copyright (C) 2026 SKALE Labs +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +import os + +from node_cli.configs import NODE_DATA_PATH, SGX_CERTS_PATH + +# File names are fixed by the sgx client library that node services use; it expects +# exactly these three entries in the certificate directory. +SGX_KEY_FILENAME = 'sgx.key' +SGX_CSR_FILENAME = 'sgx.csr' +SGX_CRT_FILENAME = 'sgx.crt' + +SGX_CERTS_BACKUP_PATH = os.path.join(NODE_DATA_PATH, 'sgx_certs_backup') + +# The SGX wallet signs certificate requests over plain HTTP on the port that follows +# its main port, which is how the sgx client library derives the address as well. +SGX_CSR_SERVER_PORT_OFFSET = 1 + +SGX_KEY_SIZE = 2048 +SGX_RPC_TIMEOUT = 60 +SGX_SIGN_POLL_INTERVAL = 10 +SGX_SIGN_TIMEOUT = 600 +SGX_CERT_EXPIRY_WARNING_DAYS = 30 + +__all__ = [ + 'SGX_CERTS_PATH', + 'SGX_CERTS_BACKUP_PATH', + 'SGX_CERT_EXPIRY_WARNING_DAYS', + 'SGX_CRT_FILENAME', + 'SGX_CSR_FILENAME', + 'SGX_CSR_SERVER_PORT_OFFSET', + 'SGX_KEY_FILENAME', + 'SGX_KEY_SIZE', + 'SGX_RPC_TIMEOUT', + 'SGX_SIGN_POLL_INTERVAL', + 'SGX_SIGN_TIMEOUT', +] diff --git a/node_cli/core/nftables.py b/node_cli/core/nftables.py index e0059fff..0d8cf0ba 100644 --- a/node_cli/core/nftables.py +++ b/node_cli/core/nftables.py @@ -27,17 +27,29 @@ from typing import Optional from node_cli.configs import ( + DEFAULT_NODE_BASE_PORT, ENV, NFTABLES_CHAIN_CONFIG_WILDCARD, NFTABLES_CHAIN_FOLDER_PATH, NFTABLES_MAIN_CONFIG_PATH, NFTABLES_SKALE_BASE_CONFIG_PATH, NFTABLES_USER_CONFIG_PATH, + NODE_CONFIG_PATH, ) -from node_cli.utils.helper import get_ssh_port, run_cmd +from node_cli.utils.helper import cleanup_dir_content, get_ssh_ports, read_json, run_cmd logger = logging.getLogger(__name__) +try: + import nftables +except (FileNotFoundError, AttributeError, ModuleNotFoundError) as err: + if 'pytest' in sys.modules or ENV == 'dev': + from collections import namedtuple # hotfix for tests + + iptc = namedtuple('nftables', ['Chain', 'Rule']) + else: + logger.error(f'Unable to import nftables due to an error {err}') + @dataclass class ServicePort: @@ -65,22 +77,89 @@ class SGXPort: LEGACY_TABLE = 'filter' CHAIN_PRIORITY = 1 HOOK = 'input' -POLICY = 'accept' +POLICY_ACCEPT = 'accept' +POLICY_DROP = 'drop' + +DYNAMIC_CHAIN_PREFIX = 'skale-' + +USER_CHAIN = 'skale_user' + +# sChain base ports are allocated as node_base_port + schain_index * 64 +# (PORTS_PER_SCHAIN in skale.py); 128 slots cover every possible allocation +PORTS_PER_SCHAIN = 64 +SCHAIN_PORTS_PER_NODE = 128 * PORTS_PER_SCHAIN +SCHAIN_BASE_PORT_ENV = 'SCHAIN_BASE_PORT' +FIREWALL_DEFAULT_DROP_ENV = 'FIREWALL_DEFAULT_DROP' +MIN_SCHAIN_BASE_PORT = 2000 +MAX_PORT = 65535 + +ICMP_ACCEPT_TYPES = ('destination-unreachable', 'time-exceeded') +ICMPV6_ACCEPT_TYPES = ( + 'destination-unreachable', + 'packet-too-big', + 'time-exceeded', + 'parameter-problem', + 'nd-router-advert', + 'nd-neighbor-solicit', + 'nd-neighbor-advert', +) -try: - import nftables -except (FileNotFoundError, AttributeError, ModuleNotFoundError) as err: - if 'pytest' in sys.modules or ENV == 'dev': - from collections import namedtuple # hotfix for tests +class NFTablesError(Exception): + pass - iptc = namedtuple('nftables', ['Chain', 'Rule']) - else: - logger.error(f'Unable to import nftables due to an error {err}') +def dport_match(protocol: str, first_port: int, last_port: int) -> dict: + right = first_port if last_port == first_port else {'range': [first_port, last_port]} + return { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': protocol, 'field': 'dport'}}, + 'right': right, + } + } -class NFTablesError(Exception): - pass + +def icmp_match(protocol: str, icmp_type: str) -> dict: + return { + 'match': { + 'left': {'payload': {'protocol': protocol, 'field': 'type'}}, + 'op': '==', + 'right': icmp_type, + } + } + + +def ip_protocol_match(protocol: str) -> dict: + return { + 'match': { + 'left': {'payload': {'protocol': 'ip', 'field': 'protocol'}}, + 'op': '==', + 'right': protocol, + } + } + + +def conntrack_accept_expr() -> list[dict]: + return [ + { + 'match': { + 'left': {'ct': {'key': 'state'}}, + 'op': 'in', + 'right': ['established', 'related'], + } + }, + {'counter': None}, + {'accept': None}, + ] + + +def loopback_accept_expr() -> list[dict]: + return [ + {'match': {'left': {'meta': {'key': 'iifname'}}, 'op': '==', 'right': 'lo'}}, + {'counter': None}, + {'accept': None}, + ] @dataclass @@ -95,10 +174,16 @@ class Rule: def __post_init__(self): if self.first_port is not None and self.last_port is None: self.last_port = self.first_port - if all( - val is None for val in (self.first_port, self.last_port, self.protocol, self.icmp_type) - ): - raise NFTablesError('Rule has no meaningful fields') + if self.protocol in ('icmp', 'icmpv6') and not self.icmp_type: + raise NFTablesError(f'{self.protocol} rule requires icmp_type') + + def to_expr(self) -> list[dict]: + matches = [] + if self.protocol in ('tcp', 'udp') and self.first_port: + matches.append(dport_match(self.protocol, self.first_port, self.last_port)) + elif self.protocol in ('icmp', 'icmpv6'): + matches.append(icmp_match(self.protocol, self.icmp_type)) + return [*matches, {'counter': None}, {self.action: None}] class NFTablesManager: @@ -144,7 +229,7 @@ def chain_exists(self, chain: str, family: Optional[str] = None) -> bool: return chain in self.get_chains(family=family) def create_chain_if_not_exists( - self, chain: str, hook: str, priority: int = CHAIN_PRIORITY, policy: str = POLICY + self, chain: str, hook: str, priority: int = CHAIN_PRIORITY, policy: str = POLICY_ACCEPT ) -> None: if not self.chain_exists(chain): cmd = { @@ -172,7 +257,7 @@ def create_chain_if_not_exists( def update_chain_policy( self, chain: str, - policy: str = POLICY, + policy: str = POLICY_ACCEPT, family: Optional[str] = None, table: Optional[str] = None, ) -> None: @@ -180,24 +265,33 @@ def update_chain_policy( family = family or self.family table = table or self.table if self.chain_exists(chain, family=family): - cmd = [ - 'nft', - 'add', - 'chain', - family, - table, - chain, - '{', - 'policy', - POLICY, - ';', - '}', - ] - run_cmd(cmd) + cmd = f'add chain {family} {table} {chain} {{ policy {policy} ; }}' + rc, output, error = self.nft.cmd(cmd) + if rc != 0: + raise NFTablesError(f'Failed to set policy {policy} on {chain}: {error}') logger.info('Updated chain policy: %s %s', chain, policy) else: logger.info('Chain %s does not exist', chain) + def get_chain_policy(self, chain: str) -> Optional[str]: + """Return the policy of a chain in the managed table or None.""" + try: + rc, output, error = self.nft.cmd(f'list chain {self.family} {self.table} {chain}') + if rc != 0: + return None + data = json.loads(output) + if not isinstance(data, dict): + return None + for item in data.get('nftables', []): + if not isinstance(item, dict): + continue + chain_data = item.get('chain') + if isinstance(chain_data, dict) and chain_data.get('name') == chain: + return chain_data.get('policy') + except (TypeError, ValueError) as e: + logger.error('Failed to get policy of chain %s: %s', chain, e) + return None + def table_exists(self) -> bool: try: rc, output, error = self.nft.cmd(f'list table {self.family} {self.table}') @@ -230,302 +324,273 @@ def get_rules(self, chain: str) -> list[dict]: logger.error('Failed to get rules: %s', e) return [] + @staticmethod + def _normalized_expr(expr: list[dict]) -> list[dict]: + return [{'counter': None} if 'counter' in statement else statement for statement in expr] + def rule_exists(self, chain: str, new_rule_expr: list[dict]) -> bool: - existing_rules = self.get_rules(chain) - - for rule in existing_rules: - expr = rule.get('expr') - for i, statement in enumerate(expr): - if 'counter' in statement: - expr[i] = {'counter': None} - rule['counter'] = None - if expr == new_rule_expr: - return True - return False + target = self._normalized_expr(new_rule_expr) + return any( + self._normalized_expr(rule.get('expr', [])) == target for rule in self.get_rules(chain) + ) - def add_drop_rule(self, rule: Rule) -> None: - expr = [] + def _execute_rule_with_op( + self, + op: str, + chain: str, + expr: Optional[list[dict]] = None, + handle: Optional[int] = None, + ) -> None: + rule: dict = {'family': self.family, 'table': self.table, 'chain': chain} + if expr is not None: + rule['expr'] = expr + if handle is not None: + rule['handle'] = handle + self.execute_cmd({'nftables': [{op: {'rule': rule}}]}) + + def _ensure_rule( + self, chain: str, expr: list[dict], op: str = 'add', label: str = 'rule' + ) -> None: + if self.rule_exists(chain, expr): + logger.info('%s already exists in chain %s', label, chain) + return + self._execute_rule_with_op(op, chain, expr=expr) + logger.info('Added %s to chain %s', label, chain) + + def _find_rule_handle(self, chain: str, expr: list[dict]) -> Optional[int]: + target = self._normalized_expr(expr) + for rule in self.get_rules(chain): + if ( + self._normalized_expr(rule.get('expr', [])) == target + and rule.get('handle') is not None + ): + return rule['handle'] + return None + + def _remove_rule_by_expr(self, chain: str, expr: list[dict]) -> bool: + handle = self._find_rule_handle(chain, expr) + if handle is None: + return False + self.delete_rule_by_handle(handle, chain=chain) + return True + @staticmethod + def _protocol_drop_expr(rule: Rule) -> list[dict]: + matches = [] if rule.first_port: - if rule.last_port == rule.first_port: - expr.append( - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, - 'right': rule.first_port, - } - } - ) - else: - expr.append( - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, - 'right': {'range': [rule.first_port, rule.last_port]}, - } - } - ) - expr.append( - { - 'match': { - 'left': {'payload': {'protocol': 'ip', 'field': 'protocol'}}, - 'op': '==', - 'right': rule.protocol, - } - }, + matches.append(dport_match('tcp', rule.first_port, rule.last_port)) + matches.append(ip_protocol_match(rule.protocol)) + return [*matches, {'counter': None}, {'drop': None}] + + def add_drop_rule(self, rule: Rule) -> None: + self._ensure_rule( + rule.chain, + self._protocol_drop_expr(rule), + label=f'{rule.protocol} drop rule', ) - expr.extend([{'counter': None}, {'drop': None}]) - if not self.rule_exists(self.chain, expr): - cmd = { - 'nftables': [ - { - 'add': { - 'rule': { - 'family': self.family, - 'table': self.table, - 'chain': rule.chain, - 'expr': expr, - } - } - } - ] - } - self.execute_cmd(cmd) - logger.info('Added drop rule %s', Rule) def remove_drop_rule(self, protocol: str) -> None: - expr = [ - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'ip', 'field': 'protocol'}}, - 'right': protocol, - } - }, - {'counter': None}, - {'drop': None}, - ] - - # Check if the drop rule exists before attempting to remove it - if self.rule_exists(self.chain, expr): - cmd = { - 'nftables': [ - { - 'delete': { - 'rule': { - 'family': self.family, - 'table': self.table, - 'chain': self.chain, - 'expr': expr, - } - } - } - ] - } - self.execute_cmd(cmd) + expr = [ip_protocol_match(protocol), {'counter': None}, {'drop': None}] + if self._remove_rule_by_expr(self.chain, expr): logger.info('Removed drop rule for %s', protocol) else: logger.info('Drop rule does not exist for %s', protocol) def add_rule(self, rule: Rule) -> None: - expr = [] - - if rule.protocol in ['tcp', 'udp']: - if rule.first_port: - if rule.last_port == rule.first_port: - expr.append( - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, - 'right': rule.first_port, - } - } - ) - else: - expr.append( - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, - 'right': {'range': [rule.first_port, rule.last_port]}, - } - } - ) - elif rule.protocol == 'icmp' and rule.icmp_type: - expr.append( - { - 'match': { - 'left': {'payload': {'protocol': 'icmp', 'field': 'type'}}, - 'op': '==', - 'right': rule.icmp_type, - } - } - ) - - expr.append({'counter': None}) - expr.append({rule.action: None}) + self._ensure_rule( + rule.chain, + rule.to_expr(), + label=f'{rule.protocol} {rule.icmp_type or rule.first_port} {rule.action} rule', + ) - if not self.rule_exists(rule.chain, expr): - cmd = { - 'nftables': [ - { - 'add': { - 'rule': { - 'family': self.family, - 'table': self.table, - 'chain': rule.chain, - 'expr': expr, - } - } - } - ] - } - self.execute_cmd(cmd) - logger.info( - 'Added new rule to chain %s: %s ports [%s, %s]', - rule.chain, - rule.protocol, - rule.first_port, - rule.last_port, - ) + def remove_rule(self, rule: Rule) -> None: + if self._remove_rule_by_expr(rule.chain, rule.to_expr()): + logger.info('Removed %s rule for %s', rule.protocol, rule.first_port) else: - logger.info( - 'Rule already exists in chain %s: %s ports [%s, %s]', - rule.chain, - rule.protocol, - rule.first_port, - rule.last_port, - ) + logger.info('No %s rule for %s to remove', rule.protocol, rule.first_port) + + def _table_listing(self) -> dict: + """Parsed json listing of the managed table; empty when absent.""" + rc, output, error = self.nft.cmd(f'list table {self.family} {self.table}') + if rc != 0: + if error and 'No such file or directory' in error: + return {} + raise NFTablesError(f'Failed to list table {self.table}: {error}') + try: + data = json.loads(output) + except (TypeError, ValueError) as err: + raise NFTablesError(f'Failed to parse table {self.table} listing: {err}') from err + if not isinstance(data, dict): + raise NFTablesError(f'Malformed table {self.table} listing') + return data + + def _table_chain_names(self) -> list[str]: + return [ + item['chain']['name'] + for item in self._table_listing().get('nftables', []) + if isinstance(item, dict) and isinstance(item.get('chain'), dict) + ] - def remove_rule(self, rule: Rule) -> None: - expr = [] - - if rule.protocol in ['tcp', 'udp']: - if rule.first_port: - if rule.last_port == rule.first_port: - expr.append( - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, - 'right': rule.first_port, - } - } - ) - else: - expr.append( - { - 'match': { - 'op': '==', - 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, - 'right': {'range': [rule.first_port, rule.last_port]}, - } - } - ) - elif rule.protocol == 'icmp' and rule.icmp_type: - expr.append( - { - 'match': { - 'left': {'payload': {'protocol': 'icmp', 'field': 'type'}}, - 'op': '==', - 'right': rule.icmp_type, - } - } - ) + def get_dynamic_chain_port_ranges(self) -> list[tuple[str, int, int]]: + """Min/max tcp dport covered by each dynamic skale-admin chain.""" + data = self._table_listing() + + ports: dict[str, list[int]] = {} + for item in data.get('nftables', []): + rule = item.get('rule') if isinstance(item, dict) else None + if not rule or not rule.get('chain', '').startswith(DYNAMIC_CHAIN_PREFIX): + continue + for statement in rule.get('expr', []): + match = statement.get('match', {}) + if match.get('left', {}).get('payload', {}).get('field') != 'dport': + continue + right = match.get('right') + chain_ports = ports.setdefault(rule['chain'], []) + if isinstance(right, dict) and 'range' in right: + chain_ports.extend(right['range']) + elif isinstance(right, int): + chain_ports.append(right) + return [(chain, min(values), max(values)) for chain, values in ports.items() if values] + + def validate_dynamic_ranges(self, envelope: tuple[int, int]) -> None: + """Ensure ports of every dynamic skale-admin chain fit into the envelope.""" + for chain, first_port, last_port in self.get_dynamic_chain_port_ranges(): + if first_port < envelope[0] or last_port > envelope[1]: + raise NFTablesError( + f'Ports {first_port}-{last_port} of dynamic chain {chain} are outside ' + f'of the allowed sChain ports range {envelope[0]}-{envelope[1]}. ' + f'Set {SCHAIN_BASE_PORT_ENV} env variable to the base port the node ' + 'was registered with and rerun the command' + ) - # Check if the rule exists before attempting to remove it - if self.rule_exists(rule.chain, expr): - cmd = { - 'nftables': [ - { - 'delete': { - 'rule': { - 'family': self.family, - 'table': self.table, - 'chain': rule.chain, - 'expr': expr, - } - } - } - ] - } - self.execute_cmd(cmd) - logger.info( - 'Removed rule from chain %s: %s ports [%s, %s]', - rule.chain, - rule.protocol, - rule.first_port, - rule.last_port, - ) - else: - logger.info( - 'Rule does not exist in chain %s: %s ports [%s, %s]', - rule.chain, - rule.protocol, - rule.first_port, - rule.last_port, - ) + def verify_critical_accepts(self) -> None: + """Ensure lockout-critical accept rules are in place before setting drop policy.""" + accepts = [('conntrack', conntrack_accept_expr())] + accepts.extend( + (f'ssh port {port}', Rule(chain=self.chain, protocol='tcp', first_port=port).to_expr()) + for port in get_ssh_ports() + ) + for name, expr in accepts: + if not self.rule_exists(self.chain, expr): + raise NFTablesError( + f'Refusing to set drop policy: {name} accept rule is missing ' + f'in chain {self.chain}' + ) - def add_connection_tracking_rule(self, chain: str) -> None: - expr = [ - { - 'match': { - 'left': {'ct': {'key': 'state'}}, - 'op': 'in', - 'right': ['established', 'related'], - } - }, + def ensure_default_drop(self) -> None: + """Switch the skale chain policy to drop after verifying the accepts.""" + self.verify_critical_accepts() + if self.get_chain_policy(self.chain) != POLICY_DROP: + self.update_chain_policy(chain=self.chain, policy=POLICY_DROP) + + def ensure_default_accept(self) -> None: + if self.get_chain_policy(self.chain) != POLICY_ACCEPT: + self.update_chain_policy(chain=self.chain, policy=POLICY_ACCEPT) + + def delete_rule_by_handle(self, handle: int, chain: Optional[str] = None) -> None: + self._execute_rule_with_op('delete', chain or self.chain, handle=handle) + + def remove_stale_envelope_rules(self, envelope: tuple[int, int]) -> None: + """Remove sChain envelope accepts anchored at a different base port.""" + envelope_spans = (SCHAIN_PORTS_PER_NODE - 1, PORTS_PER_SCHAIN - 1) + for rule in self.get_rules(self.chain): + expr = rule.get('expr', []) + if {'accept': None} not in expr: + continue + for statement in expr: + match = statement.get('match', {}) + right = match.get('right') + if ( + match.get('left', {}).get('payload', {}).get('field') == 'dport' + and isinstance(right, dict) + and 'range' in right + and right['range'][1] - right['range'][0] in envelope_spans + and tuple(right['range']) != envelope + and rule.get('handle') is not None + ): + logger.info('Removing stale envelope rule %s', right['range']) + self.delete_rule_by_handle(rule['handle']) + + def remove_misordered_udp_drop(self) -> None: + """Delete the blanket udp drop when it shadows the udp DNS accept.""" + udp_drop = [ip_protocol_match('udp'), {'counter': None}, {'drop': None}] + udp_dns_accept = [ + dport_match('udp', ServicePort.DNS, ServicePort.DNS), {'counter': None}, {'accept': None}, ] - - if not self.rule_exists(chain, expr): + drop_handle, drop_index, accept_index = None, None, None + for index, rule in enumerate(self.get_rules(self.chain)): + expr = self._normalized_expr(rule.get('expr', [])) + if expr == udp_drop: + drop_handle, drop_index = rule.get('handle'), index + elif expr == udp_dns_accept: + accept_index = index + if drop_handle is not None and (accept_index is None or drop_index < accept_index): + logger.info('Removing misordered udp drop rule') + self.delete_rule_by_handle(drop_handle) + + def remove_source_quench_rule(self) -> None: + """Remove the legacy icmp source-quench accept (deprecated by RFC 6633).""" + expr = [icmp_match('icmp', 'source-quench'), {'counter': None}, {'accept': None}] + if self._remove_rule_by_expr(self.chain, expr): + logger.info('Removed legacy source-quench rule') + + def create_user_chain_if_not_exists(self) -> None: + """Create the regular chain holding user.conf rules.""" + if not self.chain_exists(USER_CHAIN): cmd = { 'nftables': [ { 'add': { - 'rule': { + 'chain': { 'family': self.family, 'table': self.table, - 'chain': chain, - 'expr': expr, + 'name': USER_CHAIN, } } } ] } self.execute_cmd(cmd) - logger.info('Added connection tracking rule to chain %s', chain) - else: - logger.info('Connection tracking rule already exists in chain %s', chain) - - def add_loopback_rule(self, chain) -> None: - expr = [ - {'match': {'left': {'meta': {'key': 'iifname'}}, 'op': '==', 'right': 'lo'}}, - {'counter': None}, - {'accept': None}, + logger.info('Created user rules chain %s', USER_CHAIN) + + def ensure_user_chain_jump(self) -> None: + expr = [{'jump': {'target': USER_CHAIN}}] + self._ensure_rule(self.chain, expr, op='insert', label='user chain jump') + + def apply_user_rules(self) -> None: + """Reload user.conf into the live user rules chain.""" + content = '' + if os.path.isfile(NFTABLES_USER_CONFIG_PATH): + with open(NFTABLES_USER_CONFIG_PATH) as user_config: + content = user_config.read() + commands = ( + f'flush chain {self.family} {self.table} {USER_CHAIN}\n' + f'table {self.family} {self.table} {{\n' + f'chain {USER_CHAIN} {{\n' + f'{content}\n' + f'}}\n' + f'}}' + ) + rc, output, error = self.nft.cmd(commands) + if rc != 0: + raise NFTablesError(f'Failed to apply user.conf rules: {error}') + + def remove_user_rules_from_main_chain(self) -> None: + """Remove user.conf rules that older saved configs loaded into the + skale chain directly. + """ + user_exprs = [ + self._normalized_expr(rule.get('expr', [])) for rule in self.get_rules(USER_CHAIN) ] - if not self.rule_exists(chain, expr): - json_cmd = { - 'nftables': [ - { - 'add': { - 'rule': { - 'family': self.family, - 'table': self.table, - 'chain': self.chain, - 'expr': expr, - } - } - } - ] - } - self.execute_cmd(json_cmd) - else: - logger.info('Loopback rule already exists in chain %s', chain) + if not user_exprs: + return + for rule in self.get_rules(self.chain): + expr = self._normalized_expr(rule.get('expr', [])) + if expr in user_exprs and rule.get('handle') is not None: + logger.info('Moving user rule out of the main chain') + self.delete_rule_by_handle(rule['handle']) def get_base_ruleset(self) -> str: self.nft.set_json_output(False) @@ -538,62 +603,104 @@ def get_base_ruleset(self) -> str: finally: self.nft.set_json_output(True) + def _setup_user_chain(self) -> None: + self.create_user_chain_if_not_exists() + self.ensure_user_chain_jump() + self.apply_user_rules() + self.remove_user_rules_from_main_chain() + + def remove_monitoring_accepts(self) -> None: + """Remove every legacy monitoring accept from the managed base chain.""" + ssh_ports = get_ssh_ports() + monitoring_exprs = [ + Rule(chain=self.chain, protocol='tcp', first_port=port).to_expr() + for port in (ServicePort.EXPORTER, ServicePort.CADVISOR) + if port not in ssh_ports + ] + for rule in self.get_rules(self.chain): + if ( + self._normalized_expr(rule.get('expr', [])) in monitoring_exprs + and rule.get('handle') is not None + ): + self.delete_rule_by_handle(rule['handle']) + + def _add_service_accepts(self) -> None: + self._ensure_rule(self.chain, conntrack_accept_expr(), label='connection tracking rule') + tcp_ports = [ + *get_ssh_ports(), + ServicePort.DNS, + ServicePort.HTTPS, + ServicePort.HTTP, + ServicePort.WATCHDOG_HTTP, + ServicePort.WATCHDOG_HTTPS, + ] + for port in tcp_ports: + self.add_rule(Rule(chain=self.chain, protocol='tcp', first_port=port)) + self.remove_misordered_udp_drop() + self.add_rule(Rule(chain=self.chain, protocol='udp', first_port=ServicePort.DNS)) + self._ensure_rule(self.chain, loopback_accept_expr(), label='loopback rule') + + def _add_icmp_accepts(self) -> None: + self.remove_source_quench_rule() + for icmp_type in ICMP_ACCEPT_TYPES: + self.add_rule(Rule(chain=self.chain, protocol='icmp', icmp_type=icmp_type)) + for icmpv6_type in ICMPV6_ACCEPT_TYPES: + self.add_rule(Rule(chain=self.chain, protocol='icmpv6', icmp_type=icmpv6_type)) + + def _ensure_envelope(self, envelope: tuple[int, int]) -> None: + # Fine-grained filtering inside the envelope is enforced by the + # dynamic skale-admin chains that run earlier (priority 0) + self.remove_stale_envelope_rules(envelope) + self.add_rule( + Rule(chain=self.chain, protocol='tcp', first_port=envelope[0], last_port=envelope[1]) + ) - def setup_firewall(self, enable_monitoring: bool = False) -> None: + def _add_drop_rules(self) -> None: + self.add_drop_rule( + Rule(chain=self.chain, protocol='tcp', first_port=SGXPort.HTTPS, last_port=SGXPort.ZMQ) + ) + self.add_drop_rule(Rule(chain=self.chain, protocol='udp')) + + def setup_firewall(self, keep_accept_policy: bool = False) -> None: """Setup firewall rules.""" logger.info('Configuring firewall rules') + default_drop = firewall_default_drop_enabled() and not keep_accept_policy try: self.create_table_if_not_exists() - - base_chains_config = {'skale': {'hook': 'input', 'policy': 'accept'}} - - for chain, config in base_chains_config.items(): - self.create_chain_if_not_exists( - chain=chain, hook=config['hook'], policy=config['policy'] - ) - - self.add_connection_tracking_rule(self.chain) - - tcp_ports = [ - get_ssh_port(), - ServicePort.DNS, - ServicePort.HTTPS, - ServicePort.HTTP, - ServicePort.WATCHDOG_HTTP, - ServicePort.WATCHDOG_HTTPS, - ] - if enable_monitoring: - tcp_ports.extend([ServicePort.EXPORTER, ServicePort.CADVISOR]) - for port in tcp_ports: - self.add_rule(Rule(chain=self.chain, protocol='tcp', first_port=port)) - - self.add_rule(Rule(chain=self.chain, protocol='udp', first_port=ServicePort.DNS)) - self.add_loopback_rule(chain=self.chain) - - icmp_types = ['destination-unreachable', 'source-quench', 'time-exceeded'] - for icmp_type in icmp_types: - self.add_rule(Rule(chain=self.chain, protocol='icmp', icmp_type=icmp_type)) - - self.add_drop_rule( - Rule( - chain=self.chain, - first_port=SGXPort.HTTPS, - last_port=SGXPort.ZMQ, - protocol='tcp', - ) - ) - - self.add_drop_rule(Rule(chain=self.chain, protocol='udp')) - logger.info('Making sure legacy chain has default policy %s', POLICY) + self.create_chain_if_not_exists(chain=self.chain, hook=HOOK, policy=POLICY_ACCEPT) + if not default_drop: + # rollback must not be blocked by any later failing step, + # including an invalid envelope configuration + self.ensure_default_accept() + + envelope = get_schain_ports_envelope() + if default_drop: + # fail fast, before any rule is touched + self.validate_dynamic_ranges(envelope) + + self._setup_user_chain() + self.remove_monitoring_accepts() + self._add_service_accepts() + self._add_icmp_accepts() + self._ensure_envelope(envelope) + self._add_drop_rules() + + logger.info('Making sure legacy chain has default policy %s', POLICY_ACCEPT) self.update_chain_policy( - chain=LEGACY_CHAIN, policy=POLICY, family=LEGACY_FAMILY, table=LEGACY_TABLE + chain=LEGACY_CHAIN, policy=POLICY_ACCEPT, family=LEGACY_FAMILY, table=LEGACY_TABLE ) + if default_drop: + self.ensure_default_drop() + except Exception as e: logger.error('Failed to setup firewall: %s', e) raise NFTablesError(e) - logger.info('Firewall rules are configured') + logger.info( + 'Firewall rules are configured, default policy: %s', + POLICY_DROP if default_drop else POLICY_ACCEPT, + ) def cleanup_legacy_rules(self, ssh: bool = False, dns: bool = False) -> None: """Cleans up all node-cli generated rules.""" @@ -608,7 +715,7 @@ def cleanup_legacy_rules(self, ssh: bool = False, dns: bool = False) -> None: ServicePort.DNS, # tcp is redundant, making sure it's removed ] if ssh: - tcp_ports.append(get_ssh_port()) + tcp_ports.extend(get_ssh_ports()) for port in tcp_ports: self.remove_rule(Rule(chain=self.chain, protocol='tcp', first_port=port)) if dns: @@ -630,6 +737,92 @@ def flush_chain(self, chain: str) -> None: logger.error(f'Failed to flush chain: {str(e)}') raise NFTablesError('Flushing chain errored') + def delete_chain(self, chain: str) -> None: + chain_spec = {'family': self.family, 'table': self.table, 'name': chain} + self.execute_cmd( + {'nftables': [{'flush': {'chain': chain_spec}}, {'delete': {'chain': chain_spec}}]} + ) + logger.info('Deleted chain %s', chain) + + def _critical_accept_exprs(self) -> list[list[dict]]: + """Rules that keep the node reachable: conntrack, loopback, ssh, DNS.""" + exprs = [conntrack_accept_expr(), loopback_accept_expr()] + try: + ssh_ports = get_ssh_ports() + except (RuntimeError, ValueError): + # the policy is accept during cleanup, so reachability is safe + # even when ssh detection is impossible + ssh_ports = [] + for port in (*ssh_ports, ServicePort.DNS): + exprs.append(Rule(chain=self.chain, protocol='tcp', first_port=port).to_expr()) + exprs.append(Rule(chain=self.chain, protocol='udp', first_port=ServicePort.DNS).to_expr()) + return exprs + + def cleanup_firewall(self) -> None: + """Reset the firewall to a minimal state that keeps the node reachable. + + Restores the accept policy, removes the user and dynamic chains and + every rule except the critical accepts: conntrack, loopback, ssh + and DNS. + """ + self.ensure_default_accept() + self._remove_rule_by_expr(self.chain, [{'jump': {'target': USER_CHAIN}}]) + for chain in self._table_chain_names(): + if chain == USER_CHAIN or chain.startswith(DYNAMIC_CHAIN_PREFIX): + self.delete_chain(chain) + keep = [self._normalized_expr(expr) for expr in self._critical_accept_exprs()] + for rule in self.get_rules(self.chain): + if ( + self._normalized_expr(rule.get('expr', [])) not in keep + and rule.get('handle') is not None + ): + self.delete_rule_by_handle(rule['handle']) + + +def firewall_default_drop_enabled() -> bool: + value = os.getenv(FIREWALL_DEFAULT_DROP_ENV, 'True') + return value.lower() not in ('false', '0', 'no', 'off') + + +def get_registered_base_port() -> Optional[tuple[int, int]]: + """Base port and envelope size from the node config.""" + if not os.path.isfile(NODE_CONFIG_PATH): + return None + try: + node_config = read_json(NODE_CONFIG_PATH) + except (OSError, ValueError) as e: + logger.warning('Failed to read node config: %s', e) + return None + if not isinstance(node_config, dict): + logger.warning('Node config is malformed') + return None + for key, size in ( + ('node_base_port', SCHAIN_PORTS_PER_NODE), + ('schain_base_port', PORTS_PER_SCHAIN), + ): + value = node_config.get(key) + if isinstance(value, int) and not isinstance(value, bool) and value > 0: + return value, size + return None + + +def get_schain_ports_envelope() -> tuple[int, int]: + """Range of ports that can be allocated to sChains on this node.""" + env_value = os.getenv(SCHAIN_BASE_PORT_ENV) + if env_value: + try: + base_port, size = int(env_value), SCHAIN_PORTS_PER_NODE + except ValueError: + raise NFTablesError(f'{SCHAIN_BASE_PORT_ENV} must be an integer, got {env_value}') + else: + base_port, size = get_registered_base_port() or ( + DEFAULT_NODE_BASE_PORT, + SCHAIN_PORTS_PER_NODE, + ) + if base_port < MIN_SCHAIN_BASE_PORT or base_port + size - 1 > MAX_PORT: + raise NFTablesError(f'Invalid sChain base port {base_port}') + return base_port, base_port + size - 1 + def prepare_directories() -> None: logger.info('Prepare directories for nftables') @@ -637,16 +830,33 @@ def prepare_directories() -> None: create_user_config_path() -def configure_nftables(enable_monitoring: bool = False) -> None: +def configure_nftables(keep_accept_policy: bool = False) -> None: prepare_directories() enable_nftables_service() nft_mgr = NFTablesManager() - nft_mgr.setup_firewall(enable_monitoring=enable_monitoring) + nft_mgr.setup_firewall(keep_accept_policy=keep_accept_policy) ruleset = nft_mgr.get_base_ruleset() save_nftables_rules(ruleset) remove_legacy_saved_rules() +def cleanup_nftables() -> None: + """Reset the firewall after node cleanup and persist the minimal state.""" + logger.info('Cleaning up firewall rules') + nft_mgr = NFTablesManager() + ruleset = '' + if nft_mgr.table_exists() and nft_mgr.chain_exists(nft_mgr.chain): + nft_mgr.cleanup_firewall() + ruleset = nft_mgr.get_base_ruleset() + if os.path.isdir(NFTABLES_CHAIN_FOLDER_PATH): + cleanup_dir_content(NFTABLES_CHAIN_FOLDER_PATH) + if os.path.isdir(os.path.dirname(NFTABLES_SKALE_BASE_CONFIG_PATH)): + # a plain snapshot with no includes - reboot restores the same + # minimal ruleset + with open(NFTABLES_SKALE_BASE_CONFIG_PATH, 'w') as base_config: + base_config.write(ruleset) + + def enable_nftables_service() -> None: logger.info('Enabling nftables services') run_cmd(['systemctl', 'enable', 'nftables']) @@ -655,8 +865,12 @@ def enable_nftables_service() -> None: def save_nftables_base_rules(ruleset: str) -> None: ruleset_lines = ruleset.split('\n') chain_include_line = f'\tinclude "{NFTABLES_CHAIN_CONFIG_WILDCARD}"' - user_include_line = f'\t\tinclude "{NFTABLES_USER_CONFIG_PATH}"' - ruleset_lines.insert(3, user_include_line) + user_chain_lines = [ + f'\tchain {USER_CHAIN} {{', + f'\t\tinclude "{NFTABLES_USER_CONFIG_PATH}"', + '\t}', + ] + ruleset_lines[1:1] = user_chain_lines ruleset_lines.insert(-2, chain_include_line) with open(NFTABLES_SKALE_BASE_CONFIG_PATH, 'w') as f: f.write('\n'.join(ruleset_lines)) diff --git a/node_cli/core/node.py b/node_cli/core/node.py index b1f22d86..ea37e640 100644 --- a/node_cli/core/node.py +++ b/node_cli/core/node.py @@ -27,6 +27,7 @@ from typing import Optional, Tuple import docker +from filelock import FileLock from node_cli.cli import __version__ from node_cli.configs import ( @@ -34,6 +35,7 @@ CONTAINER_CONFIG_PATH, FILESTORAGE_MAPPING, LOG_PATH, + NODE_CONFIG_PATH, RESTORE_SLEEP_TIMEOUT, SCHAINS_MNT_DIR_REGULAR, SCHAINS_MNT_DIR_SINGLE_CHAIN, @@ -52,6 +54,7 @@ passive_skale, passive_fair, ) +from node_cli.core.nftables import get_registered_base_port from node_cli.migrations.focal_to_jammy import migrate as migrate_2_6 from node_cli.operations import ( cleanup_skale_op, @@ -79,6 +82,8 @@ error_exit, get_request, post_request, + read_json, + save_json, ) from node_cli.utils.meta import CliMetaManager from node_cli.utils.node_type import NodeType, NodeMode @@ -145,12 +150,39 @@ def register_node(name, p2p_ip, public_ip, port, domain_name): msg = TEXTS['node']['registered'] logger.info(msg) print(msg) + try: + save_registered_base_port(port) + logger.info('Reconfiguring firewall for the registered base port %d', port) + configure_nftables() + except Exception: + logger.exception('Post-registration firewall reconfiguration failed') + error_exit( + 'Node is successfully registered in SKALE manager, but firewall ' + 'reconfiguration failed. Run < skale node configure-firewall > ' + 'to complete the setup', + exit_code=CLIExitCodes.OPERATION_EXECUTION_ERROR, + ) else: error_msg = payload logger.error(f'Registration error {error_msg}') error_exit(error_msg, exit_code=CLIExitCodes.BAD_API_RESPONSE) +def save_registered_base_port(port: int) -> None: + """Persist the node base port to the node config. + + Kept separate from schain_base_port, which holds an already-allocated + sChain port in passive mode. skale-admin saves node_base_port during + registration as well + """ + lock = FileLock(f'{NODE_CONFIG_PATH}.lock') + with lock: + node_config = read_json(NODE_CONFIG_PATH) if os.path.isfile(NODE_CONFIG_PATH) else {} + if node_config.get('node_base_port') != port: + node_config['node_base_port'] = port + save_json(NODE_CONFIG_PATH, node_config) + + @check_not_inited def init(config_file: str, node_type: NodeType) -> None: node_mode = NodeMode.ACTIVE @@ -217,9 +249,33 @@ def init_passive( time.sleep(TM_INIT_TIMEOUT) if not is_base_containers_alive(node_type=NodeType.SKALE, node_mode=node_mode): error_exit('Containers are not running', exit_code=CLIExitCodes.OPERATION_EXECUTION_ERROR) + enable_firewall_default_drop_when_port_available() logger.info('Passive node initialized successfully') +def enable_firewall_default_drop_when_port_available(timeout: int = 300, interval: int = 5) -> None: + """Flip the firewall to default drop once admin saves the base port. + + Passive init configures nftables before skale-admin computes the mirrored + chain's base port, so the drop policy is deferred until the port is known. + """ + start = time.monotonic() + while time.monotonic() - start < timeout: + if get_registered_base_port() is not None: + configure_nftables() + return + time.sleep(interval) + logger.warning( + 'Node base port is not available after %d seconds - firewall default ' + 'drop is postponed until the next node update', + timeout, + ) + print( + 'Firewall default drop policy is postponed: the chain base port is not ' + 'known yet. It will be applied on the next < skale node update-passive >' + ) + + @check_inited @check_user def update_passive(config_file: str) -> None: @@ -523,5 +579,5 @@ def run_checks( print_failed_requirements_checks(failed_checks) -def configure_firewall_rules(enable_monitoring: bool = False) -> None: - configure_nftables(enable_monitoring=enable_monitoring) +def configure_firewall_rules() -> None: + configure_nftables() diff --git a/node_cli/core/sgx.py b/node_cli/core/sgx.py new file mode 100644 index 00000000..31cfdf99 --- /dev/null +++ b/node_cli/core/sgx.py @@ -0,0 +1,365 @@ +# -*- coding: utf-8 -*- +# +# This file is part of node-cli +# +# Copyright (C) 2026 SKALE Labs +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU Affero General Public License as published by +# the Free Software Foundation, either version 3 of the License, or +# (at your option) any later version. +# +# This program is distributed in the hope that it will be useful, +# but WITHOUT ANY WARRANTY; without even the implied warranty of +# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +# GNU Affero General Public License for more details. +# +# You should have received a copy of the GNU Affero General Public License +# along with this program. If not, see . + +"""SGX wallet options and client certificate management. + +Node services authenticate to the SGX wallet with a client certificate kept in +``node_data/sgx_certs``. The sgx client library inside the containers expects exactly +three files there (``sgx.key``, ``sgx.csr`` and ``sgx.crt``) and issues a certificate on +its own only when one of them is missing, so this module never leaves that directory +incomplete: new material is prepared next to it and moved into place with per-file +atomic renames, and the current certificate is kept until the new one is signed. +Certificate operations talk to the SGX server directly. Server options are queried +through the authenticated node API. +""" + +import datetime +import logging +import os +import secrets +import shutil +import tempfile +import time +import warnings +from collections.abc import Callable +from pathlib import Path +from urllib.parse import urlparse + +import requests +from cryptography import x509 +from cryptography.exceptions import UnsupportedAlgorithm +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.x509.oid import NameOID +from urllib3.exceptions import InsecureRequestWarning + +from node_cli.configs.sgx import ( + SGX_CERT_EXPIRY_WARNING_DAYS, + SGX_CERTS_BACKUP_PATH, + SGX_CERTS_PATH, + SGX_CRT_FILENAME, + SGX_CSR_FILENAME, + SGX_CSR_SERVER_PORT_OFFSET, + SGX_KEY_FILENAME, + SGX_KEY_SIZE, + SGX_RPC_TIMEOUT, + SGX_SIGN_POLL_INTERVAL, + SGX_SIGN_TIMEOUT, +) +from node_cli.utils.helper import get_request + +logger = logging.getLogger(__name__) + +FILE_MODES = {'key': 0o600, 'csr': 0o644, 'crt': 0o644} +INSTALL_ORDER = ('key', 'csr', 'crt') +SIGNING_PENDING = 1 + +Logger = Callable[[str], None] + + +class SgxCertificateError(Exception): + """Raised when the SGX client certificate cannot be read, issued or installed.""" + + +def get_server_options() -> tuple[str, str | dict]: + """Query SGX options through the admin API using the node's CLI credential.""" + return get_request(blueprint='info', method='sgx-options') + + +def certificate_paths(directory: str | Path | None = None) -> dict[str, Path]: + directory = Path(directory or SGX_CERTS_PATH) + return { + 'key': directory / SGX_KEY_FILENAME, + 'csr': directory / SGX_CSR_FILENAME, + 'crt': directory / SGX_CRT_FILENAME, + } + + +def csr_server_url(sgx_url: str) -> str: + """Address of the wallet's signing service, derived the way the sgx client does.""" + parsed = urlparse(sgx_url) + if parsed.scheme not in ('http', 'https') or not parsed.hostname or parsed.port is None: + raise SgxCertificateError(f'SGX URL must look like https://host:port, got {sgx_url!r}') + host = f'[{parsed.hostname}]' if ':' in parsed.hostname else parsed.hostname + return f'http://{host}:{parsed.port + SGX_CSR_SERVER_PORT_OFFSET}' + + +def get_certificate_status(directory: str | Path | None = None) -> dict: + """Describe the certificate files without contacting the SGX server.""" + paths = certificate_paths(directory) + present = {name: path.is_file() for name, path in paths.items()} + status: dict = { + 'directory': str(paths['crt'].parent), + 'present': present, + 'complete': all(present.values()), + } + if not present['crt']: + return status + cert = _load_certificate(paths['crt']) + now = datetime.datetime.now(datetime.timezone.utc) + try: + not_before = cert.not_valid_before_utc + not_after = cert.not_valid_after_utc + except AttributeError: + not_before = cert.not_valid_before.replace(tzinfo=datetime.timezone.utc) + not_after = cert.not_valid_after.replace(tzinfo=datetime.timezone.utc) + days_left = (not_after - now).days + status.update( + { + 'subject': _common_name(cert.subject), + 'issuer': _common_name(cert.issuer), + 'serial_number': format(cert.serial_number, 'x'), + 'not_valid_before': not_before.isoformat(timespec='seconds'), + 'not_valid_after': not_after.isoformat(timespec='seconds'), + 'days_left': days_left, + 'expired': now >= not_after, + 'not_yet_valid': now < not_before, + 'expires_soon': now < not_after and days_left < SGX_CERT_EXPIRY_WARNING_DAYS, + 'fingerprint_sha256': cert.fingerprint(hashes.SHA256()).hex(':'), + 'key_matches': _key_matches(paths['key'], cert) if present['key'] else None, + } + ) + return status + + +def check_certificate(sgx_url: str, directory: str | Path | None = None) -> str: + """Return the SGX server version obtained while authenticating with the local files.""" + paths = certificate_paths(directory) + missing = [name for name in ('key', 'crt') if not paths[name].is_file()] + if missing: + raise SgxCertificateError( + f'Cannot check the certificate, missing files: {", ".join(missing)}' + ) + return _server_version(sgx_url, (str(paths['crt']), str(paths['key']))) + + +def renew_certificate( + sgx_url: str, + *, + directory: str | Path | None = None, + backup_root: str | Path | None = None, + timeout: int = SGX_SIGN_TIMEOUT, + verify: bool = True, + log: Logger | None = None, +) -> dict: + """Issue a new client certificate and install it, keeping the current one until then. + + The new certificate is tested against the SGX server before it replaces the current + files unless ``verify`` is off. Previous files are copied under ``backup_root``. + """ + paths = certificate_paths(directory) + certs_dir = paths['crt'].parent + csr_url = csr_server_url(sgx_url) + try: + os.makedirs(certs_dir, exist_ok=True) + staging = Path(tempfile.mkdtemp(prefix='.sgx_certs.', dir=_staging_parent(certs_dir))) + except OSError as err: + raise SgxCertificateError( + f'Cannot prepare certificate files in {certs_dir}: {err}' + ) from err + + _say(log, 'Generating a new RSA key and certificate signing request ...') + key_pem, csr_pem = _generate_key_and_csr() + try: + staged = {name: staging / path.name for name, path in paths.items()} + _write_file(staged['key'], key_pem, FILE_MODES['key']) + _write_file(staged['csr'], csr_pem, FILE_MODES['csr']) + crt_pem = _request_signed_certificate(csr_url, csr_pem.decode('ascii'), timeout, log) + _write_file(staged['crt'], crt_pem.encode('utf-8'), FILE_MODES['crt']) + if not _key_matches(staged['key'], _load_certificate(staged['crt'])): + raise SgxCertificateError( + 'The certificate returned by the SGX server does not match the generated key' + ) + version = None + if verify: + _say(log, f'Verifying the new certificate against {sgx_url} ...') + version = _server_version(sgx_url, (str(staged['crt']), str(staged['key']))) + _say(log, f'SGX server (version {version}) accepted the new certificate') + try: + backup = _backup_existing(paths, Path(backup_root or SGX_CERTS_BACKUP_PATH)) + _install(staged, paths) + except OSError as err: + raise SgxCertificateError(f'Cannot install the new certificate files: {err}') from err + finally: + shutil.rmtree(staging, ignore_errors=True) + _say(log, f'Installed the new certificate in {certs_dir}') + return { + 'backup': str(backup) if backup else None, + 'server_version': version, + **get_certificate_status(certs_dir), + } + + +def _say(log: Logger | None, message: str) -> None: + logger.info(message) + if log is not None: + log(message) + + +def _rpc( + url: str, method: str, params: dict | None = None, cert: tuple[str, str] | None = None +) -> dict: + payload = {'id': 0, 'jsonrpc': '2.0', 'method': method, 'params': params or {}} + try: + with warnings.catch_warnings(): + # The SGX wallet presents a self-signed server certificate; node services + # skip its verification the same way and rely on client authentication. + warnings.simplefilter('ignore', InsecureRequestWarning) + response = requests.post( + url, json=payload, cert=cert, verify=False, timeout=SGX_RPC_TIMEOUT + ) + response.raise_for_status() + data = response.json() + except requests.exceptions.SSLError as err: + raise SgxCertificateError(f'SGX server {url} rejected the TLS connection: {err}') from err + except (requests.exceptions.RequestException, ValueError) as err: + raise SgxCertificateError(f'Cannot call {method} on {url}: {err}') from err + if not isinstance(data, dict): + raise SgxCertificateError(f'{method} on {url} returned an unexpected response') + if data.get('error'): + error = data['error'] + message = error.get('message', error) if isinstance(error, dict) else error + raise SgxCertificateError(f'{method} on {url} failed: {message}') + result = data.get('result') + if not isinstance(result, dict): + raise SgxCertificateError(f'{method} on {url} returned an unexpected response') + return result + + +def _raise_on_status(result: dict, method: str) -> None: + if result.get('status', 0) != 0: + message = result.get('errorMessage') or f'status {result.get("status")}' + raise SgxCertificateError(f'SGX server refused {method}: {message}') + + +def _server_version(sgx_url: str, cert: tuple[str, str]) -> str: + result = _rpc(sgx_url, 'getServerVersion', cert=cert) + _raise_on_status(result, 'getServerVersion') + return str(result.get('version') or 'unknown') + + +def _request_signed_certificate( + csr_url: str, csr_pem: str, timeout: int, log: Logger | None +) -> str: + result = _rpc(csr_url, 'signCertificate', {'certificate': csr_pem}) + _raise_on_status(result, 'signCertificate') + csr_hash = result.get('hash') + if not csr_hash: + raise SgxCertificateError('SGX server did not return a hash for the signing request') + _say(log, f'Signing request accepted by {csr_url}, hash: {csr_hash}') + + deadline = time.monotonic() + timeout + waiting = False + while True: + result = _rpc(csr_url, 'getCertificate', {'hash': csr_hash}) + if result.get('status', 0) == 0 and result.get('cert'): + return str(result['cert']) + if result.get('status', 0) not in (0, SIGNING_PENDING): + _raise_on_status(result, 'getCertificate') + if not waiting: + _say( + log, + 'Waiting for the SGX server to sign the request. ' + 'If the server requires manual confirmation, approve the hash above there.', + ) + waiting = True + if time.monotonic() >= deadline: + raise SgxCertificateError( + f'SGX server did not sign the request within {timeout} seconds ' + f'(request hash {csr_hash}); the current certificate is unchanged' + ) + time.sleep(SGX_SIGN_POLL_INTERVAL) + + +def _generate_key_and_csr() -> tuple[bytes, bytes]: + key = rsa.generate_private_key(public_exponent=65537, key_size=SGX_KEY_SIZE) + # The sgx client library names its requests after a random hex string as well. + subject = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, secrets.token_hex(32))]) + csr = x509.CertificateSigningRequestBuilder().subject_name(subject).sign(key, hashes.SHA256()) + key_pem = key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + return key_pem, csr.public_bytes(serialization.Encoding.PEM) + + +def _load_certificate(path: Path) -> x509.Certificate: + try: + return x509.load_pem_x509_certificate(path.read_bytes()) + except (OSError, ValueError) as err: + raise SgxCertificateError(f'Cannot read certificate {path}: {err}') from err + + +def _key_matches(key_path: Path, cert: x509.Certificate) -> bool: + try: + key = serialization.load_pem_private_key(key_path.read_bytes(), password=None) + except (OSError, ValueError, TypeError, UnsupportedAlgorithm): + return False + encoding = serialization.Encoding.DER + fmt = serialization.PublicFormat.SubjectPublicKeyInfo + return key.public_key().public_bytes(encoding, fmt) == cert.public_key().public_bytes( + encoding, fmt + ) + + +def _common_name(name: x509.Name) -> str: + attributes = name.get_attributes_for_oid(NameOID.COMMON_NAME) + return str(attributes[0].value) if attributes else name.rfc4514_string() + + +def _staging_parent(certs_dir: Path) -> Path: + """Directory on the same filesystem as the certificates, so renames stay atomic.""" + parent = certs_dir.parent + if certs_dir.stat().st_dev == parent.stat().st_dev: + return parent + return certs_dir + + +def _write_file(path: Path, data: bytes, mode: int) -> None: + fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_EXCL, mode) + with os.fdopen(fd, 'wb') as target: + os.fchmod(target.fileno(), mode) + target.write(data) + target.flush() + os.fsync(target.fileno()) + + +def _backup_existing(paths: dict[str, Path], backup_root: Path) -> Path | None: + existing = [path for path in paths.values() if path.is_file()] + if not existing: + return None + backup_root.mkdir(mode=0o700, parents=True, exist_ok=True) + os.chmod(backup_root, 0o700) + stamp = datetime.datetime.now(datetime.timezone.utc).strftime('%Y%m%dT%H%M%SZ') + destination = Path(tempfile.mkdtemp(prefix=f'{stamp}.', dir=backup_root)) + for path in existing: + shutil.copy2(path, destination / path.name) + return destination + + +def _install(staged: dict[str, Path], paths: dict[str, Path]) -> None: + """Move the staged files over the current ones, preserving their ownership.""" + as_root = os.geteuid() == 0 + for name in INSTALL_ORDER: + target = paths[name] + if as_root and target.exists(): + info = target.stat() + os.chown(staged[name], info.st_uid, info.st_gid) + os.replace(staged[name], target) diff --git a/node_cli/main.py b/node_cli/main.py index 9016f79d..b04b3869 100644 --- a/node_cli/main.py +++ b/node_cli/main.py @@ -36,6 +36,7 @@ from node_cli.cli.schains import schains_cli from node_cli.cli.wallet import wallet_cli from node_cli.cli.ssl import ssl_cli +from node_cli.cli.sgx import sgx_cli from node_cli.cli.passive_node import passive_node_cli from node_cli.cli.fair_boot import fair_boot_cli from node_cli.cli.fair_node import fair_node_cli @@ -106,6 +107,7 @@ def get_command_groups() -> List[click.Group]: staking_cli, wallet_cli, ssl_cli, + sgx_cli, ] else: return [ # type: ignore @@ -116,6 +118,7 @@ def get_command_groups() -> List[click.Group]: passive_node_cli, wallet_cli, ssl_cli, + sgx_cli, exit_cli, lvmpy_cli, ] diff --git a/node_cli/migrations/focal_to_jammy.py b/node_cli/migrations/focal_to_jammy.py index 030312c1..7ff5c04a 100644 --- a/node_cli/migrations/focal_to_jammy.py +++ b/node_cli/migrations/focal_to_jammy.py @@ -23,7 +23,7 @@ LEGACY_CHAIN, LEGACY_FAMILY, LEGACY_TABLE, - POLICY, + POLICY_ACCEPT, NFTablesManager, remove_legacy_saved_rules, ) @@ -121,7 +121,7 @@ def remove_old_iptables_rules() -> None: def migrate() -> None: nft = NFTablesManager(family=LEGACY_FAMILY, table=LEGACY_TABLE, chain=LEGACY_CHAIN) logger.info('Making sure legacy chain has default policy accept') - nft.update_chain_policy(chain=LEGACY_CHAIN, policy=POLICY) + nft.update_chain_policy(chain=LEGACY_CHAIN, policy=POLICY_ACCEPT) logger.info('Running migration from focal to jammy') remove_old_iptables_rules() diff --git a/node_cli/operations/base.py b/node_cli/operations/base.py index 8b777ef9..d2bb5c57 100644 --- a/node_cli/operations/base.py +++ b/node_cli/operations/base.py @@ -31,7 +31,6 @@ CONTAINER_CONFIG_PATH, CONTAINER_CONFIG_TMP_PATH, GLOBAL_SKALE_DIR, - NFTABLES_CHAIN_FOLDER_PATH, SKALE_DIR, ) from node_cli.core.checks import CheckType @@ -41,7 +40,7 @@ ensure_btrfs_kernel_module_autoloaded, prepare_host, ) -from node_cli.core.nftables import configure_nftables +from node_cli.core.nftables import cleanup_nftables, configure_nftables from node_cli.core.nginx import generate_nginx_config from node_cli.core.node_options import ( mark_active_node, @@ -74,7 +73,7 @@ rm_legacy_containers, system_prune, ) -from node_cli.utils.helper import cleanup_dir_content, rm_dir +from node_cli.utils.helper import rm_dir from node_cli.utils.meta import CliMetaManager, FairCliMetaManager from node_cli.utils.node_type import NodeMode, NodeType from node_cli.utils.print_formatters import print_failed_requirements_checks @@ -133,7 +132,7 @@ def update(settings: BaseNodeSettings, compose_env: dict, node_mode: NodeMode) - if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() lvmpy_install(settings.block_device) generate_nginx_config() @@ -177,7 +176,7 @@ def init(settings: BaseNodeSettings, compose_env: dict, node_mode: NodeMode) -> if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() prepare_host(env_type=settings.env_type) @@ -223,7 +222,7 @@ def init_passive( if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables(keep_accept_policy=True) prepare_host(env_type=settings.env_type) save_internal_settings(node_type=NodeType.SKALE, node_mode=NodeMode.PASSIVE) @@ -284,7 +283,7 @@ def update_passive(settings: BaseNodeSettings, compose_env: dict) -> bool: if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() ensure_filestorage_mapping() @@ -364,7 +363,7 @@ def turn_on( if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() save_internal_settings(node_type=node_type, node_mode=node_mode, backup_run=backup_run) logger.info('Launching containers on the node...') @@ -398,7 +397,7 @@ def restore( if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() lvmpy_install(settings.block_device) init_shared_space_volume(settings.env_type) @@ -432,6 +431,7 @@ def restore( def cleanup_passive(compose_env: dict, schain_name: str) -> None: turn_off(compose_env, node_type=NodeType.SKALE, node_mode=NodeMode.PASSIVE) + cleanup_nftables() cleanup_no_lvm_datadir(chain_name=schain_name) rm_dir(GLOBAL_SKALE_DIR) rm_dir(SKALE_DIR) @@ -441,6 +441,7 @@ def cleanup( node_mode: NodeMode, compose_env: dict, schain_name: Optional[str] = None, prune: bool = False ) -> None: turn_off(compose_env, node_type=NodeType.SKALE, node_mode=node_mode) + cleanup_nftables() if prune: system_prune() if node_mode == NodeMode.PASSIVE: @@ -449,5 +450,4 @@ def cleanup( cleanup_lvm_datadir() rm_dir(GLOBAL_SKALE_DIR) rm_dir(SKALE_DIR) - cleanup_dir_content(NFTABLES_CHAIN_FOLDER_PATH) cleanup_docker_configuration() diff --git a/node_cli/operations/fair.py b/node_cli/operations/fair.py index efba278b..f2e8746c 100644 --- a/node_cli/operations/fair.py +++ b/node_cli/operations/fair.py @@ -30,14 +30,13 @@ from node_cli.configs import ( CONTAINER_CONFIG_PATH, GLOBAL_SKALE_DIR, - NFTABLES_CHAIN_FOLDER_PATH, SKALE_DIR, ) from node_cli.core.checks import CheckType from node_cli.core.checks import run_checks as run_host_checks from node_cli.core.docker_config import cleanup_docker_configuration, configure_docker from node_cli.core.host import ensure_btrfs_kernel_module_autoloaded, prepare_host -from node_cli.core.nftables import configure_nftables +from node_cli.core.nftables import cleanup_nftables, configure_nftables from node_cli.core.nginx import generate_nginx_config from node_cli.core.schains import cleanup_no_lvm_datadir from node_cli.core.static_config import get_fair_chain_name @@ -70,7 +69,7 @@ system_prune, wait_for_container, ) -from node_cli.utils.helper import cleanup_dir_content, rm_dir +from node_cli.utils.helper import rm_dir from node_cli.utils.meta import FairCliMetaManager from node_cli.utils.print_formatters import print_failed_requirements_checks from node_cli.utils.node_type import NodeMode, NodeType @@ -98,7 +97,7 @@ def init_fair_boot( if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() prepare_host(env_type=settings.env_type) save_internal_settings(node_type=NodeType.FAIR, node_mode=NodeMode.ACTIVE) @@ -212,7 +211,7 @@ def update_fair_boot( if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() generate_nginx_config() fair_settings = get_settings((FairSettings, FairBaseSettings)) @@ -346,7 +345,7 @@ def restore( if not settings.skip_docker_config: configure_docker() - configure_nftables(enable_monitoring=settings.monitoring_containers) + configure_nftables() meta_manager = FairCliMetaManager() meta_manager.update_meta( @@ -375,12 +374,12 @@ def restore( def cleanup(node_mode: NodeMode, compose_env: dict, prune: bool = False) -> None: turn_off(compose_env, node_type=NodeType.FAIR, node_mode=node_mode) + cleanup_nftables() if prune: system_prune() cleanup_no_lvm_datadir() rm_dir(GLOBAL_SKALE_DIR) rm_dir(SKALE_DIR) - cleanup_dir_content(NFTABLES_CHAIN_FOLDER_PATH) cleanup_docker_configuration() diff --git a/node_cli/utils/api_auth.py b/node_cli/utils/api_auth.py new file mode 100644 index 00000000..b9c48ae1 --- /dev/null +++ b/node_cli/utils/api_auth.py @@ -0,0 +1,86 @@ +"""Provision and read the per-node admin API credential without logging it.""" + +import os +import pwd +import re +import secrets +import stat +import tempfile +from pathlib import Path + +from node_cli.configs import ADMIN_API_TOKEN_PATH, G_CONF_USER + + +class APIAuthError(RuntimeError): + pass + + +def read_api_token() -> str | None: + try: + fd = os.open(ADMIN_API_TOKEN_PATH, os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK) + except FileNotFoundError: + # Allows the updated CLI to talk to the old API during the first upgrade. + return None + except OSError as err: + raise APIAuthError('Cannot read the admin API credential') from err + try: + with os.fdopen(fd, 'r', encoding='ascii') as token_file: + info = os.fstat(token_file.fileno()) + if not stat.S_ISREG(info.st_mode) or stat.S_IMODE(info.st_mode) != 0o600: + raise ValueError('Invalid token file permissions') + token = token_file.read(66).removesuffix('\n') + if not re.fullmatch(r'[0-9a-f]{64}', token): + raise ValueError('Invalid token file contents') + except (OSError, ValueError) as err: + raise APIAuthError( + 'Admin API credential must be a valid token in a mode 0600 file' + ) from err + return token + + +def ensure_api_token() -> None: + """Publish a complete credential atomically; never overwrite an existing token.""" + owner = pwd.getpwnam(G_CONF_USER) + if os.geteuid() not in (0, owner.pw_uid): + raise APIAuthError('Only the configured node user or root can provision the API credential') + + path = Path(ADMIN_API_TOKEN_PATH) + path.parent.mkdir(mode=0o700, parents=True, exist_ok=True) + # Restrict the dedicated directory as well as the token. In particular, a + # root-run CLI must leave it traversable by the configured node user. + try: + directory_fd = os.open(path.parent, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW) + try: + if os.geteuid() == 0: + os.fchown(directory_fd, owner.pw_uid, owner.pw_gid) + os.fchmod(directory_fd, 0o700) + finally: + os.close(directory_fd) + except OSError as err: + raise APIAuthError('Cannot secure the admin API auth directory') from err + + if read_api_token() is not None: + return + + fd, tmp_path = tempfile.mkstemp(prefix='.admin-api-', dir=path.parent) + try: + with os.fdopen(fd, 'w', encoding='ascii') as token_file: + if os.geteuid() == 0: + os.fchown(token_file.fileno(), owner.pw_uid, owner.pw_gid) + token_file.write(secrets.token_hex(32) + '\n') + token_file.flush() + os.fsync(token_file.fileno()) + try: + os.link(tmp_path, path) + except FileExistsError: + # Another CLI process may have provisioned it concurrently. + pass + finally: + os.unlink(tmp_path) + if read_api_token() is None: + raise APIAuthError('Admin API credential disappeared during provisioning') + + +def get_api_headers() -> dict[str, str]: + token = read_api_token() + return {'Authorization': f'Bearer {token}'} if token is not None else {} diff --git a/node_cli/utils/docker_utils.py b/node_cli/utils/docker_utils.py index 6e4ad0a7..b9b8e547 100644 --- a/node_cli/utils/docker_utils.py +++ b/node_cli/utils/docker_utils.py @@ -38,6 +38,7 @@ REMOVED_CONTAINERS_FOLDER_PATH, ) from node_cli.core.node_options import active_fair, active_skale, passive_fair, passive_skale +from node_cli.utils.api_auth import ensure_api_token from node_cli.utils.helper import run_cmd from node_cli.utils.node_type import NodeMode, NodeType @@ -94,10 +95,6 @@ **REDIS_SERVICE_DICT, } -MONITORING_COMPOSE_SERVICES = { - 'node-exporter': 'monitor_node_exporter', - 'advisor': 'monitor_cadvisor', -} TELEGRAF_SERVICES = ('telegraf',) NOTIFICATION_COMPOSE_SERVICES = ('celery',) COMPOSE_TIMEOUT = 10 @@ -345,6 +342,7 @@ def compose_up( is_fair_boot: bool = False, services: list[str] | None = None, ): + ensure_api_token() env['PASSIVE_NODE'] = str(node_mode == NodeMode.PASSIVE) if passive_skale(node_type, node_mode) or passive_fair(node_type, node_mode): logger.info('Running containers for passive node') @@ -389,17 +387,6 @@ def compose_up( env=env, ) - if settings.monitoring_containers: - logger.info('Running monitoring containers') - run_cmd( - cmd=get_up_compose_cmd( - node_type=NodeType.SKALE, - node_mode=node_mode, - services=list(MONITORING_COMPOSE_SERVICES), - ), - env=env, - ) - def restart_nginx_container(dutils=None): dutils = dutils or docker_client() diff --git a/node_cli/utils/helper.py b/node_cli/utils/helper.py index 8a642aee..2ae9ab50 100644 --- a/node_cli/utils/helper.py +++ b/node_cli/utils/helper.py @@ -58,6 +58,7 @@ STREAM_LOG_FORMAT, ) from node_cli.configs.routes import get_route +from node_cli.utils.api_auth import APIAuthError, get_api_headers from node_cli.utils.exit_codes import CLIExitCodes from node_cli.utils.global_config import get_system_user, read_g_config from node_cli.utils.print_formatters import print_err_response @@ -66,6 +67,11 @@ HOST = f'http://{ADMIN_HOST}:{ADMIN_PORT}' +# Local operator credentials must not be sent through environment proxies or +# replaced by netrc credentials. +api_session = requests.Session() +api_session.trust_env = False + DEFAULT_ERROR_DATA = { 'status': 'error', 'payload': 'Request failed. Check API container logs', @@ -195,8 +201,12 @@ def post_request(blueprint, method, json=None, files=None): route = get_route(blueprint, method) url = construct_url(route) try: - response = requests.post(url, json=json, files=files) + response = api_session.post( + url, json=json, files=files, headers=get_api_headers(), allow_redirects=False + ) data = response.json() + except APIAuthError as err: + return 'error', str(err) except Exception as err: logger.exception('Request failed', exc_info=err) data = DEFAULT_ERROR_DATA @@ -211,8 +221,12 @@ def get_request( route = get_route(blueprint, method) url = construct_url(route) try: - response = requests.get(url, params=params) + response = api_session.get( + url, params=params, headers=get_api_headers(), allow_redirects=False + ) data = response.json() + except APIAuthError as err: + return 'error', str(err) except Exception as err: logger.exception('Request failed', exc_info=err) data = DEFAULT_ERROR_DATA @@ -408,7 +422,61 @@ def get_tmp_path(path: str | Path) -> str: return base + salt + '.tmp' + ext -def get_ssh_port(ssh_service_name='ssh'): +SSH_PORTS_ERROR = 'Cannot determine valid SSH ports. Set SSH_PORT to an integer from 1 to 65535.' + + +def _effective_sshd_config_ports() -> list[str]: + """Port values from `sshd -T`; ListenAddress entries take precedence.""" + try: + # explicit pipes instead of capture_output: environments that wrap + # subprocess.run with their own stdout/stderr reject the combination + result = subprocess.run( + [shutil.which('sshd') or '/usr/sbin/sshd', '-T'], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + check=True, + timeout=10, + ) + except (OSError, subprocess.SubprocessError) as err: + raise RuntimeError( + 'Cannot determine SSH ports from sshd -T. Set SSH_PORT to the ' + 'SSH listening port before configuring the firewall.' + ) from err + ports, listen_ports = [], [] + for line in result.stdout.splitlines(): + fields = line.split() + if len(fields) < 2: + continue + if fields[0] == 'port': + ports.append(fields[1]) + elif fields[0] == 'listenaddress': + # sshd -T expands addresses to IPv4:port or [IPv6]:port. + listen_ports.append(fields[1].rsplit(':', 1)[-1]) + return listen_ports or ports + + +def _validated_ssh_ports(values: list[str]) -> list[int]: + try: + ports = sorted({int(value) for value in values}) + except ValueError as err: + raise ValueError(SSH_PORTS_ERROR) from err + if not ports or any(not 1 <= port <= 65535 for port in ports): + raise ValueError(SSH_PORTS_ERROR) + return ports + + +def get_ssh_ports() -> list[int]: + """Return SSH_PORT or the ports from the effective default sshd config.""" + override = os.getenv('SSH_PORT') + values = [override] if override is not None else _effective_sshd_config_ports() + return _validated_ssh_ports(values) + + +def get_ssh_port(ssh_service_name='ssh') -> int: + """Return the first SSH port; firewall callers must use get_ssh_ports().""" + if ssh_service_name == 'ssh': + return get_ssh_ports()[0] try: return socket.getservbyname(ssh_service_name) except OSError: diff --git a/node_cli/utils/settings.py b/node_cli/utils/settings.py index a2a4858e..378ef4df 100644 --- a/node_cli/utils/settings.py +++ b/node_cli/utils/settings.py @@ -30,6 +30,7 @@ InternalSettings, SkalePassiveSettings, SkaleSettings, + get_settings, write_internal_settings_file, write_node_settings_file, ) @@ -85,3 +86,9 @@ def save_internal_settings( InternalSettings.model_validate(data) _remove_if_exists(INTERNAL_SETTINGS_PATH) write_internal_settings_file(path=INTERNAL_SETTINGS_PATH, data=data) + + +def get_sgx_url() -> str | None: + """SGX server URL of the node, or None when its mode has no SGX server (passive).""" + sgx_url = getattr(get_settings(), 'sgx_url', None) + return str(sgx_url).rstrip('/') if sgx_url else None diff --git a/pyproject.toml b/pyproject.toml index 1401ad34..38f32873 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "node-cli" -version = "3.2.0" +version = "3.3.1" description = "Node CLI tools" readme = "README.md" requires-python = ">=3.13" diff --git a/scripts/run_nftables_test.sh b/scripts/run_nftables_test.sh index e4bf8563..ac665408 100755 --- a/scripts/run_nftables_test.sh +++ b/scripts/run_nftables_test.sh @@ -10,5 +10,5 @@ docker run \ -e GLOBAL_SKALE_DIR="$PROJECT_DIR/tests/etc/skale" \ -e DOTENV_FILEPATH='tests/test-env' \ --cap-add=NET_ADMIN --cap-add=NET_RAW \ - --name ncli-tester ncli-tester py.test tests/core/migration_test.py tests/core/nftables_test.py $@ + --name ncli-tester ncli-tester py.test tests/core/migration_test.py tests/core/nftables_test.py tests/core/monitoring_firewall_test.py $@ diff --git a/scripts/run_tests.sh b/scripts/run_tests.sh index 1d3a3fe6..27bf2aca 100755 --- a/scripts/run_tests.sh +++ b/scripts/run_tests.sh @@ -5,4 +5,4 @@ PROJECT_DIR=$(dirname $DIR) . "$DIR/export_env.sh" -py.test --cov=$PROJECT_DIR/ --ignore=tests/core/nftables_test.py --ignore=tests/core/migration_test.py tests/ $@ +py.test --cov=$PROJECT_DIR/ --ignore=tests/core/nftables_test.py --ignore=tests/core/migration_test.py --ignore=tests/core/monitoring_firewall_test.py tests/ $@ diff --git a/tests/cli/exit_test.py b/tests/cli/exit_test.py index a3b26ed0..83b3f077 100644 --- a/tests/cli/exit_test.py +++ b/tests/cli/exit_test.py @@ -9,7 +9,7 @@ def test_exit_status(): resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) result = run_command_mock( - 'node_cli.utils.helper.requests.get', resp_mock, status, ['--format', 'json'] + 'node_cli.utils.helper.api_session.get', resp_mock, status, ['--format', 'json'] ) assert result.exit_code == 0 assert ( diff --git a/tests/cli/health_test.py b/tests/cli/health_test.py index 404d9a73..1d7b58db 100644 --- a/tests/cli/health_test.py +++ b/tests/cli/health_test.py @@ -47,7 +47,7 @@ def test_containers(): resp_mock = response_mock(requests.codes.ok, json_data=OK_LS_RESPONSE_DATA) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, containers) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, containers) assert result.exit_code == 0 assert ( result.output @@ -73,7 +73,7 @@ def test_checks(): } ] resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, schains) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, schains) assert result.exit_code == 0 assert ( @@ -81,7 +81,7 @@ def test_checks(): == 'sChain Name Config directory DKG Config file Volume Container IMA Firewall RPC Blocks\n-------------------------------------------------------------------------------------------------------------\ntest_schain True False False False False False False False False \n' # noqa ) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, schains, ['--json']) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, schains, ['--json']) assert result.exit_code == 0 assert ( @@ -99,7 +99,7 @@ def test_sgx_status(): 'status_https': True, } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, sgx) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, sgx) assert result.exit_code == 0 assert ( diff --git a/tests/cli/node_test.py b/tests/cli/node_test.py index a0c064f0..f24ef487 100644 --- a/tests/cli/node_test.py +++ b/tests/cli/node_test.py @@ -31,6 +31,7 @@ _turn_on, backup_node, cleanup_node, + configure_firewall, node_info, register_node, remove_node_from_maintenance, @@ -56,11 +57,29 @@ init_default_logger() +def test_configure_firewall_without_monitoring_option(): + with mock.patch('node_cli.cli.node.configure_firewall_rules') as configure: + result = run_command(configure_firewall, ['--yes']) + assert result.exit_code == 0, result.output + configure.assert_called_once_with() + + configure.reset_mock() + result = run_command(configure_firewall, ['--yes', '--monitoring']) + assert result.exit_code == 2 + assert 'No such option: --monitoring' in result.output + configure.assert_not_called() + + def test_register_node(inited_node, resource_alloc, mocked_g_config): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) - with mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True): + with ( + mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True), + mock.patch('node_cli.core.node.save_registered_base_port'), + mock.patch('node_cli.core.node.configure_nftables'), + mock.patch('node_cli.core.node.get_settings'), + ): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, register_node, ['--name', 'test-node', '--ip', '0.0.0.0', '--port', '8080', '-d', 'skale.test'], @@ -72,6 +91,37 @@ def test_register_node(inited_node, resource_alloc, mocked_g_config): ) # noqa +def test_register_node_firewall_failure(inited_node, resource_alloc, mocked_g_config): + """Post-registration firewall errors fail the command but report the registration.""" + resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) + with ( + mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True), + mock.patch( + 'node_cli.core.node.save_registered_base_port', + side_effect=OSError('disk error'), + ), + mock.patch('node_cli.core.node.configure_nftables'), + mock.patch('node_cli.core.node.get_settings'), + ): + result = run_command_mock( + 'node_cli.utils.helper.api_session.post', + resp_mock, + register_node, + ['--name', 'test-node', '--ip', '0.0.0.0', '--port', '8080', '-d', 'skale.test'], + ) + assert result.exit_code == CLIExitCodes.OPERATION_EXECUTION_ERROR.value + assert result.output == ( + 'Node registered in SKALE manager.\nFor more info run < skale node info >\n' + 'Command failed with following errors:\n' + '--------------------------------------------------\n' + 'Node is successfully registered in SKALE manager, but firewall ' + 'reconfiguration failed. Run < skale node configure-firewall > ' + 'to complete the setup\n' + '--------------------------------------------------\n' + f'You can find more info in {G_CONF_HOME}.skale/.skale-cli-log/debug-node-cli.log\n' + ) + + def test_register_node_with_error(inited_node, resource_alloc, mocked_g_config): resp_mock = response_mock( requests.codes.ok, @@ -79,7 +129,7 @@ def test_register_node_with_error(inited_node, resource_alloc, mocked_g_config): ) with mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, register_node, ['--name', 'test-node2', '--ip', '0.0.0.0', '--port', '80', '-d', 'skale.test'], @@ -93,9 +143,14 @@ def test_register_node_with_error(inited_node, resource_alloc, mocked_g_config): def test_register_node_with_prompted_ip(inited_node, resource_alloc, mocked_g_config): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) - with mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True): + with ( + mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True), + mock.patch('node_cli.core.node.save_registered_base_port'), + mock.patch('node_cli.core.node.configure_nftables'), + mock.patch('node_cli.core.node.get_settings'), + ): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, register_node, ['--name', 'test-node', '--port', '8080', '-d', 'skale.test'], @@ -110,9 +165,14 @@ def test_register_node_with_prompted_ip(inited_node, resource_alloc, mocked_g_co def test_register_node_with_default_port(inited_node, resource_alloc, mocked_g_config): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) - with mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True): + with ( + mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True), + mock.patch('node_cli.core.node.save_registered_base_port'), + mock.patch('node_cli.core.node.configure_nftables'), + mock.patch('node_cli.core.node.get_settings'), + ): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, register_node, ['--name', 'test-node', '-d', 'skale.test'], @@ -128,7 +188,7 @@ def test_register_node_with_default_port(inited_node, resource_alloc, mocked_g_c def test_register_with_no_alloc(mocked_g_config): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, register_node, ['--name', 'test-node', '-d', 'skale.test'], @@ -161,7 +221,7 @@ def test_node_info_node_info(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, node_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, node_info) assert result.exit_code == 0 assert ( result.output @@ -189,7 +249,7 @@ def test_node_info_node_info_not_created(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, node_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, node_info) assert result.exit_code == 0 assert result.output == 'This SKALE node is not registered on SKALE Manager yet\n' @@ -214,7 +274,7 @@ def test_node_info_node_info_frozen(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, node_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, node_info) assert result.exit_code == 0 assert ( result.output @@ -242,7 +302,7 @@ def test_node_info_node_info_left(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, node_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, node_info) assert result.exit_code == 0 assert ( result.output @@ -270,7 +330,7 @@ def test_node_info_node_info_leaving(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, node_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, node_info) assert result.exit_code == 0 assert ( result.output @@ -298,7 +358,7 @@ def test_node_info_node_info_in_maintenance(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, node_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, node_info) assert result.exit_code == 0 assert ( result.output @@ -310,7 +370,7 @@ def test_node_signature(): signature_sample = '0x1231231231' response_data = {'status': 'ok', 'payload': {'signature': signature_sample}} resp_mock = response_mock(requests.codes.ok, json_data=response_data) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, signature, ['1']) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, signature, ['1']) assert result.exit_code == 0 assert result.output == f'Signature: {signature_sample}\n' @@ -363,7 +423,7 @@ def test_restore(request, node_type, node_mode, test_user_conf, mocked_g_config, def test_maintenance_on(): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) result = run_command_mock( - 'node_cli.utils.helper.requests.post', resp_mock, set_node_in_maintenance, ['--yes'] + 'node_cli.utils.helper.api_session.post', resp_mock, set_node_in_maintenance, ['--yes'] ) assert result.exit_code == 0 assert ( @@ -375,7 +435,7 @@ def test_maintenance_on(): def test_maintenance_off(mocked_g_config): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) result = run_command_mock( - 'node_cli.utils.helper.requests.post', resp_mock, remove_node_from_maintenance + 'node_cli.utils.helper.api_session.post', resp_mock, remove_node_from_maintenance ) assert result.exit_code == 0 assert ( @@ -384,7 +444,9 @@ def test_maintenance_off(mocked_g_config): ) -def test_turn_off_maintenance_on(mocked_g_config, regular_user_conf, active_node_option, skale_active_settings): +def test_turn_off_maintenance_on( + mocked_g_config, regular_user_conf, active_node_option, skale_active_settings +): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) with ( mock.patch('subprocess.run', new=subprocess_run_mock), @@ -393,7 +455,7 @@ def test_turn_off_maintenance_on(mocked_g_config, regular_user_conf, active_node mock.patch('node_cli.cli.node.TYPE', NodeType.SKALE), ): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, _turn_off, ['--maintenance-on', '--yes'], @@ -406,7 +468,7 @@ def test_turn_off_maintenance_on(mocked_g_config, regular_user_conf, active_node assert result.exit_code == 0 with mock.patch('node_cli.utils.docker_utils.is_container_running', return_value=True): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, _turn_off, ['--maintenance-on', '--yes'], @@ -415,7 +477,9 @@ def test_turn_off_maintenance_on(mocked_g_config, regular_user_conf, active_node assert result.exit_code == CLIExitCodes.UNSAFE_UPDATE -def test_turn_on_maintenance_off(mocked_g_config, regular_user_conf, active_node_option, skale_active_settings): +def test_turn_on_maintenance_off( + mocked_g_config, regular_user_conf, active_node_option, skale_active_settings +): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) with ( mock.patch('subprocess.run', new=subprocess_run_mock), @@ -425,7 +489,7 @@ def test_turn_on_maintenance_off(mocked_g_config, regular_user_conf, active_node mock.patch('node_cli.cli.node.TYPE', NodeType.SKALE), ): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, _turn_on, [regular_user_conf.as_posix(), '--maintenance-off', '--sync-schains', '--yes'], @@ -443,7 +507,7 @@ def test_set_domain_name(): with mock.patch('node_cli.utils.decorators.is_node_inited', return_value=True): result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, _set_domain_name, ['-d', 'skale.test', '--yes'], diff --git a/tests/cli/passive_node_test.py b/tests/cli/passive_node_test.py index e8419220..3375c9d6 100644 --- a/tests/cli/passive_node_test.py +++ b/tests/cli/passive_node_test.py @@ -42,6 +42,7 @@ def test_init_passive(mocked_g_config, clean_node_options, passive_user_conf): mock.patch('subprocess.run', new=subprocess_run_mock), mock.patch('node_cli.core.node.init_passive_op'), mock.patch('node_cli.core.node.is_base_containers_alive', return_value=True), + mock.patch('node_cli.core.node.enable_firewall_default_drop_when_port_available'), mock.patch('node_cli.core.resources.get_disk_size', return_value=BIG_DISK_SIZE), mock.patch('node_cli.operations.base.configure_nftables'), mock.patch('node_cli.utils.decorators.is_node_inited', return_value=False), @@ -60,6 +61,7 @@ def test_init_passive_archive(mocked_g_config, clean_node_options, passive_user_ pathlib.Path(NODE_DATA_PATH).mkdir(parents=True, exist_ok=True) with ( mock.patch('node_cli.core.node.is_base_containers_alive', return_value=True), + mock.patch('node_cli.core.node.enable_firewall_default_drop_when_port_available'), mock.patch('node_cli.operations.base.cleanup_volume_artifacts'), mock.patch('node_cli.operations.base.download_skale_node'), mock.patch('node_cli.operations.base.sync_skale_node'), diff --git a/tests/cli/schains_test.py b/tests/cli/schains_test.py index ffefdea8..a15a699f 100644 --- a/tests/cli/schains_test.py +++ b/tests/cli/schains_test.py @@ -79,7 +79,7 @@ def test_ls(): }, ] resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, ls) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, ls) assert result.exit_code == 0 assert ( result.output @@ -100,14 +100,14 @@ def test_dkg(): } ] resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, dkg) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, dkg) assert result.exit_code == 0 assert ( result.output == ' sChain Name DKG Status Added At sChain Status\n---------------------------------------------------------------------\nmelodic-aldhibah IN_PROGRESS Jan 08 2020 15:26:52 Exists \n' # noqa ) - result = run_command_mock('node_cli.utils.helper.requests.get', resp_mock, dkg, ['--all']) + result = run_command_mock('node_cli.utils.helper.api_session.get', resp_mock, dkg, ['--all']) assert result.exit_code == 0 assert ( result.output @@ -164,7 +164,7 @@ def test_get_schain_config(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) result = run_command_mock( - 'node_cli.utils.helper.requests.get', resp_mock, get_schain_config, ['test1'] + 'node_cli.utils.helper.api_session.get', resp_mock, get_schain_config, ['test1'] ) assert result.exit_code == 0 assert ( @@ -189,7 +189,7 @@ def test_schain_rules(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) result = run_command_mock( - 'node_cli.utils.helper.requests.get', resp_mock, show_rules, ['schain-test'] + 'node_cli.utils.helper.api_session.get', resp_mock, show_rules, ['schain-test'] ) assert result.exit_code == 0 print(repr(result.output)) @@ -221,7 +221,7 @@ def test_info(): } resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'ok'}) result = run_command_mock( - 'node_cli.utils.helper.requests.get', resp_mock, info_, ['attractive-ed-asich'] + 'node_cli.utils.helper.api_session.get', resp_mock, info_, ['attractive-ed-asich'] ) assert ( result.output @@ -232,7 +232,7 @@ def test_info(): payload = ['error'] resp_mock = response_mock(requests.codes.ok, json_data={'payload': payload, 'status': 'error'}) result = run_command_mock( - 'node_cli.utils.helper.requests.get', resp_mock, info_, ['schain not found'] + 'node_cli.utils.helper.api_session.get', resp_mock, info_, ['schain not found'] ) assert ( result.output diff --git a/tests/cli/sgx_test.py b/tests/cli/sgx_test.py new file mode 100644 index 00000000..6e5c2c47 --- /dev/null +++ b/tests/cli/sgx_test.py @@ -0,0 +1,197 @@ +import json +from unittest.mock import Mock + +import pytest +import requests_mock + +from node_cli.cli.sgx import cert_status, options, renew, sgx_cli +from node_cli.core import sgx as core_sgx +from node_cli.utils import api_auth, helper +from node_cli.utils.exit_codes import CLIExitCodes +from tests.fixtures.settings import NODE_SKALE_ACTIVE +from tests.fixtures.sgx import FakeSgxWallet +from tests.helper import run_command + +SETTINGS_SGX_URL = NODE_SKALE_ACTIVE['sgx_url'] + + +@pytest.fixture +def certs_dir(tmp_path, monkeypatch): + directory = tmp_path / 'node_data' / 'sgx_certs' + directory.mkdir(parents=True) + monkeypatch.setattr(core_sgx, 'SGX_CERTS_PATH', str(directory)) + monkeypatch.setattr(core_sgx, 'SGX_CERTS_BACKUP_PATH', str(directory.parent / 'backup')) + monkeypatch.setattr(core_sgx, 'SGX_SIGN_POLL_INTERVAL', 0) + return directory + + +@pytest.fixture +def rpc(): + with requests_mock.Mocker() as mock: + yield mock + + +@pytest.fixture +def api_token(tmp_path, monkeypatch): + path = tmp_path / 'admin-api.token' + path.write_text('ab' * 32 + '\n') + path.chmod(0o600) + monkeypatch.setattr(api_auth, 'ADMIN_API_TOKEN_PATH', path) + return 'ab' * 32 + + +@pytest.mark.parametrize('json_format', [False, True]) +def test_options_through_admin( + rpc, api_token, inited_node, skale_active_settings, mocked_g_config, json_format +): + payload = { + 'flags': {'auto_sign': False, 'log_level': 2}, + 'effective': {'rpc_port': 1026, 'rpc_client_certificate_required': True}, + 'build': None, + } + rpc.get( + helper.construct_url('/api/v1/info/sgx-options'), + json={'status': 'ok', 'payload': payload}, + ) + args = ['sgx', 'options'] + (['--json'] if json_format else []) + result = run_command(sgx_cli, args) + assert result.exit_code == 0, result.output + if json_format: + assert json.loads(result.output) == payload + else: + assert 'SGX option' in result.output + assert 'flags.auto_sign' in result.output + assert 'false' in result.output + assert len(rpc.request_history) == 1 + assert rpc.last_request.method == 'GET' + assert rpc.last_request.headers['Authorization'] == f'Bearer {api_token}' + + +@pytest.mark.parametrize( + ('code', 'message'), + [(401, 'A valid node CLI credential is required'), (503, 'SGX server unavailable')], +) +def test_options_reports_api_errors( + rpc, api_token, inited_node, skale_active_settings, mocked_g_config, code, message +): + rpc.get( + helper.construct_url('/api/v1/info/sgx-options'), + status_code=code, + json={'status': 'error', 'payload': message}, + ) + result = run_command(options, ['--json']) + assert result.exit_code == CLIExitCodes.BAD_API_RESPONSE.value + assert message in result.output + + +def test_options_needs_an_sgx_node( + rpc, api_token, inited_node, skale_passive_settings, mocked_g_config +): + result = run_command(options) + assert result.exit_code == CLIExitCodes.NODE_STATE_ERROR.value + assert 'no SGX server configured' in result.output + assert not rpc.called + + +def test_cert_status_without_certificate(certs_dir): + result = run_command(sgx_cli, ['sgx', 'cert-status']) + assert result.exit_code == 0 + assert 'Private key' in result.output + # three file rows plus the notice + assert result.output.count('missing') == 4 + assert 'skale sgx renew' in result.output + + +def test_cert_status_shows_certificate_details(certs_dir, rpc): + FakeSgxWallet(rpc).issue_files(certs_dir) + result = run_command(cert_status) + assert result.exit_code == 0 + assert 'sgx-wallet-ca' in result.output + assert 'yes' in result.output + assert 'renew' not in result.output + + +def test_cert_status_json(certs_dir, rpc): + FakeSgxWallet(rpc).issue_files(certs_dir) + result = run_command(cert_status, ['--json']) + assert result.exit_code == 0 + data = json.loads(result.output) + assert data['complete'] is True + assert data['key_matches'] is True + assert data['issuer'] == 'sgx-wallet-ca' + + +def test_cert_status_check_uses_configured_sgx_url(certs_dir, rpc, skale_active_settings): + wallet = FakeSgxWallet(rpc, SETTINGS_SGX_URL) + wallet.issue_files(certs_dir) + result = run_command(cert_status, ['--check']) + assert result.exit_code == 0, result.output + assert 'accepted the certificate, version 1.83.0' in result.output + assert wallet.calls == [('sgx', 'getServerVersion')] + + +def test_cert_status_check_reports_rejection(certs_dir, rpc, skale_active_settings): + FakeSgxWallet(rpc, SETTINGS_SGX_URL, reject_clients=True).issue_files(certs_dir) + result = run_command(cert_status, ['--check', '--json']) + assert result.exit_code == CLIExitCodes.OPERATION_EXECUTION_ERROR.value + assert 'rejected the TLS connection' in result.output + + +def test_cert_status_check_needs_an_sgx_node(certs_dir, skale_passive_settings): + result = run_command(cert_status, ['--check']) + assert result.exit_code == CLIExitCodes.NODE_STATE_ERROR.value + assert 'no SGX server configured' in result.output + + +def test_renew_asks_for_confirmation(certs_dir, monkeypatch): + core = Mock(side_effect=AssertionError('must not run')) + monkeypatch.setattr('node_cli.cli.sgx.renew_certificate', core) + result = run_command(renew, input='n\n') + assert result.exit_code == 1 + assert 'Aborted' in result.output + core.assert_not_called() + + +def test_renew_replaces_certificate( + certs_dir, rpc, inited_node, skale_active_settings, mocked_g_config +): + wallet = FakeSgxWallet(rpc, SETTINGS_SGX_URL) + wallet.issue_files(certs_dir) + before = {path.name: path.read_bytes() for path in certs_dir.iterdir()} + result = run_command(renew, ['--yes']) + assert result.exit_code == 0, result.output + assert 'hash: ' in result.output + assert 'accepted the new certificate' in result.output + assert 'Previous certificate files were copied to' in result.output + assert 'New SGX client certificate installed' in result.output + after = {path.name: path.read_bytes() for path in certs_dir.iterdir()} + assert sorted(after) == ['sgx.crt', 'sgx.csr', 'sgx.key'] + assert after != before + assert wallet.calls[-1] == ('sgx', 'getServerVersion') + + +def test_renew_skip_verify_and_timeout_flags( + certs_dir, rpc, inited_node, skale_active_settings, mocked_g_config +): + wallet = FakeSgxWallet(rpc, SETTINGS_SGX_URL, reject_clients=True) + result = run_command(renew, ['--yes', '--skip-verify', '--timeout', '5']) + assert result.exit_code == 0, result.output + assert ('sgx', 'getServerVersion') not in wallet.calls + assert 'Previous certificate files' not in result.output + + +def test_renew_needs_an_sgx_node(certs_dir, inited_node, skale_passive_settings, mocked_g_config): + result = run_command(renew, ['--yes']) + assert result.exit_code == CLIExitCodes.NODE_STATE_ERROR.value + assert 'no SGX server configured' in result.output + + +def test_renew_reports_failures( + certs_dir, inited_node, skale_active_settings, mocked_g_config, monkeypatch +): + core = Mock(side_effect=core_sgx.SgxCertificateError('SGX server refused signCertificate')) + monkeypatch.setattr('node_cli.cli.sgx.renew_certificate', core) + result = run_command(renew, ['--yes']) + assert result.exit_code == CLIExitCodes.OPERATION_EXECUTION_ERROR.value + assert 'SGX server refused signCertificate' in result.output + core.assert_called_once_with(SETTINGS_SGX_URL, timeout=600, verify=True, log=print) diff --git a/tests/cli/wallet_test.py b/tests/cli/wallet_test.py index 234f4f1a..306cd823 100644 --- a/tests/cli/wallet_test.py +++ b/tests/cli/wallet_test.py @@ -35,7 +35,7 @@ def test_wallet_info(): response_mock = MagicMock() response_mock.status_code = requests.codes.ok response_mock.json = Mock(return_value=response_data) - result = run_command_mock('node_cli.utils.helper.requests.get', response_mock, wallet_info) + result = run_command_mock('node_cli.utils.helper.api_session.get', response_mock, wallet_info) assert result.exit_code == 0 expected = ( '--------------------------------------------------\n' @@ -47,7 +47,7 @@ def test_wallet_info(): assert result.output == expected result = run_command_mock( - 'node_cli.utils.helper.requests.get', response_mock, wallet_info, ['--format', 'json'] + 'node_cli.utils.helper.api_session.get', response_mock, wallet_info, ['--format', 'json'] ) assert result.exit_code == 0 expected = '{"address": "simple_address", "eth_balance": 13, "skale_balance": 123}\n' @@ -57,7 +57,7 @@ def test_wallet_info(): def test_wallet_send(): resp_mock = response_mock(requests.codes.ok, {'status': 'ok', 'payload': None}) result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, send, ['0x00000000000000000000000000000000', '10', '--yes'], @@ -72,7 +72,7 @@ def test_wallet_send_with_error(): {'status': 'error', 'payload': ['Strange error']}, ) result = run_command_mock( - 'node_cli.utils.helper.requests.post', + 'node_cli.utils.helper.api_session.post', resp_mock, send, ['0x00000000000000000000000000000000', '10', '--yes'], diff --git a/tests/core/core_node_test.py b/tests/core/core_node_test.py index b28719d4..4c707c67 100644 --- a/tests/core/core_node_test.py +++ b/tests/core/core_node_test.py @@ -299,7 +299,7 @@ def test_update_node(regular_user_conf, mocked_g_config, resource_file, inited_n ), ): with mock.patch( - 'node_cli.utils.helper.requests.get', return_value=safe_update_api_response() + 'node_cli.utils.helper.api_session.get', return_value=safe_update_api_response() ): # noqa result = update( regular_user_conf.as_posix(), @@ -320,7 +320,7 @@ def test_update_node(regular_user_conf, mocked_g_config, resource_file, inited_n ) @mock.patch('node_cli.core.node.is_admin_running', return_value=False) @mock.patch('node_cli.core.node.is_api_running', return_value=False) -@mock.patch('node_cli.utils.helper.requests.get') +@mock.patch('node_cli.utils.helper.api_session.get') def test_is_update_safe_when_admin_and_api_not_running( mock_requests_get, mock_is_api_running, mock_is_admin_running, node_type, node_mode ): @@ -330,7 +330,7 @@ def test_is_update_safe_when_admin_and_api_not_running( @mock.patch('node_cli.core.node.is_admin_running', return_value=False) @mock.patch('node_cli.core.node.is_api_running', return_value=True) -@mock.patch('node_cli.utils.helper.requests.get') +@mock.patch('node_cli.utils.helper.api_session.get') def test_is_update_safe_when_admin_not_running_for_passive( mock_requests_get, mock_is_api_running, mock_is_admin_running ): @@ -352,7 +352,7 @@ def test_is_update_safe_when_admin_not_running_for_passive( ids=['api_safe', 'api_unsafe'], ) @mock.patch('node_cli.core.node.is_admin_running', return_value=True) -@mock.patch('node_cli.utils.helper.requests.get') +@mock.patch('node_cli.utils.helper.api_session.get') def test_is_update_safe_when_admin_running( mock_requests_get, mock_is_admin_running, api_is_safe, expected_result, node_type, node_mode ): @@ -369,7 +369,7 @@ def test_is_update_safe_when_admin_running( ) @mock.patch('node_cli.core.node.is_admin_running', return_value=False) @mock.patch('node_cli.core.node.is_api_running', return_value=True) -@mock.patch('node_cli.utils.helper.requests.get') +@mock.patch('node_cli.utils.helper.api_session.get') def test_is_update_safe_when_only_api_running_for_regular( mock_requests_get, mock_is_api_running, @@ -392,7 +392,7 @@ def test_is_update_safe_when_only_api_running_for_regular( ], ) @mock.patch('node_cli.core.node.is_admin_running', return_value=True) -@mock.patch('node_cli.utils.helper.requests.get') +@mock.patch('node_cli.utils.helper.api_session.get') def test_is_update_safe_when_api_call_fails( mock_requests_get, mock_is_admin_running, node_type, node_mode ): diff --git a/tests/core/core_sgx_test.py b/tests/core/core_sgx_test.py new file mode 100644 index 00000000..618e9606 --- /dev/null +++ b/tests/core/core_sgx_test.py @@ -0,0 +1,287 @@ +import datetime +import json +import stat +from pathlib import Path + +import pytest +import requests_mock + +from node_cli.core import sgx +from tests.fixtures.sgx import SGX_URL, FakeSgxWallet + +CERT_FILES = ['sgx.crt', 'sgx.csr', 'sgx.key'] + + +@pytest.fixture(autouse=True) +def fast_polling(monkeypatch): + monkeypatch.setattr(sgx, 'SGX_SIGN_POLL_INTERVAL', 0) + + +@pytest.fixture +def rpc(): + with requests_mock.Mocker() as mock: + yield mock + + +@pytest.fixture +def certs_dir(tmp_path): + directory = tmp_path / 'node_data' / 'sgx_certs' + directory.mkdir(parents=True) + return directory + + +def snapshot(directory: Path) -> dict[str, bytes]: + return {path.name: path.read_bytes() for path in directory.iterdir()} + + +def names(directory: Path) -> list[str]: + return sorted(path.name for path in directory.iterdir()) + + +@pytest.mark.parametrize( + 'sgx_url, expected', + [ + ('https://sgx.example.com:1026', 'http://sgx.example.com:1027'), + ('http://127.0.0.1:2026/', 'http://127.0.0.1:2027'), + ('https://[fd00::1]:1026', 'http://[fd00::1]:1027'), + ], +) +def test_csr_server_url(sgx_url, expected): + assert sgx.csr_server_url(sgx_url) == expected + + +@pytest.mark.parametrize('sgx_url', ['https://sgx.example.com', 'ftp://host:1026', 'nonsense']) +def test_csr_server_url_rejects_incomplete_urls(sgx_url): + with pytest.raises(sgx.SgxCertificateError): + sgx.csr_server_url(sgx_url) + + +def test_status_reports_missing_files(certs_dir): + status = sgx.get_certificate_status(certs_dir) + assert status['complete'] is False + assert status['present'] == {'key': False, 'csr': False, 'crt': False} + assert 'subject' not in status + assert sgx.get_certificate_status(certs_dir / 'absent')['complete'] is False + + +def test_status_describes_certificate(certs_dir, rpc): + wallet = FakeSgxWallet(rpc) + wallet.issue_files(certs_dir) + status = sgx.get_certificate_status(certs_dir) + assert status['complete'] is True + assert status['subject'] == 'ab' * 32 + assert status['issuer'] == 'sgx-wallet-ca' + assert status['key_matches'] is True + assert status['expired'] is False + assert status['not_yet_valid'] is False + assert status['expires_soon'] is False + assert 363 <= status['days_left'] <= 365 + assert status['fingerprint_sha256'] == wallet.last_issued.fingerprint(sgx.hashes.SHA256()).hex( + ':' + ) + json.dumps(status) + + +@pytest.mark.parametrize( + 'starts_in, ends_in, expired, not_yet_valid, expires_soon', + [ + (-1, 365, False, False, False), + (-2, -1, True, False, False), + (1, 365, False, True, False), + (-1, 1, False, False, True), + ], +) +def test_status_with_legacy_certificate_dates( + certs_dir, rpc, monkeypatch, starts_in, ends_in, expired, not_yet_valid, expires_soon +): + wallet = FakeSgxWallet(rpc) + wallet.issue_files(certs_dir) + now = datetime.datetime.now(datetime.timezone.utc).replace(microsecond=0) + start = now + datetime.timedelta(days=starts_in) + end = now + datetime.timedelta(days=ends_in) + + class LegacyCertificate: + not_valid_before = start.replace(tzinfo=None) + not_valid_after = end.replace(tzinfo=None) + + def __getattr__(self, name): + if name in ('not_valid_before_utc', 'not_valid_after_utc'): + raise AttributeError(name) + return getattr(wallet.last_issued, name) + + monkeypatch.setattr(sgx, '_load_certificate', lambda _: LegacyCertificate()) + status = sgx.get_certificate_status(certs_dir) + assert status['not_valid_before'] == start.isoformat(timespec='seconds') + assert status['not_valid_after'] == end.isoformat(timespec='seconds') + assert status['expired'] is expired + assert status['not_yet_valid'] is not_yet_valid + assert status['expires_soon'] is expires_soon + assert status['key_matches'] is True + + +def test_status_detects_key_mismatch_and_partial_sets(certs_dir, rpc): + wallet = FakeSgxWallet(rpc) + wallet.issue_files(certs_dir) + other = certs_dir.parent / 'other' + wallet.issue_files(other) + (certs_dir / 'sgx.key').write_bytes((other / 'sgx.key').read_bytes()) + assert sgx.get_certificate_status(certs_dir)['key_matches'] is False + (certs_dir / 'sgx.key').write_text('not a key') + assert sgx.get_certificate_status(certs_dir)['key_matches'] is False + (certs_dir / 'sgx.key').unlink() + status = sgx.get_certificate_status(certs_dir) + assert status['key_matches'] is None + assert status['complete'] is False + + +def test_status_rejects_unreadable_certificate(certs_dir): + (certs_dir / 'sgx.crt').write_text('garbage') + with pytest.raises(sgx.SgxCertificateError, match='Cannot read certificate'): + sgx.get_certificate_status(certs_dir) + + +def test_check_certificate(certs_dir, rpc): + wallet = FakeSgxWallet(rpc) + with pytest.raises(sgx.SgxCertificateError, match='missing files: key, crt'): + sgx.check_certificate(SGX_URL, certs_dir) + wallet.issue_files(certs_dir) + assert sgx.check_certificate(SGX_URL, certs_dir) == '1.83.0' + assert wallet.client_certs == [(str(certs_dir / 'sgx.crt'), str(certs_dir / 'sgx.key'))] + + +def test_check_certificate_reports_rejection(certs_dir, rpc): + FakeSgxWallet(rpc, reject_clients=True).issue_files(certs_dir) + with pytest.raises(sgx.SgxCertificateError, match='rejected the TLS connection'): + sgx.check_certificate(SGX_URL, certs_dir) + + +def test_renew_installs_new_certificate_and_keeps_backup(certs_dir, rpc): + wallet = FakeSgxWallet(rpc) + wallet.issue_files(certs_dir) + before = snapshot(certs_dir) + backups = certs_dir.parent / 'sgx_certs_backup' + messages: list[str] = [] + + result = sgx.renew_certificate( + SGX_URL, directory=certs_dir, backup_root=backups, log=messages.append + ) + + assert names(certs_dir) == CERT_FILES + assert names(certs_dir.parent) == ['sgx_certs', 'sgx_certs_backup'] + after = snapshot(certs_dir) + assert all(after[name] != before[name] for name in CERT_FILES) + assert stat.S_IMODE((certs_dir / 'sgx.key').stat().st_mode) == 0o600 + assert stat.S_IMODE((certs_dir / 'sgx.crt').stat().st_mode) == 0o644 + assert stat.S_IMODE((certs_dir / 'sgx.csr').stat().st_mode) == 0o644 + + status = sgx.get_certificate_status(certs_dir) + assert status['complete'] and status['key_matches'] and status['issuer'] == 'sgx-wallet-ca' + assert status['fingerprint_sha256'] == wallet.last_issued.fingerprint(sgx.hashes.SHA256()).hex( + ':' + ) + assert result['server_version'] == '1.83.0' + assert result['fingerprint_sha256'] == status['fingerprint_sha256'] + + backup = Path(result['backup']) + assert backup.parent == backups + assert stat.S_IMODE(backups.stat().st_mode) == 0o700 + assert snapshot(backup) == before + + assert wallet.calls == [ + ('csr', 'signCertificate'), + ('csr', 'getCertificate'), + ('sgx', 'getServerVersion'), + ] + staged_crt, staged_key = wallet.client_certs[0] + assert Path(staged_crt).parent != certs_dir + assert Path(staged_crt).parent.name.startswith('.sgx_certs.') + assert not Path(staged_crt).exists() and not Path(staged_key).exists() + assert any('hash: ' in message for message in messages) + assert not any('Waiting' in message for message in messages) + + +def test_renew_waits_for_manual_approval(certs_dir, rpc): + wallet = FakeSgxWallet(rpc, pending_polls=2) + messages: list[str] = [] + sgx.renew_certificate(SGX_URL, directory=certs_dir, log=messages.append, timeout=30) + assert wallet.calls.count(('csr', 'getCertificate')) == 3 + assert sum('Waiting' in message for message in messages) == 1 + assert sgx.get_certificate_status(certs_dir)['complete'] + + +def test_renew_times_out_without_touching_current_files(certs_dir, rpc): + wallet = FakeSgxWallet(rpc, pending_polls=10**6) + wallet.issue_files(certs_dir) + before = snapshot(certs_dir) + backups = certs_dir.parent / 'sgx_certs_backup' + with pytest.raises(sgx.SgxCertificateError, match='did not sign the request within 0'): + sgx.renew_certificate(SGX_URL, directory=certs_dir, backup_root=backups, timeout=0) + assert snapshot(certs_dir) == before + assert names(certs_dir.parent) == ['sgx_certs'] + assert ('sgx', 'getServerVersion') not in wallet.calls + + +def test_renew_aborts_when_server_rejects_new_certificate(certs_dir, rpc): + wallet = FakeSgxWallet(rpc, reject_clients=True) + wallet.issue_files(certs_dir) + before = snapshot(certs_dir) + with pytest.raises(sgx.SgxCertificateError, match='rejected the TLS connection'): + sgx.renew_certificate(SGX_URL, directory=certs_dir) + assert snapshot(certs_dir) == before + assert names(certs_dir.parent) == ['sgx_certs'] + + +def test_renew_can_skip_server_verification(certs_dir, rpc): + wallet = FakeSgxWallet(rpc, reject_clients=True) + result = sgx.renew_certificate(SGX_URL, directory=certs_dir, verify=False) + assert result['server_version'] is None + assert result['backup'] is None + assert ('sgx', 'getServerVersion') not in wallet.calls + assert names(certs_dir) == CERT_FILES + assert names(certs_dir.parent) == ['sgx_certs'] + + +def test_renew_rejects_certificate_issued_for_another_key(certs_dir, rpc): + wallet = FakeSgxWallet(rpc, wrong_key=True) + wallet.issue_files(certs_dir) + before = snapshot(certs_dir) + with pytest.raises(sgx.SgxCertificateError, match='does not match the generated key'): + sgx.renew_certificate(SGX_URL, directory=certs_dir) + assert snapshot(certs_dir) == before + + +def test_renew_reports_signing_refusal(certs_dir, rpc): + wallet = FakeSgxWallet(rpc, sign_error='CSR rejected by policy') + wallet.issue_files(certs_dir) + before = snapshot(certs_dir) + with pytest.raises(sgx.SgxCertificateError, match='CSR rejected by policy'): + sgx.renew_certificate(SGX_URL, directory=certs_dir) + assert snapshot(certs_dir) == before + assert wallet.calls == [('csr', 'signCertificate')] + + +def test_renew_reports_transport_errors(certs_dir, rpc): + rpc.post('http://127.0.0.1:1027/', status_code=502) + with pytest.raises(sgx.SgxCertificateError, match='Cannot call signCertificate'): + sgx.renew_certificate(SGX_URL, directory=certs_dir) + assert names(certs_dir) == [] + assert names(certs_dir.parent) == ['sgx_certs'] + + +def test_renew_creates_missing_directory(tmp_path, rpc): + FakeSgxWallet(rpc) + directory = tmp_path / 'node_data' / 'sgx_certs' + result = sgx.renew_certificate(SGX_URL, directory=directory) + assert result['backup'] is None + assert names(directory) == CERT_FILES + + +def test_renew_reports_unwritable_directory(tmp_path, rpc): + FakeSgxWallet(rpc) + parent = tmp_path / 'node_data' + parent.mkdir(mode=0o500) + try: + with pytest.raises(sgx.SgxCertificateError, match='Cannot prepare certificate files'): + sgx.renew_certificate(SGX_URL, directory=parent / 'sgx_certs') + finally: + parent.chmod(0o700) diff --git a/tests/core/iptables_test.py b/tests/core/iptables_test.py index 2926b917..19b34255 100644 --- a/tests/core/iptables_test.py +++ b/tests/core/iptables_test.py @@ -1,12 +1,89 @@ import socket +import subprocess import mock +import pytest -from node_cli.utils.helper import get_ssh_port +from node_cli.utils.helper import get_ssh_port, get_ssh_ports -def test_get_ssh_port(): - assert get_ssh_port() == 22 +@pytest.fixture(autouse=True) +def clear_ssh_override(monkeypatch): + monkeypatch.delenv('SSH_PORT', raising=False) + + +def test_service_lookup_compatibility(): assert get_ssh_port('http') == 80 with mock.patch.object(socket, 'getservbyname', side_effect=OSError): - assert get_ssh_port() == 22 + assert get_ssh_port('missing-service') == 22 + + +@pytest.mark.parametrize( + 'config,expected', + [ + ('port 22\nlistenaddress 0.0.0.0:22\nlistenaddress [::]:22\n', [22]), + ('port 2222\n', [2222]), + ('port 2222\nport 2200\nport 2222\n', [2200, 2222]), + ('port 22\nlistenaddress 192.0.2.1:2222\nlistenaddress [::1]:2200\n', [2200, 2222]), + ], +) +def test_effective_sshd_ports(config, expected): + with ( + mock.patch('node_cli.utils.helper.shutil.which', return_value='/usr/sbin/sshd'), + mock.patch('node_cli.utils.helper.subprocess.run') as run, + ): + run.return_value.stdout = config + assert get_ssh_ports() == expected + run.assert_called_once_with( + ['/usr/sbin/sshd', '-T'], + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + check=True, + timeout=10, + ) + + +def test_sshd_invocation_arguments_are_valid(): + """Run a real process: mocks cannot catch invalid subprocess.run + argument combinations such as capture_output with stdout/stderr.""" + with mock.patch('node_cli.utils.helper.shutil.which', return_value='/bin/echo'): + with pytest.raises(ValueError, match='SSH_PORT'): + get_ssh_ports() + + +def test_ssh_port_override(monkeypatch): + monkeypatch.setenv('SSH_PORT', '2222') + with mock.patch('node_cli.utils.helper.subprocess.run') as run: + assert get_ssh_ports() == [2222] + assert get_ssh_port() == 2222 + run.assert_not_called() + + +@pytest.mark.parametrize('value', ['', 'ssh', '0', '-1', '65536']) +def test_invalid_ssh_port_override(monkeypatch, value): + monkeypatch.setenv('SSH_PORT', value) + with pytest.raises(ValueError, match='SSH_PORT'): + get_ssh_ports() + + +@pytest.mark.parametrize( + 'error', + [ + FileNotFoundError(), + subprocess.CalledProcessError(1, 'sshd'), + subprocess.TimeoutExpired('sshd', 10), + ], +) +def test_ssh_detection_failure_does_not_guess(error): + with mock.patch('node_cli.utils.helper.subprocess.run', side_effect=error): + with pytest.raises(RuntimeError, match='SSH_PORT'): + get_ssh_port() + + +@pytest.mark.parametrize('config', ['', 'port invalid\n', 'port 65536\n']) +def test_invalid_sshd_output(config): + with mock.patch('node_cli.utils.helper.subprocess.run') as run: + run.return_value.stdout = config + with pytest.raises(ValueError, match='SSH_PORT'): + get_ssh_ports() diff --git a/tests/core/monitoring_firewall_test.py b/tests/core/monitoring_firewall_test.py new file mode 100644 index 00000000..d85f56b2 --- /dev/null +++ b/tests/core/monitoring_firewall_test.py @@ -0,0 +1,80 @@ +from pathlib import Path +from unittest.mock import patch + +import pytest + +import node_cli.core.nftables as firewall +from node_cli.core.nftables import NFTablesManager, Rule + + +@pytest.fixture +def monitoring_firewall(monkeypatch, tmp_path): + """Run only in the isolated nftables test container (requires NET_ADMIN).""" + monkeypatch.setenv('SSH_PORT', '22') + monkeypatch.delenv('SCHAIN_BASE_PORT', raising=False) + monkeypatch.delenv('FIREWALL_DEFAULT_DROP', raising=False) + monkeypatch.setattr(firewall, 'NODE_CONFIG_PATH', str(tmp_path / 'node.json')) + monkeypatch.setattr(firewall, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'user.conf')) + monkeypatch.setattr(firewall, 'NFTABLES_SKALE_BASE_CONFIG_PATH', str(tmp_path / 'base.conf')) + monkeypatch.setattr(firewall, 'NFTABLES_CHAIN_CONFIG_WILDCARD', str(tmp_path / 'chains/*')) + (tmp_path / 'chains').mkdir() + (tmp_path / 'user.conf').touch() + manager = NFTablesManager(table='monitoring_test') + manager.create_table_if_not_exists() + manager.create_chain_if_not_exists(manager.chain, hook='input') + try: + yield manager + finally: + manager.execute_cmd( + {'nftables': [{'delete': {'table': {'family': manager.family, 'name': manager.table}}}]} + ) + + +@pytest.mark.parametrize('legacy_rule_copies', [0, 1, 2]) +def test_monitoring_accepts_absent_after_setup_and_reload(monitoring_firewall, legacy_rule_copies): + manager = monitoring_firewall + expressions = [Rule(manager.chain, 'tcp', port).to_expr() for port in (8080, 9100)] + for _ in range(legacy_rule_copies): + for expr in expressions: + manager._execute_rule_with_op('add', manager.chain, expr) + + for _ in range(2): + manager.setup_firewall() + assert manager.get_chain_policy(manager.chain) == 'drop' + for expr in expressions: + assert not manager.rule_exists(manager.chain, expr) + manager.verify_critical_accepts() + + firewall.save_nftables_base_rules(manager.get_base_ruleset()) + manager.execute_cmd( + {'nftables': [{'delete': {'table': {'family': manager.family, 'name': manager.table}}}]} + ) + rc, _, error = manager.nft.cmd(f'include "{firewall.NFTABLES_SKALE_BASE_CONFIG_PATH}"') + assert rc == 0, error + assert manager.get_chain_policy(manager.chain) == 'drop' + for expr in expressions: + assert not manager.rule_exists(manager.chain, expr) + manager.verify_critical_accepts() + + +def test_monitoring_cleanup_preserves_explicit_user_rules(monitoring_firewall): + manager = monitoring_firewall + Path(firewall.NFTABLES_USER_CONFIG_PATH).write_text('tcp dport 8080 counter accept\n') + manager.add_rule(Rule(manager.chain, 'tcp', 8080)) + manager.setup_firewall() + expr = Rule(manager.chain, 'tcp', 8080).to_expr() + assert not manager.rule_exists(manager.chain, expr) + assert manager.rule_exists(firewall.USER_CHAIN, expr) + + +def test_monitoring_cleanup_preserves_custom_ssh_port(monitoring_firewall, monkeypatch): + monkeypatch.setenv('SSH_PORT', '9100') + manager = monitoring_firewall + manager.add_rule(Rule(manager.chain, 'tcp', 9100)) + with patch.object( + manager, 'delete_rule_by_handle', wraps=manager.delete_rule_by_handle + ) as delete: + manager.setup_firewall() + delete.assert_not_called() + manager.verify_critical_accepts() + assert manager.rule_exists(manager.chain, Rule(manager.chain, 'tcp', 9100).to_expr()) diff --git a/tests/core/nftables_envelope_test.py b/tests/core/nftables_envelope_test.py new file mode 100644 index 00000000..1eccc602 --- /dev/null +++ b/tests/core/nftables_envelope_test.py @@ -0,0 +1,35 @@ +from unittest.mock import patch + +import pytest + +from node_cli.core.nftables import NFTablesManager, Rule + + +@pytest.mark.parametrize( + 'old_range,new_range', + [ + ((20128, 20191), (30000, 30063)), + ((10000, 18191), (30000, 38191)), + ((10000, 18191), (10000, 10063)), + ((10000, 10063), (10000, 18191)), + ], +) +def test_replace_envelope_preserves_current_range_and_unrelated_rules(old_range, new_range): + manager = NFTablesManager.__new__(NFTablesManager) + manager.chain = 'skale' + current = Rule('skale', 'tcp', *new_range) + rules = [ + {'handle': 1, 'expr': Rule('skale', 'tcp', *old_range).to_expr()}, + {'handle': 2, 'expr': current.to_expr()}, + {'handle': 3, 'expr': Rule('skale', 'tcp', 1026, 1031, action='drop').to_expr()}, + {'handle': 4, 'expr': Rule('skale', 'tcp', 5000, 5010).to_expr()}, + {'handle': 5, 'expr': Rule('skale', 'tcp', *old_range, action='drop').to_expr()}, + ] + with ( + patch.object(manager, 'get_rules', return_value=rules), + patch.object(manager, 'delete_rule_by_handle') as delete, + patch.object(manager, 'add_rule') as add, + ): + manager._ensure_envelope(new_range) + delete.assert_called_once_with(1) + add.assert_called_once_with(current) diff --git a/tests/core/nftables_test.py b/tests/core/nftables_test.py index b6490b08..c5e30617 100644 --- a/tests/core/nftables_test.py +++ b/tests/core/nftables_test.py @@ -4,7 +4,15 @@ import nftables -from node_cli.core.nftables import NFTablesManager, Rule +import node_cli.core.nftables as nftables_core +from node_cli.core.nftables import ( + NFTablesError, + NFTablesManager, + Rule, + conntrack_accept_expr, + dport_match, + get_schain_ports_envelope, +) @pytest.fixture(scope='module') @@ -17,6 +25,11 @@ def nft_manager(): manager.flush() +@pytest.fixture(autouse=True) +def ssh_ports(monkeypatch): + monkeypatch.setattr(nftables_core, 'get_ssh_ports', lambda: [22]) + + @pytest.fixture def mock_nft_output(): """Fixture for mocking nftables output.""" @@ -102,6 +115,50 @@ def test_create_chain_if_not_exists(mock_exists, mock_execute, nft_manager): mock_execute.assert_called_once() +@patch('nftables.Nftables.cmd') +def test_update_chain_policy_uses_given_policy(mock_cmd, nft_manager): + """Test that policy update applies the requested policy.""" + mock_cmd.return_value = (0, '', '') + with patch.object(NFTablesManager, 'chain_exists', return_value=True): + nft_manager.update_chain_policy(chain='skale', policy='drop') + assert mock_cmd.call_args[0][0] == 'add chain inet filter skale { policy drop ; }' + + mock_cmd.return_value = (1, '', 'some error') + with patch.object(NFTablesManager, 'chain_exists', return_value=True): + with pytest.raises(NFTablesError): + nft_manager.update_chain_policy(chain='skale', policy='drop') + + +@patch('nftables.Nftables.cmd') +def test_get_chain_policy(mock_cmd, nft_manager): + listing = { + 'nftables': [ + { + 'chain': { + 'family': 'inet', + 'table': 'filter', + 'name': 'skale', + 'hook': 'input', + 'policy': 'drop', + } + } + ] + } + mock_cmd.return_value = (0, json.dumps(listing), '') + assert nft_manager.get_chain_policy('skale') == 'drop' + + mock_cmd.return_value = (1, '', 'No such file or directory') + assert nft_manager.get_chain_policy('skale') is None + + mock_cmd.return_value = (0, 'not-json', '') + assert nft_manager.get_chain_policy('skale') is None + + # malformed but valid JSON must not raise - rollback depends on it + for output in ('[]', 'null', '{"nftables": [{"chain": null}]}', '{"nftables": "x"}'): + mock_cmd.return_value = (0, output, '') + assert nft_manager.get_chain_policy('skale') is None + + @pytest.mark.parametrize( 'rule_data', [ @@ -122,16 +179,786 @@ def test_add_rule(mock_exists, mock_execute, nft_manager, rule_data): @patch.object(NFTablesManager, 'execute_cmd') -def test_setup_firewall(mock_execute, nft_manager): +@patch.object(NFTablesManager, 'rule_exists') +def test_add_rule_udp_uses_udp_payload(mock_exists, mock_execute, nft_manager): + """Test that udp rules match udp dport, not tcp.""" + mock_exists.return_value = False + + nft_manager.add_rule(Rule(chain='INPUT', protocol='udp', first_port=53)) + expr = mock_execute.call_args[0][0]['nftables'][0]['add']['rule']['expr'] + assert expr[0]['match']['left']['payload'] == {'protocol': 'udp', 'field': 'dport'} + + +@patch.object(NFTablesManager, 'execute_cmd') +@patch.object(NFTablesManager, 'rule_exists') +def test_add_rule_icmpv6(mock_exists, mock_execute, nft_manager): + """Test icmpv6 rule addition.""" + mock_exists.return_value = False + + nft_manager.add_rule(Rule(chain='INPUT', protocol='icmpv6', icmp_type='nd-neighbor-solicit')) + expr = mock_execute.call_args[0][0]['nftables'][0]['add']['rule']['expr'] + assert expr[0]['match']['left']['payload'] == {'protocol': 'icmpv6', 'field': 'type'} + assert expr[0]['match']['right'] == 'nd-neighbor-solicit' + + +@patch('nftables.Nftables.cmd') +def test_get_dynamic_chain_port_ranges(mock_cmd, nft_manager): + """Test collection of port ranges covered by skale-admin chains.""" + listing = { + 'nftables': [ + {'chain': {'family': 'inet', 'table': 'filter', 'name': 'skale'}}, + {'chain': {'family': 'inet', 'table': 'filter', 'name': 'skale-test'}}, + { + 'rule': { + 'family': 'inet', + 'table': 'filter', + 'chain': 'skale', + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, + 'right': 22, + } + }, + {'accept': None}, + ], + } + }, + { + 'rule': { + 'family': 'inet', + 'table': 'filter', + 'chain': 'skale-test', + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'ip', 'field': 'saddr'}}, + 'right': '1.2.3.4', + } + }, + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, + 'right': 10001, + } + }, + {'accept': None}, + ], + } + }, + { + 'rule': { + 'family': 'inet', + 'table': 'filter', + 'chain': 'skale-test', + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, + 'right': {'range': [10000, 10063]}, + } + }, + {'drop': None}, + ], + } + }, + ] + } + mock_cmd.return_value = (0, json.dumps(listing), '') + assert nft_manager.get_dynamic_chain_port_ranges() == [('skale-test', 10000, 10063)] + + mock_cmd.return_value = (1, '', 'No such file or directory') + assert nft_manager.get_dynamic_chain_port_ranges() == [] + + mock_cmd.return_value = (1, '', 'some other error') + with pytest.raises(NFTablesError): + nft_manager.get_dynamic_chain_port_ranges() + + # unparseable or wrong-shaped output keeps the typed contract + for output in ('not-json', 'null', '[]'): + mock_cmd.return_value = (0, output, '') + with pytest.raises(NFTablesError): + nft_manager.get_dynamic_chain_port_ranges() + + +def test_validate_dynamic_ranges(nft_manager): + """Test envelope validation against dynamic chain ranges.""" + with patch.object( + NFTablesManager, + 'get_dynamic_chain_port_ranges', + return_value=[('skale-test', 10064, 10127)], + ): + nft_manager.validate_dynamic_ranges((10000, 18191)) + with pytest.raises(NFTablesError): + nft_manager.validate_dynamic_ranges((10128, 18191)) + + +def test_verify_critical_accepts(nft_manager): + with patch.object(NFTablesManager, 'rule_exists', return_value=True): + nft_manager.verify_critical_accepts() + with patch.object(NFTablesManager, 'rule_exists', return_value=False): + with pytest.raises(NFTablesError): + nft_manager.verify_critical_accepts() + + +def test_ensure_default_drop(nft_manager): + with patch.multiple( + NFTablesManager, + verify_critical_accepts=Mock(), + get_chain_policy=Mock(return_value='accept'), + update_chain_policy=Mock(), + ): + nft_manager.ensure_default_drop() + NFTablesManager.verify_critical_accepts.assert_called_once() + NFTablesManager.update_chain_policy.assert_called_once_with(chain='skale', policy='drop') + + with patch.multiple( + NFTablesManager, + verify_critical_accepts=Mock(), + get_chain_policy=Mock(return_value='drop'), + update_chain_policy=Mock(), + ): + nft_manager.ensure_default_drop() + NFTablesManager.update_chain_policy.assert_not_called() + + +def test_ensure_default_accept(nft_manager): + with patch.multiple( + NFTablesManager, + get_chain_policy=Mock(return_value='drop'), + update_chain_policy=Mock(), + ): + nft_manager.ensure_default_accept() + NFTablesManager.update_chain_policy.assert_called_once_with(chain='skale', policy='accept') + + with patch.multiple( + NFTablesManager, + get_chain_policy=Mock(return_value='accept'), + update_chain_policy=Mock(), + ): + nft_manager.ensure_default_accept() + NFTablesManager.update_chain_policy.assert_not_called() + + # unreadable policy must not skip the rollback + with patch.multiple( + NFTablesManager, + get_chain_policy=Mock(return_value=None), + update_chain_policy=Mock(), + ): + nft_manager.ensure_default_accept() + NFTablesManager.update_chain_policy.assert_called_once_with(chain='skale', policy='accept') + + +def test_remove_stale_envelope_rules(nft_manager): + stale_rule = { + 'handle': 7, + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, + 'right': {'range': [10000, 18191]}, + } + }, + {'counter': None}, + {'accept': None}, + ], + } + current_rule = { + 'handle': 8, + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, + 'right': {'range': [30000, 38191]}, + } + }, + {'counter': None}, + {'accept': None}, + ], + } + sgx_drop_rule = { + 'handle': 9, + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'tcp', 'field': 'dport'}}, + 'right': {'range': [1026, 1031]}, + } + }, + {'counter': None}, + {'drop': None}, + ], + } + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=[stale_rule, current_rule, sgx_drop_rule]), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_stale_envelope_rules((30000, 38191)) + NFTablesManager.delete_rule_by_handle.assert_called_once_with(7) + + def envelope_rule(handle, first_port, last_port): + return { + 'handle': handle, + 'expr': [ + dport_match('tcp', first_port, last_port), + {'counter': None}, + {'accept': None}, + ], + } + + # passive node: the single-chain envelope moved to a new base port + rules = [envelope_rule(11, 58128, 58191), envelope_rule(12, 58192, 58255)] + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=rules), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_stale_envelope_rules((58192, 58255)) + NFTablesManager.delete_rule_by_handle.assert_called_once_with(11) + + # upgrade transition: full-node envelope replaced by a single-chain one + rules = [envelope_rule(13, 10000, 18191), envelope_rule(14, 58128, 58191)] + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=rules), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_stale_envelope_rules((58128, 58191)) + NFTablesManager.delete_rule_by_handle.assert_called_once_with(13) + + +def test_get_schain_ports_envelope_default(monkeypatch, tmp_path): + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(tmp_path / 'nonexistent.json')) + assert get_schain_ports_envelope() == (10000, 18191) + + +def test_get_schain_ports_envelope_env_override(monkeypatch): + monkeypatch.setenv('SCHAIN_BASE_PORT', '30000') + assert get_schain_ports_envelope() == (30000, 38191) + + monkeypatch.setenv('SCHAIN_BASE_PORT', 'not-a-port') + with pytest.raises(NFTablesError): + get_schain_ports_envelope() + + monkeypatch.setenv('SCHAIN_BASE_PORT', '65000') + with pytest.raises(NFTablesError): + get_schain_ports_envelope() + + monkeypatch.setenv('SCHAIN_BASE_PORT', '1000') + with pytest.raises(NFTablesError): + get_schain_ports_envelope() + + +def test_get_schain_ports_envelope_malformed_node_config(monkeypatch, tmp_path): + config_path = tmp_path / 'node_config.json' + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(config_path)) + + for content in ( + 'not-json', + '[1, 2]', + '{"node_base_port": "not-a-port"}', + '{"node_base_port": true}', + '{"node_base_port": -1}', + ): + config_path.write_text(content) + # falls back to the default base port + assert get_schain_ports_envelope() == (10000, 18191) + + # an invalid primary value must not mask a valid fallback, which then + # carries its own single-chain envelope size + config_path.write_text(json.dumps({'node_base_port': 'bad', 'schain_base_port': 20128})) + assert get_schain_ports_envelope() == (20128, 20191) + + +def test_get_schain_ports_envelope_from_node_config(monkeypatch, tmp_path): + config_path = tmp_path / 'node_config.json' + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(config_path)) + + # active node: node_base_port saved at registration wins, full allocation + config_path.write_text( + json.dumps({'node_id': 1, 'node_base_port': 20128, 'schain_base_port': 30000}) + ) + assert get_schain_ports_envelope() == (20128, 28319) + + # passive/fair node: schain_base_port is one allocated chain's base port + # and reserves only that chain's range + config_path.write_text(json.dumps({'node_id': 1, 'schain_base_port': 20128})) + assert get_schain_ports_envelope() == (20128, 20191) + + # a chain at the top of a high node allocation is valid on passive nodes + config_path.write_text(json.dumps({'node_id': 1, 'schain_base_port': 58128})) + assert get_schain_ports_envelope() == (58128, 58191) + + # while a full node allocation must fit below the port maximum + config_path.write_text(json.dumps({'node_id': 1, 'node_base_port': 58128})) + with pytest.raises(NFTablesError): + get_schain_ports_envelope() + + config_path.write_text(json.dumps({'node_id': 1, 'schain_base_port': 65500})) + with pytest.raises(NFTablesError): + get_schain_ports_envelope() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_setup_firewall(mock_execute, nft_manager, monkeypatch, tmp_path): """Test complete firewall setup.""" + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(tmp_path / 'nonexistent.json')) + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'user.conf')) with patch.multiple( NFTablesManager, table_exists=Mock(return_value=False), chain_exists=Mock(return_value=False), rule_exists=Mock(return_value=False), + verify_critical_accepts=Mock(), + get_dynamic_chain_port_ranges=Mock(return_value=[]), + get_chain_policy=Mock(return_value='accept'), + update_chain_policy=Mock(), + apply_user_rules=Mock(), ): nft_manager.setup_firewall() assert mock_execute.called + NFTablesManager.apply_user_rules.assert_called_once() + + added_exprs = [ + call.args[0]['nftables'][0]['add']['rule']['expr'] + for call in mock_execute.call_args_list + if 'rule' in call.args[0]['nftables'][0].get('add', {}) + ] + envelope_exprs = [ + expr + for expr in added_exprs + if expr[0].get('match', {}).get('right') == {'range': [10000, 18191]} + and {'accept': None} in expr + ] + assert len(envelope_exprs) == 1 + icmpv6_exprs = [ + expr + for expr in added_exprs + if expr[0].get('match', {}).get('left', {}).get('payload', {}).get('protocol') + == 'icmpv6' + ] + assert len(icmpv6_exprs) == len(nftables_core.ICMPV6_ACCEPT_TYPES) + assert not any( + expr[0].get('match', {}).get('right') == 'source-quench' for expr in added_exprs + ) + + NFTablesManager.update_chain_policy.assert_any_call( + chain='INPUT', policy='accept', family='ip', table='filter' + ) + NFTablesManager.update_chain_policy.assert_any_call(chain='skale', policy='drop') + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_setup_firewall_default_drop_disabled(mock_execute, nft_manager, monkeypatch, tmp_path): + """Test that FIREWALL_DEFAULT_DROP=False keeps the accept policy. + + Rollback must not be blocked by envelope validation. + """ + monkeypatch.setenv('FIREWALL_DEFAULT_DROP', 'False') + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(tmp_path / 'nonexistent.json')) + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'user.conf')) + with patch.multiple( + NFTablesManager, + table_exists=Mock(return_value=True), + chain_exists=Mock(return_value=True), + rule_exists=Mock(return_value=True), + validate_dynamic_ranges=Mock(), + ensure_default_drop=Mock(), + ensure_default_accept=Mock(), + update_chain_policy=Mock(), + apply_user_rules=Mock(), + ): + nft_manager.setup_firewall() + NFTablesManager.validate_dynamic_ranges.assert_not_called() + NFTablesManager.ensure_default_drop.assert_not_called() + NFTablesManager.ensure_default_accept.assert_called_once() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_setup_firewall_keep_accept_policy(mock_execute, nft_manager, monkeypatch, tmp_path): + """Test that keep_accept_policy skips the drop flip (passive init).""" + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(tmp_path / 'nonexistent.json')) + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'user.conf')) + with patch.multiple( + NFTablesManager, + table_exists=Mock(return_value=True), + chain_exists=Mock(return_value=True), + rule_exists=Mock(return_value=True), + validate_dynamic_ranges=Mock(), + ensure_default_drop=Mock(), + ensure_default_accept=Mock(), + update_chain_policy=Mock(), + apply_user_rules=Mock(), + ): + nft_manager.setup_firewall(keep_accept_policy=True) + NFTablesManager.ensure_default_drop.assert_not_called() + NFTablesManager.ensure_default_accept.assert_called_once() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_setup_firewall_rollback_flips_accept_early( + mock_execute, nft_manager, monkeypatch, tmp_path +): + """A failing step must not block the rollback to accept.""" + monkeypatch.setenv('FIREWALL_DEFAULT_DROP', 'False') + monkeypatch.setattr(nftables_core, 'NODE_CONFIG_PATH', str(tmp_path / 'nonexistent.json')) + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'user.conf')) + with patch.multiple( + NFTablesManager, + table_exists=Mock(return_value=True), + chain_exists=Mock(return_value=True), + rule_exists=Mock(return_value=True), + update_chain_policy=Mock(), + ensure_default_accept=Mock(), + apply_user_rules=Mock(side_effect=NFTablesError('bad user rule')), + ): + with pytest.raises(NFTablesError): + nft_manager.setup_firewall() + NFTablesManager.ensure_default_accept.assert_called_once() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_setup_firewall_rollback_survives_bad_envelope( + mock_execute, nft_manager, monkeypatch, tmp_path +): + """An invalid base port configuration must not block the rollback.""" + monkeypatch.setenv('FIREWALL_DEFAULT_DROP', 'False') + monkeypatch.setenv('SCHAIN_BASE_PORT', 'not-a-port') + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'user.conf')) + with patch.multiple( + NFTablesManager, + table_exists=Mock(return_value=True), + chain_exists=Mock(return_value=True), + ensure_default_accept=Mock(), + ): + with pytest.raises(NFTablesError): + nft_manager.setup_firewall() + NFTablesManager.ensure_default_accept.assert_called_once() + + +def test_remove_source_quench_rule(nft_manager): + source_quench_rule = { + 'handle': 11, + 'expr': [ + { + 'match': { + 'left': {'payload': {'protocol': 'icmp', 'field': 'type'}}, + 'op': '==', + 'right': 'source-quench', + } + }, + {'counter': {'packets': 0, 'bytes': 0}}, + {'accept': None}, + ], + } + other_rule = { + 'handle': 12, + 'expr': [ + { + 'match': { + 'left': {'payload': {'protocol': 'icmp', 'field': 'type'}}, + 'op': '==', + 'right': 'destination-unreachable', + } + }, + {'counter': {'packets': 0, 'bytes': 0}}, + {'accept': None}, + ], + } + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=[source_quench_rule, other_rule]), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_source_quench_rule() + NFTablesManager.delete_rule_by_handle.assert_called_once_with(11, chain='skale') + + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=[other_rule]), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_source_quench_rule() + NFTablesManager.delete_rule_by_handle.assert_not_called() + + +def test_remove_misordered_udp_drop(nft_manager): + udp_drop = { + 'handle': 5, + 'expr': [ + { + 'match': { + 'left': {'payload': {'protocol': 'ip', 'field': 'protocol'}}, + 'op': '==', + 'right': 'udp', + } + }, + {'counter': {'packets': 0, 'bytes': 0}}, + {'drop': None}, + ], + } + udp_dns_accept = { + 'handle': 6, + 'expr': [ + { + 'match': { + 'op': '==', + 'left': {'payload': {'protocol': 'udp', 'field': 'dport'}}, + 'right': 53, + } + }, + {'counter': {'packets': 0, 'bytes': 0}}, + {'accept': None}, + ], + } + # drop shadows the accept -> removed + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=[udp_drop, udp_dns_accept]), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_misordered_udp_drop() + NFTablesManager.delete_rule_by_handle.assert_called_once_with(5) + + # accept missing -> drop removed so the accept can land above it + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=[udp_drop]), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_misordered_udp_drop() + NFTablesManager.delete_rule_by_handle.assert_called_once_with(5) + + # correct order -> untouched + with patch.multiple( + NFTablesManager, + get_rules=Mock(return_value=[udp_dns_accept, udp_drop]), + delete_rule_by_handle=Mock(), + ): + nft_manager.remove_misordered_udp_drop() + NFTablesManager.delete_rule_by_handle.assert_not_called() + + +@patch('nftables.Nftables.cmd') +def test_apply_user_rules(mock_cmd, nft_manager, monkeypatch, tmp_path): + user_conf = tmp_path / 'user.conf' + # inline comments and multiline rules must reach the nft parser verbatim, + # exactly as the boot include would read them + content = ( + '# custom services\n' + 'tcp dport 5000 counter accept # legacy exporter\n' + 'tcp dport {\n' + ' 6000,\n' + ' 6001,\n' + '} counter accept\n' + ) + user_conf.write_text(content) + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(user_conf)) + + mock_cmd.return_value = (0, '', '') + nft_manager.apply_user_rules() + assert mock_cmd.call_args[0][0] == ( + 'flush chain inet filter skale_user\n' + 'table inet filter {\n' + 'chain skale_user {\n' + f'{content}\n' + '}\n' + '}' + ) + + # missing file still flushes, so removed rules disappear + monkeypatch.setattr(nftables_core, 'NFTABLES_USER_CONFIG_PATH', str(tmp_path / 'absent')) + nft_manager.apply_user_rules() + assert mock_cmd.call_args[0][0].startswith('flush chain inet filter skale_user') + + mock_cmd.return_value = (1, '', 'syntax error') + with pytest.raises(NFTablesError): + nft_manager.apply_user_rules() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_ensure_user_chain_jump(mock_execute, nft_manager): + with patch.object(NFTablesManager, 'rule_exists', return_value=False): + nft_manager.ensure_user_chain_jump() + cmd = mock_execute.call_args[0][0]['nftables'][0] + assert cmd['insert']['rule']['expr'] == [{'jump': {'target': 'skale_user'}}] + + mock_execute.reset_mock() + with patch.object(NFTablesManager, 'rule_exists', return_value=True): + nft_manager.ensure_user_chain_jump() + mock_execute.assert_not_called() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_create_user_chain_if_not_exists(mock_execute, nft_manager): + with patch.object(NFTablesManager, 'chain_exists', return_value=False): + nft_manager.create_user_chain_if_not_exists() + chain = mock_execute.call_args[0][0]['nftables'][0]['add']['chain'] + assert chain == {'family': 'inet', 'table': 'filter', 'name': 'skale_user'} + # regular chain: no hook, priority or policy + + mock_execute.reset_mock() + with patch.object(NFTablesManager, 'chain_exists', return_value=True): + nft_manager.create_user_chain_if_not_exists() + mock_execute.assert_not_called() + + +def test_remove_user_rules_from_main_chain(nft_manager): + user_rule_expr = Rule(chain='skale_user', protocol='tcp', first_port=5000).to_expr() + chains = { + # the reloaded user chain is the parsed form of user.conf; native + # comments live outside expr, so matching ignores them + 'skale_user': [{'handle': 3, 'expr': user_rule_expr, 'comment': 'service #1'}], + 'skale': [ + {'handle': 7, 'expr': user_rule_expr}, + {'handle': 8, 'expr': Rule(chain='skale', protocol='tcp', first_port=22).to_expr()}, + ], + } + with ( + patch.object(NFTablesManager, 'get_rules', side_effect=lambda chain: chains[chain]), + patch.object(NFTablesManager, 'delete_rule_by_handle') as mock_delete, + ): + nft_manager.remove_user_rules_from_main_chain() + mock_delete.assert_called_once_with(7) + + # nothing loaded from user.conf - main chain untouched + chains['skale_user'] = [] + with ( + patch.object(NFTablesManager, 'get_rules', side_effect=lambda chain: chains[chain]), + patch.object(NFTablesManager, 'delete_rule_by_handle') as mock_delete, + ): + nft_manager.remove_user_rules_from_main_chain() + mock_delete.assert_not_called() + + +@patch.object(NFTablesManager, 'execute_cmd') +def test_delete_chain(mock_execute, nft_manager): + nft_manager.delete_chain('skale-test') + chain_spec = {'family': 'inet', 'table': 'filter', 'name': 'skale-test'} + assert mock_execute.call_args[0][0] == { + 'nftables': [{'flush': {'chain': chain_spec}}, {'delete': {'chain': chain_spec}}] + } + + +def test_cleanup_firewall(nft_manager): + critical_rule = {'handle': 1, 'expr': conntrack_accept_expr()} + envelope_rule = { + 'handle': 2, + 'expr': [dport_match('tcp', 10000, 18191), {'counter': None}, {'accept': None}], + } + watchdog_rule = { + 'handle': 3, + 'expr': Rule(chain='skale', protocol='tcp', first_port=3009).to_expr(), + } + ssh_rule = {'handle': 4, 'expr': Rule(chain='skale', protocol='tcp', first_port=22).to_expr()} + with patch.multiple( + NFTablesManager, + ensure_default_accept=Mock(), + _remove_rule_by_expr=Mock(), + _table_chain_names=Mock(return_value=['skale', 'skale_user', 'skale-mychain']), + delete_chain=Mock(), + get_rules=Mock(return_value=[critical_rule, envelope_rule, watchdog_rule, ssh_rule]), + delete_rule_by_handle=Mock(), + ): + nft_manager.cleanup_firewall() + NFTablesManager.ensure_default_accept.assert_called_once() + NFTablesManager._remove_rule_by_expr.assert_called_once_with( + 'skale', [{'jump': {'target': 'skale_user'}}] + ) + assert sorted(call.args[0] for call in NFTablesManager.delete_chain.call_args_list) == [ + 'skale-mychain', + 'skale_user', + ] + deleted = sorted( + call.args[0] for call in NFTablesManager.delete_rule_by_handle.call_args_list + ) + assert deleted == [2, 3] + + +def test_cleanup_firewall_without_ssh_detection(nft_manager, monkeypatch): + def raise_runtime_error(): + raise RuntimeError('SSH_PORT required') + + monkeypatch.setattr(nftables_core, 'get_ssh_ports', raise_runtime_error) + ssh_rule = {'handle': 4, 'expr': Rule(chain='skale', protocol='tcp', first_port=22).to_expr()} + with patch.multiple( + NFTablesManager, + ensure_default_accept=Mock(), + _remove_rule_by_expr=Mock(), + _table_chain_names=Mock(return_value=['skale']), + delete_chain=Mock(), + get_rules=Mock(return_value=[ssh_rule]), + delete_rule_by_handle=Mock(), + ): + # without detection the ssh rule is not in the keep set, which is + # safe because the policy is accept by then + nft_manager.cleanup_firewall() + NFTablesManager.delete_rule_by_handle.assert_called_once_with(4) + + +def test_cleanup_nftables(monkeypatch, tmp_path): + chains_dir = tmp_path / 'chains' + chains_dir.mkdir() + (chains_dir / 'skale-x.conf').write_text('chain skale-x {\n}\n') + base_conf = tmp_path / 'base.conf' + base_conf.write_text('old content') + monkeypatch.setattr(nftables_core, 'NFTABLES_CHAIN_FOLDER_PATH', str(chains_dir)) + monkeypatch.setattr(nftables_core, 'NFTABLES_SKALE_BASE_CONFIG_PATH', str(base_conf)) + + with patch.multiple( + NFTablesManager, + table_exists=Mock(return_value=True), + chain_exists=Mock(return_value=True), + cleanup_firewall=Mock(), + get_base_ruleset=Mock(return_value='table inet firewall {\n\tchain skale {\n\t}\n}'), + ): + nftables_core.cleanup_nftables() + NFTablesManager.cleanup_firewall.assert_called_once() + assert list(chains_dir.iterdir()) == [] + assert base_conf.read_text() == 'table inet firewall {\n\tchain skale {\n\t}\n}' + + # nothing configured: only the persisted state is cleared + base_conf.write_text('old content') + with patch.multiple( + NFTablesManager, + table_exists=Mock(return_value=False), + cleanup_firewall=Mock(), + ): + nftables_core.cleanup_nftables() + NFTablesManager.cleanup_firewall.assert_not_called() + assert base_conf.read_text() == '' + + +def test_save_nftables_base_rules(monkeypatch, tmp_path): + base_conf = tmp_path / 'base.conf' + monkeypatch.setattr(nftables_core, 'NFTABLES_SKALE_BASE_CONFIG_PATH', str(base_conf)) + ruleset = ( + 'table inet firewall {\n' + '\tchain skale {\n' + '\t\ttype filter hook input priority filter + 1; policy drop;\n' + '\t\tjump skale_user\n' + '\t\tct state established,related counter accept\n' + '\t}\n' + '}' + ) + nftables_core.save_nftables_base_rules(ruleset) + saved = base_conf.read_text() + + # user chain is declared before the skale chain that jumps to it + assert saved.index('chain skale_user {') < saved.index('chain skale {') + assert nftables_core.NFTABLES_USER_CONFIG_PATH in saved + assert nftables_core.NFTABLES_CHAIN_CONFIG_WILDCARD in saved + # the include lives only inside the user chain, not in the skale chain + skale_chain_part = saved[saved.index('chain skale {') :] + assert 'include "' + nftables_core.NFTABLES_USER_CONFIG_PATH not in skale_chain_part def test_invalid_protocol(nft_manager): diff --git a/tests/core/ssh_firewall_test.py b/tests/core/ssh_firewall_test.py new file mode 100644 index 00000000..7689c881 --- /dev/null +++ b/tests/core/ssh_firewall_test.py @@ -0,0 +1,53 @@ +from unittest.mock import Mock, patch + +import pytest + +from node_cli.core.nftables import NFTablesError, NFTablesManager, Rule, conntrack_accept_expr + + +@pytest.fixture +def manager(): + # These checks exercise rule generation without accessing the host firewall. + manager = NFTablesManager.__new__(NFTablesManager) + manager.chain = 'skale' + return manager + + +def test_allow_all_ssh_ports(manager): + rules = [] + with ( + patch('node_cli.core.nftables.get_ssh_ports', return_value=[2200, 2222]), + patch.object(manager, '_ensure_rule'), + patch.object(manager, 'remove_misordered_udp_drop'), + patch.object(manager, 'add_rule', side_effect=rules.append), + ): + manager._add_service_accepts() + allowed_tcp_ports = {rule.first_port for rule in rules if rule.protocol == 'tcp'} + assert {2200, 2222} <= allowed_tcp_ports + assert 22 not in allowed_tcp_ports + + +def test_refuse_drop_until_all_ssh_ports_are_allowed(manager): + installed = [conntrack_accept_expr(), Rule('skale', 'tcp', 2200).to_expr()] + with ( + patch('node_cli.core.nftables.get_ssh_ports', return_value=[2200, 2222]), + patch.object(manager, 'get_rules', side_effect=lambda _: [{'expr': e} for e in installed]), + patch.object(manager, 'get_chain_policy', return_value='accept'), + patch.object(manager, 'update_chain_policy') as update_policy, + ): + with pytest.raises(NFTablesError, match='ssh port 2222'): + manager.ensure_default_drop() + update_policy.assert_not_called() + installed.append(Rule('skale', 'tcp', 2222).to_expr()) + manager.ensure_default_drop() + update_policy.assert_called_once_with(chain='skale', policy='drop') + + +def test_detection_failure_prevents_drop(manager): + manager.update_chain_policy = Mock() + with patch( + 'node_cli.core.nftables.get_ssh_ports', side_effect=RuntimeError('SSH_PORT required') + ): + with pytest.raises(RuntimeError, match='SSH_PORT'): + manager.ensure_default_drop() + manager.update_chain_policy.assert_not_called() diff --git a/tests/fixtures/sgx.py b/tests/fixtures/sgx.py new file mode 100644 index 00000000..14d26680 --- /dev/null +++ b/tests/fixtures/sgx.py @@ -0,0 +1,145 @@ +"""Test double for the SGX wallet ports used by the SGX certificate commands.""" + +import datetime +import hashlib +from pathlib import Path + +import requests +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.x509.oid import NameOID + +from node_cli.core import sgx + +SGX_URL = 'https://127.0.0.1:1026' +ONE_DAY = datetime.timedelta(days=1) + + +def _response(result: dict) -> dict: + return {'id': 0, 'jsonrpc': '2.0', 'result': result} + + +def _pem(obj) -> bytes: + return obj.public_bytes(serialization.Encoding.PEM) + + +def _fingerprint(cert: x509.Certificate) -> bytes: + return cert.fingerprint(hashes.SHA256()) + + +class FakeSgxWallet: + """Serves the CSR signing port and the main port of an SGX wallet through requests_mock. + + The main port emulates TLS client authentication: it only answers when the presented + certificate was issued by this wallet and matches the presented key. + """ + + def __init__( + self, + mock, + sgx_url: str = SGX_URL, + *, + pending_polls: int = 0, + sign_error: str | None = None, + reject_clients: bool = False, + wrong_key: bool = False, + version: str = '1.83.0', + ): + self.sgx_url = sgx_url + self.pending_polls = pending_polls + self.sign_error = sign_error + self.reject_clients = reject_clients + self.wrong_key = wrong_key + self.version = version + self.ca_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + self.ca_name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, 'sgx-wallet-ca')]) + self.calls: list[tuple[str, str]] = [] + self.client_certs: list[tuple[str, str]] = [] + self.pending: dict[str, x509.CertificateSigningRequest] = {} + self.issued: list[x509.Certificate] = [] + mock.post(sgx_url.rstrip('/') + '/', json=self._main_port) + mock.post(sgx.csr_server_url(sgx_url) + '/', json=self._csr_port) + + @property + def last_issued(self) -> x509.Certificate: + return self.issued[-1] + + def issue_files(self, directory: Path) -> None: + """Write a valid key, request and certificate the way node services would have.""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + subject = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, 'ab' * 32)]) + csr = ( + x509.CertificateSigningRequestBuilder().subject_name(subject).sign(key, hashes.SHA256()) + ) + cert = self._sign(csr) + directory.mkdir(parents=True, exist_ok=True) + key_pem = key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + (directory / 'sgx.key').write_bytes(key_pem) + (directory / 'sgx.csr').write_bytes(_pem(csr)) + (directory / 'sgx.crt').write_bytes(_pem(cert)) + + def _sign(self, csr: x509.CertificateSigningRequest, public_key=None) -> x509.Certificate: + now = datetime.datetime.now(datetime.timezone.utc) + cert = ( + x509.CertificateBuilder() + .subject_name(csr.subject) + .issuer_name(self.ca_name) + .public_key(public_key or csr.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - ONE_DAY) + .not_valid_after(now + 365 * ONE_DAY) + .sign(self.ca_key, hashes.SHA256()) + ) + self.issued.append(cert) + return cert + + def _csr_port(self, request, context): + body = request.json() + method = body['method'] + self.calls.append(('csr', method)) + if method == 'signCertificate': + if self.sign_error: + return _response({'status': 1, 'errorMessage': self.sign_error}) + csr_pem = body['params']['certificate'] + csr = x509.load_pem_x509_csr(csr_pem.encode()) + assert csr.is_signature_valid + digest = hashlib.sha256(csr_pem.encode()).hexdigest() + self.pending[digest] = csr + return _response({'status': 0, 'hash': digest}) + if method == 'getCertificate': + csr = self.pending[body['params']['hash']] + if self.pending_polls > 0: + self.pending_polls -= 1 + return _response( + {'status': 1, 'cert': '', 'errorMessage': 'Certificate is not signed yet'} + ) + public_key = None + if self.wrong_key: + public_key = rsa.generate_private_key(65537, 2048).public_key() + cert = self._sign(csr, public_key) + return _response({'status': 0, 'cert': _pem(cert).decode()}) + raise AssertionError(f'Unexpected method on the CSR port: {method}') + + def _main_port(self, request, context): + body = request.json() + self.calls.append(('sgx', body['method'])) + assert request.verify is False + self.client_certs.append(request.cert) + crt_path, key_path = request.cert + cert = x509.load_pem_x509_certificate(Path(crt_path).read_bytes()) + key = serialization.load_pem_private_key(Path(key_path).read_bytes(), password=None) + known = {_fingerprint(issued) for issued in self.issued} + if ( + self.reject_clients + or _fingerprint(cert) not in known + or key.public_key().public_numbers() != cert.public_key().public_numbers() + ): + raise requests.exceptions.SSLError('tlsv1 alert unknown ca') + if body['method'] == 'getServerVersion': + return _response({'status': 0, 'version': self.version}) + raise AssertionError(f'Unexpected method on the main port: {body["method"]}') diff --git a/tests/routes_test.py b/tests/routes_test.py index 9a346748..51b9c5d8 100644 --- a/tests/routes_test.py +++ b/tests/routes_test.py @@ -21,6 +21,7 @@ '/api/v1/health/containers', '/api/v1/health/schains', '/api/v1/info/sgx', + '/api/v1/info/sgx-options', '/api/v1/schains/config', '/api/v1/schains/list', '/api/v1/schains/dkg-statuses', diff --git a/tests/utils/api_auth_test.py b/tests/utils/api_auth_test.py new file mode 100644 index 00000000..0549d5b6 --- /dev/null +++ b/tests/utils/api_auth_test.py @@ -0,0 +1,163 @@ +import os +import pwd +import stat +from concurrent.futures import ThreadPoolExecutor +from unittest.mock import Mock + +import pytest +import requests_mock + +from node_cli.utils import api_auth, helper + + +@pytest.fixture +def token_path(tmp_path, monkeypatch): + path = tmp_path / '.skale' / 'auth' / 'admin-api.token' + monkeypatch.setattr(api_auth, 'ADMIN_API_TOKEN_PATH', path) + monkeypatch.setattr(api_auth, 'G_CONF_USER', pwd.getpwuid(os.geteuid()).pw_name) + return path + + +def test_provisioning_preserves_token_and_restricts_permissions(token_path): + api_auth.ensure_api_token() + original = token_path.read_bytes() + assert len(api_auth.read_api_token()) == 64 + assert stat.S_IMODE(token_path.stat().st_mode) == 0o600 + assert token_path.stat().st_uid == os.geteuid() + assert stat.S_IMODE(token_path.parent.stat().st_mode) == 0o700 + assert token_path.parent.stat().st_uid == os.geteuid() + api_auth.ensure_api_token() + assert token_path.read_bytes() == original + assert list(token_path.parent.iterdir()) == [token_path] + + +def test_existing_auth_directory_is_secured(token_path): + api_auth.ensure_api_token() + original = token_path.read_bytes() + token_path.parent.chmod(0o755) + api_auth.ensure_api_token() + assert stat.S_IMODE(token_path.parent.stat().st_mode) == 0o700 + assert token_path.read_bytes() == original + + +def test_auth_directory_cannot_point_back_to_node_data(token_path): + node_data = token_path.parent.parent / 'node_data' + node_data.mkdir(parents=True) + token_path.parent.symlink_to(node_data, target_is_directory=True) + with pytest.raises(api_auth.APIAuthError): + api_auth.ensure_api_token() + assert list(node_data.iterdir()) == [] + + +def test_concurrent_provisioning_publishes_one_complete_token(token_path): + with ThreadPoolExecutor(max_workers=8) as pool: + list(pool.map(lambda _: api_auth.ensure_api_token(), range(16))) + assert len(api_auth.read_api_token()) == 64 + assert list(token_path.parent.iterdir()) == [token_path] + + +@pytest.mark.parametrize('bad_file', ['empty', 'oversized', 'non_ascii', 'permissions', 'symlink']) +def test_invalid_credentials_are_not_replaced(token_path, bad_file): + api_auth.ensure_api_token() + if bad_file == 'permissions': + token_path.chmod(0o644) + elif bad_file == 'symlink': + token_path.unlink() + token_path.symlink_to(token_path.with_suffix('.missing')) + else: + token_path.write_text( + {'empty': '', 'oversized': 'ab' * 32 + '\n\nmore', 'non_ascii': 'é'}[bad_file] + ) + with pytest.raises(api_auth.APIAuthError): + api_auth.ensure_api_token() + + +def test_missing_token_supports_upgrade_from_old_api(token_path): + assert api_auth.get_api_headers() == {} + with requests_mock.Mocker() as mock: + mock.get( + helper.construct_url('/api/v1/node/update-safe'), json={'status': 'ok', 'payload': {}} + ) + assert helper.get_request('node', 'update-safe') == ('ok', {}) + assert 'Authorization' not in mock.last_request.headers + assert not token_path.exists() + + +def test_headers_sent_for_get_post_and_upload(token_path): + api_auth.ensure_api_token() + expected = f'Bearer {api_auth.read_api_token()}' + with requests_mock.Mocker() as mock: + mock.get( + helper.construct_url('/api/v1/node/signature'), json={'status': 'ok', 'payload': {}} + ) + mock.post( + helper.construct_url('/api/v1/wallet/send-eth'), json={'status': 'ok', 'payload': {}} + ) + mock.post(helper.construct_url('/api/v1/ssl/upload'), json={'status': 'ok', 'payload': {}}) + assert helper.get_request('node', 'signature') == ('ok', {}) + assert helper.post_request('wallet', 'send-eth', json={'amount': 1}) == ('ok', {}) + assert helper.post_request('ssl', 'upload', files={'ssl_cert': ('cert', b'cert')}) == ( + 'ok', + {}, + ) + assert all(req.headers['Authorization'] == expected for req in mock.request_history) + assert 'multipart/form-data' in mock.last_request.headers['Content-Type'] + + +def test_http_errors_reach_cli(token_path): + api_auth.ensure_api_token() + with requests_mock.Mocker() as mock: + mock.post( + helper.construct_url('/api/v1/wallet/send-eth'), + status_code=401, + json={'status': 'error', 'payload': 'A valid node CLI credential is required'}, + ) + assert helper.post_request('wallet', 'send-eth') == ( + 'error', + 'A valid node CLI credential is required', + ) + + +def test_does_not_follow_redirects_or_use_proxy_credentials(token_path, monkeypatch): + api_auth.ensure_api_token() + monkeypatch.setenv('HTTP_PROXY', 'http://proxy.invalid:8080') + netrc = Mock(side_effect=AssertionError('Must not read netrc')) + monkeypatch.setattr('requests.sessions.get_netrc_auth', netrc) + with requests_mock.Mocker() as mock: + mock.get( + helper.construct_url('/api/v1/node/signature'), + status_code=302, + headers={'Location': 'http://other.invalid/steal'}, + json={'status': 'error', 'payload': 'Redirect'}, + ) + assert helper.get_request('node', 'signature') == ('error', 'Redirect') + assert len(mock.request_history) == 1 + assert mock.last_request.proxies == {} + netrc.assert_not_called() + + +def test_provisions_before_compose_starts(token_path, monkeypatch): + from node_cli.utils import docker_utils + from node_cli.utils.node_type import NodeMode, NodeType + + def run(cmd, env): + assert api_auth.read_api_token() is not None + + start = Mock(side_effect=run) + monkeypatch.setattr(docker_utils, 'run_cmd', start) + settings = Mock(tg_api_key=None) + docker_utils.compose_up({}, settings, NodeType.SKALE, NodeMode.ACTIVE) + start.assert_called_once() + + +def test_invalid_token_prevents_starting_services(token_path, monkeypatch): + from node_cli.utils import docker_utils + from node_cli.utils.node_type import NodeMode, NodeType + + api_auth.ensure_api_token() + token_path.chmod(0o644) + start = Mock() + monkeypatch.setattr(docker_utils, 'run_cmd', start) + with pytest.raises(api_auth.APIAuthError): + docker_utils.compose_up({}, Mock(), NodeType.SKALE, NodeMode.ACTIVE) + start.assert_not_called() diff --git a/text.yml b/text.yml index 6013f2d0..652adf7e 100644 --- a/text.yml +++ b/text.yml @@ -102,3 +102,27 @@ fair: exit: help: Remove node from Fair manager prompt: Are you sure you want to remove the node from Fair manager? + +sgx: + help: SGX wallet options and client certificate commands + no_sgx: This node has no SGX server configured (passive nodes do not use SGX) + options: + help: Show SGX server options through the node API + status: + help: Show the SGX client certificate used by node services + check: Also verify that the SGX server accepts the certificate + missing: |- + SGX client certificate files are missing. + Node services create them on their first SGX request, or run < skale sgx renew > + expired: SGX client certificate has expired. Run < skale sgx renew > + expires_soon: "SGX client certificate expires in {days} days. Consider running < skale sgx renew >" + key_mismatch: Private key does not match the certificate. Run < skale sgx renew > + renew: + help: Issue a new SGX client certificate and install it for node services + prompt: Are you sure you want to replace the SGX client certificate of this node? + timeout: Seconds to wait for the SGX server to sign the request + skip_verify: Install the certificate without testing it against the SGX server first + backup: "Previous certificate files were copied to {path}" + done: |- + New SGX client certificate installed. + Node services use it on their next SGX request; no restart is needed.