[ISSUE #3949] Support acl

This commit is contained in:
zhouxiang
2022-07-13 11:29:13 +08:00
parent b8b753ba97
commit 15d181cd70
10 changed files with 675 additions and 5 deletions
+8
View File
@@ -19,6 +19,10 @@
<name>rocketmq-acl ${project.version}</name>
<dependencies>
<dependency>
<groupId>${project.groupId}</groupId>
<artifactId>rocketmq-proto</artifactId>
</dependency>
<dependency>
<groupId>${project.groupId}</groupId>
<artifactId>rocketmq-remoting</artifactId>
@@ -62,6 +66,10 @@
<groupId>commons-validator</groupId>
<artifactId>commons-validator</artifactId>
</dependency>
<dependency>
<groupId>com.google.protobuf</groupId>
<artifactId>protobuf-java-util</artifactId>
</dependency>
</dependencies>
</project>
@@ -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.
*
@@ -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() + ")";
}
}
@@ -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;
}
}
@@ -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);
@@ -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;
}
}
@@ -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;
}
}
@@ -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<AccessValidator> 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();
@@ -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<AccessValidator> accessValidatorList;
public AuthenticationInterceptor(List<AccessValidator> accessValidatorList) {
this.accessValidatorList = accessValidatorList;
}
@Override
public <ReqT, RespT> ServerCall.Listener<ReqT> interceptCall(ServerCall<ReqT, RespT> call, Metadata headers,
ServerCallHandler<ReqT, RespT> next) {
return new ForwardingServerCallListener.SimpleForwardingServerCallListener<ReqT>(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));
}
}
};
}
}
@@ -62,4 +62,7 @@ public class InterceptorConstants {
public static final Metadata.Key<String> RPC_NAME
= Metadata.Key.of("x-mq-rpc-name", Metadata.ASCII_STRING_MARSHALLER);
public static final Metadata.Key<String> SESSION_TOKEN
= Metadata.Key.of("x-mq-session-token", Metadata.ASCII_STRING_MARSHALLER);
}