From d05b18b616276780637bb2fbeffa602006546257 Mon Sep 17 00:00:00 2001 From: Ronan Abhamon Date: Fri, 2 Oct 2026 22:32:32 +0200 Subject: [PATCH 1/4] feat(tests/core): add `TestSocketPort` test (#6) Signed-off-by: Ronan Abhamon --- tests/network/test_socket.py | 15 +++++++++++++++ 1 file changed, 15 insertions(+) diff --git a/tests/network/test_socket.py b/tests/network/test_socket.py index 6312e07..0cb9942 100644 --- a/tests/network/test_socket.py +++ b/tests/network/test_socket.py @@ -30,6 +30,7 @@ format_address, get_ip_address, get_socket_family_str, + get_socket_port, Socket, socket_receive, socket_send, @@ -204,6 +205,20 @@ def test_get_socket_family_str( # ------------------------------------------------------------------------------ +class TestSocketPort: + def test_get_socket_port(self) -> None: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + port = get_socket_port(sock) + assert port is not None and port > 0 + + def test_get_socket_port_closed_socket(self) -> None: + with socket.socket() as sock: + sock.close() + assert get_socket_port(sock) is None + +# ------------------------------------------------------------------------------ + @pytest.fixture def mock_sock() -> MagicMock: return MagicMock(spec=socket.socket) From 0e7a3e0b92d3a110de83acd7d54f8c6aed0fd202 Mon Sep 17 00:00:00 2001 From: Ronan Abhamon Date: Mon, 22 Jun 2026 16:02:24 +0200 Subject: [PATCH 2/4] feat(core): add TCP client/server modules (#6) Signed-off-by: Ronan Abhamon --- pyproject.toml | 2 +- src/xcp_storage/network/socket.py | 10 + src/xcp_storage/network/tcp_client.py | 145 ++++++++ src/xcp_storage/network/tcp_server.py | 245 +++++++++++++ src/xcp_storage/utils/asyncio.py | 37 ++ src/xcp_storage/utils/sync.py | 38 ++ tests/__init__.py | 0 tests/network/__init__.py | 26 ++ tests/network/conftest.py | 99 ++++++ tests/network/test_tcp_client_server.py | 446 ++++++++++++++++++++++++ tests/utils/test_sync.py | 71 ++++ 11 files changed, 1118 insertions(+), 1 deletion(-) create mode 100644 src/xcp_storage/network/tcp_client.py create mode 100644 src/xcp_storage/network/tcp_server.py create mode 100644 src/xcp_storage/utils/asyncio.py create mode 100644 src/xcp_storage/utils/sync.py create mode 100644 tests/__init__.py create mode 100644 tests/network/__init__.py create mode 100644 tests/network/conftest.py create mode 100644 tests/network/test_tcp_client_server.py create mode 100644 tests/utils/test_sync.py diff --git a/pyproject.toml b/pyproject.toml index a980279..94c3dc4 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -57,7 +57,7 @@ disallow_untyped_defs = false [tool.pyrefly] python-version = "3.14" -search-path = ["stubs"] +search-path = [".", "stubs"] project-includes = [ "src", "tests" diff --git a/src/xcp_storage/network/socket.py b/src/xcp_storage/network/socket.py index 06a226e..5c8a1c2 100644 --- a/src/xcp_storage/network/socket.py +++ b/src/xcp_storage/network/socket.py @@ -302,6 +302,12 @@ def socket_receive(sock: socket.socket, buffer: bytearray, size: Optional[int] = def get_socket_family_str(sock: socket.socket) -> str: return _FAMILY_TO_STR.get(sock.family, "Unknown") +def get_socket_port(sock: socket.socket) -> Optional[int]: + try: + return sock.getsockname()[1] + except OSError: + return None + # ------------------------------------------------------------------------------ class Socket(contextlib.AbstractContextManager): @@ -346,6 +352,10 @@ def close(self) -> None: def family_str(self) -> str: return get_socket_family_str(self.sock) + @property + def port(self) -> Optional[int]: + return get_socket_port(self.sock) + @property def timeout(self) -> Optional[float]: return self.sock.gettimeout() diff --git a/src/xcp_storage/network/tcp_client.py b/src/xcp_storage/network/tcp_client.py new file mode 100644 index 0000000..519a108 --- /dev/null +++ b/src/xcp_storage/network/tcp_client.py @@ -0,0 +1,145 @@ +# 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 +import ssl +from types import TracebackType + +from xcp_storage.network.socket import ( + create_client_sock, + Socket, + SocketDisconnectedError, +) +from xcp_storage.utils.sync import wait_for_condition + +from xcp_storage.typing import ( + Optional, + override, + Type, +) + +# ============================================================================== + +class TcpClientError(Exception): + def __init__(self, message: str) -> None: + super().__init__(message) + +# ------------------------------------------------------------------------------ + +class TcpClient(contextlib.AbstractContextManager): + def __init__( + self, + address: str, + port: int, + *, + ssl_context: Optional[ssl.SSLContext] = None, + client_timeout: float = 120 + ) -> None: + self._address = address + self._port = port + self._ssl_context = ssl_context + self._client_timeout = client_timeout + self._socket: Optional[Socket] = None + + self._entered_count = 0 + + def __del__(self) -> None: + # This test is required for a specific case: + # rare but possible if `__new__` was executed without `__init__`. + if getattr(self, "_socket", None): + self.disconnect() + + @override + def __enter__(self) -> "TcpClient": + self.connect() + self._entered_count += 1 + return self + + @override + def __exit__( + self, + exc_type: Optional[Type[BaseException]], + exc_value: Optional[BaseException], + traceback: Optional[TracebackType] + ) -> None: + self._entered_count -= 1 + if self._entered_count == 0: + self.disconnect() + + @property + def socket(self) -> Optional[Socket]: + return self._socket + + @property + def connected(self) -> bool: + return self._socket is not None + + def connect(self, timeout: Optional[float] = None) -> None: + if self._socket: + return + + if timeout is None: + # If the `connect` timeout is not set, we use the client one. + timeout = self._client_timeout + + connect_timeout = timeout if timeout is not None else self._client_timeout + + error: Optional[Exception] = None + def connect_impl() -> bool: + nonlocal error + try: + client_socket = Socket(create_client_sock( + self._address, + self._port, + reuse_address=True, + keep_alive=True, + timeout=connect_timeout, + ssl_context=self._ssl_context + )) + except Exception as e: + error = e + return False + + # Ensure that client timeout is used after this point, not the connection timeout. + client_socket.timeout = self._client_timeout + self._socket = client_socket + return True + + if not wait_for_condition(connect_impl, timeout=timeout, interval=1): + raise TcpClientError("Unable to connect to server.") from error + + def disconnect(self) -> None: + if self._socket: + self._socket.close() + self._socket = None + + def send(self, buffer: bytes, size: Optional[int] = None) -> None: + if not self._socket: + raise TcpClientError("Cannot send. Not connected.") + + try: + self._socket.send(buffer, size) + except SocketDisconnectedError: + self.disconnect() + raise + + def receive(self, buffer: bytearray, size: Optional[int] = None) -> None: + if not self._socket: + raise TcpClientError("Cannot receive. Not connected.") + + try: + self._socket.receive(buffer, size) + except SocketDisconnectedError: + self.disconnect() + raise diff --git a/src/xcp_storage/network/tcp_server.py b/src/xcp_storage/network/tcp_server.py new file mode 100644 index 0000000..043637b --- /dev/null +++ b/src/xcp_storage/network/tcp_server.py @@ -0,0 +1,245 @@ +# 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 abc import ABC, abstractmethod +import asyncio +import contextlib +import ssl +import sys +import threading + +import xcp_storage.log as log +from xcp_storage.network.socket import create_server_sock, Socket +from xcp_storage.utils.asyncio import cancel_event_loop_tasks, close_stream_writer + +from xcp_storage.typing import ( + Any, + Optional, + override, + Set, +) + +# ============================================================================== + +logger = log.get_logger() # Use default logger. + +# ------------------------------------------------------------------------------ + +class TcpServerError(Exception): + def __init__(self, message: str) -> None: + super().__init__(message) + +# ------------------------------------------------------------------------------ + +class TcpServer(ABC): + class Client: + def __init__( + self, + peername: Any, # noqa: ANN401 + reader: asyncio.StreamReader, + writer: asyncio.StreamWriter + ) -> None: + self.peername: Any = peername + self.reader = reader + self.writer = writer + + @override + def __str__(self) -> str: + return str(self.peername) + + def __init__( + self, + address: str, + port: int, + *, + ssl_context: Optional[ssl.SSLContext] = None + ) -> None: + self._address = address + self._port = port + self._ssl_context = ssl_context + + self._running = False + + self._event_loop: Optional[asyncio.AbstractEventLoop] = None + self._startup_event = threading.Event() + self._started = False + self._shutdown_event = threading.Event() + + self._server_socket: Optional[Socket] = None + self._server: Optional[asyncio.AbstractServer] = None + + self._clients: Set[TcpServer.Client] = set() + + @property + def address(self) -> str: + return self._address + + @property + def port(self) -> int: + # Return trivial port if it's not 0. As reminder, 0 = dynamic binding. + if self._port or not self._server_socket: + return self._port + + port = self._server_socket.port + return port if port is not None else 0 + + def run(self) -> None: # noqa: C901 + if self._running: + raise TcpServerError("Server is already running.") + + self._startup_event.clear() + self._shutdown_event.clear() + self._started = False + self._running = True + + old_event_loop = None + try: + logger.info("Running TCP server on `%s:%d`...", self._address, self._port) + self._server_socket = Socket(create_server_sock( + self._address, + self._port, + reuse_address=True, + keep_alive=True, + timeout=None + # No `ssl_context` here: asyncio refuses an `SSLSocket` and handles TLS itself. + )) + + with contextlib.suppress(RuntimeError): + old_event_loop = asyncio.get_event_loop() + + self._event_loop = asyncio.new_event_loop() + asyncio.set_event_loop(self._event_loop) + + server_params = { + "client_connected_cb": self._handle_client, + "sock": self._server_socket.sock, + "ssl": self._ssl_context + } + + if sys.version_info < (3, 8): + # TODO(XCPNG-3032): Workaround for old python versions. Remove me later. + # In fact we must give `loop` param for these versions and the entire + # event loop management is also required just for it. + server_params["loop"] = self._event_loop + + self._server = self._event_loop.run_until_complete( + asyncio.start_server(**server_params) # type: ignore[arg-type] + ) + # Signal the startup from the loop itself, so `async_stop` is guaranteed + # to see a running loop once `wait_for_startup` returns. + self._event_loop.call_soon(self._notify_startup) + + logger.info("TCP server started!") + self._event_loop.run_forever() + except KeyboardInterrupt: + logger.info("Closing server because break signal has been received...") + finally: + # Always unblock `wait_for_startup`, even if the startup failed. + # Useful in case of exception or keyboard interrupt. + self._startup_event.set() + + if self._server: + self._server.close() + # TODO(XCPNG-3032): Use `abort_clients` only. + for client in list(self._clients): + try: + client.writer.transport.abort() + except Exception as e: # noqa: PERF203 + logger.debug("Failed to abort client %s transport: `%s`.", client, e) + self._server = None + + if self._event_loop: + try: + cancel_event_loop_tasks(self._event_loop) + finally: + self._event_loop.close() + self._event_loop = None + asyncio.set_event_loop(old_event_loop) + + if self._server_socket: + self._server_socket.close() + self._server_socket = None + + self._clients.clear() + self._running = False + self._shutdown_event.set() + + def stop(self, *, timeout: Optional[float] = None) -> bool: + if self.async_stop(): + self.wait_for_shutdown(timeout=timeout) + return True + return False + + def async_stop(self) -> bool: + if self._event_loop and self._event_loop.is_running(): + self._event_loop.call_soon_threadsafe(self._event_loop.stop) + return True + return False + + def wait_for_startup(self, *, timeout: Optional[float] = None) -> bool: + return self._startup_event.wait(timeout) and self._started + + def wait_for_shutdown(self, *, timeout: Optional[float] = None) -> None: + self._shutdown_event.wait(timeout) + + def _notify_startup(self) -> None: + self._started = True + self._startup_event.set() + + async def _handle_client( + self, + client_reader: asyncio.StreamReader, + client_writer: asyncio.StreamWriter + ) -> None: + client = self.Client(client_writer.get_extra_info("peername"), client_reader, client_writer) + logger.info("New client %s connected.", client) + self._clients.add(client) + + connected = False + try: + if not await self._handle_client_connect(client): + return + connected = True + while not client_writer.transport.is_closing(): + if not await self._handle_client_request(client): + break + logger.info("Client %s has terminated.", client) + except asyncio.TimeoutError as e: + logger.warning("Timeout reached for client %s: `%s`.", client, e) + except asyncio.IncompleteReadError as e: + logger.warning("Connection closed for client %s: `%s`.", client, e) + except Exception as e: + logger.error("Unhandled exception for client %s: `%s`.", client, e) + finally: + if connected: + try: + await self._handle_client_disconnect(client) + except Exception as e: + logger.error("Unhandled exception for client %s during disconnect: `%s`.", client, e) + + await close_stream_writer(client_writer) + logger.info("Client %s disconnected.", client) + self._clients.remove(client) + + @abstractmethod + async def _handle_client_connect(self, client: Client) -> bool: + return False + + @abstractmethod + async def _handle_client_disconnect(self, client: Client) -> None: + pass + + @abstractmethod + async def _handle_client_request(self, client: Client) -> bool: + return False diff --git a/src/xcp_storage/utils/asyncio.py b/src/xcp_storage/utils/asyncio.py new file mode 100644 index 0000000..11be5fe --- /dev/null +++ b/src/xcp_storage/utils/asyncio.py @@ -0,0 +1,37 @@ +# 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 asyncio +import contextlib + +# ============================================================================== + +def cancel_event_loop_tasks(event_loop: asyncio.AbstractEventLoop) -> None: + try: + tasks = asyncio.all_tasks(event_loop) + except AttributeError: + # TODO(XCPNG-3032): Workaround for python 3.6. Remove me later. + tasks = asyncio.Task.all_tasks(event_loop) # type: ignore + + for task in tasks: + task.cancel() + + event_loop.run_until_complete(asyncio.tasks.gather(*tasks, return_exceptions=True)) + event_loop.run_until_complete(event_loop.shutdown_asyncgens()) + +async def close_stream_writer(stream: asyncio.StreamWriter) -> None: + with contextlib.suppress(Exception): + if not stream.transport.is_closing(): + stream.close() + await stream.wait_closed() diff --git a/src/xcp_storage/utils/sync.py b/src/xcp_storage/utils/sync.py new file mode 100644 index 0000000..41ba08d --- /dev/null +++ b/src/xcp_storage/utils/sync.py @@ -0,0 +1,38 @@ +# 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 time + +from xcp_storage.typing import ( + Callable, + TypeVar, +) + +T = TypeVar("T") + +# ============================================================================== + +def wait_for_condition(function: Callable[[], T], timeout: float, interval: float) -> T: + if timeout <= 0: + return function() + + deadline = time.monotonic() + timeout + while True: + result = function() + if result: + return result + remaining = deadline - time.monotonic() + if remaining <= 0: + return result + time.sleep(min(interval, remaining)) diff --git a/tests/__init__.py b/tests/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/network/__init__.py b/tests/network/__init__.py new file mode 100644 index 0000000..f23ec7e --- /dev/null +++ b/tests/network/__init__.py @@ -0,0 +1,26 @@ +# 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 socket + +from xcp_storage.network.socket import get_socket_port + +from xcp_storage.typing import cast + +# ============================================================================== + +def find_free_tcp_port() -> int: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + return cast(int, get_socket_port(sock)) diff --git a/tests/network/conftest.py b/tests/network/conftest.py new file mode 100644 index 0000000..3fd3f36 --- /dev/null +++ b/tests/network/conftest.py @@ -0,0 +1,99 @@ +# 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 datetime +import ipaddress +import pathlib +import ssl + +import pytest + +from xcp_storage.typing import ( + Final, + NamedTuple, + Optional, +) + +# ============================================================================== + +TLS_SERVER_CERTIFICATE_NAME: Final = "server.crt" +TLS_SERVER_KEY_NAME: Final = "server.key" + +class TlsContexts(NamedTuple): + server: ssl.SSLContext + client: ssl.SSLContext + +@pytest.fixture +def tls_contexts(tmp_path: pathlib.Path) -> TlsContexts: + """ + Server/client TLS contexts built around a throwaway self-signed certificate, + valid for `127.0.0.1` only. + """ + + pytest.importorskip("cryptography") + from cryptography import x509 + from cryptography.hazmat.primitives import hashes, serialization + from cryptography.hazmat.primitives.asymmetric import ec + from cryptography.x509.oid import NameOID + + key = ec.generate_private_key(ec.SECP256R1()) + name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "localhost")]) + now = datetime.datetime.now(datetime.timezone.utc) + delta = datetime.timedelta(days=1) + cert = ( + x509.CertificateBuilder() + .subject_name(name) + .issuer_name(name) + .public_key(key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - delta) + .not_valid_after(now + delta) + .add_extension(x509.SubjectAlternativeName([ + x509.IPAddress(ipaddress.ip_address("127.0.0.1")) + ]), critical=False) + .sign(key, hashes.SHA256()) + ) + + cert_path = tmp_path / TLS_SERVER_CERTIFICATE_NAME + key_path = tmp_path / TLS_SERVER_KEY_NAME + cert_path.write_bytes(cert.public_bytes(serialization.Encoding.PEM)) + key_path.write_bytes(key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption() + )) + + server_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + server_context.load_cert_chain(str(cert_path), str(key_path)) + client_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT) + client_context.load_verify_locations(str(cert_path)) + return TlsContexts(server_context, client_context) + +@pytest.fixture +def ssl_contexts(request: pytest.FixtureRequest) -> Optional[TlsContexts]: + """ + `None` by default (plain TCP). Parametrize it indirectly with `True` (see `over_plain_and_tls`) + to run a test over TLS; `cryptography` is then required. + """ + + if getattr(request, "param", False): + return request.getfixturevalue("tls_contexts") + return None + +@pytest.fixture +def client_ssl_context(ssl_contexts: Optional[TlsContexts]) -> Optional[ssl.SSLContext]: + return ssl_contexts.client if ssl_contexts else None + +# Run a test twice: over plain TCP then over TLS. +over_plain_and_tls = pytest.mark.parametrize("ssl_contexts", [False, True], indirect=True, ids=["plain", "tls"]) diff --git a/tests/network/test_tcp_client_server.py b/tests/network/test_tcp_client_server.py new file mode 100644 index 0000000..d9f678f --- /dev/null +++ b/tests/network/test_tcp_client_server.py @@ -0,0 +1,446 @@ +# 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 asyncio +import contextlib +import socket +import ssl +import threading +import time + +import pytest + +from tests.network import find_free_tcp_port +from tests.network.conftest import over_plain_and_tls, TlsContexts +from xcp_storage.network.socket import ( + get_socket_port, + SocketDisconnectedError, + SocketError, +) +from xcp_storage.network.tcp_client import TcpClient, TcpClientError +from xcp_storage.network.tcp_server import TcpServer, TcpServerError + +from xcp_storage.typing import ( + Callable, + cast, + Final, + Generator, + Optional, + override, +) + +# ============================================================================== + +SERVER_STARTUP_TIMEOUT: Final = 5.0 +SERVER_SHUTDOWN_TIMEOUT: Final = 5.0 + +SERVER_THREAD_SHUTDOWN_TIMEOUT: Final = 5.0 + +CLIENT_TIMEOUT: Final = 5.0 + +# ------------------------------------------------------------------------------ + +@contextlib.contextmanager +def run_threaded_server(tcp_server: TcpServer) -> Generator[TcpServer, None, None]: + server_thread = threading.Thread(target=tcp_server.run, daemon=True) + server_thread.start() + try: + assert tcp_server.wait_for_startup(timeout=SERVER_STARTUP_TIMEOUT), "Server not started." + yield tcp_server + finally: + tcp_server.stop(timeout=SERVER_SHUTDOWN_TIMEOUT) + server_thread.join(timeout=SERVER_THREAD_SHUTDOWN_TIMEOUT) + assert not server_thread.is_alive(), "Server thread still alive." + +@pytest.fixture +def tcp_server( + request: pytest.FixtureRequest, + ssl_contexts: Optional[TlsContexts] +) -> Generator[TcpServer, None, None]: + server_class = request.param + server = server_class("127.0.0.1", 0, ssl_context=ssl_contexts.server if ssl_contexts else None) + with run_threaded_server(server) as threaded_server: + yield threaded_server + +TcpClientFactory = Callable[[], TcpClient] + +@pytest.fixture +def tcp_client_factory(tcp_server: TcpServer, client_ssl_context: Optional[ssl.SSLContext]) -> TcpClientFactory: + def factory() -> TcpClient: + return TcpClient( + tcp_server.address, + tcp_server.port, + client_timeout=CLIENT_TIMEOUT, + ssl_context=client_ssl_context + ) + return factory + +# ------------------------------------------------------------------------------ + +class EchoServer(TcpServer): + @override + async def _handle_client_connect(self, client: TcpServer.Client) -> bool: + return True + + @override + async def _handle_client_disconnect(self, client: TcpServer.Client) -> None: + pass + + @override + async def _handle_client_request(self, client: TcpServer.Client) -> bool: + try: + data = await client.reader.read(1024) + if not data: + return False + client.writer.write(data) + await client.writer.drain() + return True + except Exception: + return False + +# ------------------------------------------------------------------------------ + +class RejectConnectionServer(TcpServer): + @override + async def _handle_client_connect(self, client: TcpServer.Client) -> bool: + return False + + @override + async def _handle_client_disconnect(self, client: TcpServer.Client) -> None: + pass + + @override + async def _handle_client_request(self, client: TcpServer.Client) -> bool: + return False + +# ------------------------------------------------------------------------------ + +class FragmentedServer(TcpServer): + @override + async def _handle_client_connect(self, client: TcpServer.Client) -> bool: + return True + + @override + async def _handle_client_disconnect(self, client: TcpServer.Client) -> None: + pass + + @override + async def _handle_client_request(self, client: TcpServer.Client) -> bool: + data = await client.reader.read(1024) + if not data: + return False + for octet in data: + client.writer.write(bytes([octet])) + await client.writer.drain() + await asyncio.sleep(0.05) + return True + +# ------------------------------------------------------------------------------ + +class SingleRequestServer(TcpServer): + @override + async def _handle_client_connect(self, client: TcpServer.Client) -> bool: + return True + + @override + async def _handle_client_disconnect(self, client: TcpServer.Client) -> None: + pass + + @override + async def _handle_client_request(self, client: TcpServer.Client) -> bool: + await client.reader.read(1024) + client.writer.write(b"ACK") + await client.writer.drain() + return False + +# ============================================================================== + +class TestTcpClient: + def test_del_without_init(self) -> None: + tcp_client = TcpClient.__new__(TcpClient) + tcp_client.__del__() # Must not raise even if `__init__` has not been called. + + def test_not_connected(self) -> None: + tcp_client = TcpClient("127.0.0.1", port=0) + assert not tcp_client.connected + with pytest.raises(TcpClientError, match="Cannot send. Not connected."): + tcp_client.send(b"") + with pytest.raises(TcpClientError, match="Cannot receive. Not connected."): + tcp_client.receive(bytearray()) + + def test_connect_timeout_failure(self) -> None: + client_timeout = 0.5 + tcp_client = TcpClient("127.0.0.1", port=0, client_timeout=client_timeout) + start_time = time.monotonic() + + with pytest.raises(TcpClientError, match="Unable to connect to server."): + tcp_client.connect() + + elapsed_time = time.monotonic() - start_time + assert elapsed_time >= client_timeout + assert elapsed_time < 1.0 + assert not tcp_client.socket + + def test_connect_zero_timeout(self) -> None: + tcp_client = TcpClient("127.0.0.1", port=0, client_timeout=CLIENT_TIMEOUT) + start_time = time.monotonic() + + with pytest.raises(TcpClientError, match="Unable to connect to server."): + tcp_client.connect(timeout=0.0) + + elapsed_time = time.monotonic() - start_time + assert elapsed_time < 1.0 + assert not tcp_client.socket + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_connect_is_idempotent(self, tcp_client_factory: TcpClientFactory) -> None: + with tcp_client_factory() as tcp_client: + assert tcp_client.connected + + client_socket = tcp_client.socket + assert client_socket + + tcp_client.connect() + assert tcp_client.socket is client_socket + + tcp_client.disconnect() + assert not tcp_client.connected + + def test_disconnect_is_idempotent(self) -> None: + tcp_client = TcpClient("127.0.0.1", port=0) + + tcp_client.disconnect() + assert not tcp_client.connected + tcp_client.disconnect() + assert not tcp_client.connected + + def test_connect_retries_until_server_is_up(self) -> None: + tcp_port = find_free_tcp_port() + tcp_server = EchoServer("127.0.0.1", tcp_port) + + def run() -> None: + time.sleep(0.5) + with run_threaded_server(tcp_server): + time.sleep(3.0) + + thread = threading.Thread(target=run, daemon=True) + thread.start() + try: + with TcpClient("127.0.0.1", tcp_port, client_timeout=CLIENT_TIMEOUT) as tcp_client: + assert tcp_client.connected + finally: + tcp_server.stop(timeout=SERVER_SHUTDOWN_TIMEOUT) + thread.join(timeout=SERVER_THREAD_SHUTDOWN_TIMEOUT) + + @pytest.mark.parametrize("ssl_contexts", [True], indirect=True, ids=["tls"]) + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_connect_with_untrusted_certificate(self, tcp_server: TcpServer) -> None: + # The default context only trusts the system CAs: the server self-signed certificate must be rejected. + tcp_client = TcpClient( + tcp_server.address, + tcp_server.port, + ssl_context=ssl.create_default_context(), + client_timeout=CLIENT_TIMEOUT + ) + with pytest.raises(TcpClientError, match="Unable to connect to server."): + tcp_client.connect() + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_nested_with(self, tcp_server: TcpServer) -> None: + tcp_client = TcpClient(tcp_server.address, tcp_server.port, client_timeout=CLIENT_TIMEOUT) + with tcp_client: + with tcp_client: + assert tcp_client.connected + assert tcp_client.connected + assert not tcp_client.connected + +# ------------------------------------------------------------------------------ + +class TestTcpServer: + def test_running_on_fixed_port(self) -> None: + tcp_port = find_free_tcp_port() + tcp_server = EchoServer("127.0.0.1", tcp_port) + assert tcp_server.port == tcp_port + with run_threaded_server(tcp_server): + assert tcp_server.port == tcp_port + + def test_run_twice(self) -> None: + with run_threaded_server(EchoServer("127.0.0.1", 0)) as threaded_server, \ + pytest.raises(TcpServerError, match="Server is already running."): + threaded_server.run() + + def test_stop_just_after_startup(self) -> None: + # Ensure we don't have a regression/race condition somewhere. + for _ in range(100): + with run_threaded_server(EchoServer("127.0.0.1", 0)): + pass + + def test_startup_on_used_port(self) -> None: + with socket.socket() as sock: + sock.bind(("127.0.0.1", 0)) + tcp_server = EchoServer("127.0.0.1", cast(int, get_socket_port(sock))) + with pytest.raises(SocketError, match="Failed to bind server sock."): + tcp_server.run() + + assert not tcp_server.wait_for_startup(timeout=SERVER_STARTUP_TIMEOUT) + tcp_server.wait_for_shutdown() + +# ------------------------------------------------------------------------------ + +@over_plain_and_tls +class TestTcpClientServer: + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_client_server_echo(self, tcp_client_factory: TcpClientFactory) -> None: + payload = b"Hello World!" + with tcp_client_factory() as tcp_client: + assert tcp_client.connected + + tcp_client.send(payload) + buffer = bytearray(len(payload)) + tcp_client.receive(buffer) + assert bytes(buffer) == payload + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_client_server_echo_with_fixed_size(self, tcp_client_factory: TcpClientFactory) -> None: + payload = b"Hello World!" + part_size = 5 + payload_part = payload[:part_size] + + with tcp_client_factory() as tcp_client: + assert tcp_client.connected + + tcp_client.send(payload, part_size) + buffer = bytearray(16) + tcp_client.receive(buffer, part_size) + assert bytes(buffer[:part_size]) == payload_part + assert bytes(buffer[part_size:]) == bytes(len(buffer) - part_size) + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_client_receive_after_close(self, tcp_client_factory: TcpClientFactory) -> None: + tcp_client = tcp_client_factory() + tcp_client.connect() + assert tcp_client.socket + tcp_client.socket.close() + + buffer = bytearray(16) + with pytest.raises(SocketDisconnectedError): + tcp_client.receive(buffer) + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_client_delete_to_close(self, tcp_client_factory: TcpClientFactory) -> None: + tcp_client = tcp_client_factory() + tcp_client.connect() + assert tcp_client.connected + + tcp_client.__del__() + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [RejectConnectionServer], indirect=True) + def test_client_rejected_by_server(self, tcp_client_factory: TcpClientFactory) -> None: + tcp_client = tcp_client_factory() + tcp_client.connect() + + buffer = bytearray(16) + with pytest.raises(SocketDisconnectedError): + tcp_client.receive(buffer) + + with pytest.raises(TcpClientError, match="Cannot send. Not connected."): + tcp_client.send(b"Moshimoshi?") + + with pytest.raises(TcpClientError, match="Cannot receive. Not connected."): + tcp_client.receive(buffer) + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [RejectConnectionServer], indirect=True) + def test_client_send_to_closed_peer(self, tcp_client_factory: TcpClientFactory) -> None: + tcp_client = tcp_client_factory() + tcp_client.connect() + + with pytest.raises(SocketDisconnectedError, match="Unable to send data."): + for _ in range(100): + tcp_client.send(b"Moshimoshi?") + + with pytest.raises(TcpClientError, match="Cannot send. Not connected."): + tcp_client.send(b"Moshimoshi?") + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [FragmentedServer], indirect=True) + def test_client_handles_fragmented_stream(self, tcp_client_factory: TcpClientFactory) -> None: + payload = b"A fragmented message." + with tcp_client_factory() as tcp_client: + tcp_client.send(payload) + buffer = bytearray(len(payload)) + tcp_client.receive(buffer) + assert bytes(buffer) == payload + + @pytest.mark.parametrize("tcp_server", [SingleRequestServer], indirect=True) + def test_server_single_request(self, tcp_client_factory: TcpClientFactory) -> None: + expected_message = b"ACK" + with tcp_client_factory() as tcp_client: + tcp_client.send(b"Ping?") + buffer = bytearray(len(expected_message)) + tcp_client.receive(buffer) + assert bytes(buffer) == expected_message + + with pytest.raises(SocketDisconnectedError): + tcp_client.receive(buffer) + + assert not tcp_client.connected + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_multiple_server_clients(self, tcp_client_factory: TcpClientFactory) -> None: + with contextlib.ExitStack() as stack: + tcp_clients = [stack.enter_context(tcp_client_factory()) for _ in range(5)] + messages = [f"client-{i}".encode() for i in range(len(tcp_clients))] + + for tcp_client, message in zip(tcp_clients, messages): + tcp_client.send(message) + for tcp_client, message in zip(tcp_clients, messages): + buffer = bytearray(len(message)) + tcp_client.receive(buffer) + assert bytes(buffer) == message + + @pytest.mark.parametrize("tcp_server", [EchoServer], indirect=True) + def test_server_stop_with_multiple_clients( + self, tcp_server: TcpServer, tcp_client_factory: TcpClientFactory + ) -> None: + with contextlib.ExitStack() as stack: + tcp_clients = [stack.enter_context(tcp_client_factory()) for _ in range(5)] + + message = b"ping" + for tcp_client in tcp_clients: + tcp_client.send(message) + buffer = bytearray(len(message)) + tcp_client.receive(buffer) + assert bytes(buffer) == message + + assert tcp_server.stop(timeout=SERVER_SHUTDOWN_TIMEOUT) + + for tcp_client in tcp_clients: + with pytest.raises(SocketDisconnectedError): + tcp_client.send(message) + tcp_client.receive(bytearray(len(message))) + + for tcp_client in tcp_clients: + assert not tcp_client.connected diff --git a/tests/utils/test_sync.py b/tests/utils/test_sync.py new file mode 100644 index 0000000..c4ac16c --- /dev/null +++ b/tests/utils/test_sync.py @@ -0,0 +1,71 @@ +# 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, Mock, patch + +import pytest + +from xcp_storage.utils.sync import wait_for_condition + +# ============================================================================== + +@patch("time.sleep") +class TestWaitForCondition: + def test_success(self, mock_sleep: MagicMock) -> None: + func = Mock(return_value=True) + assert wait_for_condition(func, timeout=5.0, interval=0.1) + func.assert_called_once() + mock_sleep.assert_not_called() + + def test_success_after_retries(self, mock_sleep: MagicMock) -> None: + func = Mock(side_effect=[False, False, True]) + assert wait_for_condition(func, timeout=5.0, interval=0.5) + assert func.call_count == 3 + assert mock_sleep.call_count == 2 + mock_sleep.assert_called_with(0.5) + + def test_no_timeout(self, mock_sleep: MagicMock) -> None: + func = Mock(return_value=False) + assert not wait_for_condition(func, timeout=0, interval=0) + func.assert_called_once() + mock_sleep.assert_not_called() + + def test_negative_timeout(self, mock_sleep: MagicMock) -> None: + func = Mock(return_value=False) + assert not wait_for_condition(func, timeout=-1.0, interval=1.0) + func.assert_called_once() + mock_sleep.assert_not_called() + + @patch("time.monotonic") + def test_timeout_reached(self, mock_time: MagicMock, mock_sleep: MagicMock) -> None: + mock_time.side_effect = [0.0, 2.5, 5.0, 5.0, 5.0] + func = Mock(return_value=False) + assert not wait_for_condition(func, timeout=5.0, interval=1.0) + assert func.call_count == 2 + assert mock_sleep.call_count == 1 + mock_sleep.assert_called_with(1.0) + + @patch("time.monotonic") + def test_sleep_is_capped_by_remaining_time(self, mock_time: MagicMock, mock_sleep: MagicMock) -> None: + mock_time.side_effect = [0.0, 4.5, 5.0, 5.0, 5.0] + func = Mock(return_value=False) + assert not wait_for_condition(func, timeout=5.0, interval=10.0) + mock_sleep.assert_called_once_with(0.5) + + def test_exception_is_propagated(self, mock_sleep: MagicMock) -> None: + exception_message = "Explosion!" + func = Mock(side_effect=Exception(exception_message)) + with pytest.raises(Exception, match=exception_message): + wait_for_condition(func, timeout=5.0, interval=1.0) + mock_sleep.assert_not_called() From 06e6970cf50d70cea0bf283d71a34641c0904ea7 Mon Sep 17 00:00:00 2001 From: Ronan Abhamon Date: Mon, 5 Oct 2026 18:31:09 +0200 Subject: [PATCH 3/4] feat(core): add `socket_wait_readable` helper (#6) Signed-off-by: Ronan Abhamon --- src/xcp_storage/network/socket.py | 19 ++++++++++ tests/network/test_socket.py | 61 ++++++++++++++++++++++++++++++- 2 files changed, 78 insertions(+), 2 deletions(-) diff --git a/src/xcp_storage/network/socket.py b/src/xcp_storage/network/socket.py index 5c8a1c2..c74f6bd 100644 --- a/src/xcp_storage/network/socket.py +++ b/src/xcp_storage/network/socket.py @@ -299,6 +299,22 @@ def socket_receive(sock: socket.socket, buffer: bytearray, size: Optional[int] = # ------------------------------------------------------------------------------ +def socket_wait_readable(sock: socket.socket, *, timeout: Optional[float] = None) -> bool: + if sock.fileno() < 0: + raise SocketDisconnectedError("Unable to wait for data. Socket is closed.") + + # An SSL socket may already hold decrypted data that `select` cannot see. + if isinstance(sock, ssl.SSLSocket) and sock.pending(): + return True + + try: + readable, _, _ = select.select([sock], [], [], timeout) + except OSError as e: + raise SocketDisconnectedError("Unable to wait for data.") from e + return bool(readable) + +# ------------------------------------------------------------------------------ + def get_socket_family_str(sock: socket.socket) -> str: return _FAMILY_TO_STR.get(sock.family, "Unknown") @@ -340,6 +356,9 @@ def send(self, buffer: bytes, size: Optional[int] = None) -> None: def receive(self, buffer: bytearray, size: Optional[int] = None) -> None: socket_receive(self.sock, buffer, size) + def wait_readable(self, *, timeout: Optional[float] = None) -> bool: + return socket_wait_readable(self.sock, timeout=timeout) + def close(self) -> None: if self.keep_open: return diff --git a/tests/network/test_socket.py b/tests/network/test_socket.py index 0cb9942..cbfe618 100644 --- a/tests/network/test_socket.py +++ b/tests/network/test_socket.py @@ -15,6 +15,7 @@ import errno import ipaddress import socket +import ssl from unittest.mock import ( call, MagicMock, @@ -34,12 +35,17 @@ Socket, socket_receive, socket_send, + socket_wait_readable, SocketDisconnectedError, SocketError, SocketTimeoutError, ) -from xcp_storage.typing import Final, Optional +from xcp_storage.typing import ( + Final, + Iterator, + Optional, +) # ============================================================================== @@ -221,7 +227,18 @@ def test_get_socket_port_closed_socket(self) -> None: @pytest.fixture def mock_sock() -> MagicMock: - return MagicMock(spec=socket.socket) + sock = MagicMock(spec=socket.socket) + sock.fileno.return_value = 3 # An open socket. + return sock + +@pytest.fixture +def mock_ssl_sock() -> MagicMock: + sock = MagicMock(spec=ssl.SSLSocket) + sock.fileno.return_value = 4 + sock.pending.return_value = 0 + return sock + +# ------------------------------------------------------------------------------ class TestSocketTransfer: MESSAGE: Final = b"hello world" @@ -339,6 +356,46 @@ def test_socket_receive_os_error_generic(self, mock_sock: MagicMock) -> None: # ------------------------------------------------------------------------------ +class TestSocketWaitReadable: + @pytest.fixture + def mock_select(self) -> Iterator[MagicMock]: + # Nothing is readable by default. + with patch("select.select", return_value=([], [], [])) as mock_select: + yield mock_select + + def test_wait_readable(self, mock_select: MagicMock, mock_sock: MagicMock) -> None: + timeout = 0.5 + assert not socket_wait_readable(mock_sock, timeout=timeout) + mock_select.assert_called_once_with([mock_sock], [], [], timeout) + + mock_select.reset_mock(return_value=True) + mock_select.return_value = ([mock_sock], [], []) + assert socket_wait_readable(mock_sock) + mock_select.assert_called_once_with([mock_sock], [], [], None) + + def test_wait_readable_on_closed_socket(self, mock_select: MagicMock, mock_sock: MagicMock) -> None: + mock_sock.fileno.return_value = -1 + with pytest.raises(SocketDisconnectedError, match="Unable to wait for data. Socket is closed."): + socket_wait_readable(mock_sock) + mock_select.assert_not_called() + + def test_wait_select_exception(self, mock_select: MagicMock, mock_sock: MagicMock) -> None: + mock_select.side_effect = OSError(errno.EBADF) + with pytest.raises(SocketDisconnectedError, match="Unable to wait for data."): + socket_wait_readable(mock_sock) + + def test_wait_readable_with_pending_ssl_data(self, mock_select: MagicMock, mock_ssl_sock: MagicMock) -> None: + mock_ssl_sock.pending.return_value = 1 + assert socket_wait_readable(mock_ssl_sock) + mock_select.assert_not_called() + + def test_wait_readable_without_pending_ssl_data(self, mock_select: MagicMock, mock_ssl_sock: MagicMock) -> None: + mock_select.return_value = ([mock_ssl_sock], [], []) + assert socket_wait_readable(mock_ssl_sock) + mock_select.assert_called_once_with([mock_ssl_sock], [], [], None) + +# ------------------------------------------------------------------------------ + class TestSocketWrapper: def test_socket_context_manager(self) -> None: mock_sock = MagicMock() From 09737e5b2b43cae436ebfbbcb8080a9a4494cca1 Mon Sep 17 00:00:00 2001 From: Ronan Abhamon Date: Mon, 5 Oct 2026 19:00:03 +0200 Subject: [PATCH 4/4] fix(tests/core): robustify `test_client_send_to_closed_peer` test (#6) Before this fix, this test was unstable and could fail. It's wiser to wait for the server to close instead of spamming calls. Signed-off-by: Ronan Abhamon --- tests/network/test_tcp_client_server.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/tests/network/test_tcp_client_server.py b/tests/network/test_tcp_client_server.py index d9f678f..eb3719c 100644 --- a/tests/network/test_tcp_client_server.py +++ b/tests/network/test_tcp_client_server.py @@ -376,8 +376,14 @@ def test_client_send_to_closed_peer(self, tcp_client_factory: TcpClientFactory) tcp_client = tcp_client_factory() tcp_client.connect() + # Wait for the server to close its side. Otherwise all the sends below could complete + # before the server has handled the connection, and nothing would fail. + assert tcp_client.socket + assert tcp_client.socket.wait_readable(timeout=2.0) + + deadline = time.monotonic() + 1.0 with pytest.raises(SocketDisconnectedError, match="Unable to send data."): - for _ in range(100): + while time.monotonic() < deadline: tcp_client.send(b"Moshimoshi?") with pytest.raises(TcpClientError, match="Cannot send. Not connected."):