Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

# Minarrow Rust

## Unreleased

## Added
* Scatter indices as write counterpart to gather indices for mask-based writes.

## 0.18.3 - 2026-10-01 Nightly Compatibility Release

## Added
Expand Down
332 changes: 332 additions & 0 deletions src/enums/array.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1438,6 +1438,237 @@ impl Array {
self.view(0, self.len()).gather_mask(mask)
}

/// Scatter elements of `src` into this array at the given positions.
///
/// Writes element `src_idx[i]` of `src` to position `dst_idx[i]` of this
/// array, including its null bit, and leaves every other position
/// unchanged. This is the in-place counterpart to
/// [`gather_indices`](Self::gather_indices). Source indices are relative
/// to the `src` window, and a repeated destination position takes the
/// last write.
///
/// ## Behaviour
///
/// - Fixed-width variants write into their typed buffers in place, so
/// the cost scales with the number of indices.
/// - String arrays rebuild their offsets and data once per call, so the
/// cost scales with the array length.
/// - Categorical arrays intern each written value into this array's
/// dictionary.
/// - This array gains a null mask only when `src` has one.
///
/// Returns an error when `src` differs in array type, datetime unit, or
/// decimal precision and scale, when `dst_idx` and `src_idx` differ in
/// length, or when the array is `Null`. Panics when an index is out of
/// bounds. Mutation is copy-on-write.
#[cfg(all(feature = "views", feature = "select"))]
pub fn scatter_indices(
&mut self,
dst_idx: &[usize],
src: &ArrayV,
src_idx: &[usize],
) -> Result<(), MinarrowError> {
if dst_idx.len() != src_idx.len() {
return Err(MinarrowError::ShapeError {
message: format!(
"scatter_indices: {} destination indices for {} source indices",
dst_idx.len(),
src_idx.len()
),
});
}
let dst_len = self.len();
let offset = src.offset;
let window_len = src.len();
assert!(
dst_idx.iter().all(|&d| d < dst_len),
"scatter_indices: destination index out of bounds"
);
assert!(
src_idx.iter().all(|&s| s < window_len),
"scatter_indices: source index outside the window"
);

let dtype = self.arrow_type();
let mismatch = || MinarrowError::TypeError {
from: "scatter source",
to: "array",
message: Some(format!(
"cannot scatter {:?} into {:?}",
src.array.arrow_type(),
dtype
)),
};

// Writes the null bit of each scattered element. A source without a
// mask marks every written position valid.
macro_rules! scatter_mask {
($dst_mask:expr, $src_mask:expr) => {{
match $src_mask {
Some(sm) => {
let mask = $dst_mask.get_or_insert_with(|| Bitmask::new_set_all(dst_len, true));
for (&d, &s) in dst_idx.iter().zip(src_idx) {
mask.set(d, sm.get(offset + s));
}
}
None => {
if let Some(mask) = $dst_mask.as_mut() {
for &d in dst_idx {
mask.set(d, true);
}
}
}
}
}};
}

// Copies primitive-typed elements from the source window into the
// typed buffer with one indexed write per element.
macro_rules! scatter_prim {
($dst:expr, $src:expr) => {{
let dst = Arc::make_mut($dst);
let src_data = &$src.data.as_slice()[offset..offset + window_len];
let data = dst.data.as_mut_slice();
for (&d, &s) in dst_idx.iter().zip(src_idx) {
data[d] = src_data[s];
}
scatter_mask!(dst.null_mask, $src.null_mask.as_ref());
}};
}

// Rebuilds a string array once, reading each position from the
// source window when scattered and from this array otherwise.
macro_rules! scatter_str {
($dst:expr, $ArrTy:ty) => {{
let rebuilt = {
let mut source: Vec<Option<usize>> = vec![None; dst_len];
for (&d, &s) in dst_idx.iter().zip(src_idx) {
source[d] = Some(s);
}
let mut values: Vec<&str> = Vec::with_capacity(dst_len);
let mut null_at: Vec<usize> = Vec::new();
for (i, pos) in source.iter().enumerate() {
let val = match pos {
Some(s) => src.get_str(*s),
None => $dst.get_str(i),
};
match val {
Some(v) => values.push(v),
None => {
null_at.push(i);
values.push("");
}
}
}
let mut new_arr = <$ArrTy>::from_vec(values, None);
for &i in &null_at {
new_arr.set_null(i);
}
new_arr
};
*$dst = Arc::new(rebuilt);
}};
}

// Interns each scattered value into the categorical dictionary.
macro_rules! scatter_cat {
($dst:expr) => {{
let dst = Arc::make_mut($dst);
for (&d, &s) in dst_idx.iter().zip(src_idx) {
match src.get_str(s) {
Some(v) => dst.set_str(d, v),
None => dst
.null_mask
.get_or_insert_with(|| Bitmask::new_set_all(dst_len, true))
.set(d, false),
}
}
}};
}

match (self, &src.array) {
(Array::NumericArray(dst), Array::NumericArray(src_num)) => match (dst, src_num) {
(NumericArray::Int32(d), NumericArray::Int32(s)) => scatter_prim!(d, s),
(NumericArray::Int64(d), NumericArray::Int64(s)) => scatter_prim!(d, s),
(NumericArray::Float32(d), NumericArray::Float32(s)) => scatter_prim!(d, s),
(NumericArray::Float64(d), NumericArray::Float64(s)) => scatter_prim!(d, s),
(NumericArray::UInt32(d), NumericArray::UInt32(s)) => scatter_prim!(d, s),
(NumericArray::UInt64(d), NumericArray::UInt64(s)) => scatter_prim!(d, s),
#[cfg(feature = "extended_numeric_types")]
(NumericArray::Int8(d), NumericArray::Int8(s)) => scatter_prim!(d, s),
#[cfg(feature = "extended_numeric_types")]
(NumericArray::Int16(d), NumericArray::Int16(s)) => scatter_prim!(d, s),
#[cfg(feature = "extended_numeric_types")]
(NumericArray::UInt8(d), NumericArray::UInt8(s)) => scatter_prim!(d, s),
#[cfg(feature = "extended_numeric_types")]
(NumericArray::UInt16(d), NumericArray::UInt16(s)) => scatter_prim!(d, s),
#[cfg(feature = "decimal")]
(NumericArray::Decimal32(d), NumericArray::Decimal32(s))
if d.precision == s.precision && d.scale == s.scale =>
{
scatter_prim!(d, s)
}
#[cfg(feature = "decimal")]
(NumericArray::Decimal64(d), NumericArray::Decimal64(s))
if d.precision == s.precision && d.scale == s.scale =>
{
scatter_prim!(d, s)
}
#[cfg(feature = "decimal")]
(NumericArray::Decimal128(d), NumericArray::Decimal128(s))
if d.precision == s.precision && d.scale == s.scale =>
{
scatter_prim!(d, s)
}
_ => return Err(mismatch()),
},
(Array::TextArray(dst), Array::TextArray(src_text)) => match (dst, src_text) {
(TextArray::String32(d), TextArray::String32(_)) => {
scatter_str!(d, StringArray<u32>)
}
#[cfg(feature = "large_string")]
(TextArray::String64(d), TextArray::String64(_)) => {
scatter_str!(d, StringArray<u64>)
}
#[cfg(feature = "default_categorical_8")]
(TextArray::Categorical8(d), TextArray::Categorical8(_)) => scatter_cat!(d),
#[cfg(feature = "extended_categorical")]
(TextArray::Categorical16(d), TextArray::Categorical16(_)) => scatter_cat!(d),
#[cfg(any(
not(feature = "default_categorical_8"),
feature = "extended_categorical"
))]
(TextArray::Categorical32(d), TextArray::Categorical32(_)) => scatter_cat!(d),
#[cfg(feature = "extended_categorical")]
(TextArray::Categorical64(d), TextArray::Categorical64(_)) => scatter_cat!(d),
_ => return Err(mismatch()),
},
(Array::BooleanArray(d), Array::BooleanArray(s)) => {
let dst = Arc::make_mut(d);
for (&di, &si) in dst_idx.iter().zip(src_idx) {
dst.data.set(di, s.data.get(offset + si));
}
scatter_mask!(dst.null_mask, s.null_mask.as_ref());
}
#[cfg(feature = "datetime")]
(Array::TemporalArray(dst), Array::TemporalArray(src_temp)) => match (dst, src_temp) {
(TemporalArray::Datetime32(d), TemporalArray::Datetime32(s))
if d.time_unit == s.time_unit =>
{
scatter_prim!(d, s)
}
(TemporalArray::Datetime64(d), TemporalArray::Datetime64(s))
if d.time_unit == s.time_unit =>
{
scatter_prim!(d, s)
}
_ => return Err(mismatch()),
},
_ => return Err(mismatch()),
}
Ok(())
}

/// Returns a reference to the inner array as type `Arc<T>`.
///
/// This is compile-time safe if `T` matches the actual payload, but will panic otherwise.
Expand Down Expand Up @@ -7501,3 +7732,104 @@ mod arr_macro_extensions_tests {
}
}
}

#[cfg(all(test, feature = "views", feature = "select"))]
mod scatter_tests {
#[test]
fn scatter_f64_writes_values_and_null_bits() {
let mut dst = arr_f64_opt![Some(1.0), Some(2.0), Some(3.0), Some(4.0)];
let src = arr_f64_opt![Some(10.0), None::<f64>, Some(30.0)];
dst.scatter_indices(&[3, 1], &src.view(0, 3), &[0, 1]).unwrap();
assert_eq!(dst, arr_f64_opt![Some(1.0), None::<f64>, Some(3.0), Some(10.0)]);
}

#[test]
fn scatter_reads_relative_to_source_window() {
let mut dst = arr_i64![0, 0, 0];
let src = arr_i64![1, 2, 3, 4, 5];
dst.scatter_indices(&[0, 2], &src.view(2, 3), &[0, 2]).unwrap();
assert_eq!(dst, arr_i64![3, 0, 5]);
}

#[test]
fn scatter_repeated_destination_takes_last_write() {
let mut dst = arr_i64![0, 0];
let src = arr_i64![7, 8, 9];
dst.scatter_indices(&[1, 1, 1], &src.view(0, 3), &[0, 1, 2]).unwrap();
assert_eq!(dst, arr_i64![0, 9]);
}

#[test]
fn scatter_source_without_mask_marks_written_positions_valid() {
let mut dst = arr_f64_opt![None::<f64>, None::<f64>, Some(3.0)];
let src = arr_f64![&[5.0]];
dst.scatter_indices(&[0], &src.view(0, 1), &[0]).unwrap();
assert!(!dst.null_mask().unwrap().get(1));
assert!(dst.null_mask().unwrap().get(0));
assert_eq!(dst.num().f64().data.as_slice()[0], 5.0);
}

#[test]
fn scatter_is_copy_on_write() {
let original = arr_i64![1, 2, 3];
let mut dst = original.clone();
let src = arr_i64![&[9]];
dst.scatter_indices(&[0], &src.view(0, 1), &[0]).unwrap();
assert_eq!(dst, arr_i64![9, 2, 3]);
assert_eq!(original, arr_i64![1, 2, 3]);
}

#[test]
fn scatter_boolean() {
let mut dst = arr_bool![false, false, false];
let src = arr_bool_opt![Some(true), None::<bool>];
dst.scatter_indices(&[2, 0], &src.view(0, 2), &[0, 1]).unwrap();
assert!(dst.bool().data.get(2));
assert!(dst.null_mask().unwrap().get(2));
assert!(!dst.null_mask().unwrap().get(0));
assert!(dst.null_mask().unwrap().get(1));
}

#[test]
fn scatter_string_rebuilds_with_varying_lengths() {
let mut dst = arr_str32!["a", "bb", "ccc"];
let src = arr_str32_opt![Some("long value"), None::<&str>];
dst.scatter_indices(&[0, 1], &src.view(0, 2), &[0, 1]).unwrap();
assert_eq!(dst.get_str(0), Some("long value"));
assert_eq!(dst.get_str(1), None);
assert_eq!(dst.get_str(2), Some("ccc"));
}

#[cfg(any(not(feature = "default_categorical_8"), feature = "extended_categorical"))]
#[test]
fn scatter_categorical_interns_new_values() {
let mut dst = arr_cat32!["x", "y", "x"];
let src = arr_cat32_opt![Some("z"), None::<&str>];
dst.scatter_indices(&[2, 0], &src.view(0, 2), &[0, 1]).unwrap();
assert_eq!(dst.get_str(0), None);
assert_eq!(dst.get_str(1), Some("y"));
assert_eq!(dst.get_str(2), Some("z"));
}

#[test]
fn scatter_type_mismatch_errors() {
let mut dst = arr_i64![1, 2];
let src = arr_f64![&[1.0]];
assert!(dst.scatter_indices(&[0], &src.view(0, 1), &[0]).is_err());
}

#[test]
fn scatter_length_mismatch_errors() {
let mut dst = arr_i64![1, 2];
let src = arr_i64![1, 2];
assert!(dst.scatter_indices(&[0, 1], &src.view(0, 2), &[0]).is_err());
}

#[test]
#[should_panic(expected = "destination index out of bounds")]
fn scatter_destination_out_of_bounds_panics() {
let mut dst = arr_i64![1, 2];
let src = arr_i64![&[1]];
let _ = dst.scatter_indices(&[2], &src.view(0, 1), &[0]);
}
}
Loading