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

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


The following commit(s) were added to refs/heads/master by this push:
     new 30f5d208a1e Prune values across the query executor instead of on the 
calling thread (#19513)
30f5d208a1e is described below

commit 30f5d208a1e58a1e1cee4afc5fc77a32b9e61f66
Author: Xiang Fu <[email protected]>
AuthorDate: Thu Sep 10 21:41:57 2026 -0700

    Prune values across the query executor instead of on the calling thread 
(#19513)
---
 .../query/pruner/BloomFilterSegmentPruner.java     |  99 +----------
 .../core/query/pruner/ValueBasedSegmentPruner.java |  90 +++++++++-
 .../query/pruner/BloomFilterSegmentPrunerTest.java | 136 +++++++++++++++
 .../query/pruner/ColumnValueSegmentPrunerTest.java | 189 +++++++++++++++++++++
 4 files changed, 411 insertions(+), 103 deletions(-)

diff --git 
a/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPruner.java
 
b/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPruner.java
index 44b7d99a784..623ad153baf 100644
--- 
a/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPruner.java
+++ 
b/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPruner.java
@@ -18,23 +18,18 @@
  */
 package org.apache.pinot.core.query.pruner;
 
-import java.util.ArrayList;
-import java.util.HashMap;
 import java.util.List;
 import java.util.Map;
-import java.util.Objects;
 import java.util.concurrent.ExecutorService;
 import java.util.function.Function;
 import javax.annotation.Nullable;
 import org.apache.pinot.common.request.context.ExpressionContext;
-import org.apache.pinot.common.request.context.FilterContext;
 import org.apache.pinot.common.request.context.predicate.EqPredicate;
 import org.apache.pinot.common.request.context.predicate.InPredicate;
 import org.apache.pinot.common.request.context.predicate.Predicate;
 import org.apache.pinot.core.query.prefetch.FetchPlanner;
 import org.apache.pinot.core.query.prefetch.FetchPlannerRegistry;
 import org.apache.pinot.core.query.request.context.QueryContext;
-import org.apache.pinot.core.util.QueryMultiThreadingUtils;
 import org.apache.pinot.segment.spi.FetchContext;
 import org.apache.pinot.segment.spi.ImmutableSegment;
 import org.apache.pinot.segment.spi.IndexSegment;
@@ -42,17 +37,10 @@ import org.apache.pinot.segment.spi.datasource.DataSource;
 import org.apache.pinot.segment.spi.datasource.DataSourceMetadata;
 import org.apache.pinot.segment.spi.index.reader.BloomFilterReader;
 import org.apache.pinot.spi.env.PinotConfiguration;
-import org.apache.pinot.spi.exception.QueryCancelledException;
 
 
-/// The `BloomFilterSegmentPruner` prunes segments based on bloom filter for 
EQUALITY filter. Because the access
-/// to bloom filter data is required, segment pruning is done in parallel when 
the number of segments is large.
+/// Prunes segments using bloom filters for EQ and IN predicates, with 
optional prefetch.
 public class BloomFilterSegmentPruner extends ValueBasedSegmentPruner {
-  // Try to schedule 10 segments for each thread, or evenly distribute them to 
all MAX_NUM_THREADS_PER_QUERY threads.
-  // TODO: make this threshold configurable? threshold 10 is also used in 
CombinePlanNode, which accesses the
-  //       dictionary data to do query planning and if segments are more than 
10, planning is done in parallel.
-  private static final int TARGET_NUM_SEGMENTS_PER_THREAD = 10;
-
   private FetchPlanner _fetchPlanner;
 
   @Override
@@ -83,80 +71,16 @@ public class BloomFilterSegmentPruner extends 
ValueBasedSegmentPruner {
 
   @Override
   public List<IndexSegment> prune(List<IndexSegment> segments, QueryContext 
query) {
-    if (segments.isEmpty()) {
-      return segments;
-    }
-    if (!query.isEnablePrefetch()) {
-      return super.prune(segments, query);
-    }
-    return prefetch(segments, query, fetchContexts -> {
-      int numSegments = segments.size();
-      FilterContext filter = Objects.requireNonNull(query.getFilter());
-      ValueCache cachedValues = new ValueCache();
-      Map<String, DataSource> dataSourceCache = new HashMap<>();
-      List<IndexSegment> selectedSegments = new ArrayList<>(numSegments);
-      for (int i = 0; i < numSegments; i++) {
-        dataSourceCache.clear();
-        IndexSegment segment = segments.get(i);
-        if (!pruneSegmentWithFetchContext(segment, fetchContexts[i], filter, 
dataSourceCache, cachedValues, query)) {
-          selectedSegments.add(segment);
-        }
-      }
-      return selectedSegments;
-    });
+    return prune(segments, query, null);
   }
 
   @Override
   public List<IndexSegment> prune(List<IndexSegment> segments, QueryContext 
query,
       @Nullable ExecutorService executorService) {
-    if (segments.isEmpty()) {
-      return segments;
-    }
-    if (executorService == null || segments.size() <= 
TARGET_NUM_SEGMENTS_PER_THREAD) {
-      // If executor is not provided, or the number of segments is small, 
prune them sequentially
-      return prune(segments, query);
-    }
-    // With executor service and large number of segments, prune them in 
parallel.
-    // NOTE: Even if numTasks=1 i.e. we get a single executor thread, still 
run it using a separate thread so that
-    //       the timeout can be honored. For example, this may happen when 
there is only one processor.
-    int numTasks = QueryMultiThreadingUtils.getNumTasks(segments.size(), 
TARGET_NUM_SEGMENTS_PER_THREAD,
-        query.getMaxExecutionThreads());
-    if (!query.isEnablePrefetch()) {
-      return pruneInParallel(numTasks, segments, query, executorService, null);
+    if (segments.isEmpty() || !query.isEnablePrefetch()) {
+      return super.prune(segments, query, executorService);
     }
-    return prefetch(segments, query,
-        fetchContexts -> pruneInParallel(numTasks, segments, query, 
executorService, fetchContexts));
-  }
-
-  private List<IndexSegment> pruneInParallel(int numTasks, List<IndexSegment> 
segments, QueryContext query,
-      ExecutorService executorService, @Nullable FetchContext[] fetchContexts) 
{
-    int numSegments = segments.size();
-    List<IndexSegment> allSelectedSegments = new ArrayList<>();
-    QueryMultiThreadingUtils.runTasksWithDeadline(numTasks, index -> {
-      FilterContext filter = Objects.requireNonNull(query.getFilter());
-      ValueCache cachedValues = new ValueCache();
-      Map<String, DataSource> dataSourceCache = new HashMap<>();
-      List<IndexSegment> selectedSegments = new ArrayList<>();
-      for (int i = index; i < numSegments; i += numTasks) {
-        dataSourceCache.clear();
-        IndexSegment segment = segments.get(i);
-        FetchContext fetchContext = fetchContexts != null ? fetchContexts[i] : 
null;
-        if (!pruneSegmentWithFetchContext(segment, fetchContext, filter, 
dataSourceCache, cachedValues, query)) {
-          selectedSegments.add(segment);
-        }
-      }
-      return selectedSegments;
-    }, taskRes -> {
-      if (taskRes != null) {
-        allSelectedSegments.addAll(taskRes);
-      }
-    }, e -> {
-      if (e instanceof InterruptedException) {
-        throw new QueryCancelledException("Cancelled while running 
BloomFilterSegmentPruner", e);
-      }
-      throw new RuntimeException("Caught exception while running 
BloomFilterSegmentPruner", e);
-    }, executorService, query.getEndTimeMs());
-    return allSelectedSegments;
+    return prefetch(segments, query, fetchContexts -> super.prune(segments, 
query, executorService, fetchContexts));
   }
 
   private List<IndexSegment> prefetch(List<IndexSegment> segments, 
QueryContext query,
@@ -185,19 +109,6 @@ public class BloomFilterSegmentPruner extends 
ValueBasedSegmentPruner {
     }
   }
 
-  private boolean pruneSegmentWithFetchContext(IndexSegment segment, @Nullable 
FetchContext fetchContext,
-      FilterContext filter, Map<String, DataSource> dataSourceCache, 
ValueCache cachedValues, QueryContext query) {
-    if (fetchContext == null) {
-      return pruneSegment(segment, filter, dataSourceCache, cachedValues, 
query);
-    }
-    segment.acquire(fetchContext);
-    try {
-      return pruneSegment(segment, filter, dataSourceCache, cachedValues, 
query);
-    } finally {
-      segment.release(fetchContext);
-    }
-  }
-
   @Override
   boolean pruneSegmentWithPredicate(IndexSegment segment, Predicate predicate, 
Map<String, DataSource> dataSourceCache,
       ValueCache cachedValues, QueryContext query) {
diff --git 
a/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/ValueBasedSegmentPruner.java
 
b/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/ValueBasedSegmentPruner.java
index f5f73d2cfcc..ca62cb6e40a 100644
--- 
a/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/ValueBasedSegmentPruner.java
+++ 
b/pinot-core/src/main/java/org/apache/pinot/core/query/pruner/ValueBasedSegmentPruner.java
@@ -24,6 +24,8 @@ import java.util.IdentityHashMap;
 import java.util.List;
 import java.util.Map;
 import java.util.Objects;
+import java.util.concurrent.ExecutorService;
+import javax.annotation.Nullable;
 import org.apache.pinot.common.request.context.ExpressionContext;
 import org.apache.pinot.common.request.context.FilterContext;
 import org.apache.pinot.common.request.context.predicate.EqPredicate;
@@ -31,13 +33,17 @@ import 
org.apache.pinot.common.request.context.predicate.InPredicate;
 import org.apache.pinot.common.request.context.predicate.Predicate;
 import org.apache.pinot.common.utils.config.QueryOptionsUtils;
 import org.apache.pinot.core.query.request.context.QueryContext;
+import org.apache.pinot.core.util.QueryMultiThreadingUtils;
 import 
org.apache.pinot.segment.local.segment.index.readers.bloom.GuavaBloomFilterReaderUtils;
+import org.apache.pinot.segment.spi.FetchContext;
 import org.apache.pinot.segment.spi.IndexSegment;
 import org.apache.pinot.segment.spi.datasource.DataSource;
 import org.apache.pinot.segment.spi.index.reader.BloomFilterReader;
 import org.apache.pinot.spi.data.FieldSpec.DataType;
 import org.apache.pinot.spi.env.PinotConfiguration;
 import org.apache.pinot.spi.exception.BadQueryRequestException;
+import org.apache.pinot.spi.exception.QueryCancelledException;
+import org.apache.pinot.spi.exception.QueryException;
 import org.apache.pinot.spi.utils.CommonConstants.Server;
 
 
@@ -45,6 +51,7 @@ import org.apache.pinot.spi.utils.CommonConstants.Server;
 @SuppressWarnings({"rawtypes", "unchecked"})
 abstract public class ValueBasedSegmentPruner implements SegmentPruner {
   public static final String IN_PREDICATE_THRESHOLD = "inpredicate.threshold";
+  private static final int TARGET_NUM_SEGMENTS_PER_THREAD = 10;
   protected int _inPredicateThreshold;
 
   @Override
@@ -110,20 +117,85 @@ abstract public class ValueBasedSegmentPruner implements 
SegmentPruner {
 
   @Override
   public List<IndexSegment> prune(List<IndexSegment> segments, QueryContext 
query) {
+    return prune(segments, query, null, null);
+  }
+
+  @Override
+  public List<IndexSegment> prune(List<IndexSegment> segments, QueryContext 
query,
+      @Nullable ExecutorService executorService) {
+    return prune(segments, query, executorService, null);
+  }
+
+  /// Each worker owns its caches. Parallel pruning may return segments in a 
different order.
+  protected List<IndexSegment> prune(List<IndexSegment> segments, QueryContext 
query,
+      @Nullable ExecutorService executorService, @Nullable FetchContext[] 
fetchContexts) {
     if (segments.isEmpty()) {
       return segments;
     }
-    FilterContext filter = Objects.requireNonNull(query.getFilter());
-    ValueCache cachedValues = new ValueCache();
-    Map<String, DataSource> dataSourceCache = new HashMap<>();
-    List<IndexSegment> selectedSegments = new ArrayList<>(segments.size());
-    for (IndexSegment segment : segments) {
-      dataSourceCache.clear();
-      if (!pruneSegment(segment, filter, dataSourceCache, cachedValues, 
query)) {
-        selectedSegments.add(segment);
+    if (executorService == null || segments.size() <= 
TARGET_NUM_SEGMENTS_PER_THREAD) {
+      FilterContext filter = Objects.requireNonNull(query.getFilter());
+      ValueCache cachedValues = new ValueCache();
+      Map<String, DataSource> dataSourceCache = new HashMap<>();
+      List<IndexSegment> selectedSegments = new ArrayList<>(segments.size());
+      int i = 0;
+      for (IndexSegment segment : segments) {
+        dataSourceCache.clear();
+        FetchContext fetchContext = fetchContexts != null ? fetchContexts[i++] 
: null;
+        if (!pruneSegmentWithFetchContext(segment, fetchContext, filter, 
dataSourceCache, cachedValues, query)) {
+          selectedSegments.add(segment);
+        }
       }
+      return selectedSegments;
+    }
+    int numSegments = segments.size();
+    int numTasks = QueryMultiThreadingUtils.getNumTasks(numSegments, 
TARGET_NUM_SEGMENTS_PER_THREAD,
+        query.getMaxExecutionThreads());
+    List<IndexSegment> allSelectedSegments = new ArrayList<>(numSegments);
+    QueryMultiThreadingUtils.runTasksWithDeadline(numTasks, index -> {
+      FilterContext filter = Objects.requireNonNull(query.getFilter());
+      ValueCache cachedValues = new ValueCache();
+      Map<String, DataSource> dataSourceCache = new HashMap<>();
+      List<IndexSegment> selectedSegments = new ArrayList<>();
+      for (int i = index; i < numSegments; i += numTasks) {
+        if (Thread.currentThread().isInterrupted()) {
+          throw new QueryCancelledException("Cancelled while running " + 
getClass().getSimpleName());
+        }
+        dataSourceCache.clear();
+        IndexSegment segment = segments.get(i);
+        FetchContext fetchContext = fetchContexts != null ? fetchContexts[i] : 
null;
+        if (!pruneSegmentWithFetchContext(segment, fetchContext, filter, 
dataSourceCache, cachedValues, query)) {
+          selectedSegments.add(segment);
+        }
+      }
+      return selectedSegments;
+    }, taskRes -> {
+      if (taskRes != null) {
+        allSelectedSegments.addAll(taskRes);
+      }
+    }, e -> {
+      Throwable cause = e.getCause();
+      if (cause instanceof QueryException) {
+        throw (QueryException) cause;
+      }
+      if (e instanceof InterruptedException) {
+        throw new QueryCancelledException("Cancelled while running " + 
getClass().getSimpleName(), e);
+      }
+      throw new RuntimeException("Caught exception while running " + 
getClass().getSimpleName(), e);
+    }, executorService, query.getEndTimeMs());
+    return allSelectedSegments;
+  }
+
+  private boolean pruneSegmentWithFetchContext(IndexSegment segment, @Nullable 
FetchContext fetchContext,
+      FilterContext filter, Map<String, DataSource> dataSourceCache, 
ValueCache cachedValues, QueryContext query) {
+    if (fetchContext == null) {
+      return pruneSegment(segment, filter, dataSourceCache, cachedValues, 
query);
+    }
+    segment.acquire(fetchContext);
+    try {
+      return pruneSegment(segment, filter, dataSourceCache, cachedValues, 
query);
+    } finally {
+      segment.release(fetchContext);
     }
-    return selectedSegments;
   }
 
   protected boolean pruneSegment(IndexSegment segment, FilterContext filter, 
Map<String, DataSource> dataSourceCache,
diff --git 
a/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPrunerTest.java
 
b/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPrunerTest.java
index 4a170576f29..2865581173c 100644
--- 
a/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPrunerTest.java
+++ 
b/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/BloomFilterSegmentPrunerTest.java
@@ -30,8 +30,13 @@ import java.util.ArrayList;
 import java.util.HashMap;
 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.TimeUnit;
+import java.util.concurrent.atomic.AtomicBoolean;
 import org.apache.commons.io.FileUtils;
+import org.apache.commons.lang3.exception.ExceptionUtils;
 import org.apache.pinot.core.query.request.context.QueryContext;
 import 
org.apache.pinot.core.query.request.context.utils.QueryContextConverterUtils;
 import 
org.apache.pinot.segment.local.indexsegment.immutable.ImmutableSegmentLoader;
@@ -39,6 +44,7 @@ import 
org.apache.pinot.segment.local.segment.creator.impl.SegmentIndexCreationD
 import org.apache.pinot.segment.local.segment.index.loader.IndexLoadingConfig;
 import 
org.apache.pinot.segment.local.segment.index.readers.bloom.OnHeapGuavaBloomFilterReader;
 import org.apache.pinot.segment.local.segment.readers.GenericRowRecordReader;
+import org.apache.pinot.segment.spi.FetchContext;
 import org.apache.pinot.segment.spi.ImmutableSegment;
 import org.apache.pinot.segment.spi.IndexSegment;
 import org.apache.pinot.segment.spi.SegmentMetadata;
@@ -54,18 +60,27 @@ import org.apache.pinot.spi.data.FieldSpec.DataType;
 import org.apache.pinot.spi.data.Schema;
 import org.apache.pinot.spi.data.readers.GenericRow;
 import org.apache.pinot.spi.env.PinotConfiguration;
+import org.apache.pinot.spi.exception.QueryCancelledException;
 import org.apache.pinot.spi.utils.ReadMode;
 import org.apache.pinot.spi.utils.builder.TableConfigBuilder;
+import org.mockito.ArgumentCaptor;
 import org.testng.annotations.BeforeClass;
 import org.testng.annotations.DataProvider;
 import org.testng.annotations.Test;
 
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.ArgumentMatchers.same;
+import static org.mockito.Mockito.doAnswer;
 import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.never;
+import static org.mockito.Mockito.times;
+import static org.mockito.Mockito.verify;
 import static org.mockito.Mockito.when;
 import static org.testng.Assert.assertEquals;
 import static org.testng.Assert.assertFalse;
 import static org.testng.Assert.assertNotNull;
 import static org.testng.Assert.assertTrue;
+import static org.testng.Assert.expectThrows;
 
 
 public class BloomFilterSegmentPrunerTest {
@@ -242,6 +257,127 @@ public class BloomFilterSegmentPrunerTest {
     assertEquals(selected.size(), 1);
   }
 
+  @DataProvider
+  public Object[][] prefetchExecutionModes() {
+    return new Object[][]{
+        {1, false, true},
+        {1, true, true},
+        {12, false, true},
+        {12, true, true},
+        {12, true, false}
+    };
+  }
+
+  @Test(dataProvider = "prefetchExecutionModes", timeOut = 10_000)
+  public void testPruningPreservesPrefetchLifecycle(int numSegments, boolean 
useExecutor, boolean enablePrefetch)
+      throws Exception {
+    List<IndexSegment> segments = new ArrayList<>();
+    List<IndexSegment> expected = new ArrayList<>();
+    for (int i = 0; i < numSegments; i++) {
+      IndexSegment segment = mockIndexSegment(new String[]{i % 2 == 0 ? "1.0" 
: "2.0"});
+      segments.add(segment);
+      if (i % 2 == 0) {
+        expected.add(segment);
+      }
+      AtomicBoolean acquired = new AtomicBoolean();
+      doAnswer(invocation -> {
+        assertFalse(acquired.getAndSet(true));
+        return null;
+      }).when(segment).acquire(any(FetchContext.class));
+      doAnswer(invocation -> {
+        acquired.set(false);
+        return null;
+      }).when(segment).release(any(FetchContext.class));
+      DataSource dataSource = segment.getDataSourceNullable("column");
+      DataSourceMetadata metadata = dataSource.getDataSourceMetadata();
+      when(dataSource.getDataSourceMetadata()).thenAnswer(invocation -> {
+        assertEquals(acquired.get(), enablePrefetch, "Prefetched data must be 
acquired before pruning");
+        return metadata;
+      });
+    }
+    QueryContext query = prefetchQuery();
+    query.setEnablePrefetch(enablePrefetch);
+    query.setMaxExecutionThreads(4);
+    ExecutorService executor = Executors.newFixedThreadPool(4);
+    try {
+      List<IndexSegment> selected = useExecutor ? PRUNER.prune(segments, 
query, executor)
+          : PRUNER.prune(segments, query);
+      assertEquals(selected.size(), expected.size());
+      assertEquals(Set.copyOf(selected), Set.copyOf(expected));
+      for (IndexSegment segment : segments) {
+        if (enablePrefetch) {
+          assertReleasedPrefetch(segment, 1);
+        } else {
+          verify(segment, never()).prefetch(any(FetchContext.class));
+          verify(segment, never()).acquire(any(FetchContext.class));
+          verify(segment, never()).release(any(FetchContext.class));
+        }
+      }
+    } finally {
+      executor.shutdownNow();
+      assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS));
+    }
+  }
+
+  @DataProvider
+  public Object[][] pruningFailures() {
+    return new Object[][]{{false}, {true}};
+  }
+
+  @Test(dataProvider = "pruningFailures", timeOut = 10_000)
+  public void testPrefetchReleasedAfterWorkerFailure(boolean interruptWorker)
+      throws Exception {
+    List<IndexSegment> segments = new ArrayList<>();
+    for (int i = 0; i < 12; i++) {
+      segments.add(mockIndexSegment(new String[]{"1.0"}));
+    }
+    DataSource dataSource = 
segments.getFirst().getDataSourceNullable("column");
+    DataSourceMetadata metadata = dataSource.getDataSourceMetadata();
+    IllegalStateException workerFailure = new IllegalStateException("Cannot 
read bloom filter metadata");
+    when(dataSource.getDataSourceMetadata()).thenAnswer(invocation -> {
+      if (interruptWorker) {
+        Thread.currentThread().interrupt();
+        return metadata;
+      }
+      throw workerFailure;
+    });
+    ExecutorService executor = Executors.newSingleThreadExecutor();
+    try {
+      RuntimeException failure = expectThrows(RuntimeException.class,
+          () -> PRUNER.prune(segments, prefetchQuery(), executor));
+      if (interruptWorker) {
+        assertTrue(ExceptionUtils.indexOfType(failure, 
QueryCancelledException.class) >= 0);
+      } else {
+        
assertTrue(ExceptionUtils.getThrowableList(failure).contains(workerFailure));
+      }
+      assertReleasedPrefetch(segments.getFirst(), 1);
+      for (int i = 1; i < segments.size(); i++) {
+        assertReleasedPrefetch(segments.get(i), 0);
+      }
+    } finally {
+      executor.shutdownNow();
+      assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS));
+    }
+  }
+
+  private QueryContext prefetchQuery() {
+    QueryContext query = QueryContextConverterUtils.getQueryContext(
+        "SELECT COUNT(*) FROM testTable WHERE column = 1.0");
+    query.setEnablePrefetch(true);
+    query.setMaxExecutionThreads(1);
+    query.setEndTimeMs(System.currentTimeMillis() + 30_000);
+    return query;
+  }
+
+  private void assertReleasedPrefetch(IndexSegment segment, int 
numAcquisitions) {
+    ArgumentCaptor<FetchContext> captor = 
ArgumentCaptor.forClass(FetchContext.class);
+    verify(segment).prefetch(captor.capture());
+    FetchContext fetchContext = captor.getValue();
+    assertFalse(fetchContext.isEmpty());
+    verify(segment, times(numAcquisitions)).acquire(same(fetchContext));
+    verify(segment, times(numAcquisitions + 1)).release(same(fetchContext));
+  }
+
   @Test
   public void testIsApplicableTo() {
     // EQ and IN (with small number of values) are applicable for bloom filter 
based pruning.
diff --git 
a/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/ColumnValueSegmentPrunerTest.java
 
b/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/ColumnValueSegmentPrunerTest.java
index c5cd96ec014..d4269141f75 100644
--- 
a/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/ColumnValueSegmentPrunerTest.java
+++ 
b/pinot-core/src/test/java/org/apache/pinot/core/query/pruner/ColumnValueSegmentPrunerTest.java
@@ -19,10 +19,22 @@
 package org.apache.pinot.core.query.pruner;
 
 import com.google.common.collect.ImmutableSet;
+import java.util.ArrayList;
 import java.util.Arrays;
+import java.util.Collections;
 import java.util.HashMap;
+import java.util.HashSet;
+import java.util.List;
 import java.util.Map;
 import java.util.Set;
+import java.util.concurrent.ConcurrentHashMap;
+import java.util.concurrent.CountDownLatch;
+import java.util.concurrent.ExecutorService;
+import java.util.concurrent.Executors;
+import java.util.concurrent.TimeUnit;
+import java.util.concurrent.TimeoutException;
+import java.util.concurrent.atomic.AtomicInteger;
+import org.apache.commons.lang3.exception.ExceptionUtils;
 import org.apache.pinot.core.query.request.context.QueryContext;
 import 
org.apache.pinot.core.query.request.context.utils.QueryContextConverterUtils;
 import org.apache.pinot.segment.spi.IndexSegment;
@@ -33,15 +45,22 @@ import 
org.apache.pinot.segment.spi.partition.PartitionFunctionFactory;
 import org.apache.pinot.spi.data.FieldSpec.DataType;
 import org.apache.pinot.spi.data.Schema;
 import org.apache.pinot.spi.env.PinotConfiguration;
+import org.apache.pinot.spi.exception.BadQueryRequestException;
+import org.apache.pinot.spi.exception.QueryCancelledException;
+import org.apache.pinot.spi.exception.QueryErrorCode;
 import org.testng.annotations.BeforeClass;
+import org.testng.annotations.DataProvider;
 import org.testng.annotations.Test;
 
 import static org.mockito.ArgumentMatchers.any;
 import static org.mockito.ArgumentMatchers.eq;
 import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.verifyNoInteractions;
 import static org.mockito.Mockito.when;
+import static org.testng.Assert.assertEquals;
 import static org.testng.Assert.assertFalse;
 import static org.testng.Assert.assertTrue;
+import static org.testng.Assert.expectThrows;
 
 
 public class ColumnValueSegmentPrunerTest {
@@ -233,6 +252,176 @@ public class ColumnValueSegmentPrunerTest {
     assertTrue(PRUNER.isApplicableTo(queryContext));
   }
 
+  @Test
+  public void testParallelPruningSelectsTheSameSegments() throws Exception {
+    int numSegments = 40;
+    Set<Thread> accessThreads = ConcurrentHashMap.newKeySet();
+    List<IndexSegment> segments = new ArrayList<>(numSegments);
+    for (int i = 0; i < numSegments; i++) {
+      // Alternate: half the segments hold values the predicate can match, 
half cannot and must be pruned.
+      segments.add(segmentWithRange(i % 2 == 0 ? 0 : 100, i % 2 == 0 ? 50 : 
150,
+          () -> accessThreads.add(Thread.currentThread())));
+    }
+    QueryContext serialQuery = QueryContextConverterUtils.getQueryContext(
+        "SELECT COUNT(*) FROM testTable WHERE column = 10");
+    serialQuery.setSchema(mock(Schema.class));
+    QueryContext parallelQuery = QueryContextConverterUtils.getQueryContext(
+        "SELECT COUNT(*) FROM testTable WHERE column = 10");
+    parallelQuery.setSchema(mock(Schema.class));
+    parallelQuery.setEndTimeMs(System.currentTimeMillis() + 30_000);
+    parallelQuery.setMaxExecutionThreads(4);
+
+    List<IndexSegment> serial = PRUNER.prune(segments, serialQuery);
+    assertEquals(accessThreads, Set.of(Thread.currentThread()));
+    accessThreads.clear();
+    ExecutorService executor = Executors.newFixedThreadPool(4);
+    try {
+      List<IndexSegment> parallel = PRUNER.prune(segments, parallelQuery, 
executor);
+      assertEquals(new HashSet<>(parallel), new HashSet<>(serial));
+      assertEquals(parallel.size(), numSegments / 2);
+      assertFalse(accessThreads.isEmpty());
+      assertFalse(accessThreads.contains(Thread.currentThread()), "Pruning 
must run on the supplied executor");
+    } finally {
+      executor.shutdownNow();
+    }
+  }
+
+  @Test
+  public void testParallelPruningFallsBackToCallingThread() {
+    Set<Thread> accessThreads = new HashSet<>();
+    IndexSegment segment = segmentWithRange(0, 50, () -> 
accessThreads.add(Thread.currentThread()));
+    QueryContext query = pruningQuery();
+    ExecutorService executor = mock(ExecutorService.class);
+
+    assertTrue(PRUNER.prune(List.of(), query, executor).isEmpty());
+    assertEquals(PRUNER.prune(Collections.nCopies(10, segment), query, 
executor).size(), 10);
+    verifyNoInteractions(executor);
+    assertEquals(PRUNER.prune(Collections.nCopies(40, segment), query, 
null).size(), 40);
+    assertEquals(accessThreads, Set.of(Thread.currentThread()));
+  }
+
+  @Test(timeOut = 10_000)
+  public void testParallelPruningRejectsExpiredDeadline() throws Exception {
+    CountDownLatch releaseWorker = new CountDownLatch(1);
+    IndexSegment segment = segmentWithRange(0, 50, () -> {
+      try {
+        releaseWorker.await();
+      } catch (InterruptedException e) {
+        Thread.currentThread().interrupt();
+        throw new QueryCancelledException("Pruning worker interrupted", e);
+      }
+    });
+    QueryContext query = pruningQuery();
+    query.setEndTimeMs(0);
+    ExecutorService executor = Executors.newSingleThreadExecutor();
+    try {
+      RuntimeException failure = expectThrows(RuntimeException.class,
+          () -> PRUNER.prune(Collections.nCopies(40, segment), query, 
executor));
+      assertTrue(ExceptionUtils.indexOfType(failure, TimeoutException.class) 
>= 0);
+    } finally {
+      releaseWorker.countDown();
+      executor.shutdownNow();
+      assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS));
+    }
+  }
+
+  @DataProvider
+  public Object[][] invalidPredicates() {
+    return new Object[][]{
+        {false, "column = 'potato'"},
+        {false, "column IN (1, 'potato')"},
+        {false, "column > 'potato'"},
+        {true, "column = 'potato'"},
+        {true, "column IN (1, 'potato')"}
+    };
+  }
+
+  @Test(dataProvider = "invalidPredicates", timeOut = 10_000)
+  public void testParallelPruningPreservesValidationErrors(boolean 
useBloomFilter, String predicate)
+      throws Exception {
+    IndexSegment segment = segmentWithRange(0, 50, () -> { });
+    QueryContext query = QueryContextConverterUtils.getQueryContext(
+        "SELECT COUNT(*) FROM testTable WHERE " + predicate);
+    query.setSchema(mock(Schema.class));
+    query.setMaxExecutionThreads(1);
+    query.setEndTimeMs(System.currentTimeMillis() + 30_000);
+    DataSource dataSource = segment.getDataSource("column", query.getSchema());
+    when(segment.getDataSourceNullable("column")).thenReturn(dataSource);
+    ValueBasedSegmentPruner pruner = useBloomFilter ? new 
BloomFilterSegmentPruner() : new ColumnValueSegmentPruner();
+    pruner.init(new PinotConfiguration());
+    List<IndexSegment> segments = Collections.nCopies(40, segment);
+    BadQueryRequestException serial = 
expectThrows(BadQueryRequestException.class, () -> pruner.prune(segments, 
query));
+    ExecutorService executor = Executors.newSingleThreadExecutor();
+    try {
+      BadQueryRequestException parallel = 
expectThrows(BadQueryRequestException.class,
+          () -> pruner.prune(segments, query, executor));
+      assertEquals(parallel.getErrorCode(), QueryErrorCode.QUERY_VALIDATION);
+      assertEquals(parallel.getMessage(), serial.getMessage());
+    } finally {
+      executor.shutdownNow();
+      assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS));
+    }
+  }
+
+  @Test(timeOut = 10_000)
+  public void testParallelPruningStopsInterruptedWorkerBetweenSegments() 
throws Exception {
+    AtomicInteger visits = new AtomicInteger();
+    IndexSegment segment = segmentWithRange(0, 50, () -> {
+      visits.incrementAndGet();
+      Thread.currentThread().interrupt();
+    });
+    ExecutorService executor = Executors.newSingleThreadExecutor();
+    try {
+      expectThrows(QueryCancelledException.class,
+          () -> PRUNER.prune(Collections.nCopies(40, segment), pruningQuery(), 
executor));
+      assertEquals(visits.get(), 1, "An interrupted worker must not keep 
visiting segments");
+    } finally {
+      executor.shutdownNow();
+      assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS));
+    }
+  }
+
+  @Test(timeOut = 10_000)
+  public void testParallelPruningPropagatesWorkerFailure() throws Exception {
+    IllegalStateException workerFailure = new IllegalStateException("Cannot 
read segment metadata");
+    IndexSegment segment = segmentWithRange(0, 50, () -> {
+      throw workerFailure;
+    });
+    ExecutorService executor = Executors.newSingleThreadExecutor();
+    try {
+      RuntimeException failure = expectThrows(RuntimeException.class,
+          () -> PRUNER.prune(Collections.nCopies(40, segment), pruningQuery(), 
executor));
+      
assertTrue(ExceptionUtils.getThrowableList(failure).contains(workerFailure));
+    } finally {
+      executor.shutdownNow();
+      assertTrue(executor.awaitTermination(5, TimeUnit.SECONDS));
+    }
+  }
+
+  private QueryContext pruningQuery() {
+    QueryContext query = QueryContextConverterUtils.getQueryContext(
+        "SELECT COUNT(*) FROM testTable WHERE column = 10");
+    query.setSchema(mock(Schema.class));
+    query.setMaxExecutionThreads(1);
+    query.setEndTimeMs(System.currentTimeMillis() + 30_000);
+    return query;
+  }
+
+  private IndexSegment segmentWithRange(int minValue, int maxValue, Runnable 
onAccess) {
+    IndexSegment indexSegment = mockIndexSegment();
+    DataSource dataSource = mock(DataSource.class);
+    when(indexSegment.getDataSource(eq("column"), 
any(Schema.class))).thenAnswer(invocation -> {
+      onAccess.run();
+      return dataSource;
+    });
+    DataSourceMetadata metadata = mock(DataSourceMetadata.class);
+    when(metadata.getDataType()).thenReturn(DataType.INT);
+    when(metadata.getMinValue()).thenReturn(minValue);
+    when(metadata.getMaxValue()).thenReturn(maxValue);
+    when(dataSource.getDataSourceMetadata()).thenReturn(metadata);
+    return indexSegment;
+  }
+
   private IndexSegment mockIndexSegment() {
     IndexSegment indexSegment = mock(IndexSegment.class);
     when(indexSegment.getColumnNames()).thenReturn(ImmutableSet.of("column"));


---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]

Reply via email to