diff --git a/chia/_tests/core/server/test_dos.py b/chia/_tests/core/server/test_dos.py index 7188eec3cb..f8ae7aa493 100644 --- a/chia/_tests/core/server/test_dos.py +++ b/chia/_tests/core/server/test_dos.py @@ -12,12 +12,15 @@ from chia_rs.sized_ints import uint8, uint16, uint64 import chia.server.server from chia._tests.util.time_out_assert import time_out_assert from chia.protocols import full_node_protocol -from chia.protocols.outbound_message import Message, make_msg +from chia.protocols.outbound_message import Message, NodeType, make_msg from chia.protocols.protocol_message_types import ProtocolMessageTypes -from chia.protocols.shared_protocol import Capability, Handshake +from chia.protocols.shared_protocol import Capability, Handshake, protocol_version from chia.server.rate_limits import RateLimiter from chia.server.server import ChiaServer -from chia.server.ws_connection import WSChiaConnection +from chia.server.ws_connection import ( + MAX_VERSION_STRING_BYTES, + WSChiaConnection, +) from chia.simulator.block_tools import BlockTools from chia.simulator.full_node_simulator import FullNodeSimulator from chia.types.peer_info import PeerInfo @@ -39,6 +42,32 @@ class FakeRateLimiter: return None +async def send_handshake_and_assert_protocol_error( + server_1: ChiaServer, + server_2: ChiaServer, + self_hostname: str, + payload: Handshake | bytes, + message_type: int = ProtocolMessageTypes.handshake.value, +) -> WSMessage: + server_1.invalid_protocol_ban_seconds = 10 + timeout = ClientTimeout(total=5) + async with ClientSession(timeout=timeout) as session: + url = f"wss://{self_hostname}:{server_1._port}/ws" + async with session.ws_connect( + url, + autoclose=True, + autoping=True, + ssl=server_2.ssl_client_context, + max_msg_size=50 * 1024 * 1024, + ) as ws: + msg = Message(uint8(message_type), None, bytes(payload)) + await ws.send_bytes(bytes(msg)) + response = await ws.receive() + assert response.type == WSMsgType.CLOSE + assert response.data == WSCloseCode.PROTOCOL_ERROR + return response + + class TestDos: @pytest.mark.anyio async def test_banned_host_can_not_connect( @@ -141,31 +170,73 @@ class TestDos: nodes, _, _ = setup_two_nodes_fixture server_1 = nodes[0].full_node.server server_2 = nodes[1].full_node.server - - server_1.invalid_protocol_ban_seconds = 10 - # Use the server_2 ssl information to connect to server_1 - timeout = ClientTimeout(total=10) - session = ClientSession(timeout=timeout) - url = f"wss://{self_hostname}:{server_1._port}/ws" - - ssl_context = server_2.ssl_client_context - ws = await session.ws_connect( - url, autoclose=True, autoping=True, ssl=ssl_context, max_msg_size=100 * 1024 * 1024 + handshake = Handshake("test", "0.0.32", "1.0.0.0", uint16(3456), uint8(1), [(uint16(1), "1")]) + response = await send_handshake_and_assert_protocol_error( + server_1, server_2, self_hostname, handshake, message_type=2 ) - - # Construct an otherwise valid handshake message - handshake: Handshake = Handshake("test", "0.0.32", "1.0.0.0", uint16(3456), uint8(1), [(uint16(1), "1")]) - outbound_handshake: Message = Message(uint8(2), None, bytes(handshake)) # 2 is an invalid ProtocolType - await ws.send_bytes(bytes(outbound_handshake)) - - response: WSMessage = await ws.receive() - print(response) - assert response.type == WSMsgType.CLOSE - assert response.data == WSCloseCode.PROTOCOL_ERROR assert response.extra == str(int(Err.INVALID_HANDSHAKE.value)) # We want INVALID_HANDSHAKE and not UNKNOWN - await ws.close() - await session.close() - await asyncio.sleep(1) # give some time for cleanup to work + + @pytest.mark.anyio + async def test_handshake_version_too_long_disconnect( + self, + setup_two_nodes_fixture: tuple[list[FullNodeSimulator], list[tuple[WalletNode, ChiaServer]], BlockTools], + self_hostname: str, + ) -> None: + nodes, _, _ = setup_two_nodes_fixture + server_1 = nodes[0].full_node.server + server_2 = nodes[1].full_node.server + long_version = "x" * (MAX_VERSION_STRING_BYTES + 1) + handshake = Handshake( + server_1._network_id, + protocol_version[NodeType.FULL_NODE], + long_version, + uint16(3456), + uint8(NodeType.FULL_NODE.value), + [(uint16(1), "1")], + ) + await send_handshake_and_assert_protocol_error(server_1, server_2, self_hostname, handshake) + + @pytest.mark.anyio + async def test_wrong_message_type_handshake( + self, + setup_two_nodes_fixture: tuple[list[FullNodeSimulator], list[tuple[WalletNode, ChiaServer]], BlockTools], + self_hostname: str, + ) -> None: + nodes, _, _ = setup_two_nodes_fixture + server_1 = nodes[0].full_node.server + server_2 = nodes[1].full_node.server + handshake = Handshake( + server_1._network_id, + protocol_version[NodeType.FULL_NODE], + "2.6.0", + uint16(3456), + uint8(NodeType.FULL_NODE.value), + [(uint16(1), "1")], + ) + response = await send_handshake_and_assert_protocol_error( + server_1, server_2, self_hostname, handshake, message_type=ProtocolMessageTypes.new_peak.value + ) + assert response.extra == str(int(Err.INVALID_HANDSHAKE.value)) + + @pytest.mark.anyio + async def test_wrong_network_id_handshake( + self, + setup_two_nodes_fixture: tuple[list[FullNodeSimulator], list[tuple[WalletNode, ChiaServer]], BlockTools], + self_hostname: str, + ) -> None: + nodes, _, _ = setup_two_nodes_fixture + server_1 = nodes[0].full_node.server + server_2 = nodes[1].full_node.server + handshake = Handshake( + "wrong-network-id", + protocol_version[NodeType.FULL_NODE], + "2.6.0", + uint16(3456), + uint8(NodeType.FULL_NODE.value), + [(uint16(1), "1")], + ) + response = await send_handshake_and_assert_protocol_error(server_1, server_2, self_hostname, handshake) + assert response.extra == str(int(Err.INCOMPATIBLE_NETWORK_ID.value)) @pytest.mark.anyio async def test_spam_tx( diff --git a/chia/server/ws_connection.py b/chia/server/ws_connection.py index cbf78152f3..8ac95f7b85 100644 --- a/chia/server/ws_connection.py +++ b/chia/server/ws_connection.py @@ -44,6 +44,9 @@ from chia.util.task_referencer import create_referenced_task # Max size 2^(8*4) which is around 4GiB LENGTH_BYTES: int = 4 +# Max length of peer version string in bytes (UTF-8) +MAX_VERSION_STRING_BYTES: int = 128 + WebSocket = WebSocketResponse | ClientWebSocketResponse ConnectionCallback = Callable[["WSChiaConnection"], Awaitable[None]] @@ -217,64 +220,47 @@ class WSChiaConnection: ), ) 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) + try: + message = await self._read_one_message() + 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) + if message is None: + raise ProtocolError(Err.INVALID_HANDSHAKE) + # Handle case of invalid ProtocolMessageType + try: inbound_handshake = Handshake.from_bytes(message.data) - if inbound_handshake.network_id != network_id: - raise ProtocolError(Err.INCOMPATIBLE_NETWORK_ID) + message_type = ProtocolMessageTypes(message.type) + except Exception: + raise ProtocolError(Err.INVALID_HANDSHAKE) - remote_node_type = NodeType(inbound_handshake.node_type) + if message_type != ProtocolMessageTypes.handshake: + raise ProtocolError(Err.INVALID_HANDSHAKE) + if inbound_handshake.network_id != network_id: + raise ProtocolError(Err.INCOMPATIBLE_NETWORK_ID) + + if ( + self.is_outbound + and 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]}" + ) + + remote_node_type = NodeType(inbound_handshake.node_type) + + if len(inbound_handshake.software_version.encode("utf-8")) > MAX_VERSION_STRING_BYTES: + self.log.debug("version string too long") + raise ProtocolError(Err.INVALID_HANDSHAKE) + + if not self.is_outbound: if ( remote_node_type in {NodeType.FARMER, NodeType.HARVESTER} and inbound_handshake.protocol_version != protocol_version[remote_node_type] @@ -298,12 +284,13 @@ class WSChiaConnection: ), ) 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.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())