Add subscribe msg (#1016)

* Added subscribe message and send after connection handshake

* Added new squeak worker

* Got itest working for subscribe squeaks peer message

* Update peer squeak subscriptions on follow profile changed.

* Remove sync squeaks worker
This commit is contained in:
Jonathan Zernik 2021-08-22 01:46:05 -07:00 committed by GitHub
parent 720a8a16f5
commit 9023131112
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
14 changed files with 350 additions and 98 deletions

View file

@ -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"

View file

@ -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

View file

@ -4,4 +4,4 @@ grpcio
grpcio-tools
importlib_resources==1.4.0
pytest
squeakpy==0.6.1
squeakpy==0.6.5

View file

@ -13,5 +13,5 @@ protobuf
psycopg2
requests
SQLAlchemy
squeakpy==0.6.4
squeakpy==0.6.5
typed-config

View file

@ -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',

View file

@ -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 = {}

View file

@ -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):

View file

@ -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()

View file

@ -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

View file

@ -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)

View file

@ -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())
)

View file

@ -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

View file

@ -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)

View file

@ -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,
)