use std::sync::Arc;
use crate::err::TranslationErrors;
use crate::tree::ast::expression::Expression;
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::ast::query::Query;
use crate::tree::ast::ParseWithErrors;
use crate::tree::options::TypeCheckOptions;
use crate::tree::tests::shared::{catalog_with_tables, type_check_with_catalog};
use crate::tree::typed_ast::command::TypedCommandKind;
use crate::type_check_with_options;
use crate::types::array::Array;
use crate::types::decimal_type::Decimal;
use crate::types::map::Map;
use crate::types::struct_type::Struct;
use crate::types::{Type, DOUBLE, INT, STRING, TIMESTAMP};
use pretty_assertions::assert_eq;
use rstest::rstest;
fn sysmon_events_schema() -> Vec<(&'static str, Vec<(&'static str, Type)>)> {
let sysmon_event_struct = Struct::default()
.with_str("code", STRING)
.with_str("original", STRING)
.into();
let winlog_struct = Struct::default()
.with_str("event_data", Map::new(STRING, STRING).into())
.into();
let host_struct = Struct::default().with_str("name", STRING).into();
vec![(
"simba:sysmon_events",
vec![
("timestamp", TIMESTAMP),
("event", sysmon_event_struct),
("winlog", winlog_struct),
("host", host_struct),
],
)]
}
#[rstest]
#[case::drop_preserves_set_fields(
sysmon_events_schema(),
r#"
FROM simba:sysmon_events
| SET reg_key = coalesce(winlog.event_data["TargetObject"], '')
| DROP winlog.event_data
| WHERE regexp_like(lower(reg_key), '(?i).*test.*')
"#,
Struct::default()
.with_str("reg_key", STRING)
.with_str("timestamp", TIMESTAMP)
.with_str("event", Struct::default().with_str("code", STRING).with_str("original", STRING).into())
.with_str("host", Struct::default().with_str("name", STRING).into())
)]
#[case::from_cte_flattens_environment(
sysmon_events_schema(),
r#"
DEF registry_events = FROM simba:sysmon_events
| WHERE event.code == "13"
| SET reg_key = coalesce(winlog.event_data["TargetObject"], '')
| SET host = host.name
| DROP winlog.event_data, event.original;
FROM x = registry_events
| SELECT timestamp, host, reg_key
"#,
Struct::default()
.with_str("timestamp", TIMESTAMP)
.with_str("host", STRING)
.with_str("reg_key", STRING)
)]
#[case::from_aliased_ctes_preserves_all_fields(
sysmon_events_schema(),
r#"
DEF registry_events = FROM simba:sysmon_events
| WHERE event.code == "13"
| SET reg_key = coalesce(winlog.event_data["TargetObject"], '')
| SET host = host.name
| DROP winlog.event_data, event.original;
FROM x = registry_events, y = registry_events
| SELECT x.timestamp, y.timestamp, x.host, y.host
"#,
Struct::default()
.with_str("x", Struct::default().with_str("timestamp", Type::Timestamp).with_str("host", STRING).into())
.with_str("y", Struct::default().with_str("timestamp", Type::Timestamp).with_str("host", STRING).into())
)]
#[case::window_finds_timestamp_after_cte_from(
sysmon_events_schema(),
r#"
DEF registry_events = FROM simba:sysmon_events
| WHERE event.code == "13"
| SET host = host.name
| DROP winlog.event_data;
DEF process_events = FROM simba:sysmon_events
| WHERE event.code == "1"
| SET host = host.name
| DROP winlog.event_data;
FROM registry_event = registry_events, process_event = process_events
| WINDOW process_event = last_value(process_event),
registry_event = last_value(registry_event)
BY host WITHIN 1s
"#,
Struct::default()
.with_str("process_event", Struct::default()
.with_str("host", STRING)
.with_str("timestamp", TIMESTAMP)
.with_str("event", Struct::default()
.with_str("code", STRING)
.with_str("original", STRING)
.into())
.into())
.with_str("registry_event", Struct::default()
.with_str("host", STRING)
.with_str("timestamp", TIMESTAMP)
.with_str("event", Struct::default()
.with_str("code", STRING)
.with_str("original", STRING)
.into())
.into())
.with_str("host", STRING)
.with_str("timestamp", TIMESTAMP)
.with_str("event", Struct::default().with_str("code", STRING).with_str("original", STRING).into())
)]
#[case::join_preserves_base_chain(
sysmon_events_schema(),
r#"
DEF enriched = FROM simba:sysmon_events
| WHERE event.code == "1"
| SET hostname = host.name
| SET process_name = winlog.event_data["Image"]
| DROP winlog.event_data;
FROM simba:sysmon_events
| WHERE event.code == "3"
| JOIN x = enriched ON host.name == x.hostname
| SELECT timestamp, x.process_name, x.hostname
"#,
Struct::default()
.with_str("timestamp", TIMESTAMP)
.with_str("x", Struct::default().with_str("process_name", STRING).with_str("hostname", STRING).into())
)]
#[case::unnest_with_set_fields(
sysmon_events_schema(),
r#"
FROM simba:sysmon_events
| WHERE event.code == "1"
| SET hostname = host.name
| SET process_info = {pid: winlog.event_data["ProcessId"], name: winlog.event_data["Image"]}
| UNNEST process_info
| SELECT hostname, pid, name
"#,
Struct::default()
.with_str("hostname", STRING)
.with_str("pid", STRING)
.with_str("name", STRING)
)]
#[case::stddev_int_returns_double(
vec![("metrics", vec![("host", STRING), ("value", INT)])],
r#"
FROM metrics
| AGG stddev_value = stddev(value) BY host
"#,
Struct::default()
.with_str("host", STRING)
.with_str("stddev_value", DOUBLE)
)]
#[case::stddev_double_returns_double(
vec![("metrics", vec![("host", STRING), ("value", DOUBLE)])],
r#"
FROM metrics
| AGG stddev_value = stddev(value) BY host
"#,
Struct::default()
.with_str("host", STRING)
.with_str("stddev_value", DOUBLE)
)]
#[case::stddev_decimal_returns_double(
vec![(
"metrics",
vec![
("host", STRING),
("value", Type::Decimal(Decimal::new(10, 2).unwrap()))
]
)],
r#"
FROM metrics
| AGG stddev_value = stddev(value) BY host
"#,
Struct::default()
.with_str("host", STRING)
.with_str("stddev_value", DOUBLE)
)]
#[case::join_then_drop_joined_struct(
vec![("events", vec![("timestamp", TIMESTAMP), ("count", INT)])],
r#"
DEF baseline = FROM events
| AGG avg = avg(count), stddev = stddev(count);
FROM events
| AGG count = count() BY timestamp = timestamp@h
| JOIN baseline ON true
| SET threshold_upper = baseline.avg + baseline.stddev
| SET threshold_lower = baseline.avg - baseline.stddev
| DROP baseline
"#,
Struct::default()
.with_str("threshold_lower", DOUBLE)
.with_str("threshold_upper", DOUBLE)
.with_str("timestamp", TIMESTAMP)
.with_str("count", INT)
)]
#[case::rebind_struct_to_scalar_shadows_children(
vec![(
"logs",
vec![
("timestamp", TIMESTAMP),
("message", STRING),
("host", Type::Struct(
Struct::default()
.with_str("name", STRING)
.with_str("ip", STRING)
.with_str("os", Type::Struct(
Struct::default()
.with_str("name", STRING)
.with_str("version", STRING)
))
))
]
)],
r#"
FROM logs
| SET host = host.name
"#,
Struct::default()
.with_str("host", STRING)
.with_str("timestamp", TIMESTAMP)
.with_str("message", STRING)
)]
#[case::rebind_nested_struct_to_scalar_shadows_children(
vec![(
"events",
vec![
("timestamp", TIMESTAMP),
("event", Type::Struct(
Struct::default()
.with_str("id", STRING)
.with_str("host", Type::Struct(
Struct::default()
.with_str("name", STRING)
.with_str("ip", STRING)
.with_str("os", Type::Struct(
Struct::default()
.with_str("name", STRING)
.with_str("version", STRING)
))
))
))
]
)],
r#"
FROM events
| SET event.host = event.host.name
"#,
Struct::default()
.with_str("event", Type::Struct(
Struct::default()
.with_str("id", STRING)
.with_str("host", STRING)
))
.with_str("timestamp", TIMESTAMP)
)]
#[case::rebind_after_array_agg(
vec![(
"events",
vec![
("id", INT),
("data", Type::Struct(
Struct::default()
.with_str("host", Type::Struct(
Struct::default()
.with_str("name", STRING)
.with_str("ip", STRING)
))
))
]
)],
r#"
FROM events
| AGG all_data = array_agg(data)
| SET first_data = all_data[0]
| SET first_data.host = first_data.host.name
"#,
Struct::default()
.with_str("first_data", Type::Struct(
Struct::default()
.with_str("host", STRING)
))
.with_str("all_data", Array::new(Type::Struct(
Struct::default()
.with_str("host", Type::Struct(
Struct::default()
.with_str("name", STRING)
.with_str("ip", STRING)
))
)).into())
)]
#[case::rebind_struct_field_from_same_level(
vec![(
"events",
vec![
("id", INT),
("name", STRING),
("ip", STRING)
]
)],
r#"
FROM events
| SET event = {host: {name: name, ip: ip}}
| SET event.host = event.host.name
"#,
Struct::default()
.with_str("event", Type::Struct(
Struct::default()
.with_str("host", STRING)
))
.with_str("id", INT)
.with_str("name", STRING)
.with_str("ip", STRING)
)]
fn test_regressions(
#[case] tables: Vec<(&str, Vec<(&str, Type)>)>,
#[case] query_str: &str,
#[case] expected_schema: Struct,
) -> Result<(), TranslationErrors> {
let query = Query::parse_result(query_str)?;
let catalog = catalog_with_tables(tables);
let typed = type_check_with_catalog(query, catalog).into_result()?;
let actual_schema = typed.schema();
assert_eq!(actual_schema, expected_schema);
Ok(())
}
#[rstest]
#[case::from_aliased("FROM x = events")]
#[case::join_aliased("FROM users | JOIN o = orders ON o.user_id == user_id")]
#[case::join_qualified_aliased("FROM users | JOIN o = schema.orders ON o.id == id")]
#[case::lookup_aliased("FROM users | LOOKUP o = orders ON o.user_id == user_id")]
fn test_display_roundtrip(#[case] input: &str) {
let original = Pipeline::parse_result(input).expect("input should parse");
let displayed = original.to_string();
let reparsed = Pipeline::parse_result(&displayed);
assert!(
reparsed.is_ok(),
"Failed to re-parse displayed pipeline.\n Input: {input}\n Displayed: {displayed}\n Error: {}",
reparsed.unwrap_err()
);
assert_eq!(
reparsed.unwrap(),
original,
"Round-tripped AST differs from original.\n Input: {input}\n Displayed: {displayed}"
);
}
#[test]
fn ui_time_range_injects_synthetic_within_after_union_as_first_command(
) -> Result<(), TranslationErrors> {
let tables = vec![
("a", vec![("timestamp", TIMESTAMP), ("x", INT)]),
("b", vec![("timestamp", TIMESTAMP), ("x", INT)]),
];
let catalog = catalog_with_tables(tables);
let time_range = Arc::new(Expression::parse_result("-1h..0h")?);
let query = Query::parse_result("UNION a, b")?;
let typed = type_check_with_options(
query,
TypeCheckOptions::builder()
.provider(catalog)
.time_range(time_range)
.maybe_default_space(Some(SimpleIdentifier::new("test")))
.build(),
)
.into_result()?;
let commands = typed
.pipeline
.commands()
.expect("expected a valid typed pipeline");
assert_eq!(commands.len(), 2, "UNION then injected WITHIN");
assert!(matches!(commands[0].kind, TypedCommandKind::Union(_)));
assert!(matches!(commands[1].kind, TypedCommandKind::Within(_)));
Ok(())
}