use std::sync::Arc;
use crate::err::TranslationErrors;
use crate::tree::ast::dataset_identifier::QualifiedDatasetIdentifier;
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::query::Query;
use crate::tree::ast::ParseWithErrors;
use crate::tree::builder::*;
use crate::tree::tests::shared::*;
use crate::types::struct_type::Struct;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, STRING, TIMESTAMP};
use pretty_assertions::assert_eq;
use rstest::rstest;
#[rstest]
#[case::simple_query(
vec![("users", vec![("id", INT), ("name", STRING)])],
"FROM users",
query().main(pipeline().from(|f| f.table_reference("users"))).build(),
Struct::default().with_str("id", INT).with_str("name", STRING)
)]
#[case::with_single_def(
vec![("users", vec![("id", INT), ("name", STRING), ("active", BOOLEAN)])],
"DEF active_users = FROM users | WHERE active == true;\nFROM active_users",
query().main(pipeline().from(|f| f.table_reference("active_users")))
.def_pipeline("active_users", pipeline().from(|f| f.table_reference("users")).where_cmd(eq(field_ref("active"), true)))
.build(),
Struct::default().with_str("id", INT).with_str("name", STRING).with_str("active", BOOLEAN)
)]
#[case::with_multiple_defs(
vec![("users", vec![("id", INT), ("name", STRING), ("active", BOOLEAN), ("value", INT)])],
"DEF active_users = FROM users | WHERE active == true;\nDEF high_value = FROM active_users | WHERE value > 1000;\nFROM high_value",
query().main(pipeline().from(|f| f.table_reference("high_value")))
.def_pipeline("active_users", pipeline().from(|f| f.table_reference("users")).where_cmd(eq(field_ref("active"), true)))
.def_pipeline("high_value", pipeline().from(|f| f.table_reference("active_users")).where_cmd(gt(field_ref("value"), 1000)))
.build(),
Struct::default().with_str("id", INT).with_str("name", STRING).with_str("active", BOOLEAN).with_str("value", INT)
)]
#[case::with_complex_pipeline(
vec![("sales", vec![("id", INT), ("year", INT), ("amount", INT), ("region", STRING)])],
"DEF sales_2024 = FROM sales | WHERE year == 2024;\nFROM sales_2024 | AGG total = sum(amount) BY region",
query().main(pipeline()
.from(|f| f.table_reference("sales_2024"))
.agg(|a| a.named_aggregate("total", call("sum").arg(field_ref("amount"))).group_by("region")))
.def_pipeline("sales_2024", pipeline().from(|f| f.table_reference("sales")).where_cmd(eq(field_ref("year"), 2024)))
.build(),
Struct::default().with_str("region", STRING).with_str("total", INT)
)]
#[case::with_def_referencing_def(
vec![("users", vec![("id", INT), ("name", STRING), ("active", BOOLEAN), ("verified", BOOLEAN)])],
"DEF step1 = FROM users | WHERE active == true;\nDEF step2 = FROM step1 | WHERE verified == true;\nFROM step2",
query().main(pipeline().from(|f| f.table_reference("step2")))
.def_pipeline("step1", pipeline().from(|f| f.table_reference("users")).where_cmd(eq(field_ref("active"), true)))
.def_pipeline("step2", pipeline().from(|f| f.table_reference("step1")).where_cmd(eq(field_ref("verified"), true)))
.build(),
Struct::default().with_str("id", INT).with_str("name", STRING).with_str("active", BOOLEAN).with_str("verified", BOOLEAN)
)]
fn test_queries_with_ctes(
#[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
#[case] input: &str,
#[case] expected_ast: Query,
#[case] expected_output_fields: Struct,
) -> Result<(), TranslationErrors> {
let (query, errors) = Query::parse_with_errors(input);
if !errors.is_empty() {
eprintln!("Parse errors:\n{}", errors);
}
assert!(errors.is_empty(), "Expected no parse errors");
assert_eq!(query.to_string(), expected_ast.to_string());
let catalog = catalog_with_tables(tables);
let typed = type_check_with_catalog::<Query>(Arc::new(query), catalog).into_result()?;
let final_schema = match typed.pipeline.kind.as_ref() {
Ok(valid_pipeline) => valid_pipeline.final_schema.clone(),
Err(e) => panic!("Type checking failed: {}", e),
};
assert_eq!(*final_schema.as_struct(), expected_output_fields);
Ok(())
}
#[rstest]
#[case::lookup_after_agg_and_set(
vec![
("events", vec![("timestamp", TIMESTAMP), ("user", STRING), ("action", STRING)]),
("users", vec![("user", STRING), ("email", STRING)])
],
"FROM events
| AGG total = count() BY user
| SET is_active = total > 10
| LOOKUP u = users ON user == u.user",
Struct::default()
.with_str("is_active", BOOLEAN)
.with_str("user", STRING)
.with_str("total", INT)
.with_str("u", Struct::default().with_str("user", STRING).with_str("email", STRING).into())
)]
#[case::lookup_preserves_group_by_through_multiple_sets(
vec![
("sales", vec![("product", STRING), ("region", STRING), ("amount", INT)]),
("products", vec![("product", STRING), ("category", STRING)])
],
"FROM sales
| AGG total_sales = sum(amount) BY product, region
| SET high_value = total_sales > 1000
| SET status = if(high_value, 'hot', 'cold')
| LOOKUP p = products ON product == p.product",
Struct::default()
.with_str("status", STRING)
.with_str("high_value", BOOLEAN)
.with_str("product", STRING)
.with_str("region", STRING)
.with_str("total_sales", INT)
.with_str("p", Struct::default().with_str("product", STRING).with_str("category", STRING).into())
)]
#[case::lookup_preserves_aggregate_fields(
vec![
("orders", vec![("customer_id", INT), ("order_date", TIMESTAMP), ("amount", INT)]),
("customers", vec![("customer_id", INT), ("name", STRING), ("tier", STRING)])
],
"FROM orders
| AGG total_orders = count(), total_amount = sum(amount), avg_amount = avg(amount) BY customer_id
| LOOKUP c = customers ON customer_id == c.customer_id",
Struct::default()
.with_str("customer_id", INT)
.with_str("total_orders", INT)
.with_str("total_amount", INT)
.with_str("avg_amount", DOUBLE)
.with_str("c", Struct::default().with_str("customer_id", INT).with_str("name", STRING).with_str("tier", STRING).into())
)]
#[case::join_after_multiple_sets(
vec![
("events", vec![("id", INT), ("user_id", INT), ("timestamp", TIMESTAMP)]),
("users", vec![("user_id", INT), ("name", STRING)])
],
"FROM events
| SET day = 1
| SET is_recent = day < 7
| JOIN u = users ON user_id == u.user_id",
Struct::default()
.with_str("is_recent", BOOLEAN)
.with_str("day", INT)
.with_str("id", INT)
.with_str("user_id", INT)
.with_str("timestamp", TIMESTAMP)
.with_str("u", Struct::default().with_str("user_id", INT).with_str("name", STRING).into())
)]
#[case::join_after_agg_preserves_all_fields(
vec![
("logs", vec![("level", STRING), ("message", STRING)]),
("levels", vec![("level", STRING), ("severity", INT)])
],
"FROM logs
| AGG count = count() BY level
| JOIN l = levels ON level == l.level",
Struct::default()
.with_str("level", STRING)
.with_str("count", INT)
.with_str("l", Struct::default().with_str("level", STRING).with_str("severity", INT).into())
)]
#[case::lookup_then_where_references_agg_fields(
vec![
("metrics", vec![("host", STRING), ("value", INT)]),
("hosts", vec![("host", STRING), ("max_value", INT)])
],
"FROM metrics
| AGG total = sum(value), avg = avg(value) BY host
| LOOKUP h = hosts ON host == h.host
| WHERE total > h.max_value",
Struct::default()
.with_str("host", STRING)
.with_str("total", INT)
.with_str("avg", DOUBLE)
.with_str("h", Struct::default().with_str("host", STRING).with_str("max_value", INT).into())
)]
#[case::multiple_lookups_preserve_all_fields(
vec![
("orders", vec![("order_id", INT), ("customer_id", INT), ("product_id", INT)]),
("customers", vec![("customer_id", INT), ("name", STRING)]),
("products", vec![("product_id", INT), ("category", STRING)])
],
"FROM orders
| LOOKUP c = customers ON customer_id == c.customer_id
| LOOKUP p = products ON product_id == p.product_id",
Struct::default()
.with_str("order_id", INT)
.with_str("customer_id", INT)
.with_str("product_id", INT)
.with_str("c", Struct::default().with_str("customer_id", INT).with_str("name", STRING).into())
.with_str("p", Struct::default().with_str("product_id", INT).with_str("category", STRING).into())
)]
#[case::lookup_after_agg_many_group_by_fields(
vec![
("sales", vec![("year", INT), ("quarter", INT), ("region", STRING), ("product", STRING), ("revenue", INT)]),
("targets", vec![("year", INT), ("quarter", INT), ("region", STRING), ("target", INT)])
],
"FROM sales
| AGG total_revenue = sum(revenue) BY year, quarter, region, product
| LOOKUP t = targets
ON year == t.year
AND quarter == t.quarter
AND region == t.region",
Struct::default()
.with_str("year", INT)
.with_str("quarter", INT)
.with_str("region", STRING)
.with_str("product", STRING)
.with_str("total_revenue", INT)
.with_str("t", Struct::default().with_str("year", INT).with_str("quarter", INT).with_str("region", STRING).with_str("target", INT).into())
)]
#[case::join_with_complex_on_clause_from_chain(
vec![
("events", vec![("user_id", INT), ("timestamp", TIMESTAMP), ("score", INT)]),
("users", vec![("user_id", INT), ("email", STRING)])
],
"FROM events
| SET is_high_score = score > 100
| AGG total_score = sum(score), event_count = count() BY user_id, is_high_score
| JOIN u = users ON user_id == u.user_id",
Struct::default()
.with_str("user_id", INT)
.with_str("is_high_score", BOOLEAN)
.with_str("total_score", INT)
.with_str("event_count", INT)
.with_str("u", Struct::default().with_str("user_id", INT).with_str("email", STRING).into())
)]
fn test_join_lookup_with_environment_chains(
#[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
#[case] input: &str,
#[case] expected_output_fields: Struct,
) -> Result<(), TranslationErrors> {
let (query, errors) = Query::parse_with_errors(input);
if !errors.is_empty() {
eprintln!("Parse errors:\n{}", errors);
}
assert!(errors.is_empty(), "Expected no parse errors");
let catalog = catalog_with_tables(tables);
let typed = type_check_with_catalog(query, catalog).into_result()?;
let final_schema = match typed.pipeline.kind.as_ref() {
Ok(valid_pipeline) => valid_pipeline.final_schema.clone(),
Err(e) => panic!("Type checking failed: {}", e),
};
assert_eq!(*final_schema.as_struct(), expected_output_fields);
Ok(())
}
#[rstest]
#[case::simple_from(
"FROM users",
vec!["users"]
)]
#[case::from_with_schema(
"FROM schema:users",
vec!["schema:users"]
)]
#[case::multiple_from_tables(
"FROM users, orders",
vec!["users", "orders"]
)]
#[case::union_tables(
"FROM users | UNION orders",
vec!["users", "orders"]
)]
#[case::join_extracts_both_tables(
"FROM users | JOIN o = orders ON id == o.id",
vec!["users", "orders"]
)]
#[case::lookup_extracts_both_tables(
"FROM users | LOOKUP o = orders ON id == o.id",
vec!["users", "orders"]
)]
#[case::with_clause_excludes_cte(
"DEF active = FROM users | WHERE active == true;\nFROM active",
vec!["users"]
)]
#[case::with_clause_multiple_ctes(
"DEF a = FROM users;\nDEF b = FROM a | JOIN o = orders ON id == o.id;\nFROM b",
vec!["users", "orders"]
)]
#[case::cte_in_join_excluded(
"DEF temp = FROM users;\nFROM orders | JOIN t = temp ON id == t.id",
vec!["users", "orders"]
)]
#[case::complex_query_all_sources(
"DEF enriched = FROM events | JOIN u = users ON user_id == u.id;\nFROM enriched | LOOKUP p = products ON product_id == p.id",
vec!["events", "users", "products"]
)]
#[case::append_extracts_target(
"FROM users | APPEND orders",
vec!["users", "orders"]
)]
#[case::append_cte_excluded(
"DEF temp = FROM users;\nFROM orders | APPEND temp",
vec!["users", "orders"]
)]
#[case::match_extracts_pattern_tables(
"MATCH a=logins b=logouts",
vec!["logins", "logouts"]
)]
#[case::match_nested_patterns(
"MATCH (a=logins b=logouts)+",
vec!["logins", "logouts"]
)]
#[case::match_cte_excluded(
"DEF temp = FROM users;\nMATCH a=temp b=logins",
vec!["users", "logins"]
)]
#[case::scalar_def_name_does_not_exclude_same_named_table(
"DEF x = 42;\nFROM x",
vec!["x"]
)]
fn test_query_datasets(#[case] input: &str, #[case] expected_tables: Vec<&str>) {
let (query, errors) = Query::parse_with_errors(input);
if !errors.is_empty() {
eprintln!("Parse errors:\n{}", errors);
}
assert!(errors.is_empty(), "Expected no parse errors");
let default_space = SimpleIdentifier::new("test");
let mut table_names: Vec<String> = query
.datasets(Some(&default_space))
.expect("datasets() should succeed")
.into_iter()
.map(|id| id.to_string())
.collect();
table_names.sort_unstable();
let mut expected: Vec<String> = expected_tables
.iter()
.map(|s| {
if s.contains(':') {
(*s).to_string()
} else {
format!("test:{s}")
}
})
.collect();
expected.sort_unstable();
assert_eq!(table_names, expected, "Mismatch for query: {}", input);
}
#[test]
fn test_query_datasets_by_command_groups_command_kinds() {
let input = "DEF enriched = FROM raw.events | JOIN dimensions.users ON true;\
DEF combined = UNION sources.first, sources.second;\
FROM enriched, combined, shared.dataset\
| JOIN shared.dataset ON true\
| LOOKUP shared.dataset ON true\
| LOOKUP dimensions.baselines ON true";
let (query, errors) = Query::parse_with_errors(input);
if !errors.is_empty() {
eprintln!("Parse errors:\n{}", errors);
}
assert!(errors.is_empty(), "Expected no parse errors");
let default_space = SimpleIdentifier::new("default");
let references = query
.datasets_by_command(Some(&default_space))
.expect("datasets_by_command() should succeed");
let names = |datasets: Vec<QualifiedDatasetIdentifier>| {
datasets
.into_iter()
.map(|dataset| dataset.to_string())
.collect::<Vec<_>>()
};
assert_eq!(
names(references.r#from),
vec!["default:shared.dataset", "default:raw.events"]
);
assert_eq!(
names(references.union),
vec!["default:sources.first", "default:sources.second"]
);
assert_eq!(
names(references.join),
vec!["default:shared.dataset", "default:dimensions.users"]
);
assert_eq!(
names(references.lookup),
vec!["default:shared.dataset", "default:dimensions.baselines"]
);
}