use arrow::datatypes::DataType;
use datafusion::common::ScalarValue;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum PartitionTransform {
Identity,
Year,
Month,
Day,
Hour,
Bucket(u32),
Unknown(String),
}
impl PartitionTransform {
pub fn parse(transform: &str) -> Self {
let trimmed = transform.trim();
match trimmed.to_ascii_lowercase().as_str() {
"identity" => return PartitionTransform::Identity,
"year" => return PartitionTransform::Year,
"month" => return PartitionTransform::Month,
"day" => return PartitionTransform::Day,
"hour" => return PartitionTransform::Hour,
_ => {},
}
if let Some(rest) = trimmed
.strip_prefix("bucket(")
.or_else(|| trimmed.strip_prefix("BUCKET("))
&& let Some(inner) = rest.strip_suffix(')')
&& let Ok(n) = inner.trim().parse::<u32>()
{
return PartitionTransform::Bucket(n);
}
PartitionTransform::Unknown(trimmed.to_string())
}
pub fn to_catalog_string(&self) -> String {
match self {
PartitionTransform::Identity => "identity".to_string(),
PartitionTransform::Year => "year".to_string(),
PartitionTransform::Month => "month".to_string(),
PartitionTransform::Day => "day".to_string(),
PartitionTransform::Hour => "hour".to_string(),
PartitionTransform::Bucket(n) => format!("bucket({n})"),
PartitionTransform::Unknown(s) => s.clone(),
}
}
pub fn is_producible(&self) -> bool {
matches!(
self,
PartitionTransform::Identity
| PartitionTransform::Year
| PartitionTransform::Month
| PartitionTransform::Day
| PartitionTransform::Hour
)
}
pub(crate) fn value_is_well_formed(&self, value: Option<&str>, column_type: &DataType) -> bool {
let Some(value) = value else {
return true;
};
match self {
PartitionTransform::Identity => {
ScalarValue::try_from_string(value.to_string(), column_type).is_ok()
},
PartitionTransform::Year
| PartitionTransform::Month
| PartitionTransform::Day
| PartitionTransform::Hour => value.trim().parse::<i64>().is_ok(),
PartitionTransform::Bucket(buckets) => {
matches!(value.trim().parse::<i64>(), Ok(b) if b >= 0 && b < i64::from(*buckets))
},
PartitionTransform::Unknown(_) => true,
}
}
pub fn source_bounds(
&self,
value: &str,
data_type: &DataType,
) -> Option<(ScalarValue, ScalarValue)> {
match self {
PartitionTransform::Identity => {
let scalar = ScalarValue::try_from_string(value.to_string(), data_type).ok()?;
Some((scalar.clone(), scalar))
},
PartitionTransform::Year => {
let year: i64 = value.trim().parse().ok()?;
year_bounds(year, data_type)
},
PartitionTransform::Month
| PartitionTransform::Day
| PartitionTransform::Hour
| PartitionTransform::Bucket(_)
| PartitionTransform::Unknown(_) => None,
}
}
}
fn year_bounds(year: i64, data_type: &DataType) -> Option<(ScalarValue, ScalarValue)> {
let (min_str, max_str) = match data_type {
DataType::Date32 | DataType::Date64 => {
(format!("{year}-01-01"), format!("{}-01-01", year + 1))
},
DataType::Timestamp(_, None) => (
format!("{year}-01-01 00:00:00"),
format!("{}-01-01 00:00:00", year + 1),
),
DataType::Int8
| DataType::Int16
| DataType::Int32
| DataType::Int64
| DataType::UInt8
| DataType::UInt16
| DataType::UInt32
| DataType::UInt64 => {
let scalar = ScalarValue::try_from_string(year.to_string(), data_type).ok()?;
return Some((scalar.clone(), scalar));
},
_ => return None,
};
let min = ScalarValue::try_from_string(min_str, data_type).ok()?;
let max = ScalarValue::try_from_string(max_str, data_type).ok()?;
Some((min, max))
}
#[derive(Debug, Clone)]
pub struct PartitionWriteSpec {
pub partition_id: i64,
pub keys: Vec<PartitionWriteKey>,
}
#[derive(Debug, Clone)]
pub struct PartitionWriteKey {
pub input_index: usize,
pub name: String,
pub transform: PartitionTransform,
}
impl PartitionWriteSpec {
pub fn resolve(
spec: &PartitionSpec,
column_ids: &[i64],
schema: &arrow::datatypes::Schema,
) -> crate::Result<PartitionWriteSpec> {
let mut keys = Vec::with_capacity(spec.columns.len());
for column in &spec.columns {
if !column.transform.is_producible() {
return Err(crate::DuckLakeError::Unsupported(format!(
"writing to a table partitioned by '{}' is not supported",
column.transform.to_catalog_string()
)));
}
let index = column_ids
.iter()
.position(|id| *id == column.column_id)
.ok_or_else(|| {
crate::DuckLakeError::Internal(format!(
"partition column_id {} not found in table schema",
column.column_id
))
})?;
let field = schema.fields().get(index).ok_or_else(|| {
crate::DuckLakeError::Internal(format!(
"partition column index {index} out of range for write schema"
))
})?;
let name = field.name().to_string();
keys.push(PartitionWriteKey {
input_index: index,
name,
transform: column.transform.clone(),
});
}
Ok(PartitionWriteSpec {
partition_id: spec.partition_id,
keys,
})
}
pub(crate) fn key_names(&self) -> Vec<String> {
self.keys.iter().map(|k| k.name.clone()).collect()
}
pub(crate) fn validate_values(
&self,
schema: &arrow::datatypes::Schema,
values: &[Option<String>],
) -> crate::Result<()> {
if values.len() != self.keys.len() {
return Err(crate::DuckLakeError::InvalidConfig(format!(
"partitioned write supplied {} value(s) for a spec with {} key(s)",
values.len(),
self.keys.len()
)));
}
for (key, value) in self.keys.iter().zip(values.iter()) {
let column_type = schema
.fields()
.get(key.input_index)
.map(|f| f.data_type())
.ok_or_else(|| {
crate::DuckLakeError::Internal(format!(
"partition key column index {} out of range for write schema",
key.input_index
))
})?;
if !key
.transform
.value_is_well_formed(value.as_deref(), column_type)
{
return Err(crate::DuckLakeError::InvalidConfig(format!(
"partition value {value:?} is not valid for key '{}' with transform '{}' on a \
{column_type} column",
key.name,
key.transform.to_catalog_string()
)));
}
}
Ok(())
}
}
pub(crate) fn hive_subpath(key_names: &[String], values: &[Option<String>]) -> String {
let mut rel = String::new();
for (i, value) in values.iter().enumerate() {
let name = key_names.get(i).map(String::as_str).unwrap_or("key");
let encoded = match value {
Some(v) => sanitize_partition_path(v),
None => "__HIVE_DEFAULT_PARTITION__".to_string(),
};
if rel.is_empty() {
rel = format!("{name}={encoded}");
} else {
rel = format!("{rel}/{name}={encoded}");
}
}
rel
}
fn sanitize_partition_path(value: &str) -> String {
value
.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || matches!(c, '-' | '_' | '.' | ':') {
c
} else {
'_'
}
})
.collect()
}
#[cfg(feature = "write")]
pub type PartitionGroup = (Vec<Option<String>>, Vec<arrow::record_batch::RecordBatch>);
#[cfg(feature = "write")]
fn transform_array(
transform: &PartitionTransform,
array: &arrow::array::ArrayRef,
) -> crate::Result<arrow::array::ArrayRef> {
use arrow::compute::{DatePart, date_part};
let part = match transform {
PartitionTransform::Identity => return Ok(std::sync::Arc::clone(array)),
PartitionTransform::Year => DatePart::Year,
PartitionTransform::Month => DatePart::Month,
PartitionTransform::Day => DatePart::Day,
PartitionTransform::Hour => DatePart::Hour,
other => {
return Err(crate::DuckLakeError::Unsupported(format!(
"partitioned write with transform '{}' is not supported",
other.to_catalog_string()
)));
},
};
Ok(date_part(array, part)?)
}
#[cfg(feature = "write")]
pub(crate) fn split_batches_by_partition(
output_schema: &arrow::datatypes::SchemaRef,
batches: &[arrow::record_batch::RecordBatch],
spec: &PartitionWriteSpec,
) -> crate::Result<Vec<PartitionGroup>> {
use arrow::array::{ArrayRef, RecordBatch, UInt32Array};
use arrow::compute::take;
use std::collections::HashMap;
let mut order: Vec<Vec<Option<String>>> = Vec::new();
let mut groups: HashMap<Vec<Option<String>>, Vec<RecordBatch>> = HashMap::new();
for batch in batches {
let num_rows = batch.num_rows();
if num_rows == 0 {
continue;
}
let mut transformed: Vec<ArrayRef> = Vec::with_capacity(spec.keys.len());
for key in &spec.keys {
transformed.push(transform_array(
&key.transform,
batch.column(key.input_index),
)?);
}
let mut per_batch: HashMap<Vec<Option<String>>, Vec<u32>> = HashMap::new();
let mut per_batch_order: Vec<Vec<Option<String>>> = Vec::new();
for row in 0..num_rows {
let mut values: Vec<Option<String>> = Vec::with_capacity(spec.keys.len());
for array in &transformed {
let scalar = ScalarValue::try_from_array(array, row)?;
let encoded = if scalar.is_null() {
None
} else {
match crate::stats_encode::encode_scalar(&scalar) {
Some(encoded) => Some(encoded),
None => {
return Err(crate::DuckLakeError::Unsupported(format!(
"partitioned write: partition-key value of type {} cannot be \
encoded; partitioning by this column type is not supported",
array.data_type()
)));
},
}
};
values.push(encoded);
}
if !per_batch.contains_key(&values) {
per_batch_order.push(values.clone());
}
per_batch.entry(values).or_default().push(row as u32);
}
for values in per_batch_order {
let indices = per_batch.remove(&values).unwrap_or_default();
if indices.is_empty() {
continue;
}
let index_array = UInt32Array::from(indices);
let columns = batch
.columns()
.iter()
.map(|c| take(c, &index_array, None))
.collect::<std::result::Result<Vec<_>, _>>()?;
let out = RecordBatch::try_new(output_schema.clone(), columns)?;
if !groups.contains_key(&values) {
order.push(values.clone());
}
groups.entry(values).or_default().push(out);
}
}
Ok(order
.into_iter()
.filter_map(|values| groups.remove(&values).map(|batches| (values, batches)))
.collect())
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PartitionSpecColumn {
pub partition_key_index: i32,
pub column_id: i64,
pub transform: PartitionTransform,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct PartitionSpec {
pub partition_id: i64,
pub columns: Vec<PartitionSpecColumn>,
pub prune_safe: bool,
}
impl PartitionSpec {
pub fn transform_for_column(&self, column_id: i64) -> Option<&PartitionTransform> {
self.columns
.iter()
.find(|c| c.column_id == column_id)
.map(|c| &c.transform)
}
pub fn from_rows(
rows: Vec<(i64, i32, i64, String)>,
prune_safe: bool,
) -> Option<PartitionSpec> {
let partition_id = rows.first()?.0;
let columns = rows
.into_iter()
.map(
|(_, partition_key_index, column_id, transform)| PartitionSpecColumn {
partition_key_index,
column_id,
transform: PartitionTransform::parse(&transform),
},
)
.collect();
Some(PartitionSpec {
partition_id,
columns,
prune_safe,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[cfg(feature = "write")]
use arrow::array::{ArrayRef, RecordBatch, StringArray};
#[cfg(feature = "write")]
use arrow::datatypes::SchemaRef;
use arrow::datatypes::{Field, Schema};
#[cfg(feature = "write")]
use std::sync::Arc;
#[cfg(feature = "write")]
fn identity_region_spec() -> PartitionWriteSpec {
PartitionWriteSpec {
partition_id: 1,
keys: vec![PartitionWriteKey {
input_index: 0,
name: "region".to_string(),
transform: PartitionTransform::Identity,
}],
}
}
#[cfg(feature = "write")]
#[test]
fn split_groups_by_identity_and_keeps_null_partition() {
let schema: SchemaRef = Arc::new(Schema::new(vec![Field::new(
"region",
DataType::Utf8,
true,
)]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(StringArray::from(vec![Some("us"), None, Some("us")])) as ArrayRef],
)
.unwrap();
let groups = split_batches_by_partition(
&schema,
std::slice::from_ref(&batch),
&identity_region_spec(),
)
.unwrap();
assert_eq!(groups.len(), 2);
let total: usize = groups
.iter()
.flat_map(|(_, b)| b)
.map(|b| b.num_rows())
.sum();
assert_eq!(total, 3);
let mut values: Vec<Option<String>> = groups.iter().map(|(v, _)| v[0].clone()).collect();
values.sort();
assert_eq!(values, vec![None, Some("us".to_string())]);
}
#[cfg(feature = "write")]
#[test]
fn split_errors_on_unencodable_non_null_value_instead_of_corrupting() {
let schema: SchemaRef = Arc::new(Schema::new(vec![Field::new(
"region",
DataType::Utf8,
true,
)]));
let batch = RecordBatch::try_new(
schema.clone(),
vec![Arc::new(StringArray::from(vec![Some("a\u{0}b")])) as ArrayRef],
)
.unwrap();
let err = split_batches_by_partition(
&schema,
std::slice::from_ref(&batch),
&identity_region_spec(),
)
.unwrap_err();
assert!(
err.to_string().to_lowercase().contains("encode"),
"expected an encode error, got: {err}"
);
}
#[test]
fn hive_subpath_encodes_keys_values_and_nulls() {
let keys = vec!["region".to_string(), "day".to_string()];
assert_eq!(
hive_subpath(&keys, &[Some("us".into()), Some("3".into())]),
"region=us/day=3"
);
assert_eq!(
hive_subpath(&keys, &[None, Some("3".into())]),
"region=__HIVE_DEFAULT_PARTITION__/day=3"
);
assert_eq!(
hive_subpath(&keys[..1], &[Some("a/../b".into())]),
"region=a_.._b"
);
assert_eq!(hive_subpath(&[], &[]), "");
}
#[test]
fn resolve_maps_column_ids_to_write_schema_indices() {
let schema = Schema::new(vec![
Field::new("id", DataType::Int64, false),
Field::new("region", DataType::Utf8, true),
]);
let spec = PartitionSpec {
partition_id: 7,
columns: vec![PartitionSpecColumn {
partition_key_index: 0,
column_id: 20,
transform: PartitionTransform::Identity,
}],
prune_safe: true,
};
let resolved = PartitionWriteSpec::resolve(&spec, &[10, 20], &schema).unwrap();
assert_eq!(resolved.partition_id, 7);
assert_eq!(resolved.keys.len(), 1);
assert_eq!(resolved.keys[0].input_index, 1);
assert_eq!(resolved.keys[0].name, "region");
assert_eq!(resolved.key_names(), vec!["region".to_string()]);
}
#[test]
fn resolve_rejects_non_producible_transform() {
let schema = Schema::new(vec![Field::new("id", DataType::Int64, false)]);
let spec = PartitionSpec {
partition_id: 1,
columns: vec![PartitionSpecColumn {
partition_key_index: 0,
column_id: 10,
transform: PartitionTransform::Bucket(8),
}],
prune_safe: true,
};
let err = PartitionWriteSpec::resolve(&spec, &[10], &schema).unwrap_err();
assert!(
err.to_string().contains("bucket(8)"),
"expected the transform in the error, got: {err}"
);
}
#[test]
fn resolve_allows_temporal_transform_on_tz_aware_timestamp() {
let tz = DataType::Timestamp(arrow::datatypes::TimeUnit::Microsecond, Some("UTC".into()));
let schema = Schema::new(vec![Field::new("ts", tz, true)]);
let spec = PartitionSpec {
partition_id: 1,
columns: vec![PartitionSpecColumn {
partition_key_index: 0,
column_id: 10,
transform: PartitionTransform::Year,
}],
prune_safe: true,
};
assert!(PartitionWriteSpec::resolve(&spec, &[10], &schema).is_ok());
}
#[test]
fn resolve_errors_when_partition_column_absent_from_write_schema() {
let schema = Schema::new(vec![Field::new("id", DataType::Int64, false)]);
let spec = PartitionSpec {
partition_id: 1,
columns: vec![PartitionSpecColumn {
partition_key_index: 0,
column_id: 99,
transform: PartitionTransform::Identity,
}],
prune_safe: true,
};
assert!(PartitionWriteSpec::resolve(&spec, &[10], &schema).is_err());
}
#[test]
fn parse_and_roundtrip() {
for (s, expected) in [
("identity", PartitionTransform::Identity),
("year", PartitionTransform::Year),
("MONTH", PartitionTransform::Month),
("day", PartitionTransform::Day),
("hour", PartitionTransform::Hour),
("bucket(8)", PartitionTransform::Bucket(8)),
] {
assert_eq!(PartitionTransform::parse(s), expected);
}
for t in [
PartitionTransform::Identity,
PartitionTransform::Year,
PartitionTransform::Month,
PartitionTransform::Day,
PartitionTransform::Hour,
PartitionTransform::Bucket(4),
] {
assert_eq!(PartitionTransform::parse(&t.to_catalog_string()), t);
}
}
#[test]
fn unknown_transform_preserved() {
let t = PartitionTransform::parse("truncate(10)");
assert_eq!(t, PartitionTransform::Unknown("truncate(10)".to_string()));
assert_eq!(t.to_catalog_string(), "truncate(10)");
assert!(!t.is_producible());
assert_eq!(t.source_bounds("x", &DataType::Utf8), None);
}
#[test]
fn identity_bounds_are_exact() {
let (min, max) = PartitionTransform::Identity
.source_bounds("42", &DataType::Int32)
.unwrap();
assert_eq!(min, ScalarValue::Int32(Some(42)));
assert_eq!(max, ScalarValue::Int32(Some(42)));
let (min, max) = PartitionTransform::Identity
.source_bounds("us", &DataType::Utf8)
.unwrap();
assert_eq!(min, ScalarValue::Utf8(Some("us".to_string())));
assert_eq!(max, ScalarValue::Utf8(Some("us".to_string())));
}
#[test]
fn year_bounds_span_the_year_for_dates() {
let (min, max) = PartitionTransform::Year
.source_bounds("2023", &DataType::Date32)
.unwrap();
let expected_min =
ScalarValue::try_from_string("2023-01-01".to_string(), &DataType::Date32).unwrap();
let expected_max =
ScalarValue::try_from_string("2024-01-01".to_string(), &DataType::Date32).unwrap();
assert_eq!(min, expected_min);
assert_eq!(max, expected_max);
}
#[test]
fn year_on_tz_aware_timestamp_has_no_bounds() {
let tz = DataType::Timestamp(arrow::datatypes::TimeUnit::Microsecond, Some("UTC".into()));
assert_eq!(PartitionTransform::Year.source_bounds("2023", &tz), None);
let naive = DataType::Timestamp(arrow::datatypes::TimeUnit::Microsecond, None);
assert!(
PartitionTransform::Year
.source_bounds("2023", &naive)
.is_some()
);
}
#[test]
fn non_order_preserving_transforms_have_no_bounds() {
for t in [
PartitionTransform::Month,
PartitionTransform::Day,
PartitionTransform::Hour,
PartitionTransform::Bucket(4),
] {
assert_eq!(t.source_bounds("6", &DataType::Date32), None);
}
}
}