1#![allow(unsafe_code)]
3
4use std::collections::HashMap;
40use std::sync::Arc;
41
42use arrow_array::{
43 Array, ArrayRef, Float32Array, Float64Array, Int8Array, Int16Array, Int32Array, Int64Array,
44 RecordBatch, UInt8Array, UInt16Array, UInt32Array, UInt64Array,
45};
46use arrow_schema::{DataType, Field, Schema};
47
48use super::{OxiGeoError, RasterBuffer, Result};
49use crate::types::{NoDataValue, RasterDataType};
50
51impl TryFrom<&RasterBuffer> for RecordBatch {
54 type Error = OxiGeoError;
55
56 fn try_from(buf: &RasterBuffer) -> Result<Self> {
65 let n = (buf.width() * buf.height()) as usize;
66 let nodata = buf.nodata();
67
68 #[inline]
71 fn is_nd(nodata: NoDataValue, v: f64) -> bool {
72 match nodata.as_f64() {
73 None => false,
74 Some(nd) => {
75 if nd.is_nan() && v.is_nan() {
76 true
77 } else {
78 (nd - v).abs() < f64::EPSILON
79 }
80 }
81 }
82 }
83
84 let (array, arrow_dt): (ArrayRef, DataType) = match buf.data_type() {
85 RasterDataType::UInt8 => {
86 let vals: &[u8] = buf.as_slice::<u8>().map_err(|e| OxiGeoError::Internal {
87 message: e.to_string(),
88 })?;
89 let arr: UInt8Array = vals
90 .iter()
91 .map(|&v| {
92 if is_nd(nodata, f64::from(v)) {
93 None
94 } else {
95 Some(v)
96 }
97 })
98 .collect();
99 (Arc::new(arr), DataType::UInt8)
100 }
101 RasterDataType::Int8 => {
102 let vals: &[i8] = buf.as_slice::<i8>().map_err(|e| OxiGeoError::Internal {
103 message: e.to_string(),
104 })?;
105 let arr: Int8Array = vals
106 .iter()
107 .map(|&v| {
108 if is_nd(nodata, f64::from(v)) {
109 None
110 } else {
111 Some(v)
112 }
113 })
114 .collect();
115 (Arc::new(arr), DataType::Int8)
116 }
117 RasterDataType::UInt16 => {
118 let vals: &[u16] = buf.as_slice::<u16>().map_err(|e| OxiGeoError::Internal {
119 message: e.to_string(),
120 })?;
121 let arr: UInt16Array = vals
122 .iter()
123 .map(|&v| {
124 if is_nd(nodata, f64::from(v)) {
125 None
126 } else {
127 Some(v)
128 }
129 })
130 .collect();
131 (Arc::new(arr), DataType::UInt16)
132 }
133 RasterDataType::Int16 => {
134 let vals: &[i16] = buf.as_slice::<i16>().map_err(|e| OxiGeoError::Internal {
135 message: e.to_string(),
136 })?;
137 let arr: Int16Array = vals
138 .iter()
139 .map(|&v| {
140 if is_nd(nodata, f64::from(v)) {
141 None
142 } else {
143 Some(v)
144 }
145 })
146 .collect();
147 (Arc::new(arr), DataType::Int16)
148 }
149 RasterDataType::UInt32 => {
150 let vals: &[u32] = buf.as_slice::<u32>().map_err(|e| OxiGeoError::Internal {
151 message: e.to_string(),
152 })?;
153 let arr: UInt32Array = vals
154 .iter()
155 .map(|&v| {
156 if is_nd(nodata, f64::from(v)) {
157 None
158 } else {
159 Some(v)
160 }
161 })
162 .collect();
163 (Arc::new(arr), DataType::UInt32)
164 }
165 RasterDataType::Int32 => {
166 let vals: &[i32] = buf.as_slice::<i32>().map_err(|e| OxiGeoError::Internal {
167 message: e.to_string(),
168 })?;
169 let arr: Int32Array = vals
170 .iter()
171 .map(|&v| {
172 if is_nd(nodata, f64::from(v)) {
173 None
174 } else {
175 Some(v)
176 }
177 })
178 .collect();
179 (Arc::new(arr), DataType::Int32)
180 }
181 RasterDataType::UInt64 => {
182 let vals: &[u64] = buf.as_slice::<u64>().map_err(|e| OxiGeoError::Internal {
183 message: e.to_string(),
184 })?;
185 let arr: UInt64Array = vals
186 .iter()
187 .map(|&v| {
188 if is_nd(nodata, v as f64) {
191 None
192 } else {
193 Some(v)
194 }
195 })
196 .collect();
197 (Arc::new(arr), DataType::UInt64)
198 }
199 RasterDataType::Int64 => {
200 let vals: &[i64] = buf.as_slice::<i64>().map_err(|e| OxiGeoError::Internal {
201 message: e.to_string(),
202 })?;
203 let arr: Int64Array = vals
204 .iter()
205 .map(|&v| {
206 if is_nd(nodata, v as f64) {
207 None
208 } else {
209 Some(v)
210 }
211 })
212 .collect();
213 (Arc::new(arr), DataType::Int64)
214 }
215 RasterDataType::Float32 => {
216 let vals: &[f32] = buf.as_slice::<f32>().map_err(|e| OxiGeoError::Internal {
217 message: e.to_string(),
218 })?;
219 let arr: Float32Array = vals
220 .iter()
221 .map(|&v| {
222 if is_nd(nodata, f64::from(v)) {
223 None
224 } else {
225 Some(v)
226 }
227 })
228 .collect();
229 (Arc::new(arr), DataType::Float32)
230 }
231 RasterDataType::Float64 => {
232 let vals: &[f64] = buf.as_slice::<f64>().map_err(|e| OxiGeoError::Internal {
233 message: e.to_string(),
234 })?;
235 let arr: Float64Array = vals
236 .iter()
237 .map(|&v| if is_nd(nodata, v) { None } else { Some(v) })
238 .collect();
239 (Arc::new(arr), DataType::Float64)
240 }
241 RasterDataType::CFloat32 | RasterDataType::CFloat64 => {
242 return Err(OxiGeoError::NotSupported {
243 operation: format!(
244 "Arrow conversion of complex type {}",
245 buf.data_type().name()
246 ),
247 });
248 }
249 };
250
251 debug_assert_eq!(array.len(), n, "array length must equal pixel count");
252
253 let mut metadata: HashMap<String, String> = HashMap::with_capacity(3);
254 metadata.insert("width".to_string(), buf.width().to_string());
255 metadata.insert("height".to_string(), buf.height().to_string());
256 metadata.insert("data_type".to_string(), buf.data_type().name().to_string());
257
258 let field = Field::new("pixel_values", arrow_dt, true);
259 let schema = Arc::new(Schema::new_with_metadata(vec![field], metadata));
260
261 RecordBatch::try_new(schema, vec![array]).map_err(|e| OxiGeoError::Internal {
262 message: format!("Arrow RecordBatch construction failed: {e}"),
263 })
264 }
265}
266
267impl TryFrom<RecordBatch> for RasterBuffer {
270 type Error = OxiGeoError;
271
272 fn try_from(batch: RecordBatch) -> Result<Self> {
288 if batch.num_columns() != 1 {
291 return Err(OxiGeoError::InvalidParameter {
292 parameter: "batch",
293 message: format!("Expected exactly 1 column, got {}", batch.num_columns()),
294 });
295 }
296
297 let schema = batch.schema();
298 let field = schema.field(0);
299 if field.name() != "pixel_values" {
300 return Err(OxiGeoError::InvalidParameter {
301 parameter: "batch",
302 message: format!(
303 "Expected column name 'pixel_values', got '{}'",
304 field.name()
305 ),
306 });
307 }
308
309 let meta = schema.metadata();
310
311 let width: u64 = meta
312 .get("width")
313 .ok_or(OxiGeoError::InvalidParameter {
314 parameter: "batch",
315 message: "Schema metadata missing 'width' key".to_string(),
316 })?
317 .parse::<u64>()
318 .map_err(|e| OxiGeoError::InvalidParameter {
319 parameter: "batch",
320 message: format!("Schema metadata 'width' is not a valid u64: {e}"),
321 })?;
322
323 let height: u64 = meta
324 .get("height")
325 .ok_or(OxiGeoError::InvalidParameter {
326 parameter: "batch",
327 message: "Schema metadata missing 'height' key".to_string(),
328 })?
329 .parse::<u64>()
330 .map_err(|e| OxiGeoError::InvalidParameter {
331 parameter: "batch",
332 message: format!("Schema metadata 'height' is not a valid u64: {e}"),
333 })?;
334
335 let dt_name = meta.get("data_type").ok_or(OxiGeoError::InvalidParameter {
336 parameter: "batch",
337 message: "Schema metadata missing 'data_type' key".to_string(),
338 })?;
339
340 let data_type = parse_data_type(dt_name)?;
341
342 let column = batch.column(0);
345 let bytes = arrow_column_to_bytes(column, data_type)?;
346
347 RasterBuffer::new(bytes, width, height, data_type, NoDataValue::None)
348 }
349}
350
351fn parse_data_type(name: &str) -> Result<RasterDataType> {
355 match name {
356 "UInt8" => Ok(RasterDataType::UInt8),
357 "Int8" => Ok(RasterDataType::Int8),
358 "UInt16" => Ok(RasterDataType::UInt16),
359 "Int16" => Ok(RasterDataType::Int16),
360 "UInt32" => Ok(RasterDataType::UInt32),
361 "Int32" => Ok(RasterDataType::Int32),
362 "UInt64" => Ok(RasterDataType::UInt64),
363 "Int64" => Ok(RasterDataType::Int64),
364 "Float32" => Ok(RasterDataType::Float32),
365 "Float64" => Ok(RasterDataType::Float64),
366 "CFloat32" | "CFloat64" => Err(OxiGeoError::NotSupported {
367 operation: format!("Arrow conversion of complex type {name}"),
368 }),
369 other => Err(OxiGeoError::InvalidParameter {
370 parameter: "data_type",
371 message: format!("Unknown data type '{other}' in schema metadata"),
372 }),
373 }
374}
375
376fn arrow_column_to_bytes(column: &dyn Array, data_type: RasterDataType) -> Result<Vec<u8>> {
380 macro_rules! downcast_to_bytes {
381 ($ArrowArray:ty, $native:ty, $column:expr) => {{
382 let arr = $column
383 .as_any()
384 .downcast_ref::<$ArrowArray>()
385 .ok_or_else(|| OxiGeoError::Internal {
386 message: format!(
387 "Expected {} array, got {:?}",
388 stringify!($ArrowArray),
389 $column.data_type()
390 ),
391 })?;
392 let mut bytes = Vec::with_capacity(arr.len() * core::mem::size_of::<$native>());
393 for i in 0..arr.len() {
394 let v: $native = if arr.is_null(i) {
395 <$native as Default>::default()
396 } else {
397 arr.value(i)
398 };
399 bytes.extend_from_slice(&v.to_ne_bytes());
400 }
401 Ok(bytes)
402 }};
403 }
404
405 match data_type {
406 RasterDataType::UInt8 => downcast_to_bytes!(UInt8Array, u8, column),
407 RasterDataType::Int8 => downcast_to_bytes!(Int8Array, i8, column),
408 RasterDataType::UInt16 => downcast_to_bytes!(UInt16Array, u16, column),
409 RasterDataType::Int16 => downcast_to_bytes!(Int16Array, i16, column),
410 RasterDataType::UInt32 => downcast_to_bytes!(UInt32Array, u32, column),
411 RasterDataType::Int32 => downcast_to_bytes!(Int32Array, i32, column),
412 RasterDataType::UInt64 => downcast_to_bytes!(UInt64Array, u64, column),
413 RasterDataType::Int64 => downcast_to_bytes!(Int64Array, i64, column),
414 RasterDataType::Float32 => downcast_to_bytes!(Float32Array, f32, column),
415 RasterDataType::Float64 => downcast_to_bytes!(Float64Array, f64, column),
416 RasterDataType::CFloat32 | RasterDataType::CFloat64 => Err(OxiGeoError::NotSupported {
417 operation: format!("Arrow conversion of complex type {}", data_type.name()),
418 }),
419 }
420}
421
422#[cfg(test)]
425mod tests {
426 #![allow(clippy::expect_used)]
427
428 use super::*;
429 use crate::buffer::RasterBuffer;
430 use crate::types::{NoDataValue, RasterDataType};
431
432 fn make_u8_4x4() -> RasterBuffer {
434 let data: Vec<u8> = (0u8..16).collect();
435 RasterBuffer::new(data, 4, 4, RasterDataType::UInt8, NoDataValue::None)
436 .expect("valid buffer")
437 }
438
439 fn make_f32_2x3() -> RasterBuffer {
441 let data: Vec<f32> = (0..6).map(|i| i as f32 * 1.5_f32).collect();
442 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_ne_bytes()).collect();
443 RasterBuffer::new(bytes, 2, 3, RasterDataType::Float32, NoDataValue::None)
444 .expect("valid buffer")
445 }
446
447 fn make_f64_2x2() -> RasterBuffer {
449 let data: Vec<f64> = vec![1.1, 2.2, 3.3, 4.4];
450 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_ne_bytes()).collect();
451 RasterBuffer::new(bytes, 2, 2, RasterDataType::Float64, NoDataValue::None)
452 .expect("valid buffer")
453 }
454
455 #[test]
459 fn test_raster_buffer_to_record_batch_u8() {
460 let buf = make_u8_4x4();
461 let batch = RecordBatch::try_from(&buf).expect("conversion should succeed");
462
463 assert_eq!(batch.num_rows(), 16);
465 assert_eq!(batch.num_columns(), 1);
466
467 assert_eq!(batch.schema().field(0).data_type(), &DataType::UInt8);
469 assert_eq!(batch.schema().field(0).name(), "pixel_values");
470
471 let col = batch
473 .column(0)
474 .as_any()
475 .downcast_ref::<UInt8Array>()
476 .expect("UInt8Array");
477 for i in 0u8..16 {
478 assert_eq!(col.value(i as usize), i, "value at index {i}");
479 assert!(!col.is_null(i as usize));
480 }
481 }
482
483 #[test]
485 fn test_raster_buffer_to_record_batch_f32_metadata() {
486 let buf = make_f32_2x3();
487 let batch = RecordBatch::try_from(&buf).expect("conversion should succeed");
488
489 let meta = batch.schema().metadata().clone();
490 assert_eq!(meta.get("width").map(String::as_str), Some("2"));
491 assert_eq!(meta.get("height").map(String::as_str), Some("3"));
492 assert_eq!(meta.get("data_type").map(String::as_str), Some("Float32"));
493
494 assert_eq!(batch.num_rows(), 6);
496 assert_eq!(batch.schema().field(0).data_type(), &DataType::Float32);
497 }
498
499 #[test]
501 fn test_record_batch_roundtrip_f64() {
502 let original = make_f64_2x2();
503 let batch = RecordBatch::try_from(&original).expect("forward conversion");
504 let recovered = RasterBuffer::try_from(batch).expect("reverse conversion");
505
506 assert_eq!(recovered.width(), original.width());
507 assert_eq!(recovered.height(), original.height());
508 assert_eq!(recovered.data_type(), RasterDataType::Float64);
509
510 for y in 0..original.height() {
511 for x in 0..original.width() {
512 let orig_val = original.get_pixel(x, y).expect("get_pixel original");
513 let rcvd_val = recovered.get_pixel(x, y).expect("get_pixel recovered");
514 assert!(
515 (orig_val - rcvd_val).abs() < f64::EPSILON,
516 "pixel ({x},{y}): expected {orig_val}, got {rcvd_val}"
517 );
518 }
519 }
520 }
521
522 #[test]
524 fn test_nodata_becomes_null() {
525 let data_f32: Vec<f32> = vec![0.0_f32, 1.5_f32, 0.0_f32];
527 let bytes: Vec<u8> = data_f32.iter().flat_map(|v| v.to_ne_bytes()).collect();
528 let buf = RasterBuffer::new(
529 bytes,
530 3,
531 1,
532 RasterDataType::Float32,
533 NoDataValue::Float(0.0),
534 )
535 .expect("valid buffer");
536
537 let batch = RecordBatch::try_from(&buf).expect("conversion should succeed");
538 let col = batch
539 .column(0)
540 .as_any()
541 .downcast_ref::<Float32Array>()
542 .expect("Float32Array");
543
544 assert_eq!(col.len(), 3);
545 assert!(col.is_null(0), "pixel 0 (nodata) should be null");
546 assert!(!col.is_null(1), "pixel 1 (1.5) should not be null");
547 assert!(col.is_null(2), "pixel 2 (nodata) should be null");
548 assert!((col.value(1) - 1.5_f32).abs() < f32::EPSILON);
549 }
550
551 #[test]
555 fn test_complex_type_returns_error() {
556 let buf = RasterBuffer::zeros(2, 2, RasterDataType::CFloat32);
558 let result = RecordBatch::try_from(&buf);
559
560 assert!(result.is_err(), "CFloat32 conversion should fail");
561 let err = result.expect_err("expected error");
562 assert!(
563 matches!(err, OxiGeoError::NotSupported { .. }),
564 "expected NotSupported, got {err:?}"
565 );
566 }
567
568 #[test]
570 fn test_complex_type_cfloat64_returns_error() {
571 let buf = RasterBuffer::zeros(2, 2, RasterDataType::CFloat64);
572 let result = RecordBatch::try_from(&buf);
573
574 assert!(result.is_err());
575 assert!(matches!(
576 result.expect_err("expected error"),
577 OxiGeoError::NotSupported { .. }
578 ));
579 }
580
581 #[test]
585 fn test_record_batch_to_buffer_wrong_schema_missing_metadata() {
586 let field = Field::new("pixel_values", DataType::UInt8, false);
588 let schema = Arc::new(Schema::new(vec![field]));
589 let array: ArrayRef = Arc::new(UInt8Array::from(vec![1u8, 2, 3, 4]));
590 let batch = RecordBatch::try_new(schema, vec![array]).expect("valid RecordBatch for test");
591
592 let result = RasterBuffer::try_from(batch);
593 assert!(result.is_err(), "missing metadata should fail");
594 assert!(
595 matches!(
596 result.expect_err("expected error"),
597 OxiGeoError::InvalidParameter { .. }
598 ),
599 "expected InvalidParameter error"
600 );
601 }
602
603 #[test]
605 fn test_record_batch_to_buffer_wrong_column_name() {
606 let mut metadata = HashMap::new();
607 metadata.insert("width".to_string(), "2".to_string());
608 metadata.insert("height".to_string(), "2".to_string());
609 metadata.insert("data_type".to_string(), "UInt8".to_string());
610
611 let field = Field::new("wrong_name", DataType::UInt8, false);
612 let schema = Arc::new(Schema::new_with_metadata(vec![field], metadata));
613 let array: ArrayRef = Arc::new(UInt8Array::from(vec![1u8, 2, 3, 4]));
614 let batch = RecordBatch::try_new(schema, vec![array]).expect("valid RecordBatch for test");
615
616 let result = RasterBuffer::try_from(batch);
617 assert!(result.is_err());
618 assert!(matches!(
619 result.expect_err("expected error"),
620 OxiGeoError::InvalidParameter { .. }
621 ));
622 }
623
624 #[test]
626 fn test_record_batch_to_buffer_too_many_columns() {
627 let mut metadata = HashMap::new();
628 metadata.insert("width".to_string(), "2".to_string());
629 metadata.insert("height".to_string(), "2".to_string());
630 metadata.insert("data_type".to_string(), "UInt8".to_string());
631
632 let field1 = Field::new("pixel_values", DataType::UInt8, false);
633 let field2 = Field::new("extra_column", DataType::UInt8, false);
634 let schema = Arc::new(Schema::new_with_metadata(vec![field1, field2], metadata));
635 let array1: ArrayRef = Arc::new(UInt8Array::from(vec![1u8, 2, 3, 4]));
636 let array2: ArrayRef = Arc::new(UInt8Array::from(vec![5u8, 6, 7, 8]));
637 let batch =
638 RecordBatch::try_new(schema, vec![array1, array2]).expect("valid RecordBatch for test");
639
640 let result = RasterBuffer::try_from(batch);
641 assert!(result.is_err());
642 assert!(matches!(
643 result.expect_err("expected error"),
644 OxiGeoError::InvalidParameter { .. }
645 ));
646 }
647
648 #[test]
652 fn test_roundtrip_i16() {
653 let data: Vec<i16> = vec![-1000_i16, 0, 1000, i16::MAX];
654 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_ne_bytes()).collect();
655 let buf = RasterBuffer::new(bytes, 2, 2, RasterDataType::Int16, NoDataValue::None)
656 .expect("valid buffer");
657
658 let batch = RecordBatch::try_from(&buf).expect("forward");
659 let recovered = RasterBuffer::try_from(batch).expect("reverse");
660
661 assert_eq!(recovered.data_type(), RasterDataType::Int16);
662 for y in 0..2u64 {
663 for x in 0..2u64 {
664 let orig = buf.get_pixel(x, y).expect("orig pixel");
665 let rcvd = recovered.get_pixel(x, y).expect("rcvd pixel");
666 assert!((orig - rcvd).abs() < f64::EPSILON);
667 }
668 }
669 }
670
671 #[test]
673 fn test_roundtrip_u64() {
674 let data: Vec<u64> = vec![0u64, 1, u64::MAX / 2, 1_000_000];
675 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_ne_bytes()).collect();
676 let buf = RasterBuffer::new(bytes, 2, 2, RasterDataType::UInt64, NoDataValue::None)
677 .expect("valid buffer");
678
679 let batch = RecordBatch::try_from(&buf).expect("forward");
680 let recovered = RasterBuffer::try_from(batch).expect("reverse");
681
682 assert_eq!(recovered.data_type(), RasterDataType::UInt64);
683 assert_eq!(recovered.as_bytes(), buf.as_bytes());
685 }
686
687 #[test]
689 fn test_integer_nodata_becomes_null() {
690 let data: Vec<i32> = vec![-9999_i32, 100, -9999, 200];
692 let bytes: Vec<u8> = data.iter().flat_map(|v| v.to_ne_bytes()).collect();
693 let buf = RasterBuffer::new(
694 bytes,
695 4,
696 1,
697 RasterDataType::Int32,
698 NoDataValue::Integer(-9999),
699 )
700 .expect("valid buffer");
701
702 let batch = RecordBatch::try_from(&buf).expect("conversion");
703 let col = batch
704 .column(0)
705 .as_any()
706 .downcast_ref::<Int32Array>()
707 .expect("Int32Array");
708
709 assert!(col.is_null(0), "index 0 (-9999) should be null");
710 assert!(!col.is_null(1), "index 1 (100) should not be null");
711 assert!(col.is_null(2), "index 2 (-9999) should be null");
712 assert!(!col.is_null(3), "index 3 (200) should not be null");
713 assert_eq!(col.value(1), 100);
714 assert_eq!(col.value(3), 200);
715 }
716}