hamelin_lib 0.15.6

Core library for Hamelin query language
Documentation
//! Shared AST traversal for rewriting [`TableReference`] nodes inside a [`Query`].

use std::sync::Arc;

use crate::tree::ast::clause::{FromClause, TableAlias, TableReference};
use crate::tree::ast::command::{
    AppendCommand, Command, CommandKind, FromCommand, JoinCommand, LookupCommand, MatchCommand,
    UnionCommand,
};
use crate::tree::ast::node::Span;
use crate::tree::ast::pattern::{NestedPattern, Pattern, PatternKind, QuantifiedPattern};
use crate::tree::ast::pipeline::Pipeline;
use crate::tree::ast::query::{DefBody, DefStatement, Query, ValidQuery};

/// Rewrite individual table references while traversing a query.
pub trait MapTableReferences {
    type Error;

    fn map_table_reference(
        &self,
        reference: &TableReference,
    ) -> Result<TableReference, Self::Error>;
}

/// Walk a validated query and apply `mapper` to every table reference.
pub fn map_table_references_in_valid_query<M: MapTableReferences>(
    valid: &ValidQuery,
    query_span: Span,
    mapper: &M,
) -> Result<Query, M::Error> {
    let mut defs = Vec::with_capacity(valid.defs.len());
    for d in &valid.defs {
        let body = match &d.body {
            DefBody::Pipeline(p) => DefBody::Pipeline(Arc::new(map_table_references_in_pipeline(
                p.as_ref(),
                mapper,
            )?)),
            DefBody::Expression(e) => DefBody::Expression(e.clone()),
        };
        defs.push(DefStatement {
            span: d.span,
            name: d.name.clone(),
            body,
        });
    }

    let main_pipeline = Arc::new(map_table_references_in_pipeline(
        valid.main_pipeline.as_ref(),
        mapper,
    )?);

    Ok(Query {
        span: query_span,
        kind: ValidQuery {
            span: valid.span,
            defs,
            main_pipeline,
        }
        .into(),
    })
}

fn map_table_references_in_from_clause<M: MapTableReferences>(
    clause: &FromClause,
    mapper: &M,
) -> Result<FromClause, M::Error> {
    match clause {
        FromClause::TableReference(tr) => Ok(FromClause::TableReference(Arc::new(
            mapper.map_table_reference(tr.as_ref())?,
        ))),
        FromClause::TableAlias(a) => Ok(FromClause::TableAlias(Arc::new(TableAlias {
            span: a.span,
            alias: a.alias.clone(),
            table: mapper.map_table_reference(&a.table)?,
        }))),
    }
}

fn map_table_references_in_pattern<M: MapTableReferences>(
    pattern: &Pattern,
    mapper: &M,
) -> Result<Pattern, M::Error> {
    let span = pattern.span;
    let kind = match &pattern.kind {
        PatternKind::Quantified(q) => QuantifiedPattern {
            span: q.span,
            from_clause: Arc::new(map_table_references_in_from_clause(
                q.from_clause.as_ref(),
                mapper,
            )?),
            quantifier: q.quantifier.clone(),
        }
        .into(),
        PatternKind::Nested(n) => {
            let mut patterns = Vec::with_capacity(n.patterns.len());
            for p in &n.patterns {
                patterns.push(Arc::new(map_table_references_in_pattern(p, mapper)?));
            }
            NestedPattern {
                span: n.span,
                patterns,
                quantifier: n.quantifier.clone(),
            }
            .into()
        }
        PatternKind::Error(e) => e.clone().into(),
    };
    Ok(Pattern { span, kind })
}

fn map_table_references_in_command<M: MapTableReferences>(
    cmd: &Command,
    mapper: &M,
) -> Result<Command, M::Error> {
    let span = cmd.span;
    let kind = match &cmd.kind {
        CommandKind::From(c) => {
            let mut clauses = Vec::with_capacity(c.clauses.len());
            for cl in &c.clauses {
                clauses.push(Arc::new(map_table_references_in_from_clause(cl, mapper)?));
            }
            FromCommand { clauses }.into()
        }
        CommandKind::Union(c) => {
            let mut clauses = Vec::with_capacity(c.clauses.len());
            for cl in &c.clauses {
                clauses.push(Arc::new(map_table_references_in_from_clause(cl, mapper)?));
            }
            UnionCommand { clauses }.into()
        }
        CommandKind::Join(c) => JoinCommand {
            other: Arc::new(map_table_references_in_from_clause(
                c.other.as_ref(),
                mapper,
            )?),
            on_condition: c.on_condition.clone(),
        }
        .into(),
        CommandKind::Lookup(c) => LookupCommand {
            other: Arc::new(map_table_references_in_from_clause(
                c.other.as_ref(),
                mapper,
            )?),
            on_condition: c.on_condition.clone(),
        }
        .into(),
        CommandKind::Append(c) => AppendCommand {
            table: Arc::new(mapper.map_table_reference(c.table.as_ref())?),
            distinct_by: c.distinct_by.clone(),
        }
        .into(),
        CommandKind::Match(c) => {
            let mut pattern = Vec::with_capacity(c.pattern.len());
            for p in &c.pattern {
                pattern.push(Arc::new(map_table_references_in_pattern(p, mapper)?));
            }
            MatchCommand {
                pattern,
                agg: c.agg.clone(),
                group_by: c.group_by.clone(),
                sort: c.sort.clone(),
                within: c.within.clone(),
            }
            .into()
        }
        CommandKind::Set(_) => cmd.kind.clone(),
        CommandKind::Where(_) => cmd.kind.clone(),
        CommandKind::Select(_) => cmd.kind.clone(),
        CommandKind::Drop(_) => cmd.kind.clone(),
        CommandKind::Limit(_) => cmd.kind.clone(),
        CommandKind::Trimstrings(_) => cmd.kind.clone(),
        CommandKind::Within(_) => cmd.kind.clone(),
        CommandKind::Sort(_) => cmd.kind.clone(),
        CommandKind::Parse(_) => cmd.kind.clone(),
        CommandKind::Agg(_) => cmd.kind.clone(),
        CommandKind::Distinct(_) => cmd.kind.clone(),
        CommandKind::Suppress(_) => cmd.kind.clone(),
        CommandKind::Window(_) => cmd.kind.clone(),
        CommandKind::Explode(_) => cmd.kind.clone(),
        CommandKind::Unnest(_) => cmd.kind.clone(),
        CommandKind::Rows(_) => cmd.kind.clone(),
        CommandKind::Nest(_) => cmd.kind.clone(),
        CommandKind::Error(_) => cmd.kind.clone(),
    };
    Ok(Command { span, kind })
}

fn map_table_references_in_pipeline<M: MapTableReferences>(
    pipeline: &Pipeline,
    mapper: &M,
) -> Result<Pipeline, M::Error> {
    let mut commands = Vec::with_capacity(pipeline.commands.len());
    for cmd in &pipeline.commands {
        commands.push(Arc::new(map_table_references_in_command(cmd, mapper)?));
    }
    Ok(Pipeline {
        span: pipeline.span,
        commands,
    })
}