Skip to main content

oximo_core/
sum.rs

1use oximo_expr::Expr;
2
3use crate::set::{FromIndexKey, Set};
4
5/// Domain over which the `sum!` macro iterates. Lets a single sum accept either
6/// a [`Set`] (with typed key decoding via [`FromIndexKey`]) or a borrowed slice
7/// of `Copy` keys, without intermediate conversions.
8///
9/// Returns an iterator (rather than taking a callback) so the trait method
10/// monomorphizes through to the loop body in [`__sum_over`], allowing inlining
11/// in hot sums. Implementations are typically one line.
12#[diagnostic::on_unimplemented(
13    message = "`{Self}` is not a valid index domain over key type `{K}`",
14    label = "the domain's keys are not `{K}`",
15    note = "the loop/closure binding type must match the domain's keys",
16    note = "for a `Set<T>` write `for x in set` (the key type is inferred) or annotate `for x: T in set`. Integer ranges yield `usize`/`i64`/`i32`. A slice/`Vec`/array yields its element type"
17)]
18pub trait SumDomain<K> {
19    fn keys(&self) -> impl Iterator<Item = K> + '_;
20}
21
22// A typed set yields exactly its own key type.
23// The single `SumDomain` impl for `Set<K>`, so `sum!`/`constraint!`
24// can infer the closure parameter without an annotation
25// (the erased `Set` defaulted to `Set<IndexKey>`).
26impl<K: FromIndexKey> SumDomain<K> for Set<K> {
27    fn keys(&self) -> impl Iterator<Item = K> + '_ {
28        self.iter().map(|k| K::from_index_key(&k))
29    }
30}
31
32impl<K: Copy> SumDomain<K> for [K] {
33    fn keys(&self) -> impl Iterator<Item = K> + '_ {
34        self.iter().copied()
35    }
36}
37
38impl<K: Copy> SumDomain<K> for Vec<K> {
39    fn keys(&self) -> impl Iterator<Item = K> + '_ {
40        self.iter().copied()
41    }
42}
43
44impl<K: Copy, const N: usize> SumDomain<K> for [K; N] {
45    fn keys(&self) -> impl Iterator<Item = K> + '_ {
46        self.iter().copied()
47    }
48}
49
50// Forward through a reference, so a domain that is itself a reference (e.g. a
51// `&Set` function parameter passed to `sum!`/`constraint!`) is accepted.
52impl<K, D: SumDomain<K> + ?Sized> SumDomain<K> for &D {
53    fn keys(&self) -> impl Iterator<Item = K> + '_ {
54        (**self).keys()
55    }
56}
57
58// Integer ranges as sum domains. Iteration is lazy, so `sum!(x[i] for i in 0..n)`
59// allocates nothing. Provided for the common integer types the `sum!`/`constraint!`
60// macros default to.
61impl SumDomain<usize> for std::ops::Range<usize> {
62    fn keys(&self) -> impl Iterator<Item = usize> + '_ {
63        self.clone()
64    }
65}
66
67impl SumDomain<i64> for std::ops::Range<i64> {
68    fn keys(&self) -> impl Iterator<Item = i64> + '_ {
69        self.clone()
70    }
71}
72
73impl SumDomain<i32> for std::ops::Range<i32> {
74    fn keys(&self) -> impl Iterator<Item = i32> + '_ {
75        self.clone()
76    }
77}
78
79/// Sum an expression over every element of a domain.
80///
81/// Reads as the mathematical `sum_{k in domain} f(k)`. The closure parameter is
82/// either decoded from the domain's [`crate::set::IndexKey`] via [`FromIndexKey`] (when
83/// the domain is a [`Set`]) or yielded directly (when the domain is a slice
84/// of `Copy` keys).
85///
86/// # Panics
87/// Panics if `domain` is empty, because there is no generated expression from
88/// which to recover an arena. Use the anchored macro form
89/// `sum!(model, body for key in domain)` when the domain may be empty.
90/// Macro-facing entry point backing the `sum!` macro. Not part of the stable
91/// public API.
92#[doc(hidden)]
93pub fn __sum_over<'a, K, D, F>(domain: &D, f: F) -> Expr<'a>
94where
95    D: SumDomain<K> + ?Sized,
96    F: FnMut(K) -> Expr<'a>,
97{
98    Expr::__sum_terms(domain.keys().map(f)).expect("sum_over on empty domain")
99}
100
101#[cfg(test)]
102mod tests {
103    use oximo_expr::extract_linear;
104
105    use super::*;
106    use crate::model::Model;
107
108    #[test]
109    fn sum_over_finishes_callbacks_before_rejecting_foreign_arena() {
110        use std::panic::{AssertUnwindSafe, catch_unwind};
111
112        let model = Model::new("local");
113        let other = Model::new("foreign");
114        let x = model.__var("x").build();
115        let y = other.__var("y").build();
116        let mut visited = Vec::new();
117        let before = model.arena().len();
118        let result = catch_unwind(AssertUnwindSafe(|| {
119            __sum_over(&(0..4_usize), |i| {
120                visited.push(i);
121                if i == 1 { y } else { x }
122            })
123        }));
124        assert!(result.is_err());
125        assert_eq!(visited, [0, 1, 2, 3]);
126        assert_eq!(model.arena().len(), before);
127    }
128
129    #[test]
130    fn sum_over_keeps_child_order_for_inline_and_spilled_sums() {
131        let model = Model::new("sum_order");
132        let keys = Set::range(0..8_usize);
133        let x = model.__indexed_var("x", &keys).build();
134        for count in [1_usize, 4, 8] {
135            let total = __sum_over(&(0..count), |i| x[count - i - 1]);
136            let arena = model.arena();
137            if count == 1 {
138                assert_eq!(total.id, x[0_usize].id);
139            } else {
140                let oximo_expr::ExprNode::Add(children) = arena.get(total.id) else {
141                    panic!("sum must keep its flat Add representation");
142                };
143                assert_eq!(
144                    children.as_slice(),
145                    &(0..count).rev().map(|i| x[i].id).collect::<Vec<_>>()
146                );
147            }
148        }
149    }
150
151    #[test]
152    fn sum_over_scalar_set() {
153        let m = Model::new("scalar");
154        let items = Set::range(0..4);
155        let x = m.__indexed_var("x", &items).lb(0.0).build();
156
157        let total = __sum_over(&items, |i: usize| x[i]);
158        let arena = m.arena();
159        let terms = extract_linear(&arena, total.id).expect("linear");
160        assert_eq!(terms.coeffs.len(), 4);
161        assert!(terms.coeffs.iter().all(|(_, c)| (c - 1.0).abs() < f64::EPSILON));
162    }
163
164    #[test]
165    fn sum_over_tuple_set() {
166        let m = Model::new("tuple");
167        let plants = Set::strings(["seattle", "san-diego"]);
168        let markets = Set::strings(["nyc", "chicago", "topeka"]);
169        let routes = &plants * &markets;
170        let x = m.__indexed_var("x", &routes).lb(0.0).build();
171
172        let total = __sum_over(&routes, |(p, q): (String, String)| x[(p, q)]);
173        let arena = m.arena();
174        let terms = extract_linear(&arena, total.id).expect("linear");
175        assert_eq!(terms.coeffs.len(), 6);
176    }
177
178    #[test]
179    fn nested_sum_over_double_sum() {
180        let m = Model::new("nested");
181        let plants = Set::strings(["a", "b"]);
182        let markets = Set::strings(["x", "y", "z"]);
183        let routes = &plants * &markets;
184        let x = m.__indexed_var("x", &routes).lb(0.0).build();
185
186        let total = __sum_over(&plants, |p: String| __sum_over(&markets, |q: String| x[(&p, q)]));
187        let arena = m.arena();
188        let terms = extract_linear(&arena, total.id).expect("linear");
189        assert_eq!(terms.coeffs.len(), 6);
190    }
191
192    #[test]
193    fn sum_over_passes_typed_usize_key() {
194        let m = Model::new("usizekey");
195        let items = Set::range(0..3);
196        let x = m.__indexed_var("x", &items).lb(0.0).build();
197
198        let total = __sum_over(&items, |i: usize| x[i]);
199        let arena = m.arena();
200        let terms = extract_linear(&arena, total.id).expect("linear");
201        assert_eq!(terms.coeffs.len(), 3);
202    }
203
204    #[test]
205    fn sum_over_slice_of_usize() {
206        let m = Model::new("slice");
207        let items = Set::range(0..5);
208        let x = m.__indexed_var("x", &items).lb(0.0).build();
209
210        let picked: &[usize] = &[0, 2, 4];
211        let total = __sum_over(picked, |i: usize| x[i]);
212        let arena = m.arena();
213        let terms = extract_linear(&arena, total.id).expect("linear");
214        assert_eq!(terms.coeffs.len(), 3);
215    }
216
217    #[test]
218    fn sum_over_vec_of_usize() {
219        let m = Model::new("vec");
220        let items = Set::range(0..5);
221        let x = m.__indexed_var("x", &items).lb(0.0).build();
222
223        let picked: Vec<usize> = vec![1, 3];
224        let total = __sum_over(&picked, |i: usize| x[i]);
225        let arena = m.arena();
226        let terms = extract_linear(&arena, total.id).expect("linear");
227        assert_eq!(terms.coeffs.len(), 2);
228    }
229
230    #[test]
231    fn sum_over_array_of_usize() {
232        let m = Model::new("array");
233        let items = Set::range(0..5);
234        let x = m.__indexed_var("x", &items).lb(0.0).build();
235
236        let picked: [usize; 4] = [0, 1, 2, 3];
237        let total = __sum_over(&picked, |i: usize| x[i]);
238        let arena = m.arena();
239        let terms = extract_linear(&arena, total.id).expect("linear");
240        assert_eq!(terms.coeffs.len(), 4);
241    }
242
243    #[test]
244    #[should_panic(expected = "sum_over on empty domain")]
245    fn sum_over_empty_set_panics() {
246        let m = Model::new("empty");
247        let empty = Set::range(0..0);
248        let _x = m.__indexed_var("x", &Set::range(0..1)).lb(0.0).build();
249        let _ = __sum_over(&empty, |_: usize| panic!("closure should not run"));
250    }
251
252    #[test]
253    #[should_panic(expected = "sum_over on empty domain")]
254    fn sum_over_empty_slice_panics() {
255        let m = Model::new("empty_slice");
256        let _x = m.__indexed_var("x", &Set::range(0..1)).lb(0.0).build();
257        let empty: &[usize] = &[];
258        let _ = __sum_over(empty, |_: usize| panic!("closure should not run"));
259    }
260}