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 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 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}