use std::any::TypeId;
use std::collections::VecDeque;
use std::fmt::{Debug, Display, Formatter};
use std::sync::Arc;
use anyhow::{anyhow, bail};
use derive_more::From;
use itertools::Itertools;
use vecmap::VecMap;
pub use crate::func::err::{
FunctionBindFailure, FunctionParameterBindingFailure, FunctionParameterBindingFailures,
FunctionTranslationFailure, MatchTestFailure, NoBindingError,
};
use crate::translation::ExpressionTranslation;
use crate::tree::ast::expression::Expression;
use crate::tree::typed_ast::context::ExpressionTranslationContext;
use crate::tree::typed_ast::expression::TypedExpression;
use crate::types::matcher::Matcher;
use crate::types::Type;
pub trait HasType {
fn typ(&self) -> &Type;
}
pub trait ParameterBindingProvider {
fn get_by_name(&self, name: &str) -> anyhow::Result<&dyn HasType>;
fn get_by_index(&self, index: usize) -> anyhow::Result<&dyn HasType>;
fn iter(&self) -> Box<dyn Iterator<Item = &dyn HasType> + '_>;
fn len(&self) -> usize;
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum SpecialPosition {
Agg,
Window,
Match,
}
impl Display for SpecialPosition {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
match self {
SpecialPosition::Agg => write!(f, "agg function"),
SpecialPosition::Window => write!(f, "window function"),
SpecialPosition::Match => write!(f, "match function"),
}
}
}
use crate::sql::expression::OrderByExpression;
use crate::sql::query::window::WindowReference;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct FunctionTranslationContext {
pub window: Option<WindowReference>,
pub specials_allowed: Vec<SpecialPosition>,
pub order_by: Vec<OrderByExpression>,
}
impl FunctionTranslationContext {
pub fn with_window(mut self, window: WindowReference) -> Self {
self.window = Some(window);
self
}
pub fn with_special_allowed(mut self, special: SpecialPosition) -> Self {
self.specials_allowed.push(special);
self
}
pub fn with_order_by(mut self, mut order_by: Vec<OrderByExpression>) -> Self {
self.order_by.append(&mut order_by);
self
}
pub fn add_order_by(&mut self, mut order_by: Vec<OrderByExpression>) {
self.order_by.append(&mut order_by);
}
}
pub trait FunctionDef: Send + Sync + 'static {
fn name(&self) -> &'static str;
fn parameters(&self) -> Parameters {
Parameters::new()
}
fn custom_bind(
&self,
_ast_binding: &ParameterBinding<Arc<Expression>>,
_ctx: &mut ExpressionTranslationContext,
) -> anyhow::Result<Option<ParameterBinding<Arc<TypedExpression>>>> {
Ok(None)
}
fn refine_binding(
&self,
binding: ParameterBinding<Arc<TypedExpression>>,
) -> anyhow::Result<ParameterBinding<Arc<TypedExpression>>> {
Ok(binding)
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type>;
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
None
}
fn is_deterministic(&self) -> bool {
true
}
fn sortable_input(&self) -> bool {
false
}
fn manages_window_clause(&self) -> bool {
false
}
}
impl Debug for dyn FunctionDef {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FunctionDef")
.field("name", &self.name())
.field("parameters", &self.parameters())
.finish()
}
}
impl Display for dyn FunctionDef {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "{}({})", self.name(), self.parameters())
}
}
pub struct DirectResolution<T: HasType> {
pub function_def: Arc<dyn FunctionDef>,
pub binding: ParameterBinding<T>,
pub typ: Type,
}
pub struct BroadcastResolution<T: HasType> {
pub function_def: Arc<dyn FunctionDef>,
pub binding: ParameterBinding<T>,
pub typ: Type,
pub broadcast_position: usize,
}
#[derive(From)]
pub enum FunctionResolution<T: HasType> {
Direct(DirectResolution<T>),
Broadcast(BroadcastResolution<T>),
}
#[derive(Debug, Default)]
pub struct Parameters {
parameters: Vec<Parameter>,
var_args: Option<Box<dyn Matcher + Send + Sync>>,
}
impl Parameters {
pub fn new() -> Self {
Self::default()
}
pub fn with<M: Matcher + Send + Sync + 'static>(mut self, name: &str, typ: M) -> Self {
self.parameters.push(Parameter::of(name, typ));
self
}
pub fn with_default<M: Matcher + Send + Sync + 'static>(
mut self,
name: &str,
typ: M,
default: Expression,
) -> Self {
self.parameters
.push(Parameter::with_default(name, typ, default));
self
}
pub fn with_var_args<M: Matcher + Send + Sync + 'static>(mut self, typ: M) -> Self {
self.var_args = Some(Box::new(typ));
self
}
pub fn add(mut self, parameter: Parameter) -> Self {
self.parameters.push(parameter);
self
}
pub fn add_var_args<M: Matcher + Send + Sync + 'static>(mut self, typ: M) -> Self {
self.var_args = Some(Box::new(typ));
self
}
pub fn bind_and_check<T>(
&self,
positional: Vec<T>,
named: impl IntoIterator<Item = (String, T)>,
) -> Result<ParameterBinding<T>, Vec<FunctionParameterBindingFailure>>
where
T: HasType + Clone + TryFrom<Expression>,
<T as TryFrom<Expression>>::Error: std::fmt::Debug,
{
let binding = self.bind(positional, named)?;
self.check(&binding)?;
Ok(binding)
}
pub fn bind<T>(
&self,
positional: Vec<T>,
named: impl IntoIterator<Item = (String, T)>,
) -> Result<ParameterBinding<T>, Vec<FunctionParameterBindingFailure>>
where
T: Clone + TryFrom<Expression>,
<T as TryFrom<Expression>>::Error: std::fmt::Debug,
{
let mut arguments = VecMap::new();
let mut variable_arguments = VecDeque::new();
let mut errors = Vec::new();
for param in &self.parameters {
let default_value = param.default_value.as_ref().map(|v| {
v.clone()
.try_into()
.expect("default value conversion should never fail")
});
arguments.insert(param.name.clone(), default_value);
}
for (name, arg) in named.into_iter() {
if !arguments.contains_key(name.as_str()) {
errors.push(FunctionParameterBindingFailure::UnknownNamedParameter(name));
} else {
arguments.insert(name, Some(arg));
}
}
let mut param_iter = self.parameters.iter();
let num_positional = positional.len();
let mut positional_args_iter = positional.into_iter();
while let Some(param) = param_iter.next() {
if let Some(arg) = positional_args_iter.next() {
arguments.insert(param.name.clone(), Some(arg));
}
}
while let Some(arg) = positional_args_iter.next() {
if let Some(_) = &self.var_args {
variable_arguments.push_back(arg);
} else {
errors.push(
FunctionParameterBindingFailure::TooManyPositionalArguments {
expected: self.parameters.len(),
got: num_positional,
},
);
break;
}
}
let mut final_arguments = VecMap::new();
for (name, arg) in arguments.into_iter() {
match (name, arg) {
(name, None) => {
errors.push(FunctionParameterBindingFailure::MissingRequiredArgument(
name,
));
}
(name, Some(arg)) => {
final_arguments.insert(name, arg);
}
}
}
if !errors.is_empty() {
return Err(errors);
}
Ok(ParameterBinding {
arguments: final_arguments,
var_args: variable_arguments,
})
}
pub fn check<T: HasType>(
&self,
binding: &ParameterBinding<T>,
) -> Result<(), Vec<FunctionParameterBindingFailure>> {
let mut errors = Vec::new();
for ((name, arg), param) in binding.arguments.iter().zip(self.parameters.iter()) {
if !param.typ.matches(arg.typ()) {
errors.push(FunctionParameterBindingFailure::TypeCheckFailure {
name: name.clone(),
expected: param.typ.to_string(),
got: arg.typ().clone(),
});
}
}
if let Some(var_args_matcher) = &self.var_args {
for arg in &binding.var_args {
if !var_args_matcher.matches(arg.typ()) {
errors.push(
FunctionParameterBindingFailure::TypeCheckFailureForVarargs {
expected: var_args_matcher.to_string(),
got: arg.typ().clone(),
},
);
}
}
}
if !errors.is_empty() {
return Err(errors);
}
Ok(())
}
pub fn autocomplete_snippet(&self) -> String {
self.parameters
.iter()
.enumerate()
.map(|(i, p)| format!("${{{}:{}}}", i + 1, p.name.clone()))
.join(", ")
}
}
impl Display for Parameters {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
let mut first = true;
for param in &self.parameters {
if first {
first = false;
} else {
write!(f, ", ")?;
}
write!(f, "{}", param)?;
}
if let Some(var_args) = &self.var_args {
if !first {
write!(f, ", ")?;
}
write!(f, " ...: {}", var_args)?;
}
Ok(())
}
}
pub struct Parameter {
name: String,
typ: Box<dyn Matcher + Send + Sync>,
default_value: Option<Expression>,
}
impl Parameter {
pub fn of<M: Matcher + Send + Sync + 'static>(name: &str, typ: M) -> Self {
debug_assert_eq!(name, name.to_lowercase());
Self {
name: name.to_string(),
typ: Box::new(typ),
default_value: None,
}
}
pub fn with_default<M: Matcher + Send + Sync + 'static>(
name: &str,
typ: M,
default: Expression,
) -> Self {
debug_assert_eq!(name, name.to_lowercase());
Self {
name: name.to_string(),
typ: Box::new(typ),
default_value: Some(default),
}
}
pub fn default_value(&self) -> Option<&Expression> {
self.default_value.as_ref()
}
}
impl Display for Parameter {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
write!(f, "{}={}", self.name, self.typ)
}
}
impl Debug for Parameter {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Parameter")
.field("name", &self.name)
.field("typ", &self.typ)
.field("default_value", &self.default_value)
.finish()
}
}
#[derive(Debug, Clone)]
pub struct ParameterBinding<T> {
arguments: VecMap<String, T>,
var_args: VecDeque<T>,
}
impl<T> ParameterBinding<T> {
pub fn get_by_name(&self, name: &str) -> anyhow::Result<&T> {
self.arguments.get(name).ok_or_else(|| {
anyhow!(NoBindingError {
name: name.to_string(),
})
})
}
pub fn get_by_index(&self, index: usize) -> anyhow::Result<&T> {
if index < self.arguments.len() {
self.arguments.values().nth(index).ok_or_else(|| {
anyhow!(NoBindingError {
name: format!("index {}", index),
})
})
} else {
let var_index = index - self.arguments.len();
self.var_args.get(var_index).ok_or_else(|| {
anyhow!(NoBindingError {
name: format!("index {}", index),
})
})
}
}
pub fn take_by_name(&mut self, name: &str) -> anyhow::Result<T> {
self.arguments.swap_remove(name).ok_or_else(|| {
anyhow!(NoBindingError {
name: name.to_string(),
})
})
}
pub fn take(&mut self) -> anyhow::Result<T> {
if !self.arguments.is_empty() {
Ok(self.arguments.remove_index(0).1)
} else {
self.var_args
.pop_front()
.ok_or_else(|| anyhow!("No more arguments to take"))
}
}
pub fn iter(&self) -> impl Iterator<Item = &T> {
self.arguments.values().chain(self.var_args.iter())
}
pub fn into_iter(self) -> impl Iterator<Item = T> {
self.arguments
.into_iter()
.map(|(_, v)| v)
.chain(self.var_args.into_iter())
}
pub fn len(&self) -> usize {
self.arguments.len() + self.var_args.len()
}
pub fn is_empty(&self) -> bool {
self.arguments.is_empty() && self.var_args.is_empty()
}
pub fn map<U, F: FnMut(T) -> U>(self, mut f: F) -> ParameterBinding<U> {
ParameterBinding {
arguments: self.arguments.into_iter().map(|(k, v)| (k, f(v))).collect(),
var_args: self.var_args.into_iter().map(&mut f).collect(),
}
}
pub fn try_map<U, E, F: FnMut(T) -> Result<U, E>>(
self,
mut f: F,
) -> Result<ParameterBinding<U>, E> {
let mut arguments = VecMap::new();
for (k, v) in self.arguments {
arguments.insert(k, f(v)?);
}
let mut var_args = VecDeque::new();
for v in self.var_args {
var_args.push_back(f(v)?);
}
Ok(ParameterBinding {
arguments,
var_args,
})
}
pub fn replace_by_name(&mut self, name: &str, value: T) -> anyhow::Result<()> {
if self.arguments.contains_key(name) {
self.arguments.insert(name.to_string(), value);
Ok(())
} else {
Err(anyhow!(NoBindingError {
name: name.to_string(),
}))
}
}
pub fn replace_by_index(mut self, index: usize, value: T) -> anyhow::Result<Self> {
if index < self.arguments.len() {
let key = self
.arguments
.keys()
.nth(index)
.ok_or_else(|| {
anyhow!(NoBindingError {
name: format!("index {}", index),
})
})?
.clone();
self.arguments.insert(key, value);
Ok(self)
} else {
let var_index = index - self.arguments.len();
if var_index < self.var_args.len() {
self.var_args[var_index] = value;
Ok(self)
} else {
Err(anyhow!(NoBindingError {
name: format!("index {}", index),
}))
}
}
}
pub fn empty() -> Self {
ParameterBinding {
arguments: VecMap::new(),
var_args: VecDeque::new(),
}
}
pub fn from_named(pairs: impl IntoIterator<Item = (String, T)>) -> Self {
ParameterBinding {
arguments: pairs.into_iter().collect(),
var_args: VecDeque::new(),
}
}
pub fn insert(&mut self, name: String, value: T) {
self.arguments.insert(name, value);
}
}
impl<T> ParameterBinding<Option<T>> {
pub fn take_first_defined(self) -> anyhow::Result<T> {
let mut found: Option<T> = None;
for opt in self.into_iter() {
if let Some(val) = opt {
if found.is_some() {
bail!("Expected exactly one variable, found two or more (multiple constants)");
}
found = Some(val);
}
}
found.ok_or_else(|| anyhow!("Expected exactly one variable, found zero (no constants)"))
}
}
impl HasType for ExpressionTranslation {
fn typ(&self) -> &Type {
&self.typ
}
}
impl HasType for Arc<TypedExpression> {
fn typ(&self) -> &Type {
&self.resolved_type
}
}
impl<T: HasType> ParameterBindingProvider for ParameterBinding<T> {
fn get_by_name(&self, name: &str) -> anyhow::Result<&dyn HasType> {
self.get_by_name(name).map(|v| v as &dyn HasType)
}
fn get_by_index(&self, index: usize) -> anyhow::Result<&dyn HasType> {
self.get_by_index(index).map(|v| v as &dyn HasType)
}
fn len(&self) -> usize {
self.len()
}
fn iter(&self) -> Box<dyn Iterator<Item = &dyn HasType> + '_> {
Box::new(self.iter().map(|v| v as &dyn HasType))
}
}