use crate::err::TranslationErrors;
use crate::tree::ast::expression::IntervalUnit;
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::tree::tests::shared::catalog_with_tables;
use crate::tree::typed_ast::pipeline::ValidPipeline;
use crate::type_check;
use crate::type_check_with_provider;
use crate::types::array::Array;
use crate::types::map::Map;
use crate::types::struct_type::Struct;
use crate::types::tuple::Tuple;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, STRING};
use ordermap::OrderMap;
use pretty_assertions::assert_eq;
use rstest::rstest;
#[rstest]
#[case::select_single_field(
None,
"SET name = 'Alice', age = 30, email = 'alice@example.com' | SELECT name",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30).named_field("email", "alice@example.com"))
.select(|s| s.field("name")),
Struct::default().with_str("name", STRING)
)]
#[case::select_multiple_fields(
None,
"SET name = 'Alice', age = 30, email = 'alice@example.com' | SELECT name, age, email",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30).named_field("email", "alice@example.com"))
.select(|s| s.field("name").field("age").field("email")),
Struct::default().with_str("name", STRING).with_str("age", INT).with_str("email", STRING)
)]
#[case::select_named_field(
None,
"SET name = 'Alice' | SELECT full_name = name",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice"))
.select(|s| s.named_field("full_name", field_ref("name"))),
Struct::default().with_str("full_name", STRING)
)]
#[case::select_expression_field(
None,
"SET age = 30 | SELECT age + 1",
pipeline()
.set_cmd(|l| l.named_field("age", 30))
.select(|s| s.named_field("age+1", add(field_ref("age"), 1))),
Struct::default().with_str("age+1", INT)
)]
#[case::select_named_expression(
None,
"SET age = 30 | SELECT next_year = age + 1",
pipeline()
.set_cmd(|l| l.named_field("age", 30))
.select(|s| s.named_field("next_year", add(field_ref("age"), 1))),
Struct::default().with_str("next_year", INT)
)]
#[case::select_mixed_fields(
None,
"SET name = 'Alice', age = 30, email = 'alice@example.com' | SELECT name, birth_year = 2024 - age, email",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30).named_field("email", "alice@example.com"))
.select(|s| s.field("name").named_field("birth_year", subtract(2024, field_ref("age"))).field("email")),
Struct::default().with_str("name", STRING).with_str("birth_year", INT).with_str("email", STRING)
)]
#[case::where_simple_comparison(
None,
"SET name = 'Alice', age = 30 | WHERE age > 18",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.where_cmd(gt(field_ref("age"), 18)),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::where_equality(
None,
"SET name = 'Alice', age = 30 | WHERE name == 'Alice'",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.where_cmd(eq(field_ref("name"), "Alice")),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::where_logical_and(
None,
"SET name = 'Alice', age = 30, active = true | WHERE age > 18 AND active == true",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30).named_field("active", true))
.where_cmd(and(gt(field_ref("age"), 18), eq(field_ref("active"), true))),
Struct::default().with_str("name", STRING).with_str("age", INT).with_str("active", BOOLEAN)
)]
#[case::where_logical_or(
None,
"SET status = 'pending' | WHERE status == 'pending' OR status == 'approved'",
pipeline()
.set_cmd(|l| l.named_field("status", "pending"))
.where_cmd(or(eq(field_ref("status"), "pending"), eq(field_ref("status"), "approved"))),
Struct::default().with_str("status", STRING)
)]
#[case::set_single_assignment(
None,
"SET name = 'Alice' | SET full_name = name",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice"))
.set_cmd(|l| l.named_field("full_name", field_ref("name"))),
Struct::default().with_str("full_name", STRING).with_str("name", STRING)
)]
#[case::set_expression_assignment(
None,
"SET age = 30 | SET age_next_year = age + 1",
pipeline()
.set_cmd(|l| l.named_field("age", 30))
.set_cmd(|l| l.named_field("age_next_year", add(field_ref("age"), 1))),
Struct::default().with_str("age_next_year", INT).with_str("age", INT)
)]
#[case::set_multiple_assignments(
None,
"SET name = 'Alice', age = 30 | SET full_name = name, birth_year = 2024 - age",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.set_cmd(|l| l.named_field("full_name", field_ref("name")).named_field("birth_year", subtract(2024, field_ref("age")))),
Struct::default().with_str("full_name", STRING).with_str("birth_year", INT).with_str("name", STRING).with_str("age", INT)
)]
#[case::set_nested_field_assignment(
None,
"SET user.address.city = 'New York'",
pipeline().set_cmd(|l| l.named_field(ident("user").dot("address").dot("city"), "New York")),
Struct::default()
.with_str("user", Struct::default()
.with_str("address", Struct::default()
.with_str("city", STRING).into()).into())
)]
#[case::set_overwrite_field(
None,
"SET age = 30 | SET age = age + 1",
pipeline()
.set_cmd(|l| l.named_field("age", 30))
.set_cmd(|l| l.named_field("age", add(field_ref("age"), 1))),
Struct::default().with_str("age", INT)
)]
#[case::agg_simple_count(
None,
"SET region = 'US', amount = 100 | AGG total = count()",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("amount", 100))
.agg(|a| a.named_aggregate("total", call("count"))),
Struct::default().with_str("total", INT)
)]
#[case::agg_sum_with_by(
None,
"SET region = 'US', amount = 100 | AGG total_amount = sum(amount) BY region",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("amount", 100))
.agg(|a| a.named_aggregate("total_amount", call("sum").arg(field_ref("amount"))).group_by("region")),
Struct::default().with_str("region", STRING).with_str("total_amount", INT)
)]
#[case::agg_multiple_aggregations(
None,
"SET category = 'A', amount = 100, price = 50 | AGG total = sum(amount), avg_price = avg(price) BY category",
pipeline()
.set_cmd(|l| l.named_field("category", "A").named_field("amount", 100).named_field("price", 50))
.agg(|a| a.named_aggregate("total", call("sum").arg(field_ref("amount"))).named_aggregate("avg_price", call("avg").arg(field_ref("price"))).group_by("category")),
Struct::default().with_str("category", STRING).with_str("total", INT).with_str("avg_price", DOUBLE)
)]
#[case::agg_named_group_by(
None,
"SET timestamp = ts('2024-01-01T00:00:00Z') | AGG count = count() BY year = year(timestamp)",
pipeline()
.set_cmd(|l| l.named_field("timestamp", call("ts").arg("2024-01-01T00:00:00Z")))
.agg(|a| a.named_aggregate("count", call("count")).named_group("year", call("year").arg(field_ref("timestamp")))),
Struct::default().with_str("year", INT).with_str("count", INT)
)]
#[case::agg_multiple_group_by_fields(
None,
"SET region = 'US', category = 'A', amount = 100 | AGG total = sum(amount) BY region, category",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("category", "A").named_field("amount", 100))
.agg(|a| a.named_aggregate("total", call("sum").arg(field_ref("amount"))).group_by("region").group_by("category")),
Struct::default().with_str("region", STRING).with_str("category", STRING).with_str("total", INT)
)]
#[case::sort_simple(
None,
"SET name = 'Alice', age = 30 | SORT name",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.sort(|s| s.by(field_ref("name"))),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::sort_ascending(
None,
"SET name = 'Alice', age = 30 | SORT name ASC",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.sort(|s| s.asc(field_ref("name"))),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::sort_descending(
None,
"SET name = 'Alice', age = 30 | SORT age DESC",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.sort(|s| s.desc(field_ref("age"))),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::sort_multiple(
None,
"SET name = 'Alice', age = 30 | SORT name ASC, age DESC",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.sort(|s| s.asc(field_ref("name")).desc(field_ref("age"))),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::sort_by_expression(
None,
"SET age = 30 | SORT age + 1 DESC",
pipeline()
.set_cmd(|l| l.named_field("age", 30))
.sort(|s| s.desc(add(field_ref("age"), 1))),
Struct::default().with_str("age", INT)
)]
#[case::window_simple(
None,
"SET region = 'US', sales = 100 | WINDOW row_num = row_number()",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("sales", 100))
.window(|w| w.named_field("row_num", call("row_number"))),
Struct::default().with_str("row_num", INT).with_str("region", STRING).with_str("sales", INT)
)]
#[case::window_with_by(
None,
"SET region = 'US', sales = 100 | WINDOW row_num = row_number() BY region",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("sales", 100))
.window(|w|
w.named_field("row_num", call("row_number"))
.group_by("region", field_ref("region"))
),
Struct::default().with_str("row_num", INT).with_str("region", STRING).with_str("sales", INT)
)]
#[case::window_with_sort(
None,
"SET age = 30, score = 85 | WINDOW row_num = row_number() SORT age DESC",
pipeline()
.set_cmd(|l| l.named_field("age", 30).named_field("score", 85))
.window(|w|
w.named_field("row_num", call("row_number"))
.sort(sort_command().desc(field_ref("age")))
),
Struct::default().with_str("row_num", INT).with_str("age", INT).with_str("score", INT)
)]
#[case::window_with_by_and_sort(
None,
"SET region = 'US', age = 30 | WINDOW row_num = row_number() BY region SORT age DESC",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("age", 30))
.window(|w|
w.named_field("row_num", call("row_number"))
.group_by("region", field_ref("region"))
.sort(sort_command().desc(field_ref("age")))
),
Struct::default().with_str("row_num", INT).with_str("region", STRING).with_str("age", INT)
)]
#[case::window_multiple_fields(
None,
"SET category = 'A', value = 100 | WINDOW row_num = row_number(), rank = rank() BY category",
pipeline()
.set_cmd(|l| l.named_field("category", "A").named_field("value", 100))
.window(|w|
w.named_field("row_num", call("row_number"))
.named_field("rank", call("rank"))
.group_by("category", field_ref("category"))
),
Struct::default().with_str("row_num", INT).with_str("rank", INT).with_str("category", STRING).with_str("value", INT)
)]
#[case::window_by_alias_is_available_downstream(
None,
"SET region = 'US', sales = 100 | WINDOW row_num = row_number() BY account_id = region | SELECT account_id, row_num",
pipeline()
.set_cmd(|l| l.named_field("region", "US").named_field("sales", 100))
.window(|w|
w.named_field("row_num", call("row_number"))
.group_by("account_id", field_ref("region"))
)
.select(|s| s.field("account_id").field("row_num")),
Struct::default().with_str("account_id", STRING).with_str("row_num", INT)
)]
#[case::limit_small(
None,
"SET name = 'Alice', age = 30 | LIMIT 10",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.limit(10),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::limit_one(
None,
"SET name = 'Alice', age = 30 | LIMIT 1",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30))
.limit(1),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::drop_single_field(
None,
"SET name = 'Alice', age = 30, email = 'alice@example.com' | DROP name",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30).named_field("email", "alice@example.com"))
.drop(|d| d.field("name")),
Struct::default().with_str("age", INT).with_str("email", STRING)
)]
#[case::drop_multiple_fields(
None,
"SET name = 'Alice', age = 30, email = 'alice@example.com' | DROP name, age",
pipeline()
.set_cmd(|l| l.named_field("name", "Alice").named_field("age", 30).named_field("email", "alice@example.com"))
.drop(|d| d.field("name").field("age")),
Struct::default().with_str("email", STRING)
)]
#[case::from_simple_table(
Some(vec![("users", vec![("id", INT), ("name", STRING)])]),
"FROM users",
pipeline().from(|f| f.table_reference("users")),
Struct::default().with_str("id", INT).with_str("name", STRING)
)]
#[case::from_table_with_alias(
Some(vec![("users", vec![("id", INT), ("name", STRING)])]),
"FROM u = users",
pipeline().from(|f| f.table_alias("u", "users")),
Struct::default().with_str("id", INT).with_str("name", STRING).with_str("u", Struct::default().with_str("id", INT).with_str("name", STRING).into())
)]
#[case::from_multiple_tables(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("orders", vec![("order_id", INT), ("user_id", INT)])
]),
"FROM users, orders",
pipeline().from(|f| f.table_reference("users").table_reference("orders")),
Struct::default().with_str("order_id", INT).with_str("user_id", INT).with_str("name", STRING)
)]
#[case::from_multiple_tables_with_aliases(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("orders", vec![("order_id", INT), ("user_id", INT)])
]),
"FROM u = users, o = orders",
pipeline().from(|f| f.table_alias("u", "users").table_alias("o", "orders")),
Struct::default()
.with_str("order_id", INT).with_str("user_id", INT).with_str("name", STRING)
.with_str("o", Struct::default().with_str("order_id", INT).with_str("user_id", INT).into())
.with_str("u", Struct::default().with_str("user_id", INT).with_str("name", STRING).into())
)]
#[case::from_multiple_tables_one_aliased(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("orders", vec![("order_id", INT), ("amount", INT)])
]),
"FROM u = users, orders",
pipeline().from(|f| f.table_alias("u", "users").table_reference("orders")),
Struct::default()
.with_str("order_id", INT).with_str("amount", INT).with_str("user_id", INT).with_str("name", STRING)
.with_str("u", Struct::default().with_str("user_id", INT).with_str("name", STRING).into())
)]
#[case::from_multiple_tables_mixed_aliases(
Some(vec![
("users", vec![("user_id", INT)]),
("orders", vec![("order_id", INT)]),
("products", vec![("product_id", INT)])
]),
"FROM users, o = orders, products",
pipeline().from(|f| f.table_reference("users").table_alias("o", "orders").table_reference("products")),
Struct::default()
.with_str("order_id", INT)
.with_str("o", Struct::default().with_str("order_id", INT).into())
.with_str("product_id", INT)
.with_str("user_id", INT)
)]
#[case::join_simple(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("orders", vec![("order_id", INT), ("user_id", INT), ("amount", INT)])
]),
"FROM users | JOIN orders ON orders.user_id == user_id",
pipeline()
.from(|f| f.table_reference("users"))
.join("orders", eq(field(field_ref("orders"), "user_id"), field_ref("user_id"))),
Struct::default()
.with_str("user_id", INT)
.with_str("name", STRING)
.with_str("orders", Struct::default().with_str("order_id", INT).with_str("user_id", INT).with_str("amount", INT).into())
)]
#[case::join_with_alias(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("orders", vec![("order_id", INT), ("user_id", INT), ("total", INT)])
]),
"FROM users | JOIN o = orders ON o.user_id == user_id",
pipeline()
.from(|f| f.table_reference("users"))
.join_aliased("o", "orders", eq(field(field_ref("o"), "user_id"), field_ref("user_id"))),
Struct::default()
.with_str("user_id", INT)
.with_str("name", STRING)
.with_str("o", Struct::default().with_str("order_id", INT).with_str("user_id", INT).with_str("total", INT).into())
)]
#[case::join_complex_condition(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("orders", vec![("order_id", INT), ("user_id", INT), ("status", STRING)])
]),
"FROM users | JOIN o = orders ON o.user_id == user_id AND o.status == 'active'",
pipeline()
.from(|f| f.table_reference("users"))
.join_aliased(
"o",
"orders",
and(
eq(field(field_ref("o"), "user_id"), field_ref("user_id")),
eq(field(field_ref("o"), "status"), "active")
)
),
Struct::default()
.with_str("user_id", INT)
.with_str("name", STRING)
.with_str("o", Struct::default().with_str("order_id", INT).with_str("user_id", INT).with_str("status", STRING).into())
)]
#[case::join_qualified_table_name(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("simba.orders", vec![("order_id", INT), ("user_id", INT), ("amount", INT)])
]),
"FROM users | JOIN simba.orders ON orders.user_id == user_id",
pipeline()
.from(|f| f.table_reference("users"))
.join(ident("simba").dot("orders"), eq(field(field_ref("orders"), "user_id"), field_ref("user_id"))),
Struct::default()
.with_str("user_id", INT)
.with_str("name", STRING)
.with_str("orders", Struct::default().with_str("order_id", INT).with_str("user_id", INT).with_str("amount", INT).into())
)]
#[case::join_qualified_table_explicit_alias(
Some(vec![
("users", vec![("user_id", INT), ("name", STRING)]),
("simba.orders", vec![("order_id", INT), ("user_id", INT), ("total", INT)])
]),
"FROM users | JOIN o = simba.orders ON o.user_id == user_id",
pipeline()
.from(|f| f.table_reference("users"))
.join_aliased("o", ident("simba").dot("orders"), eq(field(field_ref("o"), "user_id"), field_ref("user_id"))),
Struct::default()
.with_str("user_id", INT)
.with_str("name", STRING)
.with_str("o", Struct::default().with_str("order_id", INT).with_str("user_id", INT).with_str("total", INT).into())
)]
#[case::join_qualified_three_level(
Some(vec![
("t1", vec![("id", INT), ("value", STRING)]),
("catalog.schema.table", vec![("id", INT), ("data", STRING)])
]),
"FROM t1 | JOIN catalog.schema.table ON table.id == id",
pipeline()
.from(|f| f.table_reference("t1"))
.join(ident("catalog").dot("schema").dot("table"), eq(field(field_ref("table"), "id"), field_ref("id"))),
Struct::default()
.with_str("id", INT)
.with_str("value", STRING)
.with_str("table", Struct::default().with_str("id", INT).with_str("data", STRING).into())
)]
#[case::lookup_qualified_table_name(
Some(vec![
("events", vec![("user_id", INT), ("action", STRING)]),
("simba.users", vec![("user_id", INT), ("name", STRING), ("email", STRING)])
]),
"FROM events | LOOKUP simba.users ON user_id == users.user_id",
pipeline()
.from(|f| f.table_reference("events"))
.lookup(ident("simba").dot("users"), |l| l.on(eq(field_ref("user_id"), field(field_ref("users"), "user_id")))),
Struct::default()
.with_str("user_id", INT)
.with_str("action", STRING)
.with_str("users", Struct::default().with_str("user_id", INT).with_str("name", STRING).with_str("email", STRING).into())
)]
#[case::from_where(
Some(vec![("users", vec![("name", STRING), ("age", INT), ("active", BOOLEAN)])]),
"FROM users | WHERE age > 18",
pipeline()
.from(|f| f.table_reference("users"))
.where_cmd(gt(field_ref("age"), 18)),
Struct::default().with_str("name", STRING).with_str("age", INT).with_str("active", BOOLEAN)
)]
#[case::from_select(
Some(vec![("users", vec![("name", STRING), ("age", INT), ("email", STRING)])]),
"FROM users | SELECT name, email",
pipeline()
.from(|f| f.table_reference("users"))
.select(|s| s.field("name").field("email")),
Struct::default().with_str("name", STRING).with_str("email", STRING)
)]
#[case::from_set(
Some(vec![("sales", vec![("price", INT), ("quantity", INT)])]),
"FROM sales | SET total = price * quantity",
pipeline()
.from(|f| f.table_reference("sales"))
.set_cmd(|l| l.named_field("total", multiply(field_ref("price"), field_ref("quantity")))),
Struct::default().with_str("total", INT).with_str("price", INT).with_str("quantity", INT)
)]
#[case::from_sort(
Some(vec![("users", vec![("name", STRING), ("age", INT)])]),
"FROM users | SORT age DESC",
pipeline()
.from(|f| f.table_reference("users"))
.sort(|s| s.desc(field_ref("age"))),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::from_limit(
Some(vec![("events", vec![("timestamp", Type::Timestamp), ("severity", INT)])]),
"FROM events | LIMIT 10",
pipeline()
.from(|f| f.table_reference("events"))
.limit(10),
Struct::default().with_str("timestamp", Type::Timestamp).with_str("severity", INT)
)]
#[case::from_where_select(
Some(vec![("users", vec![("name", STRING), ("age", INT), ("email", STRING)])]),
"FROM users | WHERE age > 18 | SELECT name, email",
pipeline()
.from(|f| f.table_reference("users"))
.where_cmd(gt(field_ref("age"), 18))
.select(|s| s.field("name").field("email")),
Struct::default().with_str("name", STRING).with_str("email", STRING)
)]
#[case::from_set_select(
Some(vec![("sales", vec![("product", STRING), ("price", INT), ("quantity", INT)])]),
"FROM sales | SET total = price * quantity | SELECT product, total",
pipeline()
.from(|f| f.table_reference("sales"))
.set_cmd(|l| l.named_field("total", multiply(field_ref("price"), field_ref("quantity"))))
.select(|s| s.field("product").field("total")),
Struct::default().with_str("product", STRING).with_str("total", INT)
)]
#[case::from_where_agg(
Some(vec![("sales", vec![("year", INT), ("region", STRING), ("amount", INT)])]),
"FROM sales | WHERE year == 2024 | AGG total = sum(amount) BY region",
pipeline()
.from(|f| f.table_reference("sales"))
.where_cmd(eq(field_ref("year"), 2024))
.agg(|a| a.named_aggregate("total", call("sum").arg(field_ref("amount"))).group_by("region")),
Struct::default().with_str("region", STRING).with_str("total", INT)
)]
#[case::from_agg_sort(
Some(vec![("sales", vec![("region", STRING), ("amount", INT)])]),
"FROM sales | AGG total = sum(amount) BY region | SORT total DESC",
pipeline()
.from(|f| f.table_reference("sales"))
.agg(|a| a.named_aggregate("total", call("sum").arg(field_ref("amount"))).group_by("region"))
.sort(|s| s.desc(field_ref("total"))),
Struct::default().with_str("region", STRING).with_str("total", INT)
)]
#[case::from_select_limit(
Some(vec![("users", vec![("name", STRING), ("age", INT), ("score", INT)])]),
"FROM users | SELECT name, score | LIMIT 5",
pipeline()
.from(|f| f.table_reference("users"))
.select(|s| s.field("name").field("score"))
.limit(5),
Struct::default().with_str("name", STRING).with_str("score", INT)
)]
#[case::four_commands(
Some(vec![("users", vec![("name", STRING), ("age", INT), ("active", BOOLEAN)])]),
"FROM users | WHERE active == true | SELECT name, age | SORT age DESC",
pipeline()
.from(|f| f.table_reference("users"))
.where_cmd(eq(field_ref("active"), true))
.select(|s| s.field("name").field("age"))
.sort(|s| s.desc(field_ref("age"))),
Struct::default().with_str("name", STRING).with_str("age", INT)
)]
#[case::five_commands(
Some(vec![("users", vec![("name", STRING), ("age", INT), ("active", BOOLEAN)])]),
"FROM users | WHERE active == true | SET age_group = age / 10 | AGG count = count() BY age_group | SORT age_group ASC",
pipeline()
.from(|f| f.table_reference("users"))
.where_cmd(eq(field_ref("active"), true))
.set_cmd(|l| l.named_field("age_group", divide(field_ref("age"), 10)))
.agg(|a| a.named_aggregate("count", call("count")).group_by("age_group"))
.sort(|s| s.asc(field_ref("age_group"))),
Struct::default().with_str("age_group", INT).with_str("count", INT)
)]
#[case::six_commands(
Some(vec![("sales", vec![("year", INT), ("region", STRING), ("price", INT), ("quantity", INT)])]),
"FROM sales | WHERE year == 2024 | SET total = price * quantity | AGG revenue = sum(total) BY region | SORT revenue DESC | LIMIT 10",
pipeline()
.from(|f| f.table_reference("sales"))
.where_cmd(eq(field_ref("year"), 2024))
.set_cmd(|l| l.named_field("total", multiply(field_ref("price"), field_ref("quantity"))))
.agg(|a| a.named_aggregate("revenue", call("sum").arg(field_ref("total"))).group_by("region"))
.sort(|s| s.desc(field_ref("revenue")))
.limit(10),
Struct::default().with_str("region", STRING).with_str("revenue", INT)
)]
#[case::full_pipeline_transformation(
Some(vec![("events", vec![("timestamp", Type::Timestamp), ("event_type", STRING), ("severity", INT)])]),
r#"
FROM events
| WHERE event_type == 'admin'
| SELECT event_type, timestamp
| SORT timestamp DESC
| LIMIT 1
"#,
pipeline()
.from(|f| f.table_reference("events"))
.where_cmd(eq(field_ref("event_type"), "admin"))
.select(|s| s.field("event_type").field("timestamp"))
.sort(|s| s.desc(field_ref("timestamp")))
.limit(1),
Struct::default().with_str("event_type", STRING).with_str("timestamp", Type::Timestamp)
)]
#[case::nested_field_creation(
None,
"SET x = 'foo', y.a = 1, y.b = 2 | SET y.a = 3",
pipeline()
.set_cmd(|l| {
l.named_field("x", "foo")
.named_field(ident("y").dot("a"), 1)
.named_field(ident("y").dot("b"), 2)
})
.set_cmd(|l| l.named_field(ident("y").dot("a"), 3)),
Struct::default()
.with_str("y", Struct::default().with_str("a", INT).with_str("b", INT).into())
.with_str("x", STRING)
)]
#[case::array_operations(
None,
"SET arr = [10, 20, 30] | SET idx = -1 | SELECT result = arr[idx]",
pipeline()
.set_cmd(|l| l.named_field("arr", array().element(10).element(20).element(30)))
.set_cmd(|l| l.named_field("idx", negate(1)))
.select(|s| s.named_field("result", index(field_ref("arr"), field_ref("idx")))),
Struct::default().with_str("result", INT)
)]
#[case::sysmon_events_query(
Some(vec![
("simba.sysmon_events", vec![
("timestamp", Type::Timestamp),
("event", Struct::default()
.with_str("code", STRING)
.with_str("action", STRING).into()),
("host", Struct::default()
.with_str("name", STRING)
.with_str("hostname", STRING).into()),
])
]),
"FROM simba.sysmon_events | WHERE event.code IN ['8', '10'] | SET timestamp_m = timestamp@m | SET computer_name = host.name AS string | SET current_hour = hour(timestamp) AS int",
pipeline()
.from(|f| f.table_reference(ident("simba").dot("sysmon_events")))
.where_cmd(in_(field(field_ref("event"), "code"), array().element("8").element("10")))
.set_cmd(|l| l.named_field("timestamp_m", at_minute(field_ref("timestamp"))))
.set_cmd(|l| l.named_field("computer_name", cast(field(field_ref("host"), "name"), Type::String)))
.set_cmd(|l| l.named_field("current_hour", cast(call("hour").arg(field_ref("timestamp")), Type::Int))),
Struct::default()
.with_str("current_hour", INT)
.with_str("computer_name", STRING)
.with_str("timestamp_m", Type::Timestamp)
.with_str("timestamp", Type::Timestamp)
.with_str("event", Struct::default().with_str("code", STRING).with_str("action", STRING).into())
.with_str("host", Struct::default().with_str("name", STRING).with_str("hostname", STRING).into())
)]
#[case::timestamp_extraction_year(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT yr = year(ts)",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("yr", call("year").arg(field_ref("ts")))),
Struct::default().with_str("yr", INT)
)]
#[case::timestamp_extraction_month(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT mon = month(ts)",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("mon", call("month").arg(field_ref("ts")))),
Struct::default().with_str("mon", INT)
)]
#[case::timestamp_extraction_day(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT d = day(ts)",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("d", call("day").arg(field_ref("ts")))),
Struct::default().with_str("d", INT)
)]
#[case::timestamp_extraction_hour(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT h = hour(ts)",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("h", call("hour").arg(field_ref("ts")))),
Struct::default().with_str("h", INT)
)]
#[case::timestamp_extraction_minute(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT m = minute(ts)",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("m", call("minute").arg(field_ref("ts")))),
Struct::default().with_str("m", INT)
)]
#[case::timestamp_extraction_second(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT s = second(ts)",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("s", call("second").arg(field_ref("ts")))),
Struct::default().with_str("s", INT)
)]
#[case::timestamp_comparison(
None,
"SET ts1 = '2024-01-15T14:30:45Z' AS timestamp, ts2 = '2024-01-16T14:30:45Z' AS timestamp | WHERE ts1 < ts2",
pipeline()
.set_cmd(|l| l
.named_field("ts1", cast("2024-01-15T14:30:45Z", Type::Timestamp))
.named_field("ts2", cast("2024-01-16T14:30:45Z", Type::Timestamp)))
.where_cmd(lt(field_ref("ts1"), field_ref("ts2"))),
Struct::default().with_str("ts1", Type::Timestamp).with_str("ts2", Type::Timestamp)
)]
#[case::timestamp_null_handling(
None,
"SET ts = null, valid_ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT result = ts IS NULL",
pipeline()
.set_cmd(|l| l
.named_field("ts", null())
.named_field("valid_ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("result", is_null(field_ref("ts")))),
Struct::default().with_str("result", BOOLEAN)
)]
#[case::timestamp_cast_to_string(
None,
"SET ts = '2024-01-15T14:30:45Z' AS timestamp | SELECT ts_str = ts AS string",
pipeline()
.set_cmd(|l| l.named_field("ts", cast("2024-01-15T14:30:45Z", Type::Timestamp)))
.select(|s| s.named_field("ts_str", cast(field_ref("ts"), Type::String))),
Struct::default().with_str("ts_str", STRING)
)]
#[case::map_with_different_array_values(
None,
"SET x = map('a': [1], 'b': [])",
pipeline()
.set_cmd(|l| l.named_field("x", call("map").arg(pair("a", array().element(1))).arg(pair("b", array())))),
Struct::default().with_str("x", Map::new(STRING, Array::new(INT).into()).into())
)]
#[case::map_from_tuple_array(
None,
"SET x = map([('a', [1]), ('b', [])])",
pipeline()
.set_cmd(|l| l.named_field("x", call("map").arg(
array()
.element(tuple().element("a").element(array().element(1)))
.element(tuple().element("b").element(array()))
))),
Struct::default().with_str("x", Map::new(STRING, Array::new(INT).into()).into())
)]
#[case::array_of_tuples_with_different_arrays(
None,
"SET x = [('a', [1]),('b', [])]",
pipeline()
.set_cmd(|l| l.named_field("x",
array()
.element(tuple().element("a").element(array().element(1)))
.element(tuple().element("b").element(array()))
)),
Struct::default().with_str("x", Array::new(Tuple::new(vec![STRING, Array::new(INT).into()]).into()).into())
)]
#[case::coalesce_tuples_with_arrays(
None,
"SET x = coalesce(('a', [1]), ('b', []))",
pipeline()
.set_cmd(|l| l.named_field("x", call("coalesce")
.arg(tuple().element("a").element(array().element(1)))
.arg(tuple().element("b").element(array()))
)),
Struct::default().with_str("x", Tuple::new(vec![STRING, Array::new(INT).into()]).into())
)]
#[case::if_expression_empty_vs_nonempty_arrays(
None,
"SET x = if(true, [], [1])",
pipeline()
.set_cmd(|l| l.named_field("x", call("if")
.arg(true)
.arg(array())
.arg(array().element(1))
)),
Struct::default().with_str("x", Array::new(INT).into())
)]
#[case::map_with_merging_struct_arrays(
None,
"SET x = map('a': [{b: 1}], 'b': [{a: 2}])",
pipeline()
.set_cmd(|l| l.named_field("x", call("map")
.arg(pair("a", array().element(struct_literal().field("b", 1))))
.arg(pair("b", array().element(struct_literal().field("a", 2))))
)),
Struct::default().with_str("x", {
let mut struct_fields = OrderMap::new();
struct_fields.insert("a".parse::<SimpleIdentifier>().unwrap(), INT);
struct_fields.insert("b".parse::<SimpleIdentifier>().unwrap(), INT);
Map::new(STRING, Array::new(Struct::new(struct_fields).into()).into()).into()
})
)]
#[case::map_with_empty_and_struct_arrays(
None,
"SET x = map('a': [], 'b': [{a: 2}])",
pipeline()
.set_cmd(|l| l.named_field("x", call("map")
.arg(pair("a", array()))
.arg(pair("b", array().element(struct_literal().field("a", 2))))
)),
Struct::default().with_str("x", {
let mut struct_fields = OrderMap::new();
struct_fields.insert("a".parse::<SimpleIdentifier>().unwrap(), INT);
Map::new(STRING, Array::new(Struct::new(struct_fields).into()).into()).into()
})
)]
#[case::match_single_pattern_within_valid(
Some(vec![("events", vec![("timestamp", Type::Timestamp), ("event_type", STRING)])]),
"MATCH e=events WITHIN 5m",
pipeline()
.match_cmd(|m| m
.named_pattern("e", "events")
.within(IntervalLiteralBuilder::new(5, IntervalUnit::Minute))
),
Struct::default()
.with_str("timestamp", Type::Timestamp)
.with_str("event_type", STRING)
.with_str("e", Struct::default()
.with_str("timestamp", Type::Timestamp)
.with_str("event_type", STRING)
.into())
)]
#[case::match_single_pattern_within_sort_timestamp(
Some(vec![("events", vec![("timestamp", Type::Timestamp), ("event_type", STRING)])]),
"MATCH e=events SORT timestamp WITHIN 5m",
pipeline()
.match_cmd(|m| m
.named_pattern("e", "events")
.within(minutes(5))
.sort(sort_command().by(field_ref("timestamp")))
),
Struct::default()
.with_str("timestamp", Type::Timestamp)
.with_str("event_type", STRING)
.with_str("e", Struct::default()
.with_str("timestamp", Type::Timestamp)
.with_str("event_type", STRING)
.into())
)]
fn test_pipelines(
#[case] tables: Option<Vec<(&str, Vec<(&str, Type)>)>>,
#[case] query: &str,
#[case] expected_ast: PipelineBuilder,
#[case] expected_schema: Struct,
) -> Result<(), TranslationErrors> {
let pipeline = Pipeline::parse_result(query)?;
assert_eq!(pipeline, expected_ast.build());
let typed = if let Some(tables) = tables {
let catalog = catalog_with_tables(tables);
type_check_with_provider(pipeline, catalog).into_result()?
} else {
type_check(pipeline).into_result()?
};
let final_schema = match typed.kind {
Ok(ValidPipeline { final_schema, .. }) => final_schema,
Err(e) => panic!("Type checking failed: {}", e),
};
assert_eq!(*final_schema.as_struct(), expected_schema);
Ok(())
}
#[rstest]
#[case::map_with_incompatible_primitive_types(
None,
"SET x = map('a': [1], 'b': ['hello'])",
"The types are not mergeable"
)]
#[case::map_with_incompatible_struct_fields(
None,
"SET x = map('a': [{a: 2}], 'b': [{a: 'hello'}])",
"could not find a matching function definition"
)]
#[case::match_single_pattern_negative_within(
Some(vec![("events", vec![("timestamp", Type::Timestamp), ("event_type", STRING)])]),
"MATCH events WITHIN -5m",
"WITHIN interval cannot be negative"
)]
#[case::match_multiple_patterns_negative_within(
Some(vec![
("alerts", vec![("timestamp", Type::Timestamp), ("severity", STRING)]),
("events", vec![("timestamp", Type::Timestamp), ("event_type", STRING)]),
("signals", vec![("timestamp", Type::Timestamp), ("signal_type", STRING)])
]),
"MATCH alerts+ events? signals WITHIN -5m",
"WITHIN interval cannot be negative"
)]
#[case::from_not_first(
None,
"SET x = 1 | FROM events",
"FROM must be the first command in a pipeline"
)]
#[case::union_not_first(
None,
"SET x = 1 | UNION events",
"UNION must be the first command in a pipeline"
)]
#[case::rows_not_first(
None,
"SET x = 1 | ROWS [1, 2, 3]",
"ROWS must be the first command in a pipeline"
)]
#[case::match_not_first(
Some(vec![("events", vec![("timestamp", Type::Timestamp), ("event_type", STRING)])]),
"FROM events | MATCH e=events WITHIN 5m",
"MATCH must be the first command in a pipeline"
)]
#[case::from_after_multiple_commands(
None,
"SET x = 1 | WHERE x > 0 | FROM events",
"FROM must be the first command in a pipeline"
)]
fn test_pipeline_type_errors(
#[case] tables: Option<Vec<(&str, Vec<(&str, Type)>)>>,
#[case] query: &str,
#[case] expected_error_fragment: &str,
) {
let pipeline = Pipeline::parse_result(query).expect("Query should parse successfully");
let result = if let Some(tables) = tables {
let catalog = catalog_with_tables(tables);
type_check_with_provider(pipeline, catalog).into_result()
} else {
type_check(pipeline).into_result()
};
assert!(
result.is_err(),
"Expected type error but query succeeded: {}",
query
);
let err = result.unwrap_err();
let err_string = err.to_string();
assert!(
err_string.contains(expected_error_fragment),
"Expected error message to contain '{}', but got: {}",
expected_error_fragment,
err_string
);
}