Skip to main content

conjure_cp_core/
domain_tightening.rs

1//! Infer tighter declaration domains from the constraint list, before rewriting.
2//!
3//! The first component looks at cardinality equalities `|s| = k` (and `|s| = |t|` when `|t|` is
4//! constant) and, when `s` is a sequence find, intersects that size into its domain. Indexed
5//! equalities `forAll i : D . |m[i]| = k_i` over a matrix of sequences are included so that each
6//! cell can later be represented as a fixed-length sequence.
7
8use 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
18/// Walk the model's constraints and tighten declaration domains where a fact can be proved.
19pub 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    /// `None` for a top-level sequence; `Some(indices)` for a matrix element.
27    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                // `forAll` parses as `And(Comprehension)`, not `And` of a list.
48                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    // First component: a single generator over a finite domain, body an equality of cards.
71    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        // Prefer the side that is *not* a known constant, so `|locs[i]| = |clues[i]|` tightens
150        // the find, not the given.
151        (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}