diff --git a/src/xcp_storage/backends/linstor/__init__.py b/src/xcp_storage/backends/linstor/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/xcp_storage/backends/linstor/controller.py b/src/xcp_storage/backends/linstor/controller.py new file mode 100644 index 0000000..ec6fa56 --- /dev/null +++ b/src/xcp_storage/backends/linstor/controller.py @@ -0,0 +1,72 @@ +# Copyright (C) 2026 Vates SAS +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU 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 General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from xcp_storage.backends.linstor.satellite import LINSTOR_SATELLITE_PORT_PLAIN, LINSTOR_SATELLITE_PORT_SSL +from xcp_storage.config.platform import get_exec_path +from xcp_storage.utils.process import run_command +from xcp_storage.utils.service import ( + is_service_active, + restart_service, + start_service, + stop_service, +) + +from xcp_storage.typing import Final, List + +# ============================================================================== + +LINSTOR_CONTROLLER_PORT_PLAIN: Final = 3370 +LINSTOR_CONTROLLER_PORT_SSL: Final = 3371 + +# ------------------------------------------------------------------------------ + +_EXEC_PATH_SS: Final = get_exec_path("/usr/sbin/ss", {"debian": "/usr/bin/ss"}) + +_SERVICE_LINSTOR_CONTROLLER: Final = "linstor-controller" + +# ------------------------------------------------------------------------------ + +class LinstorController: + @staticmethod + def get_addresses() -> List[str]: + stdout = run_command([ + _EXEC_PATH_SS, "-tnpH", "state", "established", + f"( sport = :{LINSTOR_SATELLITE_PORT_PLAIN} or sport = :{LINSTOR_SATELLITE_PORT_SSL} )" + ], expected_ret_code=0) + return [ + line.split()[3].rsplit(":", 1)[0] + for line in stdout.splitlines() + ] + + @classmethod + def get_uri(cls) -> str: + # TODO(XCPNG-3033): On caller side, check that an IP address from the current pool is returned. + addresses = cls.get_addresses() + return "linstor://" + addresses[0] if addresses else "" + + @staticmethod + def is_running() -> bool: + return is_service_active(_SERVICE_LINSTOR_CONTROLLER) + + @staticmethod + def start() -> None: + start_service(_SERVICE_LINSTOR_CONTROLLER) + + @staticmethod + def stop() -> None: + stop_service(_SERVICE_LINSTOR_CONTROLLER) + + @staticmethod + def restart() -> None: + restart_service(_SERVICE_LINSTOR_CONTROLLER) diff --git a/src/xcp_storage/backends/linstor/satellite.py b/src/xcp_storage/backends/linstor/satellite.py new file mode 100644 index 0000000..8fcb62d --- /dev/null +++ b/src/xcp_storage/backends/linstor/satellite.py @@ -0,0 +1,45 @@ +# Copyright (C) 2026 Vates SAS +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU 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 General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from xcp_storage.utils.service import ( + disable_and_stop_service, + enable_and_start_service, + is_service_active, +) + +from xcp_storage.typing import Final + +# ============================================================================== + +LINSTOR_SATELLITE_PORT_PLAIN: Final = 3366 +LINSTOR_SATELLITE_PORT_SSL: Final = 3367 + +# ------------------------------------------------------------------------------ + +_SERVICE_LINSTOR_SATELLITE: Final = "linstor-satellite" + +# ------------------------------------------------------------------------------ + +class LinstorSatellite: + @staticmethod + def is_running() -> bool: + return is_service_active(_SERVICE_LINSTOR_SATELLITE) + + @staticmethod + def enable_and_start() -> None: + enable_and_start_service(_SERVICE_LINSTOR_SATELLITE) + + @staticmethod + def disable_and_stop() -> None: + disable_and_stop_service(_SERVICE_LINSTOR_SATELLITE) diff --git a/src/xcp_storage/config/platform.py b/src/xcp_storage/config/platform.py index 984598a..4bb9c03 100644 --- a/src/xcp_storage/config/platform.py +++ b/src/xcp_storage/config/platform.py @@ -12,7 +12,10 @@ # You should have received a copy of the GNU General Public License # along with this program. If not, see . -from xcp_storage.typing import Final +from functools import lru_cache +from pathlib import Path + +from xcp_storage.typing import Final, Mapping, Tuple # ============================================================================== # Attributes that depend on the execution environment. @@ -21,3 +24,35 @@ # ============================================================================== DEFAULT_FIREWALL_INPUT_CHAIN: Final = "xapi-INPUT" + +# ------------------------------------------------------------------------------ + +_OS_RELEASE_PATH: Final = "/etc/os-release" + +@lru_cache(maxsize=None) +def get_os_ids() -> Tuple[str, ...]: + """ + Get the IDs of the current distribution, most specific first: `ID` followed + by the values of `ID_LIKE` (e.g. `("ubuntu", "debian")`). The result is + cached. An empty tuple is returned if `/etc/os-release` can't be read. + """ + + try: + lines = Path(_OS_RELEASE_PATH).read_text(encoding="utf-8").splitlines() + except OSError: + return () + + values = {} + for line in lines: + key, _separator, value = line.partition("=") + values[key.strip()] = value.strip().strip("\"'") + + return tuple(values.get("ID", "").split() + values.get("ID_LIKE", "").split()) + +def get_exec_path(default: str, by_os_id: Mapping[str, str]) -> str: + """ + Get the path of an executable: the one registered for the first matching OS + ID, otherwise `default`. + """ + + return next((by_os_id[os_id] for os_id in get_os_ids() if os_id in by_os_id), default) diff --git a/src/xcp_storage/network/socket.py b/src/xcp_storage/network/socket.py index c74f6bd..0b1e098 100644 --- a/src/xcp_storage/network/socket.py +++ b/src/xcp_storage/network/socket.py @@ -198,13 +198,18 @@ def create_client_sock( reuse_address: bool = True, keep_alive: bool = True, timeout: Optional[float] = None, - ssl_context: Optional[ssl.SSLContext] = None + ssl_context: Optional[ssl.SSLContext] = None, + source_address: Optional[str] = None, + source_port: int = 0 ) -> socket.socket: family, connect = format_address(address, port) sock = _create_stream_sock(address, family, bind=False, reuse_address=reuse_address, ssl_context=ssl_context) _normalize_and_set_sock_timeout(sock, timeout) try: + if source_address: + _, bind = format_address(source_address, source_port) + sock.bind(bind) sock.connect(connect) except OSError as e: with contextlib.suppress(Exception): @@ -318,6 +323,12 @@ def socket_wait_readable(sock: socket.socket, *, timeout: Optional[float] = None def get_socket_family_str(sock: socket.socket) -> str: return _FAMILY_TO_STR.get(sock.family, "Unknown") +def get_socket_address(sock: socket.socket) -> Optional[str]: + try: + return sock.getsockname()[0] + except OSError: + return None + def get_socket_port(sock: socket.socket) -> Optional[int]: try: return sock.getsockname()[1] @@ -371,6 +382,10 @@ def close(self) -> None: def family_str(self) -> str: return get_socket_family_str(self.sock) + @property + def address(self) -> Optional[str]: + return get_socket_address(self.sock) + @property def port(self) -> Optional[int]: return get_socket_port(self.sock) diff --git a/tests/backends/test_linstor_controller.py b/tests/backends/test_linstor_controller.py new file mode 100644 index 0000000..ad04ca7 --- /dev/null +++ b/tests/backends/test_linstor_controller.py @@ -0,0 +1,189 @@ +# Copyright (C) 2026 Vates SAS +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU 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 General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +import contextlib +from unittest.mock import ( + MagicMock, + patch, +) + +import pytest + +from xcp_storage.backends.linstor.controller import _SERVICE_LINSTOR_CONTROLLER, LinstorController +from xcp_storage.network.socket import ( + create_client_sock, + create_server_sock, + Socket, + SocketError, +) + +from xcp_storage.typing import ( + Callable, + cast, + Final, + Iterator, + NamedTuple, + Optional, + Tuple, +) + +# ============================================================================== + +MODULE_NAME: Final = "xcp_storage.backends.linstor.controller" + +# ------------------------------------------------------------------------------ + +class TestLinstorControllerGetAddresses: + class LoopbackSpec(NamedTuple): + listener_address: str + client_address: str + uri_client_address: str + + LOOPBACK_ADDRESSES: Final = ["127.0.0.2", "127.0.0.3", "127.0.0.4"] + + IPV4: Final = LoopbackSpec("127.0.0.1", LOOPBACK_ADDRESSES[0], LOOPBACK_ADDRESSES[0]) + IPV6: Final = LoopbackSpec("::1", "::1", "[::1]") + + ListenFunc = Callable[..., Socket] + ConnectFunc = Callable[..., Tuple[Socket, Socket]] + PatchListenerPortFunc = Callable[[Socket], "contextlib.AbstractContextManager[None]"] + + @staticmethod + @contextlib.contextmanager + def patch_satellite_ports(plain_port: Optional[int], ssl_port: Optional[int]) -> Iterator[None]: + with patch.multiple(MODULE_NAME, LINSTOR_SATELLITE_PORT_PLAIN=plain_port, LINSTOR_SATELLITE_PORT_SSL=ssl_port): + yield + + @pytest.fixture(params=[IPV4, IPV6], ids=["ipv4", "ipv6"]) + def loopback_spec(self, request: pytest.FixtureRequest) -> LoopbackSpec: + spec = cast(TestLinstorControllerGetAddresses.LoopbackSpec, request.param) + if spec is self.IPV6: + try: + create_server_sock(spec.listener_address, 0).close() + except SocketError: + pytest.skip("IPv6 is not available.") + return spec + + @pytest.fixture + def sockets(self) -> Iterator[contextlib.ExitStack]: + with contextlib.ExitStack() as stack: + yield stack + + @pytest.fixture + def listen(self, sockets: contextlib.ExitStack) -> ListenFunc: + def impl(address: str = "127.0.0.1") -> Socket: + return sockets.enter_context(Socket(create_server_sock(address, 0))) + return impl + + @pytest.fixture + def connect(self, sockets: contextlib.ExitStack) -> ConnectFunc: + def impl(listener: Socket, client_address: str) -> Tuple[Socket, Socket]: + client = sockets.enter_context(Socket(create_client_sock( + cast(str, listener.address), + cast(int, listener.port), + source_address=client_address, + keep_alive=False + ))) + accepted, _ = listener.sock.accept() + return client, sockets.enter_context(Socket(accepted)) + return impl + + @pytest.fixture(params=["plain", "ssl"]) + def patch_listener_port(self, request: pytest.FixtureRequest, listen: ListenFunc) -> PatchListenerPortFunc: + unused_port = listen().port + + def impl(listener: Socket) -> "contextlib.AbstractContextManager[None]": + ports = (listener.port, unused_port) + return self.patch_satellite_ports(*(ports if request.param == "plain" else reversed(ports))) + return impl + + def test_no_connection(self, listen: ListenFunc, loopback_spec: LoopbackSpec) -> None: + listener = listen(loopback_spec.listener_address) + + with self.patch_satellite_ports(listener.port, listener.port): + assert LinstorController.get_addresses() == [] + assert LinstorController.get_uri() == "" + + def test_closed_connection( + self, + listen: ListenFunc, + connect: ConnectFunc, + patch_listener_port: PatchListenerPortFunc, + loopback_spec: LoopbackSpec + ) -> None: + listener = listen(loopback_spec.listener_address) + client, accepted = connect(listener, loopback_spec.client_address) + + accepted.close() + client.close() + + with patch_listener_port(listener): + assert LinstorController.get_addresses() == [] + + def test_established_connection( + self, + listen: ListenFunc, + connect: ConnectFunc, + patch_listener_port: PatchListenerPortFunc, + loopback_spec: LoopbackSpec + ) -> None: + listener = listen(loopback_spec.listener_address) + connect(listener, loopback_spec.client_address) + + with patch_listener_port(listener): + assert LinstorController.get_addresses() == [loopback_spec.uri_client_address] + assert LinstorController.get_uri() == f"linstor://{loopback_spec.uri_client_address}" + +# ------------------------------------------------------------------------------ + +@patch(f"{MODULE_NAME}.LinstorController.get_addresses") +class TestLinstorControllerGetUri: + def test_no_address(self, mock_get_addresses: MagicMock) -> None: + mock_get_addresses.return_value = [] + assert LinstorController.get_uri() == "" + + def test_one_address(self, mock_get_addresses: MagicMock) -> None: + mock_get_addresses.return_value = ["10.10.0.13"] + assert LinstorController.get_uri() == "linstor://10.10.0.13" + + def test_two_addresses(self, mock_get_addresses: MagicMock) -> None: + mock_get_addresses.return_value = ["10.10.0.13", "10.10.0.14"] + assert LinstorController.get_uri() == "linstor://10.10.0.13" + + def test_ipv6_address(self, mock_get_addresses: MagicMock) -> None: + mock_get_addresses.return_value = ["[fe80::2]"] + assert LinstorController.get_uri() == "linstor://[fe80::2]" + +# ------------------------------------------------------------------------------ + +class TestLinstorControllerService: + @patch(f"{MODULE_NAME}.is_service_active", return_value=True) + def test_running(self, mock_is_service_active: MagicMock) -> None: + assert LinstorController.is_running() + mock_is_service_active.assert_called_once_with(_SERVICE_LINSTOR_CONTROLLER) + + @patch(f"{MODULE_NAME}.is_service_active", return_value=False) + def test_not_running(self, mock_is_service_active: MagicMock) -> None: + assert not LinstorController.is_running() + mock_is_service_active.assert_called_once_with(_SERVICE_LINSTOR_CONTROLLER) + + @pytest.mark.parametrize("service_name, function_name", [ + ("start", "start_service"), + ("stop", "stop_service"), + ("restart", "restart_service") + ]) + def test_service_command(self, service_name: str, function_name: str) -> None: + with patch(f"{MODULE_NAME}.{function_name}") as mock_function_name: + getattr(LinstorController, service_name)() + mock_function_name.assert_called_once_with(_SERVICE_LINSTOR_CONTROLLER) diff --git a/tests/backends/test_linstor_satellite.py b/tests/backends/test_linstor_satellite.py new file mode 100644 index 0000000..f0e4d80 --- /dev/null +++ b/tests/backends/test_linstor_satellite.py @@ -0,0 +1,50 @@ +# Copyright (C) 2026 Vates SAS +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU 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 General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from unittest.mock import ( + MagicMock, + patch, +) + +import pytest + +from xcp_storage.backends.linstor.satellite import _SERVICE_LINSTOR_SATELLITE, LinstorSatellite + +from xcp_storage.typing import Final + +# ============================================================================== + +MODULE_NAME: Final = "xcp_storage.backends.linstor.satellite" + +# ------------------------------------------------------------------------------ + +class TestLinstorSatelliteService: + @patch(f"{MODULE_NAME}.is_service_active", return_value=True) + def test_running(self, mock_is_service_active: MagicMock) -> None: + assert LinstorSatellite.is_running() + mock_is_service_active.assert_called_once_with(_SERVICE_LINSTOR_SATELLITE) + + @patch(f"{MODULE_NAME}.is_service_active", return_value=False) + def test_not_running(self, mock_is_service_active: MagicMock) -> None: + assert not LinstorSatellite.is_running() + mock_is_service_active.assert_called_once_with(_SERVICE_LINSTOR_SATELLITE) + + @pytest.mark.parametrize("service_name, function_name", [ + ("enable_and_start", "enable_and_start_service"), + ("disable_and_stop", "disable_and_stop_service") + ]) + def test_service_command(self, service_name: str, function_name: str) -> None: + with patch(f"{MODULE_NAME}.{function_name}") as mock_function_name: + getattr(LinstorSatellite, service_name)() + mock_function_name.assert_called_once_with(_SERVICE_LINSTOR_SATELLITE) diff --git a/tests/config/test_platform.py b/tests/config/test_platform.py new file mode 100644 index 0000000..631dcdc --- /dev/null +++ b/tests/config/test_platform.py @@ -0,0 +1,82 @@ +# Copyright (C) 2026 Vates SAS +# +# This program is free software: you can redistribute it and/or modify +# it under the terms of the GNU 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 General Public License for more details. +# +# You should have received a copy of the GNU General Public License +# along with this program. If not, see . + +from pathlib import Path +from unittest.mock import patch + +import pytest + +from xcp_storage.config.platform import get_exec_path, get_os_ids + +from xcp_storage.typing import Final, Iterator, Tuple + +# ============================================================================== + +@pytest.fixture(autouse=True) +def clear_os_ids_cache() -> Iterator[None]: + get_os_ids.cache_clear() + yield + get_os_ids.cache_clear() + +@pytest.fixture +def os_release_path(tmp_path: Path) -> Iterator[Path]: + path = tmp_path / "os-release" + with patch("xcp_storage.config.platform._OS_RELEASE_PATH", str(path)): + yield path + +# ------------------------------------------------------------------------------ + +class TestGetOsIds: + @pytest.mark.parametrize(("content", "expected"), [ + ("ID=alpine\n", ("alpine",)), + ('NAME="Ubuntu"\nID=ubuntu\nID_LIKE=debian\n', ("ubuntu", "debian")), + ('ID="rhel"\nID_LIKE="fedora centos"\n', ("rhel", "fedora", "centos")), + ("NAME=Foo\n", ()), + ]) + def test_parsing(self, os_release_path: Path, content: str, expected: Tuple[str, ...]) -> None: + os_release_path.write_text(content) + assert get_os_ids() == expected + + def test_without_os_release_file(self, os_release_path: Path) -> None: + assert not os_release_path.exists() + assert get_os_ids() == () + + def test_os_release_file_is_cached(self, os_release_path: Path) -> None: + os_release_path.write_text("ID=alpine\n") + expected = get_os_ids() + os_release_path.write_text("ID=ubuntu\n") + assert get_os_ids() == expected + +# ------------------------------------------------------------------------------ + +class TestGetExecPath: + DEFAULT_PATH: Final = "/opt/default/tool" + ALT_PATH: Final = "/opt/alt/tool" + BY_OS_ID: Final = {"debian": ALT_PATH} + + def test_with_mapping_id_match(self, os_release_path: Path) -> None: + os_release_path.write_text("ID=debian\nID_LIKE=unknown\n") + assert get_exec_path(self.DEFAULT_PATH, self.BY_OS_ID) == self.ALT_PATH + + def test_with_mapping_id_like_match(self, os_release_path: Path) -> None: + os_release_path.write_text("ID=ubuntu\nID_LIKE=debian\n") + assert get_exec_path(self.DEFAULT_PATH, self.BY_OS_ID) == self.ALT_PATH + + def test_without_mapping_match(self, os_release_path: Path) -> None: + os_release_path.write_text("ID=alpine\n") + assert get_exec_path(self.DEFAULT_PATH, self.BY_OS_ID) == self.DEFAULT_PATH + + def test_without_os_release_file(self, os_release_path: Path) -> None: + assert not os_release_path.exists() + assert get_exec_path(self.DEFAULT_PATH, self.BY_OS_ID) == self.DEFAULT_PATH diff --git a/tests/network/test_socket.py b/tests/network/test_socket.py index cbfe618..c2f6839 100644 --- a/tests/network/test_socket.py +++ b/tests/network/test_socket.py @@ -30,6 +30,7 @@ create_server_sock, format_address, get_ip_address, + get_socket_address, get_socket_family_str, get_socket_port, Socket, @@ -211,6 +212,17 @@ def test_get_socket_family_str( # ------------------------------------------------------------------------------ +class TestSocketAddress: + def test_get_socket_address(self) -> None: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + assert get_socket_address(sock) == "127.0.0.1" + + def test_get_socket_address_closed_socket(self) -> None: + with socket.socket() as sock: + sock.close() + assert get_socket_address(sock) is None + class TestSocketPort: def test_get_socket_port(self) -> None: with socket.socket() as sock: