diff --git a/common/src/main/java/org/tron/core/config/Parameter.java b/common/src/main/java/org/tron/core/config/Parameter.java index 0f9402641e9..951fa6e70ab 100644 --- a/common/src/main/java/org/tron/core/config/Parameter.java +++ b/common/src/main/java/org/tron/core/config/Parameter.java @@ -98,6 +98,7 @@ public class NodeConstant { public class NetConstants { public static final long SYNC_FETCH_BATCH_NUM = 2000; + public static final long HELLO_TIME_OUT = 10000L; public static final long ADV_TIME_OUT = 20000L; public static final long SYNC_TIME_OUT = 5000L; public static final long NET_MAX_TRX_PER_SECOND = 700L; diff --git a/framework/src/main/java/org/tron/core/net/P2pEventHandlerImpl.java b/framework/src/main/java/org/tron/core/net/P2pEventHandlerImpl.java index 9dd950ae57b..075332b8f46 100644 --- a/framework/src/main/java/org/tron/core/net/P2pEventHandlerImpl.java +++ b/framework/src/main/java/org/tron/core/net/P2pEventHandlerImpl.java @@ -127,6 +127,18 @@ public void onMessage(Channel c, byte[] data) { return; } + if (data == null || data.length == 0) { + peerConnection.disconnect(Protocol.ReasonCode.BAD_PROTOCOL); + return; + } + + if (peerConnection.getHelloMessageReceive() == null + && data[0] != MessageTypes.P2P_HELLO.asByte() + && data[0] != MessageTypes.P2P_DISCONNECT.asByte()) { + peerConnection.disconnect(Protocol.ReasonCode.BAD_PROTOCOL); + return; + } + if (MessageTypes.PBFT_MSG.asByte() == data[0]) { PbftMessage message = null; try { diff --git a/framework/src/main/java/org/tron/core/net/message/handshake/HelloMessage.java b/framework/src/main/java/org/tron/core/net/message/handshake/HelloMessage.java index 68123c93db6..c629ec91877 100755 --- a/framework/src/main/java/org/tron/core/net/message/handshake/HelloMessage.java +++ b/framework/src/main/java/org/tron/core/net/message/handshake/HelloMessage.java @@ -4,7 +4,6 @@ import lombok.Getter; import org.apache.commons.lang3.StringUtils; import org.tron.common.utils.ByteArray; -import org.tron.common.utils.DecodeUtil; import org.tron.common.utils.Sha256Hash; import org.tron.common.utils.StringUtil; import org.tron.core.ChainBaseManager; @@ -13,6 +12,7 @@ import org.tron.core.net.message.MessageTypes; import org.tron.core.net.message.TronMessage; import org.tron.p2p.discover.Node; +import org.tron.p2p.utils.NetUtil; import org.tron.program.Version; import org.tron.protos.Discover.Endpoint; import org.tron.protos.Protocol; @@ -20,6 +20,8 @@ public class HelloMessage extends TronMessage { + private static final int MAX_BYTE_SIZE = 200; + @Getter private Protocol.HelloMessage helloMessage; @@ -121,6 +123,10 @@ public Class getAnswerMessage() { @Override public String toString() { + if (!valid()) { + return "P2P_HELLO: invalid hello message"; + } + StringBuilder builder = new StringBuilder(); builder.append(super.toString()) @@ -156,6 +162,10 @@ public Protocol.HelloMessage getInstance() { } public boolean valid() { + if (!validEndPoint()) { + return false; + } + byte[] genesisBlockByte = this.helloMessage.getGenesisBlockId().getHash().toByteArray(); if (genesisBlockByte.length != Sha256Hash.LENGTH) { return false; @@ -171,25 +181,45 @@ public boolean valid() { return false; } - int maxByteSize = 200; ByteString address = this.helloMessage.getAddress(); - if (!address.isEmpty() && address.toByteArray().length > maxByteSize) { + if (!address.isEmpty() && address.toByteArray().length > MAX_BYTE_SIZE) { return false; } ByteString sig = this.helloMessage.getSignature(); - if (!sig.isEmpty() && sig.toByteArray().length > maxByteSize) { + if (!sig.isEmpty() && sig.toByteArray().length > MAX_BYTE_SIZE) { return false; } ByteString codeVersion = this.helloMessage.getCodeVersion(); - if (!codeVersion.isEmpty() && codeVersion.toByteArray().length > maxByteSize) { + if (!codeVersion.isEmpty() && codeVersion.toByteArray().length > MAX_BYTE_SIZE) { return false; } return true; } + public boolean validEndPoint() { + Endpoint from = this.helloMessage.getFrom(); + ByteString ipv4 = from.getAddress(); + ByteString ipv6 = from.getAddressIpv6(); + if (from.getPort() <= 0 || from.getPort() > 0xFFFF + || ipv4.size() > MAX_BYTE_SIZE || ipv6.size() > MAX_BYTE_SIZE) { + return false; + } + if (ipv4.isEmpty() && ipv6.isEmpty()) { + return false; + } + // Validate raw literals before getFrom() constructs a Node during logging. + if (!ipv4.isEmpty() && !NetUtil.validIpV4(ByteArray.toStr(ipv4.toByteArray()))) { + return false; + } + if (!ipv6.isEmpty() && !NetUtil.validIpV6(ByteArray.toStr(ipv6.toByteArray()))) { + return false; + } + return true; + } + public static Endpoint getEndpointFromNode(Node node) { Endpoint.Builder builder = Endpoint.newBuilder() .setPort(node.getPort()); diff --git a/framework/src/main/java/org/tron/core/net/peer/PeerConnection.java b/framework/src/main/java/org/tron/core/net/peer/PeerConnection.java index 7d7457cf2fc..694a2e702ef 100644 --- a/framework/src/main/java/org/tron/core/net/peer/PeerConnection.java +++ b/framework/src/main/java/org/tron/core/net/peer/PeerConnection.java @@ -107,7 +107,7 @@ public class PeerConnection { @Setter @Getter - private HelloMessage helloMessageReceive; + private volatile HelloMessage helloMessageReceive; @Setter @Getter diff --git a/framework/src/main/java/org/tron/core/net/peer/PeerStatusCheck.java b/framework/src/main/java/org/tron/core/net/peer/PeerStatusCheck.java index 04eac202484..3a770a5d2e6 100644 --- a/framework/src/main/java/org/tron/core/net/peer/PeerStatusCheck.java +++ b/framework/src/main/java/org/tron/core/net/peer/PeerStatusCheck.java @@ -56,6 +56,14 @@ public void statusCheck() { isDisconnected = true; } + if (!isDisconnected) { + isDisconnected = peer.getHelloMessageReceive() == null + && peer.getChannel().getStartTime() <= now - NetConstants.HELLO_TIME_OUT; + if (isDisconnected) { + logger.warn("Peer {} hello message timeout", peer.getInetAddress()); + } + } + if (!isDisconnected) { isDisconnected = peer.getAdvInvRequest().values().stream() .anyMatch(time -> time < now - NetConstants.ADV_TIME_OUT); diff --git a/framework/src/main/java/org/tron/core/net/service/handshake/HandshakeService.java b/framework/src/main/java/org/tron/core/net/service/handshake/HandshakeService.java index 070a9f56406..f8cec9b9d17 100644 --- a/framework/src/main/java/org/tron/core/net/service/handshake/HandshakeService.java +++ b/framework/src/main/java/org/tron/core/net/service/handshake/HandshakeService.java @@ -4,9 +4,7 @@ import lombok.extern.slf4j.Slf4j; import org.springframework.beans.factory.annotation.Autowired; import org.springframework.stereotype.Component; -import org.tron.common.utils.ByteArray; import org.tron.core.ChainBaseManager; -import org.tron.core.ChainBaseManager.NodeType; import org.tron.core.config.args.Args; import org.tron.core.net.TronNetService; import org.tron.core.net.message.handshake.HelloMessage; @@ -49,15 +47,17 @@ public void processHelloMessage(PeerConnection peer, HelloMessage msg) { } if (!msg.valid()) { - logger.warn("Peer {} invalid hello message parameters, GenesisBlockId: {}, SolidBlockId: {}, " - + "HeadBlockId: {}, address: {}, sig: {}, codeVersion: {}", + logger.warn("Peer {} invalid hello message parameters, genesisHashLength: {}, " + + "solidHashLength: {}, headHashLength: {}, address: {}, sig: {}, codeVersion: {}, " + + "endpointValid: {}", peer.getInetSocketAddress(), - ByteArray.toHexString(msg.getInstance().getGenesisBlockId().getHash().toByteArray()), - ByteArray.toHexString(msg.getInstance().getSolidBlockId().getHash().toByteArray()), - ByteArray.toHexString(msg.getInstance().getHeadBlockId().getHash().toByteArray()), - msg.getInstance().getAddress().toByteArray().length, - msg.getInstance().getSignature().toByteArray().length, - msg.getInstance().getCodeVersion().toByteArray().length); + msg.getInstance().getGenesisBlockId().getHash().size(), + msg.getInstance().getSolidBlockId().getHash().size(), + msg.getInstance().getHeadBlockId().getHash().size(), + msg.getInstance().getAddress().size(), + msg.getInstance().getSignature().size(), + msg.getInstance().getCodeVersion().size(), + msg.validEndPoint()); peer.disconnect(ReasonCode.INCOMPATIBLE_PROTOCOL); return; } diff --git a/framework/src/test/java/org/tron/core/net/P2pHelloGateTest.java b/framework/src/test/java/org/tron/core/net/P2pHelloGateTest.java new file mode 100644 index 00000000000..64a1a3aaf1f --- /dev/null +++ b/framework/src/test/java/org/tron/core/net/P2pHelloGateTest.java @@ -0,0 +1,202 @@ +package org.tron.core.net; + +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.times; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.verifyNoInteractions; +import static org.mockito.Mockito.when; + +import java.util.Collections; +import org.junit.After; +import org.junit.Before; +import org.junit.Test; +import org.mockito.InjectMocks; +import org.mockito.Mock; +import org.mockito.MockedStatic; +import org.mockito.Mockito; +import org.mockito.MockitoAnnotations; +import org.tron.common.utils.Sha256Hash; +import org.tron.consensus.pbft.message.PbftMessage; +import org.tron.core.ChainBaseManager; +import org.tron.core.ChainBaseManager.NodeType; +import org.tron.core.capsule.BlockCapsule.BlockId; +import org.tron.core.net.message.MessageTypes; +import org.tron.core.net.message.PbftMessageFactory; +import org.tron.core.net.message.TronMessageFactory; +import org.tron.core.net.message.adv.InventoryMessage; +import org.tron.core.net.message.base.DisconnectMessage; +import org.tron.core.net.message.handshake.HelloMessage; +import org.tron.core.net.message.keepalive.PingMessage; +import org.tron.core.net.messagehandler.BlockMsgHandler; +import org.tron.core.net.messagehandler.ChainInventoryMsgHandler; +import org.tron.core.net.messagehandler.FetchInvDataMsgHandler; +import org.tron.core.net.messagehandler.InventoryMsgHandler; +import org.tron.core.net.messagehandler.PbftDataSyncHandler; +import org.tron.core.net.messagehandler.PbftMsgHandler; +import org.tron.core.net.messagehandler.SyncBlockChainMsgHandler; +import org.tron.core.net.messagehandler.TransactionsMsgHandler; +import org.tron.core.net.peer.PeerConnection; +import org.tron.core.net.peer.PeerManager; +import org.tron.core.net.service.handshake.HandshakeService; +import org.tron.core.net.service.keepalive.KeepAliveService; +import org.tron.core.net.service.statistics.NodeStatistics; +import org.tron.core.net.service.statistics.PeerStatistics; +import org.tron.p2p.connection.Channel; +import org.tron.p2p.discover.Node; +import org.tron.protos.Protocol.Inventory.InventoryType; +import org.tron.protos.Protocol.ReasonCode; + +public class P2pHelloGateTest { + + @InjectMocks + private P2pEventHandlerImpl handler; + @Mock + private HandshakeService handshakeService; + @Mock + private KeepAliveService keepAliveService; + @Mock + private SyncBlockChainMsgHandler syncBlockChainMsgHandler; + @Mock + private ChainInventoryMsgHandler chainInventoryMsgHandler; + @Mock + private InventoryMsgHandler inventoryMsgHandler; + @Mock + private FetchInvDataMsgHandler fetchInvDataMsgHandler; + @Mock + private BlockMsgHandler blockMsgHandler; + @Mock + private TransactionsMsgHandler transactionsMsgHandler; + @Mock + private PbftDataSyncHandler pbftDataSyncHandler; + @Mock + private PbftMsgHandler pbftMsgHandler; + @Mock + private PeerConnection peer; + @Mock + private Channel channel; + + private AutoCloseable mocks; + private MockedStatic peerManager; + + @Before + public void setUp() { + mocks = MockitoAnnotations.openMocks(this); + peerManager = Mockito.mockStatic(PeerManager.class); + peerManager.when(() -> PeerManager.getPeerConnection(channel)).thenReturn(peer); + } + + @After + public void tearDown() throws Exception { + peerManager.close(); + mocks.close(); + } + + @Test + public void testRejectAllOtherMessageTypesBeforeParsing() { + try (MockedStatic tronFactory = + Mockito.mockStatic(TronMessageFactory.class); + MockedStatic pbftFactory = + Mockito.mockStatic(PbftMessageFactory.class)) { + int rejected = 0; + for (MessageTypes type : MessageTypes.values()) { + if (type != MessageTypes.P2P_HELLO && type != MessageTypes.P2P_DISCONNECT) { + handler.onMessage(channel, new byte[] {type.asByte()}); + rejected++; + } + } + verify(peer, times(rejected)).disconnect(ReasonCode.BAD_PROTOCOL); + tronFactory.verifyNoInteractions(); + pbftFactory.verifyNoInteractions(); + } + verify(peer, never()).getPeerStatistics(); + verifyNoInteractions(handshakeService, keepAliveService, syncBlockChainMsgHandler, + chainInventoryMsgHandler, inventoryMsgHandler, fetchInvDataMsgHandler, blockMsgHandler, + transactionsMsgHandler, pbftDataSyncHandler, pbftMsgHandler); + } + + @Test + public void testEmptyFramesAreRejected() { + handler.onMessage(channel, new byte[0]); + handler.onMessage(channel, null); + verify(peer, times(2)).disconnect(ReasonCode.BAD_PROTOCOL); + } + + @Test + public void testUnknownPeerIsIgnored() { + peerManager.when(() -> PeerManager.getPeerConnection(channel)).thenReturn(null); + handler.onMessage(channel, new PingMessage().getSendBytes()); + verifyNoInteractions(peer, keepAliveService); + } + + @Test + public void testHelloIsDispatchedBeforeHandshake() throws Exception { + ChainBaseManager chain = mock(ChainBaseManager.class); + when(chain.getGenesisBlockId()).thenReturn(new BlockId()); + when(chain.getSolidBlockId()).thenReturn(new BlockId()); + when(chain.getHeadBlockId()).thenReturn(new BlockId()); + when(chain.getNodeType()).thenReturn(NodeType.FULL); + HelloMessage hello = new HelloMessage( + new Node(new byte[64], "127.0.0.1", null, 18888), 0, chain); + when(peer.getPeerStatistics()).thenReturn(new PeerStatistics()); + + handler.onMessage(channel, hello.getSendBytes()); + + verify(handshakeService).processHelloMessage(eq(peer), any(HelloMessage.class)); + verify(peer, never()).disconnect(any()); + } + + @Test + public void testDisconnectIsDispatchedBeforeHandshake() { + NodeStatistics stats = mock(NodeStatistics.class); + when(peer.getPeerStatistics()).thenReturn(new PeerStatistics()); + when(peer.getP2pRateLimiter()).thenReturn(new P2pRateLimiter()); + when(peer.getChannel()).thenReturn(channel); + when(peer.getNodeStatistics()).thenReturn(stats); + + handler.onMessage(channel, new DisconnectMessage(ReasonCode.PEER_QUITING).getSendBytes()); + + verify(channel).close(); + verify(stats).nodeDisconnectedRemote(ReasonCode.PEER_QUITING); + verify(peer, never()).disconnect(any()); + } + + @Test + public void testKeepAliveIsDispatchedAfterHello() throws Exception { + when(peer.getHelloMessageReceive()).thenReturn(mock(HelloMessage.class)); + when(peer.getPeerStatistics()).thenReturn(new PeerStatistics()); + + handler.onMessage(channel, new PingMessage().getSendBytes()); + + verify(keepAliveService).processMessage(eq(peer), any(PingMessage.class)); + verify(peer, never()).disconnect(any()); + } + + @Test + public void testInventoryIsDispatchedAfterHello() throws Exception { + when(peer.getHelloMessageReceive()).thenReturn(mock(HelloMessage.class)); + when(peer.getPeerStatistics()).thenReturn(new PeerStatistics()); + InventoryMessage inventory = new InventoryMessage( + Collections.singletonList(Sha256Hash.ZERO_HASH), InventoryType.BLOCK); + + handler.onMessage(channel, inventory.getSendBytes()); + + verify(inventoryMsgHandler).processMessage(eq(peer), any(InventoryMessage.class)); + verify(peer, never()).disconnect(any()); + } + + @Test + public void testPbftIsDispatchedAfterHello() throws Exception { + when(peer.getHelloMessageReceive()).thenReturn(mock(HelloMessage.class)); + byte[] data = new byte[] {MessageTypes.PBFT_MSG.asByte()}; + PbftMessage message = mock(PbftMessage.class); + try (MockedStatic factory = Mockito.mockStatic(PbftMessageFactory.class)) { + factory.when(() -> PbftMessageFactory.create(data)).thenReturn(message); + handler.onMessage(channel, data); + verify(pbftMsgHandler).processMessage(peer, message); + verify(peer, never()).disconnect(any()); + } + } +} diff --git a/framework/src/test/java/org/tron/core/net/message/handshake/HelloMessageTest.java b/framework/src/test/java/org/tron/core/net/message/handshake/HelloMessageTest.java new file mode 100644 index 00000000000..feabc2160c9 --- /dev/null +++ b/framework/src/test/java/org/tron/core/net/message/handshake/HelloMessageTest.java @@ -0,0 +1,180 @@ +package org.tron.core.net.message.handshake; + +import static org.mockito.Mockito.doThrow; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.spy; +import static org.mockito.Mockito.verify; + +import ch.qos.logback.classic.AsyncAppender; +import ch.qos.logback.classic.Level; +import ch.qos.logback.classic.LoggerContext; +import ch.qos.logback.classic.spi.ILoggingEvent; +import ch.qos.logback.classic.util.LogbackMDCAdapter; +import ch.qos.logback.core.read.ListAppender; +import com.google.protobuf.ByteString; +import java.net.InetSocketAddress; +import org.apache.commons.lang3.StringUtils; +import org.junit.After; +import org.junit.Assert; +import org.junit.Before; +import org.junit.Test; +import org.mockito.MockedStatic; +import org.mockito.Mockito; +import org.tron.p2p.P2pConfig; +import org.tron.p2p.base.Parameter; +import org.tron.p2p.utils.NetUtil; +import org.tron.protos.Discover.Endpoint; +import org.tron.protos.Protocol; + +public class HelloMessageTest { + + private P2pConfig savedConfig; + + @Before + public void setUp() { + savedConfig = Parameter.p2pConfig; + Parameter.p2pConfig = new P2pConfig(); + Parameter.p2pConfig.setIp("127.0.0.1"); + Parameter.p2pConfig.setIpv6("::1"); + } + + @After + public void tearDown() { + Parameter.p2pConfig = savedConfig; + } + + @Test(timeout = 5000) + public void testAcceptIpLiteralsAndPortBoundaries() throws Exception { + for (String[] hosts : new String[][] { + {"192.0.2.1", ""}, {"", "2001:db8::1"}, {"192.0.2.1", "2001:db8::1"}, + {"", "::1"}, {"", "::ffff:192.0.2.1"}}) { + for (int port : new int[] {1, 18888, 65535}) { + HelloMessage message = spy(hello(hosts[0], hosts[1], port)); + Assert.assertTrue(message.validEndPoint()); + Assert.assertTrue(message.valid()); + verify(message, never()).getFrom(); + } + } + } + + @Test(timeout = 5000) + public void testRejectHostnamesAndMalformedAddresses() throws Exception { + for (String[] hosts : new String[][] { + {"hello.invalid", ""}, {"", "hello.invalid"}, + {"hello.invalid", "2001:db8::1"}, {"192.0.2.1", "hello.invalid"}, + {"256.0.0.1", ""}, {"192.0.2.1\n", ""}, {" 192.0.2.1", ""}, + {"2001:db8::1", ""}, {"", "192.0.2.1"}, {"", "2001:db8:::1"}}) { + assertRejectedBeforeNodeConstruction(hello(hosts[0], hosts[1], 18888)); + } + } + + @Test(timeout = 5000) + public void testRejectMissingAddresses() throws Exception { + assertRejectedBeforeNodeConstruction(hello("", "", 18888)); + HelloMessage missingFrom = hello("192.0.2.1", "", 18888); + missingFrom.setHelloMessage(missingFrom.getInstance().toBuilder().clearFrom().build()); + assertRejectedBeforeNodeConstruction(missingFrom); + } + + @Test(timeout = 5000) + public void testRejectInvalidPorts() throws Exception { + for (int port : new int[] {Integer.MIN_VALUE, -1, 0, 65536, Integer.MAX_VALUE}) { + assertRejectedBeforeNodeConstruction(hello("192.0.2.1", "2001:db8::1", port)); + } + } + + @Test(timeout = 5000) + public void testRejectOversizedAddressesBeforeLiteralValidation() throws Exception { + String oversized = StringUtils.repeat('a', 201); + try (MockedStatic netUtil = Mockito.mockStatic(NetUtil.class)) { + assertRejectedBeforeNodeConstruction(hello(oversized, "", 18888)); + assertRejectedBeforeNodeConstruction(hello("192.0.2.1", oversized, 18888)); + netUtil.verifyNoInteractions(); + } + } + + @Test(timeout = 5000) + public void testIpv6AddressLengthBoundary() throws Exception { + String scoped = "fe80::1%" + StringUtils.repeat('a', 192); + Assert.assertEquals(200, scoped.length()); + Assert.assertTrue(NetUtil.validIpV6(scoped)); + HelloMessage message = spy(hello("", scoped, 18888)); + Assert.assertTrue(message.valid()); + verify(message, never()).getFrom(); + // The syntax remains valid, so rejection must come from the byte limit. + Assert.assertTrue(NetUtil.validIpV6(scoped + "a")); + assertRejectedBeforeNodeConstruction(hello("", scoped + "a", 18888)); + } + + @Test(timeout = 5000) + public void testValidAddressLogFormat() throws Exception { + for (String[] hosts : new String[][] { + {"192.0.2.1", "", "192.0.2.1"}, + {"", "2001:db8::1", "2001:db8::1"}, + {"192.0.2.1", "2001:db8::1", "192.0.2.1"}}) { + HelloMessage message = hello(hosts[0], hosts[1], 18888); + Assert.assertTrue(message.valid()); + // InetSocketAddress adds IPv6 brackets on JDK 17, but not on JDK 8. + InetSocketAddress expectedEndpoint = new InetSocketAddress(hosts[2], 18888); + Assert.assertFalse(expectedEndpoint.isUnresolved()); + String formatted = message.toString(); + Assert.assertTrue(formatted, formatted.contains("from: " + expectedEndpoint + "\n")); + Assert.assertTrue(formatted.contains("timestamp: 123\n")); + Assert.assertTrue(formatted.contains("headBlockId: " + + message.getHeadBlockId().getString() + "\n")); + Assert.assertTrue(formatted.contains("nodeType: 0\nlowestBlockNum: 0\n")); + } + } + + @Test(timeout = 5000) + public void testAsyncLoggingRejectsHostnameBeforeNodeConstruction() throws Exception { + HelloMessage message = spy(hello("hello.invalid", "", 18888)); + doThrow(new AssertionError("Invalid endpoint reached Node construction")) + .when(message).getFrom(); + LoggerContext context = new LoggerContext(); + context.setMDCAdapter(new LogbackMDCAdapter()); + ListAppender sink = new ListAppender<>(); + sink.setContext(context); + sink.start(); + AsyncAppender async = new AsyncAppender(); + async.setContext(context); + async.addAppender(sink); + async.setDiscardingThreshold(0); + async.start(); + ch.qos.logback.classic.Logger logger = context.getLogger("hello.validation.test"); + logger.setLevel(Level.INFO); + logger.setAdditive(false); + logger.addAppender(async); + try { + logger.info("Receive HELLO {}", message); + } finally { + async.stop(); + context.stop(); + } + Assert.assertEquals(1, sink.list.size()); + Assert.assertEquals("Receive HELLO P2P_HELLO: invalid hello message", + sink.list.get(0).getFormattedMessage()); + verify(message, never()).getFrom(); + } + + private static void assertRejectedBeforeNodeConstruction(HelloMessage message) { + HelloMessage checked = spy(message); + doThrow(new AssertionError("Invalid endpoint reached Node construction")) + .when(checked).getFrom(); + Assert.assertFalse(checked.validEndPoint()); + Assert.assertFalse(checked.valid()); + Assert.assertEquals("P2P_HELLO: invalid hello message", checked.toString()); + verify(checked, never()).getFrom(); + } + + private static HelloMessage hello(String ipv4, String ipv6, int port) throws Exception { + Protocol.HelloMessage.BlockId block = Protocol.HelloMessage.BlockId.newBuilder() + .setHash(ByteString.copyFrom(new byte[32])).build(); + Endpoint endpoint = Endpoint.newBuilder().setAddress(ByteString.copyFromUtf8(ipv4)) + .setAddressIpv6(ByteString.copyFromUtf8(ipv6)) + .setNodeId(ByteString.copyFrom(new byte[64])).setPort(port).build(); + return new HelloMessage(Protocol.HelloMessage.newBuilder().setFrom(endpoint) + .setGenesisBlockId(block).setSolidBlockId(block).setHeadBlockId(block) + .setTimestamp(123).build().toByteArray()); + } +} diff --git a/framework/src/test/java/org/tron/core/net/messagehandler/MessageHandlerTest.java b/framework/src/test/java/org/tron/core/net/messagehandler/MessageHandlerTest.java index c8205b6b721..615fb573962 100644 --- a/framework/src/test/java/org/tron/core/net/messagehandler/MessageHandlerTest.java +++ b/framework/src/test/java/org/tron/core/net/messagehandler/MessageHandlerTest.java @@ -24,6 +24,7 @@ import org.tron.core.config.args.Args; import org.tron.core.net.P2pEventHandlerImpl; import org.tron.core.net.TronNetService; +import org.tron.core.net.message.handshake.HelloMessage; import org.tron.core.net.message.keepalive.PingMessage; import org.tron.core.net.peer.PeerConnection; import org.tron.core.net.peer.PeerManager; @@ -82,6 +83,7 @@ public void testPbft() { Assert.assertFalse(c1.isDisconnect()); peer = PeerManager.getPeers().get(0); + peer.setHelloMessageReceive(mock(HelloMessage.class)); BlockCapsule blockCapsule = new BlockCapsule(1, Sha256Hash.ZERO_HASH, System.currentTimeMillis(), ByteString.EMPTY); PbftMessage pbftMessage = PbftMessage.fullNodePrePrepareBlockMsg(blockCapsule, 0L); @@ -102,7 +104,7 @@ public void testPing() { Channel c1 = mock(Channel.class); Mockito.when(c1.getInetSocketAddress()).thenReturn(a1); Mockito.when(c1.getInetAddress()).thenReturn(a1.getAddress()); - PeerManager.add(ctx, c1); + PeerManager.add(ctx, c1).setHelloMessageReceive(mock(HelloMessage.class)); PingMessage pingMessage = new PingMessage(); p2pEventHandler.onMessage(c1, pingMessage.getSendBytes()); diff --git a/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckMockTest.java b/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckMockTest.java index d2ee4be5b87..ac8462df604 100644 --- a/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckMockTest.java +++ b/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckMockTest.java @@ -2,26 +2,60 @@ import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.eq; +import static org.mockito.Mockito.clearInvocations; +import static org.mockito.Mockito.doNothing; import static org.mockito.Mockito.doThrow; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; import static org.mockito.Mockito.spy; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; +import java.util.Collections; import java.util.concurrent.ScheduledExecutorService; import java.util.concurrent.TimeUnit; import org.junit.After; +import org.junit.Before; import org.junit.Test; import org.mockito.Mockito; import org.tron.common.utils.ReflectUtils; +import org.tron.core.capsule.BlockCapsule.BlockId; +import org.tron.core.config.Parameter.NetConstants; +import org.tron.core.net.TronNetDelegate; +import org.tron.core.net.message.handshake.HelloMessage; +import org.tron.p2p.connection.Channel; +import org.tron.protos.Protocol.ReasonCode; public class PeerStatusCheckMockTest { + + private PeerStatusCheck peerStatusCheck; + private PeerConnection peer; + private Channel channel; + + @Before + public void setUp() { + peerStatusCheck = spy(new PeerStatusCheck()); + peer = spy(new PeerConnection()); + channel = mock(Channel.class); + ReflectUtils.setFieldValue(peer, "channel", channel); + doNothing().when(peer).disconnect(any()); + TronNetDelegate delegate = mock(TronNetDelegate.class); + when(delegate.getActivePeer()).thenReturn(Collections.singletonList(peer)); + ReflectUtils.setFieldValue(peerStatusCheck, "tronNetDelegate", delegate); + } + @After - public void clearMocks() { - Mockito.framework().clearInlineMocks(); + public void tearDown() { + try { + peerStatusCheck.close(); + } finally { + Mockito.framework().clearInlineMocks(); + } } @Test public void testInitException() { - PeerStatusCheck peerStatusCheck = spy(new PeerStatusCheck()); + peerStatusCheck.close(); ScheduledExecutorService executor = mock(ScheduledExecutorService.class); ReflectUtils.setFieldValue(peerStatusCheck, "peerStatusCheckExecutor", executor); doThrow(new RuntimeException("test exception")).when(peerStatusCheck).statusCheck(); @@ -41,4 +75,52 @@ public void testInitException() { Mockito.verify(peerStatusCheck).statusCheck(); } + @Test + public void testHelloTimeoutWithoutSync() { + peer.setNeedSyncFromPeer(false); + peer.setNeedSyncFromUs(false); + when(channel.getStartTime()).thenReturn(System.currentTimeMillis()); + assertTimeout(false); + + when(channel.getStartTime()).thenReturn(System.currentTimeMillis() - 10_000); + assertTimeout(true); + } + + @Test + public void testHelloTimeoutDependsOnHandshakeCompletion() { + when(channel.getStartTime()).thenReturn(System.currentTimeMillis() - 60_000); + peer.setLastInteractiveTime(System.currentTimeMillis()); + peer.setBlockBothHave(new BlockId()); + // Recent traffic and block progress must not extend the HELLO deadline. + assertTimeout(true); + + peer.setHelloMessageReceive(mock(HelloMessage.class)); + peer.setNeedSyncFromPeer(false); + assertTimeout(false); + } + + @Test + public void testHelloCheckPreservesExistingTimeouts() { + when(channel.getStartTime()).thenReturn(System.currentTimeMillis()); + ReflectUtils.setFieldValue(peer, "blockBothHaveUpdateTime", 0L); + // A fresh connection must still be subject to the existing sync timeout. + assertTimeout(true); + + peer.setHelloMessageReceive(mock(HelloMessage.class)); + peer.setNeedSyncFromPeer(false); + peer.getSyncBlockRequested().put(new BlockId(), + System.currentTimeMillis() - NetConstants.SYNC_TIME_OUT - 1_000); + assertTimeout(true); + } + + private void assertTimeout(boolean expected) { + clearInvocations(peer); + peerStatusCheck.statusCheck(); + if (expected) { + verify(peer).disconnect(ReasonCode.TIME_OUT); + } else { + verify(peer, never()).disconnect(any()); + } + } + } diff --git a/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckTest.java b/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckTest.java index 2d734f45215..7a4fb2fd403 100644 --- a/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckTest.java +++ b/framework/src/test/java/org/tron/core/net/peer/PeerStatusCheckTest.java @@ -13,6 +13,7 @@ import org.tron.core.capsule.BlockCapsule.BlockId; import org.tron.core.config.Parameter.NetConstants; import org.tron.core.config.args.Args; +import org.tron.core.net.message.handshake.HelloMessage; import org.tron.p2p.connection.Channel; @@ -44,7 +45,7 @@ public void testCheck() { ReflectUtils.setFieldValue(c1, "ctx", spy(ChannelHandlerContext.class)); Mockito.doNothing().when(c1).send((byte[]) any()); - PeerManager.add(context, c1); + PeerManager.add(context, c1).setHelloMessageReceive(Mockito.mock(HelloMessage.class)); } PeerManager.getPeers().get(0).getSyncBlockRequested() diff --git a/framework/src/test/java/org/tron/core/net/service/handshake/HandshakeHeadCheckTest.java b/framework/src/test/java/org/tron/core/net/service/handshake/HandshakeHeadCheckTest.java new file mode 100644 index 00000000000..9e7734aecb2 --- /dev/null +++ b/framework/src/test/java/org/tron/core/net/service/handshake/HandshakeHeadCheckTest.java @@ -0,0 +1,281 @@ +package org.tron.core.net.service.handshake; + +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; + +import java.net.InetSocketAddress; +import org.junit.After; +import org.junit.Assert; +import org.junit.Before; +import org.junit.Test; +import org.mockito.MockedStatic; +import org.mockito.Mockito; +import org.tron.common.utils.ReflectUtils; +import org.tron.common.utils.Sha256Hash; +import org.tron.core.ChainBaseManager; +import org.tron.core.ChainBaseManager.NodeType; +import org.tron.core.capsule.BlockCapsule.BlockId; +import org.tron.core.net.TronNetService; +import org.tron.core.net.message.handshake.HelloMessage; +import org.tron.core.net.peer.PeerConnection; +import org.tron.core.net.peer.PeerManager; +import org.tron.core.net.service.effective.EffectiveCheckService; +import org.tron.core.net.service.relay.RelayService; +import org.tron.p2p.P2pService; +import org.tron.p2p.connection.Channel; +import org.tron.p2p.discover.Node; +import org.tron.protos.Protocol; +import org.tron.protos.Protocol.ReasonCode; + +public class HandshakeHeadCheckTest { + + private HandshakeService service; + private ChainBaseManager chain; + private PeerConnection peer; + private Channel channel; + private EffectiveCheckService effectiveCheckService; + private MockedStatic netService; + private MockedStatic peerManager; + private BlockId genesis; + private BlockId solid; + private BlockId localHead; + + @Before + public void setUp() { + service = new HandshakeService(); + chain = mock(ChainBaseManager.class); + peer = mock(PeerConnection.class); + channel = mock(Channel.class); + effectiveCheckService = mock(EffectiveCheckService.class); + RelayService relayService = mock(RelayService.class); + ReflectUtils.setFieldValue(service, "chainBaseManager", chain); + ReflectUtils.setFieldValue(service, "relayService", relayService); + ReflectUtils.setFieldValue(service, "effectiveCheckService", effectiveCheckService); + when(peer.getChannel()).thenReturn(channel); + when(peer.getInetSocketAddress()).thenReturn(new InetSocketAddress("127.0.0.1", 18888)); + when(relayService.checkHelloMessage(any(), any())).thenReturn(true); + + genesis = blockId(0, 1); + solid = blockId(80, 3); + localHead = blockId(100, 2); + when(chain.getGenesisBlockId()).thenReturn(genesis); + when(chain.getSolidBlockId()).thenReturn(solid); + when(chain.getHeadBlockId()).thenReturn(localHead); + when(chain.getHeadBlockNum()).thenReturn(100L); + when(chain.getLowestBlockNum()).thenReturn(10L); + when(chain.getNodeType()).thenReturn(NodeType.FULL); + when(chain.containBlockInMainChain(genesis)).thenReturn(true); + when(chain.containBlockInMainChain(solid)).thenReturn(true); + when(chain.containBlockInMainChain(localHead)).thenReturn(true); + + P2pService p2p = mock(P2pService.class); + netService = Mockito.mockStatic(TronNetService.class); + netService.when(TronNetService::getP2pService).thenReturn(p2p); + peerManager = Mockito.mockStatic(PeerManager.class); + } + + @After + public void tearDown() { + peerManager.close(); + netService.close(); + } + + @Test + public void testAcceptUnknownHeadAtLocalHeight() { + BlockId unknown = blockId(100, 4); + HelloMessage hello = hello(unknown); + Assert.assertTrue(hello.valid()); + Assert.assertEquals(chain.getSolidBlockId(), hello.getSolidBlockId()); + Assert.assertNotEquals(localHead, unknown); + Assert.assertFalse(chain.containBlockInMainChain(unknown)); + + service.processHelloMessage(peer, hello); + + assertAccepted(hello); + } + + @Test + public void testAcceptCachedForkHeadAtLocalHeight() { + BlockId forkHead = blockId(100, 12); + when(chain.containBlock(forkHead)).thenReturn(true); + Assert.assertTrue(chain.containBlock(forkHead)); + Assert.assertFalse(chain.containBlockInMainChain(forkHead)); + HelloMessage hello = hello(forkHead); + + service.processHelloMessage(peer, hello); + + assertAccepted(hello); + } + + @Test + public void testAcceptMainChainHeadAtLocalHeight() { + HelloMessage hello = hello(localHead); + + service.processHelloMessage(peer, hello); + + assertAccepted(hello); + } + + @Test + public void testAcceptOlderMainChainHead() { + BlockId older = blockId(90, 5); + when(chain.containBlockInMainChain(older)).thenReturn(true); + HelloMessage hello = hello(older); + + service.processHelloMessage(peer, hello); + + assertAccepted(hello); + } + + @Test + public void testAcceptForkHeadWhenSolidBlockMatches() { + BlockId forkHead = blockId(90, 6); + HelloMessage hello = hello(forkHead); + Assert.assertEquals(chain.getSolidBlockId(), hello.getSolidBlockId()); + // Keep the connection available so a short-lived fork can converge through sync/broadcast. + service.processHelloMessage(peer, hello); + + verify(chain).containBlockInMainChain(solid); + assertAccepted(hello); + } + + @Test + public void testAcceptOlderForkHeadBelowLocalSolid() { + when(chain.getSolidBlockId()).thenReturn(blockId(95, 13)); + HelloMessage hello = hello(blockId(90, 7)); + Assert.assertTrue(hello.getSolidBlockId().getNum() < hello.getHeadBlockId().getNum()); + Assert.assertTrue(hello.getHeadBlockId().getNum() < chain.getSolidBlockId().getNum()); + + service.processHelloMessage(peer, hello); + + verify(chain).containBlockInMainChain(solid); + assertAccepted(hello); + } + + @Test + public void testAcceptKnownHeadAtLowestRetainedHeight() { + when(chain.getLowestBlockNum()).thenReturn(solid.getNum()); + HelloMessage hello = hello(solid); + + service.processHelloMessage(peer, hello); + + verify(chain).containBlockInMainChain(solid); + assertAccepted(hello); + } + + @Test + public void testHigherHeadContinuesToSync() { + BlockId higher = blockId(101, 8); + HelloMessage hello = hello(higher); + + service.processHelloMessage(peer, hello); + + verify(chain, never()).containBlockInMainChain(higher); + assertAccepted(hello); + } + + @Test + public void testUnverifiableSolidBelowRetainedHistoryIsRejected() { + BlockId prunedSolid = blockId(9, 9); + HelloMessage hello = hello(blockId(90, 14)); + hello.setHelloMessage(hello.getInstance().toBuilder() + .setSolidBlockId(protoBlockId(prunedSolid)).build()); + + service.processHelloMessage(peer, hello); + + verify(chain).containBlockInMainChain(prunedSolid); + assertRejected(ReasonCode.LIGHT_NODE_SYNC_FAIL); + } + + @Test + public void testGenesisMismatchIsStillRejected() { + HelloMessage hello = hello(localHead); + hello.setHelloMessage(hello.getInstance().toBuilder() + .setGenesisBlockId(protoBlockId(blockId(0, 10))).build()); + + service.processHelloMessage(peer, hello); + + assertRejected(ReasonCode.INCOMPATIBLE_CHAIN); + } + + @Test + public void testVersionMismatchIsStillRejected() { + HelloMessage hello = hello(localHead); + hello.setHelloMessage(hello.getInstance().toBuilder() + .setVersion(hello.getVersion() + 1).build()); + + service.processHelloMessage(peer, hello); + + assertRejected(ReasonCode.INCOMPATIBLE_VERSION); + } + + @Test + public void testEffectiveCheckStillRejectsOlderKnownHead() { + BlockId older = blockId(90, 5); + when(chain.containBlockInMainChain(older)).thenReturn(true); + InetSocketAddress address = peer.getInetSocketAddress(); + when(effectiveCheckService.getCur()).thenReturn(address); + + service.processHelloMessage(peer, hello(older)); + + assertRejected(ReasonCode.BELOW_THAN_ME); + } + + @Test + public void testSolidForkIsStillRejected() { + BlockId differentSolid = blockId(80, 11); + HelloMessage hello = hello(localHead); + hello.setHelloMessage(hello.getInstance().toBuilder() + .setSolidBlockId(protoBlockId(differentSolid)).build()); + + service.processHelloMessage(peer, hello); + + verify(chain).containBlockInMainChain(differentSolid); + assertRejected(ReasonCode.FORKED); + } + + @Test + public void testDuplicateHelloIsStillRejected() { + HelloMessage hello = hello(localHead); + when(peer.getHelloMessageReceive()).thenReturn(hello); + + service.processHelloMessage(peer, hello); + + assertRejected(ReasonCode.BAD_PROTOCOL); + } + + private HelloMessage hello(BlockId head) { + HelloMessage hello = new HelloMessage( + new Node(new byte[64], "127.0.0.1", null, 18888), 0, chain); + hello.setHelloMessage(hello.getInstance().toBuilder() + .setSolidBlockId(protoBlockId(solid)) + .setHeadBlockId(protoBlockId(head)).build()); + return hello; + } + + private static BlockId blockId(long number, int seed) { + return new BlockId(Sha256Hash.of(true, new byte[] {(byte) seed}), number); + } + + private static Protocol.HelloMessage.BlockId protoBlockId(BlockId id) { + return Protocol.HelloMessage.BlockId.newBuilder() + .setNumber(id.getNum()).setHash(id.getByteString()).build(); + } + + private void assertAccepted(HelloMessage hello) { + verify(peer).setHelloMessageReceive(hello); + verify(peer).onConnect(); + verify(peer, never()).disconnect(any()); + peerManager.verify(PeerManager::sortPeers); + } + + private void assertRejected(ReasonCode reason) { + verify(peer).disconnect(reason); + verify(peer, never()).setHelloMessageReceive(any()); + verify(peer, never()).onConnect(); + peerManager.verifyNoInteractions(); + } +} diff --git a/framework/src/test/java/org/tron/core/net/services/HandShakeServiceTest.java b/framework/src/test/java/org/tron/core/net/services/HandShakeServiceTest.java index dce5ccb851f..94a89ae5ec3 100644 --- a/framework/src/test/java/org/tron/core/net/services/HandShakeServiceTest.java +++ b/framework/src/test/java/org/tron/core/net/services/HandShakeServiceTest.java @@ -1,13 +1,29 @@ package org.tron.core.net.services; +import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.never; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; import static org.tron.core.net.message.handshake.HelloMessage.getEndpointFromNode; +import ch.qos.logback.classic.Level; +import ch.qos.logback.classic.Logger; +import ch.qos.logback.classic.spi.ILoggingEvent; +import ch.qos.logback.core.Appender; +import ch.qos.logback.core.read.ListAppender; import com.google.protobuf.ByteString; import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; import java.net.InetSocketAddress; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Iterator; +import java.util.List; import java.util.Random; +import java.util.concurrent.CopyOnWriteArrayList; +import java.util.stream.Collectors; +import org.apache.commons.lang3.StringUtils; import org.junit.After; import org.junit.AfterClass; import org.junit.Assert; @@ -15,7 +31,9 @@ import org.junit.ClassRule; import org.junit.Test; import org.junit.rules.TemporaryFolder; +import org.mockito.MockedStatic; import org.mockito.Mockito; +import org.slf4j.LoggerFactory; import org.springframework.context.ApplicationContext; import org.tron.common.TestConstants; import org.tron.common.application.TronApplicationContext; @@ -33,6 +51,7 @@ import org.tron.core.net.peer.PeerManager; import org.tron.core.net.service.handshake.HandshakeService; import org.tron.p2p.P2pConfig; +import org.tron.p2p.P2pService; import org.tron.p2p.base.Parameter; import org.tron.p2p.connection.Channel; import org.tron.p2p.discover.Node; @@ -41,6 +60,7 @@ import org.tron.protos.Discover.Endpoint; import org.tron.protos.Protocol; import org.tron.protos.Protocol.HelloMessage.Builder; +import org.tron.protos.Protocol.ReasonCode; public class HandShakeServiceTest { @@ -316,6 +336,114 @@ public void testProcessHelloMessage() { } } + @Test + public void testInvalidHelloLogsHashLengths() throws Exception { + int largeHashLength = Parameter.MAX_MESSAGE_LENGTH - 1024; + Endpoint validEndpoint = endpoint("127.0.0.1", "", 18888); + for (int[] lengths : new int[][] { + {32, 32, 32}, {0, 32, 32}, {32, 31, 32}, {32, 32, 33}, {0, 31, 33}, + {largeHashLength, 32, 32}, {32, largeHashLength, 32}, {32, 32, largeHashLength}}) { + // The oversized signature also exercises logging when all three hashes are 32 bytes. + assertInvalidHelloLog(lengths[0], lengths[1], lengths[2], validEndpoint, 201, true); + } + } + + @Test + public void testInvalidHelloLogsEndpointValidity() throws Exception { + for (Endpoint invalidEndpoint : new Endpoint[] { + endpoint("127.0.0.1", "", 0), endpoint("127.0.0.1", "", 65536), + Endpoint.getDefaultInstance(), endpoint("", "", 18888), + endpoint("hello.invalid", "", 18888), endpoint("127.0.0.1", "2001:db8:::1", 18888), + endpoint("127.0.0.1\nforged", "", 18888), + endpoint(StringUtils.repeat('a', Parameter.MAX_MESSAGE_LENGTH - 1024), "", 18888), + endpoint("127.0.0.1", StringUtils.repeat('a', 201), 18888)}) { + // Keep all other HELLO fields valid so the endpoint alone causes rejection. + assertInvalidHelloLog(32, 32, 32, invalidEndpoint, 0, false); + } + assertInvalidHelloLog(32, 32, 32, endpoint("", "2001:db8::1", 18888), 201, true); + } + + private void assertInvalidHelloLog(int genesisLength, int solidLength, int headLength, + Endpoint endpoint, int signatureLength, boolean expectedEndpointValid) throws Exception { + Protocol.HelloMessage proto = Protocol.HelloMessage.newBuilder() + .setFrom(endpoint) + .setGenesisBlockId(blockId(genesisLength, (byte) 0x11)) + .setSolidBlockId(blockId(solidLength, (byte) 0x22)) + .setHeadBlockId(blockId(headLength, (byte) 0x33)) + .setSignature(ByteString.copyFrom(new byte[signatureLength])) + .build(); + Assert.assertTrue(proto.getSerializedSize() + 1 < Parameter.MAX_MESSAGE_LENGTH); + HelloMessage hello = new HelloMessage(proto.toByteArray()); + Assert.assertEquals(expectedEndpointValid, hello.validEndPoint()); + Assert.assertFalse(hello.valid()); + + PeerConnection testPeer = mock(PeerConnection.class); + when(testPeer.getChannel()).thenReturn(mock(Channel.class)); + when(testPeer.getInetSocketAddress()).thenReturn(new InetSocketAddress("127.0.0.1", 18888)); + P2pService p2p = mock(P2pService.class); + Logger logger = (Logger) LoggerFactory.getLogger("net"); + Level originalLevel = logger.getLevel(); + boolean originalAdditive = logger.isAdditive(); + List> originalAppenders = new ArrayList<>(); + Iterator> iterator = logger.iteratorForAppenders(); + while (iterator.hasNext()) { + originalAppenders.add(iterator.next()); + } + ListAppender sink = new ListAppender<>(); + // The shared application context can also log from background threads. + sink.list = new CopyOnWriteArrayList<>(); + sink.setContext(logger.getLoggerContext()); + sink.start(); + try (MockedStatic netService = Mockito.mockStatic(TronNetService.class)) { + netService.when(TronNetService::getP2pService).thenReturn(p2p); + for (Appender appender : originalAppenders) { + logger.detachAppender(appender); + } + logger.setAdditive(false); + logger.setLevel(Level.WARN); + logger.addAppender(sink); + + new HandshakeService().processHelloMessage(testPeer, hello); + + verify(testPeer).disconnect(ReasonCode.INCOMPATIBLE_PROTOCOL); + verify(testPeer, never()).setHelloMessageReceive(any()); + verify(testPeer, never()).onConnect(); + List warnings = sink.list.stream() + .filter(event -> event.getMessage() + .startsWith("Peer {} invalid hello message parameters")) + .collect(Collectors.toList()); + Assert.assertEquals(1, warnings.size()); + ILoggingEvent event = warnings.get(0); + Assert.assertEquals(Level.WARN, event.getLevel()); + String text = event.getFormattedMessage(); + Assert.assertTrue("Invalid HELLO warning must remain bounded", text.length() < 512); + Assert.assertEquals("Peer /127.0.0.1:18888 invalid hello message parameters, " + + "genesisHashLength: " + genesisLength + ", solidHashLength: " + solidLength + + ", headHashLength: " + headLength + ", address: 0, sig: " + signatureLength + + ", codeVersion: 0, endpointValid: " + expectedEndpointValid, text); + } finally { + logger.detachAppender(sink); + sink.stop(); + for (Appender appender : originalAppenders) { + logger.addAppender(appender); + } + logger.setLevel(originalLevel); + logger.setAdditive(originalAdditive); + } + } + + private static Endpoint endpoint(String ipv4, String ipv6, int port) { + return Endpoint.newBuilder().setNodeId(ByteString.copyFrom(new byte[64])) + .setAddress(ByteString.copyFromUtf8(ipv4)).setAddressIpv6(ByteString.copyFromUtf8(ipv6)) + .setPort(port).build(); + } + + private static Protocol.HelloMessage.BlockId blockId(int length, byte value) { + byte[] hash = new byte[length]; + Arrays.fill(hash, value); + return Protocol.HelloMessage.BlockId.newBuilder().setHash(ByteString.copyFrom(hash)).build(); + } + private Protocol.HelloMessage.Builder getHelloMessageBuilder(Node from, long timestamp, ChainBaseManager chainBaseManager) { Endpoint fromEndpoint = getEndpointFromNode(from);