diff --git a/itests/test.sh b/itests/test.sh index 53c616c6..52dffb05 100644 --- a/itests/test.sh +++ b/itests/test.sh @@ -9,3 +9,4 @@ pytest -s tests #pytest -s tests -k "test_share_single_squeak" #pytest -s tests -k "test_delete_squeak" #pytest -s tests -k "test_get_squeak_by_lookup" +#pytest -s tests -k "test_subscribe_squeaks" diff --git a/itests/tests/test_squeak_node.py b/itests/tests/test_squeak_node.py index d01eca74..b089f147 100644 --- a/itests/tests/test_squeak_node.py +++ b/itests/tests/test_squeak_node.py @@ -17,7 +17,6 @@ from tests.util import delete_profile from tests.util import delete_squeak from tests.util import download_offers from tests.util import download_squeak -from tests.util import download_squeaks from tests.util import get_connected_peer from tests.util import get_connected_peers from tests.util import get_hash @@ -790,15 +789,52 @@ def test_connect_peer(admin_stub, other_admin_stub): print(item) assert len(item) == 0 +# TODO: Re-enable after DownloadSqueaks RPC method supports params. +# def test_get_squeak_by_lookup( +# admin_stub, +# other_admin_stub, +# connected_tcp_peer_id, +# lightning_client, +# signing_profile_id, +# saved_squeak_hash, +# ): +# # Get the squeak profile +# squeak_profile = get_squeak_profile(admin_stub, signing_profile_id) +# squeak_profile_address = squeak_profile.address +# squeak_profile_name = squeak_profile.profile_name -def test_get_squeak_by_lookup( +# # Add the contact profile to the other server and set the profile to be following +# contact_profile_id = create_contact_profile( +# other_admin_stub, squeak_profile_name, squeak_profile_address) +# other_admin_stub.SetSqueakProfileFollowing( +# squeak_admin_pb2.SetSqueakProfileFollowingRequest( +# profile_id=contact_profile_id, +# following=True, +# ) +# ) + +# # Get the squeak display item +# squeak_display_entry = get_squeak_display( +# other_admin_stub, saved_squeak_hash) +# assert squeak_display_entry is None + +# # Sync squeaks +# download_squeaks(other_admin_stub) +# time.sleep(5) + +# # Get the squeak display item +# squeak_display_entry = get_squeak_display( +# other_admin_stub, saved_squeak_hash) +# assert squeak_display_entry.squeak_hash == saved_squeak_hash + + +def test_subscribe_squeaks( admin_stub, other_admin_stub, - connected_tcp_peer_id, - lightning_client, signing_profile_id, saved_squeak_hash, ): + # Get the squeak profile squeak_profile = get_squeak_profile(admin_stub, signing_profile_id) squeak_profile_address = squeak_profile.address @@ -819,11 +855,28 @@ def test_get_squeak_by_lookup( other_admin_stub, saved_squeak_hash) assert squeak_display_entry is None - # Sync squeaks - download_squeaks(other_admin_stub) - time.sleep(5) + with open_peer_connection( + other_admin_stub, + "test_peer", + "squeaknode", + 18777, + ): + time.sleep(2) - # Get the squeak display item - squeak_display_entry = get_squeak_display( - other_admin_stub, saved_squeak_hash) - assert squeak_display_entry.squeak_hash == saved_squeak_hash + # Get the squeak display item + squeak_display_entry = get_squeak_display( + other_admin_stub, saved_squeak_hash) + assert squeak_display_entry is not None + + # Make a new squeak + new_squeak_hash = make_squeak( + admin_stub, + signing_profile_id, + "Hello again!", + ) + time.sleep(2) + + # Get the squeak display item for the new squeak + squeak_display_entry = get_squeak_display( + other_admin_stub, new_squeak_hash) + assert squeak_display_entry is not None diff --git a/requirements-itest.txt b/requirements-itest.txt index 60dcc800..04c90c39 100644 --- a/requirements-itest.txt +++ b/requirements-itest.txt @@ -4,4 +4,4 @@ grpcio grpcio-tools importlib_resources==1.4.0 pytest -squeakpy==0.6.1 +squeakpy==0.6.5 diff --git a/requirements.txt b/requirements.txt index e44897b5..6f98b65c 100644 --- a/requirements.txt +++ b/requirements.txt @@ -13,5 +13,5 @@ protobuf psycopg2 requests SQLAlchemy -squeakpy==0.6.4 +squeakpy==0.6.5 typed-config diff --git a/setup.py b/setup.py index e567ac62..1331929f 100644 --- a/setup.py +++ b/setup.py @@ -82,7 +82,7 @@ setup( include_package_data=True, zip_safe=False, install_requires=[ - 'squeakpy>=0.6.1', + 'squeakpy>=0.6.5', 'importlib_resources', 'argparse', 'googleapis-common-protos', diff --git a/squeaknode/network/connection_manager.py b/squeaknode/network/connection_manager.py index c5944e01..da3c354f 100644 --- a/squeaknode/network/connection_manager.py +++ b/squeaknode/network/connection_manager.py @@ -1,5 +1,9 @@ import logging import threading +from typing import Dict + +from squeaknode.core.peer_address import PeerAddress +from squeaknode.network.peer import Peer MIN_PEERS = 5 @@ -15,7 +19,7 @@ class ConnectionManager(object): """ def __init__(self): - self._peers = {} + self._peers: Dict[PeerAddress, Peer] = {} self.peers_lock = threading.Lock() self.peers_changed_callbacks = {} diff --git a/squeaknode/network/network_manager.py b/squeaknode/network/network_manager.py index 936bc68f..d1b0b220 100644 --- a/squeaknode/network/network_manager.py +++ b/squeaknode/network/network_manager.py @@ -1,5 +1,6 @@ import logging import socket +from typing import List import squeak.params from squeak.messages import MsgSerializable @@ -62,7 +63,7 @@ class NetworkManager(object): def disconnect_peer(self, peer_address: PeerAddress) -> None: self.connection_manager.stop_connection(peer_address) - def get_connected_peer(self, peer_address: PeerAddress): + def get_connected_peer(self, peer_address: PeerAddress) -> List[Peer]: return self.connection_manager.get_peer(peer_address) def get_connected_peers(self): diff --git a/squeaknode/network/peer.py b/squeaknode/network/peer.py index 521a4aba..fef2ba70 100644 --- a/squeaknode/network/peer.py +++ b/squeaknode/network/peer.py @@ -8,6 +8,7 @@ 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 @@ -50,6 +51,8 @@ class Peer(object): self._last_recv_ping_time = None self._recv_msg_queue = queue.Queue() + self._subscription = None + self.handshake_complete = threading.Event() self.ping_started = threading.Event() self.ping_complete = threading.Event() @@ -65,6 +68,10 @@ class Peer(object): def address(self): return self._address + @property + def subscription(self): + return self._subscription + @property def peer_address(self): # TODO: Just return the peer address object @@ -229,6 +236,16 @@ 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_subscription(self, subscription): + self._subscription = subscription + def handle_messages(self, squeak_controller): peer_message_handler = PeerMessageHandler( self, squeak_controller) @@ -246,6 +263,7 @@ class Peer(object): args=(), ).start() self.handshake(squeak_controller) + self.update_subscription(squeak_controller) yield self finally: self.close() diff --git a/squeaknode/network/peer_handler.py b/squeaknode/network/peer_handler.py index a977a7e2..3061f06a 100644 --- a/squeaknode/network/peer_handler.py +++ b/squeaknode/network/peer_handler.py @@ -23,57 +23,3 @@ class PeerHandler(): target=self.handle_connection_fn, args=(self.squeak_controller, peer_socket, address, outgoing,), ).start() - - -# class PeerListener(PeerMessageHandler): -# """Handles receiving messages from a peer. -# """ - -# def __init__(self, peer_message_handler) -> None: -# self.peer_message_handler = peer_message_handler - -# def listen_msgs(self): -# while True: -# try: -# self.peer_message_handler.handle_msgs() -# except Exception as e: -# logger.exception('Error in handle_msgs: {}'.format(e)) -# return - - -# class PeerHandshaker(Connection): -# """Handles the peer handshake. -# """ - -# def __init__(self, peer, connection_manager, peer_server, squeaks_access) -> None: -# super().__init__(peer, connection_manager, peer_server, squeaks_access) - -# def hanshake(self): -# # Initiate handshake with the peer if the connection is outgoing. -# if self.peer.outgoing: -# self.initiate_handshake() - -# # Sleep for 10 seconds. -# time.sleep(10) - -# # Disconnect from peer if handshake is not complete. -# if self.peer.has_handshake_timeout(): -# logger.info('Closing peer because of handshake timeout {}'.format(self.peer)) -# self.peer.close() - - -# class PeerPingChecker(): -# """Handles receiving messages from a peer. -# """ - -# def __init__(self, peer_message_handler) -> None: -# super().__init__() -# self.peer_message_handler = peer_message_handler - -# def handle_msgs(self): -# while True: -# try: -# self.peer_message_handler.handle_msgs() -# except Exception as e: -# logger.exception('Error in handle_msgs: {}'.format(e)) -# return diff --git a/squeaknode/network/peer_message_handler.py b/squeaknode/network/peer_message_handler.py index 90b38c0a..cf802448 100644 --- a/squeaknode/network/peer_message_handler.py +++ b/squeaknode/network/peer_message_handler.py @@ -78,6 +78,8 @@ class PeerMessageHandler: self.handle_notfound(msg) if msg.command == b'offer': self.handle_offer(msg) + if msg.command == b'subscribe': + self.handle_subscribe(msg) def handle_ping(self, msg): nonce = msg.nonce @@ -143,23 +145,7 @@ class PeerMessageHandler: pass def handle_getsqueaks(self, msg): - # TODO: Maybe combine all invs into a single send_msg. - for interest in msg.locator.vInterested: - min_block = interest.nMinBlockHeight if interest.nMinBlockHeight != -1 else None - max_block = interest.nMaxBlockHeight if interest.nMaxBlockHeight != -1 else None - reply_to_hash = interest.hashReplySqk if interest.hashReplySqk != EMPTY_HASH else None - squeak_hashes = self.squeak_controller.lookup_squeaks_for_interest( - address=[str(address) for address in interest.addresses], - min_block=min_block, - max_block=max_block, - reply_to_hash=reply_to_hash, - ) - invs = [ - CInv(type=1, hash=squeak_hash) - for squeak_hash in squeak_hashes] - if invs: - inv_msg = msg_inv(inv=invs) - self.peer.send_msg(inv_msg) + self._send_reply_invs(msg.locator) def handle_squeak(self, msg): squeak = msg.squeak @@ -183,3 +169,27 @@ class PeerMessageHandler: peer_address=self.peer.peer_address, ) self.squeak_controller.save_offer(decoded_offer) + + def handle_subscribe(self, msg): + logger.info("Received subscribe msg: {}".format(msg)) + self._send_reply_invs(msg.locator) + self.peer.set_subscription(msg) + + def _send_reply_invs(self, locator): + # TODO: Maybe combine all invs into a single send_msg. + for interest in locator.vInterested: + min_block = interest.nMinBlockHeight if interest.nMinBlockHeight != -1 else None + max_block = interest.nMaxBlockHeight if interest.nMaxBlockHeight != -1 else None + reply_to_hash = interest.hashReplySqk if interest.hashReplySqk != EMPTY_HASH else None + squeak_hashes = self.squeak_controller.lookup_squeaks_for_interest( + address=[str(address) for address in interest.addresses], + min_block=min_block, + max_block=max_block, + reply_to_hash=reply_to_hash, + ) + invs = [ + CInv(type=1, hash=squeak_hash) + for squeak_hash in squeak_hashes] + if invs: + inv_msg = msg_inv(inv=invs) + self.peer.send_msg(inv_msg) diff --git a/squeaknode/node/new_squeak_listener.py b/squeaknode/node/new_squeak_listener.py new file mode 100644 index 00000000..4742be2a --- /dev/null +++ b/squeaknode/node/new_squeak_listener.py @@ -0,0 +1,84 @@ +import logging +import queue +import threading +import uuid +from contextlib import contextmanager + +logger = logging.getLogger(__name__) + +DEFAULT_MAX_QUEUE_SIZE = 1000 +DEFAULT_UPDATE_INTERVAL_S = 1 + + +class NewSqueakListener: + + def __init__(self): + self.callbacks = {} + + def handle_new_squeak(self, squeak): + # logger.info("Handling new squeak: {!r}".format( + # get_hash(squeak).hex(), + # )) + for callback in self.callbacks.values(): + callback(squeak) + + def add_callback(self, name, callback): + self.callbacks[name] = callback + + def remove_callback(self, name): + del self.callbacks[name] + + +class NewSqueakSubscriptionClient: + + def __init__( + self, + new_squeak_listener: NewSqueakListener, + stopped: threading.Event, + max_queue_size=DEFAULT_MAX_QUEUE_SIZE, + ): + self.new_squeak_listener = new_squeak_listener + self.stopped = stopped + self.q: queue.Queue = queue.Queue(max_queue_size) + + @contextmanager + def open_subscription(self): + threading.Thread( + target=self.wait_for_stopped, + ).start() + # Register the callback to populate the queue + callback_name = "new_squeak_callback_{}".format(uuid.uuid1()), + try: + self.new_squeak_listener.add_callback( + name=callback_name, + callback=self.enqueue_squeak, + ) + logger.info("Before yielding new squeaks client...") + yield self + logger.info("After yielding new squeaks client...") + finally: + logger.info("Stopping new squeaks client...") + self.new_squeak_listener.remove_callback( + name=callback_name, + ) + logger.info("Stopped new squeaks client...") + + def enqueue_squeak(self, squeak): + self.q.put(squeak) + + def wait_for_stopped(self): + self.stopped.wait() + self.q.put(None) + + def get_squeak(self): + while True: + item = self.q.get() + if item is None: + logger.debug("Poison pill swallowed.") + return + yield item + self.q.task_done() + logger.info( + "Removed item from queue. Size: {}".format( + self.q.qsize()) + ) diff --git a/squeaknode/node/new_squeak_worker.py b/squeaknode/node/new_squeak_worker.py new file mode 100644 index 00000000..ba6becc7 --- /dev/null +++ b/squeaknode/node/new_squeak_worker.py @@ -0,0 +1,102 @@ +import logging +import threading + +from squeak.core import CSqueak +from squeak.messages import msg_inv +from squeak.net import CInterested +from squeak.net import CInv + +from squeaknode.core.util import get_hash +from squeaknode.network.network_manager import NetworkManager +from squeaknode.network.peer import Peer +from squeaknode.node.squeak_controller import SqueakController + + +logger = logging.getLogger(__name__) + +DEFAULT_MAX_QUEUE_SIZE = 1000 +DEFAULT_UPDATE_INTERVAL_S = 1 + +HASH_LENGTH = 32 +EMPTY_HASH = b'\x00' * HASH_LENGTH + + +class NewSqueakWorker: + + def __init__(self, + squeak_controller: SqueakController, + network_manager: NetworkManager, + ): + self.squeak_controller = squeak_controller + self.network_manager = network_manager + self.stopped = threading.Event() + + def start_running(self): + threading.Thread( + target=self.handle_new_squeaks, + name="new_squeaks_worker_thread", + ).start() + + def stop_running(self): + self.stopped.set() + + def handle_new_squeaks(self): + logger.debug("Starting NewSqueakWorker...") + for squeak in self.squeak_controller.subscribe_new_squeaks( + self.stopped, + ): + logger.debug("Handling new squeak: {!r}".format( + get_hash(squeak).hex(), + )) + self.forward_squeak(squeak) + + def forward_squeak(self, squeak): + logger.debug("Forward new squeak: {!r}".format( + get_hash(squeak).hex(), + )) + for peer in self.network_manager.get_connected_peers(): + if self.should_forward(squeak, peer): + logger.debug("Forwarding to peer: {}".format( + peer, + )) + squeak_hash = get_hash(squeak) + inv = CInv(type=1, hash=squeak_hash) + inv_msg = msg_inv(inv=[inv]) + peer.send_msg(inv_msg) + logger.debug("Finished checking peers to forward.") + + def should_forward(self, squeak: CSqueak, peer: Peer) -> bool: + logger.debug(""" + Checking if should forward for peer: {} + with subscription: {} + and squeak: {} + with squeak address: {} + """.format( + peer, + peer.subscription, + squeak, + str(squeak.GetAddress()), + )) + if peer.subscription is None: + return False + locator = peer.subscription.locator + for interest in locator.vInterested: + if self.squeak_matches_interest(squeak, interest): + logger.debug("Found a match!") + return True + return False + + def squeak_matches_interest(self, squeak: CSqueak, interest: CInterested) -> bool: + if len(interest.addresses) > 0 \ + and squeak.GetAddress() not in interest.addresses: + return False + # if interest.nMinBlockHeight != -1 \ + # and squeak.nBlockHeight < interest.nMinBlockHeight: + # return False + # if interest.nMaxBlockHeight != -1 \ + # and squeak.nBlockHeight > interest.nMaxBlockHeight: + # return False + if interest.hashReplySqk != EMPTY_HASH \ + and squeak.hashReplySqk != interest.hashReplySqk: + return False + return True diff --git a/squeaknode/node/squeak_controller.py b/squeaknode/node/squeak_controller.py index 56317e75..c16d8e74 100644 --- a/squeaknode/node/squeak_controller.py +++ b/squeaknode/node/squeak_controller.py @@ -11,6 +11,7 @@ from squeak.core.signing import CSigningKey from squeak.core.signing import CSqueakAddress from squeak.messages import msg_getdata from squeak.messages import msg_getsqueaks +from squeak.messages import msg_subscribe from squeak.messages import MsgSerializable from squeak.net import CInterested from squeak.net import CInv @@ -29,6 +30,8 @@ from squeaknode.core.squeak_peer import SqueakPeer from squeaknode.core.squeak_profile import SqueakProfile from squeaknode.core.util import get_hash from squeaknode.core.util import is_address_valid +from squeaknode.node.new_squeak_listener import NewSqueakListener +from squeaknode.node.new_squeak_listener import NewSqueakSubscriptionClient from squeaknode.node.received_payments_subscription_client import ReceivedPaymentsSubscriptionClient @@ -51,6 +54,7 @@ class SqueakController: self.squeak_rate_limiter = squeak_rate_limiter self.payment_processor = payment_processor self.network_manager = network_manager + self.new_squeak_listener = NewSqueakListener() self.config = config def save_squeak( @@ -79,6 +83,7 @@ class SqueakController: inserted_squeak_hash, decryption_key, ) + self.new_squeak_listener.handle_new_squeak(squeak) # Return the squeak hash. return inserted_squeak_hash @@ -136,7 +141,7 @@ class SqueakController: profile_name=profile_name, private_key=signing_key_bytes, address=str(address), - following=True, + following=False, profile_image=None, ) return self.squeak_db.insert_profile(squeak_profile) @@ -173,7 +178,7 @@ class SqueakController: profile_name=profile_name, private_key=None, address=squeak_address, - following=True, + following=False, profile_image=None, ) return self.squeak_db.insert_profile(squeak_profile) @@ -195,12 +200,14 @@ class SqueakController: def set_squeak_profile_following(self, profile_id: int, following: bool) -> None: self.squeak_db.set_profile_following(profile_id, following) + self.update_subscriptions() def rename_squeak_profile(self, profile_id: int, profile_name: str) -> None: self.squeak_db.set_profile_name(profile_id, profile_name) def delete_squeak_profile(self, profile_id: int) -> None: self.squeak_db.delete_profile(profile_id) + self.update_subscriptions() def set_squeak_profile_image(self, profile_id: int, profile_image: bytes) -> None: self.squeak_db.set_profile_image(profile_id, profile_image) @@ -516,12 +523,9 @@ class SqueakController: # ) return ret - def sync_timeline(self): + def get_interested_locator(self): block_range = self.get_block_range() - logger.info("Syncing timeline with block range: {}".format(block_range)) followed_addresses = self.get_followed_addresses() - logger.debug("Syncing timeline with followed addresses: {}".format( - followed_addresses)) interests = [ CInterested( addresses=[CSqueakAddress(address) @@ -530,14 +534,15 @@ class SqueakController: nMaxBlockHeight=block_range.max_block, ) ] - locator = CSqueakLocator( + return CSqueakLocator( vInterested=interests, ) + + def sync_timeline(self): + locator = self.get_interested_locator() getsqueaks_msg = msg_getsqueaks( locator=locator, ) - # for peer in self.connection_manager.peers: - # peer.send_msg(getsqueaks_msg) self.broadcast_msg(getsqueaks_msg) def download_single_squeak(self, squeak_hash: bytes): @@ -619,3 +624,20 @@ class SqueakController: # for payment in client.get_received_payments(): # yield payment return self.network_manager.subscribe_connected_peers(stopped) + + def subscribe_new_squeaks(self, stopped: threading.Event): + subscription_client = NewSqueakSubscriptionClient( + self.new_squeak_listener, + stopped, + ) + with subscription_client.open_subscription(): + for result in subscription_client.get_squeak(): + yield result + + def update_subscriptions(self): + locator = self.get_interested_locator() + for peer in self.network_manager.get_connected_peers(): + subscribe_msg = msg_subscribe( + locator=locator, + ) + peer.send_msg(subscribe_msg) diff --git a/squeaknode/node/squeak_node.py b/squeaknode/node/squeak_node.py index d2280048..510dff71 100644 --- a/squeaknode/node/squeak_node.py +++ b/squeaknode/node/squeak_node.py @@ -13,6 +13,7 @@ from squeaknode.db.db_engine import get_engine from squeaknode.db.squeak_db import SqueakDb from squeaknode.lightning.lnd_lightning_client import LNDLightningClient from squeaknode.network.network_manager import NetworkManager +from squeaknode.node.new_squeak_worker import NewSqueakWorker from squeaknode.node.payment_processor import PaymentProcessor from squeaknode.node.peer_connection_worker import PeerConnectionWorker from squeaknode.node.process_received_payments_worker import ProcessReceivedPaymentsWorker @@ -53,6 +54,7 @@ class SqueakNode: self.initialize_peer_sync_worker() self.initialize_squeak_deletion_worker() self.initialize_offer_expiry_worker() + self.initialize_new_squeak_worker() def start_running(self): self._initialize() @@ -64,16 +66,19 @@ class SqueakNode: self.admin_web_server.start() self.received_payment_processor_worker.start_running() self.peer_connection_worker.start() - if self.config.sync.enabled: - self.peer_sync_worker.start() + # TODO: Delete peer_sync_worker, subscribe replaces it + # if self.config.sync.enabled: + # self.peer_sync_worker.start() self.squeak_deletion_worker.start() self.offer_expiry_worker.start() + self.new_squeak_worker.start_running() def stop_running(self): self.admin_web_server.stop() self.admin_rpc_server.stop() self.network_manager.stop() self.received_payment_processor_worker.stop_running() + self.new_squeak_worker.stop_running() def initialize_network(self): # load the network @@ -196,3 +201,9 @@ class SqueakNode: self.squeak_controller, self.config.core.offer_deletion_interval_s, ) + + def initialize_new_squeak_worker(self): + self.new_squeak_worker = NewSqueakWorker( + self.squeak_controller, + self.network_manager, + )