Skip to content

Commit 5c1eada

Browse files
authored
refactor(rust/sedona-raster-functions): adopt start_raster_from / copy_into in the structural raster functions (#1193)
1 parent 0376f91 commit 5c1eada

3 files changed

Lines changed: 59 additions & 87 deletions

File tree

rust/sedona-raster-functions/src/rs_dim_band.rs

Lines changed: 42 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -22,10 +22,9 @@ use datafusion_common::cast::as_string_view_array;
2222
use datafusion_common::error::Result;
2323
use datafusion_common::exec_err;
2424
use datafusion_expr::{ColumnarValue, Volatility};
25-
use sedona_common::sedona_internal_datafusion_err;
2625
use sedona_expr::scalar_udf::{SedonaScalarKernel, SedonaScalarUDF};
27-
use sedona_raster::builder::{RasterBuilder, StartBandArgs};
28-
use sedona_raster::traits::RasterRef;
26+
use sedona_raster::builder::{RasterBuilder, RasterOverrides, StartBandArgs};
27+
use sedona_raster::traits::{BandOverrides, RasterRef};
2928
use sedona_schema::datatypes::SedonaType;
3029
use sedona_schema::matchers::ArgMatcher;
3130

@@ -90,16 +89,7 @@ impl SedonaScalarKernel for RsDimToBand {
9089

9190
require_any_band_has_dim(raster, name, "RS_DimToBand")?;
9291

93-
let t: [f64; 6] = raster.transform().try_into().map_err(|_| {
94-
sedona_internal_datafusion_err!("raster transform is not 6 elements")
95-
})?;
96-
let spatial_dims = raster.spatial_dims();
97-
new_builder.start_raster_nd(
98-
&t,
99-
&spatial_dims,
100-
raster.spatial_shape(),
101-
raster.crs(),
102-
)?;
92+
new_builder.start_raster_from(raster, RasterOverrides::default())?;
10393

10494
for band_idx in 0..raster.num_bands() {
10595
let band = raster.band(band_idx)?;
@@ -108,16 +98,7 @@ impl SedonaScalarKernel for RsDimToBand {
10898
match maybe_dim_idx {
10999
None => {
110100
// Band doesn't have this dimension -- pass through
111-
let dim_names = band.dim_names();
112-
let band_name = raster.band_name(band_idx);
113-
new_builder.start_band(StartBandArgs {
114-
name: band_name,
115-
nodata: band.nodata(),
116-
..StartBandArgs::new(&dim_names, band.shape(), band.data_type())
117-
})?;
118-
let ndb = band.nd_buffer()?;
119-
let data = ndb.as_contiguous()?;
120-
new_builder.band_data_writer().append_value(data);
101+
band.copy_into(&mut new_builder, BandOverrides::default())?;
121102
new_builder.finish_band()?;
122103
}
123104
Some(dim_idx) => {
@@ -290,16 +271,7 @@ impl SedonaScalarKernel for RsBandToDim {
290271

291272
let nodata = ref_nodata.as_deref();
292273

293-
let t: [f64; 6] = raster.transform().try_into().map_err(|_| {
294-
sedona_internal_datafusion_err!("raster transform is not 6 elements")
295-
})?;
296-
let spatial_dims = raster.spatial_dims();
297-
new_builder.start_raster_nd(
298-
&t,
299-
&spatial_dims,
300-
raster.spatial_shape(),
301-
raster.crs(),
302-
)?;
274+
new_builder.start_raster_from(raster, RasterOverrides::default())?;
303275
new_builder.start_band(StartBandArgs {
304276
nodata,
305277
..StartBandArgs::new(&new_dim_names, &new_shape, ref_data_type)
@@ -379,6 +351,43 @@ mod tests {
379351
assert_eq!(raster.band_name(2), Some("temp_time_2"));
380352
}
381353

354+
#[test]
355+
fn dimtoband_passes_through_bands_without_dim_preserving_names() {
356+
// Heterogeneous raster: the band carrying `time` expands into one band
357+
// per index, while the band without it is passed through untouched.
358+
// The pass-through derives the output band via `BandRef::copy_into`,
359+
// which inherits the name — assert it survives, since a dropped name is
360+
// silent otherwise.
361+
let udf: ScalarUDF = rs_dimtoband_udf().into();
362+
let tester = ScalarUdfTester::new(udf, vec![RASTER, SedonaType::Arrow(DataType::Utf8)]);
363+
364+
let rasters = RasterSpec::nd(&["time", "y", "x"], &[2, 2, 2])
365+
.crs(None)
366+
.band_nd(&["y", "x"], &[2, 2], BandDataType::UInt8)
367+
.name("elevation")
368+
.band(BandDataType::UInt8)
369+
.name("temperature")
370+
.build();
371+
372+
let result = tester
373+
.invoke_array_scalar(Arc::new(rasters), "time")
374+
.unwrap();
375+
376+
let result_struct = result.as_any().downcast_ref::<StructArray>().unwrap();
377+
let raster_array = RasterStructArray::try_new(result_struct).unwrap();
378+
let raster = raster_array.get(0).unwrap();
379+
380+
// Pass-through band keeps its own name; the expanded band contributes
381+
// one suffixed band per `time` index.
382+
assert_eq!(raster.num_bands(), 3);
383+
assert_eq!(raster.band_name(0), Some("elevation"));
384+
assert_eq!(raster.band_name(1), Some("temperature_time_0"));
385+
assert_eq!(raster.band_name(2), Some("temperature_time_1"));
386+
// The pass-through band is emitted unchanged, not expanded.
387+
assert_eq!(raster.band(0).unwrap().dim_names(), vec!["y", "x"]);
388+
assert_eq!(raster.band(0).unwrap().shape(), &[2, 2]);
389+
}
390+
382391
#[test]
383392
fn dimtoband_null_raster() {
384393
let udf: ScalarUDF = rs_dimtoband_udf().into();

rust/sedona-raster-functions/src/rs_ensure_loaded.rs

Lines changed: 3 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -40,7 +40,7 @@ use datafusion_expr::{
4040
};
4141
use sedona_common::{sedona_internal_datafusion_err, sedona_internal_err};
4242
use sedona_raster::array::RasterStructArray;
43-
use sedona_raster::builder::{RasterBuilder, StartBandArgs};
43+
use sedona_raster::builder::{RasterBuilder, RasterOverrides, StartBandArgs};
4444
use sedona_raster::raster_loader::{
4545
AsyncRasterLoader, RasterLoadRequest, RasterLoaderConfig, RasterLoaderRegistry,
4646
};
@@ -294,27 +294,11 @@ where
294294
)
295295
})?;
296296

297-
// Owned per-row metadata so the borrows don't span the per-band
298-
// `await` points further down.
299-
let transform: [f64; 6] = raster.transform().try_into().map_err(|_| {
300-
sedona_internal_datafusion_err!(
301-
"RS_EnsureLoaded: raster row {raster_idx} transform is not 6 elements"
302-
)
303-
})?;
304-
let spatial_dims_owned: Vec<String> = raster
305-
.spatial_dims()
306-
.iter()
307-
.map(|s| s.to_string())
308-
.collect();
309-
let spatial_dims: Vec<&str> = spatial_dims_owned.iter().map(String::as_str).collect();
310-
let spatial_shape: Vec<i64> = raster.spatial_shape().to_vec();
311-
let crs: Option<String> = raster.crs().map(|s| s.to_string());
312-
313297
builder
314-
.start_raster_nd(&transform, &spatial_dims, &spatial_shape, crs.as_deref())
298+
.start_raster_from(&raster, RasterOverrides::default())
315299
.map_err(|e| {
316300
sedona_internal_datafusion_err!(
317-
"RS_EnsureLoaded: start_raster_nd failed at row {raster_idx}: {e}"
301+
"RS_EnsureLoaded: start_raster_from failed at row {raster_idx}: {e}"
318302
)
319303
})?;
320304

rust/sedona-raster-functions/src/rs_slice.rs

Lines changed: 14 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -22,10 +22,9 @@ use datafusion_common::cast::{as_int64_array, as_string_array};
2222
use datafusion_common::error::Result;
2323
use datafusion_common::exec_err;
2424
use datafusion_expr::{ColumnarValue, Volatility};
25-
use sedona_common::sedona_internal_datafusion_err;
2625
use sedona_expr::scalar_udf::{SedonaScalarKernel, SedonaScalarUDF};
27-
use sedona_raster::builder::{RasterBuilder, StartBandArgs};
28-
use sedona_raster::traits::{BandRef, RasterRef};
26+
use sedona_raster::builder::{RasterBuilder, RasterOverrides, StartBandArgs};
27+
use sedona_raster::traits::{BandOverrides, BandRef, RasterRef};
2928
use sedona_schema::datatypes::SedonaType;
3029
use sedona_schema::matchers::ArgMatcher;
3130

@@ -98,11 +97,7 @@ impl SedonaScalarKernel for RsSlice {
9897
}
9998
validate_not_spatial(raster, name, "RS_Slice")?;
10099

101-
let t: [f64; 6] = raster.transform().try_into().map_err(|_| {
102-
sedona_internal_datafusion_err!("raster transform is not 6 elements")
103-
})?;
104-
let spatial_dims = raster.spatial_dims();
105-
new_builder.start_raster_nd(&t, &spatial_dims, raster.spatial_shape(), raster.crs())?;
100+
new_builder.start_raster_from(raster, RasterOverrides::default())?;
106101

107102
require_any_band_has_dim(raster, name, "RS_Slice")?;
108103

@@ -114,16 +109,7 @@ impl SedonaScalarKernel for RsSlice {
114109
// RS_DimToBand, and matches xarray's `isel` — variables
115110
// without the indexed dim are left alone.
116111
let Some(dim_idx) = band.dim_index(name) else {
117-
let dim_names = band.dim_names();
118-
let band_name = raster.band_name(band_idx);
119-
new_builder.start_band(StartBandArgs {
120-
name: band_name,
121-
nodata: band.nodata(),
122-
..StartBandArgs::new(&dim_names, band.shape(), band.data_type())
123-
})?;
124-
let ndb = band.nd_buffer()?;
125-
let data = ndb.as_contiguous()?;
126-
new_builder.band_data_writer().append_value(data);
112+
band.copy_into(&mut new_builder, BandOverrides::default())?;
127113
new_builder.finish_band()?;
128114
continue;
129115
};
@@ -257,11 +243,7 @@ impl SedonaScalarKernel for RsSliceRange {
257243
);
258244
}
259245

260-
let t: [f64; 6] = raster.transform().try_into().map_err(|_| {
261-
sedona_internal_datafusion_err!("raster transform is not 6 elements")
262-
})?;
263-
let spatial_dims = raster.spatial_dims();
264-
new_builder.start_raster_nd(&t, &spatial_dims, raster.spatial_shape(), raster.crs())?;
246+
new_builder.start_raster_from(raster, RasterOverrides::default())?;
265247

266248
require_any_band_has_dim(raster, name, "RS_SliceRange")?;
267249

@@ -272,16 +254,7 @@ impl SedonaScalarKernel for RsSliceRange {
272254
// dimension are emitted unchanged. Same convention as
273255
// RS_Slice and RS_DimToBand.
274256
let Some(dim_idx) = band.dim_index(name) else {
275-
let dim_names = band.dim_names();
276-
let band_name = raster.band_name(band_idx);
277-
new_builder.start_band(StartBandArgs {
278-
name: band_name,
279-
nodata: band.nodata(),
280-
..StartBandArgs::new(&dim_names, band.shape(), band.data_type())
281-
})?;
282-
let ndb = band.nd_buffer()?;
283-
let data = ndb.as_contiguous()?;
284-
new_builder.band_data_writer().append_value(data);
257+
band.copy_into(&mut new_builder, BandOverrides::default())?;
285258
new_builder.finish_band()?;
286259
continue;
287260
};
@@ -420,7 +393,9 @@ mod tests {
420393
RasterSpec::nd(&["time", "y", "x"], &[3, 2, 3])
421394
.crs(None)
422395
.band_nd(&["y", "x"], &[2, 3], BandDataType::UInt8)
396+
.name("elevation")
423397
.band(BandDataType::UInt8)
398+
.name("temperature")
424399
.build()
425400
}
426401

@@ -449,7 +424,9 @@ mod tests {
449424
let expected = RasterSpec::nd(&["time", "y", "x"], &[3, 2, 3])
450425
.crs(None)
451426
.band_values_nd(&["y", "x"], &[2, 3], &(0u8..6).collect::<Vec<u8>>())
452-
.band_values_nd(&["y", "x"], &[2, 3], &(6u8..12).collect::<Vec<u8>>());
427+
.name("elevation")
428+
.band_values_nd(&["y", "x"], &[2, 3], &(6u8..12).collect::<Vec<u8>>())
429+
.name("temperature");
453430
assert_rasters_equal(&result, &[Some(expected)]);
454431
}
455432

@@ -481,11 +458,13 @@ mod tests {
481458
let expected = RasterSpec::nd(&["time", "y", "x"], &[3, 2, 3])
482459
.crs(None)
483460
.band_values_nd(&["y", "x"], &[2, 3], &(0u8..6).collect::<Vec<u8>>())
461+
.name("elevation")
484462
.band_values_nd(
485463
&["time", "y", "x"],
486464
&[2, 2, 3],
487465
&(6u8..18).collect::<Vec<u8>>(),
488-
);
466+
)
467+
.name("temperature");
489468
assert_rasters_equal(&result, &[Some(expected)]);
490469
}
491470

0 commit comments

Comments
 (0)