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 pub fn source_names_read(&self) -> std::collections::BTreeSet<String> {
205 let mut out = std::collections::BTreeSet::new();
206 self.walk_sources(&mut |source| out.extend(source.names_read()));
207 out
208 }
209
210 fn walk_sources(&self, visit: &mut impl FnMut(&super::source::Source)) {
212 match self {
213 Comprehension::Clause { source, .. } => visit(source),
214 Comprehension::Cartesian { children }
215 | Comprehension::Zip { children, .. }
216 | Comprehension::Union { children } => {
217 for c in children {
218 c.walk_sources(visit);
219 }
220 }
221 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
222 child.walk_sources(visit);
223 }
224 }
225 }
226
227 fn collect_coordinate_specs(
228 &self,
229 acc: &mut Vec<(String, String)>,
230 seen: &mut std::collections::HashSet<String>,
231 ) {
232 match self {
233 Comprehension::Clause { name, source } => {
234 if seen.insert(name.clone()) {
235 let spec_text = source.to_text().unwrap_or_else(|| "<source>".to_string());
236 acc.push((name.clone(), spec_text));
237 }
238 }
239 Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
240 for c in children {
241 c.collect_coordinate_specs(acc, seen);
242 }
243 }
244 Comprehension::Union { children } => {
245 for c in children {
246 c.collect_coordinate_specs(acc, seen);
247 }
248 }
249 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
250 child.collect_coordinate_specs(acc, seen);
251 }
252 }
253 }
254
255 fn collect_coordinate_names(&self, acc: &mut Vec<String>) {
256 match self {
257 Comprehension::Clause { name, .. } => {
258 if !acc.contains(name) {
259 acc.push(name.clone());
260 }
261 }
262 Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
263 for c in children {
264 c.collect_coordinate_names(acc);
265 }
266 }
267 Comprehension::Union { children } => {
268 if let Some(first) = children.first() {
271 first.collect_coordinate_names(acc);
272 }
273 }
274 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
275 child.collect_coordinate_names(acc);
276 }
277 }
278 }
279
280 pub fn is_clause(&self) -> bool {
282 matches!(self, Comprehension::Clause { .. })
283 }
284
285 pub fn is_combinator(&self) -> bool {
287 matches!(
288 self,
289 Comprehension::Cartesian { .. }
290 | Comprehension::Zip { .. }
291 | Comprehension::Union { .. }
292 )
293 }
294
295 pub fn is_modifier(&self) -> bool {
297 matches!(
298 self,
299 Comprehension::Filter { .. } | Comprehension::Order { .. }
300 )
301 }
302
303 pub fn children(&self) -> Box<dyn Iterator<Item = &Comprehension> + '_> {
306 match self {
307 Comprehension::Clause { .. } => Box::new(std::iter::empty()),
308 Comprehension::Cartesian { children }
309 | Comprehension::Zip { children, .. }
310 | Comprehension::Union { children } => Box::new(children.iter()),
311 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
312 Box::new(std::iter::once(child.as_ref()))
313 }
314 }
315 }
316
317 pub fn node_count(&self) -> usize {
321 1 + self.children().map(|c| c.node_count()).sum::<usize>()
322 }
323
324 pub fn depth(&self) -> usize {
328 1 + self.children().map(|c| c.depth()).max().unwrap_or(0)
329 }
330}
331
332#[cfg(test)]
333mod tests {
334 use super::*;
335 use crate::comprehension::source::{LiteralValue, Source};
336
337 fn lit_int_clause(name: &str, values: &[i64]) -> Comprehension {
338 Comprehension::clause(
339 name,
340 Source::Literal {
341 values: values.iter().map(|n| LiteralValue::Int(*n)).collect(),
342 },
343 )
344 }
345
346 #[test]
347 fn clause_coordinates() {
348 let c = lit_int_clause("k", &[1, 2, 3]);
349 assert_eq!(c.coordinate_names(), vec!["k"]);
350 assert!(c.is_clause());
351 assert!(!c.is_combinator());
352 assert!(!c.is_modifier());
353 }
354
355 #[test]
356 fn continuous_interval_spec_text_round_trips_as_float() {
357 use crate::comprehension::cardinality::{Interval, ProductMeasure};
364 let c = Comprehension::clause(
365 "ef",
366 Source::ContinuousInterval {
367 interval: Interval {
368 lo: 1.0,
369 hi: 5.0,
370 lo_open: false,
371 hi_open: true,
372 },
373 measure: ProductMeasure::Uniform,
374 },
375 );
376 let (var, spec_text) = c.coordinate_specs().into_iter().next().unwrap();
377 assert_eq!(var, "ef");
378 let reparsed = crate::comprehension::spec::parse_source(&spec_text).unwrap();
381 assert!(
382 matches!(reparsed, Source::ContinuousInterval { .. }),
383 "reconstructed '{spec_text}' re-parsed to {reparsed:?}, expected ContinuousInterval"
384 );
385 }
386
387 #[test]
388 fn referenced_source_names_grammar_based() {
389 let bare = Comprehension::clause(
392 "eh",
393 Source::Generator {
394 expr: "eh_values".into(),
395 cardinality_hint: None,
396 },
397 );
398 let got: Vec<String> = bare.referenced_source_names().into_iter().collect();
399 assert_eq!(got, vec!["eh_values"]);
400
401 let call = Comprehension::clause(
405 "nbo",
406 Source::Generator {
407 expr: "concat(nbo_v_values)".into(),
408 cardinality_hint: None,
409 },
410 );
411 let got: Vec<String> = call.referenced_source_names().into_iter().collect();
412 assert_eq!(got, vec!["nbo_v_values"]);
413
414 let wpl = Comprehension::clause(
417 "p",
418 Source::WorkloadParamList {
419 name: "profiles".into(),
420 len_hint: None,
421 },
422 );
423 let got: Vec<String> = wpl.referenced_source_names().into_iter().collect();
424 assert_eq!(got, vec!["profiles"]);
425
426 let lit = lit_int_clause("k", &[1, 2, 3]);
428 assert!(lit.referenced_source_names().is_empty());
429
430 let cart = Comprehension::cartesian(vec![bare, call]);
432 let got: Vec<String> = cart.referenced_source_names().into_iter().collect();
433 assert_eq!(got, vec!["eh_values", "nbo_v_values"]);
434 }
435
436 #[test]
437 fn cartesian_coordinates_in_declaration_order() {
438 let c = Comprehension::cartesian(vec![
439 lit_int_clause("k", &[1, 2]),
440 lit_int_clause("limit", &[10, 20, 30]),
441 ]);
442 assert_eq!(c.coordinate_names(), vec!["k", "limit"]);
443 assert!(c.is_combinator());
444 }
445
446 #[test]
447 fn zip_coordinates() {
448 let c = Comprehension::zip(
449 vec![
450 lit_int_clause("x", &[1, 2, 3]),
451 lit_int_clause("y", &[10, 20, 30]),
452 ],
453 ZipMode::Strict,
454 );
455 assert_eq!(c.coordinate_names(), vec!["x", "y"]);
456 }
457
458 #[test]
459 fn union_takes_first_childs_shape() {
460 let a = Comprehension::cartesian(vec![
461 lit_int_clause("k", &[10]),
462 lit_int_clause("limit", &[10, 20]),
463 ]);
464 let b = Comprehension::cartesian(vec![
465 lit_int_clause("k", &[100]),
466 lit_int_clause("limit", &[100, 200]),
467 ]);
468 let u = Comprehension::union(vec![a, b]);
469 assert_eq!(u.coordinate_names(), vec!["k", "limit"]);
470 }
471
472 #[test]
473 fn filter_and_order_pass_through_coordinates() {
474 let inner = Comprehension::cartesian(vec![
475 lit_int_clause("k", &[1, 2]),
476 lit_int_clause("limit", &[10]),
477 ]);
478 let filtered = Comprehension::filter(inner.clone(), "{k} > 0");
479 assert_eq!(filtered.coordinate_names(), vec!["k", "limit"]);
480 assert!(filtered.is_modifier());
481
482 let ordered = Comprehension::order(inner, StrategyName::Lex, Some(5));
483 assert_eq!(ordered.coordinate_names(), vec!["k", "limit"]);
484 assert!(ordered.is_modifier());
485 }
486
487 #[test]
488 fn node_count_and_depth() {
489 let inner = Comprehension::cartesian(vec![
490 lit_int_clause("k", &[1, 2]),
491 lit_int_clause("limit", &[10]),
492 ]);
493 assert_eq!(inner.node_count(), 3);
495 assert_eq!(inner.depth(), 2);
496
497 let filtered = Comprehension::filter(inner, "{k} > 0");
498 assert_eq!(filtered.node_count(), 4);
500 assert_eq!(filtered.depth(), 3);
501 }
502
503 #[test]
504 fn round_trip_serde() {
505 let c = Comprehension::order(
506 Comprehension::filter(
507 Comprehension::cartesian(vec![
508 lit_int_clause("k", &[1, 2, 3]),
509 lit_int_clause("limit", &[10, 20]),
510 ]),
511 "{k} * {limit} > 5",
512 ),
513 StrategyName::Halton,
514 Some(10),
515 );
516 let json = serde_json::to_string(&c).unwrap();
517 let back: Comprehension = serde_json::from_str(&json).unwrap();
518 assert_eq!(c, back);
519 }
520
521 #[test]
524 fn an_orders_seed_round_trips_and_defaults_to_none() {
525 let c = Comprehension::order_seeded(
526 Comprehension::clause(
527 "k",
528 Source::IntRange {
529 lo: 1,
530 hi: 4,
531 step: 1,
532 },
533 ),
534 StrategyName::Shuffle,
535 Some(2),
536 Some(42),
537 );
538 let json = serde_json::to_string(&c).unwrap();
539 assert!(json.contains("\"seed\":42"), "{json}");
540 let back: Comprehension = serde_json::from_str(&json).unwrap();
541 assert_eq!(back, c);
542 let unseeded = json.replace(",\"seed\":42", "");
543 let back: Comprehension = serde_json::from_str(&unseeded).unwrap();
544 assert!(
545 matches!(back, Comprehension::Order { seed: None, .. }),
546 "{back:?}"
547 );
548 }
549}