use rudb_common::{Error, Field, LogicalType, Result, Value};
use rudb_plan::{Expr, ExprRef};
use crate::binder::Binder;
use crate::fold;
pub(crate) const STRUCT_PACK: &str = "struct_pack";
pub(crate) const STRUCT_EXTRACT: &str = "struct_extract";
impl Binder<'_> {
pub(crate) fn pack_struct(&mut self, names: &[String], values: &[ExprRef]) -> Result<ExprRef> {
for (at, name) in names.iter().enumerate() {
if !name.is_empty()
&& names[..at].iter().any(|earlier| earlier.eq_ignore_ascii_case(name))
{
return Err(Error::binder(format!(
"Duplicate named argument \"{name}\" in function call to '\"struct_pack\"'"
)));
}
}
if values.is_empty() {
return Ok(self.add_constant(Value::Struct(Vec::new())));
}
let fields = names
.iter()
.zip(values)
.map(|(name, &value)| Field::new(name.clone(), self.plan().expr_type(value).clone()))
.collect();
let args = self.plan_mut().add_expr_list(values);
let recorded = self.plan_mut().intern(STRUCT_PACK);
Ok(self.add_expr(Expr::Function { name: recorded, args }, LogicalType::Struct(fields)))
}
pub(crate) fn struct_field(
&mut self,
written: &str,
bound: &[ExprRef],
) -> Result<Option<ExprRef>> {
let extract = rudb_catalog::same_name(written, STRUCT_EXTRACT);
let subscript = ["array_extract", "list_extract", "list_element"]
.iter()
.any(|name| rudb_catalog::same_name(written, name));
let &[input, key] = bound else {
return Ok(None);
};
let LogicalType::Struct(fields) = self.plan().expr_type(input).clone() else {
return Ok(None);
};
if !extract && !subscript {
return Ok(None);
}
let at = match fold::value_of(self.plan(), key) {
Ok(Some(Value::Varchar(name))) => self.field_named(&fields, &name)?,
Ok(Some(value)) if value.logical_type().is_integer() && Field::unnamed(&fields) => {
let index = value.as_i64().unwrap_or(0);
if index < 1 || index > fields.len() as i64 {
return Err(Error::binder(format!(
"Key index {index} for struct_extract out of range - expected an index \
between 1 and {}",
fields.len()
)));
}
index as usize - 1
}
Ok(Some(value)) if value.logical_type().is_integer() => {
return Err(Error::binder(
"struct_extract with an integer key can only be used on unnamed structs, use \
a string key instead",
));
}
_ => {
return Err(Error::binder(
"Key name for struct_extract needs to be a constant string",
));
}
};
let key = self.add_constant(Value::BigInt(at as i64 + 1));
let args = self.plan_mut().add_expr_list(&[input, key]);
let recorded = self.plan_mut().intern(STRUCT_EXTRACT);
Ok(Some(self.add_expr(Expr::Function { name: recorded, args }, fields[at].ty.clone())))
}
fn field_named(&self, fields: &[Field], name: &str) -> Result<usize> {
fields.iter().position(|field| field.name.eq_ignore_ascii_case(name)).ok_or_else(|| {
let entries: Vec<String> =
fields.iter().map(|field| format!("\"{}\"", field.name)).collect();
Error::binder(format!(
"Could not find key \"{name}\" in struct\n\nCandidate Entries: {}",
entries.join(", ")
))
})
}
pub(crate) fn struct_call(
&mut self,
written: &str,
bound: &[ExprRef],
) -> Result<Option<ExprRef>> {
let name = written.to_ascii_lowercase();
let types: Vec<LogicalType> =
bound.iter().map(|&arg| self.plan().expr_type(arg).clone()).collect();
let fields_of = |at: usize| match types.get(at) {
Some(LogicalType::Struct(fields)) => Some(fields.clone()),
_ => None,
};
let returns = match name.as_str() {
"struct_insert" | "struct_update" => {
let Some(mut fields) = fields_of(0) else {
return Ok(None);
};
let Some(added) = fields_of(1).filter(|_| bound.len() == 2) else {
if bound.len() == 1 && name == "struct_insert" {
return Err(Error::invalid_input("Can't insert nothing into a STRUCT"));
}
return Err(Error::binder(format!(
"Need named argument for struct {}, e.g., a := b",
&name[7..]
)));
};
for field in added {
let same =
fields.iter().position(|one| one.name.eq_ignore_ascii_case(&field.name));
match same {
Some(_) if name == "struct_insert" => {
return Err(Error::binder(format!(
"Duplicate struct entry name \"\"{}\"\"",
field.name
)));
}
Some(at) => fields[at] = field,
None => fields.push(field),
}
}
LogicalType::Struct(fields)
}
"struct_concat" => {
let mut fields: Vec<Field> = Vec::new();
let mut unnamed = None;
for (at, ty) in types.iter().enumerate() {
let LogicalType::Struct(held) = ty else {
return Err(Error::invalid_input(format!(
"struct_concat: Argument at position \"{}\" is not a STRUCT",
at + 1
)));
};
let this = Field::unnamed(held);
if *unnamed.get_or_insert(this) != this {
return Err(Error::invalid_input(
"struct_concat: Cannot mix named and unnamed STRUCTs",
));
}
for field in held {
if !this
&& fields.iter().any(|one| one.name.eq_ignore_ascii_case(&field.name))
{
return Err(Error::invalid_input(format!(
"struct_concat: Arguments contain duplicate STRUCT entry \"{}\"",
field.name
)));
}
fields.push(field.clone());
}
}
if fields.is_empty() {
return Ok(None);
}
LogicalType::Struct(fields)
}
"struct_keys" | "struct_values" => {
let Some(fields) = fields_of(0).filter(|_| bound.len() == 1) else {
return Ok(None);
};
if name == "struct_keys" {
if Field::unnamed(&fields) {
return Err(Error::invalid_input(
"struct_keys() expects a STRUCT argument",
));
}
LogicalType::list(LogicalType::Varchar)
} else {
let unnamed = fields.iter().map(|field| Field::new("", field.ty.clone()));
LogicalType::Struct(unnamed.collect())
}
}
"struct_contains" | "struct_position" => {
let Some(fields) = fields_of(0).filter(|_| bound.len() == 2) else {
return Ok(None);
};
if !Field::unnamed(&fields) {
return Err(Error::binder(format!(
"\"{name}\" can only be used on unnamed structs"
)));
}
if name == "struct_contains" { LogicalType::Boolean } else { LogicalType::Integer }
}
_ => return Ok(None),
};
let args = self.plan_mut().add_expr_list(bound);
let recorded = self.plan_mut().intern(&name);
Ok(Some(self.add_expr(Expr::Function { name: recorded, args }, returns)))
}
}