Browse Source

Merge pull request #9084 from smatvienko-tb/feature/math-node-sequential-non-blocking

Math node sequential non blocking
pull/9094/head
Andrew Shvayka 3 years ago
committed by GitHub
parent
commit
21a55630f0
No known key found for this signature in database GPG Key ID: 4AEE18F83AFDEB23
  1. 112
      rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/math/TbMathNode.java
  2. 155
      rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/math/TbMathNodeTest.java

112
rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/math/TbMathNode.java

@ -19,6 +19,8 @@ import com.fasterxml.jackson.databind.node.ObjectNode;
import com.google.common.util.concurrent.Futures; import com.google.common.util.concurrent.Futures;
import com.google.common.util.concurrent.ListenableFuture; import com.google.common.util.concurrent.ListenableFuture;
import com.google.common.util.concurrent.MoreExecutors; import com.google.common.util.concurrent.MoreExecutors;
import lombok.Data;
import lombok.RequiredArgsConstructor;
import lombok.extern.slf4j.Slf4j; import lombok.extern.slf4j.Slf4j;
import net.objecthunter.exp4j.Expression; import net.objecthunter.exp4j.Expression;
import net.objecthunter.exp4j.ExpressionBuilder; import net.objecthunter.exp4j.ExpressionBuilder;
@ -44,6 +46,8 @@ import java.math.BigDecimal;
import java.math.RoundingMode; import java.math.RoundingMode;
import java.util.List; import java.util.List;
import java.util.Optional; import java.util.Optional;
import java.util.Queue;
import java.util.concurrent.ConcurrentLinkedQueue;
import java.util.concurrent.ConcurrentMap; import java.util.concurrent.ConcurrentMap;
import java.util.concurrent.Semaphore; import java.util.concurrent.Semaphore;
import java.util.concurrent.TimeUnit; import java.util.concurrent.TimeUnit;
@ -79,9 +83,8 @@ import java.util.stream.Collectors;
) )
public class TbMathNode implements TbNode { public class TbMathNode implements TbNode {
private static final ConcurrentMap<EntityId, Semaphore> semaphores = new ConcurrentReferenceHashMap<>(); private static final ConcurrentMap<EntityId, SemaphoreWithQueue<TbMsgTbContext>> locks = new ConcurrentReferenceHashMap<>(16, ConcurrentReferenceHashMap.ReferenceType.WEAK);
private final ThreadLocal<Expression> customExpression = new ThreadLocal<>(); private final ThreadLocal<Expression> customExpression = new ThreadLocal<>();
private TbMathNodeConfiguration config; private TbMathNodeConfiguration config;
private boolean msgBodyToJsonConversionRequired; private boolean msgBodyToJsonConversionRequired;
@ -106,42 +109,71 @@ public class TbMathNode implements TbNode {
@Override @Override
public void onMsg(TbContext ctx, TbMsg msg) { public void onMsg(TbContext ctx, TbMsg msg) {
var originator = msg.getOriginator(); var semaphoreWithQueue = locks.computeIfAbsent(msg.getOriginator(), SemaphoreWithQueue::new);
var originatorSemaphore = semaphores.computeIfAbsent(originator, tmp -> new Semaphore(1, true)); semaphoreWithQueue.getQueue().add(new TbMsgTbContext(msg, ctx));
boolean acquired = tryAcquire(originator, originatorSemaphore);
if (!acquired) { tryProcessQueue(semaphoreWithQueue);
ctx.tellFailure(msg, new RuntimeException("Failed to process message for originator synchronously")); }
return;
}
try { void tryProcessQueue(SemaphoreWithQueue<TbMsgTbContext> lockAndQueue) {
var arguments = config.getArguments(); final Semaphore semaphore = lockAndQueue.getSemaphore();
Optional<ObjectNode> msgBodyOpt = convertMsgBodyIfRequired(msg); final Queue<TbMsgTbContext> queue = lockAndQueue.getQueue();
var argumentValues = Futures.allAsList(arguments.stream() while (!queue.isEmpty()) {
.map(arg -> resolveArguments(ctx, msg, msgBodyOpt, arg)).collect(Collectors.toList())); // The semaphore have to be acquired before EACH poll and released before NEXT poll.
ListenableFuture<TbMsg> resultMsgFuture = Futures.transformAsync(argumentValues, args -> // Otherwise, some message will remain unprocessed in queue
updateMsgAndDb(ctx, msg, msgBodyOpt, calculateResult(ctx, msg, args)), ctx.getDbCallbackExecutor()); if (!semaphore.tryAcquire()) {
DonAsynchron.withCallback(resultMsgFuture, resultMsg -> { return;
try { }
ctx.tellSuccess(resultMsg); TbMsgTbContext tbMsgTbContext = null;
} finally { try {
originatorSemaphore.release(); tbMsgTbContext = queue.poll();
if (tbMsgTbContext == null) {
semaphore.release();
continue;
} }
}, t -> { final TbMsg msg = tbMsgTbContext.getMsg();
try { if (!msg.getCallback().isMsgValid()) {
ctx.tellFailure(msg, t); log.trace("[{}] Skipping non-valid message [{}]", lockAndQueue.getEntityId(), msg);
} finally { semaphore.release();
originatorSemaphore.release(); continue;
} }
}, ctx.getDbCallbackExecutor()); //DO PROCESSING
} catch (Throwable e) { final TbContext ctx = tbMsgTbContext.getCtx();
originatorSemaphore.release(); final ListenableFuture<TbMsg> resultMsgFuture = processMsgAsync(ctx, msg);
log.warn("[{}] Failed to process message: {}", originator, msg, e); DonAsynchron.withCallback(resultMsgFuture, resultMsg -> {
throw e; try {
ctx.tellSuccess(resultMsg);
} finally {
lockAndQueue.getSemaphore().release();
tryProcessQueue(lockAndQueue);
}
}, t -> {
try {
ctx.tellFailure(msg, t);
} finally {
lockAndQueue.getSemaphore().release();
tryProcessQueue(lockAndQueue);
}
}, ctx.getDbCallbackExecutor());
} catch (Throwable e) {
semaphore.release();
log.warn("[{}] Failed to process message: {}", lockAndQueue.getEntityId(), tbMsgTbContext == null ? null : tbMsgTbContext.getMsg(), e);
throw e;
}
break; //submitted async exact one task. next poll will try on callback
} }
} }
ListenableFuture<TbMsg> processMsgAsync(TbContext ctx, TbMsg msg) {
var arguments = config.getArguments();
Optional<ObjectNode> msgBodyOpt = convertMsgBodyIfRequired(msg);
var argumentValues = Futures.allAsList(arguments.stream()
.map(arg -> resolveArguments(ctx, msg, msgBodyOpt, arg)).collect(Collectors.toList()));
ListenableFuture<TbMsg> resultMsgFuture = Futures.transformAsync(argumentValues, args ->
updateMsgAndDb(ctx, msg, msgBodyOpt, calculateResult(args)), ctx.getDbCallbackExecutor());
return resultMsgFuture;
}
private boolean tryAcquire(EntityId originator, Semaphore originatorSemaphore) { private boolean tryAcquire(EntityId originator, Semaphore originatorSemaphore) {
boolean acquired; boolean acquired;
try { try {
@ -248,7 +280,7 @@ public class TbMathNode implements TbNode {
return TbMsg.transformMsg(msg, md); return TbMsg.transformMsg(msg, md);
} }
private double calculateResult(TbContext ctx, TbMsg msg, List<TbMathArgumentValue> args) { private double calculateResult(List<TbMathArgumentValue> args) {
switch (config.getOperation()) { switch (config.getOperation()) {
case ADD: case ADD:
return apply(args.get(0), args.get(1), Double::sum); return apply(args.get(0), args.get(1), Double::sum);
@ -394,4 +426,20 @@ public class TbMathNode implements TbNode {
@Override @Override
public void destroy() { public void destroy() {
} }
@Data
@RequiredArgsConstructor
static public class SemaphoreWithQueue<T> {
final EntityId entityId;
final Semaphore semaphore = new Semaphore(1);
final Queue<T> queue = new ConcurrentLinkedQueue<>();
}
@Data
@RequiredArgsConstructor
static public class TbMsgTbContext {
final TbMsg msg;
final TbContext ctx;
}
} }

155
rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/math/TbMathNodeTest.java

@ -15,17 +15,18 @@
*/ */
package org.thingsboard.rule.engine.math; package org.thingsboard.rule.engine.math;
import com.datastax.oss.driver.api.core.uuid.Uuids;
import com.google.common.util.concurrent.Futures; import com.google.common.util.concurrent.Futures;
import org.junit.After; import lombok.extern.slf4j.Slf4j;
import org.junit.Assert; import org.junit.Assert;
import org.junit.Before; import org.junit.jupiter.api.AfterEach;
import org.junit.Test; import org.junit.jupiter.api.BeforeEach;
import org.junit.runner.RunWith; import org.junit.jupiter.api.Test;
import org.junit.jupiter.api.extension.ExtendWith;
import org.mockito.ArgumentCaptor; import org.mockito.ArgumentCaptor;
import org.mockito.Mock; import org.mockito.Mock;
import org.mockito.Mockito; import org.mockito.Mockito;
import org.mockito.junit.MockitoJUnitRunner; import org.mockito.junit.jupiter.MockitoExtension;
import org.mockito.verification.Timeout;
import org.thingsboard.common.util.AbstractListeningExecutor; import org.thingsboard.common.util.AbstractListeningExecutor;
import org.thingsboard.common.util.JacksonUtil; import org.thingsboard.common.util.JacksonUtil;
import org.thingsboard.rule.engine.api.RuleEngineTelemetryService; import org.thingsboard.rule.engine.api.RuleEngineTelemetryService;
@ -47,21 +48,36 @@ import org.thingsboard.server.dao.attributes.AttributesService;
import org.thingsboard.server.dao.timeseries.TimeseriesService; import org.thingsboard.server.dao.timeseries.TimeseriesService;
import java.util.Arrays; import java.util.Arrays;
import java.util.List;
import java.util.Optional; import java.util.Optional;
import java.util.UUID;
import java.util.concurrent.CountDownLatch;
import java.util.concurrent.TimeUnit;
import java.util.stream.Collectors;
import java.util.stream.IntStream;
import static org.assertj.core.api.Assertions.assertThat;
import static org.junit.jupiter.api.Assertions.assertThrows; import static org.junit.jupiter.api.Assertions.assertThrows;
import static org.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.any;
import static org.mockito.ArgumentMatchers.anyDouble; import static org.mockito.ArgumentMatchers.anyDouble;
import static org.mockito.ArgumentMatchers.anyString; import static org.mockito.ArgumentMatchers.anyString;
import static org.mockito.ArgumentMatchers.argThat;
import static org.mockito.ArgumentMatchers.eq;
import static org.mockito.BDDMockito.willAnswer;
import static org.mockito.Mockito.lenient; import static org.mockito.Mockito.lenient;
import static org.mockito.Mockito.never;
import static org.mockito.Mockito.spy;
import static org.mockito.Mockito.times; import static org.mockito.Mockito.times;
import static org.mockito.Mockito.verify; import static org.mockito.Mockito.verify;
@RunWith(MockitoJUnitRunner.class) @ExtendWith(MockitoExtension.class)
@Slf4j
public class TbMathNodeTest { public class TbMathNodeTest {
private EntityId originator = new DeviceId(Uuids.timeBased()); static final int RULE_DISPATCHER_POOL_SIZE = 2;
private TenantId tenantId = TenantId.fromUUID(Uuids.timeBased()); static final int DB_CALLBACK_POOL_SIZE = 3;
private final EntityId originator = DeviceId.fromString("ccd71696-0586-422d-940e-755a41ec3b0d");
private final TenantId tenantId = TenantId.fromUUID(UUID.fromString("e7f46b23-0c7d-42f5-9b06-fc35ab17af8a"));
@Mock @Mock
private TbContext ctx; private TbContext ctx;
@ -71,35 +87,31 @@ public class TbMathNodeTest {
private TimeseriesService tsService; private TimeseriesService tsService;
@Mock @Mock
private RuleEngineTelemetryService telemetryService; private RuleEngineTelemetryService telemetryService;
private AbstractListeningExecutor dbExecutor; private AbstractListeningExecutor dbCallbackExecutor;
private AbstractListeningExecutor ruleEngineDispatcherExecutor;
@Before @BeforeEach
public void before() { public void before() {
dbExecutor = new AbstractListeningExecutor() { dbCallbackExecutor = new DBCallbackExecutor();
@Override dbCallbackExecutor.init();
protected int getThreadPollSize() { ruleEngineDispatcherExecutor = new RuleDispatcherExecutor();
return 3; ruleEngineDispatcherExecutor.init();
}
}; lenient().when(ctx.getAttributesService()).thenReturn(attributesService);
dbExecutor.init(); lenient().when(ctx.getTelemetryService()).thenReturn(telemetryService);
initMocks(); lenient().when(ctx.getTimeseriesService()).thenReturn(tsService);
lenient().when(ctx.getTenantId()).thenReturn(tenantId);
lenient().when(ctx.getDbCallbackExecutor()).thenReturn(dbCallbackExecutor);
} }
@After @AfterEach
public void after() { public void after() {
dbExecutor.destroy(); ruleEngineDispatcherExecutor.executor().shutdownNow();
dbCallbackExecutor.executor().shutdownNow();
} }
private void initMocks() { private void initMocks() {
Mockito.reset(ctx); Mockito.clearInvocations(ctx, attributesService, tsService, telemetryService);
Mockito.reset(attributesService);
Mockito.reset(tsService);
Mockito.reset(telemetryService);
lenient().when(ctx.getAttributesService()).thenReturn(attributesService);
lenient().when(ctx.getTelemetryService()).thenReturn(telemetryService);
lenient().when(ctx.getTimeseriesService()).thenReturn(tsService);
lenient().when(ctx.getTenantId()).thenReturn(tenantId);
lenient().when(ctx.getDbCallbackExecutor()).thenReturn(dbExecutor);
} }
private TbMathNode initNode(TbRuleNodeMathFunctionType operation, TbMathResult result, TbMathArgument... arguments) { private TbMathNode initNode(TbRuleNodeMathFunctionType operation, TbMathResult result, TbMathArgument... arguments) {
@ -496,4 +508,87 @@ public class TbMathNodeTest {
}); });
Assert.assertNotNull(thrown.getMessage()); Assert.assertNotNull(thrown.getMessage());
} }
@Test
public void testExp4j_concurrent() {
TbMathNode node = spy(initNodeWithCustomFunction("2a+3b",
new TbMathResult(TbMathArgumentType.MESSAGE_BODY, "result", 2, false, false, null),
new TbMathArgument(TbMathArgumentType.MESSAGE_BODY, "a"),
new TbMathArgument(TbMathArgumentType.MESSAGE_BODY, "b")
));
EntityId originatorSlow = DeviceId.fromString("7f01170d-6bba-419c-b95c-2b4c3ba32f30");
EntityId originatorFast = DeviceId.fromString("c45360ff-7906-4102-a2ae-3495a86168d0");
CountDownLatch slowProcessingLatch = new CountDownLatch(1);
List<TbMsg> slowMsgList = IntStream.range(0, 5)
.mapToObj(x -> TbMsg.newMsg("TEST", originatorSlow, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()))
.collect(Collectors.toList());
List<TbMsg> fastMsgList = IntStream.range(0, 2)
.mapToObj(x -> TbMsg.newMsg("TEST", originatorFast, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()))
.collect(Collectors.toList());
assertThat(slowMsgList.size()).as("slow msgs >= rule-dispatcher pool size").isGreaterThanOrEqualTo(RULE_DISPATCHER_POOL_SIZE);
log.debug("rule-dispatcher [{}], db-callback [{}], slowMsg [{}], fastMsg [{}]", RULE_DISPATCHER_POOL_SIZE, DB_CALLBACK_POOL_SIZE, slowMsgList.size(), fastMsgList.size());
willAnswer(invocation -> {
TbMsg msg = invocation.getArgument(1);
log.debug("\uD83D\uDC0C processMsgAsync slow originator [{}][{}]", msg.getOriginator(), msg);
try {
assertThat(slowProcessingLatch.await(30, TimeUnit.SECONDS)).as("await on slowProcessingLatch").isTrue();
} catch (InterruptedException e) {
throw new RuntimeException(e);
}
return invocation.callRealMethod();
}).given(node).processMsgAsync(eq(ctx), argThat(slowMsgList::contains));
willAnswer(invocation -> {
TbMsg msg = invocation.getArgument(1);
log.debug("\u26A1\uFE0F processMsgAsync FAST originator [{}][{}]", msg.getOriginator(), msg);
return invocation.callRealMethod();
}).given(node).processMsgAsync(eq(ctx), argThat(fastMsgList::contains));
willAnswer(invocation -> {
TbMsg msg = invocation.getArgument(1);
log.debug("submit slow originator onMsg [{}][{}]", msg.getOriginator(), msg);
return invocation.callRealMethod();
}).given(node).onMsg(eq(ctx), argThat(slowMsgList::contains));
willAnswer(invocation -> {
TbMsg msg = invocation.getArgument(1);
log.debug("submit FAST originator onMsg [{}][{}]", msg.getOriginator(), msg);
return invocation.callRealMethod();
}).given(node).onMsg(eq(ctx), argThat(fastMsgList::contains));
// submit slow msg may block all rule engine dispatcher threads
slowMsgList.forEach(msg -> ruleEngineDispatcherExecutor.executeAsync(() -> node.onMsg(ctx, msg)));
// wait until dispatcher threads started with all slowMsg
verify(node, new Timeout(TimeUnit.SECONDS.toMillis(5), times(slowMsgList.size()))).onMsg(eq(ctx), argThat(slowMsgList::contains));
// submit fast have to return immediately
fastMsgList.forEach(msg -> ruleEngineDispatcherExecutor.executeAsync(() -> node.onMsg(ctx, msg)));
// wait until all fast messages processed
verify(ctx, new Timeout(TimeUnit.SECONDS.toMillis(5), times(fastMsgList.size()))).tellSuccess(any());
slowProcessingLatch.countDown();
verify(ctx, new Timeout(TimeUnit.SECONDS.toMillis(5), times(fastMsgList.size() + slowMsgList.size()))).tellSuccess(any());
verify(ctx, never()).tellFailure(any(), any());
}
static class RuleDispatcherExecutor extends AbstractListeningExecutor {
@Override
protected int getThreadPollSize() {
return RULE_DISPATCHER_POOL_SIZE;
}
}
static class DBCallbackExecutor extends AbstractListeningExecutor {
@Override
protected int getThreadPollSize() {
return DB_CALLBACK_POOL_SIZE;
}
}
} }

Loading…
Cancel
Save