-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtmp_network_socket_server.py
More file actions
65 lines (57 loc) · 14.5 KB
/
Copy pathtmp_network_socket_server.py
File metadata and controls
65 lines (57 loc) · 14.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
"""
===============================================================================
RemoteDesk Pro
File: network/socket_server.py
Implements the server-side socket handling for RemoteDesk Pro.
Manages incoming client connections and message routing.
===============================================================================
"""
from __future__ import annotations
import socket
import threading
import time
import struct
import uuid
from typing import Optional, Callable, Dict, Tuple, List
from core.constants import DEFAULT_HOST, DEFAULT_PORT, SOCKET_TIMEOUT
from core.logger import logger
from network.packet_system import PacketSystem
from network.protocol import RemoteDeskMessage, MessageFactory, MessageType
class SocketServer:
"""
Server class that handles incoming socket connections.
Supports multiple concurrent client connections.
"""
def __init__(
self,
host: str = DEFAULT_HOST,
port: int = DEFAULT_PORT,
max_clients: int = 10,
on_client_connect: Optional[Callable[[str, Tuple[str, int]], None]] = None,
on_client_disconnect: Optional[Callable[[str], None]] = None,
on_message_received: Optional[Callable[[str, RemoteDeskMessage], None]] = None,
) -> None:
"""
Initialize the socket server.
Args:
host: Server host address
port: Server port number
max_clients: Maximum number of concurrent clients
on_client_connect: Callback for client connection (client_id, address_tuple)
on_client_disconnect: Callback for client disconnection (client_id)
on_message_received: Callback for incoming messages (client_id, message_object)
"""
self._host = host
self._port = port
self._max_clients = max_clients
self._on_client_connect = on_client_connect
self._on_client_disconnect = on_client_disconnect
self._on_message_received = on_message_received
self._server_socket: Optional[socket.socket] = None
self._clients: Dict[str, socket.socket] = {}\n self._client_addresses: Dict[str, Tuple[str, int]] = {}\n self._client_threads: Dict[str, threading.Thread] = {}\n self._running = False\n self._lock = threading.Lock()\n self._packet_system = PacketSystem()\n \n def start(self) -> bool:\n \"\"\"\n Start the server and begin listening for connections.\n \n Returns:\n True if server started successfully, False otherwise\n \"\"\"\n with self._lock:\n if self._running:\n logger.warning(\"Server is already running.\")\n return True\n \n try:\n self._server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n self._server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)\n self._server_socket.bind((self._host, self._port))\n self._server_socket.listen(self._max_clients)\n self._server_socket.settimeout(1.0) # Allows periodic loop checks\n \n self._running = True\n \n accept_thread = threading.Thread(\n target=self._accept_clients_loop,\n name=\"ServerAcceptThread\",\n daemon=True\n )\n accept_thread.start()\n \n logger.info(f\"Server started on {self._host}:{self._port}\")\n return True\n \n except Exception as e:\n logger.error(f\"Failed to start server on {self._host}:{self._port}: {e}\", exc_info=True)\n self._cleanup_server()\n return False\n \n def _accept_clients_loop(self) -> None:\n \"\"\"\n Continuously accepts incoming client connections in a separate thread.\n \"\"\"\n while self._running:\n try:\n client_socket, address = self._server_socket.accept()\n client_id = str(uuid.uuid4()) # Generate a unique ID for each client\n \n with self._lock:\n if len(self._clients) >= self._max_clients:\n logger.warning(f\"Connection rejected from {address}: Max clients reached.\")\n client_socket.close()\n continue\n\n self._clients[client_id] = client_socket\n self._client_addresses[client_id] = address\n \n # Start a dedicated thread for client communication\n client_thread = threading.Thread(\n target=self._handle_client_communication,\n args=(client_id, client_socket),\n name=f\"ClientHandler-{client_id}\",\n daemon=True\n )\n self._client_threads[client_id] = client_thread\n client_thread.start()\n \n logger.info(f\"Client {client_id} connected from {address}\")\n if self._on_client_connect:\n self._on_client_connect(client_id, address)\n\n except socket.timeout:\n # No new connections, continue loop\n pass\n except Exception as e:\n if self._running:\n logger.error(f\"Error in accept loop: {e}\", exc_info=True)\n break # Exit loop if server socket fails\n\n def _handle_client_communication(self, client_id: str, client_socket: socket.socket) -> None:\n \"\"\"\n Handles receiving and processing messages from a specific client.\n \n Args:\n client_id: Unique identifier for the client\n client_socket: The socket object for this client\n \"\"\"\n client_socket.settimeout(SOCKET_TIMEOUT)\n buffer = b\"\"\n \n try:\n while self._running:\n try:\n data = client_socket.recv(4096)\n if not data:\n logger.info(f\"Client {client_id} closed connection gracefully.\")\n break\n \n buffer += data\n \n while len(buffer) >= 4: # Check for length prefix\n length = struct.unpack(\'!I\', buffer[:4])[0]\n if len(buffer) < 4 + length: # Incomplete packet\n break\n \n packet_data = buffer[4 : 4 + length]\n buffer = buffer[4 + length:] # Remove processed packet from buffer\n \n message_dict = self._packet_system.parse_packet(packet_data)\n if message_dict:\n rd_message = RemoteDeskMessage.from_dict(message_dict)\n if rd_message.is_valid() and self._on_message_received:\n self._on_message_received(client_id, rd_message)\n else:\n logger.warning(f\"Invalid or unhandled message from {client_id}: {message_dict}\")\n\n except socket.timeout:\n # No data for a while, continue checking\n pass\n except ConnectionResetError:\n logger.info(f\"Client {client_id} reset connection.\")\n break\n except Exception as e:\n logger.error(f\"Error handling communication with {client_id}: {e}\", exc_info=True)\n break\n\n finally:\n self._disconnect_client(client_id)\n \n def send_message(self, client_id: str, message: RemoteDeskMessage) -> bool:\n \"\"\"\n Send a message to a specific client.\n \n Args:\n client_id: Target client identifier\n message: RemoteDeskMessage object to send\n \n Returns:\n True if sent successfully, False otherwise\n \"\"\"\n with self._lock:\n client_socket = self._clients.get(client_id)\n if not client_socket:\n logger.warning(f\"Attempted to send message to unknown client: {client_id}\")\n return False\n \n try:\n packet_data = self._packet_system.create_packet(message.to_dict())\n client_socket.sendall(packet_data)\n logger.debug(f\"Sent {message.message_type} message to {client_id}\")\n return True\n except Exception as e:\n logger.error(f\"Failed to send message to {client_id}: {e}\", exc_info=True)\n self._disconnect_client(client_id) # Consider disconnecting on send failure\n return False\n\n def broadcast_message(self, message: RemoteDeskMessage, exclude_client_id: Optional[str] = None) -> int:\n \"\"\"\n Send a message to all connected clients, optionally excluding one.\n \n Args:\n message: RemoteDeskMessage object to broadcast\n exclude_client_id: Optional client ID to exclude from broadcast\n \n Returns:\n Number of clients that successfully received the message\n \"\"\"\n sent_count = 0\n with self._lock:\n clients_to_send = [\n (cid, sock)\n for cid, sock in self._clients.items()\n if cid != exclude_client_id\n ]\n packet_data = self._packet_system.create_packet(message.to_dict())\n\n for client_id, client_socket in clients_to_send:\n try:\n client_socket.sendall(packet_data)\n sent_count += 1\n except Exception as e:\n logger.error(f\"Broadcast failed to {client_id}: {e}\", exc_info=True)\n self._disconnect_client(client_id)\n logger.debug(f\"Broadcasted {message.message_type} to {sent_count} clients\")\n return sent_count\n\n def _disconnect_client(self, client_id: str) -> None:\n \"\"\"\n Disconnect a client, remove from lists, and notify callbacks.\n \"\"\"\
with self._lock:\n if client_id in self._clients:\n try:\n self._clients[client_id].shutdown(socket.SHUT_RDWR) # Signal client to close\n self._clients[client_id].close()\n except Exception as e:\n logger.debug(f\"Error closing socket for {client_id}: {e}\")\n del self._clients[client_id]\n del self._client_addresses[client_id]\n if client_id in self._client_threads:\n # It's a daemon thread, so no need to explicitly join, it will exit.\n del self._client_threads[client_id]\n\n logger.info(f\"Client {client_id} disconnected.\")\n if self._on_client_disconnect:\n self._on_client_disconnect(client_id)\n\n def stop(self) -> None:\n \"\"\"\n Stop the server and gracefully disconnect all clients.\n \"\"\"\
with self._lock:\n if not self._running:\n logger.warning(\"Server is not running.\")\n return\n \n logger.info(\"Stopping server...\")\n self._running = False\n \n # Disconnect all active clients\n for client_id in list(self._clients.keys()): # Iterate over copy of keys\n self._disconnect_client(client_id)\n \n self._cleanup_server()\n logger.info(\"Server stopped successfully.\")\n\n def _cleanup_server(self) -> None:\n \"\"\"\n Close the server socket and clear any remaining resources.\n \"\"\"\
if self._server_socket:\n try:\n self._server_socket.close()\n except Exception as e:\n logger.error(f\"Error closing server socket: {e}\")\n self._server_socket = None\n self._clients.clear()\n self._client_addresses.clear()\n self._client_threads.clear()\n\n def get_client_count(self) -> int:\n \"\"\"\n Returns the number of currently connected clients.\n \"\"\"\
with self._lock:\n return len(self._clients)\n\n def get_client_ids(self) -> List[str]:\n \"\"\"\n Returns a list of unique identifiers for all connected clients.\n \"\"\"\
with self._lock:\n return list(self._clients.keys())\n\n def get_client_address(self, client_id: str) -> Optional[Tuple[str, int]]:\n \"\"\"\n Returns the (host, port) tuple for a given client ID.\n \"\"\"\
with self._lock:\n return self._client_addresses.get(client_id)\n\n\n# Standalone test / demo\nif __name__ == \"__main__\":\n def on_server_client_connect(client_id: str, address: Tuple[str, int]):\n logger.info(f\"DEMO: Server received connection from {client_id} ({address})\")\n server_instance.send_message(client_id, MessageFactory.create_data(\"Welcome to the server!\"))\n\n def on_server_client_disconnect(client_id: str):\n logger.info(f\"DEMO: Server detected disconnection from {client_id}\")\n\n def on_server_message_received(client_id: str, message: RemoteDeskMessage):\n logger.info(f\"DEMO: Server received message from {client_id}: {message.to_dict()}\")\n # Echo back for demo\n server_instance.send_message(client_id, MessageFactory.create_data(f\"Echo: {message.payload.get('message', 'No message')}\"))\n\n server_instance = SocketServer(\n host=\"127.0.0.1\",\n port=DEFAULT_PORT,\n on_client_connect=on_server_client_connect,\n on_client_disconnect=on_server_client_disconnect,\n on_message_received=on_server_message_received,\n )\n\n if server_instance.start():\n logger.info(\"Server demo started. Waiting for clients...\")\n try:\n while True:\n time.sleep(1) # Keep main thread alive\n except KeyboardInterrupt:\n logger.info(\"Server demo interrupted.\")\n finally:\n server_instance.stop()\n logger.info(\"Server demo finished.\")\n else:\n logger.error(\"Failed to start server demo.\")