Skip to main content

conjure_cp_essence_parser/parser/
parse_exprs.rs

1use super::ParseContext;
2use super::util::{get_expr_tree, query_toplevel};
3use crate::diagnostics::source_map::SourceMap;
4use crate::errors::FatalParseError;
5use crate::expression::parse_expression;
6use crate::util::TypecheckingContext;
7use crate::util::node_is_expression;
8use conjure_cp_core::ast::{Expression, SymbolTablePtr};
9use std::collections::BTreeMap;
10#[allow(unused)]
11use uniplate::Uniplate;
12
13pub fn parse_expr(
14    src: &str,
15    symbols_ptr: SymbolTablePtr,
16) -> Result<Option<Expression>, FatalParseError> {
17    let exprs = parse_exprs(src, symbols_ptr)?;
18    if exprs.len() != 1 {
19        return Ok(None);
20    }
21    Ok(Some(exprs[0].clone()))
22}
23
24pub fn parse_exprs(
25    src: &str,
26    symbols_ptr: SymbolTablePtr,
27) -> Result<Vec<Expression>, FatalParseError> {
28    let Some((tree, source_code)) = get_expr_tree(src) else {
29        return Ok(Vec::new());
30    };
31
32    let root = tree.root_node();
33    let mut source_map = SourceMap::default();
34    let mut decl_spans = BTreeMap::new();
35    let mut errors = Vec::new();
36    let mut ctx = ParseContext::new(
37        &source_code,
38        &root,
39        Some(symbols_ptr),
40        &mut errors,
41        &mut source_map,
42        &mut decl_spans,
43    );
44    let mut ans = Vec::new();
45    for expr in query_toplevel(&root, &node_is_expression) {
46        ctx.typechecking_context = TypecheckingContext::Unknown;
47        ctx.inner_typechecking_context = TypecheckingContext::Unknown;
48        let Some(expr) = parse_expression(&mut ctx, expr)? else {
49            continue;
50        };
51        ans.push(expr);
52    }
53    Ok(ans)
54}
55
56mod test {
57    #[allow(unused)]
58    use super::{parse_expr, parse_exprs};
59    #[allow(unused)]
60    use conjure_cp_core::ast::SymbolTablePtr;
61    #[allow(unused)]
62    use conjure_cp_core::ast::{
63        Atom, DeclarationPtr, Domain, Expression, Literal, Metadata, Moo, Name, ReturnType,
64        SymbolTable, Typeable,
65    };
66    #[allow(unused)]
67    use std::collections::HashMap;
68    #[allow(unused)]
69    use std::sync::Arc;
70    #[allow(unused)]
71    use tree_sitter::Range;
72
73    #[test]
74    pub fn test_parse_constant() {
75        let symbols = SymbolTablePtr::new();
76
77        assert_eq!(
78            parse_expr("42", symbols.clone()).unwrap().unwrap(),
79            Expression::Atomic(Metadata::new(), Atom::Literal(Literal::Int(42)))
80        );
81        assert_eq!(
82            parse_expr("true", symbols.clone()).unwrap().unwrap(),
83            Expression::Atomic(Metadata::new(), Atom::Literal(Literal::Bool(true)))
84        );
85        assert_eq!(
86            parse_expr("false", symbols).unwrap().unwrap(),
87            Expression::Atomic(Metadata::new(), Atom::Literal(Literal::Bool(false)))
88        )
89    }
90
91    #[test]
92    pub fn test_parse_expressions() {
93        let src = "x >= 5, y = a / 2";
94        let symbols = SymbolTablePtr::new();
95        let x = DeclarationPtr::new_find(
96            Name::User("x".into()),
97            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(0, 10)]),
98        );
99
100        let y = DeclarationPtr::new_find(
101            Name::User("y".into()),
102            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(0, 10)]),
103        );
104
105        let a = DeclarationPtr::new_find(
106            Name::User("a".into()),
107            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(0, 10)]),
108        );
109
110        // Clone the Rc when inserting!
111        symbols
112            .write()
113            .insert(x.clone())
114            .expect("x should not exist in the symbol-table yet, so we should be able to add it");
115
116        symbols
117            .write()
118            .insert(y.clone())
119            .expect("y should not exist in the symbol-table yet, so we should be able to add it");
120
121        symbols
122            .write()
123            .insert(a.clone())
124            .expect("a should not exist in the symbol-table yet, so we should be able to add it");
125
126        let exprs = parse_exprs(src, symbols).unwrap();
127        assert_eq!(exprs.len(), 2);
128
129        assert_eq!(
130            exprs[0],
131            Expression::Geq(
132                Metadata::new(),
133                Moo::new(Expression::Atomic(Metadata::new(), Atom::new_ref(x))),
134                Moo::new(Expression::Atomic(Metadata::new(), 5.into()))
135            )
136        );
137
138        assert_eq!(
139            exprs[1],
140            Expression::Eq(
141                Metadata::new(),
142                Moo::new(Expression::Atomic(Metadata::new(), Atom::new_ref(y))),
143                Moo::new(Expression::UnsafeDiv(
144                    Metadata::new(),
145                    Moo::new(Expression::Atomic(Metadata::new(), Atom::new_ref(a))),
146                    Moo::new(Expression::Atomic(Metadata::new(), 2.into()))
147                ))
148            )
149        );
150    }
151
152    #[test]
153    fn bars_distinguish_set_cardinality_from_integer_absolute_value() {
154        let symbols = SymbolTablePtr::new();
155        let set = DeclarationPtr::new_find(
156            Name::User("s".into()),
157            Domain::set(
158                conjure_cp_core::ast::SetAttr::new_max_size(2),
159                Domain::int(vec![conjure_cp_core::ast::Range::Bounded(1, 2)]),
160            ),
161        );
162        let integer = DeclarationPtr::new_find(
163            Name::User("x".into()),
164            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(-2, 2)]),
165        );
166        symbols.write().insert(set).unwrap();
167        symbols.write().insert(integer).unwrap();
168
169        let set_expr = parse_expr("|s| = 1", symbols.clone()).unwrap().unwrap();
170        let Expression::Eq(_, set_left, _) = set_expr else {
171            panic!("expected set cardinality comparison");
172        };
173        assert!(matches!(*set_left, Expression::Card(..)));
174
175        let int_expr = parse_expr("|x| = 1", symbols).unwrap().unwrap();
176        let Expression::Eq(_, int_left, _) = int_expr else {
177            panic!("expected integer absolute-value comparison");
178        };
179        assert!(matches!(*int_left, Expression::Abs(..)));
180    }
181
182    #[test]
183    pub fn test_parse_expression_annotations() {
184        let symbols = SymbolTablePtr::new();
185        let x = DeclarationPtr::new_find(
186            Name::User("x".into()),
187            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(0, 10)]),
188        );
189        symbols
190            .write()
191            .insert(x.clone())
192            .expect("x should not exist in the symbol-table yet, so we should be able to add it");
193
194        let domain_annotation = parse_expr("x : int(1..3)", symbols.clone())
195            .unwrap()
196            .unwrap();
197        assert!(matches!(
198            domain_annotation,
199            Expression::DomainAnnotation(_, _, _)
200        ));
201        assert_eq!(domain_annotation.to_string(), "x : int(1..3)");
202
203        let type_annotation = parse_expr("x :: int", symbols.clone()).unwrap().unwrap();
204        assert!(matches!(
205            type_annotation,
206            Expression::TypeAnnotation(_, _, _)
207        ));
208        assert_eq!(type_annotation.return_type(), ReturnType::Int);
209        assert_eq!(type_annotation.to_string(), "x :: int");
210    }
211
212    #[test]
213    pub fn test_parse_set_representation_preference() {
214        let symbols = SymbolTablePtr::new();
215        let x = DeclarationPtr::new_find(
216            Name::User("x".into()),
217            Domain::set(
218                conjure_cp_core::ast::SetAttr::new_max_size(3),
219                Domain::int(vec![conjure_cp_core::ast::Range::Bounded(1, 4)]),
220            ),
221        );
222        symbols.write().insert(x).unwrap();
223
224        let find_domain = parse_expr("x : set (representation packed) of int", symbols.clone())
225            .unwrap()
226            .unwrap();
227        let Expression::DomainAnnotation(_, _, domain) = find_domain else {
228            panic!("expected domain annotation");
229        };
230        assert_eq!(domain.representation_preference(), Some("packed"));
231        assert_eq!(
232            domain.to_string(),
233            "set (representation packed) of int(-2147483647..2147483647)"
234        );
235
236        let type_ann = parse_expr(
237            "x :: set (representation occurrence) of int",
238            symbols.clone(),
239        )
240        .unwrap()
241        .unwrap();
242        assert_eq!(
243            type_ann.to_string(),
244            "x :: set (representation occurrence) of int"
245        );
246        let Expression::TypeAnnotation(_, _, ty_domain) = type_ann else {
247            panic!("expected type annotation");
248        };
249        assert_eq!(ty_domain.representation_preference(), Some("occurrence"));
250
251        let nested = parse_expr(
252            "x : set (representation explicit) of set (representation occurrence) of int",
253            symbols,
254        )
255        .unwrap()
256        .unwrap();
257        let Expression::DomainAnnotation(_, _, nested_domain) = nested else {
258            panic!("expected domain annotation");
259        };
260        assert_eq!(nested_domain.representation_preference(), Some("explicit"));
261        let (_, inner) = nested_domain.as_set().unwrap();
262        assert_eq!(inner.representation_preference(), Some("occurrence"));
263        assert_eq!(
264            nested_domain.to_string(),
265            "set (representation explicit) of set (representation occurrence) of int(-2147483647..2147483647)"
266        );
267    }
268
269    #[test]
270    pub fn test_parse_nested_mset_representation_preference() {
271        let symbols = SymbolTablePtr::new();
272        let x = DeclarationPtr::new_find(Name::User("x".into()), Domain::bool());
273        symbols.write().insert(x).unwrap();
274
275        let annotation = parse_expr(
276            "x : matrix indexed by [int(1..2)] of record { before: mset (representation repetition, maxSize 6) of int(1..9) }",
277            symbols,
278        )
279        .unwrap()
280        .unwrap();
281        let Expression::DomainAnnotation(_, _, domain) = annotation else {
282            panic!("expected domain annotation");
283        };
284
285        assert!(domain.has_representation_preference());
286        let (record, _) = domain.as_matrix().expect("expected matrix domain");
287        let fields = record.as_record().expect("expected record domain");
288        let before = &fields
289            .iter()
290            .find(|field| field.name == Name::User("before".into()))
291            .expect("expected before field")
292            .value;
293        assert_eq!(before.representation_preference(), Some("repetition"));
294        assert_eq!(
295            domain.to_string(),
296            "matrix indexed by [int(1..2)] of record {before: mset (representation repetition, maxSize 6) of int(1..9)}"
297        );
298    }
299
300    #[test]
301    pub fn test_expression_annotations_bind_tighter_than_addition() {
302        let symbols = SymbolTablePtr::new();
303        let x = DeclarationPtr::new_find(
304            Name::User("x".into()),
305            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(0, 10)]),
306        );
307        symbols
308            .write()
309            .insert(x)
310            .expect("x should not exist in the symbol-table yet, so we should be able to add it");
311
312        let expr = parse_expr("x + 1 : int", symbols).unwrap().unwrap();
313        let Expression::Sum(_, terms) = expr else {
314            panic!("expected a sum expression");
315        };
316        let terms = (*terms).clone().unwrap_list().unwrap();
317        assert_eq!(terms.len(), 2);
318        assert!(matches!(terms[1], Expression::DomainAnnotation(_, _, _)));
319        assert_eq!(terms[1].to_string(), "1 : int(-2147483647..2147483647)");
320
321        let symbols = SymbolTablePtr::new();
322        let x = DeclarationPtr::new_find(
323            Name::User("x".into()),
324            Domain::int(vec![conjure_cp_core::ast::Range::Bounded(0, 10)]),
325        );
326        symbols
327            .write()
328            .insert(x)
329            .expect("x should not exist in the symbol-table yet, so we should be able to add it");
330
331        let expr = parse_expr("(x + 1) : int", symbols).unwrap().unwrap();
332        let Expression::DomainAnnotation(_, inner, _) = expr else {
333            panic!("expected a domain annotation");
334        };
335        assert!(matches!(*inner, Expression::Sum(_, _)));
336    }
337
338    #[test]
339    pub fn test_parse_in_with_repr_annotation() {
340        let symbols = SymbolTablePtr::new();
341        let x = DeclarationPtr::new_find(
342            Name::User("x".into()),
343            Domain::set(
344                conjure_cp_core::ast::SetAttr::new_max_size(3),
345                Domain::int(vec![conjure_cp_core::ast::Range::Bounded(1, 4)]),
346            ),
347        );
348        symbols.write().insert(x).unwrap();
349
350        let expr = parse_expr(
351            "1 in x :: set (representation packed) of int",
352            symbols.clone(),
353        )
354        .unwrap()
355        .unwrap();
356        println!("no parens: {expr}");
357        assert!(
358            matches!(expr, Expression::In(_, _, _)),
359            "expected In, got {expr:?}"
360        );
361        let Expression::In(_, _, rhs) = &expr else {
362            unreachable!()
363        };
364        assert!(
365            matches!(rhs.as_ref(), Expression::TypeAnnotation(_, _, _)),
366            "expected type annotation on in-rhs, got {rhs:?}"
367        );
368
369        let expr = parse_expr("1 in (x :: set (representation packed) of int)", symbols)
370            .unwrap()
371            .unwrap();
372        println!("with parens: {expr}");
373        let Expression::In(_, _, rhs) = expr else {
374            panic!("expected In, got {expr:?}");
375        };
376        // Parentheses may wrap as Atomic-ish structure; accept TypeAnnotation directly or inside.
377        let rhs_str = rhs.to_string();
378        assert!(
379            rhs_str.contains("set (representation packed)")
380                || matches!(rhs.as_ref(), Expression::TypeAnnotation(_, _, _)),
381            "expected annotated set on rhs, got {rhs:?}"
382        );
383    }
384}