diff --git a/acl/pom.xml b/acl/pom.xml index 686a398540..c80cc24b5c 100644 --- a/acl/pom.xml +++ b/acl/pom.xml @@ -19,6 +19,10 @@ rocketmq-acl ${project.version} + + ${project.groupId} + rocketmq-proto + ${project.groupId} rocketmq-remoting @@ -62,6 +66,10 @@ commons-validator commons-validator + + com.google.protobuf + protobuf-java-util + diff --git a/acl/src/main/java/org/apache/rocketmq/acl/AccessValidator.java b/acl/src/main/java/org/apache/rocketmq/acl/AccessValidator.java index 167fa26e88..8602525b36 100644 --- a/acl/src/main/java/org/apache/rocketmq/acl/AccessValidator.java +++ b/acl/src/main/java/org/apache/rocketmq/acl/AccessValidator.java @@ -17,9 +17,10 @@ package org.apache.rocketmq.acl; +import com.google.protobuf.GeneratedMessageV3; import java.util.List; import java.util.Map; - +import org.apache.rocketmq.acl.common.MetadataHeader; import org.apache.rocketmq.common.AclConfig; import org.apache.rocketmq.common.DataVersion; import org.apache.rocketmq.common.PlainAccessConfig; @@ -36,6 +37,14 @@ public interface AccessValidator { */ AccessResource parse(RemotingCommand request, String remoteAddr); + /** + * Parse to get the AccessResource from gRPC protocol + * @param messageV3 + * @param header + * @return Plain access resource + */ + AccessResource parse(GeneratedMessageV3 messageV3, MetadataHeader header); + /** * Validate the access resource. * diff --git a/acl/src/main/java/org/apache/rocketmq/acl/common/AuthorizationHeader.java b/acl/src/main/java/org/apache/rocketmq/acl/common/AuthorizationHeader.java new file mode 100644 index 0000000000..5fd053a5f7 --- /dev/null +++ b/acl/src/main/java/org/apache/rocketmq/acl/common/AuthorizationHeader.java @@ -0,0 +1,137 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.rocketmq.acl.common; + +import org.apache.commons.codec.DecoderException; +import org.apache.commons.codec.binary.Base64; +import org.apache.commons.codec.binary.Hex; + +public class AuthorizationHeader { + private static final String HEADER_SEPARATOR = " "; + private static final String CREDENTIALS_SEPARATOR = "/"; + private static final int AUTH_HEADER_KV_LENGTH = 2; + private static final int CREDENTIALS_LENGTH = 3; + private static final String CREDENTIAL = "Credential"; + private static final String SIGNED_HEADERS = "SignedHeaders"; + private static final String SIGNATURE = "Signature"; + private String method; + private String accessKey; + private String regionId; + private String channelKey; + private String[] signedHeaders; + private String signature; + + /** + * Parse authorization from gRPC header. + * + * @param header gRPC header string. + * @throws Exception exception. + */ + public AuthorizationHeader(String header) throws DecoderException { + String[] result = header.split(HEADER_SEPARATOR, 2); + if (result.length != 2) { + throw new DecoderException("authorization header is incorrect"); + } + this.method = result[0]; + String[] keyValues = result[1].split(","); + for (String keyValue : keyValues) { + String[] kv = keyValue.trim().split("=", 2); + int kvLength = kv.length; + if (kv.length != AUTH_HEADER_KV_LENGTH) { + throw new DecoderException("authorization keyValues length is incorrect, actual length=" + kvLength); + } + String authItem = kv[0]; + if (CREDENTIAL.equals(authItem)) { + String[] credential = kv[1].split(CREDENTIALS_SEPARATOR); + int credentialActualLength = credential.length; + if (credentialActualLength < CREDENTIALS_LENGTH) { + throw new DecoderException("authorization credential length is incorrect, actual length=" + credentialActualLength); + } + this.accessKey = credential[0]; + this.regionId = credential[1]; + this.channelKey = credential[2]; + continue; + } + if (SIGNED_HEADERS.equals(authItem)) { + this.signedHeaders = kv[1].split(";"); + continue; + } + if (SIGNATURE.equals(authItem)) { + this.signature = this.hexToBase64(kv[1]); + } + } + } + + public String hexToBase64(String input) throws DecoderException { + byte[] bytes = Hex.decodeHex(input); + return Base64.encodeBase64String(bytes); + } + + public String getMethod() { + return this.method; + } + + public String getAccessKey() { + return this.accessKey; + } + + public String getRegionId() { + return this.regionId; + } + + public String getChannelKey() { + return this.channelKey; + } + + public String[] getSignedHeaders() { + return this.signedHeaders; + } + + public String getSignature() { + return this.signature; + } + + public void setMethod(final String method) { + this.method = method; + } + + public void setAccessKey(final String accessKey) { + this.accessKey = accessKey; + } + + public void setRegionId(final String regionId) { + this.regionId = regionId; + } + + public void setChannelKey(final String channelKey) { + this.channelKey = channelKey; + } + + public void setSignedHeaders(final String[] signedHeaders) { + this.signedHeaders = signedHeaders; + } + + public void setSignature(final String signature) { + this.signature = signature; + } + + @java.lang.Override + public java.lang.String toString() { + return "GrpcAuthHeader(method=" + this.getMethod() + ", accessKey=" + this.getAccessKey() + ", regionId=" + this.getRegionId() + ", channelKey=" + this.getChannelKey() + ", signedHeaders=" + java.util.Arrays.deepToString(this.getSignedHeaders()) + ", signature=" + this.getSignature() + ")"; + } +} diff --git a/acl/src/main/java/org/apache/rocketmq/acl/common/MetadataHeader.java b/acl/src/main/java/org/apache/rocketmq/acl/common/MetadataHeader.java new file mode 100644 index 0000000000..a6824918a3 --- /dev/null +++ b/acl/src/main/java/org/apache/rocketmq/acl/common/MetadataHeader.java @@ -0,0 +1,233 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.rocketmq.acl.common; + +public class MetadataHeader { + private String remoteAddress; + private String tenantId; + private String namespace; + private String authorization; + private String datetime; + private String sessionToken; + private String requestId; + private String language; + private String clientVersion; + private String protocol; + private int requestCode; + + MetadataHeader(final String remoteAddress, final String tenantId, final String namespace, + final String authorization, final String datetime, final String sessionToken, final String requestId, + final String language, final String clientVersion, final String protocol, final int requestCode) { + this.remoteAddress = remoteAddress; + this.tenantId = tenantId; + this.namespace = namespace; + this.authorization = authorization; + this.datetime = datetime; + this.sessionToken = sessionToken; + this.requestId = requestId; + this.language = language; + this.clientVersion = clientVersion; + this.protocol = protocol; + this.requestCode = requestCode; + } + + public static class MetadataHeaderBuilder { + private String remoteAddress; + private String tenantId; + private String namespace; + private String authorization; + private String datetime; + private String sessionToken; + private String requestId; + private String language; + private String clientVersion; + private String protocol; + private int requestCode; + + MetadataHeaderBuilder() { + } + + public MetadataHeader.MetadataHeaderBuilder remoteAddress(final String remoteAddress) { + this.remoteAddress = remoteAddress; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder tenantId(final String tenantId) { + this.tenantId = tenantId; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder namespace(final String namespace) { + this.namespace = namespace; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder authorization(final String authorization) { + this.authorization = authorization; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder datetime(final String datetime) { + this.datetime = datetime; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder sessionToken(final String sessionToken) { + this.sessionToken = sessionToken; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder requestId(final String requestId) { + this.requestId = requestId; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder language(final String language) { + this.language = language; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder clientVersion(final String clientVersion) { + this.clientVersion = clientVersion; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder protocol(final String protocol) { + this.protocol = protocol; + return this; + } + + public MetadataHeader.MetadataHeaderBuilder requestCode(final int requestCode) { + this.requestCode = requestCode; + return this; + } + + public MetadataHeader build() { + return new MetadataHeader(this.remoteAddress, this.tenantId, this.namespace, this.authorization, + this.datetime, this.sessionToken, this.requestId, this.language, this.clientVersion, this.protocol, + this.requestCode); + } + + @Override public String toString() { + return "MetadataHeaderBuilder{" + "remoteAddress='" + remoteAddress + '\'' + + ", tenantId='" + tenantId + '\'' + + ", namespace='" + namespace + '\'' + + ", authorization='" + authorization + '\'' + + ", datetime='" + datetime + '\'' + + ", sessionToken='" + sessionToken + '\'' + + ", requestId='" + requestId + '\'' + + ", language='" + language + '\'' + + ", clientVersion='" + clientVersion + '\'' + + ", protocol='" + protocol + '\'' + + ", requestCode=" + requestCode + + '}'; + } + } + + public static MetadataHeader.MetadataHeaderBuilder builder() { + return new MetadataHeader.MetadataHeaderBuilder(); + } + + public String getRemoteAddress() { + return this.remoteAddress; + } + + public String getTenantId() { + return this.tenantId; + } + + public String getNamespace() { + return this.namespace; + } + + public String getAuthorization() { + return this.authorization; + } + + public String getDatetime() { + return this.datetime; + } + + public String getSessionToken() { + return this.sessionToken; + } + + public String getRequestId() { + return this.requestId; + } + + public String getLanguage() { + return this.language; + } + + public String getClientVersion() { + return this.clientVersion; + } + + public String getProtocol() { + return this.protocol; + } + + public int getRequestCode() { + return this.requestCode; + } + + public void setRemoteAddress(final String remoteAddress) { + this.remoteAddress = remoteAddress; + } + + public void setTenantId(final String tenantId) { + this.tenantId = tenantId; + } + + public void setNamespace(final String namespace) { + this.namespace = namespace; + } + + public void setAuthorization(final String authorization) { + this.authorization = authorization; + } + + public void setDatetime(final String datetime) { + this.datetime = datetime; + } + + public void setSessionToken(final String sessionToken) { + this.sessionToken = sessionToken; + } + + public void setRequestId(final String requestId) { + this.requestId = requestId; + } + + public void setLanguage(final String language) { + this.language = language; + } + + public void setClientVersion(final String clientVersion) { + this.clientVersion = clientVersion; + } + + public void setProtocol(final String protocol) { + this.protocol = protocol; + } + + public void setRequestCode(int requestCode) { + this.requestCode = requestCode; + } +} diff --git a/acl/src/main/java/org/apache/rocketmq/acl/plain/PlainAccessValidator.java b/acl/src/main/java/org/apache/rocketmq/acl/plain/PlainAccessValidator.java index 83f43ef7eb..6e1f78463e 100644 --- a/acl/src/main/java/org/apache/rocketmq/acl/plain/PlainAccessValidator.java +++ b/acl/src/main/java/org/apache/rocketmq/acl/plain/PlainAccessValidator.java @@ -16,20 +16,36 @@ */ package org.apache.rocketmq.acl.plain; +import apache.rocketmq.v1.AckMessageRequest; +import apache.rocketmq.v1.EndTransactionRequest; +import apache.rocketmq.v1.ForwardMessageToDeadLetterQueueRequest; +import apache.rocketmq.v1.HeartbeatRequest; +import apache.rocketmq.v1.NackMessageRequest; +import apache.rocketmq.v1.PullMessageRequest; +import apache.rocketmq.v1.QueryOffsetRequest; +import apache.rocketmq.v1.ReceiveMessageRequest; +import apache.rocketmq.v1.Resource; +import apache.rocketmq.v1.SendMessageRequest; +import com.google.protobuf.GeneratedMessageV3; +import java.nio.charset.StandardCharsets; import java.util.List; import java.util.Map; import java.util.SortedMap; import java.util.TreeMap; +import org.apache.commons.codec.DecoderException; import org.apache.rocketmq.acl.AccessResource; import org.apache.rocketmq.acl.AccessValidator; import org.apache.rocketmq.acl.common.AclException; import org.apache.rocketmq.acl.common.AclUtils; +import org.apache.rocketmq.acl.common.AuthorizationHeader; +import org.apache.rocketmq.acl.common.MetadataHeader; import org.apache.rocketmq.acl.common.Permission; import org.apache.rocketmq.acl.common.SessionCredentials; import org.apache.rocketmq.common.AclConfig; import org.apache.rocketmq.common.DataVersion; import org.apache.rocketmq.common.MixAll; import org.apache.rocketmq.common.PlainAccessConfig; +import org.apache.rocketmq.common.protocol.NamespaceUtil; import org.apache.rocketmq.common.protocol.RequestCode; import org.apache.rocketmq.common.protocol.header.GetConsumerListByGroupRequestHeader; import org.apache.rocketmq.common.protocol.header.UnregisterClientRequestHeader; @@ -136,6 +152,102 @@ public class PlainAccessValidator implements AccessValidator { return accessResource; } + @Override public AccessResource parse(GeneratedMessageV3 messageV3, MetadataHeader header) { + PlainAccessResource accessResource = new PlainAccessResource(); + String remoteAddress = header.getRemoteAddress(); + if (remoteAddress != null && remoteAddress.contains(":")) { + accessResource.setWhiteRemoteAddress(remoteAddress.substring(0, remoteAddress.lastIndexOf(':'))); + } else { + accessResource.setWhiteRemoteAddress(remoteAddress); + } + try { + AuthorizationHeader authorizationHeader = new AuthorizationHeader(header.getAuthorization()); + accessResource.setAccessKey(authorizationHeader.getAccessKey()); + accessResource.setSignature(authorizationHeader.getSignature()); + } catch (DecoderException e) { + throw new AclException(e.getMessage(), e); + } + accessResource.setSecretToken(header.getSessionToken()); + accessResource.setRequestCode(header.getRequestCode()); + accessResource.setContent(header.getDatetime().getBytes(StandardCharsets.UTF_8)); + + try { + String rpcFullName = messageV3.getDescriptorForType().getFullName(); + if (HeartbeatRequest.getDescriptor().getFullName().equals(rpcFullName)) { + HeartbeatRequest request = (HeartbeatRequest) messageV3; + if (request.hasProducerData()) { + Resource group = request.getProducerData() + .getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + } else if (request.hasConsumerData()) { + Resource group = request.getConsumerData() + .getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + } + } else if (SendMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + SendMessageRequest request = (SendMessageRequest) messageV3; + Resource topic = request.getMessage().getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.PUB); + } else if (ReceiveMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + ReceiveMessageRequest request = (ReceiveMessageRequest) messageV3; + Resource group = request.getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + Resource topic = request.getPartition().getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.SUB); + } else if (AckMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + AckMessageRequest request = (AckMessageRequest) messageV3; + Resource group = request.getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + Resource topic = request.getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.SUB); + } else if (NackMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + NackMessageRequest request = (NackMessageRequest) messageV3; + Resource group = request.getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + Resource topic = request.getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.SUB); + } else if (ForwardMessageToDeadLetterQueueRequest.getDescriptor().getFullName().equals(rpcFullName)) { + ForwardMessageToDeadLetterQueueRequest request = (ForwardMessageToDeadLetterQueueRequest) messageV3; + Resource group = request.getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + Resource topic = request.getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.SUB); + } else if (EndTransactionRequest.getDescriptor().getFullName().equals(rpcFullName)) { + EndTransactionRequest request = (EndTransactionRequest) messageV3; + Resource group = request.getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.PUB); + } else if (QueryOffsetRequest.getDescriptor().getFullName().equals(rpcFullName)) { + QueryOffsetRequest request = (QueryOffsetRequest) messageV3; + Resource topic = request.getPartition().getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.SUB); + } else if (PullMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + PullMessageRequest request = (PullMessageRequest) messageV3; + Resource group = request.getGroup(); + String groupName = NamespaceUtil.wrapNamespace(group.getResourceNamespace(), group.getName()); + accessResource.addResourceAndPerm(groupName, Permission.SUB); + Resource topic = request.getPartition().getTopic(); + String topicName = NamespaceUtil.wrapNamespace(topic.getResourceNamespace(), topic.getName()); + accessResource.addResourceAndPerm(topicName, Permission.SUB); + } + } catch (Throwable t) { + throw new AclException(t.getMessage(), t); + } + return accessResource; + } + @Override public void validate(AccessResource accessResource) { aclPlugEngine.validate((PlainAccessResource) accessResource); diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/common/RequestMapping.java b/proxy/src/main/java/org/apache/rocketmq/proxy/common/RequestMapping.java new file mode 100644 index 0000000000..7a2ae4dd22 --- /dev/null +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/common/RequestMapping.java @@ -0,0 +1,82 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.rocketmq.proxy.common; + +import apache.rocketmq.v1.AckMessageRequest; +import apache.rocketmq.v1.ChangeInvisibleDurationRequest; +import apache.rocketmq.v1.EndTransactionRequest; +import apache.rocketmq.v1.ForwardMessageToDeadLetterQueueResponse; +import apache.rocketmq.v1.HealthCheckRequest; +import apache.rocketmq.v1.HeartbeatRequest; +import apache.rocketmq.v1.NackMessageRequest; +import apache.rocketmq.v1.NotifyClientTerminationRequest; +import apache.rocketmq.v1.PullMessageRequest; +import apache.rocketmq.v1.QueryAssignmentRequest; +import apache.rocketmq.v1.QueryOffsetRequest; +import apache.rocketmq.v1.QueryRouteRequest; +import apache.rocketmq.v1.ReceiveMessageRequest; +import apache.rocketmq.v1.SendMessageRequest; +import org.apache.rocketmq.common.protocol.RequestCode; + +public class RequestMapping { + public static int map(String rpcFullName) { + if (QueryRouteRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.GET_ROUTEINFO_BY_TOPIC; + } + if (HeartbeatRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.HEART_BEAT; + } + if (HealthCheckRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.HEART_BEAT; + } + if (SendMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.SEND_MESSAGE_V2; + } + if (QueryAssignmentRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.GET_ROUTEINFO_BY_TOPIC; + } + if (ReceiveMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.PULL_MESSAGE; + } + if (AckMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.UPDATE_CONSUMER_OFFSET; + } + if (NackMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.CONSUMER_SEND_MSG_BACK; + } + if (ForwardMessageToDeadLetterQueueResponse.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.CONSUMER_SEND_MSG_BACK; + } + if (EndTransactionRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.END_TRANSACTION; + } + if (QueryOffsetRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.SEARCH_OFFSET_BY_TIMESTAMP; + } + if (PullMessageRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.PULL_MESSAGE; + } + if (NotifyClientTerminationRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.UNREGISTER_CLIENT; + } + if (ChangeInvisibleDurationRequest.getDescriptor().getFullName().equals(rpcFullName)) { + return RequestCode.CONSUMER_SEND_MSG_BACK; + } + return RequestCode.HEART_BEAT; + } +} diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/config/ProxyConfig.java b/proxy/src/main/java/org/apache/rocketmq/proxy/config/ProxyConfig.java index cb23d66c6b..494bf75177 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/config/ProxyConfig.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/config/ProxyConfig.java @@ -82,6 +82,8 @@ public class ProxyConfig { private int retryDelayLevelDelta = 3; private String messageDelayLevel = "1s 5s 10s 30s 1m 2m 3m 4m 5m 6m 7m 8m 9m 10m 20m 30m 1h 2h"; + private boolean enableACL = false; + public Integer getHealthCheckPort() { return healthCheckPort; } @@ -385,4 +387,12 @@ public class ProxyConfig { public void setMessageDelayLevel(String messageDelayLevel) { this.messageDelayLevel = messageDelayLevel; } + + public boolean isEnableACL() { + return enableACL; + } + + public void setEnableACL(boolean enableACL) { + this.enableACL = enableACL; + } } diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/GrpcServer.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/GrpcServer.java index c5e318a245..dc36a5e13f 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/GrpcServer.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/GrpcServer.java @@ -27,12 +27,16 @@ import io.grpc.netty.shaded.io.netty.handler.ssl.util.InsecureTrustManagerFactor import java.io.FileInputStream; import java.io.IOException; import java.io.InputStream; +import java.util.List; import java.util.concurrent.ThreadPoolExecutor; import java.util.concurrent.TimeUnit; +import org.apache.rocketmq.acl.AccessValidator; +import org.apache.rocketmq.broker.util.ServiceProvider; import org.apache.rocketmq.common.constant.LoggerName; import org.apache.rocketmq.common.thread.ThreadPoolMonitor; import org.apache.rocketmq.proxy.common.StartAndShutdown; import org.apache.rocketmq.proxy.config.ConfigurationManager; +import org.apache.rocketmq.proxy.grpc.interceptor.AuthenticationInterceptor; import org.apache.rocketmq.proxy.grpc.interceptor.ContextInterceptor; import org.apache.rocketmq.proxy.grpc.interceptor.HeaderInterceptor; import org.apache.rocketmq.proxy.grpc.service.GrpcForwardService; @@ -86,14 +90,22 @@ public class GrpcServer implements StartAndShutdown { int workerLoopNum = ConfigurationManager.getProxyConfig().getGrpcWorkerLoopNum(); int maxInboundMessageSize = ConfigurationManager.getProxyConfig().getGrpcMaxInboundMessageSize(); - this.server = serverBuilder - .maxInboundMessageSize(maxInboundMessageSize) + serverBuilder.maxInboundMessageSize(maxInboundMessageSize) .bossEventLoopGroup(new NioEventLoopGroup(bossLoopNum)) .workerEventLoopGroup(new NioEventLoopGroup(workerLoopNum)) .channelType(NioServerSocketChannel.class) .addService(messagingProcessor) - .executor(this.executor) - .intercept(new ContextInterceptor()) + .executor(this.executor); + + if (ConfigurationManager.getProxyConfig().isEnableACL()) { + List accessValidators = ServiceProvider.load(ServiceProvider.ACL_VALIDATOR_ID, AccessValidator.class); + if (accessValidators.isEmpty()) { + throw new IllegalArgumentException("Load AccessValidator failed"); + } + serverBuilder.intercept(new AuthenticationInterceptor(accessValidators)); + } + + this.server = serverBuilder.intercept(new ContextInterceptor()) .intercept(new HeaderInterceptor()) .build(); diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/AuthenticationInterceptor.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/AuthenticationInterceptor.java new file mode 100644 index 0000000000..f691399c85 --- /dev/null +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/AuthenticationInterceptor.java @@ -0,0 +1,64 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.rocketmq.proxy.grpc.interceptor; + +import com.google.protobuf.GeneratedMessageV3; +import io.grpc.Context; +import io.grpc.ForwardingServerCallListener; +import io.grpc.Metadata; +import io.grpc.ServerCall; +import io.grpc.ServerCallHandler; +import io.grpc.ServerInterceptor; +import java.util.List; +import org.apache.rocketmq.acl.AccessValidator; +import org.apache.rocketmq.acl.common.MetadataHeader; +import org.apache.rocketmq.proxy.common.RequestMapping; + +public class AuthenticationInterceptor implements ServerInterceptor { + private final List accessValidatorList; + + public AuthenticationInterceptor(List accessValidatorList) { + this.accessValidatorList = accessValidatorList; + } + + @Override + public ServerCall.Listener interceptCall(ServerCall call, Metadata headers, + ServerCallHandler next) { + return new ForwardingServerCallListener.SimpleForwardingServerCallListener(next.startCall(call, headers)) { + @Override + public void onMessage(ReqT message) { + GeneratedMessageV3 messageV3 = (GeneratedMessageV3) message; + MetadataHeader metadataHeader = MetadataHeader.builder() + .remoteAddress(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.REMOTE_ADDRESS)) + .namespace(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.NAMESPACE_ID)) + .authorization(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.AUTHORIZATION)) + .datetime(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.DATE_TIME)) + .sessionToken(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.SESSION_TOKEN)) + .requestId(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.REQUEST_ID)) + .language(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.LANGUAGE)) + .clientVersion(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.CLIENT_VERSION)) + .protocol(InterceptorConstants.METADATA.get(Context.current()).get(InterceptorConstants.PROTOCOL_VERSION)) + .requestCode(RequestMapping.map(messageV3.getDescriptorForType().getFullName())) + .build(); + for (AccessValidator accessValidator : accessValidatorList) { + accessValidator.validate(accessValidator.parse(messageV3, metadataHeader)); + } + } + }; + } +} diff --git a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/InterceptorConstants.java b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/InterceptorConstants.java index 5a672f43ea..73578a39d7 100644 --- a/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/InterceptorConstants.java +++ b/proxy/src/main/java/org/apache/rocketmq/proxy/grpc/interceptor/InterceptorConstants.java @@ -62,4 +62,7 @@ public class InterceptorConstants { public static final Metadata.Key RPC_NAME = Metadata.Key.of("x-mq-rpc-name", Metadata.ASCII_STRING_MARSHALLER); + + public static final Metadata.Key SESSION_TOKEN + = Metadata.Key.of("x-mq-session-token", Metadata.ASCII_STRING_MARSHALLER); }