use std::error::Error;
use ordermap::OrderMap;
use vecmap::VecMap;
use crate::err::NonMergeableTypes;
use crate::sql::expression::apply::{FunctionCallApply, Lambda};
use crate::sql::expression::identifier::{CompoundIdentifier, Identifier, SimpleIdentifier};
use crate::sql::expression::literal::{ColumnReference, NullLiteral, RowLiteral};
use crate::sql::expression::Cast;
use crate::sql::expression::{Leaf, SQLExpression};
use crate::sql::query::projection::{Binding, ColumnProjection, Projection};
use crate::sql::types::SQLRowType;
use crate::tree::ast::identifier::SimpleIdentifier as AstSimpleIdentifier;
use crate::types::array::Array;
use crate::types::struct_type::Struct;
use crate::types::Type;
#[derive(Clone, Debug)]
pub enum RowLiteralEntry {
Leaf(Type, SQLExpression),
Node(Box<ProjectionBuilder>),
}
impl RowLiteralEntry {
pub fn leaf(type_: Type, expression: SQLExpression) -> Self {
RowLiteralEntry::Leaf(type_, expression)
}
pub fn node(builder: ProjectionBuilder) -> Self {
RowLiteralEntry::Node(Box::new(builder))
}
}
#[derive(Clone, Debug, Default)]
pub struct ProjectionBuilder {
bindings: VecMap<SimpleIdentifier, RowLiteralEntry>,
}
impl ProjectionBuilder {
pub fn deep_initialize_from_struct(name: Option<Identifier>, hst: Struct) -> Self {
let mut builder = Self::default();
for (ast_id, typ) in hst.into_iter() {
let id: SimpleIdentifier = ast_id.into();
let root_ident = name
.clone()
.map(|n| CompoundIdentifier::from_two_idents(n, id.clone().into()).into())
.unwrap_or(id.clone().into());
if let Type::Struct(s) = typ {
builder.bindings.insert(
id.clone(),
RowLiteralEntry::node(Self::deep_initialize_from_struct(Some(root_ident), s)),
);
} else {
builder.bindings.insert(
id,
RowLiteralEntry::leaf(typ.clone(), ColumnReference::new(root_ident).into()),
);
}
}
builder
}
pub fn expand_repeated_struct(
array: SQLExpression,
from: &Type,
to: &Type,
) -> Result<SQLExpression, ExpandError> {
Self::expand_repeated_struct_inner(array, from, to, false)
}
fn expand_repeated_struct_inner(
array: SQLExpression,
from: &Type,
to: &Type,
anonymous_row: bool,
) -> Result<SQLExpression, ExpandError> {
if let Type::Array(Array {
element_type: from_elt,
}) = from
{
if let Type::Array(Array {
element_type: to_elt,
}) = to
{
if let Type::Struct(from_struct) = &**from_elt {
if let Type::Struct(to_struct) = &**to_elt {
let e = SimpleIdentifier::new("e");
let expanded = Self::deep_initialize_from_struct(
Some(e.clone().into()),
from_struct.clone(),
)
.expand(to_struct.clone())?;
let lambda_body: SQLExpression = if anonymous_row {
expanded.build_literal().into()
} else {
expanded
.build_cast()
.map_err(|e| ExpandError::Fatal(e.into()))?
.into()
};
return Ok(FunctionCallApply::with_two(
"transform",
array,
Lambda::new(vec![e], lambda_body).into(),
)
.into());
}
}
}
}
Err(ExpandError::NonMergeable(NonMergeableTypes::new(
from.clone(),
to.clone(),
)))
}
pub fn expand(&self, to: Struct) -> Result<Self, ExpandError> {
self.expand_inner(to, false)
}
fn expand_inner(&self, to: Struct, anonymous_array_casts: bool) -> Result<Self, ExpandError> {
let mut ret = Self::default();
for (ast_key, to_type) in to.into_iter() {
let to_key: SimpleIdentifier = ast_key.into();
if let Some(from_value) = self.bindings.get(&to_key) {
match from_value {
RowLiteralEntry::Leaf(t, expression) if t == &to_type => {
ret.bind(to_key.clone().into(), expression.clone(), t.clone());
}
RowLiteralEntry::Leaf(t, expression) => {
ret.bind(
to_key.clone().into(),
Self::expand_repeated_struct_inner(
expression.clone(),
t,
&to_type,
anonymous_array_casts,
)?,
to_type.clone(),
);
}
RowLiteralEntry::Node(nested_pb) => {
if let Type::Struct(to_type) = to_type {
if nested_pb.clone().build_hamelin_type() == to_type {
ret.bindings.insert(to_key.clone(), from_value.clone());
} else {
ret.bindings.insert(
to_key.clone(),
RowLiteralEntry::node(nested_pb.expand_inner(to_type, true)?),
);
}
} else {
return Err(ExpandError::ExpectedStruct(to_key.clone(), to_type));
}
}
}
} else {
ret.bind(
to_key.clone().into(),
NullLiteral::default().into(),
to_type.clone(),
);
}
}
Ok(ret)
}
pub fn bind(&mut self, id: Identifier, value: SQLExpression, type_: Type) {
match id {
Identifier::Simple(simple_id) => {
self.bindings
.insert(simple_id.clone(), RowLiteralEntry::leaf(type_, value));
}
Identifier::Compound(compound_id) => {
if let Some(RowLiteralEntry::Node(node)) =
self.bindings.get_mut(compound_id.first())
{
node.bind(compound_id.rest(), value, type_);
} else {
let mut bldr = ProjectionBuilder::default();
bldr.bind(compound_id.rest(), value, type_);
self.bindings
.insert(compound_id.first().clone(), RowLiteralEntry::node(bldr));
}
}
}
}
pub fn with_binding(mut self, id: Identifier, value: SQLExpression, type_: Type) -> Self {
self.bind(id, value, type_);
self
}
pub fn bind_all(&self, other: &Self) -> Self {
let mut ret = self.clone();
for (key, value) in other.bindings.iter() {
match value {
RowLiteralEntry::Leaf(typ, expression) => {
ret.bind(key.clone().into(), expression.clone(), typ.clone());
}
RowLiteralEntry::Node(bldr) => {
ret.bindings
.insert(key.clone(), RowLiteralEntry::node(*bldr.clone()));
}
}
}
ret
}
pub fn bind_column_reference(&mut self, ident: SimpleIdentifier, type_: Type) {
self.bindings.insert(
ident.clone(),
RowLiteralEntry::leaf(type_, ColumnReference::new(ident.into()).into()),
);
}
pub fn initialize_key(&mut self, ident: SimpleIdentifier, struct_type: Struct) {
self.bindings.insert(
ident.clone(),
RowLiteralEntry::node(Self::deep_initialize_from_struct(
Some(ident.into()),
struct_type,
)),
);
}
pub fn is_present(&self, id: &Identifier) -> bool {
match id {
Identifier::Simple(simple_id) => self.bindings.contains_key(simple_id),
Identifier::Compound(compound_id) => {
if let Some(RowLiteralEntry::Node(node)) = self.bindings.get(compound_id.first()) {
node.is_present(&compound_id.rest())
} else {
false
}
}
}
}
pub fn build_literal(&self) -> RowLiteral {
let mut entries = vec![];
for (_, value) in self.bindings.iter() {
match value {
RowLiteralEntry::Leaf(_, expression) => {
entries.push(expression.clone());
}
RowLiteralEntry::Node(n) => {
entries.push(n.build_literal().into());
}
}
}
RowLiteral::new(entries)
}
pub fn build_sql_type(&self) -> anyhow::Result<SQLRowType> {
let mut fields = OrderMap::new();
for (key, value) in self.bindings.iter() {
match value {
RowLiteralEntry::Leaf(t, _) => {
fields.insert(key.clone(), t.clone().to_sql()?);
}
RowLiteralEntry::Node(n) => {
fields.insert(key.clone(), n.build_sql_type()?.into());
}
}
}
Ok(SQLRowType::new(fields))
}
pub fn build_cast(&self) -> anyhow::Result<Cast> {
Ok(Cast::new(
self.build_literal().into(),
self.build_sql_type()?.into(),
))
}
pub fn build_hamelin_type(self) -> Struct {
let mut fields: VecMap<AstSimpleIdentifier, Type> = VecMap::new();
for (key, value) in self.bindings.into_iter() {
let ast_key: AstSimpleIdentifier = key.into();
match value {
RowLiteralEntry::Leaf(t, _) => {
fields.insert(ast_key, t);
}
RowLiteralEntry::Node(n) => {
fields.insert(ast_key, n.build_hamelin_type().into());
}
}
}
Struct::new(fields)
}
pub fn build_projections(self) -> anyhow::Result<Vec<Projection>> {
let mut projections = vec![];
for (key, value) in self.bindings.into_iter() {
match value {
RowLiteralEntry::Leaf(
_,
SQLExpression::Leaf(Leaf::ColumnReference(ColumnReference { identifier })),
) if identifier == key.clone().into() => {
projections.push(ColumnProjection::new(identifier).into());
}
RowLiteralEntry::Leaf(_, expression) => {
projections.push(Binding::new(key.clone(), expression).into());
}
RowLiteralEntry::Node(n) => {
projections.push(Binding::new(key.clone(), n.build_cast()?.into()).into());
}
}
}
Ok(projections)
}
}
#[derive(Debug, thiserror::Error)]
pub enum ExpandError {
#[error("While expanding, expected key {0} to be struct, but found {1:?}")]
ExpectedStruct(SimpleIdentifier, Type),
#[error(transparent)]
NonMergeable(#[from] NonMergeableTypes),
#[error("While expanding, encountered fatal error: {0}")]
Fatal(Box<dyn Error + Send + Sync>),
}
#[cfg(test)]
mod tests {
use crate::{
sql::{
expression::literal::{IntegerLiteral, StringLiteral},
types::SQLBaseType,
},
types::{INT, STRING},
};
use super::*;
#[test]
pub fn test_empty() {
let pb = ProjectionBuilder::default();
assert_eq!(pb.build_literal(), RowLiteral::new(vec![]));
assert_eq!(
pb.build_sql_type().unwrap(),
SQLRowType::new(OrderMap::new())
);
assert_eq!(pb.clone().build_hamelin_type(), Struct::new(VecMap::new()));
assert_eq!(pb.build_projections().unwrap(), vec![]);
}
#[test]
pub fn test_flat() {
let mut pb = ProjectionBuilder::default();
pb.bind(
SimpleIdentifier::new("a").into(),
IntegerLiteral::new("1").into(),
INT,
);
pb.bind(
SimpleIdentifier::new("b").into(),
StringLiteral::new("hello").into(),
STRING,
);
assert_eq!(
RowLiteral::new(vec![
IntegerLiteral::new("1").into(),
StringLiteral::new("hello").into()
]),
pb.build_literal()
);
assert_eq!(
SQLRowType::default()
.with_str("a", SQLBaseType::BigInt.into())
.with_str("b", SQLBaseType::VarChar.into()),
pb.build_sql_type().unwrap()
);
assert_eq!(
Struct::default().with_str("a", INT).with_str("b", STRING),
pb.build_hamelin_type()
);
}
#[test]
pub fn test_nested() {
let mut pb = ProjectionBuilder::default();
pb.bind(
SimpleIdentifier::new("a").into(),
IntegerLiteral::new("1").into(),
INT,
);
pb.bind(
SimpleIdentifier::new("b").into(),
StringLiteral::new("hello").into(),
STRING,
);
pb.bind(
CompoundIdentifier::from_two_str("c", "d").into(),
IntegerLiteral::new("2").into(),
INT,
);
assert_eq!(
RowLiteral::new(vec![
IntegerLiteral::new("1").into(),
StringLiteral::new("hello").into(),
RowLiteral::new(vec![IntegerLiteral::new("2").into()]).into()
]),
pb.build_literal()
);
assert_eq!(
SQLRowType::default()
.with_str("a", SQLBaseType::BigInt.into())
.with_str("b", SQLBaseType::VarChar.into())
.with_str(
"c",
SQLRowType::default()
.with_str("d", SQLBaseType::BigInt.into())
.into()
),
pb.build_sql_type().unwrap()
);
assert_eq!(
Struct::default()
.with_str("a", INT)
.with_str("b", STRING)
.with_str("c", Struct::default().with_str("d", INT).into()),
pb.build_hamelin_type()
);
}
#[test]
pub fn test_shadow() {
let mut pb = ProjectionBuilder::default();
pb.bind(
SimpleIdentifier::new("a").into(),
IntegerLiteral::new("1").into(),
INT,
);
pb.bind(
SimpleIdentifier::new("b").into(),
StringLiteral::new("hello").into(),
STRING,
);
pb.bind(
CompoundIdentifier::from_two_str("c", "d").into(),
IntegerLiteral::new("2").into(),
INT,
);
pb.bind(
CompoundIdentifier::from_two_str("c", "e").into(),
IntegerLiteral::new("3").into(),
INT,
);
pb.bind(
SimpleIdentifier::new("c").into(),
StringLiteral::new("world").into(),
STRING,
);
assert_eq!(
RowLiteral::new(vec![
IntegerLiteral::new("1").into(),
StringLiteral::new("hello").into(),
StringLiteral::new("world").into()
]),
pb.build_literal()
);
assert_eq!(
SQLRowType::default()
.with_str("a", SQLBaseType::BigInt.into())
.with_str("b", SQLBaseType::VarChar.into())
.with_str("c", SQLBaseType::VarChar.into()),
pb.build_sql_type().unwrap()
);
assert_eq!(
Struct::default()
.with_str("a", INT)
.with_str("b", STRING)
.with_str("c", STRING),
pb.build_hamelin_type()
);
}
#[test]
pub fn test_expand() {
let mut pb = ProjectionBuilder::default();
pb.bind(
SimpleIdentifier::new("a").into(),
IntegerLiteral::new("1").into(),
INT,
);
pb.bind(
SimpleIdentifier::new("b").into(),
StringLiteral::new("hello").into(),
STRING,
);
let expanded_struct = Struct::default()
.with_str("a", INT)
.with_str("b", STRING)
.with_str("c", INT);
let expanded = pb.expand(expanded_struct.clone()).unwrap();
assert_eq!(
RowLiteral::new(vec![
IntegerLiteral::new("1").into(),
StringLiteral::new("hello").into(),
NullLiteral::default().into()
]),
expanded.build_literal()
);
assert_eq!(
SQLRowType::default()
.with_str("a", SQLBaseType::BigInt.into())
.with_str("b", SQLBaseType::VarChar.into())
.with_str("c", SQLBaseType::BigInt.into()),
expanded.build_sql_type().unwrap()
);
assert_eq!(expanded_struct, expanded.clone().build_hamelin_type());
let contracted_struct = Struct::default().with_str("b", STRING).with_str("c", INT);
let contracted = expanded.expand(contracted_struct.clone()).unwrap();
assert_eq!(
RowLiteral::new(vec![
StringLiteral::new("hello").into(),
NullLiteral::default().into()
]),
contracted.build_literal()
);
assert_eq!(
SQLRowType::default()
.with_str("b", SQLBaseType::VarChar.into())
.with_str("c", SQLBaseType::BigInt.into()),
contracted.build_sql_type().unwrap()
);
assert_eq!(contracted_struct, contracted.build_hamelin_type());
let overlapping_struct = Struct::default().with_str("new", INT).with_str("b", STRING);
let overlapping = pb.expand(overlapping_struct.clone()).unwrap();
assert_eq!(
RowLiteral::new(vec![
NullLiteral::default().into(),
StringLiteral::new("hello").into(),
]),
overlapping.build_literal()
);
assert_eq!(
SQLRowType::default()
.with_str("new", SQLBaseType::BigInt.into())
.with_str("b", SQLBaseType::VarChar.into()),
overlapping.build_sql_type().unwrap()
);
assert_eq!(overlapping_struct, overlapping.build_hamelin_type());
}
}