From 97ee45be24de93e38bd71c054d6904eb6a6efddf Mon Sep 17 00:00:00 2001 From: Sergey Matvienko Date: Tue, 15 Aug 2023 14:15:21 +0200 Subject: [PATCH 1/3] TbMathNode: refactored for easier testing. Semaphores - WEAK reference type. calculateResult method - removed unused args. --- .../rule/engine/math/TbMathNode.java | 21 ++++++++++++------- 1 file changed, 13 insertions(+), 8 deletions(-) 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..f48e723494 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 @@ -79,7 +79,7 @@ import java.util.stream.Collectors; ) public class TbMathNode implements TbNode { - private static final ConcurrentMap semaphores = new ConcurrentReferenceHashMap<>(); + private static final ConcurrentMap semaphores = new ConcurrentReferenceHashMap<>(16, ConcurrentReferenceHashMap.ReferenceType.WEAK); private final ThreadLocal customExpression = new ThreadLocal<>(); private TbMathNodeConfiguration config; @@ -116,12 +116,7 @@ public class TbMathNode implements TbNode { } 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()); + ListenableFuture resultMsgFuture = processMsgAsync(ctx, msg); DonAsynchron.withCallback(resultMsgFuture, resultMsg -> { try { ctx.tellSuccess(resultMsg); @@ -142,6 +137,16 @@ public class TbMathNode implements TbNode { } } + 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 +253,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); From 16fdfc518d2a4cee00784fdf1633ff31ae1da992 Mon Sep 17 00:00:00 2001 From: Sergey Matvienko Date: Tue, 15 Aug 2023 14:43:56 +0200 Subject: [PATCH 2/3] TbMathNode: test added for concurrent calls by the same originator utilizing the whole rule-dispatcher pool. 1 failed. non-blocking implementation wanted; Additional refactoring: JUnit5 and mock init --- .../rule/engine/math/TbMathNodeTest.java | 124 ++++++++++++++---- 1 file changed, 98 insertions(+), 26 deletions(-) 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..4c36d51a7d 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,34 @@ 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 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 +85,41 @@ 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() { + dbCallbackExecutor = new AbstractListeningExecutor() { @Override protected int getThreadPollSize() { - return 3; + return DB_CALLBACK_POOL_SIZE; } }; - dbExecutor.init(); - initMocks(); + dbCallbackExecutor.init(); + ruleEngineDispatcherExecutor = new AbstractListeningExecutor() { + @Override + protected int getThreadPollSize() { + return RULE_DISPATCHER_POOL_SIZE; + } + }; + 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 +516,56 @@ 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 = List.of( + TbMsg.newMsg("TEST", originatorSlow, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()), + TbMsg.newMsg("TEST", originatorSlow, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()) + ); + List fastMsgList = List.of( + TbMsg.newMsg("TEST", originatorFast, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()), + TbMsg.newMsg("TEST", originatorFast, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()) + ); + + log.debug("rule-dispatcher [{}], db-callback [{}], slowMsg [{}], fastMsg [{}]", RULE_DISPATCHER_POOL_SIZE, DB_CALLBACK_POOL_SIZE, slowMsgList.size(), fastMsgList.size()); + + willAnswer(invocation -> { + TbContext ctx = invocation.getArgument(0); + TbMsg msg = invocation.getArgument(1); + log.debug("awaiting on slowProcessingLatch [{}]", 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)); + + // 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()); + } + } From 44ea477b7bb185d7eca65761d7131bedb2c40aba Mon Sep 17 00:00:00 2001 From: Sergey Matvienko Date: Tue, 15 Aug 2023 22:40:01 +0200 Subject: [PATCH 3/3] TbMathNode: refactored to act in non-blocking style. All messages go through queue by originator with single semaphore and never wait on tryAcquire. Test refactored to provide more details on how slaw and fast messages being submitted and processed --- .../rule/engine/math/TbMathNode.java | 95 ++++++++++++++----- .../rule/engine/math/TbMathNodeTest.java | 67 ++++++++----- 2 files changed, 114 insertions(+), 48 deletions(-) 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 f48e723494..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<>(16, ConcurrentReferenceHashMap.ReferenceType.WEAK); + private static final ConcurrentMap> locks = new ConcurrentReferenceHashMap<>(16, ConcurrentReferenceHashMap.ReferenceType.WEAK); private final ThreadLocal customExpression = new ThreadLocal<>(); - private TbMathNodeConfiguration config; private boolean msgBodyToJsonConversionRequired; @@ -106,34 +109,58 @@ 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 { - ListenableFuture resultMsgFuture = processMsgAsync(ctx, msg); - 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 } } @@ -399,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 4c36d51a7d..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 @@ -53,6 +53,8 @@ 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; @@ -90,19 +92,9 @@ public class TbMathNodeTest { @BeforeEach public void before() { - dbCallbackExecutor = new AbstractListeningExecutor() { - @Override - protected int getThreadPollSize() { - return DB_CALLBACK_POOL_SIZE; - } - }; + dbCallbackExecutor = new DBCallbackExecutor(); dbCallbackExecutor.init(); - ruleEngineDispatcherExecutor = new AbstractListeningExecutor() { - @Override - protected int getThreadPollSize() { - return RULE_DISPATCHER_POOL_SIZE; - } - }; + ruleEngineDispatcherExecutor = new RuleDispatcherExecutor(); ruleEngineDispatcherExecutor.init(); lenient().when(ctx.getAttributesService()).thenReturn(attributesService); @@ -528,21 +520,20 @@ public class TbMathNodeTest { EntityId originatorFast = DeviceId.fromString("c45360ff-7906-4102-a2ae-3495a86168d0"); CountDownLatch slowProcessingLatch = new CountDownLatch(1); - List slowMsgList = List.of( - TbMsg.newMsg("TEST", originatorSlow, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()), - TbMsg.newMsg("TEST", originatorSlow, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()) - ); - List fastMsgList = List.of( - TbMsg.newMsg("TEST", originatorFast, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()), - TbMsg.newMsg("TEST", originatorFast, new TbMsgMetaData(), JacksonUtil.newObjectNode().put("a", 2).put("b", 2).toString()) - ); + 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 -> { - TbContext ctx = invocation.getArgument(0); TbMsg msg = invocation.getArgument(1); - log.debug("awaiting on slowProcessingLatch [{}]", msg); + 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) { @@ -551,6 +542,24 @@ public class TbMathNodeTest { 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 @@ -568,4 +577,18 @@ public class TbMathNodeTest { 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; + } + } + }