1use std::collections::HashMap;
9
10use crate::ast::{
11 Atom, DeclarationKind, DeclarationPtr, DomainPtr, Expression as Expr, GroundDomain, Literal,
12 Metadata, Model, Moo, Name, Range,
13 comprehension::{Comprehension, ComprehensionQualifier},
14 eval_constant,
15 matrix::shape_of_dom,
16};
17
18pub fn tighten_domains_from_constraints(model: &mut Model) {
20 let facts = collect_sequence_size_facts(model);
21 apply_sequence_size_facts(facts);
22}
23
24struct SequenceSizeFact {
25 decl: DeclarationPtr,
26 indices: Option<Vec<Literal>>,
28 size: i32,
29}
30
31fn collect_sequence_size_facts(model: &Model) -> Vec<SequenceSizeFact> {
32 let mut facts = Vec::new();
33 for constraint in model.constraints() {
34 collect_from_expr(constraint, &mut facts);
35 }
36 facts
37}
38
39fn collect_from_expr(expr: &Expr, facts: &mut Vec<SequenceSizeFact>) {
40 match expr {
41 Expr::And(_, inner) => {
42 if let Some(children) = inner.unwrap_list_ref() {
43 for child in children {
44 collect_from_expr(child, facts);
45 }
46 } else {
47 collect_from_expr(inner.as_ref(), facts);
49 }
50 }
51 Expr::Comprehension(_, comprehension) => {
52 collect_from_comprehension(comprehension, facts);
53 }
54 Expr::Eq(_, lhs, rhs) => collect_from_card_eq(lhs, rhs, facts),
55 _ => {}
56 }
57}
58
59fn collect_from_comprehension(comprehension: &Comprehension, facts: &mut Vec<SequenceSizeFact>) {
60 let generators: Vec<&DeclarationPtr> = comprehension
61 .qualifiers
62 .iter()
63 .filter_map(|qualifier| match qualifier {
64 ComprehensionQualifier::Generator { ptr } => Some(ptr),
65 ComprehensionQualifier::ExpressionGenerator { .. }
66 | ComprehensionQualifier::Condition(_) => None,
67 })
68 .collect();
69
70 if generators.len() != 1 {
72 collect_from_expr(&comprehension.return_expression, facts);
73 return;
74 }
75
76 let generator = generators[0];
77 let Some(domain) = generator.domain().and_then(|domain| domain.resolve().ok()) else {
78 collect_from_expr(&comprehension.return_expression, facts);
79 return;
80 };
81 let Ok(values) = domain.values() else {
82 collect_from_expr(&comprehension.return_expression, facts);
83 return;
84 };
85 let values: Vec<Literal> = values.collect();
86 if values.is_empty() {
87 return;
88 }
89
90 let Expr::Eq(_, lhs, rhs) = &comprehension.return_expression else {
91 collect_from_expr(&comprehension.return_expression, facts);
92 return;
93 };
94
95 for value in values {
96 let Some(size) = card_eq_size_at_generator(lhs, rhs, generator, &value) else {
97 continue;
98 };
99 let Some(subject) = card_eq_subject_at_generator(lhs, rhs, generator, &value) else {
100 continue;
101 };
102 facts.push(subject.with_size(size));
103 }
104}
105
106fn collect_from_card_eq(lhs: &Expr, rhs: &Expr, facts: &mut Vec<SequenceSizeFact>) {
107 let Some(size) = card_eq_constant_size(lhs, rhs) else {
108 return;
109 };
110 if let Some(subject) = sequence_subject(lhs).or_else(|| sequence_subject(rhs)) {
111 facts.push(subject.with_size(size));
112 }
113}
114
115fn card_eq_constant_size(lhs: &Expr, rhs: &Expr) -> Option<i32> {
116 match (as_card(lhs), as_card(rhs)) {
117 (Some(_), Some(_)) => eval_card(lhs).or_else(|| eval_card(rhs)),
118 (Some(_), None) => eval_int(rhs),
119 (None, Some(_)) => eval_int(lhs),
120 (None, None) => None,
121 }
122}
123
124fn card_eq_size_at_generator(
125 lhs: &Expr,
126 rhs: &Expr,
127 generator: &DeclarationPtr,
128 value: &Literal,
129) -> Option<i32> {
130 let lhs_size = eval_card_with_index(lhs, generator, value);
131 let rhs_size = eval_card_with_index(rhs, generator, value);
132 match (as_card(lhs), as_card(rhs), lhs_size, rhs_size) {
133 (Some(_), Some(_), Some(size), _) | (Some(_), Some(_), None, Some(size)) => Some(size),
134 (Some(_), None, _, _) => eval_int(rhs),
135 (None, Some(_), _, _) => eval_int(lhs),
136 _ => None,
137 }
138}
139
140fn card_eq_subject_at_generator(
141 lhs: &Expr,
142 rhs: &Expr,
143 generator: &DeclarationPtr,
144 value: &Literal,
145) -> Option<SequenceSubject> {
146 let lhs_eval = eval_card_with_index(lhs, generator, value).is_some();
147 let rhs_eval = eval_card_with_index(rhs, generator, value).is_some();
148 match (as_card(lhs), as_card(rhs), lhs_eval, rhs_eval) {
149 (Some(_), Some(_), false, true) => sequence_subject_at(lhs, generator, value),
152 (Some(_), Some(_), true, false) => sequence_subject_at(rhs, generator, value),
153 (Some(_), None, _, _) => sequence_subject_at(lhs, generator, value),
154 (None, Some(_), _, _) => sequence_subject_at(rhs, generator, value),
155 _ => None,
156 }
157}
158
159struct SequenceSubject {
160 decl: DeclarationPtr,
161 indices: Option<Vec<Literal>>,
162}
163
164impl SequenceSubject {
165 fn with_size(self, size: i32) -> SequenceSizeFact {
166 SequenceSizeFact {
167 decl: self.decl,
168 indices: self.indices,
169 size,
170 }
171 }
172}
173
174fn sequence_subject(expr: &Expr) -> Option<SequenceSubject> {
175 let collection = as_card(expr)?;
176 match collection {
177 Expr::Atomic(_, Atom::Reference(reference)) => Some(SequenceSubject {
178 decl: reference.ptr().clone(),
179 indices: None,
180 }),
181 Expr::UnsafeIndex(_, subject, indices) | Expr::SafeIndex(_, subject, indices) => {
182 let Expr::Atomic(_, Atom::Reference(reference)) = subject.as_ref() else {
183 return None;
184 };
185 let indices: Vec<Literal> = indices.iter().map(eval_constant).collect::<Option<_>>()?;
186 Some(SequenceSubject {
187 decl: reference.ptr().clone(),
188 indices: Some(indices),
189 })
190 }
191 _ => None,
192 }
193}
194
195fn sequence_subject_at(
196 expr: &Expr,
197 generator: &DeclarationPtr,
198 value: &Literal,
199) -> Option<SequenceSubject> {
200 let collection = as_card(expr)?;
201 match collection {
202 Expr::Atomic(_, Atom::Reference(reference)) => Some(SequenceSubject {
203 decl: reference.ptr().clone(),
204 indices: None,
205 }),
206 Expr::UnsafeIndex(_, subject, indices) | Expr::SafeIndex(_, subject, indices) => {
207 let Expr::Atomic(_, Atom::Reference(reference)) = subject.as_ref() else {
208 return None;
209 };
210 let indices: Vec<Literal> = indices
211 .iter()
212 .map(|index| eval_index_at(index, generator, value))
213 .collect::<Option<_>>()?;
214 Some(SequenceSubject {
215 decl: reference.ptr().clone(),
216 indices: Some(indices),
217 })
218 }
219 _ => None,
220 }
221}
222
223fn as_card(expr: &Expr) -> Option<&Expr> {
224 match expr {
225 Expr::Card(_, collection) => Some(collection.as_ref()),
226 _ => None,
227 }
228}
229
230fn eval_card(expr: &Expr) -> Option<i32> {
231 match eval_constant(expr)? {
232 Literal::Int(size) if size >= 0 => Some(size),
233 _ => None,
234 }
235}
236
237fn eval_int(expr: &Expr) -> Option<i32> {
238 match eval_constant(expr)? {
239 Literal::Int(size) if size >= 0 => Some(size),
240 _ => None,
241 }
242}
243
244fn eval_card_with_index(expr: &Expr, generator: &DeclarationPtr, value: &Literal) -> Option<i32> {
245 let collection = as_card(expr)?;
246 let instantiated = instantiate_index_expr(collection, generator, value)?;
247 eval_card(&Expr::Card(Metadata::new(), Moo::new(instantiated)))
248}
249
250fn instantiate_index_expr(
251 expr: &Expr,
252 generator: &DeclarationPtr,
253 value: &Literal,
254) -> Option<Expr> {
255 match expr {
256 Expr::UnsafeIndex(meta, subject, indices) => {
257 let indices = instantiate_indices(indices, generator, value)?;
258 Some(Expr::UnsafeIndex(meta.clone(), subject.clone(), indices))
259 }
260 Expr::SafeIndex(meta, subject, indices) => {
261 let indices = instantiate_indices(indices, generator, value)?;
262 Some(Expr::SafeIndex(meta.clone(), subject.clone(), indices))
263 }
264 _ => None,
265 }
266}
267
268fn instantiate_indices(
269 indices: &[Expr],
270 generator: &DeclarationPtr,
271 value: &Literal,
272) -> Option<Vec<Expr>> {
273 indices
274 .iter()
275 .map(|index| {
276 eval_index_at(index, generator, value)
277 .map(|lit| Expr::Atomic(Metadata::new(), Atom::Literal(lit)))
278 })
279 .collect()
280}
281
282fn eval_index_at(index: &Expr, generator: &DeclarationPtr, value: &Literal) -> Option<Literal> {
283 if let Expr::Atomic(_, Atom::Reference(reference)) = index
284 && reference.ptr() == generator
285 {
286 return Some(value.clone());
287 }
288 eval_constant(index)
289}
290
291fn apply_sequence_size_facts(facts: Vec<SequenceSizeFact>) {
292 let mut by_decl: HashMap<Name, (DeclarationPtr, Vec<SequenceSizeFact>)> = HashMap::new();
293 for fact in facts {
294 if fact.size < 0 {
295 continue;
296 }
297 let name = fact.decl.name().clone();
298 by_decl
299 .entry(name)
300 .or_insert_with(|| (fact.decl.clone(), Vec::new()))
301 .1
302 .push(fact);
303 }
304
305 for (decl, decl_facts) in by_decl.into_values() {
306 apply_facts_to_declaration(decl, decl_facts);
307 }
308}
309
310fn apply_facts_to_declaration(mut decl: DeclarationPtr, facts: Vec<SequenceSizeFact>) {
311 if !matches!(
312 &*decl.kind(),
313 DeclarationKind::Find(_) | DeclarationKind::FindAuxiliary(_)
314 ) {
315 return;
316 }
317
318 let Some(domain) = decl.domain() else {
319 return;
320 };
321 let Ok(ground) = domain.resolve() else {
322 return;
323 };
324
325 let top_level: Vec<i32> = facts
326 .iter()
327 .filter(|fact| fact.indices.is_none())
328 .map(|fact| fact.size)
329 .collect();
330 let indexed: Vec<(Vec<Literal>, i32)> = facts
331 .iter()
332 .filter_map(|fact| {
333 fact.indices
334 .as_ref()
335 .map(|indices| (indices.clone(), fact.size))
336 })
337 .collect();
338
339 if let GroundDomain::Sequence(attr, _) = ground.as_ref()
340 && let Some(&size) = top_level.first()
341 && top_level.iter().all(|&other| other == size)
342 && let Some(new_size) = merge_to_size(&attr.size, size)
343 && let Some(new_domain) = with_sequence_size(ground.as_ref(), new_size)
344 && let Some(mut var) = decl.as_find_mut()
345 {
346 var.domain = new_domain.into();
347 return;
348 }
349
350 let GroundDomain::Matrix(inner, _) = ground.as_ref() else {
351 return;
352 };
353 let GroundDomain::Sequence(_, _) = inner.as_ref() else {
354 return;
355 };
356 if indexed.is_empty() {
357 return;
358 }
359
360 let Ok(shape) = shape_of_dom(ground.as_ref()) else {
361 return;
362 };
363 let Some(index_tuples) = index_tuples(&shape.idx_doms) else {
364 return;
365 };
366 if index_tuples.len() != shape.size {
367 return;
368 }
369
370 let mut size_by_index: HashMap<Vec<Literal>, i32> = HashMap::new();
371 for (indices, size) in indexed {
372 if let Some(existing) = size_by_index.get(&indices)
373 && *existing != size
374 {
375 return;
376 }
377 size_by_index.insert(indices, size);
378 }
379 if size_by_index.len() != index_tuples.len() {
380 return;
381 }
382
383 let sizes: Vec<i32> = index_tuples
384 .iter()
385 .map(|indices| size_by_index[indices])
386 .collect();
387 let min_size = *sizes.iter().min().unwrap_or(&0);
388 let max_size = *sizes.iter().max().unwrap_or(&0);
389 let Some(new_inner_size) = merge_to_bounds(
390 match inner.as_ref() {
391 GroundDomain::Sequence(attr, _) => &attr.size,
392 _ => return,
393 },
394 min_size,
395 max_size,
396 ) else {
397 return;
398 };
399
400 let element_domains: Option<Vec<DomainPtr>> = sizes
401 .iter()
402 .copied()
403 .map(|size| {
404 let size_range = merge_to_size(
405 match inner.as_ref() {
406 GroundDomain::Sequence(attr, _) => &attr.size,
407 _ => return None,
408 },
409 size,
410 )?;
411 Some(with_sequence_size(inner.as_ref(), size_range)?.into())
412 })
413 .collect();
414 let Some(element_domains) = element_domains else {
415 return;
416 };
417
418 let Some(new_inner) = with_sequence_size(inner.as_ref(), new_inner_size) else {
419 return;
420 };
421 let GroundDomain::Matrix(_, idx_doms) = ground.as_ref() else {
422 return;
423 };
424 let new_domain = GroundDomain::Matrix(Moo::new(new_inner), idx_doms.clone());
425 if let Some(mut var) = decl.as_find_mut() {
426 var.domain = new_domain.into();
427 var.element_domains = Some(element_domains);
428 }
429}
430
431fn with_sequence_size(domain: &GroundDomain, size: Range<i32>) -> Option<GroundDomain> {
432 let GroundDomain::Sequence(attr, inner) = domain else {
433 return None;
434 };
435 let mut attr = attr.clone();
436 attr.size = size;
437 Some(GroundDomain::Sequence(attr, inner.clone()))
438}
439
440fn merge_to_size(current: &Range<i32>, size: i32) -> Option<Range<i32>> {
441 current.contains(&size).then_some(Range::Single(size))
442}
443
444fn merge_to_bounds(current: &Range<i32>, min: i32, max: i32) -> Option<Range<i32>> {
445 let lo = current.low().copied().unwrap_or(min).max(min);
446 let hi = current.high().copied().unwrap_or(max).min(max);
447 if lo > hi {
448 return None;
449 }
450 Some(Range::new(Some(lo), Some(hi)))
451}
452
453fn index_tuples(idx_doms: &[Moo<GroundDomain>]) -> Option<Vec<Vec<Literal>>> {
454 let lists: Vec<Vec<Literal>> = idx_doms
455 .iter()
456 .map(|domain| domain.values().ok().map(Iterator::collect))
457 .collect::<Option<_>>()?;
458 Some(lists.into_iter().fold(vec![vec![]], |prefixes, list| {
459 prefixes
460 .into_iter()
461 .flat_map(|prefix| {
462 list.iter().map(move |item| {
463 let mut next = prefix.clone();
464 next.push(item.clone());
465 next
466 })
467 })
468 .collect()
469 }))
470}
471
472#[cfg(test)]
473mod tests {
474 use super::*;
475 use crate::ast::{
476 AbstractLiteral, DeclarationPtr, Domain, JectivityAttr, Name, Reference, SequenceAttr,
477 SymbolTablePtr, comprehension::ComprehensionBuilder,
478 };
479 use crate::{domain_int, matrix_expr, range};
480
481 fn sequence_attr(size: Range<i32>) -> SequenceAttr {
482 SequenceAttr {
483 size,
484 jectivity: JectivityAttr::None,
485 representation: None,
486 }
487 }
488
489 fn int_lit(value: i32) -> Expr {
490 Expr::Atomic(Metadata::new(), Atom::Literal(Literal::Int(value)))
491 }
492
493 fn ref_expr(decl: &DeclarationPtr) -> Expr {
494 Expr::Atomic(
495 Metadata::new(),
496 Atom::Reference(Reference::new(decl.clone())),
497 )
498 }
499
500 fn card(expr: Expr) -> Expr {
501 Expr::Card(Metadata::new(), Moo::new(expr))
502 }
503
504 fn eq(lhs: Expr, rhs: Expr) -> Expr {
505 Expr::Eq(Metadata::new(), Moo::new(lhs), Moo::new(rhs))
506 }
507
508 fn sequence_size(decl: &DeclarationPtr) -> Range<i32> {
509 let domain = decl.domain().unwrap();
510 let GroundDomain::Sequence(attr, _) = domain.as_ground().unwrap() else {
511 panic!("expected a sequence domain");
512 };
513 attr.size.clone()
514 }
515
516 #[test]
517 fn tightens_a_top_level_sequence_from_a_cardinality_equality() {
518 let mut model = Model::new(Default::default());
519 let s = DeclarationPtr::new_find(
520 Name::user("s"),
521 Domain::sequence(sequence_attr(range!(0..10)), domain_int!(1..3)),
522 );
523 model.symbols_mut().insert(s.clone()).unwrap();
524 model.add_constraint(eq(card(ref_expr(&s)), int_lit(4)));
525
526 tighten_domains_from_constraints(&mut model);
527
528 let s = model.symbols().lookup_local(&Name::user("s")).unwrap();
529 assert_eq!(sequence_size(&s), Range::Single(4));
530 }
531
532 #[test]
533 fn leaves_the_domain_alone_when_the_size_is_outside_the_declared_range() {
534 let mut model = Model::new(Default::default());
535 let s = DeclarationPtr::new_find(
536 Name::user("s"),
537 Domain::sequence(sequence_attr(range!(0..3)), domain_int!(1..3)),
538 );
539 model.symbols_mut().insert(s.clone()).unwrap();
540 model.add_constraint(eq(card(ref_expr(&s)), int_lit(4)));
541
542 tighten_domains_from_constraints(&mut model);
543
544 let s = model.symbols().lookup_local(&Name::user("s")).unwrap();
545 assert_eq!(sequence_size(&s), range!(0..3));
546 }
547
548 #[test]
549 fn tightens_matrix_of_sequence_cells_from_a_forall_cardinality_equality() {
550 let mut model = Model::new(Default::default());
551 let clues_lit = Literal::AbstractLiteral(AbstractLiteral::Matrix(
552 vec![
553 Literal::AbstractLiteral(AbstractLiteral::Sequence(vec![
554 Literal::Int(1),
555 Literal::Int(2),
556 Literal::Int(3),
557 ])),
558 Literal::AbstractLiteral(AbstractLiteral::Sequence(vec![
559 Literal::Int(1),
560 Literal::Int(1),
561 Literal::Int(1),
562 Literal::Int(1),
563 Literal::Int(1),
564 ])),
565 ],
566 Moo::new(GroundDomain::Int(vec![range!(1..2)])),
567 ));
568 let clues = DeclarationPtr::new_value_letting(Name::user("clues"), Expr::from(clues_lit));
569 let locs = DeclarationPtr::new_find(
570 Name::user("locs"),
571 Domain::matrix(
572 Domain::sequence(sequence_attr(range!(0..10)), domain_int!(1..9)),
573 vec![domain_int!(1..2)],
574 ),
575 );
576 model.symbols_mut().insert(clues.clone()).unwrap();
577 model.symbols_mut().insert(locs.clone()).unwrap();
578
579 let builder = ComprehensionBuilder::new(SymbolTablePtr::new());
580 let i_template = DeclarationPtr::new_find(Name::user("i"), domain_int!(1..2));
581 let mut builder = builder.generator(i_template);
582 let i = builder
583 .generator_symboltable()
584 .read()
585 .lookup_local(&Name::user("i"))
586 .expect("generator i");
587 let body = eq(
588 card(Expr::UnsafeIndex(
589 Metadata::new(),
590 Moo::new(ref_expr(&locs)),
591 vec![ref_expr(&i)],
592 )),
593 card(Expr::UnsafeIndex(
594 Metadata::new(),
595 Moo::new(ref_expr(&clues)),
596 vec![ref_expr(&i)],
597 )),
598 );
599 let comprehension = builder.with_return_value(body);
600 model.add_constraint(Expr::And(
601 Metadata::new(),
602 Moo::new(matrix_expr![Expr::Comprehension(
603 Metadata::new(),
604 Moo::new(comprehension)
605 )]),
606 ));
607
608 tighten_domains_from_constraints(&mut model);
609
610 let locs = model.symbols().lookup_local(&Name::user("locs")).unwrap();
611 let var = locs.as_find().unwrap();
612 let domains = var.element_domains.as_ref().expect("per-cell domains");
613 assert_eq!(domains.len(), 2);
614 let sizes: Vec<Range<i32>> = domains
615 .iter()
616 .map(|domain| {
617 let GroundDomain::Sequence(attr, _) = domain.as_ground().unwrap() else {
618 panic!("expected a sequence element domain");
619 };
620 attr.size.clone()
621 })
622 .collect();
623 assert_eq!(sizes, vec![Range::Single(3), Range::Single(5)]);
624 }
625}