use toasty_core::{
driver::Capability,
schema::db,
stmt::{self, VisitMut},
};
use super::{Param, Ty};
pub(super) fn extract_values(
stmt: &mut stmt::Statement,
params: &mut Vec<Param>,
capability: &Capability,
) {
struct Extract<'a> {
params: &'a mut Vec<Param>,
bind_list_param: bool,
glob_starts_with: bool,
binary_like_starts_with: bool,
}
impl stmt::VisitMut for Extract<'_> {
fn visit_expr_mut(&mut self, expr: &mut stmt::Expr) {
match expr {
stmt::Expr::AnyOp(e) => {
self.visit_expr_mut(&mut e.lhs);
if let Some(arg) = extract_array_operand(&mut e.rhs, self.params) {
*e.rhs = arg;
} else {
self.visit_expr_mut(&mut e.rhs);
}
return;
}
stmt::Expr::AllOp(e) => {
self.visit_expr_mut(&mut e.lhs);
if let Some(arg) = extract_array_operand(&mut e.rhs, self.params) {
*e.rhs = arg;
} else {
self.visit_expr_mut(&mut e.rhs);
}
return;
}
stmt::Expr::InList(e) => {
self.visit_expr_mut(&mut e.expr);
if let stmt::Expr::Value(stmt::Value::List(_)) = e.list.as_ref() {
let stmt::Expr::Value(stmt::Value::List(items)) =
std::mem::replace(e.list.as_mut(), stmt::Expr::null())
else {
unreachable!()
};
let items = items
.into_iter()
.map(|v| value_to_extracted_expr(v, self.params, false))
.collect();
*e.list = stmt::Expr::List(stmt::ExprList { items });
} else {
self.visit_expr_mut(&mut e.list);
}
return;
}
stmt::Expr::StartsWith(e)
if self.glob_starts_with || self.binary_like_starts_with =>
{
self.visit_expr_mut(&mut e.expr);
let stmt::Expr::Value(stmt::Value::String(prefix)) = e.prefix.as_ref() else {
panic!("starts_with prefix must be a string literal");
};
let pattern = if self.glob_starts_with {
glob_prefix_pattern(prefix)
} else {
binary_like_prefix_pattern(prefix)
};
let position = self.params.len();
self.params.push(Param {
value: stmt::Value::String(pattern),
ty: Ty::Inferred(db::Type::Text),
});
*e.prefix = stmt::Expr::arg(position);
return;
}
_ => {}
}
if self.bind_list_param
&& is_scalar_list(expr)
&& let Some(arg) = extract_array_operand(expr, self.params)
{
*expr = arg;
return;
}
stmt::visit_mut::visit_expr_mut(self, expr);
match expr {
stmt::Expr::Value(value) if is_extractable_scalar(value) => {
let ty = infer_ty(value);
let position = self.params.len();
let value = std::mem::replace(value, stmt::Value::Null);
self.params.push(Param { value, ty });
*expr = stmt::Expr::arg(position);
}
stmt::Expr::Value(value @ stmt::Value::Object(_)) => {
let owned = std::mem::replace(value, stmt::Value::Null);
let position = self.params.len();
self.params.push(Param {
value: owned,
ty: Ty::Unknown,
});
*expr = stmt::Expr::arg(position);
}
stmt::Expr::Value(value @ (stmt::Value::Record(_) | stmt::Value::List(_))) => {
let owned = std::mem::replace(value, stmt::Value::Null);
*expr = value_to_extracted_expr(owned, self.params, self.bind_list_param);
}
_ => {}
}
}
}
Extract {
params,
bind_list_param: capability.bind_list_param,
glob_starts_with: capability.glob_starts_with,
binary_like_starts_with: capability.binary_like_starts_with,
}
.visit_mut(stmt);
}
fn is_extractable_scalar_expr(expr: &stmt::Expr) -> bool {
matches!(expr, stmt::Expr::Value(v) if is_extractable_scalar(v))
}
fn is_scalar_list(expr: &stmt::Expr) -> bool {
match expr {
stmt::Expr::List(list) => list.items.iter().all(is_extractable_scalar_expr),
stmt::Expr::Value(stmt::Value::List(items)) => items.iter().all(is_extractable_scalar),
_ => false,
}
}
fn extract_array_operand(expr: &mut stmt::Expr, params: &mut Vec<Param>) -> Option<stmt::Expr> {
let items: Vec<stmt::Value> = match expr {
stmt::Expr::Value(stmt::Value::List(_)) => {
let stmt::Expr::Value(stmt::Value::List(items)) =
std::mem::replace(expr, stmt::Expr::null())
else {
unreachable!()
};
items
}
stmt::Expr::List(list) if list.items.iter().all(|i| matches!(i, stmt::Expr::Value(_))) => {
let stmt::Expr::List(list) = std::mem::replace(expr, stmt::Expr::null()) else {
unreachable!()
};
list.items
.into_iter()
.map(|e| match e {
stmt::Expr::Value(v) => v,
_ => unreachable!(),
})
.collect()
}
_ => return None,
};
let value = stmt::Value::List(items);
let ty = infer_ty(&value);
let position = params.len();
params.push(Param { value, ty });
Some(stmt::Expr::arg(position))
}
fn value_to_extracted_expr(
value: stmt::Value,
params: &mut Vec<Param>,
bind_list_param: bool,
) -> stmt::Expr {
match value {
stmt::Value::Null => stmt::Expr::Value(stmt::Value::Null),
stmt::Value::Record(record) => {
let fields = record
.fields
.into_iter()
.map(|f| value_to_extracted_expr(f, params, bind_list_param))
.collect();
stmt::Expr::Record(stmt::ExprRecord::from_vec(fields))
}
stmt::Value::List(values)
if bind_list_param
&& values.iter().all(|v| {
is_extractable_scalar(v) || matches!(v, stmt::Value::Object(_))
}) =>
{
let value = stmt::Value::List(values);
let ty = infer_ty(&value);
let position = params.len();
params.push(Param { value, ty });
stmt::Expr::arg(position)
}
stmt::Value::List(values) => {
let items = values
.into_iter()
.map(|v| value_to_extracted_expr(v, params, bind_list_param))
.collect();
stmt::Expr::List(stmt::ExprList { items })
}
scalar => {
let ty = infer_ty(&scalar);
let position = params.len();
params.push(Param { value: scalar, ty });
stmt::Expr::arg(position)
}
}
}
fn is_extractable_scalar(value: &stmt::Value) -> bool {
!matches!(
value,
stmt::Value::Null | stmt::Value::Record(_) | stmt::Value::List(_) | stmt::Value::Object(_)
)
}
fn infer_ty(value: &stmt::Value) -> Ty {
use stmt::Value;
match value {
Value::Bool(_) => Ty::Inferred(db::Type::Boolean),
Value::I8(_) => Ty::Inferred(db::Type::Integer(1)),
Value::I16(_) => Ty::Inferred(db::Type::Integer(2)),
Value::I32(_) => Ty::Inferred(db::Type::Integer(4)),
Value::I64(_) => Ty::Inferred(db::Type::Integer(8)),
Value::U8(_) => Ty::Inferred(db::Type::UnsignedInteger(1)),
Value::U16(_) => Ty::Inferred(db::Type::UnsignedInteger(2)),
Value::U32(_) => Ty::Inferred(db::Type::UnsignedInteger(4)),
Value::U64(_) => Ty::Inferred(db::Type::UnsignedInteger(8)),
Value::String(_) => Ty::Inferred(db::Type::Text),
Value::Uuid(_) => Ty::Inferred(db::Type::Uuid),
Value::Bytes(_) => Ty::Inferred(db::Type::Blob),
#[cfg(feature = "rust_decimal")]
Value::Decimal(_) => Ty::Inferred(db::Type::Numeric(None)),
#[cfg(feature = "jiff")]
Value::Timestamp(_) => Ty::Inferred(db::Type::Timestamp(6)),
#[cfg(feature = "jiff")]
Value::Date(_) => Ty::Inferred(db::Type::Date),
#[cfg(feature = "jiff")]
Value::Time(_) => Ty::Inferred(db::Type::Time(6)),
#[cfg(feature = "jiff")]
Value::DateTime(_) => Ty::Inferred(db::Type::DateTime(6)),
Value::List(items) => {
let elem = items
.iter()
.find(|v| !v.is_null())
.map(infer_ty)
.unwrap_or(Ty::Unknown);
Ty::List(Box::new(elem))
}
_ => Ty::Unknown,
}
}
pub(super) fn glob_prefix_pattern(prefix: &str) -> String {
let mut pattern = String::with_capacity(prefix.len() * 3 + 1);
for c in prefix.chars() {
match c {
'*' => pattern.push_str("[*]"),
'?' => pattern.push_str("[?]"),
'[' => pattern.push_str("[[]"),
c => pattern.push(c),
}
}
pattern.push('*');
pattern
}
pub(super) fn binary_like_prefix_pattern(prefix: &str) -> String {
let mut pattern = String::with_capacity(prefix.len() + 1);
for c in prefix.chars() {
match c {
'!' | '%' | '_' => {
pattern.push('!');
pattern.push(c);
}
c => pattern.push(c),
}
}
pattern.push('%');
pattern
}