diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandler.java b/proxy/src/main/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandler.java index 4ab0a01f70c..a12a5f20d27 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandler.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandler.java @@ -85,6 +85,9 @@ public boolean match(ByteBuf in) { if (!ConfigurationManager.getProxyConfig().isEnableRemotingLocalProxyGrpc()) { return false; } + if (in.readableBytes() < Integer.BYTES) { + return false; + } // If starts with 'PRI ' return in.getInt(in.readerIndex()) == PRI_INT; diff --git a/proxy/src/test/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandlerTest.java b/proxy/src/test/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandlerTest.java index 4a417ea68a2..ee0596e589a 100644 --- a/proxy/src/test/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandlerTest.java +++ b/proxy/src/test/java/org/apache/rocketmq/proxy/remoting/protocol/http2proxy/Http2ProtocolProxyHandlerTest.java @@ -17,22 +17,30 @@ package org.apache.rocketmq.proxy.remoting.protocol.http2proxy; +import io.netty.buffer.ByteBuf; +import io.netty.buffer.Unpooled; import io.netty.channel.Channel; import io.netty.channel.ChannelPipeline; import io.netty.handler.codec.haproxy.HAProxyMessageEncoder; +import org.apache.rocketmq.proxy.config.ConfigurationManager; +import org.apache.rocketmq.proxy.config.InitConfigTest; +import org.junit.After; import org.junit.Before; import org.junit.Test; import org.junit.runner.RunWith; import org.mockito.Mock; import org.mockito.junit.MockitoJUnitRunner; +import static org.junit.Assert.assertFalse; +import static org.junit.Assert.assertTrue; import static org.mockito.ArgumentMatchers.any; import static org.mockito.Mockito.when; @RunWith(MockitoJUnitRunner.class) -public class Http2ProtocolProxyHandlerTest { +public class Http2ProtocolProxyHandlerTest extends InitConfigTest { private Http2ProtocolProxyHandler http2ProtocolProxyHandler; + private boolean originalEnableRemotingLocalProxyGrpc; @Mock private Channel inboundChannel; @Mock @@ -44,9 +52,16 @@ public class Http2ProtocolProxyHandlerTest { @Before public void setUp() throws Exception { + originalEnableRemotingLocalProxyGrpc = ConfigurationManager.getProxyConfig().isEnableRemotingLocalProxyGrpc(); + ConfigurationManager.getProxyConfig().setEnableRemotingLocalProxyGrpc(true); http2ProtocolProxyHandler = new Http2ProtocolProxyHandler(); } + @After + public void tearDown() { + ConfigurationManager.getProxyConfig().setEnableRemotingLocalProxyGrpc(originalEnableRemotingLocalProxyGrpc); + } + @Test public void configPipeline() { when(inboundChannel.pipeline()).thenReturn(inboundPipeline); @@ -55,4 +70,19 @@ public void configPipeline() { when(outboundPipeline.addFirst(any(HAProxyMessageEncoder.class))).thenReturn(outboundPipeline); http2ProtocolProxyHandler.configPipeline(inboundChannel, outboundChannel); } -} \ No newline at end of file + + @Test + public void matchReturnsFalseForShortBuffers() { + assertFalse(http2ProtocolProxyHandler.match(Unpooled.EMPTY_BUFFER)); + + ByteBuf shortBuffer = Unpooled.wrappedBuffer(new byte[] {'P', 'R', 'I'}); + assertFalse(http2ProtocolProxyHandler.match(shortBuffer)); + } + + @Test + public void matchReturnsTrueForHttp2PrefacePrefix() { + ByteBuf http2Prefix = Unpooled.wrappedBuffer(new byte[] {'P', 'R', 'I', ' '}); + + assertTrue(http2ProtocolProxyHandler.match(http2Prefix)); + } +}