hamelin_lib 0.21.6

Core library for Hamelin query language
Documentation
use std::borrow::Cow;
use std::cell::RefCell;
use std::rc::Rc;

use crate::antlr::hamelinlexer;
use crate::antlr::hamelinlexer::HamelinLexer;
use crate::antlr::hamelinparser::{HamelinParser, HamelinParserContextType};
use crate::antlr::trinotypelexer::TrinoTypeLexer;
use crate::antlr::trinotypeparser::{TrinoTypeParser, TrinoTypeParserContextType};
use crate::err::{Context, LanguageArea, TranslationError, TranslationErrors};
use antlr_rust::atn_config_set::ATNConfigSet;
use antlr_rust::common_token_stream::CommonTokenStream;
use antlr_rust::dfa::DFA;
use antlr_rust::error_listener::ErrorListener;
use antlr_rust::errors::ANTLRError;
use antlr_rust::int_stream::IntStream;
use antlr_rust::recognizer::Recognizer;
use antlr_rust::token::Token;
use antlr_rust::{BailErrorStrategy, DefaultErrorStrategy, InputStream, Parser};

#[derive(Debug)]
pub struct UnicodeInputStream {
    inner: InputStream<Box<[u32]>>,
}

impl UnicodeInputStream {
    /// Create a new `UnicodeInputStream` from a Rust `String`.
    ///
    /// Internally, we convert `String -> Vec<u32> -> Box<[u32]>`,
    /// then hand that `Box<[u32]>` into `InputStream::new_owned(...)`.
    pub fn new_from_string(input: String) -> Self {
        let code_points: Vec<u32> = input.chars().map(|c| c as u32).collect();
        let boxed: Box<[u32]> = code_points.into_boxed_slice();
        let inner_stream = InputStream::new_owned(boxed);
        UnicodeInputStream {
            inner: inner_stream,
        }
    }
}

impl IntStream for UnicodeInputStream {
    fn consume(&mut self) {
        self.inner.consume()
    }
    fn la(&mut self, offset: isize) -> isize {
        self.inner.la(offset)
    }
    fn mark(&mut self) -> isize {
        self.inner.mark()
    }
    fn release(&mut self, marker: isize) {
        self.inner.release(marker)
    }
    fn index(&self) -> isize {
        self.inner.index()
    }
    fn seek(&mut self, index: isize) {
        self.inner.seek(index)
    }
    fn size(&self) -> isize {
        self.inner.size()
    }
    fn get_source_name(&self) -> String {
        self.inner.get_source_name()
    }
}

/// Now implement `CharStream<Cow<'static, str>>` for `UnicodeInputStream`.
/// Whenever ANTLR calls `get_text(start, stop)`, we grab the `&[u32]` slice out
/// of `self.inner` and build a UTF‐8 `String` (or `Cow::Owned`) from those code points.
impl antlr_rust::char_stream::CharStream<Cow<'static, str>> for UnicodeInputStream {
    fn get_text(&self, start: isize, stop: isize) -> Cow<'static, str> {
        let codepoints_in_range: Vec<u32> = self.inner.get_text(start, stop);

        // 2) Convert Vec<u32> -> String (replace invalid codepoints with U+FFFD):
        let mut s = String::with_capacity(codepoints_in_range.len());
        for cp in codepoints_in_range.into_iter() {
            if let Some(ch) = char::from_u32(cp) {
                s.push(ch);
            } else {
                // invalid code point → replacement character
                s.push('\u{FFFD}');
            }
        }
        Cow::Owned(s)
    }
}

// ANTLR also needs the “tidable” implementation so that `HamelinLexer` sees the stream’s TID:
antlr_rust::tid! { impl<'a> TidAble<'a> for UnicodeInputStream }

pub type HamelinStringParser = HamelinParser<
    'static,
    CommonTokenStream<'static, HamelinLexer<'static, UnicodeInputStream>>,
    DefaultErrorStrategy<'static, HamelinParserContextType>,
>;

pub type TrinoTypeStringParser = TrinoTypeParser<
    'static,
    CommonTokenStream<'static, TrinoTypeLexer<'static, InputStream<Box<str>>>>,
    BailErrorStrategy<'static, TrinoTypeParserContextType>,
>;

#[derive(Debug, Clone)]
pub struct CustomErrorListener {
    pub errors: Rc<RefCell<TranslationErrors>>,
}

impl Default for CustomErrorListener {
    fn default() -> Self {
        Self {
            errors: Rc::new(RefCell::new(TranslationErrors::default())),
        }
    }
}

impl<T: Recognizer<'static>> ErrorListener<'static, T> for CustomErrorListener {
    fn syntax_error(
        &self,
        _recognizer: &T,
        offending_symbol: Option<
            &<T::TF as antlr_rust::token_factory::TokenFactory<'static>>::Inner,
        >,
        _line: isize,
        _column: isize,
        msg: &str,
        e: Option<&ANTLRError>,
    ) {
        let range = match offending_symbol {
            Some(symbol) => symbol.get_start() as usize..=symbol.get_stop() as usize,
            None => match e {
                Some(e) => match e {
                    ANTLRError::LexerNoAltError { start_index } => {
                        *start_index as usize..=*start_index as usize
                    }
                    ANTLRError::NoAltError(_) => 0..=0,
                    ANTLRError::InputMismatchError(_) => 0..=0,
                    ANTLRError::PredicateError(_) => 0..=0,
                    ANTLRError::IllegalStateError(_) => 0..=0,
                    ANTLRError::FallThrough(_) => 0..=0,
                    ANTLRError::OtherError(_) => 0..=0,
                },
                None => 0..=0,
            },
        };

        if let Some(symbol) = offending_symbol {
            if symbol.get_token_type() == hamelinlexer::SIMPLE_COMMENT
                || symbol.get_token_type() == hamelinlexer::BRACKETED_COMMENT
                || symbol.get_token_type() == hamelinlexer::WS
            {
                return;
            }
        }

        let te = TranslationError::new(Context::new(range, msg)).with_area(LanguageArea::Parsing);
        self.errors.borrow_mut().add(te);
    }

    fn report_ambiguity(
        &self,
        _recognizer: &T,
        _dfa: &DFA,
        _start_index: isize,
        _stop_index: isize,
        _exact: bool,
        _ambig_alts: &bit_set::BitSet,
        _configs: &ATNConfigSet,
    ) {
        // Handle ambiguities here
        /*
        println!(
            "Ambiguity detected from index {} to {} in DFA state {:?}",
            _start_index, _stop_index, _dfa
        );
         */
    }

    fn report_attempting_full_context(
        &self,
        _recognizer: &T,
        _dfa: &DFA,
        _start_index: isize,
        _stop_index: isize,
        _conflicting_alts: &bit_set::BitSet,
        _configs: &ATNConfigSet,
    ) {
        // Handle full context attempts here
        /*
        println!(
            "Attempting full context from index {} to {} in DFA state {:?}",
            _start_index, _stop_index, _dfa
        );
         */
    }

    fn report_context_sensitivity(
        &self,
        _recognizer: &T,
        _dfa: &DFA,
        _start_index: isize,
        _stop_index: isize,
        _prediction: isize,
        _configs: &ATNConfigSet,
    ) {
        // Handle resolved context sensitivity here
        /*
        println!(
            "Context sensitivity resolved from index {} to {} with prediction {} in DFA state {:?}",
            _start_index, _stop_index, _prediction, _dfa
        );
         */
    }
}

pub fn make_hamelin_parser_from_input(
    input: String,
) -> (HamelinStringParser, Rc<RefCell<TranslationErrors>>) {
    let listener = CustomErrorListener::default();

    let mut lexer = HamelinLexer::new(UnicodeInputStream::new_from_string(input));
    lexer.remove_error_listeners();
    lexer.add_error_listener(Box::new(listener.clone()));

    let token_source = CommonTokenStream::new(lexer);
    let mut parser = HamelinParser::new(token_source);
    let errors = listener.errors.clone();
    parser.remove_error_listeners();
    parser.add_error_listener(Box::new(listener));
    (parser, errors)
}

pub fn make_trinotype_parser_from_input(input: String) -> TrinoTypeStringParser {
    let mut lexer = TrinoTypeLexer::new(InputStream::new_owned(input.into_boxed_str()));
    lexer.remove_error_listeners();

    let token_source = CommonTokenStream::new(lexer);
    let mut parser = TrinoTypeParser::with_strategy(token_source, BailErrorStrategy::new());
    parser.remove_error_listeners();
    parser
}

// Add unit tests here that assess whether antlr ignores comments
#[cfg(test)]
mod tests {

    use rstest::rstest;

    use super::*;

    #[rstest]
    #[case(r#"


        /* what the fuck */ SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
        "#
    )]
    #[case(r#"/* what the fuck */ SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
        "#
    )]
    #[case(r#"// what the fuck
        SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
        "#
    )]
    #[case(r#"
        // what the fuck
        SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
        "#
    )]
    #[case(
        r#"
        SET
            map = map('one': 1, 'two': 2),
            struct = {one: 1, two: 2},
            // what the fuck
            variant = parse_json('{"one": 1}')
        "#
    )]
    fn test_ignores_comments(#[case] input: &str) {
        let (mut parser, errors) = make_hamelin_parser_from_input(input.to_string());

        assert!(parser.queryEOF().is_ok());
        assert!(errors.borrow().is_empty());
    }

    #[rstest]
    #[case("SET resultat = 'café'")]
    #[case("SET chinese = '测试'")]
    #[case("SET emoji = '🚀'")]
    #[case("SET mixed = 'Hello 世界 🌍'")]
    fn test_unicode_handling(#[case] input: &str) {
        let (mut parser, errors) = make_hamelin_parser_from_input(input.to_string());

        assert!(parser.queryEOF().is_ok());
        assert!(errors.borrow().is_empty());
    }
}