use std::sync::Arc;
use arrow_array::cast::AsArray;
use arrow_array::{Array, ArrayRef, RecordBatch};
use arrow_schema::{ArrowError, DataType, Schema, TimeUnit};
use geopackage_core::datetime::{Date, DateTime};
use geopackage_core::types::{ColumnType, GeometryType};
use crate::{Error, Layer, Result};
use super::schema::{EXTENSION_METADATA_KEY, EXTENSION_NAME_KEY, GEOARROW_WKB, epsg_code};
struct RowLayout {
fid: Option<usize>,
geometry: Option<usize>,
values: Vec<Option<usize>>,
}
struct ArrowRow {
batch: Arc<RecordBatch>,
layout: Arc<RowLayout>,
row: usize,
}
enum ArrowRowResult {
Row(ArrowRow),
Failed(Error),
}
impl crate::writer::WritableRow for ArrowRowResult {
fn write(self, writer: &mut crate::FeatureWriter<'_>) -> Result<(i64, Option<[f64; 4]>)> {
match self {
Self::Row(row) => row.write(writer),
Self::Failed(error) => Err(error),
}
}
}
impl crate::writer::WritableRow for ArrowRow {
fn write(self, writer: &mut crate::FeatureWriter<'_>) -> Result<(i64, Option<[f64; 4]>)> {
let fid = match self
.layout
.fid
.and_then(|index| self.batch.columns().get(index))
{
Some(column) => read_i64(column, self.row)?,
None => None,
};
let mut values = Vec::with_capacity(self.layout.values.len());
for (position, index) in self.layout.values.iter().enumerate() {
let bound = match index.and_then(|index| self.batch.columns().get(index)) {
Some(column) => bind_value(column, self.row, position, &self.batch)?,
None => rusqlite::types::ToSqlOutput::Borrowed(rusqlite::types::ValueRef::Null),
};
values.push(bound);
}
let geometry = self
.layout
.geometry
.and_then(|index| self.batch.columns().get(index));
match geometry {
Some(column) if !column.is_null(self.row) => {
let wkb = binary_at(column, self.row)?;
writer.insert_wkb_bound(fid, wkb, &values)
}
_ => writer.insert_row_bound(fid, &values).map(|fid| (fid, None)),
}
}
}
impl Layer<'_> {
pub fn write_arrow<R>(&self, batches: R, batch_size: usize) -> Result<Vec<i64>>
where
R: IntoIterator<Item = std::result::Result<RecordBatch, ArrowError>>,
{
self.write_arrow_with(batches, batch_size, crate::BulkIndexOptions::default())
}
pub fn write_arrow_with<R>(
&self,
batches: R,
batch_size: usize,
options: crate::BulkIndexOptions,
) -> Result<Vec<i64>>
where
R: IntoIterator<Item = std::result::Result<RecordBatch, ArrowError>>,
{
let geometry_column = self.geometry_column().map(|g| g.column_name.clone());
let value_columns: Vec<String> = self
.value_columns()
.iter()
.map(|column| column.name.clone())
.collect();
let primary_key = self.primary_key_column().map(str::to_owned);
let rows = batches.into_iter().flat_map(move |batch| {
let taken = batch.map_err(Error::Arrow).and_then(|batch| {
let layout = layout_of(
&batch,
primary_key.as_deref(),
geometry_column.as_deref(),
&value_columns,
)?;
Ok((Arc::new(batch), Arc::new(layout)))
});
let (batch, error) = match taken {
Ok(batch) => (Some(batch), None),
Err(error) => (None, Some(error)),
};
batch
.into_iter()
.flat_map(|(batch, layout)| {
(0..batch.num_rows()).map(move |row| {
ArrowRowResult::Row(ArrowRow {
batch: Arc::clone(&batch),
layout: Arc::clone(&layout),
row,
})
})
})
.chain(error.into_iter().map(ArrowRowResult::Failed))
});
self.write_all_impl(rows, batch_size, options, crate::bulk::no_fault)
}
}
fn layout_of(
batch: &RecordBatch,
primary_key: Option<&str>,
geometry: Option<&str>,
value_columns: &[String],
) -> Result<RowLayout> {
let schema = batch.schema();
for field in schema.fields() {
let known = Some(field.name().as_str()) == primary_key
|| Some(field.name().as_str()) == geometry
|| value_columns.iter().any(|name| name == field.name());
if !known {
return Err(Error::NoSuchColumn {
table_name: String::new(),
column_name: field.name().clone(),
});
}
}
let index_of = |name: &str| schema.fields().iter().position(|f| f.name() == name);
Ok(RowLayout {
fid: primary_key.and_then(index_of),
geometry: geometry.and_then(index_of),
values: value_columns.iter().map(|name| index_of(name)).collect(),
})
}
fn binary_at(column: &ArrayRef, row: usize) -> Result<&[u8]> {
if let Some(binary) = column.as_binary_opt::<i32>() {
return Ok(binary.value(row));
}
if let Some(binary) = column.as_binary_opt::<i64>() {
return Ok(binary.value(row));
}
Err(Error::ArrowValueMismatch {
column: String::new(),
expected: "Binary or LargeBinary",
found: "another Arrow type",
})
}
fn read_i64(column: &ArrayRef, row: usize) -> Result<Option<i64>> {
if column.is_null(row) {
return Ok(None);
}
let values = column
.as_primitive_opt::<arrow_array::types::Int64Type>()
.ok_or_else(|| Error::ArrowValueMismatch {
column: String::new(),
expected: "Int64",
found: "other",
})?;
Ok(Some(values.value(row)))
}
fn bind_value<'a>(
column: &'a ArrayRef,
row: usize,
position: usize,
batch: &RecordBatch,
) -> Result<rusqlite::types::ToSqlOutput<'a>> {
use arrow_array::types::{
Date32Type, Float32Type, Float64Type, Int8Type, Int16Type, Int32Type, Int64Type,
TimestampMicrosecondType, TimestampMillisecondType,
};
use rusqlite::types::{ToSqlOutput, Value as SqlV, ValueRef};
let borrowed = |value: ValueRef<'a>| Ok(ToSqlOutput::Borrowed(value));
let owned = |value: SqlV| Ok(ToSqlOutput::Owned(value));
if column.is_null(row) {
return borrowed(ValueRef::Null);
}
let name = || {
batch
.schema()
.fields()
.get(position)
.map(|field| field.name().clone())
.unwrap_or_default()
};
let mismatch = |expected: &'static str| Error::ArrowValueMismatch {
column: name(),
expected,
found: "an array of another type",
};
let out_of_range = |source| Error::InvalidDateTimeValue {
column: name(),
text: "an Arrow date or timestamp outside the representable range".to_owned(),
source,
};
match column.data_type() {
DataType::Boolean => owned(SqlV::Integer(i64::from(
column
.as_boolean_opt()
.ok_or_else(|| mismatch("Boolean"))?
.value(row),
))),
DataType::Int8 => owned(SqlV::Integer(i64::from(
column
.as_primitive_opt::<Int8Type>()
.ok_or_else(|| mismatch("Int8"))?
.value(row),
))),
DataType::Int16 => owned(SqlV::Integer(i64::from(
column
.as_primitive_opt::<Int16Type>()
.ok_or_else(|| mismatch("Int16"))?
.value(row),
))),
DataType::Int32 => owned(SqlV::Integer(i64::from(
column
.as_primitive_opt::<Int32Type>()
.ok_or_else(|| mismatch("Int32"))?
.value(row),
))),
DataType::Int64 => borrowed(ValueRef::Integer(
column
.as_primitive_opt::<Int64Type>()
.ok_or_else(|| mismatch("Int64"))?
.value(row),
)),
DataType::Float32 => owned(SqlV::Real(f64::from(
column
.as_primitive_opt::<Float32Type>()
.ok_or_else(|| mismatch("Float32"))?
.value(row),
))),
DataType::Float64 => borrowed(ValueRef::Real(
column
.as_primitive_opt::<Float64Type>()
.ok_or_else(|| mismatch("Float64"))?
.value(row),
)),
DataType::Utf8 => borrowed(ValueRef::Text(
column
.as_string_opt::<i32>()
.ok_or_else(|| mismatch("Utf8"))?
.value(row)
.as_bytes(),
)),
DataType::LargeUtf8 => borrowed(ValueRef::Text(
column
.as_string_opt::<i64>()
.ok_or_else(|| mismatch("LargeUtf8"))?
.value(row)
.as_bytes(),
)),
DataType::Binary => borrowed(ValueRef::Blob(
column
.as_binary_opt::<i32>()
.ok_or_else(|| mismatch("Binary"))?
.value(row),
)),
DataType::LargeBinary => borrowed(ValueRef::Blob(
column
.as_binary_opt::<i64>()
.ok_or_else(|| mismatch("LargeBinary"))?
.value(row),
)),
DataType::Date32 => owned(SqlV::Text(
Date::from_days_since_epoch(
column
.as_primitive_opt::<Date32Type>()
.ok_or_else(|| mismatch("Date32"))?
.value(row),
)
.map_err(out_of_range)?
.to_string(),
)),
DataType::Timestamp(TimeUnit::Microsecond, _) => owned(SqlV::Text(
DateTime::from_micros_since_epoch(
column
.as_primitive_opt::<TimestampMicrosecondType>()
.ok_or_else(|| mismatch("Timestamp"))?
.value(row),
)
.map_err(out_of_range)?
.to_string(),
)),
DataType::Timestamp(TimeUnit::Millisecond, _) => owned(SqlV::Text(
DateTime::from_micros_since_epoch(
column
.as_primitive_opt::<TimestampMillisecondType>()
.ok_or_else(|| mismatch("Timestamp"))?
.value(row)
.saturating_mul(1_000),
)
.map_err(out_of_range)?
.to_string(),
)),
other => Err(Error::UnsupportedArrowType {
data_type: other.to_string(),
}),
}
}
impl crate::TableSchemaBuilder {
pub fn from_arrow_schema(self, schema: &Schema) -> Result<Self> {
let mut builder = self;
for field in schema.fields() {
if *field.name() == builder.primary_key_name() {
continue;
}
if field.metadata().get(EXTENSION_NAME_KEY).map(String::as_str) == Some(GEOARROW_WKB) {
let srs_id = field
.metadata()
.get(EXTENSION_METADATA_KEY)
.and_then(|json| epsg_code(json))
.unwrap_or(0);
builder = builder.geometry(
crate::GeometrySpec::new(GeometryType::Geometry, srs_id)
.column_name(field.name()),
);
continue;
}
let column_type = column_type_for(field.data_type())?;
let mut column = crate::ColumnSpec::new(field.name(), column_type);
if !field.is_nullable() {
column = column.not_null();
}
builder = builder.column(column);
}
Ok(builder)
}
}
fn column_type_for(data_type: &DataType) -> Result<ColumnType> {
Ok(match data_type {
DataType::Boolean => ColumnType::Boolean,
DataType::Int8 | DataType::UInt8 => ColumnType::TinyInt,
DataType::Int16 | DataType::UInt16 => ColumnType::SmallInt,
DataType::Int32 | DataType::UInt32 => ColumnType::MediumInt,
DataType::Int64 | DataType::UInt64 => ColumnType::Integer,
DataType::Float32 => ColumnType::Float,
DataType::Float64 => ColumnType::Double,
DataType::Utf8 | DataType::LargeUtf8 => ColumnType::Text(None),
DataType::Binary | DataType::LargeBinary => ColumnType::Blob(None),
DataType::Date32 => ColumnType::Date,
DataType::Timestamp(_, _) => ColumnType::DateTime,
other => {
return Err(Error::UnsupportedArrowType {
data_type: other.to_string(),
});
}
})
}