use crate::expressions::{parse_sql, Expression, Scalar};
use crate::schema::{DataType, StructField, StructType};
use crate::transforms::{transform_output_type, SchemaTransform};
use crate::{DeltaResult, Error};
#[derive(Debug, Clone, PartialEq)]
pub struct ColumnDefault<'a> {
raw_sql: String,
data_type: &'a DataType,
parsed_sql: Option<Expression>,
}
impl<'a> ColumnDefault<'a> {
pub(crate) fn new(raw_sql: String, data_type: &'a DataType) -> DeltaResult<Self> {
let is_null = raw_sql.trim().eq_ignore_ascii_case("null");
if matches!(data_type, DataType::Variant(_)) && !is_null {
return Err(Error::schema(format!(
"a Variant column's default must be NULL, got {raw_sql:?}"
)));
}
let parsed_sql = parse_sql(&raw_sql, data_type).ok();
Ok(Self {
raw_sql,
data_type,
parsed_sql,
})
}
pub fn raw_sql(&self) -> &str {
&self.raw_sql
}
pub fn data_type(&self) -> &DataType {
self.data_type
}
pub fn to_scalar(&self) -> DeltaResult<Option<Scalar>> {
match &self.parsed_sql {
None => Ok(None),
Some(Expression::Literal(scalar)) => Ok(Some(scalar.clone())),
Some(other) => Err(Error::generic(format!(
"kernel cannot evaluate non-literal column default expression: {other:?}"
))),
}
}
pub(crate) fn is_kernel_parsable_literal(&self) -> bool {
matches!(self.parsed_sql, Some(Expression::Literal(_)))
}
}
pub(crate) fn try_collect_column_defaults(
schema: &StructType,
) -> DeltaResult<Vec<(String, ColumnDefault<'_>)>> {
let mut collector = ColumnDefaultCollector {
path: Vec::new(),
defaults: Vec::new(),
};
collector.transform_struct(schema)?;
Ok(collector.defaults)
}
struct ColumnDefaultCollector<'a> {
path: Vec<String>,
defaults: Vec<(String, ColumnDefault<'a>)>,
}
impl<'a> ColumnDefaultCollector<'a> {
fn descend(&mut self, segment: &str, element: &'a DataType) -> DeltaResult<()> {
self.path.push(segment.to_string());
let result = self.transform(element);
self.path.pop();
result
}
}
impl<'a> SchemaTransform<'a> for ColumnDefaultCollector<'a> {
transform_output_type!(|'a, T| DeltaResult<()>);
fn transform_struct_field(&mut self, field: &'a StructField) -> DeltaResult<()> {
self.path.push(field.name().clone());
if let Some(column_default) = field.column_default()? {
self.defaults.push((self.path.join("."), column_default));
}
let result = self.recurse_into_struct_field(field);
self.path.pop();
result
}
fn transform_array_element(&mut self, etype: &'a DataType) -> DeltaResult<()> {
self.descend("element", etype)
}
fn transform_map_key(&mut self, ktype: &'a DataType) -> DeltaResult<()> {
self.descend("key", ktype)
}
fn transform_map_value(&mut self, vtype: &'a DataType) -> DeltaResult<()> {
self.descend("value", vtype)
}
fn transform_variant(&mut self, _stype: &'a StructType) -> DeltaResult<()> {
Ok(())
}
}
pub(crate) fn validate_column_defaults_metadata(schema: &StructType) -> DeltaResult<()> {
try_collect_column_defaults(schema)?;
Ok(())
}
#[cfg(test)]
pub(crate) fn field_with_default(
name: &str,
data_type: impl Into<DataType>,
raw_sql: &str,
) -> StructField {
use crate::schema::{ColumnMetadataKey, MetadataValue};
StructField::nullable(name, data_type).add_metadata([(
ColumnMetadataKey::CurrentDefault.as_ref().to_string(),
MetadataValue::String(raw_sql.to_string()),
)])
}
#[cfg(test)]
pub(crate) fn field_with_invalid_default(name: &str) -> StructField {
use crate::schema::{ColumnMetadataKey, MetadataValue};
StructField::nullable(name, DataType::INTEGER).add_metadata([(
ColumnMetadataKey::CurrentDefault.as_ref().to_string(),
MetadataValue::Number(7),
)])
}
#[cfg(test)]
mod tests {
use chrono::{DateTime, NaiveDate, TimeZone, Utc};
use rstest::rstest;
use super::*;
use crate::schema::{ArrayType, MapType, StructField};
fn struct_ty() -> DataType {
DataType::try_struct_type([StructField::nullable("a", DataType::INTEGER)]).unwrap()
}
fn struct_with_inner_default() -> DataType {
DataType::try_struct_type([field_with_default("inner", DataType::INTEGER, "42")]).unwrap()
}
fn date_days(year: i32, month: u32, day: u32) -> i32 {
let nd = NaiveDate::from_ymd_opt(year, month, day)
.unwrap()
.and_hms_opt(0, 0, 0)
.unwrap();
Utc.from_utc_datetime(&nd)
.signed_duration_since(DateTime::UNIX_EPOCH)
.num_days() as i32
}
#[derive(Debug)]
enum Expect {
Parsed(Scalar),
ParsedNull,
Unparsable,
NewErr(&'static str),
}
#[rstest]
#[case::integer("42", DataType::INTEGER, Expect::Parsed(Scalar::Integer(42)))]
#[case::string("'hello'", DataType::STRING, Expect::Parsed(Scalar::String("hello".into())))]
#[case::boolean("TRUE", DataType::BOOLEAN, Expect::Parsed(Scalar::Boolean(true)))]
#[case::date(
"DATE '2024-01-01'",
DataType::DATE,
Expect::Parsed(Scalar::Date(date_days(2024, 1, 1)))
)]
#[case::null_primitive("NULL", DataType::INTEGER, Expect::ParsedNull)]
#[case::null_array(
"NULL",
DataType::from(ArrayType::new(DataType::INTEGER, true)),
Expect::ParsedNull
)]
#[case::null_map(
"NULL",
DataType::from(MapType::new(DataType::STRING, DataType::INTEGER, true)),
Expect::ParsedNull
)]
#[case::null_struct("NULL", struct_ty(), Expect::ParsedNull)]
#[case::null_variant("NULL", DataType::unshredded_variant(), Expect::ParsedNull)]
#[case::function_call("current_timestamp()", DataType::TIMESTAMP, Expect::Unparsable)]
#[case::type_mismatch("'not an int'", DataType::INTEGER, Expect::Unparsable)]
#[case::arithmetic("1 + 1", DataType::INTEGER, Expect::Unparsable)]
#[case::non_primitive_array(
"ARRAY(1)",
DataType::from(ArrayType::new(DataType::INTEGER, true)),
Expect::Unparsable
)]
#[case::non_primitive_map(
"MAP('k', 1)",
DataType::from(MapType::new(DataType::STRING, DataType::INTEGER, true)),
Expect::Unparsable
)]
#[case::non_primitive_struct("STRUCT(1)", struct_ty(), Expect::Unparsable)]
#[case::non_null_variant("1", DataType::unshredded_variant(), Expect::NewErr("Variant"))]
fn column_default_from_new(
#[case] raw_sql: &str,
#[case] data_type: DataType,
#[case] expect: Expect,
) {
match (ColumnDefault::new(raw_sql.into(), &data_type), expect) {
(Ok(d), Expect::Parsed(scalar)) => {
assert_eq!(d.raw_sql(), raw_sql);
assert_eq!(d.data_type(), &data_type);
assert_eq!(d.to_scalar().unwrap(), Some(scalar));
}
(Ok(d), Expect::ParsedNull) => {
assert_eq!(
d.to_scalar().unwrap(),
Some(Scalar::Null(data_type.clone()))
);
}
(Ok(d), Expect::Unparsable) => {
assert_eq!(d.to_scalar().unwrap(), None);
assert_eq!(d.raw_sql(), raw_sql);
}
(Err(e), Expect::NewErr(needle)) => {
assert!(e.to_string().contains(needle), "got: {e}");
}
(result, expect) => {
panic!("unexpected outcome for {raw_sql:?}: {result:?} vs {expect:?}")
}
}
}
#[rstest]
#[case::well_formed_default(
vec![
field_with_default("c", DataType::INTEGER, "42"),
StructField::nullable("no_default", DataType::STRING),
],
None
)]
#[case::no_defaults(vec![StructField::nullable("c", DataType::INTEGER)], None)]
#[case::non_string_metadata(vec![field_with_invalid_default("c")], Some("non-string"))]
#[case::non_null_default_on_array_tolerated(
vec![field_with_default("arr", ArrayType::new(DataType::INTEGER, true), "ARRAY(1)")],
None
)]
#[case::non_null_default_on_variant_rejected(
vec![field_with_default("v", DataType::unshredded_variant(), "1")],
Some("Variant")
)]
#[case::nested_default(
vec![StructField::nullable(
"s",
DataType::try_struct_type([field_with_default("inner", DataType::INTEGER, "42")]).unwrap(),
)],
None
)]
fn validate_column_defaults_cases(
#[case] fields: Vec<StructField>,
#[case] expected_error: Option<&str>,
) {
let schema = StructType::try_new(fields).unwrap();
match (validate_column_defaults_metadata(&schema), expected_error) {
(Ok(()), None) => {}
(Err(e), Some(needle)) => assert!(e.to_string().contains(needle), "got: {e}"),
(result, expected) => panic!("unexpected outcome: {result:?} vs {expected:?}"),
}
}
#[rstest]
#[case::array_element(
DataType::from(ArrayType::new(struct_with_inner_default(), true)),
["arr", "element", "inner"]
)]
#[case::map_value(
DataType::from(MapType::new(DataType::STRING, struct_with_inner_default(), true)),
["arr", "value", "inner"]
)]
#[case::map_key(
DataType::from(MapType::new(struct_with_inner_default(), DataType::INTEGER, true)),
["arr", "key", "inner"]
)]
fn collect_nested_container_default_path(
#[case] container: DataType,
#[case] expected_path: [&str; 3],
) {
let schema = StructType::try_new([StructField::nullable("arr", container)]).unwrap();
let defaults = try_collect_column_defaults(&schema).unwrap();
let [(path, _)] = defaults.try_into().expect("exactly one default");
assert_eq!(path, expected_path.join("."));
}
#[test]
fn to_scalar_errors_on_non_literal_parsed_expression() {
let int_ty = DataType::INTEGER;
let d = ColumnDefault {
raw_sql: "x".into(),
data_type: &int_ty,
parsed_sql: Some(Expression::column(["x"])),
};
let err = d
.to_scalar()
.expect_err("non-literal parsed expression must error")
.to_string();
assert!(err.contains("non-literal"), "got: {err}");
}
}