use chrono::Months;
use hamelin_lib::tree::{
ast::clause::SortOrder,
typed_ast::expression::{FieldReferenceResolution, FoldExpressionAlgebra, TypedFieldReference},
};
use super::*;
#[derive(Clone, Copy, Debug)]
pub struct InputRange {
pub start: Timestamp,
pub end: Option<Timestamp>,
}
impl InputRange {
pub fn union(self, other: Self) -> Self {
Self {
start: self.start.min(other.start),
end: match (self.end, other.end) {
(Some(a), Some(b)) => Some(a.max(b)),
_ => None,
},
}
}
}
pub fn compute_input_range_backward(
pipeline: &TypedPipeline,
requirements: &mut HashMap<Identifier, InputRange>,
mut required: InputRange,
field: &SimpleIdentifier,
prevalidated_incremental: bool,
) -> Result<InputRange, IncrementalAnalysisError> {
let commands = &pipeline
.kind
.as_ref()
.map_err(|e| IncrementalAnalysisError::TreeHadError(e.as_ref().clone()))?
.commands;
let mut inputs: Option<InputRange> = None;
for command in commands.iter().rev() {
match &command.kind {
TypedCommandKind::Set(c) => {
required = projection_input_range(&c.projections, required, field)?
}
TypedCommandKind::Select(c) => {
if c.projections.lookup(&field.clone().into()).is_none() {
return Err(IncrementalAnalysisError::TimestampLineageError);
}
required = projection_input_range(&c.projections, required, field)?
}
TypedCommandKind::Window(c) => {
if c.sort_by.len() != 1
|| c.sort_by[0].order != SortOrder::Asc
|| !matches!(
&c.sort_by[0].expression.ast.kind,
ExpressionKind::FieldReference(reference)
if reference.field_name.valid_ref().ok() == Some(field)
)
|| !c.sort_by[0].expression.fold(&mut OnlyTimestamp(field))
{
return Err(IncrementalAnalysisError::TimestampLineageError);
}
required = projection_input_range(&c.projections, required, field)?;
inputs = Some(inputs.map_or(required, |range| range.union(required)));
let within = c
.within
.as_ref()
.ok_or(IncrementalAnalysisError::UnboundRange)?;
check_expression_for_non_deterministic(within)?;
let value = eval(within, &Environment::new())
.map_err(IncrementalAnalysisError::ExpressionEvaluationError)?;
required = expand_window(required, value)?;
inputs = Some(inputs.map_or(required, |range| range.union(required)));
required = projection_input_range(&c.group_by, required, field)?;
}
TypedCommandKind::Agg(c) => {
if c.group_by.lookup(&field.clone().into()).is_none() {
return Err(IncrementalAnalysisError::AggWithoutTimestampGroupBy);
}
required = projection_input_range(&c.group_by, required, field)?;
}
TypedCommandKind::Distinct(c) => {
if c.keys.lookup(&field.clone().into()).is_none() {
return Err(IncrementalAnalysisError::TimestampLineageError);
}
required = projection_input_range(&c.keys, required, field)?;
}
TypedCommandKind::Suppress(c) => {
required = projection_input_range(&c.group_by, required, field)?;
let ExpressionKind::IntervalLiteral(interval) = &c.interval.ast.kind else {
return Err(IncrementalAnalysisError::ValueNotInterval);
};
let (unit, multiplier) = interval
.trunc_unit_and_multiplier()
.ok_or(IncrementalAnalysisError::ValueNotInterval)?;
required.start =
*truncate_timestamp(&TimestampValue::utc(required.start), &unit, multiplier)
.map_err(IncrementalAnalysisError::ExpressionEvaluationError)?
.instant();
required.end = required
.end
.map(|end| {
let truncated =
truncate_timestamp(&TimestampValue::utc(end), &unit, multiplier)?;
next_truncation_boundary(&truncated, &unit, multiplier)
.map(|v| *v.instant())
})
.transpose()
.map_err(IncrementalAnalysisError::ExpressionEvaluationError)?;
}
TypedCommandKind::Drop(c) => {
if c.dropped_fields.contains(&field.clone().into()) {
return Err(IncrementalAnalysisError::TimestampLineageError);
}
}
TypedCommandKind::Where(_)
| TypedCommandKind::Sort(_)
| TypedCommandKind::Within(_) => {}
TypedCommandKind::From(TypedFromCommand { clauses })
| TypedCommandKind::Union(TypedUnionCommand { clauses }) => {
for reference in collect_static_refs(clauses)? {
if let Some(name) = dataset_ref_as_def_name(&reference) {
requirements
.entry(name)
.and_modify(|start| *start = start.union(required))
.or_insert(required);
}
}
return Ok(inputs.map_or(required, |range| range.union(required)));
}
TypedCommandKind::Match(c) if prevalidated_incremental => {
if let Some(within) = &c.within {
let value = eval(within, &Environment::new())
.map_err(IncrementalAnalysisError::ExpressionEvaluationError)?;
let offset = match value {
Value::Interval(delta) => Value::Interval(-delta),
Value::CalendarInterval(months) => Value::CalendarInterval(
months
.checked_neg()
.ok_or(IncrementalAnalysisError::UnboundRange)?,
),
_ => return Err(IncrementalAnalysisError::ValueNotInterval),
};
required.start = shift_timestamp(required.start, offset)?;
}
return Ok(inputs.map_or(required, |range| range.union(required)));
}
_ if prevalidated_incremental => {}
_ => {
return Err(IncrementalAnalysisError::CommandNotSupported(
command.ast.kind.command_name().to_string(),
));
}
}
inputs = Some(inputs.map_or(required, |range| range.union(required)));
}
Err(IncrementalAnalysisError::EmptyPipeline)
}
fn projection_input_range(
projections: &Projections,
output: InputRange,
field: &SimpleIdentifier,
) -> Result<InputRange, IncrementalAnalysisError> {
let Some(projection) = projections.lookup(&field.clone().into()) else {
return Ok(output);
};
check_expression_for_non_deterministic(&projection.expression)?;
if !projection.expression.fold(&mut OnlyTimestamp(field)) {
return Err(IncrementalAnalysisError::TimestampLineageError);
}
let constraint = Constraint::Range {
min: Some(TimestampValue::utc(output.start).into()),
max: output.end.map(|end| TimestampValue::utc(end).into()),
};
match reverse_eval(&projection.expression, constraint, &Environment::new()) {
Ok(Some(Constraint::Range {
min: Some(Value::Timestamp(start)),
max,
})) => {
let end = match max {
Some(Value::Timestamp(end)) if output.end.is_some() => Some(*end.instant()),
None if output.end.is_none() => None,
_ => return Err(IncrementalAnalysisError::TimestampLineageError),
};
Ok(InputRange {
start: *start.instant(),
end,
})
}
_ => Err(IncrementalAnalysisError::TimestampLineageError),
}
}
fn expand_window(range: InputRange, value: Value) -> Result<InputRange, IncrementalAnalysisError> {
let (lower, upper) = match value {
Value::Range(value) => (value.lower.clone(), value.upper.clone()),
value => (Some(value.clone()), Some(value)),
};
let start = range.start.min(shift_timestamp(
range.start,
lower.ok_or(IncrementalAnalysisError::UnboundRange)?,
)?);
let end = range
.end
.map(|end| {
shift_timestamp(end, upper.ok_or(IncrementalAnalysisError::UnboundRange)?)
.map(|shifted| end.max(shifted))
})
.transpose()?;
Ok(InputRange { start, end })
}
struct OnlyTimestamp<'a>(&'a SimpleIdentifier);
impl FoldExpressionAlgebra<bool> for OnlyTimestamp<'_> {
fn combine(&mut self, children: impl IntoIterator<Item = bool>) -> bool {
children.into_iter().all(|v| v)
}
fn field_reference(&mut self, node: &TypedFieldReference, _: &TypedExpression) -> bool {
matches!(node.resolution, FieldReferenceResolution::Pipeline)
&& node.field_name.valid_ref().ok() == Some(self.0)
}
fn leaf(&mut self, expr: &TypedExpression) -> bool {
match &expr.ast.kind {
ExpressionKind::FieldReference(reference) => {
reference.field_name.valid_ref().ok() == Some(self.0)
}
_ => true,
}
}
}
fn shift_timestamp(start: Timestamp, value: Value) -> Result<Timestamp, IncrementalAnalysisError> {
match value {
Value::Interval(delta) => start
.checked_add_signed(delta)
.ok_or(IncrementalAnalysisError::UnboundRange),
Value::CalendarInterval(months) => {
let offset = Months::new(months.unsigned_abs());
if months < 0 {
start.checked_sub_months(offset)
} else {
start.checked_add_months(offset)
}
.ok_or(IncrementalAnalysisError::UnboundRange)
}
_ => Err(IncrementalAnalysisError::ValueNotInterval),
}
}