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