diff --git a/clients/src/main/java/org/apache/kafka/clients/consumer/internals/Fetcher.java b/clients/src/main/java/org/apache/kafka/clients/consumer/internals/Fetcher.java index 5c31ca1469378..159d236b8daf5 100644 --- a/clients/src/main/java/org/apache/kafka/clients/consumer/internals/Fetcher.java +++ b/clients/src/main/java/org/apache/kafka/clients/consumer/internals/Fetcher.java @@ -97,7 +97,21 @@ import static org.apache.kafka.common.serialization.ExtendedDeserializer.Wrapper.ensureExtended; /** - * This class manage the fetching process with the brokers. + * This class manages the fetching process with the brokers. + *

+ * Thread-safety: + * Requests and responses of Fetcher may be processed by different threads since heartbeat + * thread may process responses. Other operations are single-threaded and invoked only from + * the thread polling the consumer. + *

*/ public class Fetcher implements SubscriptionState.Listener, Closeable { private final Logger log; @@ -233,7 +247,7 @@ public void clearBufferedDataForPausedPartitions() { * an in-flight fetch or pending fetch data. * @return number of fetches sent */ - public int sendFetches() { + public synchronized int sendFetches() { Map fetchRequestMap = prepareFetchRequests(); for (Map.Entry entry : fetchRequestMap.entrySet()) { final Node fetchTarget = entry.getKey(); @@ -251,39 +265,43 @@ public int sendFetches() { .addListener(new RequestFutureListener() { @Override public void onSuccess(ClientResponse resp) { - FetchResponse response = (FetchResponse) resp.responseBody(); - FetchSessionHandler handler = sessionHandlers.get(fetchTarget.id()); - if (handler == null) { - log.error("Unable to find FetchSessionHandler for node {}. Ignoring fetch response.", - fetchTarget.id()); - return; + synchronized (Fetcher.this) { + FetchResponse response = (FetchResponse) resp.responseBody(); + FetchSessionHandler handler = sessionHandler(fetchTarget.id()); + if (handler == null) { + log.error("Unable to find FetchSessionHandler for node {}. Ignoring fetch response.", + fetchTarget.id()); + return; + } + if (!handler.handleResponse(response)) { + return; + } + + Set partitions = new HashSet<>(response.responseData().keySet()); + FetchResponseMetricAggregator metricAggregator = new FetchResponseMetricAggregator(sensors, partitions); + + for (Map.Entry> entry : response.responseData().entrySet()) { + TopicPartition partition = entry.getKey(); + long fetchOffset = data.sessionPartitions().get(partition).fetchOffset; + FetchResponse.PartitionData fetchData = entry.getValue(); + + log.debug("Fetch {} at offset {} for partition {} returned fetch data {}", + isolationLevel, fetchOffset, partition, fetchData); + completedFetches.add(new CompletedFetch(partition, fetchOffset, fetchData, metricAggregator, + resp.requestHeader().apiVersion())); + } + + sensors.fetchLatency.record(resp.requestLatencyMs()); } - if (!handler.handleResponse(response)) { - return; - } - - Set partitions = new HashSet<>(response.responseData().keySet()); - FetchResponseMetricAggregator metricAggregator = new FetchResponseMetricAggregator(sensors, partitions); - - for (Map.Entry> entry : response.responseData().entrySet()) { - TopicPartition partition = entry.getKey(); - long fetchOffset = data.sessionPartitions().get(partition).fetchOffset; - FetchResponse.PartitionData fetchData = entry.getValue(); - - log.debug("Fetch {} at offset {} for partition {} returned fetch data {}", - isolationLevel, fetchOffset, partition, fetchData); - completedFetches.add(new CompletedFetch(partition, fetchOffset, fetchData, metricAggregator, - resp.requestHeader().apiVersion())); - } - - sensors.fetchLatency.record(resp.requestLatencyMs()); } @Override public void onFailure(RuntimeException e) { - FetchSessionHandler handler = sessionHandlers.get(fetchTarget.id()); - if (handler != null) { - handler.handleError(e); + synchronized (Fetcher.this) { + FetchSessionHandler handler = sessionHandler(fetchTarget.id()); + if (handler != null) { + handler.handleError(e); + } } } }); @@ -935,7 +953,7 @@ private Map prepareFetchRequests() { // if there is a leader and no in-flight requests, issue a new fetch FetchSessionHandler.Builder builder = fetchable.get(node); if (builder == null) { - FetchSessionHandler handler = sessionHandlers.get(node.id()); + FetchSessionHandler handler = sessionHandler(node.id()); if (handler == null) { handler = new FetchSessionHandler(logContext, node.id()); sessionHandlers.put(node.id(), handler); @@ -1140,6 +1158,11 @@ public void filterUnassignedPartitions(Set assignedPartitions) { } } + // Visibilty for testing + protected FetchSessionHandler sessionHandler(int node) { + return sessionHandlers.get(node); + } + public static Sensor throttleTimeSensor(Metrics metrics, FetcherMetricsRegistry metricsRegistry) { Sensor fetchThrottleTimeSensor = metrics.sensor("fetch-throttle-time"); fetchThrottleTimeSensor.add(metrics.metricInstance(metricsRegistry.fetchThrottleTimeAvg), new Avg()); diff --git a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/FetcherTest.java b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/FetcherTest.java index a4df571c38c3a..88cd8584d19f2 100644 --- a/clients/src/test/java/org/apache/kafka/clients/consumer/internals/FetcherTest.java +++ b/clients/src/test/java/org/apache/kafka/clients/consumer/internals/FetcherTest.java @@ -19,6 +19,7 @@ import org.apache.kafka.clients.ApiVersions; import org.apache.kafka.clients.ClientRequest; import org.apache.kafka.clients.ClientUtils; +import org.apache.kafka.clients.FetchSessionHandler; import org.apache.kafka.clients.Metadata; import org.apache.kafka.clients.MockClient; import org.apache.kafka.clients.NetworkClient; @@ -88,6 +89,7 @@ import org.junit.Test; import java.io.DataOutputStream; +import java.lang.reflect.Field; import java.nio.ByteBuffer; import java.nio.charset.StandardCharsets; import java.util.ArrayList; @@ -100,6 +102,13 @@ import java.util.List; import java.util.Map; import java.util.Set; +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.function.Function; +import java.util.stream.Collectors; import static java.util.Arrays.asList; import static java.util.Collections.singleton; @@ -111,6 +120,7 @@ import static org.junit.Assert.assertTrue; import static org.junit.Assert.fail; + @SuppressWarnings("deprecation") public class FetcherTest { private ConsumerRebalanceListener listener = new NoOpConsumerRebalanceListener(); @@ -149,38 +159,30 @@ public class FetcherTest { private Fetcher fetcher = createFetcher(subscriptions, metrics); private Metrics fetcherMetrics = new Metrics(time); private Fetcher fetcherNoAutoReset = createFetcher(subscriptionsNoAutoReset, fetcherMetrics); + private ExecutorService executorService; @Before public void setup() throws Exception { metadata.update(cluster, Collections.emptySet(), time.milliseconds()); client.setNode(node); - MemoryRecordsBuilder builder = MemoryRecords.builder(ByteBuffer.allocate(1024), CompressionType.NONE, TimestampType.CREATE_TIME, 1L); - builder.append(0L, "key".getBytes(), "value-1".getBytes()); - builder.append(0L, "key".getBytes(), "value-2".getBytes()); - builder.append(0L, "key".getBytes(), "value-3".getBytes()); - records = builder.build(); - - builder = MemoryRecords.builder(ByteBuffer.allocate(1024), CompressionType.NONE, TimestampType.CREATE_TIME, 4L); - builder.append(0L, "key".getBytes(), "value-4".getBytes()); - builder.append(0L, "key".getBytes(), "value-5".getBytes()); - nextRecords = builder.build(); - - builder = MemoryRecords.builder(ByteBuffer.allocate(1024), CompressionType.NONE, TimestampType.CREATE_TIME, 0L); - emptyRecords = builder.build(); - - builder = MemoryRecords.builder(ByteBuffer.allocate(1024), CompressionType.NONE, TimestampType.CREATE_TIME, 4L); - builder.append(0L, "key".getBytes(), "value-0".getBytes()); - partialRecords = builder.build(); + records = buildRecords(1L, 3, 1); + nextRecords = buildRecords(4L, 2, 4); + emptyRecords = buildRecords(0L, 0, 0); + partialRecords = buildRecords(4L, 1, 0); partialRecords.buffer().putInt(Records.SIZE_OFFSET, 10000); } @After - public void teardown() { + public void teardown() throws Exception { this.metrics.close(); this.fetcherMetrics.close(); this.fetcher.close(); this.fetcherMetrics.close(); + if (executorService != null) { + executorService.shutdownNow(); + assertTrue(executorService.awaitTermination(5, TimeUnit.SECONDS)); + } } @Test @@ -2456,6 +2458,142 @@ public void testConsumingViaIncrementalFetchRequests() { assertEquals(5, records.get(1).offset()); } + @Test + public void testFetcherConcurrency() throws Exception { + int numPartitions = 20; + Set topicPartitions = new HashSet<>(); + for (int i = 0; i < numPartitions; i++) + topicPartitions.add(new TopicPartition(topicName, i)); + cluster = TestUtils.singletonCluster(topicName, numPartitions); + metadata.update(cluster, Collections.emptySet(), time.milliseconds()); + client.setNode(node); + fetchSize = 10000; + + Fetcher fetcher = new Fetcher( + new LogContext(), + consumerClient, + minBytes, + maxBytes, + maxWaitMs, + fetchSize, + 2 * numPartitions, + true, + false, + new ByteArrayDeserializer(), + new ByteArrayDeserializer(), + metadata, + subscriptions, + metrics, + metricsRegistry, + time, + retryBackoffMs, + requestTimeoutMs, + IsolationLevel.READ_UNCOMMITTED) { + @Override + protected FetchSessionHandler sessionHandler(int id) { + final FetchSessionHandler handler = super.sessionHandler(id); + if (handler == null) + return null; + else { + return new FetchSessionHandler(new LogContext(), id) { + @Override + public Builder newBuilder() { + verifySessionPartitions(); + return handler.newBuilder(); + } + + @Override + public boolean handleResponse(FetchResponse response) { + verifySessionPartitions(); + return handler.handleResponse(response); + } + + @Override + public void handleError(Throwable t) { + verifySessionPartitions(); + handler.handleError(t); + } + + // Verify that session partitions can be traversed safely. + private void verifySessionPartitions() { + try { + Field field = FetchSessionHandler.class.getDeclaredField("sessionPartitions"); + field.setAccessible(true); + LinkedHashMap sessionPartitions = + (LinkedHashMap) field.get(handler); + for (Map.Entry entry : sessionPartitions.entrySet()) { + // If `sessionPartitions` are modified on another thread, Thread.yield will increase the + // possibility of ConcurrentModificationException if appropriate synchronization is not used. + Thread.yield(); + } + } catch (Exception e) { + throw new RuntimeException(e); + } + } + }; + } + } + }; + + subscriptions.assignFromUser(topicPartitions); + topicPartitions.forEach(tp -> subscriptions.seek(tp, 0L)); + + AtomicInteger fetchesRemaining = new AtomicInteger(1000); + executorService = Executors.newSingleThreadExecutor(); + Future future = executorService.submit(() -> { + while (fetchesRemaining.get() > 0) { + synchronized (consumerClient) { + if (!client.requests().isEmpty()) { + ClientRequest request = client.requests().peek(); + FetchRequest fetchRequest = (FetchRequest) request.requestBuilder().build(); + LinkedHashMap> responseMap = new LinkedHashMap<>(); + for (Map.Entry entry : fetchRequest.fetchData().entrySet()) { + TopicPartition tp = entry.getKey(); + long offset = entry.getValue().fetchOffset; + responseMap.put(tp, new FetchResponse.PartitionData<>(Errors.NONE, offset + 2L, offset + 2, + 0L, null, buildRecords(offset, 2, offset))); + } + client.respondToRequest(request, new FetchResponse<>(Errors.NONE, responseMap, 0, 123)); + consumerClient.poll(time.timer(0)); + } + } + } + return fetchesRemaining.get(); + }); + Map nextFetchOffsets = topicPartitions.stream() + .collect(Collectors.toMap(Function.identity(), t -> 0L)); + while (fetchesRemaining.get() > 0 && !future.isDone()) { + if (fetcher.sendFetches() == 1) { + synchronized (consumerClient) { + consumerClient.poll(time.timer(0)); + } + } + if (fetcher.hasCompletedFetches()) { + Map>> fetchedRecords = fetcher.fetchedRecords(); + if (!fetchedRecords.isEmpty()) { + fetchesRemaining.decrementAndGet(); + fetchedRecords.entrySet().forEach(entry -> { + TopicPartition tp = entry.getKey(); + List> records = entry.getValue(); + assertEquals(2, records.size()); + long nextOffset = nextFetchOffsets.get(tp); + assertEquals(nextOffset, records.get(0).offset()); + assertEquals(nextOffset + 1, records.get(1).offset()); + nextFetchOffsets.put(tp, nextOffset + 2); + }); + } + } + } + assertEquals(0, future.get()); + } + + private MemoryRecords buildRecords(long baseOffset, int count, long firstMessageId) { + MemoryRecordsBuilder builder = MemoryRecords.builder(ByteBuffer.allocate(1024), CompressionType.NONE, TimestampType.CREATE_TIME, baseOffset); + for (int i = 0; i < count; i++) + builder.append(0L, "key".getBytes(), ("value-" + (firstMessageId + i)).getBytes()); + return builder.build(); + } + private int appendTransactionalRecords(ByteBuffer buffer, long pid, long baseOffset, int baseSequence, SimpleRecord... records) { MemoryRecordsBuilder builder = MemoryRecords.builder(buffer, RecordBatch.CURRENT_MAGIC_VALUE, CompressionType.NONE, TimestampType.CREATE_TIME, baseOffset, time.milliseconds(), pid, (short) 0, baseSequence, true,