1use serde::{Deserialize, Serialize};
21
22use super::source::Source;
23use super::strategy::{StrategyName, ZipMode};
24
25#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
36#[serde(tag = "op", rename_all = "snake_case")]
37pub enum Comprehension {
38 Clause {
41 name: String,
43 source: Source,
45 },
46
47 Cartesian {
50 children: Vec<Comprehension>,
52 },
53
54 Zip {
57 children: Vec<Comprehension>,
59 mode: ZipMode,
61 },
62
63 Union {
67 children: Vec<Comprehension>,
69 },
70
71 Filter {
76 child: Box<Comprehension>,
78 predicate: String,
80 },
81
82 Order {
88 child: Box<Comprehension>,
90 strategy: StrategyName,
92 truncation: Option<u64>,
94 #[serde(default, skip_serializing_if = "Option::is_none")]
96 seed: Option<u64>,
97 },
98}
99
100impl Comprehension {
101 pub fn clause<S: Into<String>>(name: S, source: Source) -> Self {
103 Comprehension::Clause {
104 name: name.into(),
105 source,
106 }
107 }
108
109 pub fn cartesian(children: Vec<Comprehension>) -> Self {
111 Comprehension::Cartesian { children }
112 }
113
114 pub fn zip(children: Vec<Comprehension>, mode: ZipMode) -> Self {
117 Comprehension::Zip { children, mode }
118 }
119
120 pub fn union(children: Vec<Comprehension>) -> Self {
122 Comprehension::Union { children }
123 }
124
125 pub fn filter<S: Into<String>>(child: Comprehension, predicate: S) -> Self {
127 Comprehension::Filter {
128 child: Box::new(child),
129 predicate: predicate.into(),
130 }
131 }
132
133 pub fn order(child: Comprehension, strategy: StrategyName, truncation: Option<u64>) -> Self {
136 Self::order_seeded(child, strategy, truncation, None)
137 }
138
139 pub fn order_seeded(
142 child: Comprehension,
143 strategy: StrategyName,
144 truncation: Option<u64>,
145 seed: Option<u64>,
146 ) -> Self {
147 Comprehension::Order {
148 child: Box::new(child),
149 strategy,
150 truncation,
151 seed,
152 }
153 }
154
155 pub fn coordinate_names(&self) -> Vec<String> {
161 let mut acc = Vec::new();
162 self.collect_coordinate_names(&mut acc);
163 acc
164 }
165
166 pub fn coordinate_specs(&self) -> Vec<(String, String)> {
176 let mut acc = Vec::new();
177 let mut seen = std::collections::HashSet::new();
178 self.collect_coordinate_specs(&mut acc, &mut seen);
179 acc
180 }
181
182 pub fn referenced_source_names(&self) -> std::collections::BTreeSet<String> {
195 let mut out = std::collections::BTreeSet::new();
196 self.walk_sources(&mut |source| out.extend(source.referenced_names()));
197 out
198 }
199
200 fn walk_sources(&self, visit: &mut impl FnMut(&super::source::Source)) {
202 match self {
203 Comprehension::Clause { source, .. } => visit(source),
204 Comprehension::Cartesian { children }
205 | Comprehension::Zip { children, .. }
206 | Comprehension::Union { children } => {
207 for c in children {
208 c.walk_sources(visit);
209 }
210 }
211 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
212 child.walk_sources(visit);
213 }
214 }
215 }
216
217 fn collect_coordinate_specs(
218 &self,
219 acc: &mut Vec<(String, String)>,
220 seen: &mut std::collections::HashSet<String>,
221 ) {
222 match self {
223 Comprehension::Clause { name, source } => {
224 if seen.insert(name.clone()) {
225 let spec_text = source.to_text().unwrap_or_else(|| "<source>".to_string());
226 acc.push((name.clone(), spec_text));
227 }
228 }
229 Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
230 for c in children {
231 c.collect_coordinate_specs(acc, seen);
232 }
233 }
234 Comprehension::Union { children } => {
235 for c in children {
236 c.collect_coordinate_specs(acc, seen);
237 }
238 }
239 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
240 child.collect_coordinate_specs(acc, seen);
241 }
242 }
243 }
244
245 fn collect_coordinate_names(&self, acc: &mut Vec<String>) {
246 match self {
247 Comprehension::Clause { name, .. } => {
248 if !acc.contains(name) {
249 acc.push(name.clone());
250 }
251 }
252 Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
253 for c in children {
254 c.collect_coordinate_names(acc);
255 }
256 }
257 Comprehension::Union { children } => {
258 if let Some(first) = children.first() {
261 first.collect_coordinate_names(acc);
262 }
263 }
264 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
265 child.collect_coordinate_names(acc);
266 }
267 }
268 }
269
270 pub fn is_clause(&self) -> bool {
272 matches!(self, Comprehension::Clause { .. })
273 }
274
275 pub fn is_combinator(&self) -> bool {
277 matches!(
278 self,
279 Comprehension::Cartesian { .. }
280 | Comprehension::Zip { .. }
281 | Comprehension::Union { .. }
282 )
283 }
284
285 pub fn is_modifier(&self) -> bool {
287 matches!(
288 self,
289 Comprehension::Filter { .. } | Comprehension::Order { .. }
290 )
291 }
292
293 pub fn children(&self) -> Box<dyn Iterator<Item = &Comprehension> + '_> {
296 match self {
297 Comprehension::Clause { .. } => Box::new(std::iter::empty()),
298 Comprehension::Cartesian { children }
299 | Comprehension::Zip { children, .. }
300 | Comprehension::Union { children } => Box::new(children.iter()),
301 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
302 Box::new(std::iter::once(child.as_ref()))
303 }
304 }
305 }
306
307 pub fn node_count(&self) -> usize {
311 1 + self.children().map(|c| c.node_count()).sum::<usize>()
312 }
313
314 pub fn depth(&self) -> usize {
318 1 + self.children().map(|c| c.depth()).max().unwrap_or(0)
319 }
320}
321
322#[cfg(test)]
323mod tests {
324 use super::*;
325 use crate::comprehension::source::{LiteralValue, Source};
326
327 fn lit_int_clause(name: &str, values: &[i64]) -> Comprehension {
328 Comprehension::clause(
329 name,
330 Source::Literal {
331 values: values.iter().map(|n| LiteralValue::Int(*n)).collect(),
332 },
333 )
334 }
335
336 #[test]
337 fn clause_coordinates() {
338 let c = lit_int_clause("k", &[1, 2, 3]);
339 assert_eq!(c.coordinate_names(), vec!["k"]);
340 assert!(c.is_clause());
341 assert!(!c.is_combinator());
342 assert!(!c.is_modifier());
343 }
344
345 #[test]
346 fn continuous_interval_spec_text_round_trips_as_float() {
347 use crate::comprehension::cardinality::{Interval, ProductMeasure};
354 let c = Comprehension::clause(
355 "ef",
356 Source::ContinuousInterval {
357 interval: Interval {
358 lo: 1.0,
359 hi: 5.0,
360 lo_open: false,
361 hi_open: true,
362 },
363 measure: ProductMeasure::Uniform,
364 },
365 );
366 let (var, spec_text) = c.coordinate_specs().into_iter().next().unwrap();
367 assert_eq!(var, "ef");
368 let reparsed = crate::comprehension::spec::parse_source(&spec_text).unwrap();
371 assert!(
372 matches!(reparsed, Source::ContinuousInterval { .. }),
373 "reconstructed '{spec_text}' re-parsed to {reparsed:?}, expected ContinuousInterval"
374 );
375 }
376
377 #[test]
378 fn referenced_source_names_grammar_based() {
379 let bare = Comprehension::clause(
382 "eh",
383 Source::Generator {
384 expr: "eh_values".into(),
385 cardinality_hint: None,
386 },
387 );
388 let got: Vec<String> = bare.referenced_source_names().into_iter().collect();
389 assert_eq!(got, vec!["eh_values"]);
390
391 let call = Comprehension::clause(
395 "nbo",
396 Source::Generator {
397 expr: "concat(nbo_v_values)".into(),
398 cardinality_hint: None,
399 },
400 );
401 let got: Vec<String> = call.referenced_source_names().into_iter().collect();
402 assert_eq!(got, vec!["nbo_v_values"]);
403
404 let wpl = Comprehension::clause(
407 "p",
408 Source::WorkloadParamList {
409 name: "profiles".into(),
410 len_hint: None,
411 },
412 );
413 let got: Vec<String> = wpl.referenced_source_names().into_iter().collect();
414 assert_eq!(got, vec!["profiles"]);
415
416 let lit = lit_int_clause("k", &[1, 2, 3]);
418 assert!(lit.referenced_source_names().is_empty());
419
420 let cart = Comprehension::cartesian(vec![bare, call]);
422 let got: Vec<String> = cart.referenced_source_names().into_iter().collect();
423 assert_eq!(got, vec!["eh_values", "nbo_v_values"]);
424 }
425
426 #[test]
427 fn cartesian_coordinates_in_declaration_order() {
428 let c = Comprehension::cartesian(vec![
429 lit_int_clause("k", &[1, 2]),
430 lit_int_clause("limit", &[10, 20, 30]),
431 ]);
432 assert_eq!(c.coordinate_names(), vec!["k", "limit"]);
433 assert!(c.is_combinator());
434 }
435
436 #[test]
437 fn zip_coordinates() {
438 let c = Comprehension::zip(
439 vec![
440 lit_int_clause("x", &[1, 2, 3]),
441 lit_int_clause("y", &[10, 20, 30]),
442 ],
443 ZipMode::Strict,
444 );
445 assert_eq!(c.coordinate_names(), vec!["x", "y"]);
446 }
447
448 #[test]
449 fn union_takes_first_childs_shape() {
450 let a = Comprehension::cartesian(vec![
451 lit_int_clause("k", &[10]),
452 lit_int_clause("limit", &[10, 20]),
453 ]);
454 let b = Comprehension::cartesian(vec![
455 lit_int_clause("k", &[100]),
456 lit_int_clause("limit", &[100, 200]),
457 ]);
458 let u = Comprehension::union(vec![a, b]);
459 assert_eq!(u.coordinate_names(), vec!["k", "limit"]);
460 }
461
462 #[test]
463 fn filter_and_order_pass_through_coordinates() {
464 let inner = Comprehension::cartesian(vec![
465 lit_int_clause("k", &[1, 2]),
466 lit_int_clause("limit", &[10]),
467 ]);
468 let filtered = Comprehension::filter(inner.clone(), "{k} > 0");
469 assert_eq!(filtered.coordinate_names(), vec!["k", "limit"]);
470 assert!(filtered.is_modifier());
471
472 let ordered = Comprehension::order(inner, StrategyName::Lex, Some(5));
473 assert_eq!(ordered.coordinate_names(), vec!["k", "limit"]);
474 assert!(ordered.is_modifier());
475 }
476
477 #[test]
478 fn node_count_and_depth() {
479 let inner = Comprehension::cartesian(vec![
480 lit_int_clause("k", &[1, 2]),
481 lit_int_clause("limit", &[10]),
482 ]);
483 assert_eq!(inner.node_count(), 3);
485 assert_eq!(inner.depth(), 2);
486
487 let filtered = Comprehension::filter(inner, "{k} > 0");
488 assert_eq!(filtered.node_count(), 4);
490 assert_eq!(filtered.depth(), 3);
491 }
492
493 #[test]
494 fn round_trip_serde() {
495 let c = Comprehension::order(
496 Comprehension::filter(
497 Comprehension::cartesian(vec![
498 lit_int_clause("k", &[1, 2, 3]),
499 lit_int_clause("limit", &[10, 20]),
500 ]),
501 "{k} * {limit} > 5",
502 ),
503 StrategyName::Halton,
504 Some(10),
505 );
506 let json = serde_json::to_string(&c).unwrap();
507 let back: Comprehension = serde_json::from_str(&json).unwrap();
508 assert_eq!(c, back);
509 }
510
511 #[test]
514 fn an_orders_seed_round_trips_and_defaults_to_none() {
515 let c = Comprehension::order_seeded(
516 Comprehension::clause(
517 "k",
518 Source::IntRange {
519 lo: 1,
520 hi: 4,
521 step: 1,
522 },
523 ),
524 StrategyName::Shuffle,
525 Some(2),
526 Some(42),
527 );
528 let json = serde_json::to_string(&c).unwrap();
529 assert!(json.contains("\"seed\":42"), "{json}");
530 let back: Comprehension = serde_json::from_str(&json).unwrap();
531 assert_eq!(back, c);
532 let unseeded = json.replace(",\"seed\":42", "");
533 let back: Comprehension = serde_json::from_str(&unseeded).unwrap();
534 assert!(
535 matches!(back, Comprehension::Order { seed: None, .. }),
536 "{back:?}"
537 );
538 }
539}