From 164a3e9637d53e8100c592f9eb38df708f7d9786 Mon Sep 17 00:00:00 2001 From: Jonathan Zernik Date: Thu, 26 Aug 2021 23:55:36 -0700 Subject: [PATCH] Add connection class (#1071) * Add connection class and got itest working * Remove open_connection context manager from peer class * Remove old methods from peer class * Handle initial sync in connection class * Remove old comments from peer class * Removed stop event from peer class * Handle received messages in connection class * Remove old comment from peer class --- squeaknode/network/connection.py | 111 ++++++++++++++++++ squeaknode/network/network_manager.py | 30 +++-- squeaknode/network/peer.py | 127 +++------------------ squeaknode/network/peer_message_handler.py | 1 + 4 files changed, 142 insertions(+), 127 deletions(-) create mode 100644 squeaknode/network/connection.py diff --git a/squeaknode/network/connection.py b/squeaknode/network/connection.py new file mode 100644 index 00000000..b935b377 --- /dev/null +++ b/squeaknode/network/connection.py @@ -0,0 +1,111 @@ +import logging +import threading +from contextlib import contextmanager + +from squeak.messages import msg_subscribe + +from squeaknode.network.peer import Peer +from squeaknode.network.peer_message_handler import PeerMessageHandler + +HANDSHAKE_TIMEOUT = 5 + + +logger = logging.getLogger(__name__) + + +class Connection(object): + """Handles lifecycle of a connection to a peer. + """ + + def __init__(self, peer: Peer, squeak_controller): + self.peer = peer + self.squeak_controller = squeak_controller + + @contextmanager + def connect(self, connection_manager): + self.start_receiving_msgs() + self.handshake() + logger.info("Adding peer.") + connection_manager.add_peer(self.peer) + try: + logger.info("Yielding peer.") + yield self + except Exception: + logger.exception("Peer connection failed.") + finally: + logger.info("Removing peer.") + connection_manager.remove_peer(self.peer) + + def handle_connection(self): + self.initial_sync() + self.handle_messages() + + def initial_sync(self): + # TODO: getaddrs from peer. + self.update_subscription() + + def update_subscription(self): + locator = self.squeak_controller.get_interested_locator() + subscribe_msg = msg_subscribe( + locator=locator, + ) + self.peer.send_msg(subscribe_msg) + + def start_receiving_msgs(self): + threading.Thread( + target=self.peer.recv_msgs, + args=(), + ).start() + + def handshake(self): + timer = HandshakeTimer( + self.peer.stop, + str(self), + ) + timer.start_timer() + + if self.peer.outgoing: + self.peer.send_version() + self.peer.receive_version() + if not self.peer.outgoing: + self.peer.send_version() + + self.peer.set_connected() + logger.info("HANDSHAKE COMPLETE-----------") + timer.stop_timer() + + def handle_messages(self): + peer_message_handler = PeerMessageHandler( + self.peer, + self.squeak_controller, + ) + peer_message_handler.handle_msgs() + + +class HandshakeTimer: + """Stop the peer if handshake is not complete before timeout. + """ + + def __init__(self, + stop_fn, + peer_name, + ): + self.stop_fn = stop_fn + self.peer_name = peer_name + self.timer = None + + def start_timer(self): + self.timer = threading.Timer( + HANDSHAKE_TIMEOUT, + self.stop_peer, + ) + self.timer.name = "handshake_timere_thread_{}".format(self.peer_name) + self.timer.start() + + def stop_timer(self): + logger.debug("Canceling handshake timer.") + self.timer.cancel() + + def stop_peer(self): + logger.info("Closing peer from handshake timer.") + self.stop_fn() diff --git a/squeaknode/network/network_manager.py b/squeaknode/network/network_manager.py index 8b7de494..bc61d840 100644 --- a/squeaknode/network/network_manager.py +++ b/squeaknode/network/network_manager.py @@ -7,6 +7,7 @@ from squeak.messages import MsgSerializable from squeaknode.core.peer_address import PeerAddress from squeaknode.network.connected_peers_subscription_client import ConnectedPeersSubscriptionClient +from squeaknode.network.connection import Connection from squeaknode.network.connection_manager import ConnectionManager from squeaknode.network.peer import Peer from squeaknode.network.peer_client import PeerClient @@ -91,23 +92,20 @@ class NetworkManager(object): This method blocks until the peer connection has stopped. """ + peer = Peer( + peer_socket, + self.local_address, + address, + outgoing, + ) + logger.debug( - 'Setting up controller for peer address {} ...'.format(address)) - with Peer( - peer_socket, - self.local_address, - address, - outgoing, - ).open_connection(squeak_controller) as peer: - self.connection_manager.add_peer(peer) - try: - peer.sync(squeak_controller) - peer.handle_messages(squeak_controller) - except Exception: - logger.exception("Peer connection failed.") - finally: - self.connection_manager.remove_peer(peer) - logger.debug('Stopped controller for peer address {}.'.format(address)) + 'Setting up connection for peer address {} ...'.format(address)) + with Connection(peer, squeak_controller).connect( + self.connection_manager + ) as connection: + connection.handle_connection() + logger.debug('Stopped connection for peer address {}.'.format(address)) @property def local_address(self) -> PeerAddress: diff --git a/squeaknode/network/peer.py b/squeaknode/network/peer.py index 71e36916..15a38a2a 100644 --- a/squeaknode/network/peer.py +++ b/squeaknode/network/peer.py @@ -3,19 +3,16 @@ import queue import socket import threading import time -from contextlib import contextmanager from io import BytesIO from bitcoin.core.serialize import SerializationTruncationError from bitcoin.net import CAddress -from squeak.messages import msg_subscribe from squeak.messages import msg_verack from squeak.messages import msg_version from squeak.messages import MsgSerializable from squeaknode.core.peer_address import PeerAddress from squeaknode.core.util import generate_version_nonce -from squeaknode.network.peer_message_handler import PeerMessageHandler from squeaknode.network.util import time_now @@ -25,7 +22,6 @@ LAST_MESSAGE_TIMEOUT = 600 PING_TIMEOUT = 10 PING_INTERVAL = 60 -HANDSHAKE_TIMEOUT = 5 UPDATE_TIME_INTERVAL = 10 HANDSHAKE_VERSION = 70002 @@ -60,10 +56,10 @@ class Peer(object): self._subscription = None - self.handshake_complete = threading.Event() - self.ping_started = threading.Event() - self.ping_complete = threading.Event() - self.stopped = threading.Event() + self.msg_receiver = MessageReceiver( + self._peer_socket, + self._recv_msg_queue, + ) @property def nVersion(self): @@ -123,14 +119,6 @@ class Peer(object): def remote_version(self, remote_version): self._remote_version = remote_version - @property - def is_handshake_complete(self): - return self.handshake_complete.is_set() - - @property - def is_open(self): - return not self.stopped.is_set() - @property def last_msg_revc_time(self): return self._last_msg_revc_time @@ -161,13 +149,8 @@ class Peer(object): logger.info('Received msg {} from {}'.format(msg, self)) return msg - def start(self): - msg_receiver = MessageReceiver( - self._peer_socket, self._recv_msg_queue, self.stopped) - threading.Thread( - target=msg_receiver.recv_msgs, - args=(), - ).start() + def recv_msgs(self): + self.msg_receiver.recv_msgs() def stop(self): logger.info("Stopping peer socket: {}".format(self._peer_socket)) @@ -187,21 +170,15 @@ class Peer(object): logger.info('Failed to send msg to {}'.format(self)) self.stop() - def handshake(self, squeak_controller): - timer = HandshakeTimer( - self.stop, - str(self), - ) - timer.start_timer() - - if self.outgoing: - local_version = self.version_pkt(squeak_controller) - self.local_version = local_version - self.send_msg(local_version) - verack = self.recv_msg() - if not isinstance(verack, msg_verack): - raise Exception('Wrong message type for verack response.') + def send_version(self): + local_version = self.version_pkt() + self.local_version = local_version + self.send_msg(local_version) + verack = self.recv_msg() + if not isinstance(verack, msg_verack): + raise Exception('Wrong message type for verack response.') + def receive_version(self): remote_version = self.recv_msg() if not isinstance(remote_version, msg_version): raise Exception('Wrong message type for version message.') @@ -209,18 +186,7 @@ class Peer(object): verack = msg_verack() self.send_msg(verack) - if not self.outgoing: - local_version = self.version_pkt(squeak_controller) - self.local_version = local_version - self.send_msg(local_version) - verack = self.recv_msg() - if not isinstance(verack, msg_verack): - raise Exception('Wrong message type for verack response.') - - logger.info("HANDSHAKE COMPLETE-----------") - timer.stop_timer() - - def version_pkt(self, squeak_controller): + def version_pkt(self): """Get the version message for this peer.""" msg = msg_version() msg.nVersion = HANDSHAKE_VERSION @@ -229,40 +195,12 @@ class Peer(object): msg.nNonce = generate_version_nonce() return msg - def update_subscription(self, squeak_controller): - locator = squeak_controller.get_interested_locator() - subscribe_msg = msg_subscribe( - locator=locator, - ) - self.send_msg(subscribe_msg) - def set_connected(self): self._connect_time = time_now() def set_subscription(self, subscription): self._subscription = subscription - def handle_messages(self, squeak_controller): - peer_message_handler = PeerMessageHandler( - self, squeak_controller) - peer_message_handler.handle_msgs() - - def sync(self, squeak_controller): - # TODO: getaddrs from peer. - self.update_subscription(squeak_controller) - - @contextmanager - def open_connection(self, squeak_controller): - logger.debug('Setting up peer {} ...'.format(self)) - try: - self.start() - self.handshake(squeak_controller) - self.set_connected() - yield self - finally: - self.stop() - logger.debug('Stopped connection to peer {} ...'.format(self)) - def __repr__(self): return "Peer(%s)" % (str(self.remote_address)) @@ -301,10 +239,9 @@ class MessageReceiver: """Reads bytes from the socket and puts messages in the receive queue. """ - def __init__(self, socket, queue, stopped_event): + def __init__(self, socket, queue): self.socket = socket self.queue = queue - self.stopped_event = stopped_event self.decoder = MessageDecoder() def _recv_msgs(self): @@ -322,41 +259,9 @@ class MessageReceiver: for msg in self.decoder.process_recv_data(recv_data): self.queue.put(msg) - if self.stopped_event.is_set(): - return def recv_msgs(self): try: self._recv_msgs() except Exception: logger.info('Failed to receive msg from {}'.format(self)) - self.stopped_event.set() - - -class HandshakeTimer: - """Stop the peer if handshake is not complete before timeout. - """ - - def __init__(self, - stop_fn, - peer_name, - ): - self.stop_fn = stop_fn - self.peer_name = peer_name - self.timer = None - - def start_timer(self): - self.timer = threading.Timer( - HANDSHAKE_TIMEOUT, - self.stop_peer, - ) - self.timer.name = "handshake_timere_thread_{}".format(self.peer_name) - self.timer.start() - - def stop_timer(self): - logger.debug("Canceling handshake timer.") - self.timer.cancel() - - def stop_peer(self): - logger.info("Closing peer from handshake timer.") - self.stop_fn() diff --git a/squeaknode/network/peer_message_handler.py b/squeaknode/network/peer_message_handler.py index b3b54e68..82d4785a 100644 --- a/squeaknode/network/peer_message_handler.py +++ b/squeaknode/network/peer_message_handler.py @@ -39,6 +39,7 @@ class PeerMessageHandler: while msg is not None: self.handle_peer_message(msg) msg = self.peer.recv_msg() + logger.info("Finished handle_msgs") def handle_peer_message(self, msg): """Handle messages from a peer with completed handshake."""