Skip to content

Commit eb29f76

Browse files
committed
feat(core): add socket module
Signed-off-by: Ronan Abhamon <ronan.abhamon@vates.tech>
1 parent b5ce436 commit eb29f76

2 files changed

Lines changed: 650 additions & 0 deletions

File tree

‎src/xcp_storage/network/socket.py‎

Lines changed: 316 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,316 @@
1+
# Copyright (C) 2026 Vates SAS
2+
#
3+
# This program is free software: you can redistribute it and/or modify
4+
# it under the terms of the GNU General Public License as published by
5+
# the Free Software Foundation, either version 3 of the License, or
6+
# (at your option) any later version.
7+
# This program is distributed in the hope that it will be useful,
8+
# but WITHOUT ANY WARRANTY; without even the implied warranty of
9+
# MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
10+
# GNU General Public License for more details.
11+
#
12+
# You should have received a copy of the GNU General Public License
13+
# along with this program. If not, see <https://www.gnu.org/licenses/>.
14+
15+
import contextlib
16+
import errno
17+
import ipaddress
18+
import select
19+
import socket
20+
import ssl
21+
from types import TracebackType
22+
23+
from xcp_storage.typing import (
24+
Any,
25+
Dict,
26+
Optional,
27+
Tuple,
28+
Type,
29+
Union,
30+
)
31+
32+
# ==============================================================================
33+
34+
_FAMILY_TO_STR = {
35+
socket.AF_INET: "IPv4",
36+
socket.AF_INET6: "IPv6",
37+
socket.AF_UNIX: "Unix"
38+
}
39+
40+
_IP_VERSION_TO_FAMILY = {
41+
4: socket.AF_INET,
42+
6: socket.AF_INET6,
43+
0: socket.AF_UNSPEC
44+
}
45+
46+
_SERVER_BACKLOG = 128
47+
48+
# ------------------------------------------------------------------------------
49+
50+
class SocketError(Exception):
51+
def __init__(self, message: str) -> None:
52+
super().__init__(message)
53+
54+
class SocketTimeoutError(SocketError):
55+
def __init__(self) -> None:
56+
super().__init__("Socket timeout.")
57+
58+
class SocketDisconnectedError(SocketError):
59+
def __init__(self, message: str) -> None:
60+
super().__init__(message)
61+
62+
# ------------------------------------------------------------------------------
63+
64+
def get_ip_address(address: str, ip_version: int = 0) -> Union[ipaddress.IPv4Address, ipaddress.IPv6Address]:
65+
# Note: `address` can be a hostname or IP.
66+
with contextlib.suppress(ValueError):
67+
return ipaddress.ip_address(address)
68+
69+
family = _IP_VERSION_TO_FAMILY.get(ip_version)
70+
if family is None:
71+
raise SocketError("Unknown IP version.")
72+
73+
try:
74+
info = socket.getaddrinfo(address or socket.gethostname(), 80, family, socket.SOCK_STREAM, socket.SOL_TCP)
75+
return ipaddress.ip_address(info[0][4][0])
76+
except socket.gaierror as e:
77+
raise SocketError("Cannot resolve IP.") from e
78+
except IndexError:
79+
raise SocketError("Cannot resolve IP: no valid address.") from None
80+
81+
# ------------------------------------------------------------------------------
82+
83+
def format_address(address: str, port: int) -> Tuple[socket.AddressFamily, Union[
84+
Tuple[str, int],
85+
Tuple[str, int, int, int]
86+
]]:
87+
if not address:
88+
raise SocketError("No hostname/IP.")
89+
90+
ip_address = get_ip_address(address)
91+
if ip_address.version == 4:
92+
family = socket.AF_INET
93+
return (family, (str(ip_address), port))
94+
if ip_address.version == 6:
95+
family = socket.AF_INET6
96+
return (family, (str(ip_address), port, 0, 0))
97+
98+
raise SocketError("Unknown IP version.")
99+
100+
# ------------------------------------------------------------------------------
101+
102+
def _create_stream_sock(
103+
address: str,
104+
family: socket.AddressFamily,
105+
*,
106+
bind: bool,
107+
reuse_address: bool,
108+
ssl_context: Optional[ssl.SSLContext]
109+
) -> socket.socket:
110+
sock = socket.socket(family, socket.SOCK_STREAM)
111+
if ssl_context:
112+
if bind:
113+
sock = ssl_context.wrap_socket(sock, server_side=True)
114+
else:
115+
sock = ssl_context.wrap_socket(sock, server_side=False, server_hostname=address or None)
116+
117+
if reuse_address:
118+
set_socket_reuseaddr(sock)
119+
120+
return sock
121+
122+
def _normalize_timeout(timeout: Optional[float]) -> Optional[float]:
123+
if timeout is not None and timeout < 0:
124+
timeout = None
125+
return timeout
126+
127+
# ------------------------------------------------------------------------------
128+
129+
def create_server_sock(
130+
address: str,
131+
port: int,
132+
*,
133+
reuse_address: bool = True,
134+
keep_alive: bool = True,
135+
timeout: Optional[float] = None,
136+
ssl_context: Optional[ssl.SSLContext] = None
137+
) -> socket.socket:
138+
family, bind = format_address(address, port)
139+
sock = _create_stream_sock(address, family, bind=True, reuse_address=reuse_address, ssl_context=ssl_context)
140+
141+
timeout = _normalize_timeout(timeout)
142+
if timeout is not None:
143+
sock.settimeout(timeout)
144+
145+
try:
146+
sock.bind(bind)
147+
sock.listen(_SERVER_BACKLOG)
148+
except OSError as e:
149+
with contextlib.suppress(Exception):
150+
sock.close()
151+
raise SocketError("Failed to bind server sock.") from e
152+
153+
if keep_alive:
154+
set_socket_keepalive(sock)
155+
156+
return sock
157+
158+
def create_client_sock(
159+
address: str,
160+
port: int,
161+
*,
162+
reuse_address: bool = True,
163+
keep_alive: bool = True,
164+
timeout: Optional[float] = None,
165+
ssl_context: Optional[ssl.SSLContext] = None
166+
) -> socket.socket:
167+
family, connect = format_address(address, port)
168+
sock = _create_stream_sock(address, family, bind=False, reuse_address=reuse_address, ssl_context=ssl_context)
169+
170+
timeout = _normalize_timeout(timeout)
171+
if timeout is not None:
172+
sock.settimeout(timeout)
173+
174+
while True:
175+
try:
176+
sock.connect(connect)
177+
break
178+
except OSError as e:
179+
if e.errno in (errno.EAGAIN, errno.EINPROGRESS):
180+
_, ready, _ = select.select([], [sock], [], timeout)
181+
if sock in ready:
182+
continue
183+
184+
with contextlib.suppress(Exception):
185+
sock.close()
186+
raise SocketError("Failed to connect client sock.") from e
187+
188+
if keep_alive:
189+
set_socket_keepalive(sock)
190+
191+
return sock
192+
193+
# ------------------------------------------------------------------------------
194+
195+
def set_socket_reuseaddr(sock: socket.socket) -> None:
196+
with contextlib.suppress(Exception):
197+
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
198+
199+
def set_socket_keepalive(sock: socket.socket) -> None:
200+
with contextlib.suppress(Exception):
201+
sock.setsockopt(socket.SOL_SOCKET, socket.SO_KEEPALIVE, 1)
202+
203+
# ------------------------------------------------------------------------------
204+
205+
def _get_buffer_size(buffer: Union[bytes, bytearray], size: Optional[int] = None) -> int:
206+
if size is None:
207+
size = len(buffer)
208+
else:
209+
assert len(buffer) >= size, "Buffer size must be greater than or equal to size."
210+
return size
211+
212+
# ------------------------------------------------------------------------------
213+
214+
def socket_send(sock: socket.socket, buffer: bytes, size: Optional[int] = None) -> None:
215+
size = _get_buffer_size(buffer, size)
216+
view = memoryview(buffer)
217+
218+
pos = 0
219+
while pos < size:
220+
try:
221+
n = sock.send(view[pos:size])
222+
if not n:
223+
break
224+
pos += n
225+
except TimeoutError:
226+
raise SocketTimeoutError() from None
227+
except OSError as e:
228+
if e.errno in (errno.EAGAIN, errno.EWOULDBLOCK):
229+
select.select([], [sock], [])
230+
continue
231+
raise SocketDisconnectedError("Unable to send data.") from e
232+
233+
if pos != size:
234+
raise SocketDisconnectedError("Not enough data sent.") from None
235+
236+
# ------------------------------------------------------------------------------
237+
238+
def socket_receive(sock: socket.socket, buffer: bytearray, size: Optional[int] = None) -> None:
239+
size = _get_buffer_size(buffer, size)
240+
view = memoryview(buffer)
241+
242+
pos = 0
243+
while pos < size:
244+
try:
245+
n = sock.recv_into(view[pos:], min(size - pos, 8192))
246+
if not n:
247+
break
248+
pos += n
249+
except TimeoutError:
250+
raise SocketTimeoutError() from None
251+
except OSError as e:
252+
if e.errno in (errno.EAGAIN, errno.EWOULDBLOCK):
253+
select.select([sock], [], [])
254+
continue
255+
raise SocketDisconnectedError("Unable to receive data.") from e
256+
257+
if pos != size:
258+
raise SocketDisconnectedError("Not enough data received.") from None
259+
260+
# ------------------------------------------------------------------------------
261+
262+
def get_socket_family_str(sock: socket.socket) -> str:
263+
return _FAMILY_TO_STR.get(sock.family, "Unknown")
264+
265+
# ------------------------------------------------------------------------------
266+
267+
class Socket:
268+
def __init__(self, sock: socket.socket, *, keep_open: bool = False) -> None:
269+
self.sock = sock
270+
self.keep_open = keep_open
271+
272+
def __del__(self) -> None:
273+
self.close()
274+
275+
def __enter__(self) -> "Socket":
276+
return self
277+
278+
def __exit__(
279+
self,
280+
exc_type: Optional[Type[BaseException]],
281+
exc_value: Optional[BaseException],
282+
traceback: Optional[TracebackType]
283+
) -> None:
284+
self.close()
285+
286+
def send(self, buffer: bytes, size: Optional[int] = None) -> None:
287+
socket_send(self.sock, buffer, size)
288+
289+
def receive(self, buffer: bytearray, size: Optional[int] = None) -> None:
290+
socket_receive(self.sock, buffer, size)
291+
292+
def close(self) -> None:
293+
if self.keep_open:
294+
return
295+
with contextlib.suppress(Exception):
296+
self.sock.shutdown(socket.SHUT_RDWR)
297+
with contextlib.suppress(Exception):
298+
self.sock.close()
299+
300+
@property
301+
def family_str(self) -> str:
302+
return get_socket_family_str(self.sock)
303+
304+
@property
305+
def timeout(self) -> Optional[float]:
306+
return self.sock.gettimeout()
307+
308+
@timeout.setter
309+
def timeout(self, value: Optional[float]) -> None:
310+
self.sock.settimeout(value)
311+
312+
@property
313+
def peer_certificate(self) -> Optional[Dict[str, Any]]:
314+
if isinstance(self.sock, ssl.SSLSocket):
315+
return self.sock.getpeercert()
316+
return None

0 commit comments

Comments
 (0)