Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -57,7 +57,7 @@ disallow_untyped_defs = false

[tool.pyrefly]
python-version = "3.14"
search-path = ["stubs"]
search-path = [".", "stubs"]
project-includes = [
"src",
"tests"
Expand Down
29 changes: 29 additions & 0 deletions src/xcp_storage/network/socket.py
Original file line number Diff line number Diff line change
Expand Up @@ -299,9 +299,31 @@ 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")

def get_socket_port(sock: socket.socket) -> Optional[int]:
try:
return sock.getsockname()[1]
except OSError:
return None

# ------------------------------------------------------------------------------

class Socket(contextlib.AbstractContextManager):
Expand Down Expand Up @@ -334,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
Expand All @@ -346,6 +371,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()
Expand Down
145 changes: 145 additions & 0 deletions src/xcp_storage/network/tcp_client.py
Original file line number Diff line number Diff line change
@@ -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 <https://www.gnu.org/licenses/>.

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
Loading
Loading