Skip to content

Commit 94d3c6d

Browse files
committed
feat(core): add TCP client/server modules
Signed-off-by: Ronan Abhamon <ronan.abhamon@vates.tech>
1 parent b1f593c commit 94d3c6d

7 files changed

Lines changed: 664 additions & 0 deletions

File tree

‎src/xcp_storage/network/socket.py‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -262,6 +262,12 @@ def socket_receive(sock: socket.socket, buffer: bytearray, size: Optional[int] =
262262
def get_socket_family_str(sock: socket.socket) -> str:
263263
return _FAMILY_TO_STR.get(sock.family, "Unknown")
264264

265+
def get_socket_port(sock: socket.socket) -> Optional[int]:
266+
try:
267+
return sock.getsockname()[1]
268+
except OSError:
269+
return None
270+
265271
# ------------------------------------------------------------------------------
266272

267273
class Socket:
@@ -301,6 +307,10 @@ def close(self) -> None:
301307
def family_str(self) -> str:
302308
return get_socket_family_str(self.sock)
303309

310+
@property
311+
def port(self) -> Optional[int]:
312+
return get_socket_port(self.sock)
313+
304314
@property
305315
def timeout(self) -> Optional[float]:
306316
return self.sock.gettimeout()
Lines changed: 132 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,132 @@
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 ssl
16+
from types import TracebackType
17+
18+
from xcp_storage.network.socket import (
19+
create_client_sock,
20+
Socket,
21+
SocketDisconnectedError,
22+
)
23+
from xcp_storage.utils.sync import wait_for_condition
24+
25+
from xcp_storage.typing import (
26+
Optional,
27+
Type,
28+
)
29+
30+
# ==============================================================================
31+
32+
class TcpClientError(Exception):
33+
def __init__(self, message: str) -> None:
34+
super().__init__(message)
35+
36+
# ------------------------------------------------------------------------------
37+
38+
class TcpClient:
39+
def __init__(
40+
self,
41+
address: str,
42+
port: int,
43+
ssl_context: Optional[ssl.SSLContext] = None,
44+
client_timeout: float = 120
45+
) -> None:
46+
self._address = address
47+
self._port = port
48+
self._ssl_context = ssl_context
49+
self._client_timeout = client_timeout
50+
self._socket: Optional[Socket] = None
51+
52+
self._entered_count = 0
53+
54+
def __del__(self) -> None:
55+
self.disconnect()
56+
57+
def __enter__(self) -> "TcpClient":
58+
if not self._socket:
59+
self.connect()
60+
self._entered_count += 1
61+
return self
62+
63+
def __exit__(
64+
self,
65+
exc_type: Optional[Type[BaseException]],
66+
exc_value: Optional[BaseException],
67+
traceback: Optional[TracebackType]
68+
) -> None:
69+
self._entered_count -= 1
70+
if self._entered_count == 0:
71+
self.disconnect()
72+
73+
@property
74+
def socket(self) -> Optional[Socket]:
75+
return self._socket
76+
77+
@property
78+
def connected(self) -> bool:
79+
return self._socket is not None
80+
81+
def connect(self, timeout: Optional[float] = None) -> None:
82+
if self._socket:
83+
return
84+
85+
if timeout is None:
86+
# If the `connect` timeout is not set, we use the client one.
87+
timeout = self._client_timeout
88+
89+
error: Optional[Exception] = None
90+
def connect_impl() -> bool:
91+
nonlocal error
92+
try:
93+
self._socket = Socket(create_client_sock(
94+
self._address,
95+
self._port,
96+
reuse_address=True,
97+
keep_alive=True,
98+
timeout=self._client_timeout,
99+
ssl_context=self._ssl_context
100+
))
101+
except Exception as e:
102+
error = e
103+
return False
104+
return True
105+
106+
if not wait_for_condition(connect_impl, timeout=timeout, interval=1):
107+
raise TcpClientError("Unable to connect to server.") from error
108+
109+
def disconnect(self) -> None:
110+
if self._socket:
111+
self._socket.close()
112+
self._socket = None
113+
114+
def send(self, buffer: bytes, size: Optional[int] = None) -> None:
115+
if self._socket:
116+
try:
117+
self._socket.send(buffer, size)
118+
except SocketDisconnectedError:
119+
self.disconnect()
120+
raise
121+
else:
122+
TcpClientError("Cannot send. Not connected.")
123+
124+
def receive(self, buffer: bytearray, size: Optional[int] = None) -> None:
125+
if self._socket:
126+
try:
127+
self._socket.receive(buffer, size)
128+
except SocketDisconnectedError:
129+
self.disconnect()
130+
raise
131+
else:
132+
TcpClientError("Cannot receive. Not connected.")
Lines changed: 212 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,212 @@
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+
from abc import ABC, abstractmethod
16+
import asyncio
17+
import ssl
18+
import sys
19+
import threading
20+
21+
import xcp_storage.log as log
22+
from xcp_storage.network.socket import create_server_sock, Socket
23+
from xcp_storage.utils.asyncio import cancel_event_loop_tasks, close_stream_writer
24+
25+
from xcp_storage.typing import (
26+
Any,
27+
Optional,
28+
override,
29+
Set,
30+
)
31+
32+
# ==============================================================================
33+
34+
logger = log.get_logger() # Use default logger.
35+
36+
# ------------------------------------------------------------------------------
37+
38+
class TcpServerError(Exception):
39+
def __init__(self, message: str) -> None:
40+
super().__init__(message)
41+
42+
# ------------------------------------------------------------------------------
43+
44+
class TcpServer(ABC):
45+
class Client:
46+
def __init__(
47+
self,
48+
peername: Any, # noqa: ANN401
49+
reader: asyncio.StreamReader,
50+
writer: asyncio.StreamWriter
51+
) -> None:
52+
self.peername: Any = peername
53+
self.reader = reader
54+
self.writer = writer
55+
56+
@override
57+
def __str__(self) -> str:
58+
return str(self.peername)
59+
60+
def __init__(self, address: str, port: int, ssl_context: Optional[ssl.SSLContext] = None) -> None:
61+
self._address = address
62+
self._port = port
63+
self._ssl_context = ssl_context
64+
65+
self._running = False
66+
67+
self._event_loop: Optional[asyncio.AbstractEventLoop] = None
68+
self._startup_event = threading.Event()
69+
70+
self._server_socket: Optional[Socket] = None
71+
self._server: Optional[asyncio.AbstractServer] = None
72+
73+
self._clients: Set[TcpServer.Client] = set()
74+
75+
@property
76+
def address(self) -> str:
77+
return self._address
78+
79+
@property
80+
def port(self) -> int:
81+
# Return trivial port if it's not 0. As reminder, 0 = dynamic binding.
82+
if self._port or not self._server_socket:
83+
return self._port
84+
85+
port = self._server_socket.port
86+
return port if port is not None else 0
87+
88+
def run(self) -> None:
89+
if self._running:
90+
raise TcpServerError("Server is already running.")
91+
92+
self._startup_event.clear()
93+
self._running = True
94+
95+
try:
96+
logger.info("Running TCP server on `%s:%d`...", self._address, self._port)
97+
self._server_socket = Socket(create_server_sock(
98+
self._address,
99+
self._port,
100+
reuse_address=True,
101+
keep_alive=True,
102+
timeout=0.0, # 0 here to set non-blocking mode.
103+
ssl_context=self._ssl_context
104+
))
105+
106+
try:
107+
old_event_loop = asyncio.get_event_loop()
108+
except RuntimeError:
109+
# No event loop.
110+
old_event_loop = None
111+
112+
self._event_loop = asyncio.new_event_loop()
113+
asyncio.set_event_loop(self._event_loop)
114+
115+
server_params = {
116+
"client_connected_cb": self._handle_client,
117+
"sock": self._server_socket.sock
118+
}
119+
120+
if sys.version_info < (3, 8):
121+
# TODO(XCPNG-3032): Workaround for old python versions. Remove me later.
122+
# In fact we must give `loop` param for these versions and the entire
123+
# event loop management is also required just for it.
124+
server_params["loop"] = self._event_loop
125+
126+
self._server = self._event_loop.run_until_complete(
127+
asyncio.start_server(**server_params) # type: ignore[arg-type]
128+
)
129+
self._startup_event.set()
130+
131+
logger.info("TCP server started!")
132+
self._event_loop.run_forever()
133+
except KeyboardInterrupt:
134+
logger.info("Closing server because break signal has been received...")
135+
except Exception:
136+
self._startup_event.set()
137+
raise
138+
finally:
139+
if self._event_loop:
140+
if self._server:
141+
self._server.close()
142+
self._event_loop.run_until_complete(self._server.wait_closed())
143+
self._server = None
144+
145+
try:
146+
cancel_event_loop_tasks(self._event_loop)
147+
finally:
148+
self._event_loop.close()
149+
self._event_loop = None
150+
if old_event_loop:
151+
asyncio.set_event_loop(old_event_loop)
152+
153+
if self._server_socket:
154+
self._server_socket.close()
155+
self._server_socket = None
156+
157+
self._clients.clear()
158+
self._running = False
159+
160+
def async_stop(self) -> None:
161+
if self._event_loop and self._event_loop.is_running():
162+
self._event_loop.call_soon_threadsafe(self._event_loop.stop)
163+
164+
def wait_for_startup(self) -> None:
165+
self._startup_event.wait()
166+
167+
async def _handle_client(
168+
self,
169+
client_reader: asyncio.StreamReader,
170+
client_writer: asyncio.StreamWriter
171+
) -> None:
172+
client = self.Client(client_writer.get_extra_info("peername"), client_reader, client_writer)
173+
logger.info("New client %s connected.", client)
174+
self._clients.add(client)
175+
176+
rejected = False
177+
try:
178+
if not await self._handle_client_connect(client):
179+
rejected = True
180+
return
181+
while not client_writer.transport.is_closing():
182+
if not await self._handle_client_request(client):
183+
break
184+
logger.info("Client %s has terminated.", client)
185+
except asyncio.TimeoutError as e:
186+
logger.warning("Timeout reached for client %s: `%s`.", client, e)
187+
except asyncio.IncompleteReadError as e:
188+
logger.warning("Connection closed for client %s: `%s`.", client, e)
189+
except Exception as e:
190+
logger.error("Unhandled exception for client %s: `%s`.", client, e)
191+
finally:
192+
if not rejected:
193+
try:
194+
await self._handle_client_disconnect(client)
195+
except Exception as e:
196+
logger.error("Unhandled exception for client %s during disconnect: `%s`.", client, e)
197+
198+
await close_stream_writer(client_writer)
199+
logger.info("Client %s disconnected.", client)
200+
self._clients.remove(client)
201+
202+
@abstractmethod
203+
async def _handle_client_connect(self, client: Client) -> bool:
204+
return False
205+
206+
@abstractmethod
207+
async def _handle_client_disconnect(self, client: Client) -> None:
208+
pass
209+
210+
@abstractmethod
211+
async def _handle_client_request(self, client: Client) -> bool:
212+
return False

0 commit comments

Comments
 (0)