hamelin_lib 0.21.4

Core library for Hamelin query language
Documentation
use std::ops::RangeInclusive;
use std::rc::Rc;

use antlr_rust::errors::ANTLRError;
use antlr_rust::recognizer::Recognizer;
use antlr_rust::token::Token;
use anyhow::anyhow;

use crate::antlr::hamelinparser::{
    CommandContextAll, CommandEOFContextAttrs, ExpressionContextAll, ExpressionEOFContextAll,
    ExpressionEOFContextAttrs, HamelinParserContextType, HamelintypeContextAll,
    IdentifierContextAll, IdentifierEOFContextAll, IdentifierEOFContextAttrs,
    PipelineEOFContextAll, QueryContextAll, QueryEOFContextAll, QueryEOFContextAttrs,
    SimpleIdentifierContextAll, SimpleIdentifierEOFContextAttrs,
};
use crate::err::{Context, TranslationError, TranslationErrors};
use crate::parser::{make_hamelin_parser_from_input, HamelinStringParser};

pub fn execute_resilient_parse<T>(
    input: String,
    parse_fn: fn(HamelinStringParser) -> Result<T, ANTLRError>,
) -> anyhow::Result<(T, TranslationErrors)> {
    let (node, errors) = {
        let (parser, errors) = make_hamelin_parser_from_input(input);
        parse_fn(parser)
            .map_err(|e| anyhow!(e.to_string()))
            .map(|node| (node, errors))?
    };
    let unwrapped = Rc::try_unwrap(errors)
        .map_err(|_| anyhow!("could not unwrap Rc<TranslationErrors> after parsing complete"))?
        .into_inner();
    Ok((node, unwrapped))
}

pub fn execute_parse<T>(
    input: String,
    parse_fn: fn(HamelinStringParser) -> Result<T, ANTLRError>,
) -> Result<T, TranslationErrors> {
    let len = input.len();
    let (node, errors) = {
        let (parser, errors) = make_hamelin_parser_from_input(input);
        parse_fn(parser)
            .map_err(|e| {
                TranslationError::new(Context::new(
                    0..=len.saturating_sub(1),
                    "failed to load grammars",
                ))
                .with_source_boxed(e.to_string().into())
                .single()
            })
            .map(|node| (node, errors))?
    };
    let unwrapped = Rc::try_unwrap(errors)
        .map_err(|_| {
            TranslationError::new(Context::new(0..=0, "internal error"))
                .with_source_boxed(
                    "could not unwrap Rc<TranslationErrors> after parsing complete".into(),
                )
                .single()
        })?
        .into_inner();
    unwrapped.or_ok(node)
}

pub fn parse_command(command: String) -> Result<Rc<CommandContextAll<'static>>, TranslationErrors> {
    execute_parse(command, |mut parser| parser.commandEOF())
        .and_then(|ctx| TranslationErrors::expect(&*ctx, ctx.command()))
}

pub fn parse_query(statement: String) -> Result<Rc<QueryContextAll<'static>>, TranslationErrors> {
    execute_parse(statement, |mut parser| parser.queryEOF())
        .and_then(|ctx| TranslationErrors::expect(&*ctx, ctx.query()))
}

pub fn resilient_parse_query(
    statement: String,
) -> anyhow::Result<(Rc<QueryEOFContextAll<'static>>, TranslationErrors)> {
    execute_resilient_parse(statement, |mut parser| parser.queryEOF())
}

pub fn parse_expression(
    expression: String,
) -> Result<Rc<ExpressionContextAll<'static>>, TranslationErrors> {
    execute_parse(expression, |mut parser| parser.expressionEOF())
        .and_then(|ctx| TranslationErrors::expect(&*ctx, ctx.expression()))
}

pub fn resilient_parse_expression(
    expression: String,
) -> anyhow::Result<(Rc<ExpressionEOFContextAll<'static>>, TranslationErrors)> {
    execute_resilient_parse(expression, |mut parser| parser.expressionEOF())
}

pub fn resilient_parse_pipeline(
    pipeline: String,
) -> anyhow::Result<(Rc<PipelineEOFContextAll<'static>>, TranslationErrors)> {
    execute_resilient_parse(pipeline, |mut parser| parser.pipelineEOF())
}

pub fn parse_identifier(
    identifier: String,
) -> Result<Rc<IdentifierContextAll<'static>>, TranslationErrors> {
    execute_parse(identifier, |mut parser| parser.identifierEOF())
        .and_then(|ctx| TranslationErrors::expect(&*ctx, ctx.identifier()))
}

pub fn resilient_parse_identifier(
    identifier: String,
) -> anyhow::Result<(Rc<IdentifierEOFContextAll<'static>>, TranslationErrors)> {
    execute_resilient_parse(identifier, |mut parser| parser.identifierEOF())
}

pub fn parse_simple_identifier(
    identifier: String,
) -> Result<Rc<SimpleIdentifierContextAll<'static>>, TranslationErrors> {
    execute_parse(identifier, |mut parser| parser.simpleIdentifierEOF())
        .and_then(|ctx| TranslationErrors::expect(&*ctx, ctx.simpleIdentifier()))
}

pub fn parse_type(
    type_string: String,
) -> Result<Rc<HamelintypeContextAll<'static>>, TranslationErrors> {
    execute_parse(type_string, |mut parser| parser.hamelintype())
}

/// A node in the concrete syntax tree, with its rule/token name and source span.
pub struct CstNode {
    pub name: String,
    /// Raw ANTLR code-point span (start, stop). Negative values or start > stop indicate
    /// error-recovery/missing nodes. None only when no token info is available at all.
    pub span: Option<(isize, isize)>,
    pub is_error: bool,
    pub children: Vec<CstNode>,
}

/// Parse a query and return its ANTLR CST as a tree of `CstNode`s with source spans.
pub fn query_cst_tree(input: String) -> anyhow::Result<(CstNode, TranslationErrors)> {
    let (mut parser, errors) = make_hamelin_parser_from_input(input);
    let tree = parser.queryEOF().map_err(|e| anyhow!(e.to_string()))?;
    let rule_names = (*parser).get_rule_names();
    let cst = build_cst_tree(&*tree, rule_names);
    drop(tree);
    drop(parser);
    let unwrapped = Rc::try_unwrap(errors)
        .map_err(|_| anyhow!("could not unwrap Rc<TranslationErrors> after parsing complete"))?
        .into_inner();
    Ok((cst, unwrapped))
}

type HamelinNode = <HamelinParserContextType as antlr_rust::parser::ParserNodeType<'static>>::Type;

fn node_span(tree: &HamelinNode) -> Option<(isize, isize)> {
    use antlr_rust::tree::{ErrorNode, TerminalNode};
    use antlr_rust::TidExt;

    if let Some(tok) = tree.downcast_ref::<TerminalNode<HamelinParserContextType>>() {
        return Some((tok.symbol.get_start(), tok.symbol.get_stop()));
    }
    if let Some(tok) = tree.downcast_ref::<ErrorNode<HamelinParserContextType>>() {
        return Some((tok.symbol.get_start(), tok.symbol.get_stop()));
    }

    Some((tree.start().get_start(), tree.stop().get_stop()))
}

fn build_cst_tree(tree: &HamelinNode, rule_names: &[&str]) -> CstNode {
    use antlr_rust::tree::ErrorNode;
    use antlr_rust::TidExt;

    let span = node_span(tree);
    let is_error_terminal = tree
        .downcast_ref::<ErrorNode<HamelinParserContextType>>()
        .is_some();
    let child_count = tree.get_child_count();
    let name = if child_count == 0 {
        tree.get_text()
    } else {
        antlr_rust::trees::get_node_text(tree, rule_names)
    };
    let has_inverted_span = matches!(span, Some((s, e)) if s > e);
    let is_error = is_error_terminal || name.is_empty() || has_inverted_span;
    let children: Vec<CstNode> = tree
        .get_children()
        .map(|child| build_cst_tree(&*child, rule_names))
        .collect();
    CstNode {
        name,
        span,
        is_error,
        children,
    }
}

#[cfg(test)]
mod resilient_parse_tests {
    use crate::antlr::hamelinparser::QueryEOFContextAttrs;

    use super::resilient_parse_query;

    #[test]
    fn def_from_incomplete_parses() {
        let (eof, _errs) = resilient_parse_query("DEF x = FROM ".to_string()).expect("parse");
        assert!(eof.query().is_some());
    }

    #[test]
    fn def_with_semicolon_is_pipeline_query_with_defs() {
        use crate::antlr::hamelinparser::{
            QueryContextAll, QueryMainContextAll, QueryWithDefsContextAttrs,
        };

        let (eof, _) = resilient_parse_query(
            "DEF x = FROM simba.okta | SELECT published ; FROM x | SELECT p".to_string(),
        )
        .expect("parse");
        let q = eof.query().expect("query");
        assert!(matches!(
            q.as_ref(),
            QueryContextAll::QueryWithDefsContext(_)
        ));
        let QueryContextAll::QueryWithDefsContext(s) = q.as_ref() else {
            unreachable!()
        };
        assert!(!s.defStatement_all().is_empty());
        assert!(matches!(
            s.queryMain().as_deref(),
            Some(QueryMainContextAll::PipelineQueryMainContext(_))
        ));
    }

    #[test]
    fn main_pipeline_may_end_with_semicolon() {
        use crate::antlr::hamelinparser::{
            QueryContextAll, QueryMainContextAll, QueryWithoutDefsContextAttrs,
        };

        let (eof, _) = resilient_parse_query("FROM t | SELECT x;".to_string()).expect("parse");
        let q = eof.query().expect("query");
        assert!(matches!(
            q.as_ref(),
            QueryContextAll::QueryWithoutDefsContext(_)
        ));
        let QueryContextAll::QueryWithoutDefsContext(s) = q.as_ref() else {
            unreachable!()
        };
        assert!(matches!(
            s.queryMain().as_deref(),
            Some(QueryMainContextAll::PipelineQueryMainContext(_))
        ));
    }

    #[test]
    fn main_expression_may_end_with_semicolon() {
        use crate::antlr::hamelinparser::{
            QueryContextAll, QueryMainContextAll, QueryWithoutDefsContextAttrs,
        };

        let (eof, _) = resilient_parse_query("1 + 1;".to_string()).expect("parse");
        let q = eof.query().expect("query");
        assert!(matches!(
            q.as_ref(),
            QueryContextAll::QueryWithoutDefsContext(_)
        ));
        let QueryContextAll::QueryWithoutDefsContext(s) = q.as_ref() else {
            unreachable!()
        };
        assert!(matches!(
            s.queryMain().as_deref(),
            Some(QueryMainContextAll::ExpressionQueryMainContext(_))
        ));
    }
}