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 {
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()
}
}
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);
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 {
s.push('\u{FFFD}');
}
}
Cow::Owned(s)
}
}
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,
) {
}
fn report_attempting_full_context(
&self,
_recognizer: &T,
_dfa: &DFA,
_start_index: isize,
_stop_index: isize,
_conflicting_alts: &bit_set::BitSet,
_configs: &ATNConfigSet,
) {
}
fn report_context_sensitivity(
&self,
_recognizer: &T,
_dfa: &DFA,
_start_index: isize,
_stop_index: isize,
_prediction: isize,
_configs: &ATNConfigSet,
) {
}
}
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
}
#[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());
}
}