diff --git a/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/math/TbMathNode.java b/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/math/TbMathNode.java index eff3917c14..0c34f4fdcc 100644 --- a/rule-engine/rule-engine-components/src/main/java/org/thingsboard/rule/engine/math/TbMathNode.java +++ b/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.ListenableFuture; import com.google.common.util.concurrent.MoreExecutors; +import lombok.Data; +import lombok.RequiredArgsConstructor; import lombok.extern.slf4j.Slf4j; import net.objecthunter.exp4j.Expression; import net.objecthunter.exp4j.ExpressionBuilder; @@ -44,6 +46,8 @@ import java.math.BigDecimal; import java.math.RoundingMode; import java.util.List; import java.util.Optional; +import java.util.Queue; +import java.util.concurrent.ConcurrentLinkedQueue; import java.util.concurrent.ConcurrentMap; import java.util.concurrent.Semaphore; import java.util.concurrent.TimeUnit; @@ -79,9 +83,8 @@ import java.util.stream.Collectors; ) public class TbMathNode implements TbNode { - private static final ConcurrentMap semaphores = new ConcurrentReferenceHashMap<>(); + private static final ConcurrentMap> locks = new ConcurrentReferenceHashMap<>(16, ConcurrentReferenceHashMap.ReferenceType.WEAK); private final ThreadLocal customExpression = new ThreadLocal<>(); - private TbMathNodeConfiguration config; private boolean msgBodyToJsonConversionRequired; @@ -106,42 +109,71 @@ public class TbMathNode implements TbNode { @Override public void onMsg(TbContext ctx, TbMsg msg) { - var originator = msg.getOriginator(); - var originatorSemaphore = semaphores.computeIfAbsent(originator, tmp -> new Semaphore(1, true)); - boolean acquired = tryAcquire(originator, originatorSemaphore); + var semaphoreWithQueue = locks.computeIfAbsent(msg.getOriginator(), SemaphoreWithQueue::new); + semaphoreWithQueue.getQueue().add(new TbMsgTbContext(msg, ctx)); - if (!acquired) { - ctx.tellFailure(msg, new RuntimeException("Failed to process message for originator synchronously")); - return; - } + tryProcessQueue(semaphoreWithQueue); + } - try { - var arguments = config.getArguments(); - Optional msgBodyOpt = convertMsgBodyIfRequired(msg); - var argumentValues = Futures.allAsList(arguments.stream() - .map(arg -> resolveArguments(ctx, msg, msgBodyOpt, arg)).collect(Collectors.toList())); - ListenableFuture resultMsgFuture = Futures.transformAsync(argumentValues, args -> - updateMsgAndDb(ctx, msg, msgBodyOpt, calculateResult(ctx, msg, args)), ctx.getDbCallbackExecutor()); - DonAsynchron.withCallback(resultMsgFuture, resultMsg -> { - try { - ctx.tellSuccess(resultMsg); - } finally { - originatorSemaphore.release(); + void tryProcessQueue(SemaphoreWithQueue lockAndQueue) { + final Semaphore semaphore = lockAndQueue.getSemaphore(); + final Queue queue = lockAndQueue.getQueue(); + while (!queue.isEmpty()) { + // The semaphore have to be acquired before EACH poll and released before NEXT poll. + // Otherwise, some message will remain unprocessed in queue + if (!semaphore.tryAcquire()) { + return; + } + TbMsgTbContext tbMsgTbContext = null; + try { + tbMsgTbContext = queue.poll(); + if (tbMsgTbContext == null) { + semaphore.release(); + continue; } - }, t -> { - try { - ctx.tellFailure(msg, t); - } finally { - originatorSemaphore.release(); + final TbMsg msg = tbMsgTbContext.getMsg(); + if (!msg.getCallback().isMsgValid()) { + log.trace("[{}] Skipping non-valid message [{}]", lockAndQueue.getEntityId(), msg); + semaphore.release(); + continue; } - }, ctx.getDbCallbackExecutor()); - } catch (Throwable e) { - originatorSemaphore.release(); - log.warn("[{}] Failed to process message: {}", originator, msg, e); - throw e; + //DO PROCESSING + final TbContext ctx = tbMsgTbContext.getCtx(); + final ListenableFuture resultMsgFuture = processMsgAsync(ctx, msg); + DonAsynchron.withCallback(resultMsgFuture, resultMsg -> { + 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 processMsgAsync(TbContext ctx, TbMsg msg) { + var arguments = config.getArguments(); + Optional msgBodyOpt = convertMsgBodyIfRequired(msg); + var argumentValues = Futures.allAsList(arguments.stream() + .map(arg -> resolveArguments(ctx, msg, msgBodyOpt, arg)).collect(Collectors.toList())); + ListenableFuture resultMsgFuture = Futures.transformAsync(argumentValues, args -> + updateMsgAndDb(ctx, msg, msgBodyOpt, calculateResult(args)), ctx.getDbCallbackExecutor()); + return resultMsgFuture; + } + private boolean tryAcquire(EntityId originator, Semaphore originatorSemaphore) { boolean acquired; try { @@ -248,7 +280,7 @@ public class TbMathNode implements TbNode { return TbMsg.transformMsg(msg, md); } - private double calculateResult(TbContext ctx, TbMsg msg, List args) { + private double calculateResult(List args) { switch (config.getOperation()) { case ADD: return apply(args.get(0), args.get(1), Double::sum); @@ -394,4 +426,20 @@ public class TbMathNode implements TbNode { @Override public void destroy() { } + + @Data + @RequiredArgsConstructor + static public class SemaphoreWithQueue { + final EntityId entityId; + final Semaphore semaphore = new Semaphore(1); + final Queue queue = new ConcurrentLinkedQueue<>(); + } + + @Data + @RequiredArgsConstructor + static public class TbMsgTbContext { + final TbMsg msg; + final TbContext ctx; + } + } diff --git a/rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/math/TbMathNodeTest.java b/rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/math/TbMathNodeTest.java index 69d1b46dbe..50424a1521 100644 --- a/rule-engine/rule-engine-components/src/test/java/org/thingsboard/rule/engine/math/TbMathNodeTest.java +++ b/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; -import com.datastax.oss.driver.api.core.uuid.Uuids; import com.google.common.util.concurrent.Futures; -import org.junit.After; +import lombok.extern.slf4j.Slf4j; import org.junit.Assert; -import org.junit.Before; -import org.junit.Test; -import org.junit.runner.RunWith; +import org.junit.jupiter.api.AfterEach; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.api.extension.ExtendWith; import org.mockito.ArgumentCaptor; import org.mockito.Mock; 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.JacksonUtil; 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 java.util.Arrays; +import java.util.List; 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.mockito.ArgumentMatchers.any; import static org.mockito.ArgumentMatchers.anyDouble; 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.never; +import static org.mockito.Mockito.spy; import static org.mockito.Mockito.times; import static org.mockito.Mockito.verify; -@RunWith(MockitoJUnitRunner.class) +@ExtendWith(MockitoExtension.class) +@Slf4j public class TbMathNodeTest { - private EntityId originator = new DeviceId(Uuids.timeBased()); - private TenantId tenantId = TenantId.fromUUID(Uuids.timeBased()); + static final int RULE_DISPATCHER_POOL_SIZE = 2; + 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 private TbContext ctx; @@ -71,35 +87,31 @@ public class TbMathNodeTest { private TimeseriesService tsService; @Mock private RuleEngineTelemetryService telemetryService; - private AbstractListeningExecutor dbExecutor; + private AbstractListeningExecutor dbCallbackExecutor; + private AbstractListeningExecutor ruleEngineDispatcherExecutor; - @Before + @BeforeEach public void before() { - dbExecutor = new AbstractListeningExecutor() { - @Override - protected int getThreadPollSize() { - return 3; - } - }; - dbExecutor.init(); - initMocks(); + dbCallbackExecutor = new DBCallbackExecutor(); + dbCallbackExecutor.init(); + ruleEngineDispatcherExecutor = new RuleDispatcherExecutor(); + ruleEngineDispatcherExecutor.init(); + + 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(dbCallbackExecutor); } - @After + @AfterEach public void after() { - dbExecutor.destroy(); + ruleEngineDispatcherExecutor.executor().shutdownNow(); + dbCallbackExecutor.executor().shutdownNow(); } private void initMocks() { - Mockito.reset(ctx); - 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); + Mockito.clearInvocations(ctx, attributesService, tsService, telemetryService); } private TbMathNode initNode(TbRuleNodeMathFunctionType operation, TbMathResult result, TbMathArgument... arguments) { @@ -496,4 +508,87 @@ public class TbMathNodeTest { }); 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 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 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; + } + } + }