viirya commented on code in PR #6534:
URL: https://github.com/apache/datafusion-comet/pull/6534#discussion_r4178440615


##########
native/core/src/execution/jni_api.rs:
##########
@@ -3000,4 +3069,295 @@ mod tests {
         assert!(producer.next_batch().unwrap().is_none());
         producer.stop().unwrap();
     }
+
+    #[test]
+    fn sort_row_partitions_sorts_in_place() {
+        let mut pointers = vec![
+            5_i64 << 40 | 3,
+            1,
+            2_i64 << 40,
+            1_i64 << 40 | 7,
+            0,
+            5_i64 << 40,
+        ];
+        let mut expected = pointers.clone();
+        // Stable, so the lower 40 bits keep their input order within a 
partition.
+        expected.sort_by_key(|pointer| (*pointer as u64) >> 40);
+        sort_row_partitions(&mut pointers);
+        assert_eq!(pointers, expected);
+        assert_eq!(&pointers[..2], &[1, 0]);
+    }
+
+    #[test]
+    fn shuffle_partition_offsets_requires_an_executed_shuffle_write() {
+        let error = shuffle_partition_offsets(None).unwrap_err();
+        assert!(
+            error
+                .to_string()
+                .contains("before the plan has been executed"),
+            "{error}"
+        );
+
+        let plan = 
datafusion::physical_plan::empty::EmptyExec::new(Arc::new(Schema::empty()));
+        let error = shuffle_partition_offsets(Some(&plan)).unwrap_err();
+        assert!(
+            error
+                .to_string()
+                .contains("only available on a native shuffle write plan"),
+            "{error}"
+        );
+    }
+
+    #[test]
+    fn shuffle_partition_offsets_after_the_write_is_drained() {
+        use datafusion::datasource::memory::MemorySourceConfig;
+        use datafusion::datasource::source::DataSourceExec;
+        use datafusion_comet_shuffle::{CometPartitioning, RoundRobinStrategy};
+
+        let batch = int_batch();
+        let num_partitions = 3;
+        let dir = tempfile::tempdir().unwrap();
+        let data_file = dir.path().join("data.out");
+        let writer = ShuffleWriterExec::try_new(
+            Arc::new(DataSourceExec::new(Arc::new(
+                MemorySourceConfig::try_new(&[vec![batch.clone()]], 
batch.schema(), None).unwrap(),
+            ))),
+            CometPartitioning::RoundRobin(num_partitions, 
RoundRobinStrategy::default()),
+            CompressionCodec::Zstd(1),
+            data_file.to_str().unwrap().to_string(),
+            false,
+            1024 * 1024,
+            None,
+        )
+        .unwrap();
+
+        let error = shuffle_partition_offsets(Some(&writer)).unwrap_err();
+        assert!(
+            error.to_string().contains("not drained to completion"),
+            "{error}"
+        );
+
+        let task_ctx = Arc::new(TaskContext::default());
+        let stream = writer.execute(0, task_ctx).unwrap();
+        Runtime::new()
+            .unwrap()
+            .block_on(datafusion::physical_plan::common::collect(stream))
+            .unwrap();
+
+        // One offset per partition plus the data file length.
+        let offsets = shuffle_partition_offsets(Some(&writer)).unwrap();
+        assert_eq!(offsets.len(), num_partitions + 1);
+        assert_eq!(offsets[0], 0);
+        assert!(offsets.windows(2).all(|pair| pair[0] <= pair[1]));
+        let file_len = std::fs::metadata(&data_file).unwrap().len() as i64;
+        assert_eq!(offsets[num_partitions], file_len);
+    }
+
+    /// Spark `UnsafeRow`s with one non-null `long` field each: an 8-byte null 
bitset, then the
+    /// value.
+    fn long_rows(values: &[i64]) -> Vec<[u8; 16]> {
+        values
+            .iter()
+            .map(|value| {
+                let mut row = [0_u8; 16];
+                row[8..].copy_from_slice(&value.to_le_bytes());
+                row
+            })
+            .collect()
+    }
+
+    /// Splits a shuffle data file into its blocks, each without the 8-byte 
length and the 8-byte
+    /// field count, so it starts at the codec tag as `read_ipc_compressed` 
expects.
+    fn shuffle_blocks(data: &[u8]) -> Vec<&[u8]> {
+        let mut blocks = Vec::new();
+        let mut pos = 0;
+        while pos < data.len() {
+            let length = u64::from_le_bytes(data[pos..pos + 
8].try_into().unwrap()) as usize;
+            blocks.push(&data[pos + 16..pos + 8 + length]);
+            pos += 8 + length;
+        }
+        blocks
+    }
+
+    /// Writes `rows` (16-byte `long` rows) with `write_sorted_file` in 
batches of 2 rows and
+    /// returns the results and the file contents.
+    fn write_long_rows(
+        rows: &[[u8; 16]],
+        sizes: &[i32],
+        path: &std::path::Path,
+        checksum_enabled: bool,
+        current_checksum: i64,
+        codec: &str,
+    ) -> CometResult<([i64; 3], Vec<u8>)> {
+        let addresses: Vec<i64> = rows.iter().map(|row| row.as_ptr() as 
i64).collect();
+        let results = unsafe {
+            write_sorted_file(
+                &addresses,
+                sizes,
+                &[DataType::Int64],
+                path.to_str().unwrap().to_string(),
+                1.0,
+                2,
+                checksum_enabled,
+                0,
+                current_checksum,
+                codec,
+                1,
+            )
+        }?;
+        Ok((results, std::fs::read(path).unwrap()))
+    }
+
+    #[test]
+    fn write_sorted_file_codecs() {
+        let values = [1_i64, 2, 3, 4, 5];
+        let rows = long_rows(&values);
+        let sizes = vec![16_i32; rows.len()];
+        let dir = tempfile::tempdir().unwrap();
+
+        // Unknown codec names fall back to lz4.
+        for (codec, tag) in [
+            ("lz4", b"LZ4_"),
+            ("zstd", b"ZSTD"),
+            ("snappy", b"SNAP"),
+            ("unknown", b"LZ4_"),
+        ] {
+            let path = dir.path().join(codec);
+            let ([written, checksum, _], data) =
+                write_long_rows(&rows, &sizes, &path, false, i64::MIN, 
codec).unwrap();
+            assert_eq!(written, data.len() as i64, "{codec}");
+            assert_eq!(checksum, i64::MIN, "no checksum without 
checksum_enabled");
+
+            let blocks = shuffle_blocks(&data);
+            assert_eq!(blocks.len(), 3, "{codec}: 5 rows in batches of 2");
+            let mut decoded = Vec::new();
+            for block in blocks {
+                assert_eq!(&block[..4], tag, "{codec}");
+                let batch = read_ipc_compressed(block).unwrap();
+                let column = 
arrow::array::AsArray::as_primitive::<arrow::datatypes::Int64Type>(
+                    batch.column(0),
+                );
+                decoded.extend(column.values().iter().copied());
+            }
+            assert_eq!(decoded, values, "{codec}");
+        }
+    }
+
+    #[test]
+    fn write_sorted_file_checksums_the_written_bytes() {

Review Comment:
   Good catch, thanks. `write_sorted_file_checksums_the_written_bytes` now runs 
for both `checksum_algo` 0 (CRC32) and 1 (Adler-32), each with a fresh checksum 
and one continued from the previous file. The Adler-32 expectation comes from a 
small reference in the test, as you described. I checked that the test fails 
for Adler-32 (and still passes for CRC32) when the `i64::MIN` branch is removed.
   



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]


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

Reply via email to