Skip to main content

hamelin_lib/
parser.rs

1use std::borrow::Cow;
2use std::cell::RefCell;
3use std::rc::Rc;
4
5use crate::antlr::hamelinlexer;
6use crate::antlr::hamelinlexer::HamelinLexer;
7use crate::antlr::hamelinparser::{HamelinParser, HamelinParserContextType};
8use crate::antlr::trinotypelexer::TrinoTypeLexer;
9use crate::antlr::trinotypeparser::{TrinoTypeParser, TrinoTypeParserContextType};
10use crate::err::{Context, LanguageArea, TranslationError, TranslationErrors};
11use antlr_rust::atn_config_set::ATNConfigSet;
12use antlr_rust::common_token_stream::CommonTokenStream;
13use antlr_rust::dfa::DFA;
14use antlr_rust::error_listener::ErrorListener;
15use antlr_rust::errors::ANTLRError;
16use antlr_rust::int_stream::IntStream;
17use antlr_rust::recognizer::Recognizer;
18use antlr_rust::token::Token;
19use antlr_rust::{BailErrorStrategy, DefaultErrorStrategy, InputStream, Parser};
20
21#[derive(Debug)]
22pub struct UnicodeInputStream {
23    inner: InputStream<Box<[u32]>>,
24}
25
26impl UnicodeInputStream {
27    /// Create a new `UnicodeInputStream` from a Rust `String`.
28    ///
29    /// Internally, we convert `String -> Vec<u32> -> Box<[u32]>`,
30    /// then hand that `Box<[u32]>` into `InputStream::new_owned(...)`.
31    pub fn new_from_string(input: String) -> Self {
32        let code_points: Vec<u32> = input.chars().map(|c| c as u32).collect();
33        let boxed: Box<[u32]> = code_points.into_boxed_slice();
34        let inner_stream = InputStream::new_owned(boxed);
35        UnicodeInputStream {
36            inner: inner_stream,
37        }
38    }
39}
40
41impl IntStream for UnicodeInputStream {
42    fn consume(&mut self) {
43        self.inner.consume()
44    }
45    fn la(&mut self, offset: isize) -> isize {
46        self.inner.la(offset)
47    }
48    fn mark(&mut self) -> isize {
49        self.inner.mark()
50    }
51    fn release(&mut self, marker: isize) {
52        self.inner.release(marker)
53    }
54    fn index(&self) -> isize {
55        self.inner.index()
56    }
57    fn seek(&mut self, index: isize) {
58        self.inner.seek(index)
59    }
60    fn size(&self) -> isize {
61        self.inner.size()
62    }
63    fn get_source_name(&self) -> String {
64        self.inner.get_source_name()
65    }
66}
67
68/// Now implement `CharStream<Cow<'static, str>>` for `UnicodeInputStream`.
69/// Whenever ANTLR calls `get_text(start, stop)`, we grab the `&[u32]` slice out
70/// of `self.inner` and build a UTF‐8 `String` (or `Cow::Owned`) from those code points.
71impl antlr_rust::char_stream::CharStream<Cow<'static, str>> for UnicodeInputStream {
72    fn get_text(&self, start: isize, stop: isize) -> Cow<'static, str> {
73        let codepoints_in_range: Vec<u32> = self.inner.get_text(start, stop);
74
75        // 2) Convert Vec<u32> -> String (replace invalid codepoints with U+FFFD):
76        let mut s = String::with_capacity(codepoints_in_range.len());
77        for cp in codepoints_in_range.into_iter() {
78            if let Some(ch) = char::from_u32(cp) {
79                s.push(ch);
80            } else {
81                // invalid code point → replacement character
82                s.push('\u{FFFD}');
83            }
84        }
85        Cow::Owned(s)
86    }
87}
88
89// ANTLR also needs the “tidable” implementation so that `HamelinLexer` sees the stream’s TID:
90antlr_rust::tid! { impl<'a> TidAble<'a> for UnicodeInputStream }
91
92pub type HamelinStringParser = HamelinParser<
93    'static,
94    CommonTokenStream<'static, HamelinLexer<'static, UnicodeInputStream>>,
95    DefaultErrorStrategy<'static, HamelinParserContextType>,
96>;
97
98pub type TrinoTypeStringParser = TrinoTypeParser<
99    'static,
100    CommonTokenStream<'static, TrinoTypeLexer<'static, InputStream<Box<str>>>>,
101    BailErrorStrategy<'static, TrinoTypeParserContextType>,
102>;
103
104#[derive(Debug, Clone)]
105pub struct CustomErrorListener {
106    pub errors: Rc<RefCell<TranslationErrors>>,
107}
108
109impl Default for CustomErrorListener {
110    fn default() -> Self {
111        Self {
112            errors: Rc::new(RefCell::new(TranslationErrors::default())),
113        }
114    }
115}
116
117impl<T: Recognizer<'static>> ErrorListener<'static, T> for CustomErrorListener {
118    fn syntax_error(
119        &self,
120        _recognizer: &T,
121        offending_symbol: Option<
122            &<T::TF as antlr_rust::token_factory::TokenFactory<'static>>::Inner,
123        >,
124        _line: isize,
125        _column: isize,
126        msg: &str,
127        e: Option<&ANTLRError>,
128    ) {
129        let range = match offending_symbol {
130            Some(symbol) => symbol.get_start() as usize..=symbol.get_stop() as usize,
131            None => match e {
132                Some(e) => match e {
133                    ANTLRError::LexerNoAltError { start_index } => {
134                        *start_index as usize..=*start_index as usize
135                    }
136                    ANTLRError::NoAltError(_) => 0..=0,
137                    ANTLRError::InputMismatchError(_) => 0..=0,
138                    ANTLRError::PredicateError(_) => 0..=0,
139                    ANTLRError::IllegalStateError(_) => 0..=0,
140                    ANTLRError::FallThrough(_) => 0..=0,
141                    ANTLRError::OtherError(_) => 0..=0,
142                },
143                None => 0..=0,
144            },
145        };
146
147        if let Some(symbol) = offending_symbol {
148            if symbol.get_token_type() == hamelinlexer::SIMPLE_COMMENT
149                || symbol.get_token_type() == hamelinlexer::BRACKETED_COMMENT
150                || symbol.get_token_type() == hamelinlexer::WS
151            {
152                return;
153            }
154        }
155
156        let te = TranslationError::new(Context::new(range, msg)).with_area(LanguageArea::Parsing);
157        self.errors.borrow_mut().add(te);
158    }
159
160    fn report_ambiguity(
161        &self,
162        _recognizer: &T,
163        _dfa: &DFA,
164        _start_index: isize,
165        _stop_index: isize,
166        _exact: bool,
167        _ambig_alts: &bit_set::BitSet,
168        _configs: &ATNConfigSet,
169    ) {
170        // Handle ambiguities here
171        /*
172        println!(
173            "Ambiguity detected from index {} to {} in DFA state {:?}",
174            _start_index, _stop_index, _dfa
175        );
176         */
177    }
178
179    fn report_attempting_full_context(
180        &self,
181        _recognizer: &T,
182        _dfa: &DFA,
183        _start_index: isize,
184        _stop_index: isize,
185        _conflicting_alts: &bit_set::BitSet,
186        _configs: &ATNConfigSet,
187    ) {
188        // Handle full context attempts here
189        /*
190        println!(
191            "Attempting full context from index {} to {} in DFA state {:?}",
192            _start_index, _stop_index, _dfa
193        );
194         */
195    }
196
197    fn report_context_sensitivity(
198        &self,
199        _recognizer: &T,
200        _dfa: &DFA,
201        _start_index: isize,
202        _stop_index: isize,
203        _prediction: isize,
204        _configs: &ATNConfigSet,
205    ) {
206        // Handle resolved context sensitivity here
207        /*
208        println!(
209            "Context sensitivity resolved from index {} to {} with prediction {} in DFA state {:?}",
210            _start_index, _stop_index, _prediction, _dfa
211        );
212         */
213    }
214}
215
216pub fn make_hamelin_parser_from_input(
217    input: String,
218) -> (HamelinStringParser, Rc<RefCell<TranslationErrors>>) {
219    let listener = CustomErrorListener::default();
220
221    let mut lexer = HamelinLexer::new(UnicodeInputStream::new_from_string(input));
222    lexer.remove_error_listeners();
223    lexer.add_error_listener(Box::new(listener.clone()));
224
225    let token_source = CommonTokenStream::new(lexer);
226    let mut parser = HamelinParser::new(token_source);
227    let errors = listener.errors.clone();
228    parser.remove_error_listeners();
229    parser.add_error_listener(Box::new(listener));
230    (parser, errors)
231}
232
233pub fn make_trinotype_parser_from_input(input: String) -> TrinoTypeStringParser {
234    let mut lexer = TrinoTypeLexer::new(InputStream::new_owned(input.into_boxed_str()));
235    lexer.remove_error_listeners();
236
237    let token_source = CommonTokenStream::new(lexer);
238    let mut parser = TrinoTypeParser::with_strategy(token_source, BailErrorStrategy::new());
239    parser.remove_error_listeners();
240    parser
241}
242
243// Add unit tests here that assess whether antlr ignores comments
244#[cfg(test)]
245mod tests {
246
247    use rstest::rstest;
248
249    use super::*;
250
251    #[rstest]
252    #[case(r#"
253
254
255        /* what the fuck */ SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
256        "#
257    )]
258    #[case(r#"/* what the fuck */ SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
259        "#
260    )]
261    #[case(r#"// what the fuck
262        SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
263        "#
264    )]
265    #[case(r#"
266        // what the fuck
267        SET map = map('one': 1, 'two': 2), struct = {one: 1, two: 2}, variant = parse_json('{"one": 1}')
268        "#
269    )]
270    #[case(
271        r#"
272        SET
273            map = map('one': 1, 'two': 2),
274            struct = {one: 1, two: 2},
275            // what the fuck
276            variant = parse_json('{"one": 1}')
277        "#
278    )]
279    fn test_ignores_comments(#[case] input: &str) {
280        let (mut parser, errors) = make_hamelin_parser_from_input(input.to_string());
281
282        assert!(parser.queryEOF().is_ok());
283        assert!(errors.borrow().is_empty());
284    }
285
286    #[rstest]
287    #[case("SET resultat = 'café'")]
288    #[case("SET chinese = '测试'")]
289    #[case("SET emoji = '🚀'")]
290    #[case("SET mixed = 'Hello 世界 🌍'")]
291    fn test_unicode_handling(#[case] input: &str) {
292        let (mut parser, errors) = make_hamelin_parser_from_input(input.to_string());
293
294        assert!(parser.queryEOF().is_ok());
295        assert!(errors.borrow().is_empty());
296    }
297}