mirror of
https://github.com/Chia-Network/chia-blockchain.git
synced 2026-08-29 10:06:27 -05:00
* Eliminate rate limits and bans for exempt peer networks * Add some tests for banning when closing connections * make sure to pass in parameter to constructor
781 lines
34 KiB
Python
781 lines
34 KiB
Python
from __future__ import annotations
|
|
|
|
import asyncio
|
|
import logging
|
|
import math
|
|
import time
|
|
import traceback
|
|
from collections.abc import Awaitable, Callable
|
|
from dataclasses import dataclass, field
|
|
from ipaddress import IPv4Network, IPv6Network
|
|
from typing import Any
|
|
|
|
from aiohttp import ClientSession, WebSocketError, WSCloseCode, WSMessage, WSMsgType
|
|
from aiohttp.client import ClientWebSocketResponse
|
|
from aiohttp.web import WebSocketResponse
|
|
from chia_rs.sized_bytes import bytes32
|
|
from chia_rs.sized_ints import int16, uint8, uint16
|
|
from packaging.version import Version
|
|
from typing_extensions import Protocol, final
|
|
|
|
from chia import __version__
|
|
from chia.protocols.outbound_message import Message, NodeType, make_msg
|
|
from chia.protocols.protocol_message_types import ProtocolMessageTypes
|
|
from chia.protocols.protocol_state_machine import message_response_ok
|
|
from chia.protocols.protocol_timing import (
|
|
API_EXCEPTION_BAN_SECONDS,
|
|
CONSENSUS_ERROR_BAN_SECONDS,
|
|
INTERNAL_PROTOCOL_ERROR_BAN_SECONDS,
|
|
RATE_LIMITER_BAN_SECONDS,
|
|
)
|
|
from chia.protocols.shared_protocol import Capability, Error, Handshake, protocol_version
|
|
from chia.server.api_protocol import ApiMetadata, ApiProtocol
|
|
from chia.server.capabilities import known_active_capabilities
|
|
from chia.server.rate_limits import RateLimiter
|
|
from chia.types.peer_info import PeerInfo
|
|
from chia.util.errors import ApiError, ConsensusError, Err, ProtocolError, TimestampError
|
|
from chia.util.log_exceptions import log_exceptions
|
|
|
|
# Each message is prepended with LENGTH_BYTES bytes specifying the length
|
|
from chia.util.network import is_in_network, is_localhost
|
|
from chia.util.streamable import Streamable
|
|
from chia.util.task_referencer import create_referenced_task
|
|
|
|
# Max size 2^(8*4) which is around 4GiB
|
|
LENGTH_BYTES: int = 4
|
|
|
|
WebSocket = WebSocketResponse | ClientWebSocketResponse
|
|
ConnectionCallback = Callable[["WSChiaConnection"], Awaitable[None]]
|
|
|
|
error_response_version = Version("0.0.35")
|
|
|
|
|
|
def create_default_last_message_time_dict() -> dict[ProtocolMessageTypes, float]:
|
|
return {message_type: -math.inf for message_type in ProtocolMessageTypes}
|
|
|
|
|
|
class ConnectionClosedCallbackProtocol(Protocol):
|
|
async def __call__(
|
|
self,
|
|
connection: WSChiaConnection,
|
|
ban_time: int,
|
|
closed_connection: bool = ...,
|
|
) -> None: ...
|
|
|
|
|
|
@final
|
|
@dataclass
|
|
class WSChiaConnection:
|
|
"""
|
|
Represents a connection to another node. Local host and port are ours, while peer host and
|
|
port are the host and port of the peer that we are connected to. Node_id and connection_type are
|
|
set after the handshake is performed in this connection.
|
|
"""
|
|
|
|
ws: WebSocket = field(repr=False)
|
|
api: ApiProtocol = field(repr=False)
|
|
local_type: NodeType
|
|
local_port: int | None
|
|
local_capabilities_for_handshake: list[tuple[uint16, str]] = field(repr=False)
|
|
local_capabilities: list[Capability]
|
|
peer_info: PeerInfo
|
|
peer_node_id: bytes32
|
|
log: logging.Logger = field(repr=False)
|
|
|
|
close_callback: ConnectionClosedCallbackProtocol | None = field(repr=False)
|
|
outbound_rate_limiter: RateLimiter
|
|
inbound_rate_limiter: RateLimiter
|
|
stub_metadata_for_type: dict[NodeType, ApiMetadata] = field(repr=False)
|
|
|
|
# connection properties
|
|
is_outbound: bool
|
|
|
|
# Messaging
|
|
received_message_callback: ConnectionCallback | None = field(repr=False)
|
|
incoming_queue: asyncio.Queue[Message] = field(default_factory=asyncio.Queue, repr=False)
|
|
outgoing_queue: asyncio.Queue[Message] = field(default_factory=asyncio.Queue, repr=False)
|
|
api_tasks: dict[bytes32, asyncio.Task[None]] = field(default_factory=dict, repr=False)
|
|
# Contains task ids of api tasks which should not be canceled
|
|
execute_tasks: set[bytes32] = field(default_factory=set, repr=False)
|
|
|
|
# ChiaConnection metrics
|
|
creation_time: float = field(default_factory=time.time)
|
|
bytes_read: int = 0
|
|
bytes_written: int = 0
|
|
last_message_time: float = 0
|
|
|
|
peer_server_port: uint16 | None = None
|
|
inbound_task: asyncio.Task[None] | None = field(default=None, repr=False)
|
|
incoming_message_task: asyncio.Task[None] | None = field(default=None, repr=False)
|
|
outbound_task: asyncio.Task[None] | None = field(default=None, repr=False)
|
|
_close_event: asyncio.Event = field(default_factory=asyncio.Event, repr=False)
|
|
session: ClientSession | None = field(default=None, repr=False)
|
|
|
|
pending_requests: dict[uint16, asyncio.Event] = field(default_factory=dict, repr=False)
|
|
request_results: dict[uint16, Message] = field(default_factory=dict, repr=False)
|
|
closed: bool = False
|
|
connection_type: NodeType | None = None
|
|
request_nonce: uint16 = uint16(0)
|
|
peer_capabilities: list[Capability] = field(default_factory=list)
|
|
# Used by the Chia Seeder.
|
|
version: str = field(default_factory=str)
|
|
protocol_version: Version = field(default_factory=lambda: Version("0"))
|
|
|
|
log_rate_limit_last_time: dict[ProtocolMessageTypes, float] = field(
|
|
default_factory=create_default_last_message_time_dict,
|
|
repr=False,
|
|
)
|
|
|
|
exempt_peer_networks: list[IPv4Network | IPv6Network] = field(
|
|
default_factory=list,
|
|
repr=False,
|
|
)
|
|
|
|
@classmethod
|
|
def create(
|
|
cls,
|
|
local_type: NodeType,
|
|
ws: WebSocket,
|
|
api: ApiProtocol,
|
|
server_port: int | None,
|
|
log: logging.Logger,
|
|
is_outbound: bool,
|
|
received_message_callback: ConnectionCallback | None,
|
|
close_callback: ConnectionClosedCallbackProtocol | None,
|
|
peer_id: bytes32,
|
|
inbound_rate_limit_percent: int,
|
|
outbound_rate_limit_percent: int,
|
|
local_capabilities_for_handshake: list[tuple[uint16, str]],
|
|
stub_metadata_for_type: dict[NodeType, ApiMetadata],
|
|
session: ClientSession | None = None,
|
|
exempt_peer_networks: list[IPv4Network | IPv6Network] = [],
|
|
) -> WSChiaConnection:
|
|
assert ws._writer is not None
|
|
peername = ws._writer.transport.get_extra_info("peername")
|
|
|
|
if peername is None:
|
|
raise ValueError(f"Was not able to get peername for {peer_id}")
|
|
|
|
if is_outbound:
|
|
request_nonce = uint16(0)
|
|
else:
|
|
# Different nonce to reduce chances of overlap. Each peer will increment the nonce by one for each
|
|
# request. The receiving peer (not is_outbound), will use 2^15 to 2^16 - 1
|
|
request_nonce = uint16(2**15)
|
|
|
|
return cls(
|
|
ws=ws,
|
|
api=api,
|
|
local_type=local_type,
|
|
local_port=server_port,
|
|
local_capabilities_for_handshake=local_capabilities_for_handshake,
|
|
local_capabilities=known_active_capabilities(local_capabilities_for_handshake),
|
|
peer_info=PeerInfo(peername[0], peername[1]),
|
|
peer_node_id=peer_id,
|
|
log=log,
|
|
close_callback=close_callback,
|
|
request_nonce=request_nonce,
|
|
outbound_rate_limiter=RateLimiter(incoming=False, percentage_of_limit=outbound_rate_limit_percent),
|
|
inbound_rate_limiter=RateLimiter(incoming=True, percentage_of_limit=inbound_rate_limit_percent),
|
|
is_outbound=is_outbound,
|
|
received_message_callback=received_message_callback,
|
|
stub_metadata_for_type=stub_metadata_for_type,
|
|
session=session,
|
|
exempt_peer_networks=exempt_peer_networks,
|
|
)
|
|
|
|
def _get_extra_info(self, name: str) -> Any | None:
|
|
writer = self.ws._writer
|
|
assert writer is not None, "websocket's ._writer is None, was .prepare() called?"
|
|
transport = writer.transport
|
|
if transport is None:
|
|
return None
|
|
try:
|
|
return transport.get_extra_info(name)
|
|
except AttributeError:
|
|
# "/usr/lib/python3.11/asyncio/sslproto.py", line 91, in get_extra_info
|
|
# return self._ssl_protocol._get_extra_info(name, default)
|
|
# AttributeError: 'NoneType' object has no attribute '_get_extra_info'
|
|
return None
|
|
|
|
async def perform_handshake(
|
|
self,
|
|
network_id: str,
|
|
server_port: int,
|
|
local_type: NodeType,
|
|
) -> None:
|
|
if self.is_outbound:
|
|
outbound_handshake = make_msg(
|
|
ProtocolMessageTypes.handshake,
|
|
Handshake(
|
|
network_id,
|
|
protocol_version[local_type],
|
|
__version__,
|
|
uint16(server_port),
|
|
uint8(local_type.value),
|
|
self.local_capabilities_for_handshake,
|
|
),
|
|
)
|
|
await self._send_message(outbound_handshake)
|
|
inbound_handshake_msg = await self._read_one_message()
|
|
if inbound_handshake_msg is None:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
inbound_handshake = Handshake.from_bytes(inbound_handshake_msg.data)
|
|
|
|
# Handle case of invalid ProtocolMessageType
|
|
try:
|
|
message_type: ProtocolMessageTypes = ProtocolMessageTypes(inbound_handshake_msg.type)
|
|
except Exception:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
|
|
if message_type != ProtocolMessageTypes.handshake:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
|
|
if inbound_handshake.network_id != network_id:
|
|
raise ProtocolError(Err.INCOMPATIBLE_NETWORK_ID)
|
|
|
|
if (
|
|
local_type in {NodeType.FARMER, NodeType.HARVESTER}
|
|
and inbound_handshake.protocol_version != protocol_version[local_type]
|
|
):
|
|
self.log.warning(
|
|
f"protocol version mismatch: "
|
|
f"local_type={local_type} "
|
|
f"incoming={inbound_handshake.protocol_version} "
|
|
f"our={protocol_version[local_type]}"
|
|
)
|
|
|
|
self.version = inbound_handshake.software_version
|
|
self.protocol_version = Version(inbound_handshake.protocol_version)
|
|
self.peer_server_port = inbound_handshake.server_port
|
|
self.connection_type = NodeType(inbound_handshake.node_type)
|
|
# "1" means capability is enabled
|
|
self.peer_capabilities = known_active_capabilities(inbound_handshake.capabilities)
|
|
else:
|
|
try:
|
|
message = await self._read_one_message()
|
|
except Exception:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
|
|
if message is None:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
|
|
# Handle case of invalid ProtocolMessageType
|
|
try:
|
|
message_type = ProtocolMessageTypes(message.type)
|
|
except Exception:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
|
|
if message_type != ProtocolMessageTypes.handshake:
|
|
raise ProtocolError(Err.INVALID_HANDSHAKE)
|
|
|
|
inbound_handshake = Handshake.from_bytes(message.data)
|
|
if inbound_handshake.network_id != network_id:
|
|
raise ProtocolError(Err.INCOMPATIBLE_NETWORK_ID)
|
|
|
|
remote_node_type = NodeType(inbound_handshake.node_type)
|
|
|
|
if (
|
|
remote_node_type in {NodeType.FARMER, NodeType.HARVESTER}
|
|
and inbound_handshake.protocol_version != protocol_version[remote_node_type]
|
|
):
|
|
self.log.warning(
|
|
f"protocol version mismatch: "
|
|
f"remote_type={remote_node_type} "
|
|
f"incoming={inbound_handshake.protocol_version} "
|
|
f"our={protocol_version[remote_node_type]}"
|
|
)
|
|
|
|
outbound_handshake = make_msg(
|
|
ProtocolMessageTypes.handshake,
|
|
Handshake(
|
|
network_id,
|
|
protocol_version[remote_node_type],
|
|
__version__,
|
|
uint16(server_port),
|
|
uint8(local_type.value),
|
|
self.local_capabilities_for_handshake,
|
|
),
|
|
)
|
|
await self._send_message(outbound_handshake)
|
|
self.version = inbound_handshake.software_version
|
|
self.protocol_version = Version(inbound_handshake.protocol_version)
|
|
self.peer_server_port = inbound_handshake.server_port
|
|
self.connection_type = remote_node_type
|
|
# "1" means capability is enabled
|
|
self.peer_capabilities = known_active_capabilities(inbound_handshake.capabilities)
|
|
|
|
self.outbound_task = create_referenced_task(self.outbound_handler())
|
|
self.inbound_task = create_referenced_task(self.inbound_handler())
|
|
self.incoming_message_task = create_referenced_task(self.incoming_message_handler())
|
|
|
|
async def close(
|
|
self,
|
|
ban_time: int = 0,
|
|
ws_close_code: WSCloseCode = WSCloseCode.OK,
|
|
error: Err | None = None,
|
|
) -> None:
|
|
"""
|
|
Closes the connection, and finally calls the close_callback on the server, so the connection gets removed
|
|
from the global list.
|
|
"""
|
|
if self.closed:
|
|
# always try to call the callback even for closed connections
|
|
with log_exceptions(self.log, consume=True):
|
|
self.log.debug(f"Closing already closed connection for {self.peer_info.host}")
|
|
if self.close_callback is not None:
|
|
await self.close_callback(self, ban_time, closed_connection=True)
|
|
self._close_event.set()
|
|
return None
|
|
self.closed = True
|
|
|
|
if error is None:
|
|
message = b""
|
|
else:
|
|
message = str(int(error.value)).encode("utf-8")
|
|
|
|
try:
|
|
if self.inbound_task is not None:
|
|
self.inbound_task.cancel()
|
|
if self.incoming_message_task is not None:
|
|
self.incoming_message_task.cancel()
|
|
if self.outbound_task is not None:
|
|
self.outbound_task.cancel()
|
|
if self.ws is not None and self.ws.closed is False:
|
|
await self.ws.close(code=ws_close_code, message=message)
|
|
if self.session is not None:
|
|
await self.session.close()
|
|
self.cancel_pending_requests()
|
|
self.cancel_tasks()
|
|
except Exception:
|
|
error_stack = traceback.format_exc()
|
|
self.log.warning(f"Exception closing socket: {error_stack}")
|
|
raise
|
|
finally:
|
|
with log_exceptions(self.log, consume=True):
|
|
if self.close_callback is not None:
|
|
await self.close_callback(self, ban_time, closed_connection=False)
|
|
self._close_event.set()
|
|
|
|
async def wait_until_closed(self) -> None:
|
|
await self._close_event.wait()
|
|
|
|
async def ban_peer_bad_protocol(self, log_err_msg: str) -> None:
|
|
"""Ban peer for protocol violation"""
|
|
ban_seconds = INTERNAL_PROTOCOL_ERROR_BAN_SECONDS
|
|
self.log.error(f"Banning peer for {ban_seconds} seconds: {self.peer_info.host} {log_err_msg}")
|
|
await self.close(ban_seconds, WSCloseCode.PROTOCOL_ERROR, Err.INVALID_PROTOCOL_MESSAGE)
|
|
|
|
def cancel_pending_requests(self) -> None:
|
|
for message_id, event in self.pending_requests.items():
|
|
try:
|
|
event.set()
|
|
except Exception as e:
|
|
self.log.error(f"Failed setting event for {message_id}: {e} {traceback.format_exc()}")
|
|
|
|
def cancel_tasks(self) -> None:
|
|
for task_id, task in self.api_tasks.copy().items():
|
|
if task_id in self.execute_tasks:
|
|
continue
|
|
task.cancel()
|
|
|
|
async def outbound_handler(self) -> None:
|
|
try:
|
|
while not self.closed:
|
|
msg = await self.outgoing_queue.get()
|
|
if msg is not None:
|
|
await self._send_message(msg)
|
|
except asyncio.CancelledError:
|
|
pass
|
|
except Exception as e:
|
|
expected_types = (BrokenPipeError, ConnectionResetError, TimeoutError)
|
|
expected = False
|
|
if isinstance(e, expected_types) or isinstance(e.__cause__, expected_types):
|
|
expected = True
|
|
elif isinstance(e, OSError):
|
|
if e.errno in {113}:
|
|
expected = True
|
|
|
|
if expected:
|
|
self.log.warning(f"{e} {self.peer_info.host}")
|
|
else:
|
|
error_stack = traceback.format_exc()
|
|
self.log.error(f"Exception: {e} with {self.peer_info.host}")
|
|
self.log.error(f"Exception Stack: {error_stack}")
|
|
|
|
async def _api_call(self, full_message: Message, task_id: bytes32) -> None:
|
|
start_time = time.time()
|
|
message_type = ""
|
|
try:
|
|
if self.received_message_callback is not None:
|
|
await self.received_message_callback(self)
|
|
self.log.debug(
|
|
f"<- {ProtocolMessageTypes(full_message.type).name} from peer {self.peer_node_id} {self.peer_info.host}"
|
|
)
|
|
|
|
if full_message.type == ProtocolMessageTypes.error.value:
|
|
error = Error.from_bytes(full_message.data)
|
|
self.log.warning(f"ApiError: {error} from {self.peer_node_id}, {self.peer_info}")
|
|
return None
|
|
|
|
bare_message_type = ProtocolMessageTypes(full_message.type)
|
|
metadata = self.api.metadata.message_type_to_request.get(bare_message_type)
|
|
message_type = bare_message_type.name
|
|
|
|
if metadata is None:
|
|
self.log.error(f"Non existing function: {message_type}")
|
|
raise ProtocolError(Err.INVALID_PROTOCOL_MESSAGE, [message_type])
|
|
|
|
if metadata is None:
|
|
self.log.error(f"Peer trying to call non api function {message_type}")
|
|
raise ProtocolError(Err.INVALID_PROTOCOL_MESSAGE, [message_type])
|
|
|
|
# If api is not ready ignore the request
|
|
if not self.api.ready():
|
|
self.log.warning(f"API not ready, ignore request: {full_message}")
|
|
return None
|
|
|
|
timeout: int | None = 600
|
|
if metadata.execute_task:
|
|
# Don't timeout on methods with execute_task decorator, these need to run fully
|
|
self.execute_tasks.add(task_id)
|
|
timeout = None
|
|
|
|
if metadata.peer_required:
|
|
coroutine = metadata.method(self.api, full_message.data, self)
|
|
else:
|
|
coroutine = metadata.method(self.api, full_message.data)
|
|
|
|
async def wrapped_coroutine() -> Message | None:
|
|
try:
|
|
result = await coroutine
|
|
return result
|
|
except asyncio.CancelledError:
|
|
pass
|
|
except ApiError as api_error:
|
|
self.log.warning(f"ApiError: {api_error} from {self.peer_node_id}, {self.peer_info}")
|
|
if self.protocol_version >= error_response_version:
|
|
return make_msg(
|
|
ProtocolMessageTypes.error,
|
|
Error(int16(api_error.code.value), api_error.message, api_error.data),
|
|
)
|
|
else:
|
|
return None
|
|
except TimestampError:
|
|
raise
|
|
except Exception as e:
|
|
tb = traceback.format_exc()
|
|
self.log.error(f"Exception: {e}, {self.get_peer_logging()}. {tb}")
|
|
raise
|
|
return None
|
|
|
|
response: Message | None = await asyncio.wait_for(wrapped_coroutine(), timeout=timeout)
|
|
self.log.debug(
|
|
f"Time taken to process {message_type} from {self.peer_node_id} is {time.time() - start_time} seconds"
|
|
)
|
|
|
|
if response is not None:
|
|
response_message = Message(response.type, full_message.id, response.data)
|
|
await self.send_message(response_message)
|
|
# todo uncomment when enabling none response capability
|
|
# check that this call needs a reply
|
|
# elif message_requires_reply(ProtocolMessageTypes(full_message.type)) and self.has_capability(
|
|
# Capability.NONE_RESPONSE
|
|
# ):
|
|
# # this peer can accept None reply's, send empty msg back, so it doesn't wait for timeout
|
|
# response_message = Message(uint8(ProtocolMessageTypes.none_response.value), full_message.id, b"")
|
|
# await self.send_message(response_message)
|
|
except TimeoutError:
|
|
self.log.error(f"Timeout error for: {message_type}")
|
|
except TimestampError:
|
|
self.log.info("Received block with timestamp too far into the future")
|
|
except Exception as e:
|
|
if not self.closed:
|
|
tb = traceback.format_exc()
|
|
self.log.error(f"Exception: {e} {type(e)}, closing connection {self.get_peer_logging()}. {tb}")
|
|
else:
|
|
self.log.debug(f"Exception: {e} while closing connection")
|
|
if isinstance(e, ConsensusError):
|
|
ban_time = CONSENSUS_ERROR_BAN_SECONDS
|
|
else:
|
|
ban_time = API_EXCEPTION_BAN_SECONDS
|
|
# TODO: actually throw one of the errors from errors.py and pass this to close
|
|
await self.close(ban_time, WSCloseCode.PROTOCOL_ERROR, Err.UNKNOWN)
|
|
finally:
|
|
if task_id in self.api_tasks:
|
|
self.api_tasks.pop(task_id)
|
|
if task_id in self.execute_tasks:
|
|
self.execute_tasks.remove(task_id)
|
|
|
|
async def incoming_message_handler(self) -> None:
|
|
while True:
|
|
message = await self.incoming_queue.get()
|
|
task_id: bytes32 = bytes32.secret()
|
|
api_task = create_referenced_task(self._api_call(message, task_id))
|
|
self.api_tasks[task_id] = api_task
|
|
|
|
async def inbound_handler(self) -> None:
|
|
try:
|
|
while not self.closed:
|
|
message = await self._read_one_message()
|
|
if message is not None:
|
|
if message.id in self.pending_requests:
|
|
self.request_results[message.id] = message
|
|
event = self.pending_requests[message.id]
|
|
event.set()
|
|
else:
|
|
await self.incoming_queue.put(message)
|
|
else:
|
|
continue
|
|
except asyncio.CancelledError:
|
|
self.log.debug("Inbound_handler task cancelled")
|
|
except Exception as e:
|
|
error_stack = traceback.format_exc()
|
|
self.log.error(f"Exception: {e}")
|
|
self.log.error(f"Exception Stack: {error_stack}")
|
|
|
|
async def send_message(self, message: Message) -> bool:
|
|
"""Send message sends a message with no tracking / callback."""
|
|
if self.closed:
|
|
return False
|
|
await self.outgoing_queue.put(message)
|
|
return True
|
|
|
|
async def call_api(
|
|
self,
|
|
request_method: Callable[..., Awaitable[Message | None]],
|
|
message: Streamable,
|
|
timeout: int = 60,
|
|
) -> Any:
|
|
if self.connection_type is None:
|
|
raise ValueError("handshake not done yet")
|
|
request_metadata = ApiMetadata.from_bound_method(request_method)
|
|
assert request_metadata is not None, f"ApiMetadata unavailable for {request_method}"
|
|
if (
|
|
request_metadata.request_type
|
|
not in self.stub_metadata_for_type[self.connection_type].message_type_to_request
|
|
):
|
|
raise AttributeError(
|
|
f"Node type {self.connection_type} does not have method {request_metadata.request_type.name}"
|
|
)
|
|
|
|
request = Message(uint8(request_metadata.request_type.value), None, bytes(message))
|
|
request_start_t = time.time()
|
|
response = await self.send_request(request, timeout)
|
|
self.log.debug(
|
|
f"Time for request {request_metadata.request_type.name}: {self.get_peer_logging()} = "
|
|
f"{time.time() - request_start_t}, None? {response is None}"
|
|
)
|
|
# todo or response.type == ProtocolMessageTypes.none_response.value when enabling none response
|
|
if response is None or response.data == b"":
|
|
return None
|
|
sent_message_type = ProtocolMessageTypes(request.type)
|
|
recv_message_type = ProtocolMessageTypes(response.type)
|
|
if recv_message_type == ProtocolMessageTypes.error:
|
|
return Error.from_bytes(response.data)
|
|
if not message_response_ok(sent_message_type, recv_message_type):
|
|
# peer protocol violation
|
|
error_message = f"WSConnection.invoke sent message {sent_message_type.name} "
|
|
f"but received {recv_message_type.name}"
|
|
await self.ban_peer_bad_protocol(error_message)
|
|
raise ProtocolError(Err.INVALID_PROTOCOL_MESSAGE, [error_message])
|
|
|
|
recv_method = self.stub_metadata_for_type[self.local_type].message_type_to_request[recv_message_type].method
|
|
receive_metadata = ApiMetadata.from_bound_method(recv_method)
|
|
assert receive_metadata is not None, f"ApiMetadata unavailable for {recv_method}"
|
|
return receive_metadata.message_class.from_bytes(response.data)
|
|
|
|
async def send_request(self, message_no_id: Message, timeout: int) -> Message | None:
|
|
"""Sends a message and waits for a response."""
|
|
if self.closed:
|
|
return None
|
|
|
|
# We will wait for this event, it will be set either by the response, or the timeout
|
|
event = asyncio.Event()
|
|
|
|
# The request nonce is an integer between 0 and 2**16 - 1, which is used to match requests to responses
|
|
# If is_outbound, 0 <= nonce < 2^15, else 2^15 <= nonce < 2^16
|
|
request_id = self.request_nonce
|
|
if self.is_outbound:
|
|
self.request_nonce = uint16(self.request_nonce + 1) if self.request_nonce != (2**15 - 1) else uint16(0)
|
|
else:
|
|
self.request_nonce = uint16(self.request_nonce + 1) if self.request_nonce != (2**16 - 1) else uint16(2**15)
|
|
|
|
message = Message(message_no_id.type, request_id, message_no_id.data)
|
|
assert message.id is not None
|
|
self.pending_requests[message.id] = event
|
|
await self.outgoing_queue.put(message)
|
|
|
|
try:
|
|
await asyncio.wait_for(event.wait(), timeout=timeout)
|
|
except asyncio.TimeoutError:
|
|
self.log.debug(f"Request timeout: {message}")
|
|
|
|
self.pending_requests.pop(message.id)
|
|
result: Message | None = None
|
|
if message.id in self.request_results:
|
|
result = self.request_results[message.id]
|
|
assert result is not None
|
|
self.log.debug(
|
|
f"<- {ProtocolMessageTypes(result.type).name} from: {self.peer_info.host}:{self.peer_info.port}"
|
|
)
|
|
self.request_results.pop(message.id)
|
|
|
|
return result
|
|
|
|
async def _wait_and_retry(self, msg: Message) -> None:
|
|
try:
|
|
await asyncio.sleep(1)
|
|
await self.outgoing_queue.put(msg)
|
|
except Exception as e:
|
|
self.log.debug(f"Exception {e} while waiting to retry sending rate limited message")
|
|
return None
|
|
|
|
async def _send_message(self, message: Message) -> None:
|
|
encoded: bytes = bytes(message)
|
|
size = len(encoded)
|
|
assert len(encoded) < (2 ** (LENGTH_BYTES * 8))
|
|
limiter_msg = self.outbound_rate_limiter.process_msg_and_check(
|
|
message, self.local_capabilities, self.peer_capabilities
|
|
)
|
|
if limiter_msg is not None:
|
|
if not is_localhost(self.peer_info.host) and not is_in_network(
|
|
self.peer_info.host, self.exempt_peer_networks
|
|
):
|
|
message_type = ProtocolMessageTypes(message.type)
|
|
last_time = self.log_rate_limit_last_time[message_type]
|
|
now = time.monotonic()
|
|
if now - last_time >= 30:
|
|
self.log_rate_limit_last_time[message_type] = now
|
|
details = ", ".join(
|
|
[
|
|
f"{message_type.name}",
|
|
f"sz: {len(message.data) / 1000:0.2f} kB",
|
|
f"peer: {self.peer_info.host}",
|
|
f"{limiter_msg}",
|
|
]
|
|
)
|
|
self.log.info(f"Rate limiting ourselves. Dropping outbound message: {details}")
|
|
|
|
# TODO: fix this special case. This function has rate limits which are too low.
|
|
if ProtocolMessageTypes(message.type) != ProtocolMessageTypes.respond_peers:
|
|
create_referenced_task(self._wait_and_retry(message), known_unreferenced=True)
|
|
|
|
return None
|
|
else:
|
|
self.log.debug(
|
|
f"Not rate limiting ourselves or exempt peers. "
|
|
f"message type: {ProtocolMessageTypes(message.type).name}, "
|
|
f"peer: {self.peer_info.host}"
|
|
)
|
|
|
|
await self.ws.send_bytes(encoded)
|
|
self.log.debug(
|
|
f"-> {ProtocolMessageTypes(message.type).name} to peer {self.peer_info.host} {self.peer_node_id}"
|
|
)
|
|
self.bytes_written += size
|
|
|
|
async def _read_one_message(self) -> Message | None:
|
|
message: WSMessage = await self.ws.receive()
|
|
|
|
if self.connection_type is not None:
|
|
connection_type_str = NodeType(self.connection_type).name.lower()
|
|
else:
|
|
connection_type_str = ""
|
|
if message.type == WSMsgType.CLOSING:
|
|
self.log.debug(
|
|
f"Closing connection to {connection_type_str} {self.peer_info.host}:"
|
|
f"{self.peer_server_port}/"
|
|
f"{self.peer_info.port}"
|
|
)
|
|
create_referenced_task(self.close(), known_unreferenced=True)
|
|
await asyncio.sleep(3)
|
|
elif message.type == WSMsgType.CLOSE:
|
|
self.log.debug(
|
|
f"Peer closed connection {connection_type_str} {self.peer_info.host}:"
|
|
f"{self.peer_server_port}/"
|
|
f"{self.peer_info.port}"
|
|
)
|
|
create_referenced_task(self.close(), known_unreferenced=True)
|
|
await asyncio.sleep(3)
|
|
elif message.type == WSMsgType.CLOSED:
|
|
if not self.closed:
|
|
create_referenced_task(self.close(), known_unreferenced=True)
|
|
await asyncio.sleep(3)
|
|
return None
|
|
elif message.type == WSMsgType.BINARY:
|
|
data = message.data
|
|
full_message_loaded: Message = Message.from_bytes(data)
|
|
self.bytes_read += len(data)
|
|
self.last_message_time = time.time()
|
|
try:
|
|
message_type = ProtocolMessageTypes(full_message_loaded.type).name
|
|
except Exception:
|
|
message_type = "Unknown"
|
|
limiter_msg = self.inbound_rate_limiter.process_msg_and_check(
|
|
full_message_loaded, self.local_capabilities, self.peer_capabilities
|
|
)
|
|
if limiter_msg is not None:
|
|
if (
|
|
self.local_type == NodeType.FULL_NODE
|
|
and not is_localhost(self.peer_info.host)
|
|
and not is_in_network(self.peer_info.host, self.exempt_peer_networks)
|
|
):
|
|
details = ", ".join([f"{self.peer_info.host}", f"message: {message_type}", limiter_msg])
|
|
self.log.error(f"Peer has been rate limited and will be disconnected: {details}")
|
|
# Only full node disconnects peers, to prevent abuse and crashing timelords, farmers, etc
|
|
create_referenced_task(self.close(RATE_LIMITER_BAN_SECONDS), known_unreferenced=True)
|
|
await asyncio.sleep(3)
|
|
return None
|
|
else:
|
|
self.log.debug(
|
|
f"Peer surpassed rate limit {self.peer_info.host}, message: {message_type}, "
|
|
f"port {self.peer_info.port} but not disconnecting"
|
|
)
|
|
return full_message_loaded
|
|
return full_message_loaded
|
|
elif message.type == WSMsgType.ERROR:
|
|
self.log.error(f"WebSocket Error: {message}")
|
|
if isinstance(message.data, WebSocketError) and message.data.code == WSCloseCode.MESSAGE_TOO_BIG:
|
|
create_referenced_task(self.close(RATE_LIMITER_BAN_SECONDS), known_unreferenced=True)
|
|
else:
|
|
create_referenced_task(self.close(), known_unreferenced=True)
|
|
await asyncio.sleep(3)
|
|
|
|
else:
|
|
self.log.error(f"Unexpected WebSocket message type: {message}")
|
|
create_referenced_task(self.close())
|
|
await asyncio.sleep(3)
|
|
return None
|
|
|
|
# Used by the Chia Seeder.
|
|
def get_version(self) -> str:
|
|
return self.version
|
|
|
|
def get_tls_version(self) -> str:
|
|
ssl_obj = self._get_extra_info("ssl_object")
|
|
if ssl_obj is not None:
|
|
return str(ssl_obj.version())
|
|
else:
|
|
return "unknown"
|
|
|
|
def get_peer_info(self) -> PeerInfo | None:
|
|
result = self._get_extra_info("peername")
|
|
if result is None:
|
|
return None
|
|
connection_host = result[0]
|
|
port = self.peer_server_port if self.peer_server_port is not None else self.peer_info.port
|
|
return PeerInfo(connection_host, port)
|
|
|
|
def get_peer_logging(self) -> PeerInfo:
|
|
info: PeerInfo | None = self.get_peer_info()
|
|
if info is None:
|
|
# in this case, we will use self.peer_info.host which is friendlier for logging
|
|
port = self.peer_server_port if self.peer_server_port is not None else self.peer_info.port
|
|
return PeerInfo(self.peer_info.host, port)
|
|
else:
|
|
return info
|
|
|
|
def has_capability(self, capability: Capability) -> bool:
|
|
return capability in self.peer_capabilities
|