Skip to content

Commit 574a3f4

Browse files
authored
fix: dynamically load correct array types given inferred schema (#26)
* dynamically load arrays as needed based on their data types. * fix typo
1 parent c04b271 commit 574a3f4

3 files changed

Lines changed: 125 additions & 50 deletions

File tree

Cargo.lock

Lines changed: 2 additions & 2 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

Cargo.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,6 @@ arrow-schema = "56.2.0"
1010
async-trait = "0.1.89"
1111
datafusion = "50.2"
1212
futures = "0.3.31"
13-
geoarrow-array = "0.6.1"
1413
geoarrow-schema = "0.6.1"
1514
icechunk = "0.3.13"
1615
object_store = "0.12.4"
@@ -24,4 +23,5 @@ zarrs_object_store = "0.5.0"
2423
zarrs_storage = { version = "0.4.0", features = ["async"] }
2524

2625
[dev-dependencies]
26+
geoarrow-array = "0.6.1"
2727
tokio = { version = "1.48", features = ["macros", "rt-multi-thread"] }

src/table_provider.rs

Lines changed: 122 additions & 47 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,11 @@
1-
use arrow_array::{ArrayRef, RecordBatch, StringViewArray, TimestampMillisecondArray};
2-
use arrow_schema::SchemaRef;
1+
use arrow_array::{
2+
ArrayRef, BinaryArray, BinaryViewArray, BooleanArray, Float32Array, Float64Array, Int8Array,
3+
Int16Array, Int32Array, Int64Array, LargeBinaryArray, LargeStringArray, RecordBatch,
4+
StringArray, StringViewArray, TimestampMicrosecondArray, TimestampMillisecondArray,
5+
TimestampNanosecondArray, TimestampSecondArray, UInt8Array, UInt16Array, UInt32Array,
6+
UInt64Array,
7+
};
8+
use arrow_schema::{DataType, Field, SchemaRef, TimeUnit};
39
use async_trait::async_trait;
410
use datafusion::catalog::Session;
511
use datafusion::datasource::{TableProvider, TableType};
@@ -13,9 +19,6 @@ use datafusion::physical_plan::{
1319
DisplayAs, DisplayFormatType, ExecutionPlan, Partitioning, PlanProperties,
1420
SendableRecordBatchStream,
1521
};
16-
use geoarrow_array::GeoArrowArray;
17-
use geoarrow_array::array::WktViewArray;
18-
use geoarrow_schema::Crs;
1922
use object_store::ObjectStore;
2023
use std::any::Any;
2124
use std::fmt::{self, Debug};
@@ -61,7 +64,6 @@ impl ZarrTableProvider {
6164
) -> ZarrDataFusionResult<Self> {
6265
let zarr_backend = IcechunkBackend::new(icechunk_session, handle);
6366
let schema = zarr_backend.infer_group_schema(group_path.into()).await?;
64-
// dbg!(schema.as_ref());
6567
Ok(Self {
6668
schema,
6769
zarr_backend: zarr_backend.into(),
@@ -133,8 +135,6 @@ impl SyncZarrBackend {
133135
}
134136
}
135137

136-
// TODO: Have an icechunk backend that stores both the icechunk session **and** the tokio runtime. Then we can ensure that loading data always happens within the correct runtime context.
137-
138138
#[derive(Clone)]
139139
struct IcechunkBackend {
140140
store: Arc<dyn AsyncReadableListableStorageTraits>,
@@ -238,18 +238,6 @@ impl From<SyncZarrBackend> for ZarrBackend {
238238
}
239239

240240
impl ZarrBackend {
241-
// fn new_filesystem<P: AsRef<std::path::Path>>(
242-
// base_path: P,
243-
// ) -> Result<Self, FilesystemStoreCreateError> {
244-
// Ok(Self::Sync(SyncZarrBackend::new_filesystem(base_path)?))
245-
// }
246-
247-
// fn new_object_store<T: ObjectStore>(store: T) -> Self {
248-
// Self::Async(AsyncZarrBackend(Arc::new(
249-
// zarrs_object_store::AsyncObjectStore::new(store),
250-
// )))
251-
// }
252-
253241
async fn load_array<T: ElementOwned + MaybeSend + MaybeSync + 'static>(
254242
&self,
255243
path: &str,
@@ -263,36 +251,123 @@ impl ZarrBackend {
263251
}
264252
}
265253

266-
async fn load_record_batch(self, schema: SchemaRef) -> ZarrDataFusionResult<RecordBatch> {
267-
let collection_data: Vec<String> = self.load_array("/meta/collection").await?;
268-
let date_data: Vec<i64> = self.load_array("/meta/date").await?;
269-
let bbox_data: Vec<String> = self.load_array("/meta/bbox").await?;
270-
271-
// Create Arrow arrays from the loaded data
272-
let collection_arrow: ArrayRef = Arc::new(StringViewArray::from(collection_data));
273-
let date_arrow: ArrayRef = Arc::new(TimestampMillisecondArray::from(date_data));
274-
let wkt_crs = Crs::from_authority_code("EPSG:4326".to_string());
275-
let wkt_metadata = Arc::new(geoarrow_schema::Metadata::new(wkt_crs, None));
276-
let wkt_arrow = WktViewArray::new(bbox_data.into(), wkt_metadata);
277-
278-
let columns = schema
279-
.fields()
280-
.iter()
281-
.map(|field| match field.name().as_str() {
282-
"collection" => collection_arrow.clone(),
283-
"date" => date_arrow.clone(),
284-
"bbox" => wkt_arrow.clone().into_array_ref(),
285-
_ => panic!("Unexpected field name: {}", field.name()),
286-
})
287-
.collect();
254+
async fn load_array_given_field(&self, field: &Field) -> ZarrDataFusionResult<ArrayRef> {
255+
// Note: we don't need to check for extension type information here, because we're only
256+
// loading the physical data, and the metadata is already held in the schema.
288257

289-
// Create the RecordBatch
290-
let record_batch = RecordBatch::try_new(schema.clone(), columns)?;
258+
// TODO: refactor so this can be stored in the ZarrBackend
259+
let group = "/meta";
260+
let name = field.name();
261+
let path = format!("{group}/{name}");
291262

292-
// dbg!(&record_batch);
293-
// dbg!("equal?", schema.as_ref() == record_batch.schema().as_ref());
263+
match field.data_type() {
264+
DataType::Boolean => {
265+
let data: Vec<bool> = self.load_array(&path).await?;
266+
Ok(Arc::new(BooleanArray::from(data)))
267+
}
268+
DataType::Int8 => {
269+
let data: Vec<i8> = self.load_array(&path).await?;
270+
Ok(Arc::new(Int8Array::from(data)))
271+
}
272+
DataType::Int16 => {
273+
let data: Vec<i16> = self.load_array(&path).await?;
274+
Ok(Arc::new(Int16Array::from(data)))
275+
}
276+
DataType::Int32 => {
277+
let data: Vec<i32> = self.load_array(&path).await?;
278+
Ok(Arc::new(Int32Array::from(data)))
279+
}
280+
DataType::Int64 => {
281+
let data: Vec<i64> = self.load_array(&path).await?;
282+
Ok(Arc::new(Int64Array::from(data)))
283+
}
284+
DataType::UInt8 => {
285+
let data: Vec<u8> = self.load_array(&path).await?;
286+
Ok(Arc::new(UInt8Array::from(data)))
287+
}
288+
DataType::UInt16 => {
289+
let data: Vec<u16> = self.load_array(&path).await?;
290+
Ok(Arc::new(UInt16Array::from(data)))
291+
}
292+
DataType::UInt32 => {
293+
let data: Vec<u32> = self.load_array(&path).await?;
294+
Ok(Arc::new(UInt32Array::from(data)))
295+
}
296+
DataType::UInt64 => {
297+
let data: Vec<u64> = self.load_array(&path).await?;
298+
Ok(Arc::new(UInt64Array::from(data)))
299+
}
300+
// DataType::Float16 => {
301+
// let data: Vec<f16> = self.load_array(&path).await?;
302+
// Ok(Arc::new(Float16Array::from(data)))
303+
// }
304+
DataType::Float32 => {
305+
let data: Vec<f32> = self.load_array(&path).await?;
306+
Ok(Arc::new(Float32Array::from(data)))
307+
}
308+
DataType::Float64 => {
309+
let data: Vec<f64> = self.load_array(&path).await?;
310+
Ok(Arc::new(Float64Array::from(data)))
311+
}
312+
DataType::Binary => {
313+
let data: Vec<Vec<u8>> = self.load_array(&path).await?;
314+
let refs: Vec<&[u8]> = data.iter().map(|v| v.as_slice()).collect();
315+
Ok(Arc::new(BinaryArray::from(refs)))
316+
}
317+
DataType::LargeBinary => {
318+
let data: Vec<Vec<u8>> = self.load_array(&path).await?;
319+
let refs: Vec<&[u8]> = data.iter().map(|v| v.as_slice()).collect();
320+
Ok(Arc::new(LargeBinaryArray::from(refs)))
321+
}
322+
DataType::BinaryView => {
323+
let data: Vec<Vec<u8>> = self.load_array(&path).await?;
324+
let refs: Vec<&[u8]> = data.iter().map(|v| v.as_slice()).collect();
325+
Ok(Arc::new(BinaryViewArray::from(refs)))
326+
}
327+
DataType::Utf8 => {
328+
let data: Vec<String> = self.load_array(&path).await?;
329+
Ok(Arc::new(StringArray::from(data)))
330+
}
331+
DataType::LargeUtf8 => {
332+
let data: Vec<String> = self.load_array(&path).await?;
333+
Ok(Arc::new(LargeStringArray::from(data)))
334+
}
335+
DataType::Utf8View => {
336+
let data: Vec<String> = self.load_array(&path).await?;
337+
Ok(Arc::new(StringViewArray::from(data)))
338+
}
339+
DataType::Timestamp(unit, _) => match unit {
340+
TimeUnit::Millisecond => {
341+
let data: Vec<i64> = self.load_array(&path).await?;
342+
Ok(Arc::new(TimestampMillisecondArray::from(data)))
343+
}
344+
TimeUnit::Microsecond => {
345+
let data: Vec<i64> = self.load_array(&path).await?;
346+
Ok(Arc::new(TimestampMicrosecondArray::from(data)))
347+
}
348+
TimeUnit::Nanosecond => {
349+
let data: Vec<i64> = self.load_array(&path).await?;
350+
Ok(Arc::new(TimestampNanosecondArray::from(data)))
351+
}
352+
TimeUnit::Second => {
353+
let data: Vec<i64> = self.load_array(&path).await?;
354+
Ok(Arc::new(TimestampSecondArray::from(data)))
355+
}
356+
},
357+
_ => Err(ZarrDataFusionError::Custom(format!(
358+
"Unsupported Arrow data type: {:?}",
359+
field.data_type()
360+
))),
361+
}
362+
}
363+
364+
async fn load_record_batch(self, schema: SchemaRef) -> ZarrDataFusionResult<RecordBatch> {
365+
let mut arrays = vec![];
366+
for field in schema.fields() {
367+
arrays.push(self.load_array_given_field(field).await?);
368+
}
294369

295-
Ok(record_batch)
370+
Ok(RecordBatch::try_new(schema.clone(), arrays)?)
296371
}
297372
}
298373

0 commit comments

Comments
 (0)