Files
chia-blockchain/chia/server/ws_connection.py
Earle LoweandGitHub 3461286e8d Eliminate rate limits and bans for exempt peer networks (#20345)
* 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
2025-12-15 09:46:40 -08:00

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