diff --git a/integration-tests/test_lance_spark.py b/integration-tests/test_lance_spark.py index 7b54f5797..01b58b074 100644 --- a/integration-tests/test_lance_spark.py +++ b/integration-tests/test_lance_spark.py @@ -915,6 +915,41 @@ def test_create_distributed_bitmap_index(self, spark): ) assert spark.sql("SELECT * FROM default.test_table").count() == 4 + def test_zonemap_partial_coverage_after_append(self, spark): + """A fragment appended after index creation must remain visible to filtered scans.""" + spark.sql(""" + CREATE TABLE default.test_table ( + id INT, + name STRING, + value DOUBLE + ) + """) + + initial = [(i, f"Name{i}", float(i)) for i in range(10)] + spark.createDataFrame(initial, ["id", "name", "value"]).writeTo( + "default.test_table" + ).append() + + spark.sql(""" + ALTER TABLE default.test_table + CREATE INDEX idx_id_zonemap USING zonemap (id) + WITH (rows_per_zone = 4) + """).collect() + + spark.createDataFrame( + [(1000, "Appended", 1000.0)], ["id", "name", "value"] + ).writeTo("default.test_table").append() + + rows = spark.sql(""" + SELECT id, name, value + FROM default.test_table + WHERE id = 1000 + """).collect() + + assert len(rows) == 1 + assert rows[0].id == 1000 + assert rows[0].name == "Appended" + def test_create_btree_index_on_nested_literal_dot_field(self, spark): """Test CREATE INDEX on nested struct fields, including literal dots.""" spark.sql(""" diff --git a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScan.java b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScan.java index e0d072c74..33dd540e7 100644 --- a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScan.java +++ b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScan.java @@ -51,6 +51,7 @@ import java.io.Serializable; import java.util.Arrays; import java.util.Collections; +import java.util.HashSet; import java.util.List; import java.util.Objects; import java.util.Set; @@ -85,6 +86,9 @@ public class LanceScan */ private final java.util.Map> zonemapStats; + /** Live fragment IDs from the same Dataset snapshot as {@link #zonemapStats}. */ + private final Set liveFragmentIds; + /** * Pre-computed surviving fragment IDs from zonemap pruning in LanceScanBuilder. When non-null, * {@link #pruneByZonemapStats} skips re-computing and uses these directly. @@ -138,6 +142,7 @@ public LanceScan( Predicate[] pushedPredicates, LanceStatistics statistics, java.util.Map> zonemapStats, + Set liveFragmentIds, Set survivingFragmentIds, List precomputedSplits, java.util.Map precomputedFragmentRowCounts, @@ -159,6 +164,9 @@ public LanceScan( : new Predicate[0]; this.statistics = statistics; this.zonemapStats = zonemapStats != null ? zonemapStats : Collections.emptyMap(); + this.liveFragmentIds = + Collections.unmodifiableSet( + new HashSet<>(Objects.requireNonNull(liveFragmentIds, "liveFragmentIds"))); this.cachedSurvivingFragmentIds = survivingFragmentIds; this.precomputedSplits = precomputedSplits; this.precomputedFragmentRowCounts = @@ -371,7 +379,8 @@ private List pruneByZonemapStats(List allSplits) { allowedIds = cachedSurvivingFragmentIds; } else if (!zonemapStats.isEmpty()) { allowedIds = - ZonemapFragmentPruner.pruneFragments(pushedPredicates, zonemapStats).orElse(null); + ZonemapFragmentPruner.pruneFragments(pushedPredicates, zonemapStats, liveFragmentIds) + .orElse(null); } else { return allSplits; } diff --git a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScanBuilder.java b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScanBuilder.java index d33a7f38e..c98d011c5 100644 --- a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScanBuilder.java +++ b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceScanBuilder.java @@ -191,6 +191,12 @@ public Scan build() { SparkLanceShardingUtils.isEmpty(shardingSpec) ? SparkLanceShardingUtils.firstShardingSpec(dataset) : shardingSpec; + + // Plan splits before zonemap analysis so live fragment IDs, zonemap stats, and the splits + // shipped to workers all come from the same Dataset snapshot. + LanceSplit.ScanPlanResult scanPlan = LanceSplit.planScan(dataset, readOptions); + Set liveFragmentIds = new HashSet<>(scanPlan.getFragmentRowCounts().keySet()); + for (ShardingField field : SparkLanceShardingUtils.fields(activeShardingSpec)) { columnsToLoad.add(SparkLanceShardingUtils.columnName(field, lanceSchema)); } @@ -215,7 +221,8 @@ public Scan build() { continue; } java.util.Optional> keys = - SparkLanceShardingUtils.detectFragmentKeys(field, lanceSchema, colStats); + SparkLanceShardingUtils.detectFragmentKeys( + field, lanceSchema, colStats, liveFragmentIds); if (keys.isPresent()) { fragmentShardingKeys = keys.get(); activeShardingExpression = SparkLanceShardingUtils.toSparkExpression(field, lanceSchema); @@ -234,7 +241,8 @@ public Scan build() { Set survivingFragmentIds = null; if (pushedPredicates.length > 0 && !zonemapStats.isEmpty()) { survivingFragmentIds = - ZonemapFragmentPruner.pruneFragments(pushedPredicates, zonemapStats).orElse(null); + ZonemapFragmentPruner.pruneFragments(pushedPredicates, zonemapStats, liveFragmentIds) + .orElse(null); } // Scale rows and full size by the zonemap fragment-pruning ratio first, then let @@ -242,10 +250,14 @@ public Scan build() { // (when the projected schema is narrower than the full schema). long projectedRows = summary.getTotalRows(); long projectedFullSize = summary.getTotalFilesSize(); - if (survivingFragmentIds != null && summary.getTotalFragments() > 0) { - double ratio = (double) survivingFragmentIds.size() / summary.getTotalFragments(); - projectedRows = (long) (projectedRows * ratio); - projectedFullSize = (long) (projectedFullSize * ratio); + if (survivingFragmentIds != null && !liveFragmentIds.isEmpty()) { + long survivingRows = + survivingFragmentIds.stream().mapToLong(scanPlan.getFragmentRowCounts()::get).sum(); + LanceStatistics postPruning = + LanceStatistics.estimatePostPruningByRows( + summary.getTotalRows(), summary.getTotalFilesSize(), survivingRows); + projectedRows = postPruning.numRows().getAsLong(); + projectedFullSize = postPruning.sizeInBytes().getAsLong(); } LanceStatistics statistics = LanceStatistics.estimateProjected(projectedRows, projectedFullSize, fullSchema, schema); @@ -254,19 +266,13 @@ public Scan build() { "Scan statistics after pruning: {} of {} fragments survive," + " estimatedSize={}, estimatedRows={} (full: size={}, rows={})", survivingFragmentIds.size(), - summary.getTotalFragments(), + liveFragmentIds.size(), statistics.sizeInBytes(), statistics.numRows(), summary.getTotalFilesSize(), summary.getTotalRows()); } - // Pre-compute splits and per-fragment row counts from the same Dataset handle that we - // already opened above. This consolidates two driver-side opens into one and lets us pin - // the resolved version onto the read options shipped to workers, providing snapshot - // isolation across all tasks of this query. The version is kept as a long end-to-end so - // long-lived high-write-frequency datasets do not silently truncate to a wrong version. - LanceSplit.ScanPlanResult scanPlan = LanceSplit.planScan(dataset, readOptions); LanceSparkReadOptions resolvedReadOptions = readOptions.withRef(scanPlan.getRef()); Optional whereCondition = @@ -282,6 +288,7 @@ public Scan build() { pushedPredicates, statistics, zonemapStats, + liveFragmentIds, survivingFragmentIds, scanPlan.getSplits(), scanPlan.getFragmentRowCounts(), diff --git a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceStatistics.java b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceStatistics.java index 3e0826c86..e40509e5c 100644 --- a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceStatistics.java +++ b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/LanceStatistics.java @@ -67,6 +67,30 @@ public static LanceStatistics estimatePostPruning( return new LanceStatistics((long) (totalRows * ratio), (long) (totalFilesSize * ratio)); } + /** + * Estimate post-pruning statistics from the exact row count of surviving fragments. + * + *

The row count is exact for the fragments selected by planning. File size remains an estimate + * because the scan plan does not carry per-fragment byte sizes, so it is scaled by the + * surviving-row ratio. Invalid or non-selective inputs conservatively retain full-table stats. + * + * @param totalRows total rows in the dataset + * @param totalFilesSize total file size in bytes + * @param survivingRows exact row count across surviving fragments + * @return row-weighted post-pruning statistics + */ + static LanceStatistics estimatePostPruningByRows( + long totalRows, long totalFilesSize, long survivingRows) { + if (totalRows <= 0 || survivingRows >= totalRows) { + return new LanceStatistics(totalRows, totalFilesSize); + } + if (survivingRows <= 0) { + return new LanceStatistics(0, 0); + } + double ratio = (double) survivingRows / totalRows; + return new LanceStatistics(survivingRows, (long) (totalFilesSize * ratio)); + } + /** * Estimate post-projection size using {@code sizeInBytes × (projectedWidths / fullWidths)}, the * same formula Spark's DSv2 {@code FileScan.estimateStatistics} applies after column pruning (see diff --git a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/ZonemapFragmentPruner.java b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/ZonemapFragmentPruner.java index 0152b1ef6..cc3503a30 100644 --- a/lance-spark-base_2.12/src/main/java/org/lance/spark/read/ZonemapFragmentPruner.java +++ b/lance-spark-base_2.12/src/main/java/org/lance/spark/read/ZonemapFragmentPruner.java @@ -31,6 +31,7 @@ import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.Objects; import java.util.Optional; import java.util.Set; @@ -61,11 +62,16 @@ private ZonemapFragmentPruner() {} * * @param pushedPredicates the V2 predicates pushed down by Spark * @param zonemapStatsByColumn map from column name to its zonemap zone stats + * @param liveFragmentIds fragment IDs in the Dataset snapshot that produced the stats * @return present with the set of fragment IDs that might match; empty if no pruning can be * derived */ public static Optional> pruneFragments( - Predicate[] pushedPredicates, Map> zonemapStatsByColumn) { + Predicate[] pushedPredicates, + Map> zonemapStatsByColumn, + Set liveFragmentIds) { + + Objects.requireNonNull(liveFragmentIds, "liveFragmentIds"); if (pushedPredicates == null || pushedPredicates.length == 0 @@ -76,7 +82,8 @@ public static Optional> pruneFragments( Set result = null; for (Predicate predicate : pushedPredicates) { - Optional> fragmentIds = analyzePredicate(predicate, zonemapStatsByColumn); + Optional> fragmentIds = + analyzePredicate(predicate, zonemapStatsByColumn, liveFragmentIds); if (fragmentIds.isPresent()) { if (result == null) { result = new HashSet<>(fragmentIds.get()); @@ -100,13 +107,15 @@ public static Optional> pruneFragments( * not aliased by any other reference. Callers may freely mutate it. */ private static Optional> analyzePredicate( - Predicate predicate, Map> statsByColumn) { + Predicate predicate, + Map> statsByColumn, + Set liveFragmentIds) { if (predicate instanceof And) { - return analyzeAnd((And) predicate, statsByColumn); + return analyzeAnd((And) predicate, statsByColumn, liveFragmentIds); } if (predicate instanceof Or) { - return analyzeOr((Or) predicate, statsByColumn); + return analyzeOr((Or) predicate, statsByColumn, liveFragmentIds); } if (predicate instanceof Not) { return Optional.empty(); @@ -116,21 +125,25 @@ private static Optional> analyzePredicate( String name = predicate.name(); switch (name) { case "=": - return analyzeComparison(children, statsByColumn, ComparisonType.EQUALS); + return analyzeComparison(children, statsByColumn, liveFragmentIds, ComparisonType.EQUALS); case "<": - return analyzeComparison(children, statsByColumn, ComparisonType.LESS_THAN); + return analyzeComparison( + children, statsByColumn, liveFragmentIds, ComparisonType.LESS_THAN); case "<=": - return analyzeComparison(children, statsByColumn, ComparisonType.LESS_THAN_OR_EQUAL); + return analyzeComparison( + children, statsByColumn, liveFragmentIds, ComparisonType.LESS_THAN_OR_EQUAL); case ">": - return analyzeComparison(children, statsByColumn, ComparisonType.GREATER_THAN); + return analyzeComparison( + children, statsByColumn, liveFragmentIds, ComparisonType.GREATER_THAN); case ">=": - return analyzeComparison(children, statsByColumn, ComparisonType.GREATER_THAN_OR_EQUAL); + return analyzeComparison( + children, statsByColumn, liveFragmentIds, ComparisonType.GREATER_THAN_OR_EQUAL); case "IN": - return analyzeIn(children, statsByColumn); + return analyzeIn(children, statsByColumn, liveFragmentIds); case "IS_NULL": - return analyzeIsNull(children, statsByColumn); + return analyzeIsNull(children, statsByColumn, liveFragmentIds); case "IS_NOT_NULL": - return analyzeIsNotNull(children, statsByColumn); + return analyzeIsNotNull(children, statsByColumn, liveFragmentIds); default: return Optional.empty(); } @@ -138,7 +151,10 @@ private static Optional> analyzePredicate( @SuppressWarnings("unchecked") private static Optional> analyzeComparison( - Expression[] children, Map> statsByColumn, ComparisonType type) { + Expression[] children, + Map> statsByColumn, + Set liveFragmentIds, + ComparisonType type) { if (children.length != 2 || !(children[0] instanceof NamedReference) @@ -162,13 +178,19 @@ private static Optional> analyzeComparison( } Set matchingFragments = new HashSet<>(); + Set indexedFragments = new HashSet<>(); for (ZoneStats zone : stats) { + if (!liveFragmentIds.contains(zone.getFragmentId())) { + continue; + } + indexedFragments.add(zone.getFragmentId()); if (zoneMatchesComparison(zone, target, type)) { matchingFragments.add(zone.getFragmentId()); } } - return Optional.of(matchingFragments); + return Optional.of( + includeUnindexedFragments(matchingFragments, indexedFragments, liveFragmentIds)); } @SuppressWarnings("unchecked") @@ -206,7 +228,9 @@ private static boolean zoneMatchesComparison( } private static Optional> analyzeIn( - Expression[] children, Map> statsByColumn) { + Expression[] children, + Map> statsByColumn, + Set liveFragmentIds) { if (children.length < 1 || !(children[0] instanceof NamedReference)) { return Optional.empty(); @@ -229,7 +253,12 @@ private static Optional> analyzeIn( } Set matchingFragments = new HashSet<>(); + Set indexedFragments = new HashSet<>(); for (ZoneStats zone : stats) { + if (!liveFragmentIds.contains(zone.getFragmentId())) { + continue; + } + indexedFragments.add(zone.getFragmentId()); for (Object value : normalizedValues) { if (value == null) { if (zone.getNullCount() > 0) { @@ -252,11 +281,14 @@ private static Optional> analyzeIn( } } - return Optional.of(matchingFragments); + return Optional.of( + includeUnindexedFragments(matchingFragments, indexedFragments, liveFragmentIds)); } private static Optional> analyzeIsNull( - Expression[] children, Map> statsByColumn) { + Expression[] children, + Map> statsByColumn, + Set liveFragmentIds) { if (children.length != 1 || !(children[0] instanceof NamedReference)) { return Optional.empty(); @@ -268,17 +300,25 @@ private static Optional> analyzeIsNull( } Set matchingFragments = new HashSet<>(); + Set indexedFragments = new HashSet<>(); for (ZoneStats zone : stats) { + if (!liveFragmentIds.contains(zone.getFragmentId())) { + continue; + } + indexedFragments.add(zone.getFragmentId()); if (zone.getNullCount() > 0) { matchingFragments.add(zone.getFragmentId()); } } - return Optional.of(matchingFragments); + return Optional.of( + includeUnindexedFragments(matchingFragments, indexedFragments, liveFragmentIds)); } private static Optional> analyzeIsNotNull( - Expression[] children, Map> statsByColumn) { + Expression[] children, + Map> statsByColumn, + Set liveFragmentIds) { if (children.length != 1 || !(children[0] instanceof NamedReference)) { return Optional.empty(); @@ -290,7 +330,12 @@ private static Optional> analyzeIsNotNull( } Set matchingFragments = new HashSet<>(); + Set indexedFragments = new HashSet<>(); for (ZoneStats zone : stats) { + if (!liveFragmentIds.contains(zone.getFragmentId())) { + continue; + } + indexedFragments.add(zone.getFragmentId()); // Zone has non-null rows if zoneLength exceeds nullCount. // Conservative: zoneLength may include gaps from deletions. if (zone.getNullCount() < zone.getZoneLength()) { @@ -298,13 +343,16 @@ private static Optional> analyzeIsNotNull( } } - return Optional.of(matchingFragments); + return Optional.of( + includeUnindexedFragments(matchingFragments, indexedFragments, liveFragmentIds)); } private static Optional> analyzeAnd( - And predicate, Map> statsByColumn) { - Optional> left = analyzePredicate(predicate.left(), statsByColumn); - Optional> right = analyzePredicate(predicate.right(), statsByColumn); + And predicate, Map> statsByColumn, Set liveFragmentIds) { + Optional> left = + analyzePredicate(predicate.left(), statsByColumn, liveFragmentIds); + Optional> right = + analyzePredicate(predicate.right(), statsByColumn, liveFragmentIds); if (left.isPresent() && right.isPresent()) { Set intersection = new HashSet<>(left.get()); @@ -317,9 +365,11 @@ private static Optional> analyzeAnd( } private static Optional> analyzeOr( - Or predicate, Map> statsByColumn) { - Optional> left = analyzePredicate(predicate.left(), statsByColumn); - Optional> right = analyzePredicate(predicate.right(), statsByColumn); + Or predicate, Map> statsByColumn, Set liveFragmentIds) { + Optional> left = + analyzePredicate(predicate.left(), statsByColumn, liveFragmentIds); + Optional> right = + analyzePredicate(predicate.right(), statsByColumn, liveFragmentIds); if (left.isPresent() && right.isPresent()) { Set union = new HashSet<>(left.get()); @@ -329,6 +379,15 @@ private static Optional> analyzeOr( return Optional.empty(); } + private static Set includeUnindexedFragments( + Set matchingFragments, Set indexedFragments, Set liveFragmentIds) { + Set candidates = new HashSet<>(matchingFragments); + Set unindexedFragments = new HashSet<>(liveFragmentIds); + unindexedFragments.removeAll(indexedFragments); + candidates.addAll(unindexedFragments); + return candidates; + } + private static String columnName(NamedReference ref) { String[] names = ref.fieldNames(); return names.length == 1 ? names[0] : String.join(".", names); diff --git a/lance-spark-base_2.12/src/main/java/org/lance/spark/sharding/SparkLanceShardingUtils.java b/lance-spark-base_2.12/src/main/java/org/lance/spark/sharding/SparkLanceShardingUtils.java index 35af1920c..7a8a4d1eb 100644 --- a/lance-spark-base_2.12/src/main/java/org/lance/spark/sharding/SparkLanceShardingUtils.java +++ b/lance-spark-base_2.12/src/main/java/org/lance/spark/sharding/SparkLanceShardingUtils.java @@ -41,9 +41,12 @@ import java.util.ArrayList; import java.util.Collections; import java.util.HashMap; +import java.util.HashSet; import java.util.List; import java.util.Map; +import java.util.Objects; import java.util.Optional; +import java.util.Set; /** Spark-facing helpers for Lance MemWAL sharding specs. */ public final class SparkLanceShardingUtils { @@ -130,14 +133,31 @@ public static Expression toSparkExpression(ShardingField field, LanceSchema sche } public static Optional> detectFragmentKeys( - ShardingField field, LanceSchema schema, List zones) { + ShardingField field, + LanceSchema schema, + List zones, + Set liveFragmentIds) { columnName(field, schema); - Map result = new HashMap<>(); + Objects.requireNonNull(liveFragmentIds, "liveFragmentIds"); + if (liveFragmentIds.isEmpty()) { + return Optional.empty(); + } + + List liveZones = new ArrayList<>(); + Set coveredFragmentIds = new HashSet<>(); for (ZoneStats zone : zones) { - result.putIfAbsent(zone.getFragmentId(), null); + if (liveFragmentIds.contains(zone.getFragmentId())) { + liveZones.add(zone); + coveredFragmentIds.add(zone.getFragmentId()); + } } - for (int fragmentId : new ArrayList<>(result.keySet())) { - Optional key = fragmentKeyFromZones(field, schema, zones, fragmentId); + if (!coveredFragmentIds.equals(liveFragmentIds)) { + return Optional.empty(); + } + + Map result = new HashMap<>(); + for (int fragmentId : liveFragmentIds) { + Optional key = fragmentKeyFromZones(field, schema, liveZones, fragmentId); if (!key.isPresent()) { return Optional.empty(); } diff --git a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/BaseBucketSpjTest.java b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/BaseBucketSpjTest.java index d50919ac7..608952296 100644 --- a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/BaseBucketSpjTest.java +++ b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/BaseBucketSpjTest.java @@ -159,4 +159,35 @@ public void testBucketJoinPlanShowsSpj() { + " partitioning must be recognized. Plan: " + plan); } + + @Test + public void testPartialZonemapCoverageDisablesSpj() { + String tableA = "bkt_partial_a_" + UUID.randomUUID().toString().replace("-", ""); + String tableB = "bkt_partial_b_" + UUID.randomUUID().toString().replace("-", ""); + createBucketedTable(tableA, 4); + createBucketedTable(tableB, 4); + + String fullA = catalogName + ".default." + tableA; + String fullB = catalogName + ".default." + tableB; + + // This fragment is not covered by the already-committed zonemap index. + spark.sql( + String.format("INSERT INTO %s (id, region, value) VALUES (1000, 'east', 1000.0)", fullA)); + + Dataset joined = + spark.sql( + String.format( + "SELECT a.id, a.region, b.value " + + "FROM %s a JOIN %s b " + + "ON a.region = b.region", + fullA, fullB)); + + long count = joined.count(); + assertEquals(310, count, "The unindexed fragment must participate in the join"); + + String plan = joined.queryExecution().executedPlan().toString(); + assertTrue( + plan.contains("Exchange"), + "Partial zonemap coverage must disable SPJ and require a shuffle. Plan: " + plan); + } } diff --git a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceScanTest.java b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceScanTest.java index 78291ff96..3e67cfc76 100644 --- a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceScanTest.java +++ b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceScanTest.java @@ -211,6 +211,7 @@ public void testOutputPartitioningWithPartitionInfo() { new Predicate[0], null, Collections.emptyMap(), + plan.getFragmentRowCounts().keySet(), null, plan.getSplits(), plan.getFragmentRowCounts(), @@ -296,6 +297,7 @@ public void testOutputPartitioningWithBucketInfo() { new Predicate[0], null, Collections.emptyMap(), + bucketPlan.getFragmentRowCounts().keySet(), null, bucketPlan.getSplits(), bucketPlan.getFragmentRowCounts(), diff --git a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceStatisticsTest.java b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceStatisticsTest.java index ad489c0fa..490ca00ae 100644 --- a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceStatisticsTest.java +++ b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/LanceStatisticsTest.java @@ -77,6 +77,30 @@ public void testEstimatePostPruningZeroTotalFragments() { assertEquals(50000, stats.sizeInBytes().getAsLong()); } + @Test + public void testEstimatePostPruningByRowsHandlesUnevenFragments() { + LanceStatistics stats = LanceStatistics.estimatePostPruningByRows(1_000_020, 10_000_200, 20); + + assertEquals(20, stats.numRows().getAsLong()); + assertEquals(200, stats.sizeInBytes().getAsLong()); + } + + @Test + public void testEstimatePostPruningByRowsKeepsEmptyDatasetSize() { + LanceStatistics stats = LanceStatistics.estimatePostPruningByRows(0, 128, 0); + + assertEquals(0, stats.numRows().getAsLong()); + assertEquals(128, stats.sizeInBytes().getAsLong()); + } + + @Test + public void testEstimatePostPruningByRowsWithNoSurvivors() { + LanceStatistics stats = LanceStatistics.estimatePostPruningByRows(100, 1_000, 0); + + assertEquals(0, stats.numRows().getAsLong()); + assertEquals(0, stats.sizeInBytes().getAsLong()); + } + @Test public void testEstimateProjectedScalesByColumnWidthRatio() { // Full schema has 9 columns; project 3 of equal width → size should scale by 3/9. diff --git a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/ZonemapFragmentPrunerTest.java b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/ZonemapFragmentPrunerTest.java index 779e60d68..edfc9eda9 100644 --- a/lance-spark-base_2.12/src/test/java/org/lance/spark/read/ZonemapFragmentPrunerTest.java +++ b/lance-spark-base_2.12/src/test/java/org/lance/spark/read/ZonemapFragmentPrunerTest.java @@ -28,6 +28,7 @@ import java.util.Map; import java.util.Optional; import java.util.Set; +import java.util.stream.Collectors; import static org.junit.jupiter.api.Assertions.*; @@ -55,12 +56,20 @@ private static Map> threeFragmentStats(String column) { return stats; } + private static Set fragmentIds(Map> stats) { + return stats.values().stream() + .flatMap(List::stream) + .map(ZoneStats::getFragmentId) + .collect(Collectors.toSet()); + } + @Test public void testEqualToMatchesOneFragment() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 150L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -70,7 +79,8 @@ public void testEqualToMatchesNoFragment() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 500L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertTrue(result.get().isEmpty()); } @@ -81,7 +91,8 @@ public void testEqualToMatchesBoundary() { // Value at exact boundary between fragment 0 and 1 Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 99L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0), result.get()); } @@ -91,7 +102,8 @@ public void testEqualToMatchesExactMin() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 100L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -102,7 +114,8 @@ public void testLessThanPrunesHighFragments() { // x < 50 → only fragment 0's min (0) < 50 Predicate[] filters = new Predicate[] {TestPredicates.lt("x", 50L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0), result.get()); } @@ -113,7 +126,8 @@ public void testLessThanOrEqualIncludesBoundary() { // x <= 100 → fragment 0 (min=0 <= 100) and fragment 1 (min=100 <= 100) Predicate[] filters = new Predicate[] {TestPredicates.lte("x", 100L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 1), result.get()); } @@ -124,7 +138,8 @@ public void testGreaterThanPrunesLowFragments() { // x > 250 → only fragment 2's max (299) > 250 Predicate[] filters = new Predicate[] {TestPredicates.gt("x", 250L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(2), result.get()); } @@ -135,7 +150,8 @@ public void testGreaterThanOrEqualIncludesBoundary() { // x >= 199 → fragment 1 (max=199 >= 199) and fragment 2 (max=299 >= 199) Predicate[] filters = new Predicate[] {TestPredicates.gte("x", 199L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1, 2), result.get()); } @@ -146,7 +162,8 @@ public void testInWithMultipleValues() { // x IN (50, 250) → fragment 0 and fragment 2 Predicate[] filters = new Predicate[] {TestPredicates.in("x", 50L, 250L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 2), result.get()); } @@ -157,7 +174,8 @@ public void testInWithNoMatchingValues() { // x IN (500, 600) → no fragments match Predicate[] filters = new Predicate[] {TestPredicates.in("x", 500L, 600L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertTrue(result.get().isEmpty()); } @@ -175,7 +193,8 @@ public void testInWithNonLiteralChildBailsOut() { }; Predicate[] filters = new Predicate[] {new Predicate("IN", children)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertFalse( result.isPresent(), "IN with a non-Literal child must bail out instead of pruning on the remaining literals"); @@ -187,7 +206,8 @@ public void testIsNullWithNoNulls() { // All zones have nullCount=0 Predicate[] filters = new Predicate[] {TestPredicates.isNull("x")}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertTrue(result.get().isEmpty()); } @@ -204,7 +224,8 @@ public void testIsNullWithSomeNulls() { Predicate[] filters = new Predicate[] {TestPredicates.isNull("x")}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -221,7 +242,8 @@ public void testIsNotNull() { Predicate[] filters = new Predicate[] {TestPredicates.isNotNull("x")}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 2), result.get()); } @@ -235,7 +257,8 @@ public void testAndIntersectsFragments() { TestPredicates.and(TestPredicates.gte("x", 50L), TestPredicates.lte("x", 150L)) }; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 1), result.get()); } @@ -249,7 +272,8 @@ public void testOrUnionsFragments() { TestPredicates.or(TestPredicates.eq("x", 50L), TestPredicates.eq("x", 250L)) }; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 2), result.get()); } @@ -263,7 +287,8 @@ public void testOrWithUnconstainedSideReturnsEmpty() { TestPredicates.or(TestPredicates.eq("x", 50L), TestPredicates.eq("name", "Alice")) }; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertFalse(result.isPresent()); } @@ -272,7 +297,8 @@ public void testNotReturnsNoPruning() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.not(TestPredicates.eq("x", 50L))}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertFalse(result.isPresent()); } @@ -282,21 +308,24 @@ public void testNonIndexedColumnReturnsNoPruning() { // Filter on column 'y' which has no zonemap stats Predicate[] filters = new Predicate[] {TestPredicates.eq("y", 50L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertFalse(result.isPresent()); } @Test public void testEmptyFilters() { Map> stats = threeFragmentStats("x"); - Optional> result = ZonemapFragmentPruner.pruneFragments(new Predicate[] {}, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(new Predicate[] {}, stats, fragmentIds(stats)); assertFalse(result.isPresent()); } @Test public void testNullFilters() { Map> stats = threeFragmentStats("x"); - Optional> result = ZonemapFragmentPruner.pruneFragments(null, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(null, stats, fragmentIds(stats)); assertFalse(result.isPresent()); } @@ -304,7 +333,8 @@ public void testNullFilters() { public void testEmptyStats() { Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 50L)}; Optional> result = - ZonemapFragmentPruner.pruneFragments(filters, Collections.emptyMap()); + ZonemapFragmentPruner.pruneFragments( + filters, Collections.emptyMap(), Collections.emptySet()); assertFalse(result.isPresent()); } @@ -316,7 +346,8 @@ public void testMultipleTopLevelFiltersIntersect() { Predicate[] filters = new Predicate[] {TestPredicates.gte("x", 50L), TestPredicates.lt("x", 150L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); // x >= 50 matches {0,1,2} (all have max >= 50) // x < 150 matches {0,1} (min < 150 for frags 0 and 1) @@ -347,7 +378,8 @@ public void testMultipleColumnsIntersect() { Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 50L), TestPredicates.eq("y", 1200L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertTrue(result.get().isEmpty()); } @@ -367,7 +399,8 @@ public void testMultipleZonesPerFragment() { // x = 75 → matches second zone of fragment 0 → fragment 0 survives Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 75L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0), result.get()); } @@ -385,7 +418,8 @@ public void testStringColumnComparison() { // name = 'foo' → falls in [eaa, hzz] → fragment 1 Predicate[] filters = new Predicate[] {TestPredicates.eq("name", "foo")}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -402,7 +436,8 @@ public void testInWithNullValue() { // x IN (null, 50) → fragment 0 (has 50) and fragment 1 (has nulls) Predicate[] filters = new Predicate[] {TestPredicates.in("x", null, 50L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 1), result.get()); } @@ -416,7 +451,8 @@ public void testContradictoryAndYieldsEmptySet() { TestPredicates.and(TestPredicates.gt("x", 300L), TestPredicates.lt("x", 0L)) }; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertTrue(result.get().isEmpty()); } @@ -435,7 +471,8 @@ public void testNestedAndInsideOr() { TestPredicates.eq("x", 250L)) }; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 2), result.get()); } @@ -452,7 +489,8 @@ public void testAllNullZoneSkippedForEqualTo() { // x = 50 → fragment 0 matches, fragment 1 (all null) does not Predicate[] filters = new Predicate[] {TestPredicates.eq("x", 50L)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0), result.get()); } @@ -462,7 +500,8 @@ public void testIntegerLiteralAgainstLongZoneStats() { Map> stats = threeFragmentStats("seq"); Predicate[] filters = new Predicate[] {TestPredicates.eq("seq", 150)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -472,7 +511,8 @@ public void testShortLiteralAgainstLongZoneStats() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.gt("x", (short) 150)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1, 2), result.get()); } @@ -482,7 +522,8 @@ public void testByteLiteralAgainstLongZoneStats() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.lte("x", (byte) 50)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0), result.get()); } @@ -499,7 +540,8 @@ public void testFloatLiteralAgainstDoubleZoneStats() { Predicate[] filters = new Predicate[] {TestPredicates.eq("f", 15.0f)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -519,7 +561,8 @@ public void testDateLiteralAgainstLongZoneStats() { Predicate[] filters = new Predicate[] {TestPredicates.eq("d", Date.valueOf("2022-04-27"))}; // day 19109 - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(1), result.get()); } @@ -529,7 +572,8 @@ public void testInListWithIntegerLiteralsAgainstLongZoneStats() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.in("x", 50, 250, 999)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 2), result.get()); } @@ -540,8 +584,122 @@ public void testInListWithMixedWidthLiteralsAgainstLongZoneStats() { Map> stats = threeFragmentStats("x"); Predicate[] filters = new Predicate[] {TestPredicates.in("x", 50, 250L, (short) 70, (byte) 5)}; - Optional> result = ZonemapFragmentPruner.pruneFragments(filters, stats); + Optional> result = + ZonemapFragmentPruner.pruneFragments(filters, stats, fragmentIds(stats)); assertTrue(result.isPresent()); assertEquals(Set.of(0, 2), result.get()); } + + @Test + public void testPartialComparisonIncludesUnindexedFragment() { + Map> stats = new HashMap<>(); + stats.put( + "x", + Arrays.asList( + new ZoneStats(0, 0, 100, 0L, 99L, 0), new ZoneStats(1, 0, 100, 100L, 199L, 0))); + + Optional> result = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] {TestPredicates.eq("x", 500L)}, stats, Set.of(0, 1, 2)); + + assertTrue(result.isPresent()); + assertEquals(Set.of(2), result.get()); + } + + @Test + public void testPartialInIncludesMatchingAndUnindexedFragments() { + Map> stats = new HashMap<>(); + stats.put( + "x", + Arrays.asList( + new ZoneStats(0, 0, 100, 0L, 99L, 0), new ZoneStats(1, 0, 100, 100L, 199L, 0))); + + Optional> result = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] {TestPredicates.in("x", 50L, 500L)}, stats, Set.of(0, 1, 2)); + + assertTrue(result.isPresent()); + assertEquals(Set.of(0, 2), result.get()); + } + + @Test + public void testPartialIsNullIncludesUnindexedFragment() { + Map> stats = new HashMap<>(); + stats.put( + "x", + Arrays.asList( + new ZoneStats(0, 0, 100, 0L, 99L, 0), new ZoneStats(1, 0, 100, 100L, 199L, 5))); + + Optional> result = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] {TestPredicates.isNull("x")}, stats, Set.of(0, 1, 2)); + + assertTrue(result.isPresent()); + assertEquals(Set.of(1, 2), result.get()); + } + + @Test + public void testPartialIsNotNullIncludesUnindexedFragment() { + Map> stats = new HashMap<>(); + stats.put( + "x", + Arrays.asList( + new ZoneStats(0, 0, 100, 0L, 99L, 0), new ZoneStats(1, 0, 100, null, null, 100))); + + Optional> result = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] {TestPredicates.isNotNull("x")}, stats, Set.of(0, 1, 2)); + + assertTrue(result.isPresent()); + assertEquals(Set.of(0, 2), result.get()); + } + + @Test + public void testPartialCoverageIsAppliedBeforeAndOrAcrossColumns() { + Map> stats = new HashMap<>(); + stats.put( + "x", + Arrays.asList( + new ZoneStats(0, 0, 100, 0L, 99L, 0), new ZoneStats(1, 0, 100, 100L, 199L, 0))); + stats.put( + "y", + Arrays.asList( + new ZoneStats(1, 0, 100, 1000L, 1099L, 0), new ZoneStats(2, 0, 100, 2000L, 2099L, 0))); + + Optional> andResult = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] { + TestPredicates.and(TestPredicates.eq("x", 50L), TestPredicates.eq("y", 2050L)) + }, + stats, + Set.of(0, 1, 2)); + Optional> orResult = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] { + TestPredicates.or(TestPredicates.eq("x", 150L), TestPredicates.eq("y", 1050L)) + }, + stats, + Set.of(0, 1, 2)); + + assertTrue(andResult.isPresent()); + assertEquals(Set.of(0, 2), andResult.get()); + assertTrue(orResult.isPresent()); + assertEquals(Set.of(0, 1, 2), orResult.get()); + } + + @Test + public void testRetiredFragmentStatsAreIgnored() { + Map> stats = new HashMap<>(); + stats.put( + "x", + Arrays.asList( + new ZoneStats(0, 0, 100, 0L, 99L, 0), new ZoneStats(9, 0, 100, 500L, 599L, 0))); + + Optional> result = + ZonemapFragmentPruner.pruneFragments( + new Predicate[] {TestPredicates.eq("x", 550L)}, stats, Set.of(0, 1)); + + assertTrue(result.isPresent()); + assertEquals(Set.of(1), result.get()); + } } diff --git a/lance-spark-base_2.12/src/test/java/org/lance/spark/sharding/SparkLanceShardingUtilsTest.java b/lance-spark-base_2.12/src/test/java/org/lance/spark/sharding/SparkLanceShardingUtilsTest.java index 776dbd7b5..fd06633a9 100644 --- a/lance-spark-base_2.12/src/test/java/org/lance/spark/sharding/SparkLanceShardingUtilsTest.java +++ b/lance-spark-base_2.12/src/test/java/org/lance/spark/sharding/SparkLanceShardingUtilsTest.java @@ -13,6 +13,7 @@ */ package org.lance.spark.sharding; +import org.lance.index.scalar.ZoneStats; import org.lance.memwal.ShardingField; import org.lance.memwal.ShardingSpec; @@ -21,8 +22,12 @@ import org.apache.spark.sql.connector.expressions.Transform; import org.junit.jupiter.api.Test; +import java.util.Arrays; import java.util.Collections; +import java.util.List; +import java.util.Map; import java.util.Optional; +import java.util.Set; import static org.junit.jupiter.api.Assertions.*; @@ -111,4 +116,56 @@ public void testSourceIdBackedFieldRequiresLanceSchema() { () -> SparkLanceShardingUtils.toSparkExpression(field, null)); assertTrue(error.getMessage().contains("requires Lance schema")); } + + @Test + public void testDetectFragmentKeysRequiresFullLiveCoverage() { + ShardingField field = + SparkLanceShardingUtils.fromSparkTransforms( + new Transform[] {Expressions.identity("region")}) + .fields() + .get(0); + List zones = + Arrays.asList( + new ZoneStats(0, 0, 10, "east", "east", 0), new ZoneStats(1, 0, 10, "west", "west", 0)); + + Optional> keys = + SparkLanceShardingUtils.detectFragmentKeys(field, null, zones, Set.of(0, 1)); + + assertTrue(keys.isPresent()); + assertEquals(Map.of(0, "east", 1, "west"), keys.get()); + } + + @Test + public void testDetectFragmentKeysRejectsPartialLiveCoverage() { + ShardingField field = + SparkLanceShardingUtils.fromSparkTransforms( + new Transform[] {Expressions.identity("region")}) + .fields() + .get(0); + List zones = Collections.singletonList(new ZoneStats(0, 0, 10, "east", "east", 0)); + + Optional> keys = + SparkLanceShardingUtils.detectFragmentKeys(field, null, zones, Set.of(0, 1)); + + assertFalse(keys.isPresent()); + } + + @Test + public void testDetectFragmentKeysIgnoresRetiredFragmentStats() { + ShardingField field = + SparkLanceShardingUtils.fromSparkTransforms( + new Transform[] {Expressions.identity("region")}) + .fields() + .get(0); + List zones = + Arrays.asList( + new ZoneStats(0, 0, 10, "east", "east", 0), + new ZoneStats(9, 0, 10, "retired", "retired", 0)); + + Optional> keys = + SparkLanceShardingUtils.detectFragmentKeys(field, null, zones, Set.of(0)); + + assertTrue(keys.isPresent()); + assertEquals(Collections.singletonMap(0, "east"), keys.get()); + } }