use std::{collections::HashMap, sync::Arc};
use chrono::{DateTime, TimeZone, Utc};
use crate::incremental::TimeRange;
use hamelin_lib::{
provider::EnvironmentProvider,
tree::{
ast::{
dataset_identifier::QualifiedDatasetIdentifier, identifier::SimpleIdentifier,
query::Query,
},
builder::{pipeline::PipelineBuilder, query as build_query},
options::TypeCheckOptions,
typed_ast::{pipeline::TypedPipeline, query::TypedStatement},
},
type_check_with_options,
types::{array::Array, struct_type::Struct, INT, STRING, TIMESTAMP},
};
pub const TIMESTAMP_FIELD: &str = "timestamp";
pub fn test_default_space() -> SimpleIdentifier {
SimpleIdentifier::new("test")
}
fn test_dataset_id(table_name: &str) -> QualifiedDatasetIdentifier {
QualifiedDatasetIdentifier::from_canonical_str(&format!("test:{table_name}"))
.expect("parse table name")
}
pub fn ts(s: &str) -> DateTime<Utc> {
let parts: Vec<&str> = s.split(' ').collect();
let date_parts: Vec<&str> = parts[0].split('-').collect();
let time_parts: Vec<&str> = parts[1].split(':').collect();
let year = date_parts[0].parse::<i32>().unwrap();
let month = date_parts[1].parse::<u32>().unwrap();
let day = date_parts[2].parse::<u32>().unwrap();
let hour = time_parts[0].parse::<u32>().unwrap();
let minute = time_parts[1].parse::<u32>().unwrap();
let second = time_parts[2].parse::<u32>().unwrap();
Utc.with_ymd_and_hms(year, month, day, hour, minute, second)
.unwrap()
}
pub fn time_range(start: &str, end: &str) -> TimeRange {
(ts(start)..=ts(end)).into()
}
pub fn time_range_exclusive(start: &str, end: &str) -> TimeRange {
(ts(start)..ts(end)).into()
}
pub fn stale_ranges(
table_name: &str,
start: &str,
end: &str,
) -> HashMap<QualifiedDatasetIdentifier, Option<TimeRange>> {
let mut map = HashMap::new();
let table_id = test_stale_key(table_name);
map.insert(table_id, Some(time_range(start, end)));
map
}
pub fn multi_stale_ranges(
ranges: &[(&str, &str, &str)],
) -> HashMap<QualifiedDatasetIdentifier, Option<TimeRange>> {
let mut map = HashMap::new();
for (table_name, start, end) in ranges {
let table_id = test_stale_key(table_name);
map.insert(table_id, Some(time_range(start, end)));
}
map
}
fn test_stale_key(table_name: &str) -> QualifiedDatasetIdentifier {
if table_name.contains(':') {
QualifiedDatasetIdentifier::from_canonical_str(table_name).expect("parse table name")
} else {
QualifiedDatasetIdentifier::from_canonical_str(&format!("test:{table_name}"))
.expect("parse table name")
}
}
pub fn timestamp_field() -> SimpleIdentifier {
TIMESTAMP_FIELD.parse().unwrap()
}
#[derive(Debug)]
pub struct MockIncrementalProvider;
impl EnvironmentProvider for MockIncrementalProvider {
fn reflect_columns(&self, name: &QualifiedDatasetIdentifier) -> anyhow::Result<Struct> {
let events = test_dataset_id("events");
let users = test_dataset_id("users");
let other_events = test_dataset_id("other_events");
if name == &events {
Ok(Struct::default()
.with_str("timestamp", TIMESTAMP)
.with_str("user", STRING)
.with_str("severity", STRING)
.with_str("amount", INT)
.with_str("user_id", STRING)
.with_str("tags", Array::new(STRING).into()))
} else if name == &users {
Ok(Struct::default()
.with_str("id", STRING)
.with_str("name", STRING))
} else if name == &other_events {
Ok(Struct::default()
.with_str("timestamp", TIMESTAMP)
.with_str("data", STRING))
} else {
anyhow::bail!("Unknown table: {:?}", name);
}
}
fn reflect_datasets(&self) -> anyhow::Result<Vec<QualifiedDatasetIdentifier>> {
Ok(vec![
test_dataset_id("events"),
test_dataset_id("users"),
test_dataset_id("other_events"),
])
}
}
pub fn build_typed_query_from_pipeline(builder: PipelineBuilder) -> TypedStatement {
build_typed_query(build_query().main(builder.build()).build())
}
pub fn build_pipeline(builder: PipelineBuilder) -> Arc<TypedPipeline> {
build_typed_query_from_pipeline(builder).pipeline
}
pub fn build_typed_query(query: Query) -> TypedStatement {
type_check_with_options(
query,
TypeCheckOptions::builder()
.provider(Arc::new(MockIncrementalProvider))
.maybe_default_space(Some(SimpleIdentifier::new("test")))
.build(),
)
.output
}