1use oximo_expr::Expr;
2
3use crate::set::{FromIndexKey, Set};
4
5#[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
22impl<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
50impl<K, D: SumDomain<K> + ?Sized> SumDomain<K> for &D {
53 fn keys(&self) -> impl Iterator<Item = K> + '_ {
54 (**self).keys()
55 }
56}
57
58impl 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#[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}