diff --git a/application/src/main/java/org/thingsboard/server/service/queue/ProcessingAttemptContext.java b/application/src/main/java/org/thingsboard/server/service/queue/ProcessingAttemptContext.java index 2073ff8c21..aefb1697cd 100644 --- a/application/src/main/java/org/thingsboard/server/service/queue/ProcessingAttemptContext.java +++ b/application/src/main/java/org/thingsboard/server/service/queue/ProcessingAttemptContext.java @@ -27,11 +27,13 @@ import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentMap; import java.util.concurrent.CountDownLatch; import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicInteger; public class ProcessingAttemptContext { private final TbRuleEngineSubmitStrategy submitStrategy; + private final AtomicInteger pendingCount; private final CountDownLatch processingTimeoutLatch = new CountDownLatch(1); @Getter private final ConcurrentMap> pendingMap; @@ -45,6 +47,7 @@ public class ProcessingAttemptContext { public ProcessingAttemptContext(TbRuleEngineSubmitStrategy submitStrategy) { this.submitStrategy = submitStrategy; this.pendingMap = submitStrategy.getPendingMap(); + this.pendingCount = new AtomicInteger(pendingMap.size()); } public boolean await(long packProcessingTimeout, TimeUnit milliseconds) throws InterruptedException { @@ -54,16 +57,12 @@ public class ProcessingAttemptContext { public void onSuccess(UUID id) { TbProtoQueueMsg msg; boolean empty = false; - synchronized (pendingMap) { - msg = pendingMap.remove(id); - if (msg != null) { - empty = pendingMap.isEmpty(); - } - } + msg = pendingMap.remove(id); if (msg != null) { + empty = pendingCount.decrementAndGet() == 0; successMap.put(id, msg); + submitStrategy.onSuccess(id); } - submitStrategy.onSuccess(id); if (empty) { processingTimeoutLatch.countDown(); } @@ -72,13 +71,9 @@ public class ProcessingAttemptContext { public void onFailure(TenantId tenantId, UUID id, RuleEngineException e) { TbProtoQueueMsg msg; boolean empty = false; - synchronized (pendingMap) { - msg = pendingMap.remove(id); - if (msg != null) { - empty = pendingMap.isEmpty(); - } - } + msg = pendingMap.remove(id); if (msg != null) { + empty = pendingCount.decrementAndGet() == 0; failedMap.put(id, msg); exceptionsMap.putIfAbsent(tenantId, e); } diff --git a/application/src/test/java/org/thingsboard/server/service/queue/ProcessingAttemptContextTest.java b/application/src/test/java/org/thingsboard/server/service/queue/ProcessingAttemptContextTest.java new file mode 100644 index 0000000000..6c3914abc8 --- /dev/null +++ b/application/src/test/java/org/thingsboard/server/service/queue/ProcessingAttemptContextTest.java @@ -0,0 +1,51 @@ +package org.thingsboard.server.service.queue; + +import lombok.extern.slf4j.Slf4j; +import org.junit.Assert; +import org.junit.Test; +import org.junit.runner.RunWith; +import org.mockito.Mockito; +import org.mockito.runners.MockitoJUnitRunner; +import org.thingsboard.server.gen.transport.TransportProtos; +import org.thingsboard.server.queue.common.TbProtoQueueMsg; +import org.thingsboard.server.service.queue.processing.TbRuleEngineSubmitStrategy; + +import java.util.UUID; +import java.util.concurrent.ConcurrentHashMap; +import java.util.concurrent.ConcurrentMap; +import java.util.concurrent.ExecutorService; +import java.util.concurrent.Executors; +import java.util.concurrent.TimeUnit; + +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.when; + +@Slf4j +@RunWith(MockitoJUnitRunner.class) +public class ProcessingAttemptContextTest { + + @Test + public void testHighConcurrencyCase() throws InterruptedException { + TbRuleEngineSubmitStrategy strategyMock = mock(TbRuleEngineSubmitStrategy.class); + int msgCount = 1000; + int parallelCount = 5; + ExecutorService executorService = Executors.newFixedThreadPool(parallelCount); + try { + ConcurrentMap> messages = new ConcurrentHashMap<>(); + for (int i = 0; i < msgCount; i++) { + messages.put(UUID.randomUUID(), new TbProtoQueueMsg<>(UUID.randomUUID(), null)); + } + when(strategyMock.getPendingMap()).thenReturn(messages); + ProcessingAttemptContext context = new ProcessingAttemptContext(strategyMock); + for (UUID uuid : messages.keySet()) { + for (int i = 0; i < parallelCount; i++) { + executorService.submit(() -> context.onSuccess(uuid)); + } + } + Assert.assertTrue(context.await(10, TimeUnit.SECONDS)); + Mockito.verify(strategyMock, Mockito.times(msgCount)).onSuccess(Mockito.any(UUID.class)); + } finally { + executorService.shutdownNow(); + } + } +}