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 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
68impl 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 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 s.push('\u{FFFD}');
83 }
84 }
85 Cow::Owned(s)
86 }
87}
88
89antlr_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 }
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 }
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 }
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#[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}