1use crate::comprehension::ast::Comprehension as AlgebraAst;
26use crate::comprehension::clause_ast::{
27 Clause as ClauseForm, ClauseSource as ClauseSourceForm, Comprehension as ClauseAst,
28 ComprehensionMode as ModeForm, Subspace as SubspaceForm, TraversalOrder as OrderForm,
29 ZipMode as ZipModeForm,
30};
31use crate::comprehension::strategy::{StrategyName, ZipMode as AlgebraZipMode};
32
33use super::source_parser::{SourceParseError, parse_source};
34
35#[derive(Debug, Clone, PartialEq)]
37pub enum ConvertError {
38 SourceParse {
40 clause_var: String,
42 source: String,
44 cause: SourceParseError,
46 },
47 EmptyComprehension,
49 EmptyUnionSubspace,
51 ParallelArityMismatch {
54 vars: usize,
56 exprs: usize,
58 },
59}
60
61impl std::fmt::Display for ConvertError {
62 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63 match self {
64 ConvertError::SourceParse {
65 clause_var,
66 source,
67 cause,
68 } => write!(
69 f,
70 "clause {clause_var:?} source {source:?} failed to parse: {cause}"
71 ),
72 ConvertError::EmptyComprehension => f.write_str("comprehension has no clauses"),
73 ConvertError::EmptyUnionSubspace => f.write_str("union has empty sub-space"),
74 ConvertError::ParallelArityMismatch { vars, exprs } => {
75 write!(f, "parallel clause vars={vars} != exprs={exprs}")
76 }
77 }
78 }
79}
80
81impl std::error::Error for ConvertError {}
82
83pub fn clauses_to_algebra(clauses: &ClauseAst) -> Result<AlgebraAst, ConvertError> {
102 let body = match &clauses.mode {
103 ModeForm::Cartesian(clauses) => convert_cartesian(clauses)?,
104 ModeForm::Union(subspaces) => convert_union(subspaces)?,
105 };
106
107 let with_filter = if let Some(pred) = &clauses.filter {
108 AlgebraAst::filter(body, pred.clone())
109 } else {
110 body
111 };
112
113 let with_order = if let Some(order) = &clauses.order {
114 let (strategy, truncation, seed) = convert_order(order)?;
115 AlgebraAst::order_seeded(with_filter, strategy, truncation, seed)
116 } else {
117 with_filter
118 };
119
120 Ok(with_order)
121}
122
123fn convert_cartesian(clauses: &[ClauseForm]) -> Result<AlgebraAst, ConvertError> {
124 if clauses.is_empty() {
125 return Err(ConvertError::EmptyComprehension);
126 }
127 let algebra_children: Vec<AlgebraAst> = clauses
128 .iter()
129 .map(convert_clause)
130 .collect::<Result<_, _>>()?;
131 if algebra_children.len() == 1 {
132 Ok(algebra_children.into_iter().next().unwrap())
136 } else {
137 Ok(AlgebraAst::cartesian(algebra_children))
138 }
139}
140
141fn convert_union(subspaces: &[SubspaceForm]) -> Result<AlgebraAst, ConvertError> {
142 if subspaces.is_empty() {
143 return Err(ConvertError::EmptyComprehension);
144 }
145 let algebra_children: Vec<AlgebraAst> = subspaces
146 .iter()
147 .map(|s| {
148 if s.is_empty() {
149 Err(ConvertError::EmptyUnionSubspace)
150 } else {
151 convert_cartesian(&s.clauses)
152 }
153 })
154 .collect::<Result<_, _>>()?;
155 if algebra_children.len() == 1 {
156 Ok(algebra_children.into_iter().next().unwrap())
157 } else {
158 Ok(AlgebraAst::union(algebra_children))
159 }
160}
161
162fn convert_clause(clause: &ClauseForm) -> Result<AlgebraAst, ConvertError> {
163 match &clause.source {
164 ClauseSourceForm::Single(source_str) => {
165 let var = clause
166 .single_var()
167 .unwrap_or_else(|| clause.first_var())
168 .to_string();
169 let source = parse_source(source_str).map_err(|cause| ConvertError::SourceParse {
170 clause_var: var.clone(),
171 source: source_str.clone(),
172 cause,
173 })?;
174 Ok(AlgebraAst::clause(var, source))
175 }
176 ClauseSourceForm::Parallel { mode, exprs } => {
177 if clause.vars.len() != exprs.len() {
178 return Err(ConvertError::ParallelArityMismatch {
179 vars: clause.vars.len(),
180 exprs: exprs.len(),
181 });
182 }
183 let mut children = Vec::with_capacity(clause.vars.len());
187 for (var, expr) in clause.vars.iter().zip(exprs.iter()) {
188 let source = parse_source(expr).map_err(|cause| ConvertError::SourceParse {
189 clause_var: var.clone(),
190 source: expr.clone(),
191 cause,
192 })?;
193 children.push(AlgebraAst::clause(var.clone(), source));
194 }
195 let zip_mode = convert_zip_mode(*mode);
196 Ok(AlgebraAst::zip(children, zip_mode))
197 }
198 }
199}
200
201fn convert_zip_mode(clauses: ZipModeForm) -> AlgebraZipMode {
202 match clauses {
203 ZipModeForm::Strict => AlgebraZipMode::Strict,
204 ZipModeForm::Truncate => AlgebraZipMode::Truncate,
205 ZipModeForm::Cycle => AlgebraZipMode::Cycle,
206 }
207}
208
209pub(crate) fn convert_order(
215 order: &OrderForm,
216) -> Result<(StrategyName, Option<u64>, Option<u64>), ConvertError> {
217 let triple = match order {
218 OrderForm::Lex { count } => (StrategyName::Lex, count.map(|n| n as u64), None),
219 OrderForm::ReverseLex { count } => {
220 (StrategyName::ReverseLex, count.map(|n| n as u64), None)
221 }
222 OrderForm::Diagonal { count } => (StrategyName::Diagonal, count.map(|n| n as u64), None),
223 OrderForm::Antidiagonal { count } => {
224 (StrategyName::Antidiagonal, count.map(|n| n as u64), None)
225 }
226 OrderForm::Extrema { strata } => (StrategyName::Extrema, strata.map(|n| n as u64), None),
227 OrderForm::Shells { depth, .. } => (StrategyName::Shells, depth.map(|n| n as u64), None),
228 OrderForm::Halton { count } => (StrategyName::Halton, count.map(|n| n as u64), None),
229 OrderForm::Sobol { count } => (StrategyName::Sobol, count.map(|n| n as u64), None),
230 OrderForm::Lhs { count, seed } => (StrategyName::Lhs, count.map(|n| n as u64), *seed),
231 OrderForm::Shuffle { count, seed } => {
232 (StrategyName::Shuffle, count.map(|n| n as u64), *seed)
233 }
234 };
235 Ok(triple)
236}
237
238#[cfg(test)]
239mod tests {
240 use super::*;
241 use crate::comprehension::source::{LiteralValue, Source};
242
243 fn clause(var: &str, source: &str) -> ClauseForm {
244 ClauseForm::new(var, source)
245 }
246
247 #[test]
248 fn cartesian_single_clause_collapses_to_clause() {
249 let clauses = ClauseAst {
250 mode: ModeForm::Cartesian(vec![clause("k", "1..10")]),
251 filter: None,
252 order: None,
253 };
254 let algebra = clauses_to_algebra(&clauses).unwrap();
255 match algebra {
256 AlgebraAst::Clause { name, source } => {
257 assert_eq!(name, "k");
258 assert!(matches!(
259 source,
260 Source::IntRange {
261 lo: 1,
262 hi: 10,
263 step: 1
264 }
265 ));
266 }
267 other => panic!("expected Clause, got {other:?}"),
268 }
269 }
270
271 #[test]
272 fn multi_clause_cartesian_becomes_algebra_cartesian() {
273 let clauses = ClauseAst {
274 mode: ModeForm::Cartesian(vec![
275 clause("k", "1..10"),
276 clause("limit", "[10, 100, 1000]"),
277 ]),
278 filter: None,
279 order: None,
280 };
281 let algebra = clauses_to_algebra(&clauses).unwrap();
282 match algebra {
283 AlgebraAst::Cartesian { children } => {
284 assert_eq!(children.len(), 2);
285 match &children[0] {
287 AlgebraAst::Clause { name, source } => {
288 assert_eq!(name, "k");
289 assert!(matches!(
290 source,
291 Source::IntRange {
292 lo: 1,
293 hi: 10,
294 step: 1
295 }
296 ));
297 }
298 other => panic!("expected Clause, got {other:?}"),
299 }
300 match &children[1] {
302 AlgebraAst::Clause { name, source } => {
303 assert_eq!(name, "limit");
304 match source {
305 Source::Literal { values } => {
306 assert_eq!(values.len(), 3);
307 assert_eq!(values[0], LiteralValue::Int(10));
308 }
309 other => panic!("expected Literal, got {other:?}"),
310 }
311 }
312 other => panic!("expected Clause, got {other:?}"),
313 }
314 }
315 other => panic!("expected Cartesian, got {other:?}"),
316 }
317 }
318
319 #[test]
320 fn filter_wraps_body() {
321 let clauses = ClauseAst {
322 mode: ModeForm::Cartesian(vec![clause("k", "1..10")]),
323 filter: Some("{k} > 5".to_string()),
324 order: None,
325 };
326 let algebra = clauses_to_algebra(&clauses).unwrap();
327 assert!(matches!(algebra, AlgebraAst::Filter { .. }));
328 }
329
330 #[test]
331 fn order_lex_with_count_round_trips() {
332 let clauses = ClauseAst {
333 mode: ModeForm::Cartesian(vec![clause("k", "1..10")]),
334 filter: None,
335 order: Some(OrderForm::Lex { count: Some(5) }),
336 };
337 let algebra = clauses_to_algebra(&clauses).unwrap();
338 match algebra {
339 AlgebraAst::Order {
340 strategy: StrategyName::Lex,
341 truncation: Some(5),
342 ..
343 } => {}
344 other => panic!("expected Order(Lex, Some(5)), got {other:?}"),
345 }
346 }
347
348 #[test]
349 fn order_halton_with_count() {
350 let clauses = ClauseAst {
351 mode: ModeForm::Cartesian(vec![clause("k", "1..10"), clause("limit", "1..100")]),
352 filter: None,
353 order: Some(OrderForm::Halton { count: Some(20) }),
354 };
355 let algebra = clauses_to_algebra(&clauses).unwrap();
356 match algebra {
357 AlgebraAst::Order {
358 strategy: StrategyName::Halton,
359 truncation: Some(20),
360 ..
361 } => {}
362 other => panic!("expected Order(Halton, Some(20)), got {other:?}"),
363 }
364 }
365
366 #[test]
367 fn union_of_subspaces() {
368 let clauses = ClauseAst {
369 mode: ModeForm::Union(vec![
370 SubspaceForm::new(vec![clause("k", "10"), clause("limit", "[1, 2, 3]")]),
371 SubspaceForm::new(vec![clause("k", "100"), clause("limit", "[10, 20, 30]")]),
372 ]),
373 filter: None,
374 order: None,
375 };
376 let algebra = clauses_to_algebra(&clauses).unwrap();
377 match algebra {
378 AlgebraAst::Union { children } => assert_eq!(children.len(), 2),
379 other => panic!("expected Union, got {other:?}"),
380 }
381 }
382
383 #[test]
384 fn parallel_clause_becomes_zip() {
385 let parallel = ClauseForm::parallel(["x", "y"], ["1..3", "10..30"]);
386 let clauses = ClauseAst {
387 mode: ModeForm::Cartesian(vec![parallel]),
388 filter: None,
389 order: None,
390 };
391 let algebra = clauses_to_algebra(&clauses).unwrap();
392 match algebra {
395 AlgebraAst::Zip {
396 children,
397 mode: AlgebraZipMode::Strict,
398 } => {
399 assert_eq!(children.len(), 2);
400 }
401 other => panic!("expected Zip, got {other:?}"),
402 }
403 }
404
405 #[test]
406 fn unparseable_source_falls_back_to_generator() {
407 let clauses = ClauseAst {
412 mode: ModeForm::Cartesian(vec![clause("k", "totally nonsense")]),
413 filter: None,
414 order: None,
415 };
416 let algebra = clauses_to_algebra(&clauses).unwrap();
417 match algebra {
418 AlgebraAst::Clause { source, .. } => match source {
419 crate::comprehension::source::Source::Generator { expr, .. } => {
420 assert_eq!(expr, "totally nonsense");
421 }
422 other => panic!("expected Generator, got {other:?}"),
423 },
424 other => panic!("expected Clause, got {other:?}"),
425 }
426 }
427
428 #[test]
435 fn a_seeded_order_keeps_its_seed_in_the_algebra() {
436 let clauses = ClauseAst {
437 mode: ModeForm::Cartesian(vec![clause("k", "1..10")]),
438 filter: None,
439 order: Some(OrderForm::Shuffle {
440 count: Some(3),
441 seed: Some(42),
442 }),
443 };
444 let algebra = clauses_to_algebra(&clauses).unwrap();
445 assert!(
446 matches!(
447 algebra,
448 AlgebraAst::Order {
449 strategy: StrategyName::Shuffle,
450 truncation: Some(3),
451 seed: Some(42),
452 ..
453 }
454 ),
455 "{algebra:?}"
456 );
457 assert_eq!(
458 convert_order(&OrderForm::Lhs {
459 count: None,
460 seed: Some(7)
461 })
462 .unwrap(),
463 (StrategyName::Lhs, None, Some(7))
464 );
465 assert_eq!(
466 convert_order(&OrderForm::Halton { count: Some(4) }).unwrap(),
467 (StrategyName::Halton, Some(4), None)
468 );
469 }
470}