Skip to content
Merged
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 @@ -158,6 +158,26 @@ public final void channelRead(ChannelHandlerContext ctx, Object msg) {
LWT willMessage =
connMsg.variableHeader().isWillFlag() ? getWillMessage(connMsg, clientInfo) : null;

CompletableFuture<AuthResult> willPermission = willMessage == null
? CompletableFuture.completedFuture(AuthResult.ok(okOrGoAway.successInfo))
: checkWillPermission(connMsg, willMessage, okOrGoAway.successInfo);
return willPermission.thenApply(
permissionResult -> new SessionPreparation(permissionResult, settings, willMessage));
}
}, ctx.executor())
.thenComposeAsync(sessionPreparation -> {
if (sessionPreparation == null) {
return CompletableFuture.completedFuture(null);
}
AuthResult okOrGoAway = sessionPreparation.permissionResult;
if (okOrGoAway.goAway != null) {
handleGoAway(okOrGoAway.goAway);
return CompletableFuture.completedFuture(null);
} else {
ClientInfo clientInfo = okOrGoAway.successInfo.clientInfo;
TenantSettings settings = sessionPreparation.settings;
LWT willMessage = sessionPreparation.willMessage;

int keepAliveSeconds = keepAliveSeconds(connMsg.variableHeader().keepAliveTimeSeconds(),
settings);
String userSessionId = userSessionId(clientInfo);
Expand Down Expand Up @@ -449,6 +469,10 @@ private CompletableFuture<ExpireResult> expireInbox(long reqId,
protected abstract CompletableFuture<AuthResult> checkConnectPermission(MqttConnectMessage message,
SuccessInfo successInfo);

protected abstract CompletableFuture<AuthResult> checkWillPermission(MqttConnectMessage message,
LWT willMessage,
SuccessInfo successInfo);

protected abstract void handleMqttMessage(MqttMessage message);

protected abstract GoAway onNoEnoughResources(MqttConnectMessage message, TenantResourceType resourceType,
Expand Down Expand Up @@ -661,6 +685,11 @@ private enum ExpireResult {
ERROR
}

private record SessionPreparation(AuthResult permissionResult,
TenantSettings settings,
LWT willMessage) {
}

public record SuccessInfo(ClientInfo clientInfo,
Optional<String> responseInfo, // mqtt5
Optional<ByteString> authData, // mqtt5
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -403,7 +403,7 @@ public final void channelInactive(ChannelHandlerContext ctx) {
resendTask.cancel(true);
}
if (noDelayLWT != null) {
addBgTask(pubWillMessage(noDelayLWT));
addBgTask(doPubLastWill(noDelayLWT));
}
cancelStallTask();
Sets.newHashSet(fgTasks).forEach(t -> t.cancel(true));
Expand Down Expand Up @@ -1513,27 +1513,6 @@ private CompletableFuture<Void> handleQoS2Pub(long reqId,
});
}

private CompletableFuture<Void> pubWillMessage(LWT willMessage) {
return authProvider.checkPermission(clientInfo(), buildPubAction(willMessage.getTopic(),
willMessage.getMessage()
.getPubQoS(),
willMessage.getMessage().getIsRetain()))
.thenCompose(checkResult -> {
assert ctx.executor().inEventLoop();
if (checkResult.hasGranted()) {
return doPubLastWill(willMessage);
} else {
sessionCtx.eventCollector.report(getLocal(PubActionDisallow.class)
.isLastWill(true)
.topic(willMessage.getTopic())
.qos(willMessage.getMessage().getPubQoS())
.isRetain(willMessage.getMessage().getIsRetain())
.clientInfo(clientInfo));
return CompletableFuture.completedFuture(null);
}
});
}

private void checkIdle() {
if (sessionCtx.nanoTime() - lastActiveAtNanos > idleTimeoutNanos) {
idleTimeoutTask.cancel(true);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
import static org.apache.bifromq.mqtt.handler.condition.ORCondition.or;
import static org.apache.bifromq.mqtt.handler.v3.MQTT3MessageUtils.toWillMessage;
import static org.apache.bifromq.mqtt.utils.AuthUtil.buildConnAction;
import static org.apache.bifromq.mqtt.utils.AuthUtil.buildPubAction;
import static org.apache.bifromq.plugin.eventcollector.ThreadLocalEventPool.getLocal;
import static org.apache.bifromq.type.MQTTClientInfoConstants.MQTT_CHANNEL_ID_KEY;
import static org.apache.bifromq.type.MQTTClientInfoConstants.MQTT_CLIENT_ADDRESS_KEY;
Expand Down Expand Up @@ -72,6 +73,7 @@
import org.apache.bifromq.plugin.clientbalancer.IClientBalancer;
import org.apache.bifromq.plugin.clientbalancer.Redirection;
import org.apache.bifromq.plugin.eventcollector.OutOfTenantResource;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.accessctrl.PubActionDisallow;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.channelclosed.AuthError;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.channelclosed.IdentifierRejected;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.channelclosed.MalformedClientIdentifier;
Expand Down Expand Up @@ -258,6 +260,49 @@ protected CompletableFuture<AuthResult> checkConnectPermission(MqttConnectMessag
});
}

@Override
protected CompletableFuture<AuthResult> checkWillPermission(MqttConnectMessage message,
LWT willMessage,
SuccessInfo successInfo) {
ClientInfo clientInfo = successInfo.clientInfo();
return authProvider.checkPermission(clientInfo, buildPubAction(willMessage.getTopic(),
willMessage.getMessage().getPubQoS(), willMessage.getMessage().getIsRetain()))
.handle((checkResult, e) -> {
if (e != null) {
return goAway(MqttMessageBuilders.connAck()
.returnCode(CONNECTION_REFUSED_SERVER_UNAVAILABLE)
.build(),
getLocal(AuthError.class)
.cause("Failed to check Will publish permission")
.peerAddress(ChannelAttrs.socketAddress(ctx.channel())));
}
switch (checkResult.getTypeCase()) {
case GRANTED -> {
return AuthResult.ok(successInfo);
}
case DENIED -> {
return goAway(MqttMessageBuilders.connAck()
.returnCode(CONNECTION_REFUSED_NOT_AUTHORIZED)
.build(),
getLocal(PubActionDisallow.class)
.isLastWill(true)
.topic(willMessage.getTopic())
.qos(willMessage.getMessage().getPubQoS())
.isRetain(willMessage.getMessage().getIsRetain())
.clientInfo(clientInfo));
}
default -> {
return goAway(MqttMessageBuilders.connAck()
.returnCode(CONNECTION_REFUSED_SERVER_UNAVAILABLE)
.build(),
getLocal(AuthError.class)
.cause("Failed to check Will publish permission")
.peerAddress(ChannelAttrs.socketAddress(ctx.channel())));
}
}
});
}

@Override
protected void handleMqttMessage(MqttMessage message) {
// never happen in MQTT3
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.toWillMessage;
import static org.apache.bifromq.mqtt.handler.v5.MQTT5MessageUtils.topicAliasMaximum;
import static org.apache.bifromq.mqtt.utils.AuthUtil.buildConnAction;
import static org.apache.bifromq.mqtt.utils.AuthUtil.buildPubAction;
import static org.apache.bifromq.mqtt.utils.MQTT5MessageSizer.MIN_CONTROL_PACKET_SIZE;
import static org.apache.bifromq.plugin.eventcollector.ThreadLocalEventPool.getLocal;
import static org.apache.bifromq.type.MQTTClientInfoConstants.MQTT_CHANNEL_ID_KEY;
Expand Down Expand Up @@ -103,6 +104,7 @@
import org.apache.bifromq.plugin.clientbalancer.IClientBalancer;
import org.apache.bifromq.plugin.clientbalancer.Redirection;
import org.apache.bifromq.plugin.eventcollector.OutOfTenantResource;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.accessctrl.PubActionDisallow;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.channelclosed.AuthError;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.channelclosed.EnhancedAuthAbortByClient;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.channelclosed.MalformedClientIdentifier;
Expand Down Expand Up @@ -379,6 +381,59 @@ protected CompletableFuture<AuthResult> checkConnectPermission(MqttConnectMessag
});
}

@Override
protected CompletableFuture<AuthResult> checkWillPermission(MqttConnectMessage message,
LWT willMessage,
SuccessInfo successInfo) {
ClientInfo clientInfo = successInfo.clientInfo();
return authProvider.checkPermission(clientInfo, buildPubAction(willMessage.getTopic(),
willMessage.getMessage().getPubQoS(), willMessage.getMessage().getIsRetain(),
toUserProperties(message.payload().willProperties())))
.handle((checkResult, e) -> {
if (e != null) {
return goAway(MqttMessageBuilders.connAck()
.properties(MQTT5MessageBuilders.connAckProperties()
.reasonString("Failed to check Will publish permission")
.build())
.returnCode(CONNECTION_REFUSED_UNSPECIFIED_ERROR)
.build(),
getLocal(AuthError.class)
.cause("Failed to check Will publish permission")
.peerAddress(ChannelAttrs.socketAddress(ctx.channel())));
}
switch (checkResult.getTypeCase()) {
case GRANTED -> {
return AuthResult.ok(successInfo);
}
case DENIED -> {
return goAway(MqttMessageBuilders.connAck()
.properties(MQTT5MessageBuilders.connAckProperties()
.reasonString("Will publish not authorized")
.build())
.returnCode(CONNECTION_REFUSED_NOT_AUTHORIZED_5)
.build(),
getLocal(PubActionDisallow.class)
.isLastWill(true)
.topic(willMessage.getTopic())
.qos(willMessage.getMessage().getPubQoS())
.isRetain(willMessage.getMessage().getIsRetain())
.clientInfo(clientInfo));
}
default -> {
return goAway(MqttMessageBuilders.connAck()
.properties(MQTT5MessageBuilders.connAckProperties()
.reasonString("Failed to check Will publish permission")
.build())
.returnCode(CONNECTION_REFUSED_UNSPECIFIED_ERROR)
.build(),
getLocal(AuthError.class)
.cause("Failed to check Will publish permission")
.peerAddress(ChannelAttrs.socketAddress(ctx.channel())));
}
}
});
}

private void extendedAuth(MQTT5ExtendedAuthData authData) {
this.isAuthing = true;
authProvider.extendedAuth(authData)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -276,6 +276,10 @@ protected void mockAuthPass(String... attrsKeyValues) {
.thenReturn(CompletableFuture.completedFuture(CheckResult.newBuilder()
.setGranted(Granted.getDefaultInstance())
.build()));
when(authProvider.checkPermission(any(ClientInfo.class), argThat(MQTTAction::hasPub)))
.thenReturn(CompletableFuture.completedFuture(CheckResult.newBuilder()
.setGranted(Granted.getDefaultInstance())
.build()));
}

protected void mockAuthReject(Reject.Code code, String reason) {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,8 +29,12 @@
import static org.apache.bifromq.plugin.eventcollector.EventType.MQTT_SESSION_START;
import static org.apache.bifromq.plugin.eventcollector.EventType.PING_REQ;
import static org.apache.bifromq.type.MQTTClientInfoConstants.MQTT_PROTOCOL_VER_KEY;
import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.argThat;
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 static org.testng.Assert.assertEquals;
import static org.testng.Assert.assertNotEquals;
import static org.testng.Assert.assertNull;
Expand All @@ -39,17 +43,24 @@
import io.netty.handler.codec.mqtt.MqttConnAckMessage;
import io.netty.handler.codec.mqtt.MqttConnectMessage;
import io.netty.handler.codec.mqtt.MqttMessage;
import java.util.concurrent.CompletableFuture;
import java.util.concurrent.TimeUnit;
import lombok.extern.slf4j.Slf4j;
import org.apache.bifromq.inbox.rpc.proto.AttachReply;
import org.apache.bifromq.inbox.rpc.proto.DetachReply;
import org.apache.bifromq.mqtt.utils.MQTTMessageUtils;
import org.apache.bifromq.plugin.authprovider.type.CheckResult;
import org.apache.bifromq.plugin.authprovider.type.Denied;
import org.apache.bifromq.plugin.authprovider.type.Error;
import org.apache.bifromq.plugin.authprovider.type.MQTTAction;
import org.apache.bifromq.plugin.authprovider.type.Reject;
import org.apache.bifromq.plugin.eventcollector.Event;
import org.apache.bifromq.plugin.eventcollector.EventType;
import org.apache.bifromq.plugin.eventcollector.mqttbroker.clientconnected.ClientConnected;
import org.apache.bifromq.type.ClientInfo;
import org.mockito.ArgumentCaptor;
import org.testng.Assert;
import org.testng.annotations.DataProvider;
import org.testng.annotations.Test;

@Slf4j
Expand Down Expand Up @@ -253,12 +264,58 @@ public void validWillTopic() {
mockInboxDetach(DetachReply.Code.OK);
MqttConnectMessage connectMessage = MQTTMessageUtils.qoSWillMqttConnectMessage(1, true);
channel.writeInbound(connectMessage);
channel.runPendingTasks();
MqttConnAckMessage ackMessage = channel.readOutbound();
// verifications
assertEquals(ackMessage.variableHeader().connectReturnCode(), CONNECTION_ACCEPTED);
verifyEvent(MQTT_SESSION_START, CLIENT_CONNECTED);
}

@Test
public void willPermissionDeniedBeforeSessionCalls() {
mockAuthPass();
when(authProvider.checkPermission(any(ClientInfo.class), argThat(MQTTAction::hasPub)))
.thenReturn(CompletableFuture.completedFuture(CheckResult.newBuilder()
.setDenied(Denied.getDefaultInstance())
.build()));

channel.writeInbound(MQTTMessageUtils.qoSWillMqttConnectMessage(1, true));
channel.advanceTimeBy(disconnectDelay, TimeUnit.MILLISECONDS);
channel.runScheduledPendingTasks();
channel.runPendingTasks();

MqttConnAckMessage connAck = channel.readOutbound();
assertEquals(connAck.variableHeader().connectReturnCode(), CONNECTION_REFUSED_NOT_AUTHORIZED);
verifyNoInteractions(inboxClient, sessionDictClient);
}

@DataProvider
public Object[][] willPermissionErrors() {
return new Object[][] {
{CompletableFuture.failedFuture(new RuntimeException("auth unavailable"))},
{CompletableFuture.completedFuture(CheckResult.newBuilder()
.setError(Error.getDefaultInstance())
.build())},
{CompletableFuture.completedFuture(CheckResult.getDefaultInstance())}
};
}

@Test(dataProvider = "willPermissionErrors")
public void willPermissionErrorBeforeSessionCalls(CompletableFuture<CheckResult> permissionResult) {
mockAuthPass();
when(authProvider.checkPermission(any(ClientInfo.class), argThat(MQTTAction::hasPub)))
.thenReturn(permissionResult);

channel.writeInbound(MQTTMessageUtils.qoSWillMqttConnectMessage(1, true));
channel.advanceTimeBy(disconnectDelay, TimeUnit.MILLISECONDS);
channel.runScheduledPendingTasks();
channel.runPendingTasks();

MqttConnAckMessage connAck = channel.readOutbound();
assertEquals(connAck.variableHeader().connectReturnCode(), CONNECTION_REFUSED_SERVER_UNAVAILABLE);
verifyNoInteractions(inboxClient, sessionDictClient);
}

@Test
public void pingAndPingResp() {
mockAuthPass();
Expand All @@ -267,6 +324,7 @@ public void pingAndPingResp() {
mockInboxDetach(DetachReply.Code.OK);
MqttConnectMessage connectMessage = MQTTMessageUtils.qoSWillMqttConnectMessage(1, true);
channel.writeInbound(connectMessage);
channel.runPendingTasks();
MqttConnAckMessage ackMessage = channel.readOutbound();
// verifications
assertEquals(ackMessage.variableHeader().connectReturnCode(), CONNECTION_ACCEPTED);
Expand Down
Loading
Loading