use std::fmt;
use glaredb_error::{DbError, Result};
use super::Expression;
use crate::arrays::datatype::{DataType, DataTypeId};
use crate::explain::context_display::{ContextDisplay, ContextDisplayMode, ContextDisplayWrapper};
use crate::functions::cast::builtin::BUILTIN_CAST_FUNCTION_SETS;
use crate::functions::cast::{CastFlatten, CastFunctionSet, PlannedCastFunction, RawCastFunction};
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct CastExpr {
pub to: DataType,
pub expr: Box<Expression>,
pub cast_function: PlannedCastFunction,
}
impl CastExpr {
pub fn new_using_default_casts(expr: impl Into<Expression>, to: DataType) -> Result<Self> {
let target_id = to.id();
let cast_set = find_cast_function_set(target_id).ok_or_else(|| {
DbError::new(format!(
"Unable to find cast function to handle target type: {target_id}"
))
})?;
let expr = expr.into();
if let Expression::Cast(existing_cast) = &expr {
let child = &existing_cast.expr;
let child_datatype = child.datatype()?;
if let Some(cast_fn) = find_cast_function(cast_set, child_datatype.id()) {
if matches!(cast_fn.flatten, CastFlatten::Safe) {
let child = match expr {
Expression::Cast(cast) => cast.expr,
_ => unreachable!("expr variant checked in outer if statement"),
};
let bind_state = cast_fn.call_bind(&child_datatype, &to)?;
let planned = PlannedCastFunction {
name: cast_set.name,
raw: cast_fn,
state: bind_state,
};
return Ok(CastExpr {
to,
expr: child,
cast_function: planned,
});
}
}
}
let src_datatype = expr.datatype()?;
let cast_fn = find_cast_function(cast_set, src_datatype.id()).ok_or_else(|| {
DbError::new(format!(
"Cast function '{}' cannot handle source type {}",
cast_set.name, src_datatype,
))
})?;
let bind_state = cast_fn.call_bind(&src_datatype, &to)?;
let planned = PlannedCastFunction {
name: cast_set.name,
raw: cast_fn,
state: bind_state,
};
Ok(CastExpr {
to,
expr: Box::new(expr),
cast_function: planned,
})
}
}
impl ContextDisplay for CastExpr {
fn fmt_using_context(
&self,
mode: ContextDisplayMode,
f: &mut fmt::Formatter<'_>,
) -> fmt::Result {
write!(
f,
"CAST({} TO {})",
ContextDisplayWrapper::with_mode(self.expr.as_ref(), mode),
self.to
)
}
}
fn find_cast_function_set(target: DataTypeId) -> Option<&'static CastFunctionSet> {
BUILTIN_CAST_FUNCTION_SETS
.iter()
.find(|&cast_set| cast_set.target == target)
}
fn find_cast_function(set: &CastFunctionSet, src: DataTypeId) -> Option<&'static RawCastFunction> {
set.functions.iter().find(|&cast_fn| cast_fn.src == src)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::expr;
#[test]
fn no_flatten_unsafe() {
let cast = CastExpr::new_using_default_casts(
CastExpr::new_using_default_casts(expr::lit("123456789e-1234"), DataType::float32())
.unwrap(),
DataType::int64(),
)
.unwrap();
assert!(matches!(cast.expr.as_ref(), Expression::Cast(_)));
}
#[test]
fn flatten_safe() {
let cast = CastExpr::new_using_default_casts(
CastExpr::new_using_default_casts(expr::lit(14_i16), DataType::int32()).unwrap(),
DataType::int64(),
)
.unwrap();
assert_eq!(Expression::from(expr::lit(14_i16)), *cast.expr);
assert_eq!(DataType::int64(), cast.to);
}
}