This is an automated email from the ASF dual-hosted git repository.

joewitt pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/nifi.git


The following commit(s) were added to refs/heads/main by this push:
     new 7298695681e NIFI-16287 - Ensure orderly shutdown of active and idle 
JMS workers This closes #11619
7298695681e is described below

commit 7298695681efcecf4d99d9d028b7d1fb98f471f7
Author: Pierre Villard <[email protected]>
AuthorDate: Thu Sep 3 15:05:52 2026 +0200

    NIFI-16287 - Ensure orderly shutdown of active and idle JMS workers
    This closes #11619
---
 .../nifi/jms/processors/AbstractJMSProcessor.java  |  92 +++---
 .../org/apache/nifi/jms/processors/JMSWorker.java  |   6 +-
 .../nifi/jms/processors/JmsWorkerLifecycle.java    | 188 ++++++++++++
 .../jms/processors/AbstractJMSProcessorTest.java   | 326 +++++++++++++++++++++
 .../apache/nifi/jms/processors/ConsumeJMSIT.java   | 247 ++++++++++++++++
 .../apache/nifi/jms/processors/JMSWorkerTest.java  | 173 +++++++++++
 .../apache/nifi/jms/processors/PublishJMSIT.java   | 259 ++++++++++++++++
 7 files changed, 1247 insertions(+), 44 deletions(-)

diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/AbstractJMSProcessor.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/AbstractJMSProcessor.java
index 2d8559e9ad7..6e343a5c452 100644
--- 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/AbstractJMSProcessor.java
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/AbstractJMSProcessor.java
@@ -17,7 +17,6 @@
 package org.apache.nifi.jms.processors;
 
 import jakarta.jms.ConnectionFactory;
-import jakarta.jms.Message;
 import org.apache.nifi.annotation.lifecycle.OnScheduled;
 import org.apache.nifi.annotation.lifecycle.OnStopped;
 import org.apache.nifi.annotation.lifecycle.OnUnscheduled;
@@ -31,7 +30,6 @@ import org.apache.nifi.expression.ExpressionLanguageScope;
 import org.apache.nifi.jms.cf.IJMSConnectionFactoryProvider;
 import org.apache.nifi.jms.cf.JMSConnectionFactoryHandler;
 import org.apache.nifi.jms.cf.JMSConnectionFactoryProperties;
-import org.apache.nifi.jms.cf.JMSConnectionFactoryProvider;
 import org.apache.nifi.jms.cf.JMSConnectionFactoryProviderDefinition;
 import org.apache.nifi.jms.cf.JndiJmsConnectionFactoryHandler;
 import org.apache.nifi.jms.cf.JndiJmsConnectionFactoryProperties;
@@ -55,10 +53,8 @@ import java.util.ArrayList;
 import java.util.Collection;
 import java.util.Collections;
 import java.util.List;
-import java.util.concurrent.BlockingQueue;
-import java.util.concurrent.LinkedBlockingQueue;
-import java.util.concurrent.atomic.AtomicBoolean;
 import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicReference;
 import java.util.stream.Collectors;
 
 /**
@@ -170,10 +166,8 @@ public abstract class AbstractJMSProcessor<T extends 
JMSWorker> extends Abstract
             .required(true)
             .build();
 
-    private volatile IJMSConnectionFactoryProvider connectionFactoryProvider;
-    private volatile BlockingQueue<T> workerPool;
+    private final AtomicReference<JmsWorkerLifecycle<T>> workerLifecycle = new 
AtomicReference<>();
     private volatile boolean runOnPrimary;
-    private final AtomicBoolean shutdownWorkers = new AtomicBoolean(false);
     private final AtomicInteger clientIdCounter = new AtomicInteger(1);
 
     protected static String getClientId(ProcessContext context) {
@@ -216,20 +210,32 @@ public abstract class AbstractJMSProcessor<T extends 
JMSWorker> extends Abstract
 
     @OnPrimaryNodeStateChange
     public void onPrimaryNodeChange(final PrimaryNodeState newState) {
-        if (isScheduled() && runOnPrimary && 
newState.equals(PrimaryNodeState.PRIMARY_NODE_REVOKED)) {
-            shutdownWorkers.set(true);
-            close();
-        } else {
-            shutdownWorkers.set(false);
+        final JmsWorkerLifecycle<T> lifecycle = workerLifecycle.get();
+        if (isScheduled() && runOnPrimary && lifecycle != null) {
+            if (newState.equals(PrimaryNodeState.PRIMARY_NODE_REVOKED)) {
+                lifecycle.retireGeneration();
+            } else if (newState.equals(PrimaryNodeState.ELECTED_PRIMARY_NODE)) 
{
+                lifecycle.activateFreshGeneration();
+            }
         }
     }
 
     @Override
     public void onTrigger(ProcessContext context, ProcessSession session) 
throws ProcessException {
-        T worker = workerPool.poll();
-        if (worker == null) {
+        final JmsWorkerLifecycle<T> lifecycle = workerLifecycle.get();
+        if (lifecycle == null) {
+            return;
+        }
+
+        final JmsWorkerLifecycle.Generation<T> generation = 
lifecycle.captureGeneration();
+        T worker = lifecycle.pollIdleWorker(generation);
+        if (worker == null && lifecycle.canCreateWorker(generation)) {
             try {
-                worker = buildTargetResource(context);
+                worker = buildTargetResource(context, 
lifecycle.getConnectionFactoryProvider());
+                if (!lifecycle.registerWorker(generation, worker)) {
+                    worker.shutdown();
+                    worker = null;
+                }
             } catch (Exception e) {
                 getLogger().error("Failed to initialize JMS Connection 
Factory", e);
                 context.yield();
@@ -237,6 +243,10 @@ public abstract class AbstractJMSProcessor<T extends 
JMSWorker> extends Abstract
             }
         }
 
+        if (worker == null) {
+            return;
+        }
+
         try {
             rendezvousWithJms(context, session, worker);
         } finally {
@@ -245,37 +255,31 @@ public abstract class AbstractJMSProcessor<T extends 
JMSWorker> extends Abstract
             //if worker is not valid anymore, don't put it back into a pool, 
try to rebuild it first, or discard.
             //this will be helpful in a situation, when JNDI has changed, or 
JMS server is not available
             //and reconnection is required.
-            if (worker == null || !worker.isValid()) {
+            if (!worker.isValid()) {
                 getLogger().debug("Worker is invalid. Will try re-create... ");
                 try {
-                    if (worker != null) {
-                        worker.shutdown();
-                    }
                     // Safe to cast. Method 
#buildTargetResource(ProcessContext context) sets only CachingConnectionFactory
                     CachingConnectionFactory currentCF = 
(CachingConnectionFactory) worker.jmsTemplate.getConnectionFactory();
-                    
connectionFactoryProvider.resetConnectionFactory(currentCF.getTargetConnectionFactory());
-                    worker = buildTargetResource(context);
+                    if (lifecycle.handleInvalidWorker(generation, worker, 
currentCF.getTargetConnectionFactory())) {
+                        final T replacementWorker = 
buildTargetResource(context, lifecycle.getConnectionFactoryProvider());
+                        if (lifecycle.registerWorker(generation, 
replacementWorker)) {
+                            lifecycle.releaseWorker(generation, 
replacementWorker);
+                        } else {
+                            replacementWorker.shutdown();
+                        }
+                    }
                 } catch (Exception e) {
-                    getLogger().error("Failed to rebuild:  {}", 
connectionFactoryProvider);
-                    worker = null;
-                }
-            }
-            if (worker != null) {
-                worker.jmsTemplate.setExplicitQosEnabled(false);
-                
worker.jmsTemplate.setDeliveryMode(Message.DEFAULT_DELIVERY_MODE);
-                worker.jmsTemplate.setTimeToLive(Message.DEFAULT_TIME_TO_LIVE);
-                worker.jmsTemplate.setPriority(Message.DEFAULT_PRIORITY);
-                if (!shutdownWorkers.get()) {
-                    workerPool.offer(worker);
-                } else {
-                    worker.shutdown();
+                    getLogger().error("Failed to rebuild:  {}", 
lifecycle.getConnectionFactoryProvider());
                 }
+            } else {
+                lifecycle.releaseWorker(generation, worker);
             }
         }
     }
 
     @OnScheduled
     public void setup(final ProcessContext context) {
+        final IJMSConnectionFactoryProvider connectionFactoryProvider;
         if (context.getProperty(CF_SERVICE).isSet()) {
             connectionFactoryProvider = 
context.getProperty(CF_SERVICE).asControllerService(JMSConnectionFactoryProviderDefinition.class);
         } else if 
(context.getProperty(JndiJmsConnectionFactoryProperties.JNDI_CONNECTION_FACTORY_NAME).isSet())
 {
@@ -286,21 +290,23 @@ public abstract class AbstractJMSProcessor<T extends 
JMSWorker> extends Abstract
             throw new ProcessException("No Connection Factory configured.");
         }
 
-        workerPool = new 
LinkedBlockingQueue<>(context.getMaxConcurrentTasks());
+        workerLifecycle.set(new 
JmsWorkerLifecycle<>(connectionFactoryProvider, 
context.getMaxConcurrentTasks(), getLogger()));
         runOnPrimary = 
context.getExecutionNode().equals(ExecutionNode.PRIMARY);
-        shutdownWorkers.set(false);
     }
 
     @OnUnscheduled
     public void shutdownConnectionFactoryProvider(final ProcessContext 
context) {
-        connectionFactoryProvider = null;
+        final JmsWorkerLifecycle<T> lifecycle = 
workerLifecycle.getAndSet(null);
+        if (lifecycle != null) {
+            lifecycle.closeCycle();
+        }
     }
 
     @OnStopped
     public void close() {
-        T worker;
-        while ((worker = workerPool.poll()) != null) {
-            worker.shutdown();
+        final JmsWorkerLifecycle<T> lifecycle = 
workerLifecycle.getAndSet(null);
+        if (lifecycle != null) {
+            lifecycle.closeCycle();
         }
     }
 
@@ -323,12 +329,12 @@ public abstract class AbstractJMSProcessor<T extends 
JMSWorker> extends Abstract
     /**
      * This method essentially performs initialization of this Processor by
      * obtaining an instance of the {@link ConnectionFactory} from the
-     * {@link JMSConnectionFactoryProvider} (ControllerService) and performing 
a
+     * {@link JMSConnectionFactoryProviderDefinition} (ControllerService) and 
performing a
      * series of {@link ConnectionFactory} adaptations which eventually results
      * in an instance of the {@link CachingConnectionFactory} used to construct
      * {@link JmsTemplate} used by this Processor.
      */
-    private T buildTargetResource(ProcessContext context) {
+    private T buildTargetResource(final ProcessContext context, final 
IJMSConnectionFactoryProvider connectionFactoryProvider) {
         final ConnectionFactory connectionFactory = 
connectionFactoryProvider.getConnectionFactory();
 
         final UserCredentialsConnectionFactoryAdapter cfCredentialsAdapter = 
new UserCredentialsConnectionFactoryAdapter();
diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JMSWorker.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JMSWorker.java
index 162b852d8d6..811f68c3c1d 100644
--- 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JMSWorker.java
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JMSWorker.java
@@ -22,6 +22,7 @@ import 
org.springframework.jms.connection.CachingConnectionFactory;
 import org.springframework.jms.core.JmsTemplate;
 
 import java.nio.channels.Channel;
+import java.util.concurrent.atomic.AtomicBoolean;
 
 /**
  * Base class for implementing publishing and consuming JMS workers.
@@ -34,6 +35,7 @@ abstract class JMSWorker {
     protected final JmsTemplate jmsTemplate;
     protected final ComponentLog processLog;
     private final CachingConnectionFactory connectionFactory;
+    private final AtomicBoolean shutdown = new AtomicBoolean(false);
     private boolean isValid = true;
 
     /**
@@ -51,7 +53,9 @@ abstract class JMSWorker {
     }
 
     public void shutdown() {
-        connectionFactory.destroy();
+        if (shutdown.compareAndSet(false, true)) {
+            connectionFactory.destroy();
+        }
     }
 
     @Override
diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JmsWorkerLifecycle.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JmsWorkerLifecycle.java
new file mode 100644
index 00000000000..64084d6da38
--- /dev/null
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/main/java/org/apache/nifi/jms/processors/JmsWorkerLifecycle.java
@@ -0,0 +1,188 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.nifi.jms.processors;
+
+import jakarta.jms.ConnectionFactory;
+import org.apache.nifi.jms.cf.IJMSConnectionFactoryProvider;
+import org.apache.nifi.logging.ComponentLog;
+import org.springframework.jms.core.JmsTemplate;
+
+import java.util.ArrayList;
+import java.util.HashSet;
+import java.util.List;
+import java.util.Set;
+import java.util.concurrent.BlockingQueue;
+import java.util.concurrent.LinkedBlockingQueue;
+
+final class JmsWorkerLifecycle<T extends JMSWorker> {
+
+    private final IJMSConnectionFactoryProvider connectionFactoryProvider;
+    private final int maxPoolSize;
+    private final ComponentLog logger;
+
+    private boolean cycleOpen = true;
+    private Generation<T> currentGeneration;
+
+    JmsWorkerLifecycle(final IJMSConnectionFactoryProvider 
connectionFactoryProvider, final int maxPoolSize, final ComponentLog logger) {
+        this.connectionFactoryProvider = connectionFactoryProvider;
+        this.maxPoolSize = maxPoolSize;
+        this.logger = logger;
+        this.currentGeneration = new Generation<>(maxPoolSize);
+    }
+
+    synchronized Generation<T> captureGeneration() {
+        return currentGeneration;
+    }
+
+    synchronized T pollIdleWorker(final Generation<T> generation) {
+        return isCurrentOpenGeneration(generation) ? 
generation.idleWorkers.poll() : null;
+    }
+
+    synchronized boolean canCreateWorker(final Generation<T> generation) {
+        return isCurrentOpenGeneration(generation);
+    }
+
+    boolean registerWorker(final Generation<T> generation, final T worker) {
+        boolean accepted;
+        synchronized (this) {
+            generation.allWorkers.add(worker);
+            accepted = isCurrentOpenGeneration(generation);
+            if (!accepted) {
+                generation.allWorkers.remove(worker);
+            }
+        }
+        return accepted;
+    }
+
+    void releaseWorker(final Generation<T> generation, final T worker) {
+        if (worker == null) {
+            return;
+        }
+
+        boolean closeWorker = true;
+        synchronized (this) {
+            if (worker.isValid() && isCurrentOpenGeneration(generation)) {
+                resetWorker(worker.jmsTemplate);
+                if (generation.idleWorkers.offer(worker)) {
+                    closeWorker = false;
+                } else {
+                    generation.allWorkers.remove(worker);
+                }
+            } else {
+                generation.allWorkers.remove(worker);
+            }
+        }
+
+        if (closeWorker) {
+            worker.shutdown();
+        }
+    }
+
+    boolean handleInvalidWorker(final Generation<T> generation, final T 
worker, final ConnectionFactory cachedConnectionFactory) {
+        synchronized (this) {
+            generation.allWorkers.remove(worker);
+        }
+
+        worker.shutdown();
+
+        synchronized (this) {
+            if (!isCurrentOpenGeneration(generation)) {
+                return false;
+            }
+
+            
connectionFactoryProvider.resetConnectionFactory(cachedConnectionFactory);
+            return true;
+        }
+    }
+
+    void retireGeneration() {
+        closeWorkers(retireCurrentGeneration());
+    }
+
+    synchronized void activateFreshGeneration() {
+        if (cycleOpen && currentGeneration.retired) {
+            currentGeneration = new Generation<>(maxPoolSize);
+        }
+    }
+
+    void closeCycle() {
+        closeWorkers(closeCurrentCycle());
+    }
+
+    IJMSConnectionFactoryProvider getConnectionFactoryProvider() {
+        return connectionFactoryProvider;
+    }
+
+    private synchronized List<T> retireCurrentGeneration() {
+        return currentGeneration.retireAndSnapshot();
+    }
+
+    private synchronized List<T> closeCurrentCycle() {
+        if (!cycleOpen) {
+            return List.of();
+        }
+
+        cycleOpen = false;
+        return currentGeneration.retireAndSnapshot();
+    }
+
+    private synchronized boolean isCurrentOpenGeneration(final Generation<T> 
generation) {
+        return cycleOpen && currentGeneration == generation && 
!generation.retired;
+    }
+
+    private void closeWorkers(final List<T> workers) {
+        for (final T worker : workers) {
+            try {
+                worker.shutdown();
+            } catch (final Exception e) {
+                logger.error("Failed to close JMS worker {}", worker, e);
+            }
+        }
+    }
+
+    private void resetWorker(final JmsTemplate jmsTemplate) {
+        jmsTemplate.setExplicitQosEnabled(false);
+        jmsTemplate.setDeliveryMode(jakarta.jms.Message.DEFAULT_DELIVERY_MODE);
+        jmsTemplate.setTimeToLive(jakarta.jms.Message.DEFAULT_TIME_TO_LIVE);
+        jmsTemplate.setPriority(jakarta.jms.Message.DEFAULT_PRIORITY);
+    }
+
+    static final class Generation<T extends JMSWorker> {
+
+        private final BlockingQueue<T> idleWorkers;
+        private final Set<T> allWorkers = new HashSet<>();
+
+        private boolean retired;
+
+        private Generation(final int maxPoolSize) {
+            idleWorkers = new LinkedBlockingQueue<>(maxPoolSize);
+        }
+
+        private List<T> retireAndSnapshot() {
+            if (retired) {
+                return List.of();
+            }
+
+            retired = true;
+            idleWorkers.clear();
+
+            final List<T> workers = new ArrayList<>(allWorkers);
+            allWorkers.clear();
+            return workers;
+        }
+    }
+}
diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/AbstractJMSProcessorTest.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/AbstractJMSProcessorTest.java
new file mode 100644
index 00000000000..d26e42a2fdf
--- /dev/null
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/AbstractJMSProcessorTest.java
@@ -0,0 +1,326 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.nifi.jms.processors;
+
+import jakarta.jms.ConnectionFactory;
+import org.apache.nifi.jms.cf.IJMSConnectionFactoryProvider;
+import org.apache.nifi.util.LogMessage;
+import org.apache.nifi.util.MockComponentLog;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.Timeout;
+
+import java.util.Arrays;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.Future;
+import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicBoolean;
+import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicReference;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertNull;
+import static org.junit.jupiter.api.Assertions.assertSame;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+import static org.mockito.Mockito.mock;
+
+public class AbstractJMSProcessorTest {
+
+    @Test
+    @Timeout(value = 10)
+    public void earlyUnscheduleShouldCloseIdleAndActiveWorkers() {
+        final RecordingConnectionFactoryProvider provider = new 
RecordingConnectionFactoryProvider();
+        final MockComponentLog logger = new MockComponentLog("processor", new 
Object());
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(provider, 2, logger);
+        final JmsWorkerLifecycle.Generation<TestWorker> generation = 
lifecycle.captureGeneration();
+
+        final TestWorker idleWorker = new TestWorker("idle");
+        final TestWorker activeWorker = new TestWorker("active");
+
+        assertTrue(lifecycle.registerWorker(generation, idleWorker));
+        lifecycle.releaseWorker(generation, idleWorker);
+        assertTrue(lifecycle.registerWorker(generation, activeWorker));
+
+        lifecycle.closeCycle();
+
+        assertEquals(1, idleWorker.getDestroyCalls());
+        assertEquals(1, activeWorker.getDestroyCalls());
+        assertNull(lifecycle.pollIdleWorker(generation));
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void retiredGenerationShouldCloseConstructedWorkerBeforeJmsUse() 
throws Exception {
+        final RecordingConnectionFactoryProvider provider = new 
RecordingConnectionFactoryProvider();
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(provider, 1, new MockComponentLog("processor", new 
Object()));
+        final JmsWorkerLifecycle.Generation<TestWorker> generation = 
lifecycle.captureGeneration();
+        final CountDownLatch workerConstructed = new CountDownLatch(1);
+        final CountDownLatch retireRequested = new CountDownLatch(1);
+        final AtomicReference<TestWorker> workerReference = new 
AtomicReference<>();
+        final ExecutorService executorService = 
Executors.newSingleThreadExecutor();
+
+        try {
+            final Future<Boolean> workerUsedFuture = executorService.submit(() 
-> {
+                assertTrue(lifecycle.canCreateWorker(generation));
+
+                final TestWorker worker = new TestWorker("constructing");
+                workerReference.set(worker);
+                workerConstructed.countDown();
+
+                assertTrue(retireRequested.await(5, TimeUnit.SECONDS));
+                final boolean accepted = lifecycle.registerWorker(generation, 
worker);
+                if (accepted) {
+                    worker.markUsed();
+                } else {
+                    worker.shutdown();
+                }
+
+                return worker.wasUsed();
+            });
+
+            assertTrue(workerConstructed.await(5, TimeUnit.SECONDS));
+            lifecycle.retireGeneration();
+            retireRequested.countDown();
+
+            assertFalse(workerUsedFuture.get(5, TimeUnit.SECONDS));
+        } finally {
+            executorService.shutdownNow();
+            executorService.awaitTermination(5, TimeUnit.SECONDS);
+        }
+
+        final TestWorker worker = workerReference.get();
+        assertNotNull(worker);
+        assertEquals(1, worker.getDestroyCalls());
+        assertFalse(worker.wasUsed());
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void returnedWorkerShouldBeClosedInsteadOfRePooledAfterRetirement() 
{
+        final RecordingConnectionFactoryProvider provider = new 
RecordingConnectionFactoryProvider();
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(provider, 1, new MockComponentLog("processor", new 
Object()));
+        final JmsWorkerLifecycle.Generation<TestWorker> generation = 
lifecycle.captureGeneration();
+        final TestWorker worker = new TestWorker("returned");
+
+        assertTrue(lifecycle.registerWorker(generation, worker));
+
+        lifecycle.retireGeneration();
+        lifecycle.releaseWorker(generation, worker);
+        lifecycle.activateFreshGeneration();
+
+        assertEquals(1, worker.getDestroyCalls());
+        assertNull(lifecycle.pollIdleWorker(generation));
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void retiredGenerationShouldNotResetOrRebuildInvalidWorker() {
+        final RecordingConnectionFactoryProvider provider = new 
RecordingConnectionFactoryProvider();
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(provider, 1, new MockComponentLog("processor", new 
Object()));
+        final JmsWorkerLifecycle.Generation<TestWorker> generation = 
lifecycle.captureGeneration();
+        final TestWorker invalidWorker = new TestWorker("invalid");
+
+        assertTrue(lifecycle.registerWorker(generation, invalidWorker));
+        invalidWorker.setValid(false);
+        lifecycle.retireGeneration();
+
+        final boolean rebuildAllowed = 
lifecycle.handleInvalidWorker(generation, invalidWorker, 
provider.getConnectionFactory());
+
+        assertFalse(rebuildAllowed);
+        assertEquals(0, provider.getResetCalls());
+        assertEquals(1, invalidWorker.getDestroyCalls());
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void oldCycleCleanupShouldNotMutateReplacementCycle() {
+        final JmsWorkerLifecycle<TestWorker> oldLifecycle = new 
JmsWorkerLifecycle<>(new RecordingConnectionFactoryProvider(), 1, new 
MockComponentLog("old", new Object()));
+        final JmsWorkerLifecycle<TestWorker> replacementLifecycle = new 
JmsWorkerLifecycle<>(new RecordingConnectionFactoryProvider(), 1, new 
MockComponentLog("replacement", new Object()));
+        final JmsWorkerLifecycle.Generation<TestWorker> oldGeneration = 
oldLifecycle.captureGeneration();
+        final JmsWorkerLifecycle.Generation<TestWorker> replacementGeneration 
= replacementLifecycle.captureGeneration();
+        final TestWorker oldWorker = new TestWorker("old-worker");
+        final TestWorker oldActiveWorker = new TestWorker("old-active-worker");
+        final TestWorker replacementWorker = new 
TestWorker("replacement-worker");
+
+        assertTrue(oldLifecycle.registerWorker(oldGeneration, oldWorker));
+        assertTrue(oldLifecycle.registerWorker(oldGeneration, 
oldActiveWorker));
+        assertTrue(replacementLifecycle.registerWorker(replacementGeneration, 
replacementWorker));
+        replacementLifecycle.releaseWorker(replacementGeneration, 
replacementWorker);
+
+        oldLifecycle.closeCycle();
+        oldLifecycle.releaseWorker(oldGeneration, oldActiveWorker);
+
+        assertEquals(1, oldWorker.getDestroyCalls());
+        assertEquals(1, oldActiveWorker.getDestroyCalls());
+        assertEquals(0, replacementWorker.getDestroyCalls());
+        assertSame(replacementWorker, 
replacementLifecycle.pollIdleWorker(replacementLifecycle.captureGeneration()));
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void 
primaryRevocationShouldRetireGenerationAndElectionShouldUseFreshGenerationOnly()
 {
+        final RecordingConnectionFactoryProvider provider = new 
RecordingConnectionFactoryProvider();
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(provider, 2, new MockComponentLog("processor", new 
Object()));
+        final JmsWorkerLifecycle.Generation<TestWorker> revokedGeneration = 
lifecycle.captureGeneration();
+        final TestWorker idleWorker = new TestWorker("idle");
+        final TestWorker activeWorker = new TestWorker("active");
+
+        assertTrue(lifecycle.registerWorker(revokedGeneration, idleWorker));
+        lifecycle.releaseWorker(revokedGeneration, idleWorker);
+        assertTrue(lifecycle.registerWorker(revokedGeneration, activeWorker));
+
+        lifecycle.retireGeneration();
+        lifecycle.activateFreshGeneration();
+
+        final JmsWorkerLifecycle.Generation<TestWorker> freshGeneration = 
lifecycle.captureGeneration();
+        final TestWorker freshWorker = new TestWorker("fresh");
+        assertTrue(lifecycle.registerWorker(freshGeneration, freshWorker));
+        lifecycle.releaseWorker(freshGeneration, freshWorker);
+        lifecycle.releaseWorker(revokedGeneration, activeWorker);
+
+        assertEquals(1, idleWorker.getDestroyCalls());
+        assertEquals(1, activeWorker.getDestroyCalls());
+        assertSame(freshWorker, lifecycle.pollIdleWorker(freshGeneration));
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void bulkCloseShouldContinueAfterShutdownFailureAndLogWorker() {
+        final RecordingConnectionFactoryProvider provider = new 
RecordingConnectionFactoryProvider();
+        final MockComponentLog logger = new MockComponentLog("processor", new 
Object());
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(provider, 3, logger);
+        final JmsWorkerLifecycle.Generation<TestWorker> generation = 
lifecycle.captureGeneration();
+        final TestWorker healthyWorker = new TestWorker("healthy");
+        final TestWorker failingWorker = new TestWorker("failing", true);
+        final TestWorker trailingWorker = new TestWorker("trailing");
+
+        assertTrue(lifecycle.registerWorker(generation, healthyWorker));
+        assertTrue(lifecycle.registerWorker(generation, failingWorker));
+        assertTrue(lifecycle.registerWorker(generation, trailingWorker));
+
+        lifecycle.closeCycle();
+
+        assertEquals(1, healthyWorker.getDestroyCalls());
+        assertEquals(1, failingWorker.getDestroyCalls());
+        assertEquals(1, trailingWorker.getDestroyCalls());
+        assertTrue(logger.getErrorMessages().stream().anyMatch(message -> 
containsWorkerReference(message, failingWorker)));
+    }
+
+    @Test
+    public void closeCycleShouldBeIdempotent() {
+        final JmsWorkerLifecycle<TestWorker> lifecycle = new 
JmsWorkerLifecycle<>(
+                new RecordingConnectionFactoryProvider(), 1, new 
MockComponentLog("processor", new Object()));
+        final JmsWorkerLifecycle.Generation<TestWorker> generation = 
lifecycle.captureGeneration();
+        final TestWorker worker = new TestWorker("worker");
+        assertTrue(lifecycle.registerWorker(generation, worker));
+
+        lifecycle.closeCycle();
+        lifecycle.closeCycle();
+
+        assertEquals(1, worker.getDestroyCalls());
+    }
+
+    private boolean containsWorkerReference(final LogMessage message, final 
TestWorker worker) {
+        final Object[] arguments = message.getArgs();
+        return arguments != null && Arrays.stream(arguments).anyMatch(argument 
-> argument == worker);
+    }
+
+    private static final class RecordingConnectionFactoryProvider implements 
IJMSConnectionFactoryProvider {
+
+        private final ConnectionFactory connectionFactory = 
mock(ConnectionFactory.class);
+        private final AtomicInteger resetCalls = new AtomicInteger();
+
+        @Override
+        public ConnectionFactory getConnectionFactory() {
+            return connectionFactory;
+        }
+
+        @Override
+        public void resetConnectionFactory(final ConnectionFactory 
cachedFactory) {
+            resetCalls.incrementAndGet();
+        }
+
+        private int getResetCalls() {
+            return resetCalls.get();
+        }
+
+    }
+
+    private static final class TestWorker extends JMSWorker {
+
+        private final DestroyTrackingCachingConnectionFactory 
cachingConnectionFactory;
+        private final AtomicBoolean used = new AtomicBoolean();
+
+        private TestWorker(final String destinationName) {
+            this(destinationName, false);
+        }
+
+        private TestWorker(final String destinationName, final boolean 
failDestroy) {
+            this(new 
DestroyTrackingCachingConnectionFactory(mock(ConnectionFactory.class), 
failDestroy), destinationName);
+        }
+
+        private TestWorker(final DestroyTrackingCachingConnectionFactory 
connectionFactory, final String destinationName) {
+            super(connectionFactory, createJmsTemplate(connectionFactory, 
destinationName), new MockComponentLog(destinationName, destinationName));
+            this.cachingConnectionFactory = connectionFactory;
+        }
+
+        private static org.springframework.jms.core.JmsTemplate 
createJmsTemplate(final DestroyTrackingCachingConnectionFactory 
connectionFactory, final String destinationName) {
+            final org.springframework.jms.core.JmsTemplate jmsTemplate = new 
org.springframework.jms.core.JmsTemplate();
+            jmsTemplate.setConnectionFactory(connectionFactory);
+            jmsTemplate.setDefaultDestinationName(destinationName);
+            return jmsTemplate;
+        }
+
+        private void markUsed() {
+            used.set(true);
+        }
+
+        private boolean wasUsed() {
+            return used.get();
+        }
+
+        private int getDestroyCalls() {
+            return cachingConnectionFactory.getDestroyCalls();
+        }
+    }
+
+    private static final class DestroyTrackingCachingConnectionFactory extends 
org.springframework.jms.connection.CachingConnectionFactory {
+
+        private final AtomicInteger destroyCalls = new AtomicInteger();
+        private final boolean failFirstDestroy;
+
+        private DestroyTrackingCachingConnectionFactory(final 
ConnectionFactory targetConnectionFactory, final boolean failFirstDestroy) {
+            super(targetConnectionFactory);
+            this.failFirstDestroy = failFirstDestroy;
+        }
+
+        @Override
+        public void destroy() {
+            final int destroyCall = destroyCalls.incrementAndGet();
+            if (failFirstDestroy && destroyCall == 1) {
+                throw new RuntimeException("destroy failed");
+            }
+        }
+
+        private int getDestroyCalls() {
+            return destroyCalls.get();
+        }
+    }
+}
diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/ConsumeJMSIT.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/ConsumeJMSIT.java
index fe31c2f36d5..3d6a79d9b63 100644
--- 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/ConsumeJMSIT.java
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/ConsumeJMSIT.java
@@ -26,6 +26,7 @@ import jakarta.jms.ConnectionFactory;
 import jakarta.jms.JMSException;
 import jakarta.jms.MapMessage;
 import jakarta.jms.Message;
+import jakarta.jms.MessageConsumer;
 import jakarta.jms.MessageProducer;
 import jakarta.jms.ObjectMessage;
 import jakarta.jms.Session;
@@ -51,8 +52,11 @@ import org.apache.nifi.processor.exception.ProcessException;
 import org.apache.nifi.processor.io.OutputStreamCallback;
 import org.apache.nifi.reporting.InitializationException;
 import org.apache.nifi.scheduling.ExecutionNode;
+import org.apache.nifi.state.MockStateManager;
 import org.apache.nifi.util.MockFlowFile;
 import org.apache.nifi.util.MockProcessContext;
+import org.apache.nifi.util.MockProcessSession;
+import org.apache.nifi.util.SharedSessionState;
 import org.apache.nifi.util.TestRunner;
 import org.apache.nifi.util.TestRunners;
 import org.junit.jupiter.api.Test;
@@ -63,13 +67,24 @@ import org.springframework.jms.core.MessageCreator;
 import org.springframework.jms.support.JmsHeaders;
 
 import java.io.IOException;
+import java.lang.reflect.InvocationTargetException;
+import java.lang.reflect.Method;
+import java.lang.reflect.Proxy;
 import java.net.URI;
 import java.net.UnknownHostException;
 import java.util.HashMap;
 import java.util.List;
 import java.util.Map;
+import java.util.Objects;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.Future;
 import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicLong;
 import java.util.concurrent.atomic.AtomicReference;
+import java.util.function.Consumer;
 import javax.net.SocketFactory;
 
 import static java.util.Arrays.asList;
@@ -78,7 +93,9 @@ import static 
org.apache.nifi.jms.processors.helpers.JMSTestUtil.createJsonRecor
 import static 
org.apache.nifi.jms.processors.helpers.JMSTestUtil.createJsonRecordSetWriterService;
 import static org.junit.jupiter.api.Assertions.assertEquals;
 import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.junit.jupiter.api.Assertions.assertInstanceOf;
 import static org.junit.jupiter.api.Assertions.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertThrows;
 import static org.junit.jupiter.api.Assertions.assertTrue;
 import static org.mockito.ArgumentMatchers.any;
 import static org.mockito.Mockito.doThrow;
@@ -690,6 +707,97 @@ public class ConsumeJMSIT {
         }
     }
 
+    @Test
+    @Timeout(value = 10000, unit = TimeUnit.MILLISECONDS)
+    public void 
activeReceiveShouldBeInterruptedOnUnscheduledWithoutFlowFileOrReplacementConnection()
 throws Exception {
+        final BrokerService broker = new BrokerService();
+        final ExecutorService executorService = 
Executors.newSingleThreadExecutor();
+        try {
+            broker.setPersistent(false);
+            broker.setBrokerName("blocked-receive-broker");
+            broker.start();
+
+            final ActiveMQConnectionFactory innerConnectionFactory = new 
ActiveMQConnectionFactory("vm://blocked-receive-broker");
+            final BlockingReceiveConnectionFactory controlledConnectionFactory 
= new BlockingReceiveConnectionFactory(innerConnectionFactory);
+            final CountingConsumeJMS processor = new CountingConsumeJMS();
+            final TestRunner runner = initializeTestRunner(processor, 
controlledConnectionFactory.getConnectionFactory(), 
"blocked-receive-destination");
+            runner.setProperty(ConsumeJMS.TIMEOUT, "30 sec");
+
+            final ProcessContext processContext = runner.getProcessContext();
+            processor.onSchedule(processContext);
+            processor.setup(processContext);
+
+            final MockProcessSession processSession = new 
MockProcessSession(new SharedSessionState(processor, new AtomicLong(0L)), 
processor);
+            final Future<Throwable> future = executorService.submit(() -> {
+                try {
+                    processor.onTrigger(processContext, processSession);
+                    return null;
+                } catch (final Throwable throwable) {
+                    return throwable;
+                }
+            });
+
+            assertTrue(controlledConnectionFactory.awaitReceiveEntered(5, 
TimeUnit.SECONDS));
+
+            processor.shutdownConnectionFactoryProvider(processContext);
+
+            final Throwable thrown = future.get(5, TimeUnit.SECONDS);
+            assertTrue(thrown == null || thrown instanceof ProcessException || 
thrown instanceof JMSException
+                            || thrown instanceof 
org.springframework.jms.IllegalStateException,
+                    "Unexpected exception type returned from blocked receive 
shutdown path: " + thrown);
+            assertTrue(controlledConnectionFactory.awaitConnectionClosed(5, 
TimeUnit.SECONDS));
+            assertEquals(0, broker.getCurrentConnections());
+            assertEquals(1, 
controlledConnectionFactory.getOpenedConnections());
+            assertEquals(1, processor.getBuildCount());
+            
assertTrue(processSession.getFlowFilesForRelationship(ConsumeJMS.REL_SUCCESS).isEmpty());
+            
assertTrue(processSession.getFlowFilesForRelationship(ConsumeJMS.REL_PARSE_FAILURE).isEmpty());
+        } finally {
+            executorService.shutdownNow();
+            broker.stop();
+        }
+    }
+
+    @Test
+    @Timeout(value = 10000, unit = TimeUnit.MILLISECONDS)
+    public void 
delayedCommitCallbackShouldFailAcknowledgeAfterShutdownAndAllowRedelivery() 
throws Exception {
+        final BrokerService broker = new BrokerService();
+        try {
+            broker.setPersistent(false);
+            broker.setBrokerName("delayed-commit-broker");
+            broker.start();
+
+            final String destinationName = "delayed-commit-destination";
+            final ActiveMQConnectionFactory connectionFactory = new 
ActiveMQConnectionFactory("vm://delayed-commit-broker");
+            publishQueueMessage(connectionFactory, destinationName, 
"delayed-ack-message");
+
+            final ConsumeJMS processor = new ConsumeJMS();
+            final TestRunner runner = initializeTestRunner(processor, 
connectionFactory, destinationName);
+            final ProcessContext processContext = runner.getProcessContext();
+            processor.onSchedule(processContext);
+            processor.setup(processContext);
+
+            final DelayedCommitProcessSession processSession = new 
DelayedCommitProcessSession(new SharedSessionState(processor, new 
AtomicLong(0L)), processor);
+
+            processor.onTrigger(processContext, processSession);
+
+            assertEquals(1, 
processSession.getFlowFilesForRelationship(ConsumeJMS.REL_SUCCESS).size());
+            assertNotNull(processSession.getDelayedSuccessCallback());
+
+            processor.shutdownConnectionFactoryProvider(processContext);
+
+            final ProcessException exception = 
assertThrows(ProcessException.class, processSession::runDelayedSuccessCallback);
+            assertNotNull(exception.getCause());
+
+            final JmsTemplate jmsTemplate = new JmsTemplate(connectionFactory);
+            jmsTemplate.setReceiveTimeout(5000L);
+            final Message redelivered = jmsTemplate.receive(destinationName);
+            assertInstanceOf(TextMessage.class, redelivered);
+            assertEquals("delayed-ack-message", ((TextMessage) 
redelivered).getText());
+        } finally {
+            broker.stop();
+        }
+    }
+
     private static ArrayNode createTestJsonInput() {
         final ObjectMapper mapper = new ObjectMapper();
 
@@ -766,4 +874,143 @@ public class ConsumeJMSIT {
         return c1Consumer;
     }
 
+    private static void publishQueueMessage(final ConnectionFactory 
connectionFactory, final String destinationName, final String messageText) {
+        final JmsTemplate jmsTemplate = new JmsTemplate(connectionFactory);
+        jmsTemplate.send(destinationName, session -> 
session.createTextMessage(messageText));
+    }
+
+    private static Object invokeTarget(final Object target, final Method 
method, final Object[] args) throws Throwable {
+        try {
+            return method.invoke(target, args);
+        } catch (final InvocationTargetException exception) {
+            throw exception.getCause();
+        }
+    }
+
+    private static final class CountingConsumeJMS extends ConsumeJMS {
+
+        private final AtomicInteger buildCount = new AtomicInteger();
+
+        @Override
+        protected JMSConsumer finishBuildingJmsWorker(final 
CachingConnectionFactory connectionFactory, final JmsTemplate jmsTemplate,
+                                                      final ProcessContext 
processContext) {
+            buildCount.incrementAndGet();
+            return super.finishBuildingJmsWorker(connectionFactory, 
jmsTemplate, processContext);
+        }
+
+        private int getBuildCount() {
+            return buildCount.get();
+        }
+    }
+
+    private static final class DelayedCommitProcessSession extends 
MockProcessSession {
+
+        private Runnable delayedSuccessCallback;
+
+        private DelayedCommitProcessSession(final SharedSessionState 
sharedState, final ConsumeJMS processor) {
+            super(sharedState, processor, new MockStateManager(processor));
+        }
+
+        @Override
+        public void commitAsync(final Runnable onSuccess, final 
Consumer<Throwable> onFailure) {
+            delayedSuccessCallback = onSuccess;
+            super.commitAsync(null, onFailure);
+        }
+
+        private Runnable getDelayedSuccessCallback() {
+            return delayedSuccessCallback;
+        }
+
+        private void runDelayedSuccessCallback() {
+            delayedSuccessCallback.run();
+        }
+    }
+
+    private static final class BlockingReceiveConnectionFactory {
+
+        private final ConnectionFactory connectionFactory;
+        private final CountDownLatch receiveEntered = new CountDownLatch(1);
+        private final CountDownLatch connectionClosed = new CountDownLatch(1);
+        private final AtomicInteger openedConnections = new AtomicInteger();
+
+        private BlockingReceiveConnectionFactory(final ConnectionFactory 
targetConnectionFactory) {
+            Objects.requireNonNull(targetConnectionFactory);
+            connectionFactory = (ConnectionFactory) Proxy.newProxyInstance(
+                    ConnectionFactory.class.getClassLoader(),
+                    new Class[] {ConnectionFactory.class},
+                    (proxy, method, args) -> {
+                        final Object result = 
invokeTarget(targetConnectionFactory, method, args);
+                        if ("createConnection".equals(method.getName())) {
+                            openedConnections.incrementAndGet();
+                            return createConnectionProxy((Connection) result);
+                        }
+                        return result;
+                    });
+        }
+
+        private ConnectionFactory getConnectionFactory() {
+            return connectionFactory;
+        }
+
+        private boolean awaitReceiveEntered(final long timeout, final TimeUnit 
timeUnit) throws InterruptedException {
+            return receiveEntered.await(timeout, timeUnit);
+        }
+
+        private boolean awaitConnectionClosed(final long timeout, final 
TimeUnit timeUnit) throws InterruptedException {
+            return connectionClosed.await(timeout, timeUnit);
+        }
+
+        private int getOpenedConnections() {
+            return openedConnections.get();
+        }
+
+        private Connection createConnectionProxy(final Connection 
targetConnection) {
+            return (Connection) Proxy.newProxyInstance(
+                    Connection.class.getClassLoader(),
+                    new Class[] {Connection.class},
+                    (proxy, method, args) -> {
+                        final Object result = invokeTarget(targetConnection, 
method, args);
+                        if ("createSession".equals(method.getName())) {
+                            return createSessionProxy((Session) result);
+                        }
+                        if ("close".equals(method.getName())) {
+                            connectionClosed.countDown();
+                        }
+                        return result;
+                    });
+        }
+
+        private Session createSessionProxy(final Session targetSession) {
+            return (Session) Proxy.newProxyInstance(
+                    Session.class.getClassLoader(),
+                    new Class[] {Session.class},
+                    (proxy, method, args) -> {
+                        final Object result = invokeTarget(targetSession, 
method, args);
+                        if (isConsumerFactoryMethod(method.getName())) {
+                            return 
createMessageConsumerProxy((MessageConsumer) result);
+                        }
+                        return result;
+                    });
+        }
+
+        private MessageConsumer createMessageConsumerProxy(final 
MessageConsumer targetConsumer) {
+            return (MessageConsumer) Proxy.newProxyInstance(
+                    MessageConsumer.class.getClassLoader(),
+                    new Class[] {MessageConsumer.class},
+                    (proxy, method, args) -> {
+                        if ("receive".equals(method.getName())) {
+                            receiveEntered.countDown();
+                        }
+                        return invokeTarget(targetConsumer, method, args);
+                    });
+        }
+
+        private boolean isConsumerFactoryMethod(final String methodName) {
+            return "createConsumer".equals(methodName)
+                    || "createDurableConsumer".equals(methodName)
+                    || "createSharedConsumer".equals(methodName)
+                    || "createSharedDurableConsumer".equals(methodName);
+        }
+    }
+
 }
diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/JMSWorkerTest.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/JMSWorkerTest.java
new file mode 100644
index 00000000000..cbb6c919796
--- /dev/null
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/JMSWorkerTest.java
@@ -0,0 +1,173 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *     http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+package org.apache.nifi.jms.processors;
+
+import jakarta.jms.ConnectionFactory;
+import org.apache.nifi.logging.ComponentLog;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.Timeout;
+import org.springframework.jms.connection.CachingConnectionFactory;
+import org.springframework.jms.core.JmsTemplate;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.ExecutionException;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.Future;
+import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicInteger;
+
+import static org.junit.jupiter.api.Assertions.assertDoesNotThrow;
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertInstanceOf;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+import static org.junit.jupiter.api.Assertions.assertTrue;
+import static org.mockito.Mockito.mock;
+
+public class JMSWorkerTest {
+
+    @Test
+    @Timeout(value = 10)
+    public void 
shutdownShouldDestroyConnectionFactoryOnceAcrossConcurrentCallers() throws 
Exception {
+        final DestroyTrackingCachingConnectionFactory cachingConnectionFactory 
= new DestroyTrackingCachingConnectionFactory(mock(ConnectionFactory.class), 
false);
+        final TestWorker worker = new TestWorker(cachingConnectionFactory);
+        final int callers = 8;
+        final CountDownLatch ready = new CountDownLatch(callers);
+        final CountDownLatch start = new CountDownLatch(1);
+        final ExecutorService executorService = 
Executors.newFixedThreadPool(callers);
+
+        try {
+            final List<Future<Void>> futures = new ArrayList<>();
+            for (int index = 0; index < callers; index++) {
+                futures.add(executorService.submit(() -> {
+                    ready.countDown();
+                    assertTrue(start.await(5, TimeUnit.SECONDS));
+                    worker.shutdown();
+                    return null;
+                }));
+            }
+
+            assertTrue(ready.await(5, TimeUnit.SECONDS));
+            start.countDown();
+
+            for (final Future<Void> future : futures) {
+                future.get(5, TimeUnit.SECONDS);
+            }
+        } finally {
+            executorService.shutdownNow();
+            executorService.awaitTermination(5, TimeUnit.SECONDS);
+        }
+
+        assertEquals(1, cachingConnectionFactory.getDestroyCalls());
+    }
+
+    @Test
+    @Timeout(value = 10)
+    public void shutdownShouldNotRetryDestroyAfterInitialFailure() throws 
Exception {
+        final DestroyTrackingCachingConnectionFactory cachingConnectionFactory 
= new DestroyTrackingCachingConnectionFactory(mock(ConnectionFactory.class), 
true);
+        final TestWorker worker = new TestWorker(cachingConnectionFactory);
+
+        final RuntimeException thrown = getRuntimeException(worker);
+
+        assertNotNull(thrown, "Expected one shutdown caller to receive the 
destroy failure");
+        assertEquals("destroy failed", thrown.getMessage());
+        assertEquals(1, cachingConnectionFactory.getDestroyCalls());
+        assertDoesNotThrow(worker::shutdown);
+        assertEquals(1, cachingConnectionFactory.getDestroyCalls());
+    }
+
+    private RuntimeException getRuntimeException(final TestWorker worker) 
throws Exception {
+        final ExecutorService executorService = 
Executors.newFixedThreadPool(2);
+        final CountDownLatch ready = new CountDownLatch(2);
+        final CountDownLatch start = new CountDownLatch(1);
+
+        try {
+            final Future<Void> firstFuture = executorService.submit(() -> {
+                ready.countDown();
+                assertTrue(start.await(5, TimeUnit.SECONDS));
+                worker.shutdown();
+                return null;
+            });
+            final Future<Void> secondFuture = executorService.submit(() -> {
+                ready.countDown();
+                assertTrue(start.await(5, TimeUnit.SECONDS));
+                worker.shutdown();
+                return null;
+            });
+
+            assertTrue(ready.await(5, TimeUnit.SECONDS));
+            start.countDown();
+
+            RuntimeException runtimeException = null;
+            runtimeException = mergeRuntimeException(runtimeException, 
firstFuture);
+            runtimeException = mergeRuntimeException(runtimeException, 
secondFuture);
+            return runtimeException;
+        } finally {
+            executorService.shutdownNow();
+            executorService.awaitTermination(5, TimeUnit.SECONDS);
+        }
+    }
+
+    private RuntimeException mergeRuntimeException(final RuntimeException 
current, final Future<Void> future) throws Exception {
+        try {
+            future.get(5, TimeUnit.SECONDS);
+            return current;
+        } catch (final ExecutionException executionException) {
+            assertInstanceOf(RuntimeException.class, 
executionException.getCause());
+            return current == null ? (RuntimeException) 
executionException.getCause() : current;
+        }
+    }
+
+    private static final class TestWorker extends JMSWorker {
+
+        private TestWorker(final DestroyTrackingCachingConnectionFactory 
connectionFactory) {
+            super(connectionFactory, createJmsTemplate(connectionFactory), 
mock(ComponentLog.class));
+        }
+
+        private static JmsTemplate createJmsTemplate(final 
CachingConnectionFactory connectionFactory) {
+            final JmsTemplate jmsTemplate = new JmsTemplate();
+            jmsTemplate.setConnectionFactory(connectionFactory);
+            jmsTemplate.setDefaultDestinationName("test-destination");
+            return jmsTemplate;
+        }
+    }
+
+    private static final class DestroyTrackingCachingConnectionFactory extends 
CachingConnectionFactory {
+
+        private final AtomicInteger destroyCalls = new AtomicInteger();
+        private final boolean failFirstDestroy;
+
+        private DestroyTrackingCachingConnectionFactory(final 
ConnectionFactory targetConnectionFactory, final boolean failFirstDestroy) {
+            super(targetConnectionFactory);
+            this.failFirstDestroy = failFirstDestroy;
+        }
+
+        @Override
+        public void destroy() {
+            final int destroyCall = destroyCalls.incrementAndGet();
+            if (failFirstDestroy && destroyCall == 1) {
+                throw new RuntimeException("destroy failed");
+            }
+        }
+
+        private int getDestroyCalls() {
+            return destroyCalls.get();
+        }
+    }
+}
diff --git 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/PublishJMSIT.java
 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/PublishJMSIT.java
index 525c99dd83b..40cc7855160 100644
--- 
a/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/PublishJMSIT.java
+++ 
b/nifi-extension-bundles/nifi-jms-bundle/nifi-jms-processors/src/test/java/org/apache/nifi/jms/processors/PublishJMSIT.java
@@ -19,9 +19,12 @@ package org.apache.nifi.jms.processors;
 import com.fasterxml.jackson.databind.ObjectMapper;
 import com.fasterxml.jackson.databind.node.ArrayNode;
 import jakarta.jms.BytesMessage;
+import jakarta.jms.Connection;
 import jakarta.jms.ConnectionFactory;
 import jakarta.jms.Message;
+import jakarta.jms.MessageProducer;
 import jakarta.jms.Queue;
+import jakarta.jms.Session;
 import jakarta.jms.TextMessage;
 import org.apache.activemq.ActiveMQConnectionFactory;
 import org.apache.activemq.broker.BrokerService;
@@ -40,6 +43,8 @@ import org.apache.nifi.provenance.ProvenanceEventType;
 import org.apache.nifi.reporting.InitializationException;
 import org.apache.nifi.util.MockFlowFile;
 import org.apache.nifi.util.MockProcessContext;
+import org.apache.nifi.util.MockProcessSession;
+import org.apache.nifi.util.SharedSessionState;
 import org.apache.nifi.util.TestRunner;
 import org.apache.nifi.util.TestRunners;
 import org.junit.jupiter.api.AfterEach;
@@ -50,12 +55,21 @@ import org.springframework.jms.core.JmsTemplate;
 import org.springframework.jms.support.JmsHeaders;
 
 import java.io.IOException;
+import java.lang.reflect.InvocationTargetException;
+import java.lang.reflect.Method;
 import java.lang.reflect.Proxy;
 import java.net.URI;
 import java.util.HashMap;
 import java.util.List;
 import java.util.Map;
+import java.util.Objects;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.Future;
 import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicInteger;
+import java.util.concurrent.atomic.AtomicLong;
 import java.util.concurrent.atomic.AtomicReference;
 import javax.net.SocketFactory;
 
@@ -550,6 +564,110 @@ public class PublishJMSIT {
         }
     }
 
+    @Test
+    @Timeout(value = 10000, unit = TimeUnit.MILLISECONDS)
+    public void activeSendShouldRouteSuccessWhenSendCompletesBeforeShutdown() 
throws Exception {
+        final BrokerService broker = new BrokerService();
+        final ExecutorService executorService = 
Executors.newSingleThreadExecutor();
+        try {
+            broker.setPersistent(false);
+            broker.setBrokerName("publisher-success-broker");
+            broker.start();
+
+            final BlockingProducerConnectionFactory 
controlledConnectionFactory = new BlockingProducerConnectionFactory(
+                    new 
ActiveMQConnectionFactory("vm://publisher-success-broker"));
+            final CountingPublishJMS processor = new CountingPublishJMS();
+            final TestRunner runner = initializeTestRunner(processor, 
controlledConnectionFactory.getConnectionFactory(), 
"publisher-success-destination");
+
+            final ProcessContext processContext = runner.getProcessContext();
+            processor.onScheduled(processContext);
+            processor.setup(processContext);
+
+            final MockProcessSession processSession = 
createEnqueuedSession(processor, "success-message");
+            final Future<Throwable> future = executorService.submit(() -> {
+                try {
+                    processor.onTrigger(processContext, processSession);
+                    return null;
+                } catch (final Throwable throwable) {
+                    return throwable;
+                }
+            });
+
+            assertTrue(controlledConnectionFactory.awaitSendEntered(5, 
TimeUnit.SECONDS));
+            controlledConnectionFactory.allowSendToProceed();
+            assertTrue(controlledConnectionFactory.awaitSendCompleted(5, 
TimeUnit.SECONDS));
+            assertTrue(processor.awaitRendezvousCompleted(5, 
TimeUnit.SECONDS));
+
+            processor.shutdownConnectionFactoryProvider(processContext);
+            processor.allowRendezvousReturn();
+
+            assertNull(future.get(5, TimeUnit.SECONDS));
+            assertEquals(1, 
processSession.getFlowFilesForRelationship(REL_SUCCESS).size());
+            
assertTrue(processSession.getFlowFilesForRelationship(REL_FAILURE).isEmpty());
+            assertEquals(1, 
controlledConnectionFactory.getOpenedConnections());
+            assertEquals(1, processor.getBuildCount());
+
+            final JmsTemplate verifyTemplate = new JmsTemplate(new 
ActiveMQConnectionFactory("vm://publisher-success-broker"));
+            verifyTemplate.setReceiveTimeout(5000L);
+            final Message message = 
verifyTemplate.receive("publisher-success-destination");
+            assertInstanceOf(BytesMessage.class, message);
+        } finally {
+            executorService.shutdownNow();
+            broker.stop();
+        }
+    }
+
+    @Test
+    @Timeout(value = 10000, unit = TimeUnit.MILLISECONDS)
+    public void 
activeSendShouldRouteFailureWithoutReplacementAfterShutdownInterruptsSend() 
throws Exception {
+        final BrokerService broker = new BrokerService();
+        final ExecutorService executorService = 
Executors.newSingleThreadExecutor();
+        try {
+            broker.setPersistent(false);
+            broker.setBrokerName("publisher-failure-broker");
+            broker.start();
+
+            final BlockingProducerConnectionFactory 
controlledConnectionFactory = new BlockingProducerConnectionFactory(
+                    new 
ActiveMQConnectionFactory("vm://publisher-failure-broker"));
+            final CountingPublishJMS processor = new CountingPublishJMS();
+            final TestRunner runner = initializeTestRunner(processor, 
controlledConnectionFactory.getConnectionFactory(), 
"publisher-failure-destination");
+
+            final ProcessContext processContext = runner.getProcessContext();
+            processor.onScheduled(processContext);
+            processor.setup(processContext);
+
+            final MockProcessSession processSession = 
createEnqueuedSession(processor, "failure-message");
+            final Future<Throwable> future = executorService.submit(() -> {
+                try {
+                    processor.onTrigger(processContext, processSession);
+                    return null;
+                } catch (final Throwable throwable) {
+                    return throwable;
+                }
+            });
+
+            assertTrue(controlledConnectionFactory.awaitSendEntered(5, 
TimeUnit.SECONDS));
+
+            processor.shutdownConnectionFactoryProvider(processContext);
+            controlledConnectionFactory.allowSendToProceed();
+            assertTrue(processor.awaitRendezvousCompleted(5, 
TimeUnit.SECONDS));
+            processor.allowRendezvousReturn();
+
+            assertNull(future.get(5, TimeUnit.SECONDS));
+            
assertTrue(processSession.getFlowFilesForRelationship(REL_SUCCESS).isEmpty());
+            assertEquals(1, 
processSession.getFlowFilesForRelationship(REL_FAILURE).size());
+            assertEquals(1, 
controlledConnectionFactory.getOpenedConnections());
+            assertEquals(1, processor.getBuildCount());
+
+            final JmsTemplate verifyTemplate = new JmsTemplate(new 
ActiveMQConnectionFactory("vm://publisher-failure-broker"));
+            verifyTemplate.setReceiveTimeout(250L);
+            
assertNull(verifyTemplate.receive("publisher-failure-destination"));
+        } finally {
+            executorService.shutdownNow();
+            broker.stop();
+        }
+    }
+
     @Test
     public void 
whenExceptionIsRaisedDuringConnectionFactoryInitializationTheProcessorShouldBeYielded()
 {
         final String nonExistentClassName = "DummyInitialContextFactoryClass";
@@ -758,6 +876,147 @@ public class PublishJMSIT {
         assertEquals(expectedDetails, event.getDetails());
     }
 
+    private MockProcessSession createEnqueuedSession(final PublishJMS 
processor, final String content) {
+        final SharedSessionState sharedSessionState = new 
SharedSessionState(processor, new AtomicLong(0L));
+        final MockProcessSession processSession = new 
MockProcessSession(sharedSessionState, processor);
+        
sharedSessionState.getFlowFileQueue().offer(processSession.createFlowFile(content.getBytes()));
+        return processSession;
+    }
+
+    private static Object invokeTarget(final Object target, final Method 
method, final Object[] args) throws Throwable {
+        try {
+            return method.invoke(target, args);
+        } catch (final InvocationTargetException exception) {
+            throw exception.getCause();
+        }
+    }
+
+    private static final class CountingPublishJMS extends PublishJMS {
+
+        private final AtomicInteger buildCount = new AtomicInteger();
+        private final CountDownLatch rendezvousCompleted = new 
CountDownLatch(1);
+        private final CountDownLatch allowRendezvousReturn = new 
CountDownLatch(1);
+
+        @Override
+        protected void rendezvousWithJms(final ProcessContext context, final 
ProcessSession processSession, final JMSPublisher publisher) throws 
ProcessException {
+            super.rendezvousWithJms(context, processSession, publisher);
+            rendezvousCompleted.countDown();
+            try {
+                assertTrue(allowRendezvousReturn.await(5, TimeUnit.SECONDS));
+            } catch (final InterruptedException e) {
+                Thread.currentThread().interrupt();
+                throw new ProcessException(e);
+            }
+        }
+
+        @Override
+        protected JMSPublisher finishBuildingJmsWorker(final 
org.springframework.jms.connection.CachingConnectionFactory connectionFactory,
+                                                       final JmsTemplate 
jmsTemplate, final ProcessContext processContext) {
+            buildCount.incrementAndGet();
+            return super.finishBuildingJmsWorker(connectionFactory, 
jmsTemplate, processContext);
+        }
+
+        private int getBuildCount() {
+            return buildCount.get();
+        }
+
+        private boolean awaitRendezvousCompleted(final long timeout, final 
TimeUnit timeUnit) throws InterruptedException {
+            return rendezvousCompleted.await(timeout, timeUnit);
+        }
+
+        private void allowRendezvousReturn() {
+            allowRendezvousReturn.countDown();
+        }
+    }
+
+    private static final class BlockingProducerConnectionFactory {
+
+        private final ConnectionFactory connectionFactory;
+        private final CountDownLatch sendEntered = new CountDownLatch(1);
+        private final CountDownLatch allowSend = new CountDownLatch(1);
+        private final CountDownLatch sendCompleted = new CountDownLatch(1);
+        private final AtomicInteger openedConnections = new AtomicInteger();
+
+        private BlockingProducerConnectionFactory(final ConnectionFactory 
targetConnectionFactory) {
+            Objects.requireNonNull(targetConnectionFactory);
+            connectionFactory = (ConnectionFactory) Proxy.newProxyInstance(
+                    ConnectionFactory.class.getClassLoader(),
+                    new Class[] {ConnectionFactory.class},
+                    (proxy, method, args) -> {
+                        final Object result = 
invokeTarget(targetConnectionFactory, method, args);
+                        if ("createConnection".equals(method.getName())) {
+                            openedConnections.incrementAndGet();
+                            return createConnectionProxy((Connection) result);
+                        }
+                        return result;
+                    });
+        }
+
+        private ConnectionFactory getConnectionFactory() {
+            return connectionFactory;
+        }
+
+        private boolean awaitSendEntered(final long timeout, final TimeUnit 
timeUnit) throws InterruptedException {
+            return sendEntered.await(timeout, timeUnit);
+        }
+
+        private boolean awaitSendCompleted(final long timeout, final TimeUnit 
timeUnit) throws InterruptedException {
+            return sendCompleted.await(timeout, timeUnit);
+        }
+
+        private void allowSendToProceed() {
+            allowSend.countDown();
+        }
+
+        private int getOpenedConnections() {
+            return openedConnections.get();
+        }
+
+        private Connection createConnectionProxy(final Connection 
targetConnection) {
+            return (Connection) Proxy.newProxyInstance(
+                    Connection.class.getClassLoader(),
+                    new Class[] {Connection.class},
+                    (proxy, method, args) -> {
+                        final Object result = invokeTarget(targetConnection, 
method, args);
+                        if ("createSession".equals(method.getName())) {
+                            return createSessionProxy((Session) result);
+                        }
+                        return result;
+                    });
+        }
+
+        private Session createSessionProxy(final Session targetSession) {
+            return (Session) Proxy.newProxyInstance(
+                    Session.class.getClassLoader(),
+                    new Class[] {Session.class},
+                    (proxy, method, args) -> {
+                        final Object result = invokeTarget(targetSession, 
method, args);
+                        if ("createProducer".equals(method.getName())) {
+                            return createProducerProxy((MessageProducer) 
result);
+                        }
+                        return result;
+                    });
+        }
+
+        private MessageProducer createProducerProxy(final MessageProducer 
targetProducer) {
+            return (MessageProducer) Proxy.newProxyInstance(
+                    MessageProducer.class.getClassLoader(),
+                    new Class[] {MessageProducer.class},
+                    (proxy, method, args) -> {
+                        if ("send".equals(method.getName())) {
+                            sendEntered.countDown();
+                            assertTrue(allowSend.await(5, TimeUnit.SECONDS));
+                            try {
+                                return invokeTarget(targetProducer, method, 
args);
+                            } finally {
+                                sendCompleted.countDown();
+                            }
+                        }
+                        return invokeTarget(targetProducer, method, args);
+                    });
+        }
+    }
+
     private static ArrayNode createTestJsonInput() {
         final ObjectMapper mapper = new ObjectMapper();
 

Reply via email to