use std::sync::Arc;
use crate::err::TranslationErrors;
use crate::tree::ast::expression::{Expression, ExpressionKind, IntLiteral, UnaryPrefixOperator};
use crate::tree::ast::node::Span;
use crate::tree::ast::ops::UnaryPrefixOp;
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::builder::*;
use crate::tree::tests::shared::catalog_with_tables;
use crate::tree::typed_ast::expression::{MapExpressionAlgebra, TypedExpression};
use crate::type_check_with_provider;
use crate::types::Type;
use crate::types::{INT, STRING, TIMESTAMP};
use pretty_assertions::assert_eq;
struct IdentityAlgebra;
impl MapExpressionAlgebra for IdentityAlgebra {
}
struct NegatingAlgebra;
impl MapExpressionAlgebra for NegatingAlgebra {
fn leaf(&mut self, expr: &TypedExpression) -> Arc<Expression> {
if let ExpressionKind::IntLiteral(IntLiteral { int }) = &expr.ast.kind {
if *int > 0 && *expr.resolved_type == INT {
return Arc::new(Expression {
span: Span::NONE,
kind: UnaryPrefixOperator {
operator: UnaryPrefixOp::Minus,
operand: expr.ast.clone(),
}
.into(),
});
}
}
expr.ast.clone()
}
}
fn verify_cata_roundtrip(pipeline: Pipeline) -> Result<(), TranslationErrors> {
verify_cata_roundtrip_with_tables(pipeline, vec![])
}
fn verify_cata_roundtrip_with_tables(
pipeline: Pipeline,
tables: Vec<(&str, Vec<(&str, Type)>)>,
) -> Result<(), TranslationErrors> {
let catalog = catalog_with_tables(tables.clone());
let typed = type_check_with_provider(pipeline.clone(), catalog).into_result()?;
let original_schema = typed.schema();
let mut alg = IdentityAlgebra;
let commands = typed.commands().expect("pipeline should be valid");
let transformed_commands: Vec<_> = commands
.iter()
.map(|cmd| cmd.cata_expressions(&mut alg))
.collect();
let transformed_pipeline = Pipeline {
span: pipeline.span.clone(),
commands: transformed_commands,
};
let catalog = catalog_with_tables(tables);
let retyped = type_check_with_provider(transformed_pipeline, catalog).into_result()?;
let retyped_schema = retyped.schema();
assert_eq!(
original_schema, retyped_schema,
"Schema mismatch after cata_expressions roundtrip"
);
Ok(())
}
#[test]
fn test_cata_where_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("severity", 5)
.named_field("active", true)
.named_field("timestamp", call("now"))
})
.where_cmd(and(
gt(field_ref("severity"), 3),
eq(field_ref("active"), true),
))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_within_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("timestamp", call("now")))
.within(hours(1))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_limit_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("x", 1))
.limit(add(10, 5))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_sort_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("severity", 5)
.named_field("timestamp", call("now"))
})
.sort(|s| s.desc(field_ref("severity")).asc(field_ref("timestamp")))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_let_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("severity", 5).named_field("message", "hello"))
.let_cmd(|l| {
l.named_field("doubled", multiply(field_ref("severity"), 2))
.named_field("upper_msg", call("upper").arg(field_ref("message")))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_select_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("severity", 5)
.named_field("message", "hello")
.named_field("timestamp", call("now"))
})
.select(|s| {
s.named_field("severity_plus_1", add(field_ref("severity"), 1))
.named_field("lower_msg", call("lower").arg(field_ref("message")))
.field("timestamp")
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_parse_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("message", "200 GET"))
.parse(|p| {
p.pattern(r"(?P<status>\d+) (?P<method>\w+)")
.identifier("status")
.identifier("method")
.source(field_ref("message"))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_agg_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("user_id", "u1").named_field("severity", 5))
.agg(|a| {
a.named_aggregate("cnt", call("count"))
.group_by("user_id")
.named_group("sev_plus_one", add(field_ref("severity"), 1))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_agg_command_with_sort() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("user_id", "u1")
.named_field("message", "hello")
.named_field("timestamp", call("now"))
})
.agg(|a| {
a.named_aggregate("msgs", call("array_agg").arg(field_ref("message")))
.sort(sort_command().asc(field_ref("timestamp")))
.group_by("user_id")
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_window_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("user_id", "u1")
.named_field("timestamp", call("now"))
})
.window(|w| {
w.named_field("rn", call("row_number"))
.group_by("user_id", field_ref("user_id"))
.sort(sort_command().asc(field_ref("timestamp")))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_window_command_with_within() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("timestamp", call("now")))
.window(|w| {
w.named_field("rn", call("row_number"))
.sort(sort_command().asc(field_ref("timestamp")))
.within(hours(1))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_explode_command() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
.explode(|e| e.named_field("item", field_ref("arr")))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_match_command_basic() -> Result<(), TranslationErrors> {
let p = pipeline()
.match_cmd(|m| {
m.named_quantified_pattern("x", "events", quantifier_exactly("2"))
.named_quantified_pattern("y", "events", quantifier_at_least_one())
.agg("cnt", call("count"))
.group_by("user_id")
.sort(sort_command().asc(field_ref("timestamp")))
})
.build();
verify_cata_roundtrip_with_tables(
p,
vec![(
"events",
vec![("timestamp", TIMESTAMP), ("user_id", STRING)],
)],
)
}
#[test]
fn test_cata_match_command_with_within() -> Result<(), TranslationErrors> {
let p = pipeline()
.match_cmd(|m| {
m.named_pattern("x", "events")
.named_quantified_pattern("y", "events", quantifier_at_least_one())
.agg("cnt", call("count"))
.group_by("user_id")
.sort(sort_command().asc(field_ref("timestamp")))
.within(hours(1))
})
.build();
verify_cata_roundtrip_with_tables(
p,
vec![(
"events",
vec![("timestamp", TIMESTAMP), ("user_id", STRING)],
)],
)
}
#[test]
fn test_cata_match_command_with_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.match_cmd(|m| {
m.named_quantified_pattern("x", "events", quantifier_at_least_one())
.agg("cnt", call("count"))
.agg("max_sev", call("max").arg(field_ref("severity")))
.group_by("user_id")
.named_group("upper_user", call("upper").arg(field_ref("user_id")))
.sort(sort_command().desc(field_ref("timestamp")))
})
.build();
verify_cata_roundtrip_with_tables(
p,
vec![(
"events",
vec![
("timestamp", TIMESTAMP),
("user_id", STRING),
("severity", INT.into()),
],
)],
)
}
#[test]
fn test_cata_complex_pipeline() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("severity", 5)
.named_field("active", true)
.named_field("message", "hello")
.named_field("timestamp", call("now"))
.named_field("user_id", "u1")
})
.where_cmd(and(
gt(field_ref("severity"), 1),
eq(field_ref("active"), true),
))
.let_cmd(|l| {
l.named_field("severity_doubled", multiply(field_ref("severity"), 2))
.named_field("msg_upper", call("upper").arg(field_ref("message")))
})
.select(|s| {
s.field("timestamp")
.field("user_id")
.field("severity_doubled")
.field("msg_upper")
})
.sort(|s| s.desc(field_ref("severity_doubled")))
.limit(100)
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_pipeline_with_agg_and_window() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("severity", 5).named_field("user_id", "u1"))
.where_cmd(gt(field_ref("severity"), 0))
.agg(|a| {
a.named_aggregate("total", call("count"))
.named_aggregate("max_sev", call("max").arg(field_ref("severity")))
.group_by("user_id")
})
.window(|w| {
w.named_field("rank", call("row_number"))
.sort(sort_command().desc(field_ref("total")))
})
.select(|s| {
s.field("user_id")
.field("total")
.field("max_sev")
.field("rank")
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_array_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("arr", array().element(1).element(2).element(3))
.named_field(
"nested",
array()
.element(array().element(1).element(2))
.element(array().element(3).element(4)),
)
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_struct_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field(
"obj",
struct_literal()
.field("a", 1)
.field("b", "hello")
.field("c", true),
)
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_tuple_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("tup", tuple().element(1).element("hello").element(true)))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_cast_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("x", cast(42, Type::String))
.named_field("y", cast("123", Type::Int))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_field_access_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("obj", struct_literal().field("a", 1).field("b", "hello")))
.let_cmd(|l| l.named_field("val", field(field_ref("obj"), "a")))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_function_calls() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field("x", call("upper").arg("hello"))
.named_field("y", call("abs").arg(negate(5)))
.named_field("z", call("coalesce").arg(null()).arg("default"))
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_lambda_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
.let_cmd(|l| {
l.named_field(
"doubled",
call("transform")
.arg(field_ref("arr"))
.arg(lambda1("x").body(multiply(field_ref("x"), 2))),
)
})
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_broadcast_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
.let_cmd(|l| l.named_field("scaled", multiply(field_ref("arr"), 10)))
.build();
verify_cata_roundtrip(p)
}
#[test]
fn test_cata_nested_expressions() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| {
l.named_field(
"result",
add(
divide(subtract(multiply(add(1, 2), 3), 4), 2),
call("abs").arg(negate(10)),
),
)
})
.build();
verify_cata_roundtrip(p)
}
fn verify_cata_transforms(pipeline: Pipeline) -> Result<(), TranslationErrors> {
verify_cata_transforms_with_tables(pipeline, vec![])
}
fn verify_cata_transforms_with_tables(
pipeline: Pipeline,
tables: Vec<(&str, Vec<(&str, Type)>)>,
) -> Result<(), TranslationErrors> {
let catalog = catalog_with_tables(tables.clone());
let typed = type_check_with_provider(pipeline.clone(), catalog).into_result()?;
let mut alg = NegatingAlgebra;
let commands = typed.commands().expect("pipeline should be valid");
let transformed_commands: Vec<_> = commands
.iter()
.map(|cmd| cmd.cata_expressions(&mut alg))
.collect();
let transformed_pipeline = Pipeline {
span: pipeline.span.clone(),
commands: transformed_commands,
};
let catalog = catalog_with_tables(tables);
type_check_with_provider(transformed_pipeline.clone(), catalog).into_result()?;
assert_ne!(
pipeline, transformed_pipeline,
"Transformation should have changed the pipeline"
);
Ok(())
}
#[test]
fn test_negating_transform_let() -> Result<(), TranslationErrors> {
let p = pipeline().let_cmd(|l| l.named_field("x", 42)).build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_where() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("x", 5))
.where_cmd(gt(field_ref("x"), 10))
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_select() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("x", 5))
.select(|s| s.named_field("y", add(field_ref("x"), 1)))
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_array() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_struct() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("obj", struct_literal().field("a", 1).field("b", 2)))
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_tuple() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("tup", tuple().element(1).element(2)))
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_function_arg() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("x", call("abs").arg(5)))
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_agg() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("x", 5).named_field("y", 10))
.agg(|a| {
a.named_aggregate("cnt", call("count"))
.named_group("bucket", divide(field_ref("x"), 10))
})
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_window() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("x", 5).named_field("timestamp", call("now")))
.window(|w| {
w.named_field("rn", call("row_number"))
.sort(sort_command().asc(add(field_ref("x"), 1)))
})
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_lambda() -> Result<(), TranslationErrors> {
let p = pipeline()
.let_cmd(|l| l.named_field("arr", array().element(1).element(2).element(3)))
.let_cmd(|l| {
l.named_field(
"doubled",
call("transform")
.arg(field_ref("arr"))
.arg(lambda1("elem").body(add(field_ref("elem"), 1))),
)
})
.build();
verify_cata_transforms(p)
}
#[test]
fn test_negating_transform_match() -> Result<(), TranslationErrors> {
let p = pipeline()
.match_cmd(|m| {
m.named_quantified_pattern("x", "events", quantifier_at_least_one())
.agg("total", add(call("count"), 1))
.group_by("user_id")
.sort(sort_command().asc(field_ref("timestamp")))
})
.build();
verify_cata_transforms_with_tables(
p,
vec![(
"events",
vec![("timestamp", TIMESTAMP), ("user_id", STRING)],
)],
)
}