use std::collections::HashSet;
use std::fmt;
use std::slice::Iter as SliceIter;
use std::vec::IntoIter as VecIntoIter;
use indexmap::IndexMap;
use substrait::proto;
use substrait::proto::expression::field_reference::ReferenceType;
use substrait::proto::expression::literal::LiteralType;
use substrait::proto::expression::{RexType, reference_segment};
use super::{Explainable, ExtensionError};
use crate::textify::expressions::Reference;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum AddendumKind {
Enhancement,
Optimization,
ExtensionTable,
}
impl AddendumKind {
pub(crate) fn prefix(self) -> &'static str {
match self {
AddendumKind::Enhancement => "Enh",
AddendumKind::Optimization => "Opt",
AddendumKind::ExtensionTable => "Ext",
}
}
}
#[derive(Debug, Clone)]
pub struct Expr(Box<proto::Expression>);
impl Expr {
pub fn field(index: i32) -> Self {
Reference(index).into()
}
pub fn as_proto(&self) -> &proto::Expression {
self.0.as_ref()
}
pub fn to_proto(&self) -> proto::Expression {
self.as_proto().clone()
}
pub fn as_direct_reference(&self) -> Option<i32> {
let Some(RexType::Selection(field_ref)) = self.as_proto().rex_type.as_ref() else {
return None;
};
let Some(ReferenceType::DirectReference(segment)) = field_ref.reference_type.as_ref()
else {
return None;
};
let Some(reference_segment::ReferenceType::StructField(field)) =
segment.reference_type.as_ref()
else {
return None;
};
if field.child.is_some() {
return None;
}
Some(field.field)
}
}
impl From<proto::Expression> for Expr {
fn from(expr: proto::Expression) -> Self {
Expr(Box::new(expr))
}
}
impl From<proto::expression::Literal> for Expr {
fn from(literal: proto::expression::Literal) -> Self {
proto::Expression {
rex_type: Some(RexType::Literal(literal)),
}
.into()
}
}
impl From<Reference> for Expr {
fn from(reference: Reference) -> Self {
proto::Expression::from(reference).into()
}
}
impl From<Expr> for proto::Expression {
fn from(expr: Expr) -> Self {
*expr.0
}
}
impl From<i64> for Expr {
fn from(value: i64) -> Self {
proto::expression::Literal {
literal_type: Some(LiteralType::I64(value)),
nullable: false,
type_variation_reference: 0,
}
.into()
}
}
impl From<f64> for Expr {
fn from(value: f64) -> Self {
proto::expression::Literal {
literal_type: Some(LiteralType::Fp64(value)),
nullable: false,
type_variation_reference: 0,
}
.into()
}
}
impl From<bool> for Expr {
fn from(value: bool) -> Self {
proto::expression::Literal {
literal_type: Some(LiteralType::Boolean(value)),
nullable: false,
type_variation_reference: 0,
}
.into()
}
}
impl From<String> for Expr {
fn from(value: String) -> Self {
proto::expression::Literal {
literal_type: Some(LiteralType::String(value)),
nullable: false,
type_variation_reference: 0,
}
.into()
}
}
impl From<&str> for Expr {
fn from(value: &str) -> Self {
value.to_string().into()
}
}
#[derive(Debug, Clone, Default)]
pub struct ExtensionArgs {
pub positional: Vec<ExtensionValue>,
pub named: IndexMap<String, ExtensionValue>,
pub output_columns: Vec<ExtensionColumn>,
}
pub struct ArgsAccess<'a> {
args: &'a ExtensionArgs,
handled: HashSet<&'a str>,
positional_handled: bool,
}
impl<'a> ArgsAccess<'a> {
pub(crate) fn new(args: &'a ExtensionArgs) -> Self {
Self {
args,
handled: HashSet::new(),
positional_handled: false,
}
}
pub fn positional(&mut self) -> &'a [ExtensionValue] {
self.positional_handled = true;
&self.args.positional
}
pub fn output_columns(&self) -> &'a [ExtensionColumn] {
&self.args.output_columns
}
pub fn get_named_arg(&mut self, name: &str) -> Option<&'a ExtensionValue> {
match self.args.named.get_key_value(name) {
Some((k, value)) => {
self.handled.insert(k);
Some(value)
}
None => None,
}
}
pub fn get_named<T>(&mut self, name: &str) -> Result<Option<T>, ExtensionError>
where
T: TryFrom<&'a ExtensionValue>,
T::Error: Into<ExtensionError>,
{
self.get_named_arg(name)
.map(|value| {
T::try_from(value).map_err(|error| ExtensionError::NamedArgumentConversion {
name: name.to_string(),
source: Box::new(error.into()),
})
})
.transpose()
}
pub fn get_named_tuple<T>(&mut self, name: &str) -> Result<Option<Vec<T>>, ExtensionError>
where
T: TryFrom<&'a ExtensionValue>,
T::Error: Into<ExtensionError>,
{
self.get_named_arg(name)
.map(|value| {
let tuple = <&TupleValue>::try_from(value)?;
tuple
.into_iter()
.map(|element| T::try_from(element).map_err(Into::into))
.collect::<Result<Vec<_>, ExtensionError>>()
})
.transpose()
.map_err(|source| ExtensionError::NamedArgumentConversion {
name: name.to_string(),
source: Box::new(source),
})
}
pub fn expect_named<T>(&mut self, name: &str) -> Result<T, ExtensionError>
where
T: TryFrom<&'a ExtensionValue>,
T::Error: Into<ExtensionError>,
{
self.get_named(name)?
.ok_or_else(|| ExtensionError::MissingArgument {
name: name.to_string(),
})
}
pub(crate) fn finish(self) -> Result<(), ExtensionError> {
if !self.positional_handled && !self.args.positional.is_empty() {
return Err(ExtensionError::InvalidArgument(format!(
"Unhandled positional arguments: {}",
self.args.positional.len()
)));
}
let mut unhandled_args = Vec::new();
for name in self.args.named.keys() {
if !self.handled.contains(name.as_str()) {
unhandled_args.push(name.as_str());
}
}
if unhandled_args.is_empty() {
Ok(())
} else {
unhandled_args.sort();
Err(ExtensionError::InvalidArgument(format!(
"Unknown named arguments: {}",
unhandled_args.join(", ")
)))
}
}
}
#[derive(Debug, Clone)]
pub struct TupleValue(Vec<ExtensionValue>);
impl TupleValue {
pub fn len(&self) -> usize {
self.0.len()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn iter(&self) -> SliceIter<'_, ExtensionValue> {
self.0.iter()
}
}
impl<'a> IntoIterator for &'a TupleValue {
type Item = &'a ExtensionValue;
type IntoIter = SliceIter<'a, ExtensionValue>;
fn into_iter(self) -> Self::IntoIter {
self.0.iter()
}
}
impl IntoIterator for TupleValue {
type Item = ExtensionValue;
type IntoIter = VecIntoIter<ExtensionValue>;
fn into_iter(self) -> Self::IntoIter {
self.0.into_iter()
}
}
impl FromIterator<ExtensionValue> for TupleValue {
fn from_iter<I: IntoIterator<Item = ExtensionValue>>(iter: I) -> Self {
TupleValue(iter.into_iter().collect())
}
}
impl From<Vec<ExtensionValue>> for TupleValue {
fn from(items: Vec<ExtensionValue>) -> Self {
TupleValue(items)
}
}
#[derive(Debug, Clone)]
pub enum ExtensionValue {
String(String),
Integer(i64),
Float(f64),
Boolean(bool),
Null,
Error(ExtensionError),
Expr(Expr),
Enum(String),
Tuple(TupleValue),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExtensionValueKind {
String,
Integer,
Float,
Boolean,
Null,
Error,
Reference,
Enum,
Tuple,
Expression,
}
impl fmt::Display for ExtensionValueKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
ExtensionValueKind::String => write!(f, "string"),
ExtensionValueKind::Integer => write!(f, "integer"),
ExtensionValueKind::Float => write!(f, "float"),
ExtensionValueKind::Boolean => write!(f, "boolean"),
ExtensionValueKind::Null => write!(f, "null"),
ExtensionValueKind::Error => write!(f, "error"),
ExtensionValueKind::Reference => write!(f, "reference"),
ExtensionValueKind::Enum => write!(f, "enum"),
ExtensionValueKind::Tuple => write!(f, "tuple"),
ExtensionValueKind::Expression => write!(f, "expression"),
}
}
}
impl ExtensionValue {
pub fn kind(&self) -> ExtensionValueKind {
match self {
ExtensionValue::String(_) => ExtensionValueKind::String,
ExtensionValue::Integer(_) => ExtensionValueKind::Integer,
ExtensionValue::Float(_) => ExtensionValueKind::Float,
ExtensionValue::Boolean(_) => ExtensionValueKind::Boolean,
ExtensionValue::Null => ExtensionValueKind::Null,
ExtensionValue::Error(_) => ExtensionValueKind::Error,
ExtensionValue::Expr(_) => ExtensionValueKind::Expression,
ExtensionValue::Enum(_) => ExtensionValueKind::Enum,
ExtensionValue::Tuple(_) => ExtensionValueKind::Tuple,
}
}
}
impl From<ExtensionError> for ExtensionValue {
fn from(error: ExtensionError) -> Self {
ExtensionValue::Error(error)
}
}
impl From<Expr> for ExtensionValue {
fn from(expr: Expr) -> Self {
ExtensionValue::Expr(expr)
}
}
impl From<proto::Expression> for ExtensionValue {
fn from(expr: proto::Expression) -> Self {
Expr::from(expr).into()
}
}
impl From<proto::expression::Literal> for ExtensionValue {
fn from(literal: proto::expression::Literal) -> Self {
Expr::from(literal).into()
}
}
impl From<Reference> for ExtensionValue {
fn from(reference: Reference) -> Self {
Expr::from(reference).into()
}
}
impl From<i64> for ExtensionValue {
fn from(value: i64) -> Self {
ExtensionValue::Integer(value)
}
}
impl From<f64> for ExtensionValue {
fn from(value: f64) -> Self {
ExtensionValue::Float(value)
}
}
impl From<bool> for ExtensionValue {
fn from(value: bool) -> Self {
ExtensionValue::Boolean(value)
}
}
impl From<String> for ExtensionValue {
fn from(value: String) -> Self {
ExtensionValue::String(value)
}
}
impl From<&str> for ExtensionValue {
fn from(value: &str) -> Self {
ExtensionValue::String(value.to_string())
}
}
impl<T> From<Vec<T>> for ExtensionValue
where
T: Into<ExtensionValue>,
{
fn from(values: Vec<T>) -> Self {
ExtensionValue::Tuple(values.into_iter().map(Into::into).collect())
}
}
impl ExtensionError {
fn invalid_type(expected: ExtensionValueKind, actual: &ExtensionValue) -> Self {
match actual {
ExtensionValue::Error(source) => Self::ArgumentConversion {
expected,
source: Box::new(source.clone()),
},
_ => Self::InvalidArgumentType {
expected,
actual: actual.kind(),
},
}
}
}
impl<'a> TryFrom<&'a ExtensionValue> for &'a str {
type Error = ExtensionError;
fn try_from(value: &'a ExtensionValue) -> Result<&'a str, Self::Error> {
match value {
ExtensionValue::String(s) => Ok(s),
v => Err(ExtensionError::invalid_type(ExtensionValueKind::String, v)),
}
}
}
impl TryFrom<&ExtensionValue> for String {
type Error = ExtensionError;
fn try_from(value: &ExtensionValue) -> Result<String, Self::Error> {
<&str>::try_from(value).map(ToOwned::to_owned)
}
}
impl TryFrom<ExtensionValue> for String {
type Error = ExtensionError;
fn try_from(value: ExtensionValue) -> Result<String, Self::Error> {
String::try_from(&value)
}
}
pub struct EnumValue(pub String);
impl<'a> TryFrom<&'a ExtensionValue> for EnumValue {
type Error = ExtensionError;
fn try_from(value: &'a ExtensionValue) -> Result<EnumValue, Self::Error> {
match value {
ExtensionValue::Enum(s) => Ok(EnumValue(s.clone())),
v => Err(ExtensionError::invalid_type(ExtensionValueKind::Enum, v)),
}
}
}
impl<'a> TryFrom<&'a ExtensionValue> for &'a TupleValue {
type Error = ExtensionError;
fn try_from(value: &'a ExtensionValue) -> Result<&'a TupleValue, Self::Error> {
match value {
ExtensionValue::Tuple(tv) => Ok(tv),
v => Err(ExtensionError::invalid_type(ExtensionValueKind::Tuple, v)),
}
}
}
impl TryFrom<&ExtensionValue> for i64 {
type Error = ExtensionError;
fn try_from(value: &ExtensionValue) -> Result<i64, Self::Error> {
match value {
ExtensionValue::Integer(i) => Ok(*i),
v => Err(ExtensionError::invalid_type(ExtensionValueKind::Integer, v)),
}
}
}
impl TryFrom<&ExtensionValue> for f64 {
type Error = ExtensionError;
fn try_from(value: &ExtensionValue) -> Result<f64, Self::Error> {
match value {
ExtensionValue::Float(f) => Ok(*f),
v => Err(ExtensionError::invalid_type(ExtensionValueKind::Float, v)),
}
}
}
impl TryFrom<&ExtensionValue> for bool {
type Error = ExtensionError;
fn try_from(value: &ExtensionValue) -> Result<bool, Self::Error> {
match value {
ExtensionValue::Boolean(b) => Ok(*b),
v => Err(ExtensionError::invalid_type(ExtensionValueKind::Boolean, v)),
}
}
}
impl TryFrom<&ExtensionValue> for Reference {
type Error = ExtensionError;
fn try_from(value: &ExtensionValue) -> Result<Reference, Self::Error> {
match value {
ExtensionValue::Expr(expr) => expr
.as_direct_reference()
.map(Reference)
.ok_or_else(|| ExtensionError::invalid_type(ExtensionValueKind::Reference, value)),
v => Err(ExtensionError::invalid_type(
ExtensionValueKind::Reference,
v,
)),
}
}
}
impl TryFrom<&ExtensionValue> for Expr {
type Error = ExtensionError;
fn try_from(value: &ExtensionValue) -> Result<Expr, Self::Error> {
match value {
ExtensionValue::Expr(e) => Ok(e.clone()),
ExtensionValue::Integer(i) => Ok(Expr::from(*i)),
ExtensionValue::Float(f) => Ok(Expr::from(*f)),
ExtensionValue::String(s) => Ok(Expr::from(s.as_str())),
ExtensionValue::Boolean(b) => Ok(Expr::from(*b)),
v => Err(ExtensionError::invalid_type(
ExtensionValueKind::Expression,
v,
)),
}
}
}
#[derive(Debug, Clone)]
pub enum ExtensionColumn {
Named {
name: String,
r#type: proto::Type,
},
Expr(Expr),
}
impl ExtensionColumn {
pub fn field(index: i32) -> Self {
Self::Expr(Expr::field(index))
}
}
impl ExtensionArgs {
pub fn parse<T>(&self) -> Result<T, ExtensionError>
where
T: Explainable,
{
let mut access = ArgsAccess::new(self);
let value = T::from_args(&mut access)?;
access.finish()?;
Ok(value)
}
pub fn push<T>(&mut self, value: T)
where
T: Into<ExtensionValue>,
{
self.positional.push(value.into());
}
pub fn insert<K, V>(&mut self, name: K, value: V) -> Option<ExtensionValue>
where
K: Into<String>,
V: Into<ExtensionValue>,
{
self.named.insert(name.into(), value.into())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_named_invalid_type(
error: &ExtensionError,
name: &str,
expected: ExtensionValueKind,
actual: ExtensionValueKind,
) {
assert!(
matches!(
error,
ExtensionError::NamedArgumentConversion {
name: actual_name,
source,
} if actual_name == name
&& matches!(
source.as_ref(),
ExtensionError::InvalidArgumentType {
expected: actual_expected,
actual: actual_actual,
} if *actual_expected == expected && *actual_actual == actual
)
),
"unexpected error: {error:?}"
);
}
#[test]
fn get_named_converts_present_values_and_returns_none_for_missing_values() {
let mut args = ExtensionArgs::default();
args.insert("count", 8_i64);
let mut access = ArgsAccess::new(&args);
assert_eq!(access.get_named::<i64>("count").unwrap(), Some(8));
assert_eq!(access.get_named::<i64>("missing").unwrap(), None);
assert!(access.finish().is_ok());
}
#[test]
fn get_named_contextualizes_conversion_errors() {
let mut args = ExtensionArgs::default();
args.insert("count", ExtensionValue::Null);
let mut access = ArgsAccess::new(&args);
let error = access
.get_named::<i64>("count")
.expect_err("null should not convert to i64");
assert_eq!(
error.to_string(),
"Invalid named argument 'count': Invalid argument: expected integer, got null"
);
assert!(access.finish().is_ok());
}
#[test]
fn get_named_tuple_converts_present_values_and_returns_none_for_missing_values() {
let mut args = ExtensionArgs::default();
args.insert(
"names",
ExtensionValue::Tuple(vec!["first".into(), "second".into()].into()),
);
let mut access = ArgsAccess::new(&args);
assert_eq!(access.get_named_tuple::<String>("missing").unwrap(), None);
assert_eq!(
access.get_named_tuple::<String>("names").unwrap(),
Some(vec!["first".to_string(), "second".to_string()])
);
assert!(access.finish().is_ok());
}
#[test]
fn get_named_tuple_contextualizes_invalid_outer_value() {
let mut args = ExtensionArgs::default();
args.insert("names", "not a tuple");
let mut access = ArgsAccess::new(&args);
let error = access
.get_named_tuple::<String>("names")
.expect_err("string should not convert to tuple");
assert_named_invalid_type(
&error,
"names",
ExtensionValueKind::Tuple,
ExtensionValueKind::String,
);
assert!(access.finish().is_ok());
}
#[test]
fn get_named_tuple_contextualizes_invalid_element() {
let mut args = ExtensionArgs::default();
args.insert(
"names",
ExtensionValue::Tuple(vec!["first".into(), 2_i64.into()].into()),
);
let mut access = ArgsAccess::new(&args);
let error = access
.get_named_tuple::<String>("names")
.expect_err("integer should not convert to string");
assert_named_invalid_type(
&error,
"names",
ExtensionValueKind::String,
ExtensionValueKind::Integer,
);
assert!(access.finish().is_ok());
}
#[test]
fn vector_encodes_as_tuple() {
let encoded: ExtensionValue = vec![1_i64, 2_i64].into();
let tuple = <&TupleValue>::try_from(&encoded).unwrap();
let values = tuple
.iter()
.map(i64::try_from)
.collect::<Result<Vec<_>, _>>()
.unwrap();
assert_eq!(values, vec![1, 2]);
}
#[test]
fn error_value_reports_expected_type_and_source_when_extracted() {
let value = ExtensionValue::Error(ExtensionError::Custom("bad value".to_string()));
let error = i64::try_from(&value).expect_err("error value should not convert");
assert_eq!(
error.to_string(),
"Cannot convert argument to integer: bad value"
);
assert!(matches!(
&error,
ExtensionError::ArgumentConversion {
expected: ExtensionValueKind::Integer,
source,
} if matches!(source.as_ref(), ExtensionError::Custom(message) if message == "bad value")
));
}
#[test]
fn expect_named_reports_missing_argument_name() {
let args = ExtensionArgs::default();
let mut access = ArgsAccess::new(&args);
let error = access
.expect_named::<i64>("count")
.expect_err("missing argument should fail");
assert_eq!(error.to_string(), "Missing required argument: count");
assert!(access.finish().is_ok());
}
}