diff --git a/pytoniq/adnl/adnl.py b/pytoniq/adnl/adnl.py index 9ab00b7..39b8dc9 100644 --- a/pytoniq/adnl/adnl.py +++ b/pytoniq/adnl/adnl.py @@ -11,7 +11,8 @@ from pytoniq_core.tl.generator import TlGenerator -from pytoniq_core.crypto.ciphers import Server, Client, AdnlChannel, get_random, aes_ctr_encrypt, aes_ctr_decrypt, get_shared_key, create_aes_ctr_sipher_from_key_n_data +from pytoniq_core.crypto.ciphers import Server, Client, AdnlChannel, get_random, aes_ctr_encrypt, aes_ctr_decrypt, \ + get_shared_key, create_aes_ctr_sipher_from_key_n_data class SocketProtocol(asyncio.DatagramProtocol): @@ -19,7 +20,7 @@ class SocketProtocol(asyncio.DatagramProtocol): def __init__(self, timeout: int = 10): # https://github.com/eerimoq/asyncudp/blob/main/asyncudp/__init__.py self._error = None - self._packets = asyncio.Queue(10000) + self._packets = asyncio.Queue(500000) self.timeout = timeout self.logger = logging.getLogger(self.__class__.__name__) @@ -27,12 +28,15 @@ def connection_made(self, transport: transports.DatagramTransport) -> None: super().connection_made(transport) def datagram_received(self, data: bytes, addr: typing.Tuple[typing.Union[str, Any], int]) -> None: - self.logger.debug(f'received {len(data)} bytes') - self._packets.put_nowait((data, addr)) + self.logger.debug(f'received {len(data)} bytes from {addr}; queue {self._packets.qsize()}') + try: + self._packets.put_nowait((data, addr)) + except asyncio.QueueFull: + self.logger.warning('Queue is full, dropping packet') super().datagram_received(data, addr) def error_received(self, exc: Exception) -> None: - raise exc + self.logger.warning(f'error received: {exc}') super().error_received(exc) async def receive(self): @@ -41,6 +45,8 @@ async def receive(self): class Node(Server): + PING_INTERVAL = 60 + def __init__( self, peer_host: str, # ipv4 host @@ -59,6 +65,15 @@ def __init__( self.pinger: asyncio.Task = None self.connected = False self.logger = logging.getLogger(self.__class__.__name__) + self._last_ping_at = time.time() + self._lost_pings = 0 + + @property + def lost_pings(self): + return self._lost_pings + + def reset_pings(self): + self._last_ping_at = time.time() self._lost_pings = 0 async def connect(self): @@ -66,7 +81,8 @@ async def connect(self): async def send_ping(self) -> None: random_id = get_random(8) - resp = await self.transport.send_query_message(tl_schema_name='dht.ping', data={'random_id': random_id}, peer=self) + resp = await self.transport.send_query_message(tl_schema_name='dht.ping', data={'random_id': random_id}, + peer=self) assert resp[0].get('random_id') == int.from_bytes(random_id, 'big', signed=True) def start_ping(self): @@ -76,15 +92,15 @@ async def ping(self): while True: try: await self.send_ping() - self._lost_pings = 0 + self.reset_pings() self.logger.debug(f'pinged {self.key_id.hex()}') except asyncio.TimeoutError: self._lost_pings += 1 - if self._lost_pings > 3: + if self._lost_pings > 3 and self._last_ping_at < time.time() - 15: if self.key_id in self.transport.peers: self.transport.peers.pop(self.key_id) await self.disconnect() - await asyncio.sleep(5) + await asyncio.sleep(self.PING_INTERVAL) async def get_signed_address_list(self): return (await self.transport.send_query_message('dht.getSignedAddressList', {}, self))[0] @@ -101,8 +117,16 @@ def inc_seqno(self): async def disconnect(self): if self.connected: + self.logger.debug(f'disconnected {self.key_id.hex()}') self.connected = False self.pinger.cancel() + self.transport.peers.pop(self.key_id, None) + for ch in self.channels: + self.transport.channels.pop(ch.server_aes_key_id, None) + self.channels = [] + self.seqno = 1 + self.confirm_seqno = 0 + self.reset_pings() class AdnlTransportError(Exception): @@ -130,6 +154,8 @@ def __init__(self, self.query_handlers: typing.Dict[str, typing.Callable] = {} self.custom_handlers: typing.Dict[str, typing.Callable] = {} self._message_parts: typing.Dict[str, dict] = {} # {'hash': {'remained': int, 'parts': list}} + self.inited = False + self.pending_channels = {} """########### connection ###########""" self.transport: asyncio.DatagramTransport = None @@ -222,8 +248,7 @@ def _decrypt_any(self, resp_packet: bytes) -> typing.Tuple[bytes, typing.Optiona :param resp_packet: bytes of received packet :return: decrypted packet and maybe `Node` """ - key_id = resp_packet[:32] - if key_id == self.client.get_key_id(): + if resp_packet.startswith(self.local_id): server_public_key = resp_packet[32:64] checksum = resp_packet[64:96] encrypted = resp_packet[96:] @@ -238,14 +263,15 @@ def _decrypt_any(self, resp_packet: bytes) -> typing.Tuple[bytes, typing.Optiona assert hashlib.sha256(decrypted).digest() == checksum, 'invalid checksum' return decrypted, None else: - for peer_id, channel in self.channels.items(): - if key_id == channel.server_aes_key_id: - checksum = resp_packet[32:64] - encrypted = resp_packet[64:] - decrypted = channel.decrypt(encrypted, checksum) - assert hashlib.sha256(decrypted).digest() == checksum, 'invalid checksum' - return decrypted, self.peers.get(peer_id) - # TODO make new connection + key_id = resp_packet[:32] + channel = self.channels.get(key_id) + if channel: + peer_id = channel.peer_id + checksum = resp_packet[32:64] + encrypted = resp_packet[64:] + decrypted = channel.decrypt(encrypted, checksum) + assert hashlib.sha256(decrypted).digest() == checksum, 'invalid checksum' + return decrypted, self.peers.get(peer_id) self.logger.debug(f'unknown key id from node: {key_id.hex()}') return b'', None @@ -253,11 +279,14 @@ def _process_outcoming_message(self, message: dict) -> typing.Optional[asyncio.F future = self.loop.create_future() type_ = message['@type'] if type_ == 'adnl.message.query': - self.tasks[message.get('query_id')[::-1].hex()] = future + id_ = message.get('query_id')[::-1].hex() + self.tasks[id_] = future elif type_ == 'adnl.message.createChannel': - self.tasks[message.get('key')] = future + id_ = message.get('key') + self.tasks[id_] = future else: return + future.id = id_ return future def _create_futures(self, data: dict) -> typing.List[asyncio.Future]: @@ -280,25 +309,27 @@ async def _receive(futures: typing.List[asyncio.Future]) -> list: async def _process_incoming_message(self, message: dict, peer: Node): if peer: - self.logger.debug(f'Received message {message} from peer {peer.get_key_id().hex()}') + self.logger.debug(f'Received message {message} from peer {peer.key_id.hex()}') if message['@type'] == 'adnl.message.answer': - future = self.tasks.pop(message.get('query_id')) - future.set_result(message['answer']) + future = self.tasks.pop(message.get('query_id'), None) + if future and not future.done(): + future.set_result(message['answer']) elif message['@type'] == 'adnl.message.confirmChannel': - if message.get('peer_key') in self.tasks: - future = self.tasks.pop(message.get('peer_key')) - future.set_result(message) + self._process_confirm_channel(message, peer) + elif message['@type'] == 'adnl.message.createChannel': + await self._process_create_channel(message, peer) elif message['@type'] == 'adnl.message.query': if peer is None: - self.logger.info(f'Received query message from unknown peer: {message}') - # not implemented, todo: make connection with new peer + self.logger.debug(f'Received query message from unknown peer: {message}') return + peer.reset_pings() await self._process_query_message(message, peer) elif message['@type'] == 'adnl.message.custom': if peer is None: # should not ever happen fixme - self.logger.info(f'Received custom message from unknown peer: {message}') + self.logger.debug(f'Received custom message from unknown peer: {message}') return + peer.reset_pings() await self._process_custom_message(message, peer) elif message['@type'] == 'adnl.message.part': hash_ = message['hash'] @@ -309,13 +340,74 @@ async def _process_incoming_message(self, message: dict, peer: Node): self._message_parts[hash_]['parts'].append(message) if self._message_parts[hash_]['remained'] == 0: - data = self._collect_adnl_message_parts(hash_) + try: + data = self._collect_adnl_message_parts(hash_) + except: + return if isinstance(data, dict) and data['@type'] != 'adnl.message.part': # to avoid infinity recursion, but should never happen await self._process_incoming_message(data, peer) else: - self.logger.info(f'unexpected message type received: {message}') + self.logger.debug(f'unexpected message type received: {message}') # raise AdnlTransportError(f'unexpected message type received: {message}') + def _store_new_channel(self, channel_client: Client, key: str, peer: Node): + channel_peer = Server(peer.host, peer.port, bytes.fromhex(key)) + channel = AdnlChannel(channel_client, channel_peer, self.local_id, peer.key_id) + channel.peer_id = peer.key_id + self.channels[channel.server_aes_key_id] = channel + peer.channels.append(channel) + + def _process_confirm_channel(self, message: dict, peer: Node): + if message.get('peer_key') in self.tasks: + future = self.tasks.pop(message.get('peer_key')) + if not future.done(): + future.set_result(message) + if peer.key_id not in self.pending_channels: + return + + channel_client = self.pending_channels.get(peer.key_id) # add channel to the object from connect_to_peer + self._store_new_channel(channel_client, message['key'], peer) + + async def _process_create_channel(self, message: dict, peer: Node): + if peer.key_id in self.peers: + return # drop packet since peer is already connected + key = message.get('key') + if key is None: + return + channel_client = Client(Client.generate_ed25519_private_key()) + + ts = int(time.time()) + + confirm_channel_message = { + '@type': 'adnl.message.confirmChannel', + 'peer_key': key, + 'key': channel_client.ed25519_public.encode().hex(), + 'date': ts + } + + data = { + 'from_short': {'id': self.local_id.hex()}, + 'message': confirm_channel_message, + 'address': { + 'addrs': [], + 'version': ts, + 'reinit_date': ts, + 'priority': 0, + 'expire_at': 0, + }, + 'recv_addr_list_version': ts, + 'reinit_date': ts, + 'dst_reinit_date': 0, + } + + self._store_new_channel(channel_client, key, peer) + + peer.connected = True + self.peers[peer.key_id] = peer + + await self.send_message_outside_channel(data, peer) + peer.start_ping() + def _collect_adnl_message_parts(self, hash_: str, deserialize_after: bool = True): if hash_ not in self._message_parts: raise AdnlTransportError(f'Provided hash not in message parts') @@ -390,35 +482,63 @@ def set_default_custom_message_handler(self, handler: typing.Callable): """ self.set_custom_message_handler(None, handler) - async def listen(self): - while True: - packet, addr = await self.protocol.receive() + async def process_packet(self, packet_data: bytes, addr: tuple): - decrypted, peer = self._decrypt_any(packet) + try: + decrypted, peer = self._decrypt_any(packet_data) if not decrypted: - continue - response = self.schemas.deserialize(decrypted)[0] - - if peer is None: - if 'from_short' in response: - peer = self.peers.get(bytes.fromhex(response['from_short']['id'])) + return + packet, _ = self.schemas.deserialize(decrypted) + if not isinstance(packet, dict): # must be deserialized + return + except: + return - if peer is not None: - received_seqno = response.get('seqno', 0) - if received_seqno > peer.confirm_seqno: - peer.confirm_seqno = received_seqno + if peer is None: + if 'from' in packet: + peer = Node(addr[0], addr[1], base64.b64encode(bytes.fromhex(packet['from']['key'])).decode(), self) + if 'from_short' in packet: + peer = self.peers.get(bytes.fromhex(packet['from_short']['id'])) + + if peer is not None: + received_seqno = packet.get('seqno', 0) + if received_seqno > peer.confirm_seqno: + peer.confirm_seqno = received_seqno + + message = packet.get('message') + messages = packet.get('messages', []) + + if message: + messages = [message] + messages + for message in messages: + try: + await self._process_incoming_message(message, peer) + finally: + continue - message = response.get('message') - messages = response.get('messages') + async def listen(self): + while True: + packet_data, addr = await self.protocol.receive() + try: + await asyncio.wait_for(self.process_packet(packet_data, addr), timeout=1) + except asyncio.TimeoutError: + self.logger.warning(f'packet processing timeout: len({packet_data}) from {addr}') + continue + except Exception as e: + self.logger.warning(f'packet processing error: {e}') + continue - if message: - await self._process_incoming_message(message, peer) - if messages: - for message in messages: - await self._process_incoming_message(message, peer) + async def _wait(self, futures: typing.List[asyncio.Future]): + try: + result = await asyncio.wait_for(self._receive(futures), self.timeout) + return result + except asyncio.TimeoutError: + raise + finally: + for f in futures: + self.tasks.pop(f.id, None) async def send_message_in_channel(self, data: dict, channel: typing.Optional[AdnlChannel] = None, peer: Node = None) -> list: - if peer is None: raise AdnlTransportError('Must provide peer') @@ -440,7 +560,7 @@ async def send_message_in_channel(self, data: dict, channel: typing.Optional[Adn res = channel.encrypt(serialized) self.transport.sendto(res, addr=peer.addr) - result = await asyncio.wait_for(self._receive(futures), self.timeout) + result = await self._wait(futures) return result @@ -476,8 +596,8 @@ async def send_message_outside_channel(self, data: dict, peer: Node) -> list: else: raise Exception(f'sending seqno {sending_seqno}, client seqno: {peer.seqno}') if futures: - result = await asyncio.wait_for(self._receive(futures), self.timeout) - return result + return await self._wait(futures) + async def start(self): self.loop = asyncio.get_running_loop() @@ -487,13 +607,15 @@ async def start(self): reuse_port=True ) self.listener = self.loop.create_task(self.listen()) + self.inited = True return def _get_default_message(self): + random_id = get_random(8) return { '@type': 'adnl.message.query', 'query_id': get_random(32), - 'query': self.schemas.get_by_name('dht.getSignedAddressList').little_id() + 'query': {'@type': 'dht.ping', 'random_id': random_id}, } async def connect_to_peer(self, peer: Node) -> list: @@ -503,6 +625,9 @@ async def connect_to_peer(self, peer: Node) -> list: :return: response dict for default message """ + if peer.key_id in self.peers: + raise AdnlTransportError(f"Peer {peer.key_id.hex()} is already connected") + ts = int(time.time()) channel_client = Client(Client.generate_ed25519_private_key()) create_channel_message = { @@ -517,6 +642,7 @@ async def connect_to_peer(self, peer: Node) -> list: data = { 'from': from_, # 'from_short': {'id': self.client.get_key_id().hex()}, + # 'message': create_channel_message, 'messages': [create_channel_message, default_message], 'address': { 'addrs': [], @@ -530,20 +656,24 @@ async def connect_to_peer(self, peer: Node) -> list: 'dst_reinit_date': 0, } - messages = await self.send_message_outside_channel(data, peer) + self.pending_channels[peer.key_id] = channel_client + self.peers[peer.key_id] = peer + + try: + messages = await self.send_message_outside_channel(data, peer) + except Exception as e: + self.peers.pop(peer.key_id) + await peer.disconnect() + raise e + finally: + self.pending_channels.pop(peer.key_id, None) confirm_channel = messages[0] assert confirm_channel.get('@type') == 'adnl.message.confirmChannel', (f'expected adnl.message.confirmChannel,' f' got {confirm_channel.get("@type")}') assert confirm_channel['peer_key'] == channel_client.ed25519_public.encode().hex() - channel_peer = Server(peer.host, peer.port, bytes.fromhex(confirm_channel['key'])) - channel = AdnlChannel(channel_client, channel_peer, self.local_id, peer.get_key_id()) - self.channels[peer.get_key_id()] = channel - peer.channels.append(channel) - peer.start_ping() peer.connected = True - self.peers[peer.key_id] = peer return messages[1] @@ -552,14 +682,15 @@ async def close(self): while not self.listener.cancelled(): await asyncio.sleep(0) self.transport.abort() + self.inited = False async def send_query_message(self, tl_schema_name: str, data: dict, peer: Node) -> typing.List[dict]: message = { '@type': 'adnl.message.query', 'query_id': get_random(32), 'query': self.schemas.serialize( - self.schemas.get_by_name(tl_schema_name), - data + self.schemas.get_by_name(tl_schema_name), + data ) } diff --git a/pytoniq/adnl/dht.py b/pytoniq/adnl/dht.py index 09404f1..5cb7d30 100644 --- a/pytoniq/adnl/dht.py +++ b/pytoniq/adnl/dht.py @@ -12,7 +12,10 @@ from pytoniq_core.tl import TlGenerator from .adnl import Node, AdnlTransport -from .overlay import OverlayNode, OverlayTransport +from .overlay.overlay import OverlayNode, OverlayTransport + + +_default_tl_schemas = TlGenerator.with_default_schemas().generate() class DhtError(Exception): @@ -50,8 +53,7 @@ def from_dict(cls, transport: AdnlTransport, data: dict, check_signature=True) - # check signature if check_signature: - schemas = TlGenerator.with_default_schemas().generate() - signed_message = schemas.serialize(schema=schemas.get_by_name('dht.node'), data=data) + signed_message = _default_tl_schemas.serialize(schema=_default_tl_schemas.get_by_name('dht.node'), data=data) if not verify_sign(pub_k, signed_message, signature): raise Exception('invalid node signature!') @@ -125,6 +127,8 @@ async def find_value(self, key: bytes, k: int = 6, timeout: int = 10): await asyncio.wait_for(node.connect(), 1) except asyncio.TimeoutError: continue + except Exception: + continue try: resp = await node.find_value(key=key, k=k) except asyncio.exceptions.TimeoutError: @@ -256,8 +260,9 @@ async def get_overlay_node(self, node: dict, overlay_transport: OverlayTransport port = node_addr['port'] pub_k = base64.b64encode(bytes.fromhex(resp['value']['key']['id']['key'])).decode() - node = OverlayNode(peer_host=host, peer_port=port, peer_pub_key=pub_k, transport=overlay_transport) - return node + onode = OverlayNode(peer_host=host, peer_port=port, peer_pub_key=pub_k, transport=overlay_transport) + onode.add_params(node['signature'], node['version']) + return onode @classmethod def from_config(cls, config: dict, adnl_transport: AdnlTransport): diff --git a/pytoniq/adnl/overlay.py b/pytoniq/adnl/overlay.py deleted file mode 100644 index 9468f15..0000000 --- a/pytoniq/adnl/overlay.py +++ /dev/null @@ -1,227 +0,0 @@ -import asyncio -import random -import time -import hashlib -import typing - -from pytoniq_core.tl.generator import TlGenerator - -from pytoniq_core import BlockIdExt, Block, Slice -from pytoniq_core.crypto.ciphers import get_random - -from .adnl import Node, AdnlTransport, AdnlTransportError - - -class OverlayTransportError(AdnlTransportError): - pass - - -class OverlayNode(Node): - - def __init__( - self, - peer_host: str, # ipv4 host - peer_port: int, # port - peer_pub_key: str, - transport: "OverlayTransport" - ): - self.transport: "OverlayTransport" = None - super().__init__(peer_host, peer_port, peer_pub_key, transport) - - async def send_ping(self) -> None: - peers = [ - self.transport.get_signed_myself() - ] - await self.transport.send_query_message('overlay.getRandomPeers', {'peers': {'nodes': peers}}, peer=self) - - -class OverlayTransport(AdnlTransport): - - def __init__(self, - private_key: bytes = None, - tl_schemas_path: str = None, - local_address: tuple = ('0.0.0.0', None), - overlay_id: typing.Union[str, bytes] = None, - *args, **kwargs - ) -> None: - - super().__init__(private_key, tl_schemas_path, local_address, *args, **kwargs) - if overlay_id is None: - raise OverlayTransportError('must provide overlay id in OverlayTransport') - - if isinstance(overlay_id, bytes): - overlay_id = overlay_id.hex() - - self.overlay_id = overlay_id - - @staticmethod - def get_overlay_id(zero_state_file_hash: typing.Union[bytes, str], - workchain: int = 0, shard: int = -9223372036854775808) -> str: - - if isinstance(zero_state_file_hash, bytes): - zero_state_file_hash = zero_state_file_hash.hex() - - schemes = TlGenerator.with_default_schemas().generate() - - sch = schemes.get_by_name('tonNode.shardPublicOverlayId') - data = { - "workchain": workchain, - "shard": shard, - "zero_state_file_hash": zero_state_file_hash - } - - key_id = hashlib.sha256(schemes.serialize(sch, data)).digest() - - sch = schemes.get_by_name('pub.overlay') - data = { - 'name': key_id - } - - key_id = schemes.serialize(sch, data) - - return hashlib.sha256(key_id).digest().hex() - - @classmethod - def get_mainnet_overlay_id(cls, workchain: int = 0, shard: int = -9223372036854775808) -> str: - return cls.get_overlay_id('5e994fcf4d425c0a6ce6a792594b7173205f740a39cd56f537defd28b48a0f6e', workchain, shard) - - @classmethod - def get_testnet_overlay_id(cls, workchain: int = 0, shard: int = -9223372036854775808) -> str: - return cls.get_overlay_id('67e20ac184b9e039a62667acc3f9c00f90f359a76738233379efa47604980ce8', workchain, shard) - - async def _process_query_message(self, message: dict, peer: OverlayNode): - query = message.get('query') - if isinstance(query, list): - if query[0]['@type'] == 'overlay.query': - assert query[0]['overlay'] == self.overlay_id, 'Unknown overlay id received' - query = query[-1] - await self._process_query_handler(message, query, peer) - - async def _process_custom_message(self, message: dict, peer: Node): - data = message.get('data') - if isinstance(data, list): - if data[0]['@type'] in ('overlay.query', 'overlay.message'): - assert data[0]['overlay'] == self.overlay_id, 'Unknown overlay id received' - data = data[-1] - if data['@type'] == 'overlay.broadcast': - # Force broadcast distributing for the network stability. Can be removed in the future. - # Note that this is almost takes no time to do and will be done in the background. - asyncio.create_task(self.distribute_broadcast(data, ignore_errors=True)) - - await self._process_custom_message_handler(data, peer) - - async def distribute_broadcast(self, message: dict, ignore_errors: bool = True): - tasks = [] - peers = random.choices(list(self.peers.items()), k=3) # https://github.com/ton-blockchain/ton/blob/e30049930a7372a3c1d28a1e59956af8eb489439/overlay/overlay-broadcast.cpp#L69 - for _, peer in peers: - tasks.append(self.send_custom_message(message, peer)) - result = await asyncio.gather(*tasks, return_exceptions=ignore_errors) - failed = 0 - for r in result: - if isinstance(r, Exception): - failed += 1 - self.logger.debug(f'Spread broadcast: {failed} failed out of {len(result)}') - - def get_signed_myself(self): - ts = int(time.time()) - - overlay_node_data = {'id': {'@type': 'pub.ed25519', 'key': self.client.ed25519_public.encode().hex()}, - 'overlay': self.overlay_id, 'version': ts, 'signature': b''} - - overlay_node_to_sign = self.schemas.serialize(self.schemas.get_by_name('overlay.node.toSign'), - {'id': {'id': self.client.get_key_id().hex()}, - 'overlay': self.overlay_id, - 'version': overlay_node_data['version']}) - signature = self.client.sign(overlay_node_to_sign) - - overlay_node = overlay_node_data | {'signature': signature} - return overlay_node - - async def send_query_message(self, tl_schema_name: str, data: dict, peer: Node) -> typing.List[typing.Union[dict, bytes]]: - """ - :param tl_schema_name: - :param data: - :param peer: - :return: dict if response was known TL schema, bytes otherwise - """ - - message = { - '@type': 'adnl.message.query', - 'query_id': get_random(32), - 'query': self.schemas.serialize(self.schemas.get_by_name('overlay.query'), data={'overlay': self.overlay_id}) - + self.schemas.serialize(self.schemas.get_by_name(tl_schema_name), data) - } - data = { - 'message': message, - } - - result = await self.send_message_in_channel(data, None, peer) - return result - - async def send_custom_message(self, message: typing.Union[dict, bytes], peer: Node) -> list: - - custom_message = { - '@type': 'adnl.message.custom', - 'data': (self.schemas.serialize(self.schemas.get_by_name('overlay.message'), data={'overlay': self.overlay_id}) + - self.schemas.serialize(self.schemas.get_by_name(message['@type']), message)) - } - - data = { - 'message': custom_message, - } - - result = await self.send_message_in_channel(data, None, peer) - return result - - def get_message_with_overlay_prefix(self, schema_name: str, data: dict) -> bytes: - return (self.schemas.serialize( - schema=self.schemas.get_by_name('overlay.query'), - data={'overlay': self.overlay_id}) - + self.schemas.serialize( - schema=self.schemas.get_by_name(schema_name), - data=data) - ) - - def _get_default_message(self): - peers = [ - self.get_signed_myself() - ] - return { - '@type': 'adnl.message.query', - 'query_id': get_random(32), - 'query': self.get_message_with_overlay_prefix('overlay.getRandomPeers', {'peers': {'nodes': peers}}) - } - - async def get_random_peers(self, peer: OverlayNode): - overlay_node = self.get_signed_myself() - - peers = [ - overlay_node - ] - return await self.send_query_message(tl_schema_name='overlay.getRandomPeers', data={'peers': {'nodes': peers}}, - peer=peer) - - async def get_capabilities(self, peer: OverlayNode): - return await self.send_query_message(tl_schema_name='tonNode.getCapabilities', data={}, peer=peer) - - async def raw_download_block(self, block: BlockIdExt, peer: OverlayNode) -> bytes: - """ - :param block: - :param peer: - :return: block boc - """ - return (await self.send_query_message(tl_schema_name='tonNode.downloadBlock', - data={'block': block.to_dict()}, peer=peer))[0] - - async def download_block(self, block: BlockIdExt, peer: OverlayNode) -> Block: - """ - :param block: - :param peer: - :return: deserialized block - """ - blk_boc = await self.raw_download_block(block, peer) - return Block.deserialize(Slice.one_from_boc(blk_boc)) - - async def prepare_block(self, block: BlockIdExt, peer: OverlayNode) -> dict: - return (await self.send_query_message(tl_schema_name='tonNode.prepareBlock', - data={'block': block.to_dict()}, peer=peer))[0] diff --git a/pytoniq/adnl/overlay/__init__.py b/pytoniq/adnl/overlay/__init__.py new file mode 100644 index 0000000..6fa74b1 --- /dev/null +++ b/pytoniq/adnl/overlay/__init__.py @@ -0,0 +1,5 @@ +from .overlay import OverlayTransport, OverlayNode, OverlayTransportError +from .broadcast import BroadcastSimple, InvalidBroadcast +from .fec_broadcast import BroadcastFecPart, BroadcastFec, InvalidBroadcastFec, create_fec_broadcast +from .overlay_manager import OverlayManager +from .shard_overlay import ShardOverlay diff --git a/pytoniq/adnl/overlay/broadcast.py b/pytoniq/adnl/overlay/broadcast.py new file mode 100644 index 0000000..9102b8b --- /dev/null +++ b/pytoniq/adnl/overlay/broadcast.py @@ -0,0 +1,144 @@ +import asyncio +import hashlib +import logging +import time + +from pytoniq_core.crypto.ciphers import Server +from pytoniq_core.crypto.signature import verify_sign + +from .overlay import OverlayTransport +from .privacy import BroadcastCheckResult + + +class InvalidBroadcast(Exception): + pass + + +class BroadcastSimple: + """ + overlay.broadcast src:PublicKey certificate:overlay.Certificate flags:int data:bytes date:int signature:bytes = overlay.Broadcast; + """ + + def __init__( + self, + overlay_transport: OverlayTransport, + data: dict, + ): + self._overlay = overlay_transport + self._data = data + self.is_valid = False + if isinstance(data['data'], dict): + self.data_bytes = self._overlay.schemas.serialize(self._data['data'].get('@type'), data['data']) + else: + self.data_bytes = self._data['data'] + self._logger = logging.getLogger(self.__class__.__name__) + + @property + def date(self) -> int: + return self._data['date'] + + @property + def flags(self) -> int: + return self._data['flags'] + + @property + def serialized(self) -> bytes: + return self._overlay.schemas.serialize('overlay.broadcast', self._data) + + @property + def hash(self) -> bytes: + return hashlib.sha256(self.serialized).digest() + + @property + def source_key(self) -> bytes: + return bytes.fromhex(self._data['src']['key']) + + @property + def broadcast_hash(self): + data_hash = hashlib.sha256(self.data_bytes).digest() + key_id = Server('', 0, self.source_key).get_key_id().hex() + return self.compute_broadcast_id(self._overlay, data_hash.hex(), key_id, self.flags) + + @staticmethod + def compute_broadcast_id(overlay: OverlayTransport, data_hash: str, src: str, flags: int) -> bytes: + if flags & 1: + src = (b'\x00' * 32).hex() + broadcast_id_data = {'data_hash': data_hash, 'src': src, 'flags': flags} + broadcast_id_serialized = overlay.schemas.serialize('overlay.broadcast.id', broadcast_id_data) + return hashlib.sha256(broadcast_id_serialized).digest() + + def check_signature(self) -> bool: + to_sign_data = {'hash': self.broadcast_hash.hex(), + 'date': self.date} + to_sign = self._overlay.schemas.serialize('overlay.broadcast.toSign', to_sign_data) + return verify_sign(self.source_key, to_sign, self._data['signature']) + + def run_checks(self) -> None: + if self.date < int(time.time()) - 20: + raise InvalidBroadcast('broadcast is too old') + + if self.date > int(time.time()) + 20: + raise InvalidBroadcast('broadcast is too new') + + if self.hash in self._overlay.broadcasts: + raise InvalidBroadcast('broadcast already received') + + r = self._overlay.check_source_eligible(self.source_key, self._data['certificate'], len(self.data_bytes), False) + + if r == BroadcastCheckResult.Forbidden: + raise InvalidBroadcast('source is not eligible') + self.is_valid = r == BroadcastCheckResult.Allowed + + if not self.check_signature(): + raise InvalidBroadcast('invalid signature') + + async def run(self): + try: + self.run_checks() + except InvalidBroadcast as e: + self._logger.debug(f'Failed to check broadcast: {e}, brcst: {self._data}') + return + self._overlay.broadcasts[self.hash] = self + source_key_id = Server('', 0, self.source_key).get_key_id() + data = self._data['data'] + if isinstance(data, bytes): + try: + data, _ = self._overlay.schemas.deserialize(data) + except: + pass + if not self.is_valid: + if await self._overlay.check_broadcast(data, source_key_id): + self.is_valid = True + if self.is_valid: + await self._overlay.handle_broadcast(data, source_key_id) + await self.distribute() + + async def distribute(self) -> None: + tasks = [] + peers = self._overlay.get_neighbours(3) + for peer in peers: + tasks.append(self._overlay.send_custom_message(self.serialized, peer)) + result = await asyncio.gather(*tasks, return_exceptions=True) + + @classmethod + def create(cls, overlay: OverlayTransport, data: bytes, flags: int = 0) -> "BroadcastSimple": + ts = int(time.time()) + data_hash = hashlib.sha256(data).digest() + to_sign_data = { + 'hash': cls.compute_broadcast_id(overlay, data_hash.hex(), overlay.client.get_key_id().hex(), flags).hex(), + 'date': ts + } + to_sign = overlay.schemas.serialize('overlay.broadcast.toSign', to_sign_data) + signature = overlay.client.sign(to_sign) + + from_ = {'@type': 'pub.ed25519', 'key': overlay.client.ed25519_public.encode().hex()} + broadcast = { + '@type': 'overlay.broadcast', + 'src': from_, + 'certificate': {'@type': 'overlay.emptyCertificate'}, + 'flags': flags, + 'data': data, + 'signature': signature, + 'date': ts + } + return cls(overlay, broadcast) diff --git a/pytoniq/adnl/overlay/fec_broadcast.py b/pytoniq/adnl/overlay/fec_broadcast.py new file mode 100644 index 0000000..42d072b --- /dev/null +++ b/pytoniq/adnl/overlay/fec_broadcast.py @@ -0,0 +1,374 @@ +import asyncio +import hashlib +import logging +import time + +from pytoniq_core.crypto.ciphers import Server +from pytoniq_core.crypto.signature import verify_sign +from pytoniq_core.tl.generator import TlSchemas + +from .overlay import OverlayTransport +from .privacy import BroadcastCheckResult +from ..rldp.raptorq import get_decoder, get_encoder + + +class InvalidBroadcastFec(Exception): + pass + + +class BroadcastFec: + + def __init__( + self, + broadcast_hash: bytes, + src: bytes, + data_hash: bytes, + flags: int, + date: int, + fec_type: dict, + overlay: OverlayTransport + ): + self.broadcast_hash = broadcast_hash + self.src = src + self.data_hash = data_hash + self.flags = flags + self.date = date + self.fec_type = fec_type + self.decoder = None + self.encode = None + self.next_seqno = 0 + self.received_parts = 0 + self.parts = {} + self.result = None + self.ready = False + self.completed_neighbours = set() + self._overlay = overlay + self._logger = logging.getLogger(self.__class__.__name__) + + self.run_checks() + self.init_fec_type() + + def run_checks(self): + if self.fec_type['data_size'] > self._overlay.max_fec_broadcast_size: + raise InvalidBroadcastFec('too big fec broadcast') + + def init_fec_type(self): + if self.fec_type['@type'] != 'fec.raptorQ': + raise InvalidBroadcastFec('unsupported fec type') + self.decoder = get_decoder( + self._overlay.raptorq_engine, + self.fec_type['data_size'], + self.fec_type['symbol_size'], + self.fec_type['symbols_count'] + ) + + def received_part(self, seqno: int) -> bool: + if seqno + 64 < self.next_seqno: + return True + if seqno >= self.next_seqno: + return False + return bool(self.received_parts & (1 << (self.next_seqno - seqno - 1))) + + def add_received_part(self, seqno: int): + if seqno < self.next_seqno: + self.received_parts |= (1 << (self.next_seqno - seqno - 1)) + else: + old = self.next_seqno + self.next_seqno = seqno + 1 + if self.next_seqno - old >= 64: + self.received_parts = 1 + else: + self.received_parts = self.received_parts << (self.next_seqno - old) + self.received_parts |= 1 + + def is_eligible_sender(self, src: bytes): + if self.flags & 1: + return True + return src == self.src + + def add_part(self, seqno: int, data: bytes, serialized: bytes) -> bool: + res = self.decoder.add_symbol(seqno, data) + self.parts[seqno] = serialized + return res + + def finish(self): + if not self.decoder.may_try_decode(): + raise Exception('need more parts') + result = self.decoder.try_decode() + if result: + if hashlib.sha256(result).digest() != self.data_hash: + raise InvalidBroadcastFec('data hash mismatch') + # self.encoder = get_encoder(self.result, self.fec_type['symbol_size']) todo + self.ready = True + del self.decoder # can be useful: so we dont need to wait for gc + self.decoder = None + return result + + def finalized(self): + return self.ready + + def add_completed(self, peer_id: bytes): + self.completed_neighbours.add(peer_id) + + async def distribute_part(self, seqno: int): + if seqno not in self.parts: + return + data = self.parts.get(seqno) + peers = self._overlay.get_neighbours(5) + tasks = [] + for peer in peers: + if peer.get_key_id() in self.completed_neighbours: # todo: short broadcasts + continue + tasks.append(self._overlay.send_custom_message(data, peer)) + result = await asyncio.gather(*tasks, return_exceptions=True) + + +class BroadcastFecPart: + """ + overlay.broadcastFec src:PublicKey certificate:overlay.Certificate data_hash:int256 data_size:int flags:int + data:bytes seqno:int fec:fec.Type date:int signature:bytes = overlay.Broadcast; + """ + + def __init__( + self, + overlay_transport: OverlayTransport, + data: dict + ): + self.brcst: BroadcastFec = None + self._overlay = overlay_transport + self._data = data + self.untrusted = False + + self._logger = logging.getLogger(self.__class__.__name__) + + @property + def date(self) -> int: + return self._data['date'] + + @property + def flags(self) -> int: + return self._data['flags'] + + @property + def seqno(self) -> int: + return self._data['seqno'] + + @property + def source_key(self) -> bytes: + return bytes.fromhex(self._data['src']['key']) + + @property + def broadcast_hash(self): + key_id = Server('', 0, self.source_key).get_key_id().hex() + return self.compute_broadcast_id(self._overlay, self._data['data_hash'], key_id, self._data['flags'], self._data['fec'], self._data['data_size']) + + @property + def part_data_hash(self): + return hashlib.sha256(self._data['data']).digest() + + @property + def part_hash(self): + return self.compute_broadcast_part_id(self._overlay.schemas, self.broadcast_hash.hex(), self.part_data_hash.hex(), self.seqno) + + @property + def serialized(self) -> bytes: + return self._overlay.schemas.serialize('overlay.broadcastFec', self._data) + + @property + def is_short(self) -> bool: + return 'Short' in self._data['@type'] + + @staticmethod + def compute_broadcast_id(overlay: OverlayTransport, data_hash: str, src: str, flags: int, fec_type: dict, size: int) -> bytes: + """ + overlay.broadcastFec.id src:int256 type:int256 data_hash:int256 size:int flags:int = overlay.broadcastFec.Id; + """ + if flags & 1: + src = (b'\x00' * 32).hex() + type_ = hashlib.sha256(overlay.schemas.serialize(fec_type['@type'], fec_type)).digest().hex() + broadcast_id_data = {'data_hash': data_hash, 'src': src, 'flags': flags, 'type': type_, 'size': size} + broadcast_id_serialized = overlay.schemas.serialize('overlay.broadcastFec.id', broadcast_id_data) + return hashlib.sha256(broadcast_id_serialized).digest() + + @staticmethod + def compute_broadcast_part_id(schemes: TlSchemas, broadcast_hash: str, data_hash: str, seqno: int): + """ + overlay.broadcastFec.partId broadcast_hash:int256 data_hash:int256 seqno:int = overlay.broadcastFec.PartId; + """ + data = {'broadcast_hash': broadcast_hash, 'data_hash': data_hash, 'seqno': seqno} + return hashlib.sha256(schemes.serialize('overlay.broadcastFec.partId', data)).digest() + + def check_signature(self) -> bool: + to_sign_data = {'hash': self.part_hash.hex(), + 'date': self.date} + to_sign = self._overlay.schemas.serialize('overlay.broadcast.toSign', to_sign_data) + return verify_sign(self.source_key, to_sign, self._data['signature']) + + def run_checks(self): + if self.date < int(time.time()) - 20: + raise InvalidBroadcastFec('broadcast is too old') + + if self.date > int(time.time()) + 20: + raise InvalidBroadcastFec('broadcast is too new') + + if self._data['fec']['@type'] != 'fec.raptorQ': + raise InvalidBroadcastFec('unsupported fec type') + + if self.brcst and self.brcst.received_part(self.seqno): + raise InvalidBroadcastFec('broadcast already received') + + r = self._overlay.check_source_eligible(self.source_key, self._data['certificate'], self._data['data_size'], True) + if r == BroadcastCheckResult.Forbidden: + raise InvalidBroadcastFec('source is not eligible') + if r == BroadcastCheckResult.NeedCheck: + self.untrusted = True + + if self.brcst: + if not self.brcst.is_eligible_sender(self.source_key): + raise InvalidBroadcastFec('source is not eligible') + + if not self.check_signature(): + raise InvalidBroadcastFec('signature is not valid') + + async def apply(self): + if not self.brcst: + self.brcst = self._overlay.fec_broadcasts.get(self.broadcast_hash) + if not self.brcst: + if self.is_short: + self._logger.debug(f'short broadcast part for incomplete broadcast') + return + b = BroadcastFec(self.broadcast_hash, self.source_key, bytes.fromhex(self._data['data_hash']), + self._data['flags'], self.date, self._data['fec'], self._overlay) + self.brcst = b + self._overlay.fec_broadcasts[self.broadcast_hash] = b + if self.brcst.received_part(self.seqno): + raise InvalidBroadcastFec('duplicate part') + self.brcst.add_received_part(self.seqno) + + if self.brcst.finalized() and self.is_short: + raise InvalidBroadcastFec('short broadcast part for incomplete broadcast') + + if not self.brcst.finalized(): + self.brcst.add_part( + self.seqno, + self._data['data'], + self._overlay.schemas.serialize(self._data['@type'], self._data) + ) + if self.brcst.decoder.may_try_decode(): + try: + r = self.brcst.finish() + except InvalidBroadcastFec: + self._logger.debug(f'failed to finish broadcast: {self._data}') + return + except Exception as e: + if 'need more parts' in str(e): + return + raise e + + try: + r, _ = self._overlay.schemas.deserialize(r) + except: + pass + + if self.untrusted: + if await self._overlay.check_broadcast(r, self.source_key): + await self._overlay.handle_broadcast(r, self.source_key) + # await self.distribute() # todo: check why we distribute only one part + else: + await self._overlay.handle_broadcast(r, self.source_key) + + async def run(self): + try: + self.run_checks() + except InvalidBroadcastFec as e: + self._logger.debug(f'Failed to check broadcast: {e}, brcst: {self._data}') + return + try: + await self.apply() + except InvalidBroadcastFec as e: + self._logger.debug(f'Failed to apply broadcast: {e}, brcst: {self._data}') + return + # if not self.untrusted: + await self.distribute() + + async def distribute(self): + await self.brcst.distribute_part(self.seqno) + + @classmethod + def create( + cls, overlay: OverlayTransport, part: bytes, data_hash: str, + seqno: int, flags: int, fec_type: dict, data_size: int, date: int + ): + broadcast_hash = cls.compute_broadcast_id( + overlay=overlay, + data_hash=data_hash, + src=overlay.client.get_key_id().hex(), + flags=flags, + fec_type=fec_type, + size=data_size + ) + part_data_hash = hashlib.sha256(part).digest() + part_hash = cls.compute_broadcast_part_id(overlay.schemas, broadcast_hash.hex(), part_data_hash.hex(), seqno) + + to_sign_data = {'hash': part_hash.hex(), + 'date': date} + to_sign = overlay.schemas.serialize('overlay.broadcast.toSign', to_sign_data) + signature = overlay.client.sign(to_sign) + + part_data = { + '@type': 'overlay.broadcastFec', + 'src': {'@type': 'pub.ed25519', 'key': overlay.client.ed25519_public.encode().hex()}, + 'certificate': overlay.get_certificate(), # todo certificates + 'data_hash': data_hash, + 'data_size': data_size, + 'flags': flags, + 'data': part, + 'seqno': seqno, + 'fec': fec_type, + 'date': date, + 'signature': signature + } + return cls(overlay, part_data) + + +async def create_fec_broadcast(overlay: OverlayTransport, data: bytes, flags: int): + if len(data) > 1 << 27: + raise InvalidBroadcastFec('too big data') + + symbol_size = overlay.max_simple_broadcast_size + to_send = int((len(data) / symbol_size + 1) * 2) + symbols_count = (len(data) + symbol_size - 1) // symbol_size + + data_hash = hashlib.sha256(data).digest() + ts = int(time.time()) + fec_type = {'@type': 'fec.raptorQ', 'data_size': len(data), 'symbol_size': symbol_size, 'symbols_count': symbols_count} + try: + encoder = get_encoder( + overlay.raptorq_engine, + data, + symbol_size + ) + except: + raise InvalidBroadcastFec('failed to create encoder') + + seqno = 0 + broadcast_hash = b'' + + while seqno < to_send: + for _ in range(4): + part = encoder.gen_symbol(seqno) + if part is None: + seqno += 1 + continue + part = BroadcastFecPart.create(overlay, part, data_hash.hex(), seqno, flags, fec_type, len(data), ts) + try: + await part.run() + except Exception as e: + logging.getLogger('create_fec_broadcast').debug(f'failed to run part: {e}') + pass + broadcast_hash = part.broadcast_hash + + seqno += 1 + await asyncio.sleep(0.01) + + return broadcast_hash diff --git a/pytoniq/adnl/overlay/overlay.py b/pytoniq/adnl/overlay/overlay.py new file mode 100644 index 0000000..d81749e --- /dev/null +++ b/pytoniq/adnl/overlay/overlay.py @@ -0,0 +1,343 @@ +import enum +import inspect +import random +import time +import hashlib +import typing + +from pytoniq_core.tl.generator import TlGenerator +from pytoniq_core.crypto.ciphers import get_random, Server + +from ..adnl import Node, AdnlTransport, AdnlTransportError +from .privacy import OverlayPrivacyRules, Certificate, BroadcastCheckResult + + +class OverlayTransportError(AdnlTransportError): + pass + + +class OverlayNode(Node): + + PING_INTERVAL = 10 + + def __init__( + self, + peer_host: str, # ipv4 host + peer_port: int, # port + peer_pub_key: str, + transport: "OverlayTransport" + ): + self.transport: "OverlayTransport" = None + self.signature = b'' + self.version = 0 + super().__init__(peer_host, peer_port, peer_pub_key, transport) + + def add_params(self, signature: bytes, version: int): + self.signature = signature + self.version = version + + def to_tl(self) -> typing.Optional[dict]: + if not self.signature: + return None + return { + '@type': 'overlay.node', + 'id': {'@type': 'pub.ed25519', 'key': self.ed25519_public.encode().hex()}, + 'overlay': self.transport.overlay_id, + 'version': self.version, + 'signature': self.signature + } + + async def send_ping(self) -> None: + peers = [ + self.transport.get_signed_myself() + ] + await self.transport.send_query_message('overlay.getRandomPeers', {'peers': {'nodes': peers}}, peer=self) + + +class OverlayTransport(AdnlTransport): + max_simple_broadcast_size = 768 + max_fec_broadcast_size = 16 << 20 + + def __init__(self, + private_key: bytes = None, + tl_schemas_path: str = None, + local_address: tuple = ('0.0.0.0', None), + overlay_id: typing.Union[str, bytes] = None, + *args, **kwargs + ) -> None: + + super().__init__(private_key, tl_schemas_path, local_address, *args, **kwargs) + if overlay_id is None: + raise OverlayTransportError('must provide overlay id in OverlayTransport') + + if isinstance(overlay_id, bytes): + overlay_id = overlay_id.hex() + + self.overlay_id = overlay_id + self.broadcasts = {} + self.fec_broadcasts = {} + self.broadcast_checkers: typing.Dict[str, typing.Callable] = {} + self.broadcast_handlers: typing.Dict[str, typing.Callable] = {} + if 'rules' in kwargs: + self.rules = kwargs['rules'] + assert isinstance(self.rules, OverlayPrivacyRules), 'rules must be instance of OverlayPrivacyRules' + else: + self.rules = OverlayPrivacyRules.default(allow_fec=kwargs.get('allow_fec', False)) + if self.rules.allow_fec: + self.raptorq_engine = kwargs.get('raptorq_engine', None) + self.max_peers = kwargs.get('max_peers', 30) + self.signed_myself = None + self.signed_myself = self.get_signed_myself() + + @staticmethod + def get_overlay_id(zero_state_file_hash: typing.Union[bytes, str], + workchain: int = 0, shard: int = -9223372036854775808) -> str: + + if isinstance(zero_state_file_hash, bytes): + zero_state_file_hash = zero_state_file_hash.hex() + + schemes = TlGenerator.with_default_schemas().generate() + + sch = schemes.get_by_name('tonNode.shardPublicOverlayId') + data = { + "workchain": workchain, + "shard": shard, + "zero_state_file_hash": zero_state_file_hash + } + + key_id = hashlib.sha256(schemes.serialize(sch, data)).digest() + + sch = schemes.get_by_name('pub.overlay') + data = { + 'name': key_id + } + + key_id = schemes.serialize(sch, data) + + return hashlib.sha256(key_id).digest().hex() + + @classmethod + def get_mainnet_overlay_id(cls, workchain: int = 0, shard: int = -1 << 63) -> str: + return cls.get_overlay_id('5e994fcf4d425c0a6ce6a792594b7173205f740a39cd56f537defd28b48a0f6e', workchain, shard) + + @classmethod + def get_testnet_overlay_id(cls, workchain: int = 0, shard: int = -1 << 63) -> str: + return cls.get_overlay_id('67e20ac184b9e039a62667acc3f9c00f90f359a76738233379efa47604980ce8', workchain, shard) + + async def _process_query_message(self, message: dict, peer: OverlayNode): + query = message.get('query') + if isinstance(query, list): + if query[0]['@type'] == 'overlay.query': + assert query[0]['overlay'] == self.overlay_id, 'Unknown overlay id received' + query = query[-1] + await self._process_query_handler(message, query, peer) + + async def _process_custom_message(self, message: dict, peer: Node): + data = message.get('data') + if isinstance(data, list): + if data[0]['@type'] in ('overlay.query', 'overlay.message'): + assert data[0]['overlay'] == self.overlay_id, 'Unknown overlay id received' + data = data[-1] + # Force broadcast distributing: Note that this is almost takes no time and will be done in the background + if data['@type'] == 'overlay.broadcast': + from .broadcast import BroadcastSimple + try: + await BroadcastSimple(self, data).run() + except Exception as e: + self.logger.debug(f'Error while processing broadcast: {type(e)}: {e}') + return + self.bcast_gc() + return + if data['@type'] == 'overlay.broadcastFec': + from .fec_broadcast import BroadcastFecPart + try: + await BroadcastFecPart(self, data).run() + except Exception as e: + self.logger.debug(f'Error while processing broadcastFec: {type(e)}: {e}') + self.bcast_gc() + return + + await self._process_custom_message_handler(data, peer) + + def get_neighbours(self, max_size: int): + if len(self.peers) <= max_size: + return list(self.peers.values()) + peers = random.choices(list(self.peers.values()), k=max_size) # https://github.com/ton-blockchain/ton/blob/e30049930a7372a3c1d28a1e59956af8eb489439/overlay/overlay-broadcast.cpp#L69 + return peers + + def get_signed_myself(self): + ts = int(time.time()) + if self.signed_myself is not None and self.signed_myself['version'] > ts - 60: + return self.signed_myself + + overlay_node_data = {'id': {'@type': 'pub.ed25519', 'key': self.client.ed25519_public.encode().hex()}, + 'overlay': self.overlay_id, 'version': ts, 'signature': b''} + + overlay_node_to_sign = self.schemas.serialize(self.schemas.get_by_name('overlay.node.toSign'), + {'id': {'id': self.local_id.hex()}, + 'overlay': self.overlay_id, + 'version': overlay_node_data['version']}) + signature = self.client.sign(overlay_node_to_sign) + + overlay_node = overlay_node_data | {'signature': signature} + self.signed_myself = overlay_node + return overlay_node + + async def send_query_message(self, tl_schema_name: str, data: dict, peer: Node) -> typing.List[typing.Union[dict, bytes]]: + """ + :param tl_schema_name: + :param data: + :param peer: + :return: dict if response was known TL schema, bytes otherwise + """ + + message = { + '@type': 'adnl.message.query', + 'query_id': get_random(32), + 'query': self.schemas.serialize(self.schemas.get_by_name('overlay.query'), data={'overlay': self.overlay_id}) + + self.schemas.serialize(self.schemas.get_by_name(tl_schema_name), data) + } + data = { + 'message': message, + } + + result = await self.send_message_in_channel(data, None, peer) + return result + + async def send_custom_message(self, message: typing.Union[dict, bytes], peer: Node) -> list: + if isinstance(message, dict): + message = self.schemas.serialize(message['@type'], message) + + custom_message = { + '@type': 'adnl.message.custom', + 'data': self.get_message_with_overlay_prefix(message, False) + } + + data = { + 'message': custom_message, + } + + result = await self.send_message_in_channel(data, None, peer) + return result + + def get_message_with_overlay_prefix(self, data: dict, query: bool) -> bytes: + if isinstance(data, dict): + data = self.schemas.serialize(schema=data['@type'], data=data) + return (self.schemas.serialize( + schema='overlay.query' if query else 'overlay.message', + data={'overlay': self.overlay_id}) + data + ) + + def _get_default_message(self): + peers = [ + self.get_signed_myself() + ] + return { + '@type': 'adnl.message.query', + 'query_id': get_random(32), + 'query': self.get_message_with_overlay_prefix( + {'@type': 'overlay.getRandomPeers', 'peers': {'nodes': peers}}, + True + ) + } + + async def get_random_peers(self, peer: OverlayNode): + known_peers = self.get_neighbours(5) + peers = [self.get_signed_myself()] + for peer in known_peers: + if peer.to_tl(): + peers.append(peer.to_tl()) + return await self.send_query_message(tl_schema_name='overlay.getRandomPeers', data={'peers': {'nodes': peers}}, + peer=peer) + + async def get_capabilities(self, peer: OverlayNode): + return await self.send_query_message(tl_schema_name='tonNode.getCapabilities', data={}, peer=peer) + + def bcast_gc(self): + i = iter(self.broadcasts.copy()) + while len(self.broadcasts) > 250: + brcst = next(i) + del self.broadcasts[brcst] + for b_hash, b in list(self.fec_broadcasts.items()): + if b.date < time.time() - 60: + del self.fec_broadcasts[b_hash] + else: + break + + def check_source_eligible(self, source: bytes, cert: dict, size: int, is_feq: bool) -> BroadcastCheckResult: + if size == 0: + return BroadcastCheckResult.Forbidden + key_id = Server('', 0, source).get_key_id() + r = self.rules.check_rules(key_id, size, is_feq) + if cert['@type'] == 'overlay.emptyCertificate' or r == BroadcastCheckResult.Allowed: + return r + r2 = Certificate(cert, self.schemas).check(key_id, bytes.fromhex(self.overlay_id), size, is_feq) + issuer_key_id = Server('', 0, bytes.fromhex(cert['issued_by']['key'])).get_key_id() + r2 = min(r2.value, self.rules.check_rules(issuer_key_id, size, is_feq).value) + return BroadcastCheckResult(max(r.value, r2)) + + async def check_broadcast(self, data: typing.Union[bytes, dict], src_key_id: bytes) -> bool: + if isinstance(data, dict): + checker = self.broadcast_checkers.get(data['@type'], self.broadcast_checkers.get(None)) + else: + checker = self.broadcast_checkers.get(None) + if checker: + if inspect.iscoroutinefunction(checker): + try: + return await checker(data, src_key_id) + except: + return False + try: + return checker(data, src_key_id) + except: + return False + return True + + async def handle_broadcast(self, data: typing.Union[bytes, dict], src_key_id: bytes): + if isinstance(data, dict): + handler = self.broadcast_handlers.get(data['@type'], self.broadcast_handlers.get(None)) + else: + handler = self.broadcast_handlers.get(None) + if handler: + if inspect.iscoroutinefunction(handler): + try: + return await handler(data, src_key_id) + except: + return False + try: + return handler(data, src_key_id) + except: + return False + return True + + def set_broadcast_checker(self, type_: str, checker: typing.Callable): + """ + :param type_: TL type of broadcast + :param checker: function to handle message. **Must** return dict or bytes or None. If + None returned than answer won't be sent to the sender. Takes two arguments: data (dict or bytes) and src_key_id (bytes) + :return: + """ + self.broadcast_checkers[type_] = checker + + def set_default_broadcast_checker(self, checker: typing.Callable): + self.set_broadcast_checker(None, checker) + + def set_broadcast_handler(self, type_: str, handler: typing.Callable): + """ + :param type_: TL type of broadcast + :param handler: function to handle message. **Must** return dict or bytes or None. If + None returned than answer won't be sent to the sender. Takes two arguments: data (dict or bytes) and src_key_id (bytes) + :return: + """ + self.broadcast_handlers[type_] = handler + + def set_default_broadcast_handler(self, handler: typing.Callable): + self.set_broadcast_handler(None, handler) + + def get_certificate(self): + data = {'@type': 'overlay.certificate', 'issued_by': {'@type': 'pub.ed25519', 'key': self.client.ed25519_public.encode().hex()}, + 'expire_at': int(time.time()) + 3600, 'max_size': 16 << 20, 'signature': b''} + to_sign = self.schemas.serialize('overlay.certificate', data) + signature = self.client.sign(to_sign) + data['signature'] = signature + return data diff --git a/pytoniq/adnl/overlay/overlay_manager.py b/pytoniq/adnl/overlay/overlay_manager.py new file mode 100644 index 0000000..e490873 --- /dev/null +++ b/pytoniq/adnl/overlay/overlay_manager.py @@ -0,0 +1,130 @@ +import logging +import asyncio +import random + +from pytoniq_core.crypto.ciphers import Server + +from .overlay import OverlayTransport, OverlayNode + + +def process_get_random_peers_request(_, overlay_client: OverlayTransport): + known_peers = overlay_client.get_neighbours(5) + peers = [overlay_client.get_signed_myself()] + for peer in known_peers: + if peer.to_tl(): + peers.append(peer.to_tl()) + return { + '@type': 'overlay.nodes', + 'nodes': peers + } + + +def process_get_capabilities_request(_): + return { + '@type': 'tonNode.capabilities', + 'version': 2, + 'capabilities': 2, + } + + +class OverlayManager: + + def __init__(self, overlay: OverlayTransport, dht_client, max_peers: int = 30): + from ..dht import DhtClient + self.overlay = overlay + self.dht: DhtClient = dht_client + self.max_peers = max_peers + self.logger = logging.getLogger(self.__class__.__name__) + self.init_handlers() + + def init_handlers(self): + self.overlay.set_query_handler(type_='overlay.getRandomPeers', + handler=lambda i: process_get_random_peers_request(i, self.overlay)) + self.overlay.set_query_handler(type_='tonNode.getCapabilities', + handler=lambda i: process_get_capabilities_request(i)) + + async def start(self): + if not self.overlay.inited: + await self.overlay.start() + self.overlay.loop.create_task(self.get_more_peers()) + + async def get_more_peers(self): + while True: + if len(self.overlay.peers) == 0: + self.logger.debug('Getting first peers! This may take some time') + try: + nodes = await self.dht.get_overlay_nodes(self.overlay.overlay_id, self.overlay) + except asyncio.TimeoutError: + nodes = [] + except Exception as e: + self.logger.warning(f'Failed to get first peers: {e}') + await asyncio.sleep(10) + continue + for node in nodes: + node: OverlayNode + if node is None: + continue + try: + await asyncio.wait_for(node.connect(), 1.5) + self.logger.debug(f'Connected to peer {node.key_id.hex()}') + except asyncio.TimeoutError: + self.overlay.peers.pop(node.key_id, None) + continue + self.logger.debug(f'Got {len(self.overlay.peers)} first peers') + await asyncio.sleep(10) + continue + dif = self.max_peers - len(self.overlay.peers) + if dif <= 1: + await asyncio.sleep(10) + continue + clients = [] + tasks = [] + peers = list(self.overlay.peers.items()) + random.shuffle(peers) + for _, peer in peers: + if dif <= 0: + break + self.logger.debug(f'getting nodes from peer {peer.get_key_id().hex()}') + if not peer.connected: + self.logger.debug(f'peer {peer.get_key_id().hex()} is already not connected') + continue + tasks.append(self.overlay.get_random_peers(peer)) + dif -= 3 # lets assume we got at least 3 alive peer from a node + result = await asyncio.gather(*tasks, return_exceptions=True) + tasks = [] + for resp in result: + if isinstance(resp, Exception): + continue + self.logger.debug(f"got {len(resp[0]['nodes'])} from peer") + for node in resp[0]['nodes']: + pub_k = bytes.fromhex(node['id']['key']) + adnl_addr = Server('', 0, pub_key=pub_k).get_key_id() + if adnl_addr not in self.overlay.peers: + tasks.append(self.dht.get_overlay_node(node, self.overlay)) + # new_client = await dht_client.get_overlay_node(node, overlay) + # if new_client is not None: + # clients.append(new_client) + result = await asyncio.gather(*tasks, return_exceptions=True) + + for new_client in result: + if isinstance(new_client, (Exception, BaseException)): + continue + if new_client is not None: + clients.append(new_client) + + async def try_connect(client): + try: + await client.connect() + return True + except asyncio.TimeoutError: + return False + + tasks = [] + for new_client in clients: + if len(self.overlay.peers) + len(tasks) >= self.max_peers: + break + if new_client is not None: + tasks.append(try_connect(new_client)) + result = await asyncio.gather(*tasks, return_exceptions=True) + self.logger.debug(f'Got {sum([1 for r in result if r is True])} more peers') + await asyncio.sleep(10) diff --git a/pytoniq/adnl/overlay/privacy.py b/pytoniq/adnl/overlay/privacy.py new file mode 100644 index 0000000..d010499 --- /dev/null +++ b/pytoniq/adnl/overlay/privacy.py @@ -0,0 +1,109 @@ +import time +import typing +import enum + +from pytoniq_core.tl import TlSchemas, TlGenerator +from pytoniq_core.crypto.signature import verify_sign + + +class BroadcastCheckResult(enum.Enum): + Forbidden = 1 + NeedCheck = 2 + Allowed = 3 + + +class OverlayPrivacyRules: + + def __init__( + self, + max_unauth_size: int, + flags: int, + authorized_keys: typing.Dict[bytes, int] + ): + """ + :param max_unauth_size: maximum broadcast without authorization bytes length + :param flags: + :param authorized_keys: {key: max_size} + """ + self.max_unauth_size = max_unauth_size + self.flags = flags + self.authorized_keys = authorized_keys + + @classmethod + def default(cls, allow_fec: bool): + from .overlay import OverlayTransport + return cls( + max_unauth_size=OverlayTransport.max_fec_broadcast_size, + flags=int(allow_fec), + authorized_keys={} + ) + + @property + def allow_fec(self) -> bool: + return bool(self.flags & 1) + + def check_rules(self, key_id: bytes, size: int, is_fec: bool) -> BroadcastCheckResult: + if key_id not in self.authorized_keys: + if size > self.max_unauth_size: + return BroadcastCheckResult.Forbidden + if not(self.flags & 1) and is_fec: + return BroadcastCheckResult.Forbidden + return BroadcastCheckResult.Allowed if self.flags & 2 else BroadcastCheckResult.NeedCheck + return BroadcastCheckResult.Allowed if size <= self.authorized_keys[key_id] else BroadcastCheckResult.Forbidden + + +class InvalidCertificate(Exception): + pass + + +class Certificate: + def __init__( + self, + data: dict, + schemes: TlSchemas = None + ): + self._data = data + if schemes is None: + schemes = TlGenerator.with_default_schemas().generate() + self._schemes = schemes + + @property + def flags(self): + return self._data.get('flags', self.get_cert_default_flags(self._data['max_size'])) + + @staticmethod + def get_cert_default_flags(max_size: int) -> int: + from .overlay import OverlayTransport + return (1 if max_size > OverlayTransport.max_simple_broadcast_size else 0) | 2 # allowFec if max_size > 768 else 0 + + def to_sign(self, overlay_id: bytes, issued_to: bytes) -> bytes: + """ + overlay.certificateId overlay_id:int256 node:int256 expire_at:int max_size:int = overlay.CertificateId; + overlay.certificateIdV2 overlay_id:int256 node:int256 expire_at:int max_size:int flags:int = overlay.CertificateId; + """ + if self.flags == self.get_cert_default_flags(self._data['max_size']): + data = {'overlay_id': overlay_id.hex(), 'node': issued_to.hex(), 'expire_at': self._data['expire_at'], 'max_size': self._data['max_size']} + return self._schemes.serialize('overlay.certificateId', data) + data = {'overlay_id': overlay_id.hex(), 'node': issued_to.hex(), 'expire_at': self._data['expire_at'], 'max_size': self._data['max_size'], 'flags': self._data['flags']} + return self._schemes.serialize('overlay.certificateIdV2', data) + + def check(self, node_id: bytes, overlay_id: bytes, size: int, is_fec: bool) -> BroadcastCheckResult: + """ + overlay.certificateV2 issued_by:PublicKey expire_at:int max_size:int flags:int signature:bytes = overlay.Certificate; + """ + if size > self._data['max_size']: + return BroadcastCheckResult.Forbidden + + if time.time() > self._data['expire_at']: + return BroadcastCheckResult.Forbidden + + if is_fec and not(self.flags & 1): + return BroadcastCheckResult.Forbidden + + pub_key = bytes.fromhex(self._data['issued_by']['key']) + + to_sign = self.to_sign(overlay_id, node_id) + if not verify_sign(pub_key, to_sign, self._data['signature']): + return BroadcastCheckResult.Forbidden + + return BroadcastCheckResult.Allowed if self.flags & 2 else BroadcastCheckResult.NeedCheck diff --git a/pytoniq/adnl/overlay/shard_overlay.py b/pytoniq/adnl/overlay/shard_overlay.py new file mode 100644 index 0000000..b57dc29 --- /dev/null +++ b/pytoniq/adnl/overlay/shard_overlay.py @@ -0,0 +1,57 @@ +import typing + +from pytoniq_core.tl import BlockIdExt + +from .overlay import OverlayTransport, OverlayNode +from .overlay_manager import OverlayManager +from .broadcast import BroadcastSimple +from .fec_broadcast import create_fec_broadcast + + +class ShardOverlay: + + def __init__( + self, + overlay_manager: OverlayManager, + external_messages_handler: typing.Callable = None, + blocks_handler: typing.Callable = None, + shard_blocks_handler: typing.Callable = None, + ): + self._overlay: OverlayTransport = overlay_manager.overlay + self._manager = overlay_manager + self.external_messages_disabled = False + self.external_messages_handler = external_messages_handler # todo: we can emulate externals, but need to store state + self.blocks_handler = blocks_handler # todo: we can check blocks via key blocks and validator signatures + self.shard_blocks_handler = shard_blocks_handler + self.init_handlers() + + def init_handlers(self): + if self.external_messages_handler is not None: + self._overlay.set_broadcast_handler('tonNode.externalMessageBroadcast', self.external_messages_handler) + if self.blocks_handler is not None: + self._overlay.set_broadcast_handler('tonNode.newShardBlockBroadcast', self.shard_blocks_handler) + self._overlay.set_broadcast_handler('tonNode.blockBroadcast', self.blocks_handler) + + async def send_external_message(self, message: bytes): + data = {'@type': 'tonNode.externalMessageBroadcast', 'message': { + 'data': message, + '@type': 'tonNode.externalMessage'}} + query = self._overlay.schemas.serialize('tonNode.externalMessageBroadcast', data) + if len(query) < self._overlay.max_simple_broadcast_size: + b = BroadcastSimple.create(self._overlay, query, 0) + await b.run() + else: + await create_fec_broadcast(self._overlay, query, 1) + + async def raw_download_block(self, block: BlockIdExt, peer: OverlayNode) -> bytes: + """ + :param block: + :param peer: + :return: block boc + """ + return (await self._overlay.send_query_message(tl_schema_name='tonNode.downloadBlock', + data={'block': block.to_dict()}, peer=peer))[0] + + async def prepare_block(self, block: BlockIdExt, peer: OverlayNode) -> dict: + return (await self._overlay.send_query_message(tl_schema_name='tonNode.prepareBlock', + data={'block': block.to_dict()}, peer=peer))[0] diff --git a/pytoniq/adnl/rldp/__init__.py b/pytoniq/adnl/rldp/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/pytoniq/adnl/rldp/raptorq.py b/pytoniq/adnl/rldp/raptorq.py new file mode 100644 index 0000000..39b483a --- /dev/null +++ b/pytoniq/adnl/rldp/raptorq.py @@ -0,0 +1,20 @@ +SYMBOL_SIZE = 768 + + +def _import_raptorq(): + try: + import pyraptorq + except ImportError: + raise ImportError('pyraptorq library is required to use RLDP, use command: `pip install "pytoniq[rldp]"`') + + +def get_encoder(engine, data: bytes, symbol_size: int): + _import_raptorq() + from pyraptorq import Encoder + return Encoder(data, symbol_size, engine) + + +def get_decoder(engine, data_len: int, symbol_size: int, symbols_count: int): + _import_raptorq() + from pyraptorq import Decoder + return Decoder(symbols_count, symbol_size, data_len, engine) diff --git a/pytoniq/liteclient/balancer.py b/pytoniq/liteclient/balancer.py index c5d6449..141e594 100644 --- a/pytoniq/liteclient/balancer.py +++ b/pytoniq/liteclient/balancer.py @@ -8,7 +8,7 @@ from pytoniq_core import BlockIdExt, Block, Address, Account, ShardAccount, SimpleAccount, ShardDescr, Transaction, Cell from pytoniq_core.tlb.block import BinTree -from .client import LiteClient, LiteClientError +from .client import LiteClient, LiteClientError, LiteServerError class BalancerError(LiteClientError): @@ -49,12 +49,11 @@ def alive_peers_num(self): def archival_peers_num(self): return len(self._archival_peers) - @property def last_mc_block(self): seqno = self._find_consensus_block() for p in self._peers: - if p.last_mc_block.seqno == seqno: + if p.last_mc_block is not None and p.last_mc_block.seqno == seqno: return p.last_mc_block return None @@ -161,6 +160,7 @@ async def _check_peers(self): self._alive_peers.add(i) else: self._alive_peers.discard(i) + continue ping_res = await self._ping_peer(client) if ping_res: self._alive_peers.add(i) @@ -222,6 +222,8 @@ def _update_mc_seqnos(self): def _find_consensus_block(self): self._update_mc_seqnos() seqnos = sorted(self._mc_blocks.values(), reverse=True) + if not seqnos: + return 0 return seqnos[len(seqnos) * 2 // 3] # block that knows at least 2/3 liteservers def _delete_unsync_peers(self): @@ -233,8 +235,10 @@ def _delete_unsync_peers(self): async def execute_method(self, method_name_: str, *args, **kwargs) -> typing.Union[dict, typing.Any]: only_archive = kwargs.pop('only_archive', False) choose_random = kwargs.pop('choose_random', False) - - for _ in range(self.max_retries): + retry = False + i = 0 + while i < self.max_retries or retry: + retry = False if not len(self._alive_peers): raise BalancerError(f'have no alive peers') @@ -267,6 +271,16 @@ async def execute_method(self, method_name_: str, *args, **kwargs) -> typing.Uni self._update_average_request_time(ind, self.timeout * 10**6) # provide milliseconds self._alive_peers.discard(ind) continue + except LiteServerError as e: + if e.message == 'timeout': + self._update_average_request_time(ind, self.timeout * 10 ** 6) + self._alive_peers.discard(ind) + continue + raise e + except ConnectionError: # if socket is dead we just try another peer and somewhere in future will reconnect to this one + self._alive_peers.discard(ind) + retry = True + continue finally: self._current_req_num[ind] -= 1 raise asyncio.TimeoutError() diff --git a/pytoniq/liteclient/client.py b/pytoniq/liteclient/client.py index 58f8693..9884c65 100644 --- a/pytoniq/liteclient/client.py +++ b/pytoniq/liteclient/client.py @@ -118,22 +118,33 @@ def encrypt(self, data: bytes) -> bytes: def decrypt(self, data: bytes) -> bytes: return aes_ctr_decrypt(self.dec_sipher, data) + async def _drain(self): + try: + await self.writer.drain() + except ConnectionError: + await self.close() + raise + async def send(self, data: bytes, qid: typing.Union[str, int, None]) -> asyncio.Future: future = self.loop.create_future() self.writer.write(data) - await self.writer.drain() + await self._drain() self.tasks[qid] = future return future async def send_and_encrypt(self, data: bytes, qid: str) -> asyncio.Future: future = self.loop.create_future() self.writer.write(self.encrypt(data)) - await self.writer.drain() + await self._drain() self.tasks[qid] = future return future async def receive(self, data_len: int) -> bytes: - data = await self.reader.readexactly(data_len) + try: + data = await self.reader.readexactly(data_len) + except ConnectionError: + await self.close() + raise return data async def receive_and_decrypt(self, data_len: int) -> bytes: @@ -188,14 +199,20 @@ async def reconnect(self) -> None: async def close(self) -> None: for i in [self.pinger, self.updater, self.listener]: + if i.done() and not i.cancelled(): + i.exception() i.cancel() while not i.done(): - await asyncio.sleep(0.001) + await asyncio.sleep(0) self.inited = False self.tasks = {} self.reader = None - self.writer.close() - await self.writer.wait_closed() + if self.writer: + self.writer.close() + try: + await self.writer.wait_closed() + except ConnectionError: + pass self.writer = None self.logger.info(msg='client has been closed') diff --git a/setup.py b/setup.py index 07e310d..fd8ace6 100644 --- a/setup.py +++ b/setup.py @@ -5,7 +5,7 @@ setuptools.setup( name="pytoniq", - version="0.1.38", + version="0.1.39", author="Maksim Kurbatov", author_email="cyrbatoff@gmail.com", description="TON Blockchain SDK", @@ -28,5 +28,6 @@ ], extras_require={ 'tvm': ['pytvm>=0.0.11'], + 'rldp': ["pyraptorq>=0.1.2"] } )