hamelin_analysis 0.23.1

Analysis utilities for Hamelin query language
Documentation
//! Shared backward propagation for bounded and open-ended output requirements.

use chrono::Months;
use hamelin_lib::tree::{
    ast::clause::SortOrder,
    typed_ast::expression::{FieldReferenceResolution, FoldExpressionAlgebra, TypedFieldReference},
};

use super::*;

/// Internal range with a genuinely open upper bound; no sentinel timestamp
/// participates in reverse evaluation or interval arithmetic.
#[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) => {
                // Interval expansion is supported only on the ascending timestamp
                // clock. Other expressions or directions need separate analysis.
                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)));
                // Grouping assignments are evaluated before the window, so reverse
                // them after expanding the frame in the window's timestamp domain.
                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)));
            }
            // Incremental callers have already checked command eligibility in
            // their forward pass (including lookup policy). The suffix API has
            // no such validation and must reject unsupported commands here.
            _ 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),
    }
}