Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -116,6 +116,10 @@ public class ProxyConfig implements ConfigFile {
* max message body size, 0 or negative number means no limit for proxy
*/
private int maxMessageSize = 4 * 1024 * 1024;
/**
* max message count in one batch send request
*/
private int batchSendMaxMsgNum = 4096;
/**
* if true, proxy will check message body size and reject msg if it's body is empty
*/
Expand Down Expand Up @@ -595,6 +599,14 @@ public void setMaxMessageSize(int maxMessageSize) {
this.maxMessageSize = maxMessageSize;
}

public int getBatchSendMaxMsgNum() {
return batchSendMaxMsgNum;
}

public void setBatchSendMaxMsgNum(int batchSendMaxMsgNum) {
this.batchSendMaxMsgNum = batchSendMaxMsgNum;
}

public int getMaxUserPropertySize() {
return maxUserPropertySize;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,13 +34,16 @@
import java.util.HashSet;
import java.util.List;
import java.util.Map;
import java.util.Objects;
import java.util.Set;
import java.util.concurrent.CompletableFuture;
import org.apache.commons.lang3.StringUtils;
import org.apache.rocketmq.client.producer.SendResult;
import org.apache.rocketmq.common.attribute.TopicMessageType;
import org.apache.rocketmq.common.message.Message;
import org.apache.rocketmq.common.message.MessageAccessor;
import org.apache.rocketmq.common.message.MessageConst;
import org.apache.rocketmq.common.message.MessageDecoder;
import org.apache.rocketmq.common.sysflag.MessageSysFlag;
import org.apache.rocketmq.proxy.common.ProxyContext;
import org.apache.rocketmq.proxy.config.ConfigurationManager;
Expand Down Expand Up @@ -75,13 +78,16 @@ public CompletableFuture<SendMessageResponse> sendMessage(ProxyContext ctx, Send
apache.rocketmq.v2.Message message = messageList.get(0);
Resource topic = message.getTopic();
validateTopic(topic);
validateBatchEncoding(messageList);
List<Message> messages = buildMessage(ctx, messageList, topic);
validateBatchMessages(messages);

future = this.messagingProcessor.sendMessage(
ctx,
new SendMessageQueueSelector(request),
topic.getName(),
buildSysFlag(message),
buildMessage(ctx, request.getMessagesList(), topic)
messages
).thenApply(result -> convertToSendMessageResponse(ctx, request, result));
} catch (Throwable t) {
future.completeExceptionally(t);
Expand All @@ -103,6 +109,60 @@ protected List<Message> buildMessage(ProxyContext context, List<apache.rocketmq.
return messageExtList;
}

protected void validateBatchEncoding(List<apache.rocketmq.v2.Message> messageList) {
if (messageList.size() <= 1) {
return;
}
for (apache.rocketmq.v2.Message message : messageList) {
if (Encoding.GZIP.equals(message.getSystemProperties().getBodyEncoding())) {
throw new GrpcProxyException(Code.MESSAGE_CORRUPTED,
"batch send does not support compressed messages");
}
}
}

protected void validateBatchMessages(List<Message> messageList) {
if (messageList.size() <= 1) {
return;
}

ProxyConfig config = ConfigurationManager.getProxyConfig();
if (messageList.size() > config.getBatchSendMaxMsgNum()) {
throw new GrpcProxyException(Code.MESSAGE_CORRUPTED,
"batch message count cannot exceed the max " + config.getBatchSendMaxMsgNum());
}

Message firstMessage = messageList.get(0);
TopicMessageType messageType = TopicMessageType.parseFromMessageProperty(firstMessage.getProperties());
if (!TopicMessageType.NORMAL.equals(messageType) && !TopicMessageType.FIFO.equals(messageType)) {
throw new GrpcProxyException(Code.MESSAGE_PROPERTY_CONFLICT_WITH_TYPE,
"batch send only supports normal or FIFO messages");
}

String messageGroup = firstMessage.getProperty(MessageConst.PROPERTY_SHARDING_KEY);
if (TopicMessageType.FIFO.equals(messageType) && StringUtils.isBlank(messageGroup)) {
throw new GrpcProxyException(Code.ILLEGAL_MESSAGE_GROUP,
"message group cannot be empty for FIFO batch messages");
}
for (Message message : messageList) {
if (!messageType.equals(TopicMessageType.parseFromMessageProperty(message.getProperties()))) {
throw new GrpcProxyException(Code.MESSAGE_PROPERTY_CONFLICT_WITH_TYPE,
"all messages in a batch must have the same message type");
}
if (TopicMessageType.FIFO.equals(messageType)
&& !Objects.equals(messageGroup, message.getProperty(MessageConst.PROPERTY_SHARDING_KEY))) {
throw new GrpcProxyException(Code.MESSAGE_PROPERTY_CONFLICT_WITH_TYPE,
"all FIFO messages in a batch must have the same message group");
}
}

int maxMessageSize = config.getMaxMessageSize();
if (maxMessageSize > 0 && MessageDecoder.encodeMessages(messageList).length > maxMessageSize) {
throw new GrpcProxyException(Code.MESSAGE_BODY_TOO_LARGE,
"batch message body cannot exceed the max " + maxMessageSize);
}
}

protected Message buildMessage(ProxyContext context, apache.rocketmq.v2.Message protoMessage, String producerGroup) {
String topicName = protoMessage.getTopic().getName();

Expand Down Expand Up @@ -333,41 +393,21 @@ protected SendMessageResponse convertToSendMessageResponse(ProxyContext ctx, Sen
SendMessageResponse.Builder builder = SendMessageResponse.newBuilder();

Set<Code> responseCodes = new HashSet<>();
for (SendResult result : resultList) {
SendResultEntry resultEntry;
switch (result.getSendStatus()) {
case FLUSH_DISK_TIMEOUT:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.MASTER_PERSISTENCE_TIMEOUT, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
case FLUSH_SLAVE_TIMEOUT:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.SLAVE_PERSISTENCE_TIMEOUT, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
case SLAVE_NOT_AVAILABLE:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.HA_NOT_AVAILABLE, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
case SEND_OK:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.OK, Code.OK.name()))
.setOffset(result.getQueueOffset())
.setMessageId(StringUtils.defaultString(result.getMsgId()))
.setTransactionId(StringUtils.defaultString(result.getTransactionId()))
.setRecallHandle(StringUtils.defaultString(result.getRecallHandle()))
.build();
break;
default:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.INTERNAL_SERVER_ERROR, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
if (request.getMessagesCount() > 1 && resultList.size() == 1) {
SendResult batchResult = resultList.get(0);
for (int i = 0; i < request.getMessagesCount(); i++) {
SendResultEntry resultEntry = convertToSendResultEntry(batchResult,
request.getMessages(i).getSystemProperties().getMessageId(), batchResult.getQueueOffset() + i);
builder.addEntries(resultEntry);
responseCodes.add(resultEntry.getStatus().getCode());
}
} else {
for (SendResult result : resultList) {
SendResultEntry resultEntry = convertToSendResultEntry(result,
StringUtils.defaultString(result.getMsgId()), result.getQueueOffset());
builder.addEntries(resultEntry);
responseCodes.add(resultEntry.getStatus().getCode());
}
builder.addEntries(resultEntry);
responseCodes.add(resultEntry.getStatus().getCode());
}
if (responseCodes.size() > 1) {
builder.setStatus(ResponseBuilder.getInstance().buildStatus(Code.MULTIPLE_RESULTS, Code.MULTIPLE_RESULTS.name()));
Expand All @@ -380,6 +420,42 @@ protected SendMessageResponse convertToSendMessageResponse(ProxyContext ctx, Sen
return builder.build();
}

protected SendResultEntry convertToSendResultEntry(SendResult result, String messageId, long queueOffset) {
SendResultEntry resultEntry;
switch (result.getSendStatus()) {
case FLUSH_DISK_TIMEOUT:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.MASTER_PERSISTENCE_TIMEOUT, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
case FLUSH_SLAVE_TIMEOUT:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.SLAVE_PERSISTENCE_TIMEOUT, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
case SLAVE_NOT_AVAILABLE:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.HA_NOT_AVAILABLE, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
case SEND_OK:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.OK, Code.OK.name()))
.setOffset(queueOffset)
.setMessageId(StringUtils.defaultString(messageId))
.setTransactionId(StringUtils.defaultString(result.getTransactionId()))
.setRecallHandle(StringUtils.defaultString(result.getRecallHandle()))
.build();
break;
default:
resultEntry = SendResultEntry.newBuilder()
.setStatus(ResponseBuilder.getInstance().buildStatus(Code.INTERNAL_SERVER_ERROR, "send message failed, sendStatus=" + result.getSendStatus()))
.build();
break;
}
return resultEntry;
}

protected static class SendMessageQueueSelector implements QueueSelector {

private final SendMessageRequest request;
Expand All @@ -392,13 +468,10 @@ public SendMessageQueueSelector(SendMessageRequest request) {
public AddressableMessageQueue select(ProxyContext ctx, MessageQueueView messageQueueView) {
try {
apache.rocketmq.v2.Message message = request.getMessages(0);
String shardingKey = null;
if (request.getMessagesCount() == 1) {
shardingKey = message.getSystemProperties().getMessageGroup();
// lite topic
if (StringUtils.isBlank(shardingKey)) {
shardingKey = message.getSystemProperties().getLiteTopic();
}
String shardingKey = message.getSystemProperties().getMessageGroup();
// lite topic
if (StringUtils.isBlank(shardingKey)) {
shardingKey = message.getSystemProperties().getLiteTopic();
}
AddressableMessageQueue targetMessageQueue;
if (StringUtils.isNotEmpty(shardingKey)) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
import com.google.protobuf.util.Timestamps;
import java.time.Duration;
import java.util.HashMap;
import java.util.List;
import java.util.Map;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.ExecutionException;
Expand Down Expand Up @@ -58,6 +59,7 @@
import org.assertj.core.util.Lists;
import org.junit.Before;
import org.junit.Test;
import org.mockito.ArgumentCaptor;

import static org.apache.rocketmq.proxy.service.route.TopicRouteService.buildPenalizerByMQFaultStrategy;
import static org.junit.Assert.assertEquals;
Expand All @@ -68,6 +70,9 @@
import static org.mockito.ArgumentMatchers.anyInt;
import static org.mockito.ArgumentMatchers.anyString;
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.when;

public class SendMessageActivityTest extends BaseActivityTest {
Expand Down Expand Up @@ -122,6 +127,115 @@ public void sendMessage() throws Exception {
assertEquals(msgId, response.getEntries(0).getMessageId());
}

@Test
@SuppressWarnings("unchecked")
public void testSendFifoBatchInOneInvocation() throws Exception {
String firstMessageId = MessageClientIDSetter.createUniqID();
String secondMessageId = MessageClientIDSetter.createUniqID();
String messageGroup = "group";
SendResult sendResult = new SendResult(SendStatus.SEND_OK, null, null, null, 10);
when(this.messagingProcessor.sendMessage(any(), any(), anyString(), anyInt(), any()))
.thenReturn(CompletableFuture.completedFuture(Lists.newArrayList(sendResult)));

SendMessageRequest request = SendMessageRequest.newBuilder()
.addMessages(createMessage(firstMessageId, MessageType.FIFO, messageGroup, 16))
.addMessages(createMessage(secondMessageId, MessageType.FIFO, messageGroup, 16))
.build();
SendMessageResponse response = this.sendMessageActivity.sendMessage(createContext(), request).get();

ArgumentCaptor<List<org.apache.rocketmq.common.message.Message>> messageListCaptor =
ArgumentCaptor.forClass(List.class);
verify(this.messagingProcessor, times(1)).sendMessage(any(), any(), anyString(), anyInt(),
messageListCaptor.capture());
assertEquals(2, messageListCaptor.getValue().size());
assertEquals(2, response.getEntriesCount());
assertEquals(firstMessageId, response.getEntries(0).getMessageId());
assertEquals(secondMessageId, response.getEntries(1).getMessageId());
assertEquals(10, response.getEntries(0).getOffset());
assertEquals(11, response.getEntries(1).getOffset());
}

@Test
public void testRejectFifoBatchWithDifferentMessageGroups() {
SendMessageRequest request = SendMessageRequest.newBuilder()
.addMessages(createMessage(MessageClientIDSetter.createUniqID(), MessageType.FIFO, "group-a", 16))
.addMessages(createMessage(MessageClientIDSetter.createUniqID(), MessageType.FIFO, "group-b", 16))
.build();

ExecutionException exception = assertThrows(ExecutionException.class,
() -> this.sendMessageActivity.sendMessage(createContext(), request).get());
GrpcProxyException cause = (GrpcProxyException) exception.getCause();
assertEquals(Code.MESSAGE_PROPERTY_CONFLICT_WITH_TYPE, cause.getCode());
}

@Test
public void testRejectCompressedBatch() {
Message compressedMessage = createMessage(
MessageClientIDSetter.createUniqID(), MessageType.NORMAL, "", 16);
compressedMessage = compressedMessage.toBuilder()
.setSystemProperties(compressedMessage.getSystemProperties().toBuilder()
.setBodyEncoding(Encoding.GZIP))
.build();
SendMessageRequest request = SendMessageRequest.newBuilder()
.addMessages(compressedMessage)
.addMessages(createMessage(MessageClientIDSetter.createUniqID(), MessageType.NORMAL, "", 16))
.build();

ExecutionException exception = assertThrows(ExecutionException.class,
() -> this.sendMessageActivity.sendMessage(createContext(), request).get());
GrpcProxyException cause = (GrpcProxyException) exception.getCause();
assertEquals(Code.MESSAGE_CORRUPTED, cause.getCode());
verify(this.messagingProcessor, never()).sendMessage(any(), any(), anyString(), anyInt(), any());
}

@Test
public void testRejectBatchWhoseEncodedBodyExceedsLimit() {
int previousMaxMessageSize = ConfigurationManager.getProxyConfig().getMaxMessageSize();
ConfigurationManager.getProxyConfig().setMaxMessageSize(80);
try {
List<org.apache.rocketmq.common.message.Message> messages = Lists.newArrayList(
new org.apache.rocketmq.common.message.Message(TOPIC, new byte[30]),
new org.apache.rocketmq.common.message.Message(TOPIC, new byte[30]));

GrpcProxyException exception = assertThrows(GrpcProxyException.class,
() -> this.sendMessageActivity.validateBatchMessages(messages));
assertEquals(Code.MESSAGE_BODY_TOO_LARGE, exception.getCode());
} finally {
ConfigurationManager.getProxyConfig().setMaxMessageSize(previousMaxMessageSize);
}
}

@Test
public void testRejectBatchWhoseMessageCountExceedsLimit() {
int previousMaxMessageCount = ConfigurationManager.getProxyConfig().getBatchSendMaxMsgNum();
ConfigurationManager.getProxyConfig().setBatchSendMaxMsgNum(1);
try {
List<org.apache.rocketmq.common.message.Message> messages = Lists.newArrayList(
new org.apache.rocketmq.common.message.Message(TOPIC, new byte[1]),
new org.apache.rocketmq.common.message.Message(TOPIC, new byte[1]));

GrpcProxyException exception = assertThrows(GrpcProxyException.class,
() -> this.sendMessageActivity.validateBatchMessages(messages));
assertEquals(Code.MESSAGE_CORRUPTED, exception.getCode());
} finally {
ConfigurationManager.getProxyConfig().setBatchSendMaxMsgNum(previousMaxMessageCount);
}
}

private Message createMessage(String messageId, MessageType messageType, String messageGroup, int bodySize) {
return Message.newBuilder()
.setTopic(Resource.newBuilder().setName(TOPIC).build())
.setSystemProperties(SystemProperties.newBuilder()
.setMessageId(messageId)
.setMessageType(messageType)
.setMessageGroup(messageGroup)
.setBornTimestamp(Timestamps.fromMillis(System.currentTimeMillis()))
.setBornHost(StringUtils.defaultString(NetworkUtil.getLocalAddress(), "127.0.0.1:1234"))
.build())
.setBody(ByteString.copyFrom(new byte[bodySize]))
.build();
}

@Test
public void testConvertToSendMessageResponse() {
{
Expand Down Expand Up @@ -362,6 +476,11 @@ public void testSendOrderMessageQueueSelector() throws Exception {

SendMessageActivity.SendMessageQueueSelector selector2 = new SendMessageActivity.SendMessageQueueSelector(
SendMessageRequest.newBuilder()
.addMessages(Message.newBuilder()
.setSystemProperties(SystemProperties.newBuilder()
.setMessageGroup(String.valueOf(1))
.build())
.build())
.addMessages(Message.newBuilder()
.setSystemProperties(SystemProperties.newBuilder()
.setMessageGroup(String.valueOf(1))
Expand Down
Loading
Loading