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

Abacn pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git


The following commit(s) were added to refs/heads/master by this push:
     new 39c60f087fe Support ACKNOWLEDGE_MODE for JmsIO (#39253)
39c60f087fe is described below

commit 39c60f087fe74a119823d5133cac8263970c9600
Author: Yi Hu <[email protected]>
AuthorDate: Mon Jul 20 13:11:06 2026 -0400

    Support ACKNOWLEDGE_MODE for JmsIO (#39253)
    
    * Support ACKNOWLEDGEMENT_MODE for JmsIO
    
    * CheckpointMark behavior in alignment with different ACKNOWLEDGEMENT_MODE
    
    * Ref count active checkpoint for quicker onClose that releases session
    
    * Fix hanging checkpoint when no incoming data in direct runner. This
      allows us to do an exact assert
    
    * Optimize long running unit test usign a short retry
    
    * Re-enable AMQP integration test after stuck unack messages resolved
    
    * Address comments
    
    * Loose withCloseTimeout to fix flaky integration test; fix label
    
    * Fix javadoc
---
 CHANGES.md                                         |   1 +
 .../direct/UnboundedReadEvaluatorFactory.java      |  13 +-
 .../apache/beam/sdk/io/jms/JmsCheckpointMark.java  | 113 ++++++++---
 .../java/org/apache/beam/sdk/io/jms/JmsIO.java     | 194 ++++++++++++++++---
 .../java/org/apache/beam/sdk/io/jms/CommonJms.java |  27 ++-
 .../java/org/apache/beam/sdk/io/jms/JmsIOIT.java   | 209 ++++++++++++++-------
 .../java/org/apache/beam/sdk/io/jms/JmsIOTest.java | 184 +++++++++++++++---
 7 files changed, 586 insertions(+), 155 deletions(-)

diff --git a/CHANGES.md b/CHANGES.md
index a6fd20ad34d..19a16e1dcae 100644
--- a/CHANGES.md
+++ b/CHANGES.md
@@ -69,6 +69,7 @@
 ## New Features / Improvements
 
 * (Python) Removed the `envoy-data-plane` (and transitive `betterproto`) 
dependency; `EnvoyRateLimiter` now uses a small vendored protobuf definition 
instead, resolving dependency conflicts for downstream projects 
([#37854](https://github.com/apache/beam/issues/37854)).
+* (Java) Supported acknowledge mode for JmsIO 
([#39253](https://github.com/apache/beam/issues/39253)).
 * X feature added (Java/Python) 
([#X](https://github.com/apache/beam/issues/X)).
 
 ## Breaking Changes
diff --git 
a/runners/direct-java/src/main/java/org/apache/beam/runners/direct/UnboundedReadEvaluatorFactory.java
 
b/runners/direct-java/src/main/java/org/apache/beam/runners/direct/UnboundedReadEvaluatorFactory.java
index e45990a760e..c4a4fe86bf9 100644
--- 
a/runners/direct-java/src/main/java/org/apache/beam/runners/direct/UnboundedReadEvaluatorFactory.java
+++ 
b/runners/direct-java/src/main/java/org/apache/beam/runners/direct/UnboundedReadEvaluatorFactory.java
@@ -173,16 +173,17 @@ class UnboundedReadEvaluatorFactory implements 
TransformEvaluatorFactory {
         } else {
           Instant watermark = reader.getWatermark();
           if (watermark.isBefore(BoundedWindow.TIMESTAMP_MAX_VALUE)) {
-            // If the reader had no elements available, but the shard is not 
done, reuse it later
-            // Might be better to finalize old checkpoint.
+            // If the reader had no elements available, but the shard is not 
done, reuse it later.
+            // Finalize old checkpoint now.
+            final CheckpointMarkT checkpoint = shard.getCheckpoint();
+            if (checkpoint != null) {
+              checkpoint.finalizeCheckpoint();
+            }
             resultBuilder.addUnprocessedElements(
                 Collections.<WindowedValue<?>>singleton(
                     WindowedValues.timestampedValueInGlobalWindow(
                         UnboundedSourceShard.of(
-                            shard.getSource(),
-                            shard.getDeduplicator(),
-                            reader,
-                            shard.getCheckpoint()),
+                            shard.getSource(), shard.getDeduplicator(), 
reader, null),
                         watermark)));
           } else {
             // End of input. Close the reader after finalizing old checkpoint.
diff --git 
a/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsCheckpointMark.java
 
b/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsCheckpointMark.java
index 276d61ae5ef..3a58d28c1d9 100644
--- 
a/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsCheckpointMark.java
+++ 
b/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsCheckpointMark.java
@@ -19,13 +19,17 @@ package org.apache.beam.sdk.io.jms;
 
 import java.io.IOException;
 import java.io.Serializable;
+import java.util.ArrayList;
+import java.util.List;
 import java.util.Objects;
+import java.util.concurrent.atomic.AtomicInteger;
 import java.util.concurrent.locks.ReentrantReadWriteLock;
 import javax.jms.JMSException;
 import javax.jms.Message;
 import javax.jms.MessageConsumer;
 import javax.jms.Session;
 import org.apache.beam.sdk.io.UnboundedSource;
+import org.apache.beam.sdk.io.jms.JmsIO.AcknowledgeMode;
 import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.annotations.VisibleForTesting;
 import org.checkerframework.checker.nullness.qual.Nullable;
 import org.joda.time.Instant;
@@ -41,29 +45,32 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
   private static final Logger LOG = 
LoggerFactory.getLogger(JmsCheckpointMark.class);
 
   private Instant oldestMessageTimestamp;
-  private transient @Nullable Message lastMessage;
+  private transient @Nullable List<Message> messages;
   private transient @Nullable MessageConsumer consumer;
   private transient @Nullable Session session;
+  private transient @Nullable AtomicInteger activeCheckpoints;
 
   private JmsCheckpointMark(
       Instant oldestMessageTimestamp,
-      @Nullable Message lastMessage,
+      @Nullable List<Message> messages,
       @Nullable MessageConsumer consumer,
-      @Nullable Session session) {
+      @Nullable Session session,
+      @Nullable AtomicInteger activeCheckpoints) {
     this.oldestMessageTimestamp = oldestMessageTimestamp;
-    this.lastMessage = lastMessage;
+    this.messages = messages;
     this.consumer = consumer;
     this.session = session;
+    this.activeCheckpoints = activeCheckpoints;
   }
 
   /** Acknowledge all outstanding message. */
   @Override
   public void finalizeCheckpoint() {
     try {
-      // Jms spec will implicitly acknowledge _all_ messaged already received 
by the same
-      // session if one message in this session is being acknowledged.
-      if (lastMessage != null) {
-        lastMessage.acknowledge();
+      if (messages != null) {
+        for (Message message : messages) {
+          message.acknowledge();
+        }
       }
     } catch (JMSException e) {
       // The effect of this is message not get acknowledged and thus will be 
redelivered. It is
@@ -93,14 +100,37 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
         LOG.info("Error closing JMS session. It may have already been 
closed.");
       }
     }
+
+    if (activeCheckpoints != null) {
+      activeCheckpoints.decrementAndGet();
+    }
+  }
+
+  @VisibleForTesting
+  @Nullable
+  List<Message> getMessages() {
+    return messages;
+  }
+
+  @VisibleForTesting
+  @Nullable
+  Session getSession() {
+    return session;
+  }
+
+  @VisibleForTesting
+  @Nullable
+  MessageConsumer getConsumer() {
+    return consumer;
   }
 
   // set an empty list to messages when deserialize
   private void readObject(java.io.ObjectInputStream stream)
       throws IOException, ClassNotFoundException {
     stream.defaultReadObject();
-    lastMessage = null;
+    messages = null;
     session = null;
+    consumer = null;
   }
 
   @Override
@@ -120,8 +150,8 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
     return Objects.hash(oldestMessageTimestamp);
   }
 
-  static Preparer newPreparer() {
-    return new Preparer();
+  static Preparer newPreparer(AcknowledgeMode acknowledgeMode) {
+    return new Preparer(acknowledgeMode);
   }
 
   /**
@@ -129,15 +159,18 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
    */
   static class Preparer {
     private Instant oldestMessageTimestamp = Instant.now();
-    private transient @Nullable Message lastMessage = null;
+    private transient List<Message> messages = new ArrayList<>();
+    private final AcknowledgeMode acknowledgeMode;
 
     @VisibleForTesting transient boolean discarded = false;
 
     @VisibleForTesting final ReentrantReadWriteLock lock = new 
ReentrantReadWriteLock();
 
-    private Preparer() {}
+    private Preparer(AcknowledgeMode acknowledgeMode) {
+      this.acknowledgeMode = acknowledgeMode;
+    }
 
-    void add(Message message) throws Exception {
+    void add(Message message) throws JMSException {
       lock.writeLock().lock();
       try {
         if (discarded) {
@@ -149,7 +182,18 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
         if (currentMessageTimestamp.isBefore(oldestMessageTimestamp)) {
           oldestMessageTimestamp = currentMessageTimestamp;
         }
-        lastMessage = message;
+        if (acknowledgeMode == AcknowledgeMode.INDIVIDUAL_ACKNOWLEDGE) {
+          messages.add(message);
+        } else {
+          // Jms spec will implicitly acknowledge _all_ messaged already 
received by the same
+          // session if one message in this session is being acknowledged. 
Only need to ack
+          // last seen one.
+          if (messages.isEmpty()) {
+            messages.add(message);
+          } else {
+            messages.set(0, message);
+          }
+        }
       } finally {
         lock.writeLock().unlock();
       }
@@ -167,6 +211,7 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
     void discard() {
       lock.writeLock().lock();
       try {
+        messages.clear();
         this.discarded = true;
       } finally {
         lock.writeLock().unlock();
@@ -175,21 +220,43 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
 
     /**
      * Create a new checkpoint mark based on the current preparer. This will 
reset the messages held
-     * by the preparer, and the owner of the preparer is responsible to create 
a new Jms session
-     * after this call.
+     * by the preparer. If AcknowledgeMode is CLIENT_ACKNOWLEDGE, the owner of 
the preparer is
+     * responsible to create a new Jms session after this call.
      */
-    JmsCheckpointMark newCheckpoint(@Nullable MessageConsumer consumer, 
@Nullable Session session) {
+    JmsCheckpointMark newCheckpoint(
+        @Nullable MessageConsumer consumer,
+        @Nullable Session session,
+        @Nullable AcknowledgeMode acknowledgeMode,
+        @Nullable AtomicInteger activeCheckpoints) {
       JmsCheckpointMark checkpointMark;
       lock.writeLock().lock();
       try {
         if (discarded) {
-          lastMessage = null;
+          messages.clear();
           checkpointMark = this.emptyCheckpoint();
         } else {
+          List<Message> messagesCopy = null;
+          MessageConsumer consumerToPass = null;
+          Session sessionToPass = null;
+          if (!messages.isEmpty()) {
+            messagesCopy = new ArrayList<>(messages);
+          }
+          if (acknowledgeMode == AcknowledgeMode.CLIENT_ACKNOWLEDGE) {
+            consumerToPass = consumer;
+            sessionToPass = session;
+          }
           checkpointMark =
-              new JmsCheckpointMark(oldestMessageTimestamp, lastMessage, 
consumer, session);
-          lastMessage = null;
+              new JmsCheckpointMark(
+                  oldestMessageTimestamp,
+                  messagesCopy,
+                  consumerToPass,
+                  sessionToPass,
+                  activeCheckpoints);
+          messages.clear();
           oldestMessageTimestamp = Instant.now();
+          if (activeCheckpoints != null) {
+            activeCheckpoints.incrementAndGet();
+          }
         }
       } finally {
         lock.writeLock().unlock();
@@ -198,11 +265,11 @@ class JmsCheckpointMark implements 
UnboundedSource.CheckpointMark, Serializable
     }
 
     JmsCheckpointMark emptyCheckpoint() {
-      return new JmsCheckpointMark(oldestMessageTimestamp, null, null, null);
+      return new JmsCheckpointMark(oldestMessageTimestamp, null, null, null, 
null);
     }
 
     boolean isEmpty() {
-      return lastMessage == null;
+      return messages.isEmpty();
     }
   }
 }
diff --git 
a/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsIO.java 
b/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsIO.java
index 5deb7a2b62c..8926d584a13 100644
--- a/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsIO.java
+++ b/sdks/java/io/jms/src/main/java/org/apache/beam/sdk/io/jms/JmsIO.java
@@ -35,6 +35,7 @@ import java.util.Optional;
 import java.util.UUID;
 import java.util.concurrent.ScheduledExecutorService;
 import java.util.concurrent.TimeUnit;
+import java.util.concurrent.atomic.AtomicInteger;
 import java.util.stream.Stream;
 import javax.jms.Connection;
 import javax.jms.ConnectionFactory;
@@ -118,6 +119,22 @@ import org.slf4j.LoggerFactory;
  *
  * }</pre>
  *
+ * <h4>Acknowledgment Modes and Client Prefetch Configuration</h4>
+ *
+ * <p>By default, {@link JmsIO} consumes messages using {@link 
AcknowledgeMode#CLIENT_ACKNOWLEDGE}
+ * where a new {@link javax.jms.Session} is created for each checkpoint to 
prevent premature
+ * acknowledgments across bundles. When using {@link 
AcknowledgeMode#CLIENT_ACKNOWLEDGE}, if your
+ * JMS broker or client library utilizes client-side message prefetch buffers 
(such as Apache
+ * ActiveMQ), you should configure {@code prefetch=0} on your {@link 
javax.jms.ConnectionFactory}
+ * (e.g., via {@code ?jms.prefetchPolicy.all=0} in the broker URL or {@code
+ * ActiveMQPrefetchPolicy.setAll(0)}). Otherwise, unconsumed messages could be 
held inside old
+ * consumers in low throughput scenario and could lead to message backlog.
+ *
+ * <p>Alternatively, if your JMS broker supports individual message 
acknowledgment (such as ActiveMQ
+ * or Amazon MQ {@code ActiveMQSession.INDIVIDUAL_ACKNOWLEDGE = 4}), you can 
specify {@link
+ * Read#withAcknowledgeMode(AcknowledgeMode)} with {@link 
AcknowledgeMode#INDIVIDUAL_ACKNOWLEDGE}.
+ * In this mode, a single shared session and consumer are reused across all 
checkpoints.
+ *
  * <h3>Writing to a JMS destination</h3>
  *
  * <p>JmsIO sink supports writing text messages to a JMS destination on a 
broker. To configure a JMS
@@ -146,10 +163,11 @@ public class JmsIO {
         .setCoder(SerializableCoder.of(JmsRecord.class))
         .setCloseTimeout(DEFAULT_CLOSE_TIMEOUT)
         .setRequiresDeduping(false)
+        .setAcknowledgeMode(AcknowledgeMode.CLIENT_ACKNOWLEDGE)
         .setMessageMapper(
             new MessageMapper<JmsRecord>() {
               @Override
-              public JmsRecord mapMessage(Message message) throws Exception {
+              public JmsRecord mapMessage(Message message) throws JMSException 
{
                 TextMessage textMessage = (TextMessage) message;
                 Map<String, Object> properties = new HashMap<>();
                 @SuppressWarnings("rawtypes")
@@ -182,6 +200,7 @@ public class JmsIO {
         .setMaxNumRecords(Long.MAX_VALUE)
         .setCloseTimeout(DEFAULT_CLOSE_TIMEOUT)
         .setRequiresDeduping(false)
+        .setAcknowledgeMode(AcknowledgeMode.CLIENT_ACKNOWLEDGE)
         .build();
   }
 
@@ -261,6 +280,10 @@ public class JmsIO {
 
     abstract boolean isRequiresDeduping();
 
+    abstract AcknowledgeMode getAcknowledgeMode();
+
+    abstract @Nullable Integer getIndividualAcknowledgeModeCode();
+
     abstract Builder<T> builder();
 
     @AutoValue.Builder
@@ -292,6 +315,10 @@ public class JmsIO {
 
       abstract Builder<T> setRequiresDeduping(boolean requiresDeduping);
 
+      abstract Builder<T> setAcknowledgeMode(AcknowledgeMode acknowledgeMode);
+
+      abstract Builder<T> setIndividualAcknowledgeModeCode(Integer 
individualAcknowledgeModeCode);
+
       abstract Read<T> build();
     }
 
@@ -484,6 +511,29 @@ public class JmsIO {
       return builder().setRequiresDeduping(true).build();
     }
 
+    /**
+     * Specify the {@link AcknowledgeMode} used for consuming and 
acknowledging JMS messages.
+     *
+     * <p>To use {@link AcknowledgeMode#INDIVIDUAL_ACKNOWLEDGE}, providers 
other than ActiveMQ,
+     * ActiveMQ Artemis, Qpid JMS, require configuring {@link 
#withIndividualAcknowledgeModeCode}
+     * explicitly.
+     */
+    public Read<T> withAcknowledgeMode(AcknowledgeMode acknowledgeMode) {
+      checkArgument(acknowledgeMode != null, "acknowledgeMode can not be 
null");
+      return builder().setAcknowledgeMode(acknowledgeMode).build();
+    }
+
+    /**
+     * Specify the custom integer code for individual message acknowledgment 
when using {@link
+     * AcknowledgeMode#INDIVIDUAL_ACKNOWLEDGE}.
+     *
+     * <p>Different JMS providers use different proprietary integer constants 
for individual
+     * acknowledgment (e.g., ActiveMQ uses 4, Qpid JMS / ActiveMQ Artemis use 
101).
+     */
+    public Read<T> withIndividualAcknowledgeModeCode(int 
individualAcknowledgeModeCode) {
+      return 
builder().setIndividualAcknowledgeModeCode(individualAcknowledgeModeCode).build();
+    }
+
     @Override
     public PCollection<T> expand(PBegin input) {
       checkArgument(
@@ -515,6 +565,7 @@ public class JmsIO {
       super.populateDisplayData(builder);
       builder.addIfNotNull(DisplayData.item("queue", getQueue()));
       builder.addIfNotNull(DisplayData.item("topic", getTopic()));
+      builder.add(DisplayData.item("acknowledgeMode", 
getAcknowledgeMode().name()));
     }
 
     
///////////////////////////////////////////////////////////////////////////////////////
@@ -530,6 +581,35 @@ public class JmsIO {
 
   private JmsIO() {}
 
+  public enum AcknowledgeMode {
+    /**
+     * Acknowledge on checkpoint finalization with session isolation across 
checkpoints. Due to that
+     * runners may or may not finalize checkpoint in timely minor, unacked 
messages in unclosed
+     * consumer could be stuck when source throughput is zero. When using this 
mode in low
+     * throughput use cases, client-side prefetch buffers should be disabled 
(e.g. set {@code
+     * prefetch=0} for ActiveMQ) in connection factory properties.
+     */
+    CLIENT_ACKNOWLEDGE,
+
+    /**
+     * CLIENT_ACKNOWLEDGE but without session isolation across checkpoints. 
Acknowledging a message
+     * implicitly acknowledge all messages received up to that point per JMS 
spec. Best for
+     * performance but not safe on worker crash or scaling down. Good when 
best effort delivery is
+     * acceptable.
+     */
+    CLIENT_ACKNOWLEDGE_UNSAFE,
+
+    /**
+     * Acknowledge messages individually on checkpoint finalization. 
Recommended for JMS providers
+     * that support individual message acknowledgment (e.g., ActiveMQ, Amazon 
MQ, Artemis, Qpid
+     * JMS).
+     */
+    INDIVIDUAL_ACKNOWLEDGE
+
+    // When adding new AcknowledgeMode enum, update 
getAckModeCode(AcknowledgeMode mode) to handle
+    // the new mode.
+  }
+
   /**
    * An interface used by {@link JmsIO.Read} for converting each jms {@link 
Message} into an element
    * of the resulting {@link PCollection}.
@@ -601,10 +681,13 @@ public class JmsIO {
     private byte[] currentID;
     private long receiveTimeoutMillis;
     private PipelineOptions options;
+    // Acknowlging messages need open consumer. Tracking active checkpoints 
allows delayed close of
+    // session and consumer.
+    private final AtomicInteger activeCheckpoints = new AtomicInteger(0);
 
     public UnboundedJmsReader(UnboundedJmsSource<T> source, PipelineOptions 
options) {
       this.source = source;
-      this.checkpointMarkPreparer = JmsCheckpointMark.newPreparer();
+      this.checkpointMarkPreparer = 
JmsCheckpointMark.newPreparer(source.spec.getAcknowledgeMode());
       this.currentMessage = null;
       this.currentID = EMPTY;
       this.options = options;
@@ -613,7 +696,8 @@ public class JmsIO {
     /** recreate session and consumer. */
     private synchronized void recreateSession() throws IOException {
       try {
-        this.session = this.connection.createSession(false, 
Session.CLIENT_ACKNOWLEDGE);
+        int ackMode = getAckModeCode(source.spec.getAcknowledgeMode());
+        this.session = this.connection.createSession(false, ackMode);
       } catch (Exception e) {
         throw new IOException("Error creating JMS session", e);
       }
@@ -634,6 +718,34 @@ public class JmsIO {
       }
     }
 
+    private int getAckModeCode(AcknowledgeMode mode) {
+      if (mode == AcknowledgeMode.CLIENT_ACKNOWLEDGE
+          || mode == AcknowledgeMode.CLIENT_ACKNOWLEDGE_UNSAFE) {
+        return Session.CLIENT_ACKNOWLEDGE;
+      } else if (mode == AcknowledgeMode.INDIVIDUAL_ACKNOWLEDGE) {
+        Integer configuredCode = 
source.spec.getIndividualAcknowledgeModeCode();
+        if (configuredCode != null) {
+          return configuredCode;
+        }
+        String connectionClassName = this.connection.getClass().getName();
+        if 
(connectionClassName.contains("org.apache.activemq.ActiveMQConnection")) {
+          return 4;
+        } else if (connectionClassName.contains("org.apache.qpid.jms")) {
+          return 101;
+        } else if 
(connectionClassName.contains("org.apache.activemq.artemis")) {
+          return 101;
+        } else {
+          throw new IllegalArgumentException(
+              String.format(
+                  "Unknown JMS provider '%s' for INDIVIDUAL_ACKNOWLEDGE. "
+                      + "Please specify the code explicitly via 
Read#withIndividualAcknowledgeModeCode(int).",
+                  connectionClassName));
+        }
+      } else {
+        throw new IllegalArgumentException(String.format("Unknown 
AcknowledgeMode: %s", mode));
+      }
+    }
+
     @Override
     public boolean start() throws IOException {
       Read<T> spec = source.spec;
@@ -752,16 +864,23 @@ public class JmsIO {
 
       MessageConsumer consumerToClose;
       Session sessionTofinalize;
+      AcknowledgeMode mode = source.spec.getAcknowledgeMode();
       synchronized (this) {
-        consumerToClose = consumer;
-        sessionTofinalize = session;
-      }
-      try {
-        recreateSession();
-      } catch (IOException e) {
-        throw new RuntimeException(e);
+        if (mode == AcknowledgeMode.CLIENT_ACKNOWLEDGE) {
+          consumerToClose = consumer;
+          sessionTofinalize = session;
+          try {
+            recreateSession();
+          } catch (IOException e) {
+            throw new RuntimeException(e);
+          }
+        } else {
+          consumerToClose = null;
+          sessionTofinalize = null;
+        }
       }
-      return checkpointMarkPreparer.newCheckpoint(consumerToClose, 
sessionTofinalize);
+      return checkpointMarkPreparer.newCheckpoint(
+          consumerToClose, sessionTofinalize, mode, activeCheckpoints);
     }
 
     @Override
@@ -783,21 +902,44 @@ public class JmsIO {
     private void doClose() {
       try {
         closeAutoscaler();
-        closeConsumer();
-        ScheduledExecutorService executorService =
-            options.as(ExecutorOptions.class).getScheduledExecutorService();
-        executorService.schedule(
-            () -> {
-              LOG.debug("Closing connection after delay {}", 
source.spec.getCloseTimeout());
-              // Discard the checkpoints and set the reader as inactive
-              checkpointMarkPreparer.discard();
-              closeSession();
-              closeConnection();
-            },
-            source.spec.getCloseTimeout().getMillis(),
-            TimeUnit.MILLISECONDS);
+        // Discard the checkpoints and set the reader as inactive
+        checkpointMarkPreparer.discard();
+        if (source.spec.getAcknowledgeMode() == 
AcknowledgeMode.CLIENT_ACKNOWLEDGE) {
+          // checkpointMark holds session in CLIENT_ACKNOWLEDGE mode. Therefore
+          // we can close consumer and session immediately.
+          closeConsumer();
+          closeSession();
+        }
+        if (activeCheckpoints.get() == 0) {
+          closeConsumer();
+          closeSession();
+          closeConnection();
+        } else {
+          ScheduledExecutorService executorService =
+              options.as(ExecutorOptions.class).getScheduledExecutorService();
+          long deadline = System.currentTimeMillis() + 
source.spec.getCloseTimeout().getMillis();
+          long pollInterval = 1L;
+          executorService.schedule(
+              new Runnable() {
+                @Override
+                public void run() {
+                  if (activeCheckpoints.get() == 0 || 
System.currentTimeMillis() >= deadline) {
+                    LOG.debug(
+                        "Closing connection after checkpoints finalized or 
timeout: {}",
+                        source.spec.getCloseTimeout());
+                    closeConsumer();
+                    closeSession();
+                    closeConnection();
+                  } else {
+                    executorService.schedule(this, pollInterval, 
TimeUnit.SECONDS);
+                  }
+                }
+              },
+              pollInterval,
+              TimeUnit.SECONDS);
+        }
       } catch (Exception e) {
-        LOG.error("Error closing reader", e);
+        LOG.warn("Error closing reader", e);
       }
     }
 
@@ -809,7 +951,7 @@ public class JmsIO {
           connection = null;
         }
       } catch (Exception e) {
-        LOG.error("Error closing connection", e);
+        LOG.warn("Error closing connection", e);
       }
     }
 
diff --git 
a/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/CommonJms.java 
b/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/CommonJms.java
index 1d1245e6877..16173a36248 100644
--- a/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/CommonJms.java
+++ b/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/CommonJms.java
@@ -26,7 +26,6 @@ import java.util.function.Supplier;
 import javax.jms.BytesMessage;
 import javax.jms.ConnectionFactory;
 import javax.jms.Message;
-import org.apache.activemq.ActiveMQConnectionFactory;
 import org.apache.activemq.broker.BrokerPlugin;
 import org.apache.activemq.broker.BrokerService;
 import org.apache.activemq.security.AuthenticationUser;
@@ -91,12 +90,16 @@ public class CommonJms implements Serializable {
     this.connectionFactoryClass = connectionFactoryClass;
   }
 
+  private boolean shouldAppendPort() {
+    return brokerPort != null && brokerPort != 0;
+  }
+
   void startBroker() throws Exception {
     broker = new BrokerService();
     broker.setUseJmx(false);
     broker.setPersistenceAdapter(new MemoryPersistenceAdapter());
     TransportFactory.registerTransportFactory("amqp", new 
AmqpTransportFactory());
-    if (connectionFactoryClass != ActiveMQConnectionFactory.class) {
+    if (shouldAppendPort()) {
       broker.addConnector(String.format("%s:%d?transport.transformer=jms", 
brokerUrl, brokerPort));
     } else {
       broker.addConnector(brokerUrl);
@@ -119,10 +122,18 @@ public class CommonJms implements Serializable {
     broker.waitUntilStarted();
   }
 
+  private String getBrokerUrlWithPort() {
+    if (shouldAppendPort()) {
+      return String.format("%s:%d", brokerUrl, brokerPort);
+    } else {
+      return brokerUrl;
+    }
+  }
+
   ConnectionFactory createConnectionFactory()
       throws NoSuchMethodException, InvocationTargetException, 
InstantiationException,
           IllegalAccessException {
-    return 
connectionFactoryClass.getConstructor(String.class).newInstance(brokerUrl);
+    return 
connectionFactoryClass.getConstructor(String.class).newInstance(getBrokerUrlWithPort());
   }
 
   ConnectionFactory createConnectionFactoryWithSyncAcksAndWithoutPrefetch()
@@ -130,13 +141,15 @@ public class CommonJms implements Serializable {
           IllegalAccessException {
     return connectionFactoryClass
         .getConstructor(String.class)
-        .newInstance(brokerUrl + BROKER_WITHOUT_PREFETCH_PARAM + 
forceAsyncAcksParam);
+        .newInstance(getBrokerUrlWithPort() + BROKER_WITHOUT_PREFETCH_PARAM + 
forceAsyncAcksParam);
   }
 
   void stopBroker() throws Exception {
-    broker.stop();
-    broker.waitUntilStopped();
-    broker = null;
+    if (broker != null) {
+      broker.stop();
+      broker.waitUntilStopped();
+      broker = null;
+    }
   }
 
   Class<? extends ConnectionFactory> getConnectionFactoryClass() {
diff --git 
a/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOIT.java 
b/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOIT.java
index 266d04342d1..3dbb20775f7 100644
--- a/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOIT.java
+++ b/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOIT.java
@@ -20,8 +20,8 @@ package org.apache.beam.sdk.io.jms;
 import static org.apache.beam.sdk.io.jms.CommonJms.PASSWORD;
 import static org.apache.beam.sdk.io.jms.CommonJms.QUEUE;
 import static org.apache.beam.sdk.io.jms.CommonJms.USERNAME;
+import static org.junit.Assert.assertEquals;
 import static org.junit.Assert.assertNotEquals;
-import static org.junit.Assert.assertTrue;
 
 import java.io.IOException;
 import java.io.Serializable;
@@ -29,8 +29,10 @@ import java.time.Instant;
 import java.util.Collection;
 import java.util.Enumeration;
 import java.util.HashSet;
+import java.util.Map;
 import java.util.Set;
 import java.util.UUID;
+import java.util.concurrent.ConcurrentHashMap;
 import java.util.function.Function;
 import javax.jms.Connection;
 import javax.jms.ConnectionFactory;
@@ -46,6 +48,7 @@ import org.apache.beam.sdk.coders.SerializableCoder;
 import org.apache.beam.sdk.io.GenerateSequence;
 import org.apache.beam.sdk.io.common.IOITHelper;
 import org.apache.beam.sdk.io.common.IOTestPipelineOptions;
+import org.apache.beam.sdk.io.common.NetworkTestHelper;
 import org.apache.beam.sdk.metrics.Counter;
 import org.apache.beam.sdk.metrics.Metrics;
 import org.apache.beam.sdk.options.Default;
@@ -63,6 +66,7 @@ import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.collect.Immuta
 import org.apache.qpid.jms.JmsConnectionFactory;
 import org.joda.time.Duration;
 import org.junit.After;
+import org.junit.AfterClass;
 import org.junit.Before;
 import org.junit.Rule;
 import org.junit.Test;
@@ -142,105 +146,174 @@ public class JmsIOIT implements Serializable {
   @Rule public transient TestPipeline pipelineWrite = TestPipeline.create();
   @Rule public transient TestPipeline pipelineRead = TestPipeline.create();
 
-  @Parameterized.Parameters(name = "with client class {3}")
+  @Parameterized.Parameters(name = "with client class {2}")
   public static Collection<Object[]> connectionFactories() {
     return ImmutableList.of(
-        new Object[] {
-          "vm://localhost", 5672, "jms.sendAcksAsync=false", 
ActiveMQConnectionFactory.class
-        });
-    // TODO(https://github.com/apache/beam/issues/26175) Test failure on 
direct runner due to
-    //  JmsIO read on amqp slow on CI (passed locally)
-    // new Object[] {
-    //   "amqp://localhost", 5672, "jms.forceAsyncAcks=false", 
JmsConnectionFactory.class
-    // });
+        new Object[] {"vm://localhost", "jms.sendAcksAsync=false", 
ActiveMQConnectionFactory.class},
+        new Object[] {"amqp://localhost", "jms.forceAsyncAcks=false", 
JmsConnectionFactory.class});
   }
 
-  private final CommonJms commonJms;
+  private static final Map<String, CommonJms> BROKERS = new 
ConcurrentHashMap<>();
+
+  private final String brokerUrl;
+  private final Integer brokerPort;
+  private final String forceAsyncAcksParam;
+  private final Class<? extends ConnectionFactory> connectionFactoryClassParam;
+  private CommonJms commonJms;
   private ConnectionFactory connectionFactory;
   private Class<? extends ConnectionFactory> connectionFactoryClass;
 
   public JmsIOIT(
       String brokerUrl,
-      Integer brokerPort,
       String forceAsyncAcksParam,
       Class<? extends ConnectionFactory> connectionFactoryClass) {
-    this.commonJms =
-        new CommonJms(
-            OPTIONS.isLocalJmsBrokerEnabled() ? brokerUrl : 
OPTIONS.getJmsBrokerHost(),
-            OPTIONS.isLocalJmsBrokerEnabled() ? brokerPort : 
OPTIONS.getJmsBrokerPort(),
-            forceAsyncAcksParam,
-            connectionFactoryClass);
+    this.brokerUrl = brokerUrl;
+    if (OPTIONS.isLocalJmsBrokerEnabled()) {
+      try {
+        this.brokerPort = NetworkTestHelper.getAvailableLocalPort();
+      } catch (IOException e) {
+        throw new RuntimeException("Failed to find available port", e);
+      }
+    } else {
+      this.brokerPort = OPTIONS.getJmsBrokerPort();
+    }
+    this.forceAsyncAcksParam = forceAsyncAcksParam;
+    this.connectionFactoryClassParam = connectionFactoryClass;
   }
 
   @Before
   public void setup() throws Exception {
     if (OPTIONS.isLocalJmsBrokerEnabled()) {
-      this.commonJms.startBroker();
-      connectionFactory = this.commonJms.createConnectionFactory();
-      connectionFactoryClass = this.commonJms.getConnectionFactoryClass();
-      // use a small number of record for local integration test
+      String key = brokerUrl + ":" + connectionFactoryClassParam.getName();
+      commonJms =
+          BROKERS.computeIfAbsent(
+              key,
+              k -> {
+                CommonJms broker =
+                    new CommonJms(
+                        brokerUrl, brokerPort, forceAsyncAcksParam, 
connectionFactoryClassParam);
+                try {
+                  broker.startBroker();
+                } catch (Exception e) {
+                  throw new RuntimeException(e);
+                }
+                return broker;
+              });
       OPTIONS.setNumberOfRecords(10000);
+    } else {
+      commonJms =
+          new CommonJms(
+              OPTIONS.getJmsBrokerHost(),
+              OPTIONS.getJmsBrokerPort(),
+              forceAsyncAcksParam,
+              connectionFactoryClassParam);
     }
   }
 
+  @AfterClass
+  public static void afterClass() throws Exception {
+    for (CommonJms broker : BROKERS.values()) {
+      try {
+        broker.stopBroker();
+      } catch (Exception e) {
+        // ignore errors on shutdown
+      }
+    }
+    BROKERS.clear();
+  }
+
+  private void setupConnection(JmsIO.AcknowledgeMode acknowledgeMode) throws 
Exception {
+    if (acknowledgeMode == JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE) {
+      connectionFactory = 
this.commonJms.createConnectionFactoryWithSyncAcksAndWithoutPrefetch();
+    } else {
+      connectionFactory = this.commonJms.createConnectionFactory();
+    }
+    connectionFactoryClass = this.commonJms.getConnectionFactoryClass();
+  }
+
   @After
   public void tearDown() throws Exception {
-    if (OPTIONS.isLocalJmsBrokerEnabled()) {
-      this.commonJms.stopBroker();
-      connectionFactory = null;
-      connectionFactoryClass = null;
-    }
+    connectionFactory = null;
+    connectionFactoryClass = null;
+  }
+
+  @Test
+  public void testPublishingThenReadingAll() throws Exception {
+    runPublishingThenReadingAll(JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE);
+  }
+
+  @Test
+  public void testPublishingThenReadingAllIndividualAcknowledge() throws 
Exception {
+    runPublishingThenReadingAll(JmsIO.AcknowledgeMode.INDIVIDUAL_ACKNOWLEDGE);
   }
 
   @Test
-  public void testPublishingThenReadingAll() throws IOException, JMSException {
-    PipelineResult writeResult = publishingMessages();
+  public void testPublishingThenReadingAllClientAcknowledgeUnsafe() throws 
Exception {
+    
runPublishingThenReadingAll(JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE_UNSAFE);
+  }
+
+  private void runPublishingThenReadingAll(JmsIO.AcknowledgeMode 
acknowledgeMode) throws Exception {
+    setupConnection(acknowledgeMode);
+    String queue = QUEUE + "_" + acknowledgeMode.name();
+    PipelineResult writeResult = publishingMessages(queue);
     PipelineResult.State writeState = writeResult.waitUntilFinish();
     assertNotEquals(PipelineResult.State.FAILED, writeState);
 
-    PipelineResult readResult = readMessages();
-    PipelineResult.State readState =
-        
readResult.waitUntilFinish(Duration.standardSeconds(OPTIONS.getReadTimeout()));
-    // A workaround to stop the pipeline for waiting for too long
+    PipelineResult readResult = readMessages(acknowledgeMode, queue);
+    MetricsReader metricsReader = new MetricsReader(readResult, NAMESPACE);
+    long startTime = System.currentTimeMillis();
+    long timeoutMillis = OPTIONS.getReadTimeout() * 1000L;
+    PipelineResult.State readState = readResult.getState();
+    while (System.currentTimeMillis() - startTime < timeoutMillis
+        && (readState == null || !readState.isTerminal())) {
+      Thread.sleep(500);
+      readState = readResult.getState();
+      if (readState != null && readState.isTerminal()) {
+        break;
+      }
+      long actualRecords = 
metricsReader.getCounterMetric(READ_ELEMENT_METRIC_NAME);
+      if (actualRecords >= OPTIONS.getNumberOfRecords()) {
+        int unackRecords = countRemain(queue);
+        if (unackRecords == 0) {
+          readResult.cancel();
+          readState = readResult.getState();
+          break;
+        }
+      }
+    }
     cancelIfTimeouted(readResult, readState);
     assertNotEquals(PipelineResult.State.FAILED, readState);
 
-    MetricsReader metricsReader = new MetricsReader(readResult, NAMESPACE);
     long actualRecords = 
metricsReader.getCounterMetric(READ_ELEMENT_METRIC_NAME);
-
-    // TODO(yathu) resolve pending messages with direct runner then we can 
simply assert
-    //   actual-records == total-records.
-    //   Due to direct runner only finalize checkpoint at very end, there are 
open consumers (may
-    //   with buffer) and O(open_consumer) message won't get delivered to 
other session.
-    int unackRecords = countRemain(QUEUE);
-    assertTrue(
-        String.format("Too many unacknowledged messages: %d", unackRecords),
-        unackRecords < OPTIONS.getNumberOfRecords() * 0.003);
-
-    // acknowledged records
-    int ackRecords = OPTIONS.getNumberOfRecords() - unackRecords;
-    assertTrue(
-        String.format(
-            "actual number of records %d smaller than expected: %d.", 
actualRecords, ackRecords),
-        ackRecords <= actualRecords);
+    int unackRecords = countRemain(queue);
+    assertEquals("All messages should be acknowledged", 0, unackRecords);
+    assertEquals("All records should be read", (long) 
OPTIONS.getNumberOfRecords(), actualRecords);
     collectAndPublishMetrics(writeResult, readResult);
   }
 
   private void cancelIfTimeouted(PipelineResult readResult, 
PipelineResult.State readState)
       throws IOException {
-    if (readState == null) {
+    if (readState == null || !readState.isTerminal()) {
       readResult.cancel();
     }
   }
 
-  private PipelineResult readMessages() {
+  private PipelineResult readMessages(JmsIO.AcknowledgeMode acknowledgeMode, 
String queue) {
     pipelineRead.getOptions().as(JmsIOITOptions.class).setStreaming(true);
     pipelineRead.getOptions().as(JmsIOITOptions.class).setBlockOnRun(false);
-    JmsIO.Read<String> jmsIORead = JmsIO.readMessage();
+    JmsIO.Read<String> jmsIORead =
+        JmsIO.<String>readMessage()
+            .withAcknowledgeMode(acknowledgeMode)
+            // Decrease withCloseTimeout to be smaller than pipeline timeout. 
Direct runner randomly
+            // closes reader causing cached pending consumer hanging until 
closeTimeout
+            .withCloseTimeout(Duration.standardSeconds(10));
     if 
(pipelineRead.getOptions().as(JmsIOITOptions.class).getUseConnectionFactoryProviderFn())
 {
       jmsIORead =
           jmsIORead.withConnectionFactoryProviderFn(
-              
CommonJms.toSerializableFunction(commonJms::createConnectionFactory));
+              CommonJms.toSerializableFunction(
+                  acknowledgeMode == JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE
+                      ? 
commonJms::createConnectionFactoryWithSyncAcksAndWithoutPrefetch
+                      : commonJms::createConnectionFactory));
     } else {
       jmsIORead = jmsIORead.withConnectionFactory(connectionFactory);
     }
@@ -248,7 +321,7 @@ public class JmsIOIT implements Serializable {
         .apply(
             "Read Messages",
             jmsIORead
-                .withQueue(QUEUE)
+                .withQueue(queue)
                 .withUsername(USERNAME)
                 .withPassword(PASSWORD)
                 .withCoder(SerializableCoder.of(String.class))
@@ -258,7 +331,7 @@ public class JmsIOIT implements Serializable {
     return pipelineRead.run();
   }
 
-  private PipelineResult publishingMessages() {
+  private PipelineResult publishingMessages(String queue) {
 
     JmsIO.Write<String> jmsIOWrite = JmsIO.write();
     if 
(pipelineWrite.getOptions().as(JmsIOITOptions.class).getUseConnectionFactoryProviderFn())
 {
@@ -276,7 +349,7 @@ public class JmsIOIT implements Serializable {
         .apply(
             "Publish to Jms Broker",
             jmsIOWrite
-                .withQueue(QUEUE)
+                .withQueue(queue)
                 .withUsername(USERNAME)
                 .withPassword(PASSWORD)
                 .withValueMapper(new TextMessageMapper()));
@@ -318,17 +391,19 @@ public class JmsIOIT implements Serializable {
   }
 
   private int countRemain(String queue) throws JMSException {
-    Connection connection = connectionFactory.createConnection(USERNAME, 
PASSWORD);
-    connection.start();
-    Session session = connection.createSession(false, 
Session.AUTO_ACKNOWLEDGE);
-    QueueBrowser browser = session.createBrowser(session.createQueue(queue));
-    Enumeration<Message> messages = browser.getEnumeration();
-    int count = 0;
-    while (messages.hasMoreElements()) {
-      messages.nextElement();
-      count++;
+    try (Connection connection = connectionFactory.createConnection(USERNAME, 
PASSWORD)) {
+      connection.start();
+      try (Session session = connection.createSession(false, 
Session.CLIENT_ACKNOWLEDGE);
+          QueueBrowser browser = 
session.createBrowser(session.createQueue(queue))) {
+        Enumeration<Message> messages = browser.getEnumeration();
+        int count = 0;
+        while (messages.hasMoreElements()) {
+          messages.nextElement();
+          count++;
+        }
+        return count;
+      }
     }
-    return count;
   }
 
   static class ToString extends DoFn<Long, String> {
diff --git 
a/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOTest.java 
b/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOTest.java
index b3233f86617..eb6fb4faec0 100644
--- a/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOTest.java
+++ b/sdks/java/io/jms/src/test/java/org/apache/beam/sdk/io/jms/JmsIOTest.java
@@ -54,7 +54,6 @@ import static org.mockito.Mockito.when;
 import java.io.IOException;
 import java.io.NotSerializableException;
 import java.io.Serializable;
-import java.lang.reflect.InvocationTargetException;
 import java.lang.reflect.Proxy;
 import java.nio.ByteBuffer;
 import java.nio.charset.StandardCharsets;
@@ -90,6 +89,7 @@ import org.apache.beam.sdk.coders.SerializableCoder;
 import org.apache.beam.sdk.coders.StringUtf8Coder;
 import org.apache.beam.sdk.io.UnboundedSource;
 import org.apache.beam.sdk.io.UnboundedSource.CheckpointMark;
+import org.apache.beam.sdk.io.common.NetworkTestHelper;
 import org.apache.beam.sdk.io.jms.JmsIO.UnboundedJmsReader;
 import org.apache.beam.sdk.metrics.MetricNameFilter;
 import org.apache.beam.sdk.metrics.MetricQueryResults;
@@ -109,6 +109,7 @@ import 
org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Throwable
 import org.apache.qpid.jms.JmsAcknowledgeCallback;
 import org.apache.qpid.jms.JmsConnectionFactory;
 import org.apache.qpid.jms.message.JmsTextMessage;
+import org.checkerframework.checker.nullness.qual.Nullable;
 import org.joda.time.Duration;
 import org.junit.After;
 import org.junit.Before;
@@ -151,39 +152,42 @@ public class JmsIOTest {
       RetryConfiguration.create(1, Duration.standardSeconds(1), null);
   @Rule public final transient TestPipeline pipeline = TestPipeline.create();
 
-  @Parameterized.Parameters(name = "with client class {3}")
+  @Parameterized.Parameters(name = "with client class {2}")
   public static Collection<Object[]> connectionFactories() {
     return Arrays.asList(
-        new Object[] {
-          "vm://localhost", 5672, "jms.sendAcksAsync=false", 
ActiveMQConnectionFactory.class
-        },
-        new Object[] {
-          "amqp://localhost", 5672, "jms.forceAsyncAcks=false", 
JmsConnectionFactory.class
-        });
+        new Object[] {"vm://localhost", "jms.sendAcksAsync=false", 
ActiveMQConnectionFactory.class},
+        new Object[] {"amqp://localhost", "jms.forceAsyncAcks=false", 
JmsConnectionFactory.class});
   }
 
-  private final CommonJms commonJms;
-  private final ConnectionFactory connectionFactory;
+  private CommonJms commonJms;
+  private ConnectionFactory connectionFactory;
   private final Class<? extends ConnectionFactory> connectionFactoryClass;
-  private final ConnectionFactory 
connectionFactoryWithSyncAcksAndWithoutPrefetch;
+  private ConnectionFactory connectionFactoryWithSyncAcksAndWithoutPrefetch;
+  private final String brokerUrl;
+  private final Integer brokerPort;
+  private final String forceAsyncAcksParam;
 
   public JmsIOTest(
       String brokerUrl,
-      Integer brokerPort,
       String forceAsyncAcksParam,
-      Class<? extends ConnectionFactory> connectionFactoryClass)
-      throws InvocationTargetException, NoSuchMethodException, 
InstantiationException,
-          IllegalAccessException {
-    this.commonJms =
-        new CommonJms(brokerUrl, brokerPort, forceAsyncAcksParam, 
connectionFactoryClass);
+      Class<? extends ConnectionFactory> connectionFactoryClass) {
+    this.brokerUrl = brokerUrl;
+    this.forceAsyncAcksParam = forceAsyncAcksParam;
     this.connectionFactoryClass = connectionFactoryClass;
-    this.connectionFactory = commonJms.createConnectionFactory();
-    this.connectionFactoryWithSyncAcksAndWithoutPrefetch =
-        commonJms.createConnectionFactoryWithSyncAcksAndWithoutPrefetch();
+    try {
+      this.brokerPort = NetworkTestHelper.getAvailableLocalPort();
+    } catch (IOException e) {
+      throw new RuntimeException("Failed to find available port", e);
+    }
   }
 
   @Before
   public void beforeEach() throws Exception {
+    this.commonJms =
+        new CommonJms(brokerUrl, brokerPort, forceAsyncAcksParam, 
connectionFactoryClass);
+    this.connectionFactory = commonJms.createConnectionFactory();
+    this.connectionFactoryWithSyncAcksAndWithoutPrefetch =
+        commonJms.createConnectionFactoryWithSyncAcksAndWithoutPrefetch();
     this.commonJms.startBroker();
   }
 
@@ -610,6 +614,12 @@ public class JmsIOTest {
 
     // get checkpoint mark after consumed 4 messages
     CheckpointMark mark = reader.getCheckpointMark();
+    JmsCheckpointMark jmsMark = (JmsCheckpointMark) mark;
+    // In CLIENT_ACKNOWLEDGE mode, session/consumer are recreated on 
checkpoint:
+    assertNotNull(jmsMark.getConsumer());
+    assertNotNull(jmsMark.getSession());
+    assertNotNull(jmsMark.getMessages());
+    assertEquals(1, jmsMark.getMessages().size());
 
     // consume two more messages after checkpoint made
     reader.advance();
@@ -625,7 +635,119 @@ public class JmsIOTest {
     assertEquals(7, count(QUEUE));
   }
 
+  @Test
+  public void testCheckpointMarkAndFinalizeSeparatelyIndividualAcknowledge() 
throws Exception {
+    UnboundedJmsReader reader = 
setupReaderForTest(JmsIO.AcknowledgeMode.INDIVIDUAL_ACKNOWLEDGE);
+
+    assertTrue(reader.start());
+    assertTrue(advanceWithRetry(reader));
+    assertTrue(advanceWithRetry(reader));
+
+    CheckpointMark mark = reader.getCheckpointMark();
+    JmsCheckpointMark jmsMark = (JmsCheckpointMark) mark;
+    assertNull(jmsMark.getConsumer());
+    assertNull(jmsMark.getSession());
+    assertNotNull(jmsMark.getMessages());
+    assertEquals(3, jmsMark.getMessages().size());
+
+    reader.advance();
+    reader.advance();
+
+    assertEquals(10, count(QUEUE));
+    mark.finalizeCheckpoint();
+
+    // Verify only checkpointed messages are acknowledged
+    assertEquals(7, count(QUEUE));
+  }
+
+  @Test
+  public void testCheckpointMarkAndFinalizeSeparatelyClientAcknowledgeUnsafe() 
throws Exception {
+    UnboundedJmsReader reader = 
setupReaderForTest(JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE_UNSAFE);
+
+    assertTrue(reader.start());
+    assertTrue(advanceWithRetry(reader));
+    assertTrue(advanceWithRetry(reader));
+
+    CheckpointMark mark = reader.getCheckpointMark();
+    JmsCheckpointMark jmsMark = (JmsCheckpointMark) mark;
+    assertNull(jmsMark.getConsumer());
+    assertNull(jmsMark.getSession());
+    assertNotNull(jmsMark.getMessages());
+    assertEquals(1, jmsMark.getMessages().size());
+
+    reader.advance();
+    reader.advance();
+
+    assertEquals(10, count(QUEUE));
+    mark.finalizeCheckpoint();
+
+    // Verify all messages consumed on the session up to checkpoint are 
acknowledged
+    assertEquals(5, count(QUEUE));
+  }
+
+  @Test
+  public void testJmsCheckpointMarkIndividualAcknowledgeAllMessages() throws 
Exception {
+    Message msg1 = Mockito.mock(Message.class);
+    Message msg2 = Mockito.mock(Message.class);
+    Message msg3 = Mockito.mock(Message.class);
+
+    JmsCheckpointMark.Preparer preparer =
+        
JmsCheckpointMark.newPreparer(JmsIO.AcknowledgeMode.INDIVIDUAL_ACKNOWLEDGE);
+    preparer.add(msg1);
+    preparer.add(msg2);
+    preparer.add(msg3);
+
+    AtomicInteger activeCheckpoints = new AtomicInteger(0);
+    JmsCheckpointMark mark =
+        preparer.newCheckpoint(
+            null, null, JmsIO.AcknowledgeMode.INDIVIDUAL_ACKNOWLEDGE, 
activeCheckpoints);
+    assertNotNull(mark.getMessages());
+    assertEquals(3, mark.getMessages().size());
+    assertNull(mark.getConsumer());
+    assertNull(mark.getSession());
+    assertEquals(1, activeCheckpoints.get());
+
+    mark.finalizeCheckpoint();
+
+    Mockito.verify(msg1, Mockito.times(1)).acknowledge();
+    Mockito.verify(msg2, Mockito.times(1)).acknowledge();
+    Mockito.verify(msg3, Mockito.times(1)).acknowledge();
+    assertEquals(0, activeCheckpoints.get());
+  }
+
+  @Test
+  public void 
testJmsCheckpointMarkClientAcknowledgeUnsafeNoSessionRecreation() throws 
Exception {
+    Message msg1 = Mockito.mock(Message.class);
+    Message msg2 = Mockito.mock(Message.class);
+
+    JmsCheckpointMark.Preparer preparer =
+        
JmsCheckpointMark.newPreparer(JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE_UNSAFE);
+    preparer.add(msg1);
+    preparer.add(msg2);
+
+    AtomicInteger activeCheckpoints = new AtomicInteger(0);
+    JmsCheckpointMark mark =
+        preparer.newCheckpoint(
+            null, null, JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE_UNSAFE, 
activeCheckpoints);
+    assertNotNull(mark.getMessages());
+    assertEquals(1, mark.getMessages().size());
+    assertNull(mark.getConsumer());
+    assertNull(mark.getSession());
+    assertEquals(1, activeCheckpoints.get());
+
+    mark.finalizeCheckpoint();
+
+    Mockito.verify(msg2, Mockito.times(1)).acknowledge();
+    Mockito.verify(msg1, Mockito.never()).acknowledge();
+    assertEquals(0, activeCheckpoints.get());
+  }
+
   private JmsIO.UnboundedJmsReader setupReaderForTest() throws JMSException {
+    return setupReaderForTest(null);
+  }
+
+  private JmsIO.UnboundedJmsReader setupReaderForTest(
+      JmsIO.@Nullable AcknowledgeMode acknowledgeMode) throws JMSException {
     // we are using no prefetch here
     // prefetch is an ActiveMQ feature: to make efficient use of network 
resources the broker
     // utilizes a 'push' model to dispatch messages to consumers. However, in 
the case of our
@@ -652,6 +774,9 @@ public class JmsIOTest {
             .withUsername(USERNAME)
             .withPassword(PASSWORD)
             .withQueue(QUEUE);
+    if (acknowledgeMode != null) {
+      spec = spec.withAcknowledgeMode(acknowledgeMode);
+    }
     JmsIO.UnboundedJmsSource source = new JmsIO.UnboundedJmsSource(spec);
     JmsIO.UnboundedJmsReader reader = 
source.createReader(PipelineOptionsFactory.create(), null);
     return reader;
@@ -768,7 +893,9 @@ public class JmsIOTest {
   /** Test the checkpoint mark default coder, which is actually AvroCoder. */
   @Test
   public void testCheckpointMarkDefaultCoder() throws Exception {
-    JmsCheckpointMark jmsCheckpointMark = 
JmsCheckpointMark.newPreparer().newCheckpoint(null, null);
+    JmsCheckpointMark jmsCheckpointMark =
+        JmsCheckpointMark.newPreparer(JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE)
+            .newCheckpoint(null, null, 
JmsIO.AcknowledgeMode.CLIENT_ACKNOWLEDGE, null);
     Coder coder = new JmsIO.UnboundedJmsSource(null).getCheckpointMarkCoder();
     CoderProperties.coderSerializable(coder);
     CoderProperties.coderDecodeEncodeEqual(coder, jmsCheckpointMark);
@@ -819,7 +946,7 @@ public class JmsIOTest {
   }
 
   @Test
-  public void testCloseWithTimeout() throws IOException {
+  public void testCloseWithTimeout() throws IOException, JMSException {
     Duration closeTimeout = Duration.millis(2000L);
     JmsIO.Read spec =
         JmsIO.read()
@@ -843,10 +970,13 @@ public class JmsIOTest {
     JmsIO.UnboundedJmsReader reader = source.createReader(options, null);
     reader.start();
     assertFalse(getDiscardedValue(reader));
+    reader.checkpointMarkPreparer.add(Mockito.mock(Message.class));
+    CheckpointMark mark = reader.getCheckpointMark();
     reader.close();
-    assertFalse(getDiscardedValue(reader));
+    assertTrue(getDiscardedValue(reader));
     verify(mockScheduledExecutorService)
-        .schedule(any(Runnable.class), eq(closeTimeout.getMillis()), 
eq(TimeUnit.MILLISECONDS));
+        .schedule(any(Runnable.class), eq(1L), eq(TimeUnit.SECONDS));
+    mark.finalizeCheckpoint();
     runnableArgumentCaptor.getValue().run();
     assertTrue(getDiscardedValue(reader));
     verifyNoMoreInteractions(mockScheduledExecutorService);
@@ -982,7 +1112,8 @@ public class JmsIOTest {
     int maxPublicationAttempts = 2;
     List<String> data = Collections.singletonList(messageText);
     RetryConfiguration retryConfiguration =
-        RetryConfiguration.create(maxPublicationAttempts, null, null);
+        RetryConfiguration.create(
+            maxPublicationAttempts, Duration.standardSeconds(5), 
Duration.millis(10L));
 
     WriteJmsResult<String> output =
         pipeline
@@ -1039,7 +1170,8 @@ public class JmsIOTest {
     List<String> data = Arrays.asList("Message 1", "Message 2", "Message 3", 
"Message 4");
 
     RetryConfiguration retryConfiguration =
-        RetryConfiguration.create(maxPublicationAttempts, null, null);
+        RetryConfiguration.create(
+            maxPublicationAttempts, Duration.standardSeconds(5), 
Duration.millis(10L));
 
     WriteJmsResult<String> output =
         pipeline

Reply via email to