1#[cfg(test)]
63use std::sync::Arc;
64
65use crate::ast::Value;
66use crate::dsl::compile::eval_const_expr_for;
67use crate::iteration::comprehension::ast::Comprehension;
68use crate::iteration::comprehension::cardinality::{Interval, ProductMeasure};
69use crate::iteration::comprehension::eval_source::{EvalContext, SourceEval};
70use crate::iteration::comprehension::measure::AxisMeasure;
71use crate::iteration::comprehension::metadata::IndexFn;
72use crate::iteration::comprehension::source::Source;
73use crate::iteration::comprehension::strategies::{EvaluatedInput, Tuple, TupleValue};
74use crate::iteration::comprehension::strategy::StrategyName;
75#[cfg(test)]
76use crate::kernel::PolydatKernel;
77use crate::kernel::interp::{Layered, Lookup, interpolate_via_kernel};
78
79pub type RuntimeTuple = Vec<(String, Value)>;
86
87struct EvaluatedNode {
97 tuples: Vec<RuntimeTuple>,
98 index_fn: Option<IndexFn>,
99}
100
101#[derive(Debug, Clone)]
103pub enum RuntimeError {
104 SourceEval {
107 var: String,
109 source: String,
111 message: String,
113 },
114 FilterEval {
116 predicate: String,
118 message: String,
120 },
121 OrderEval {
123 strategy: StrategyName,
125 message: String,
127 },
128 StrategyRejectsInput {
131 strategy: StrategyName,
133 index_fn: Option<IndexFn>,
135 },
136 UnsupportedShape(String),
139}
140
141impl std::fmt::Display for RuntimeError {
142 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
143 match self {
144 RuntimeError::SourceEval {
145 var,
146 source,
147 message,
148 } => {
149 write!(f, "for_each clause '{var} in {source}': {message}")
150 }
151 RuntimeError::FilterEval { predicate, message } => {
152 write!(f, "comprehension filter '{predicate}': {message}")
153 }
154 RuntimeError::OrderEval { strategy, message } => {
155 write!(f, "order strategy {strategy:?}: {message}")
156 }
157 RuntimeError::StrategyRejectsInput { strategy, index_fn } => write!(
158 f,
159 "order strategy {strategy:?} rejects input shape {index_fn:?} \
160 (V4: per-strategy IndexFn contract; see spec §3.6's strategy table)"
161 ),
162 RuntimeError::UnsupportedShape(msg) => write!(f, "{msg}"),
163 }
164 }
165}
166
167impl std::error::Error for RuntimeError {}
168
169pub fn evaluate_for_iteration(
181 comp: &Comprehension,
182 scope: &dyn Lookup,
183) -> Result<Vec<RuntimeTuple>, RuntimeError> {
184 evaluate_for_iteration_reported(comp, scope).map(|e| e.tuples)
185}
186
187pub fn evaluate_for_iteration_reported(
202 comp: &Comprehension,
203 scope: &dyn Lookup,
204) -> Result<EvaluatedIteration, RuntimeError> {
205 let mut state = EvalState::new(comp, scope);
206 let tuples = state.evaluate_node(comp, &[])?.tuples;
207 Ok(EvaluatedIteration {
208 tuples,
209 clauses: state.yields,
210 })
211}
212
213fn fast_predicate(predicate: &str, tuple: &RuntimeTuple) -> Option<bool> {
218 let p = predicate.trim();
219 if p.eq_ignore_ascii_case("true") {
220 return Some(true);
221 }
222 if p.eq_ignore_ascii_case("false") {
223 return Some(false);
224 }
225 if let Some(inner) = p.strip_prefix('!') {
226 return fast_predicate(inner, tuple).map(|b| !b);
227 }
228 if let Some(parts) = split_top(p, "||") {
229 let mut any = false;
230 for part in parts {
231 any |= fast_predicate(&part, tuple)?;
232 }
233 return Some(any);
234 }
235 if let Some(parts) = split_top(p, "&&") {
236 let mut all = true;
237 for part in parts {
238 all &= fast_predicate(&part, tuple)?;
239 }
240 return Some(all);
241 }
242 if let Some(pos) = p.find(" in ") {
243 let name = curly(p[..pos].trim())?;
244 let list = p[pos + 4..].trim().strip_prefix('[')?.strip_suffix(']')?;
245 let needle = tuple_scalar(tuple, &name)?;
246 let mut hit = false;
247 for item in list.split(',') {
248 let lit = literal(item.trim())?;
249 hit |= scalar_eq(&needle, &lit);
250 }
251 return Some(hit);
252 }
253 for op in ["==", "!=", "<=", ">=", "<", ">"] {
254 if let Some((lhs, rhs)) = split_op(p, op) {
255 let lhs = lhs.trim();
256 let rhs = rhs.trim();
257 let a = operand(tuple, lhs)?;
258 let b = operand(tuple, rhs)?;
259 return Some(match op {
260 "==" => scalar_eq(&a, &b),
261 "!=" => !scalar_eq(&a, &b),
262 "<" => scalar_cmp(&a, &b)? == std::cmp::Ordering::Less,
263 ">" => scalar_cmp(&a, &b)? == std::cmp::Ordering::Greater,
264 "<=" => scalar_cmp(&a, &b)? != std::cmp::Ordering::Greater,
265 _ => scalar_cmp(&a, &b)? != std::cmp::Ordering::Less,
266 });
267 }
268 }
269 None
270}
271
272#[derive(Debug, Clone, PartialEq)]
273enum Scalar {
274 Int(i128),
275 Float(f64),
276 Str(String),
277 Bool(bool),
278}
279
280fn operand(tuple: &RuntimeTuple, text: &str) -> Option<Scalar> {
281 match curly(text) {
282 Some(name) => tuple_scalar(tuple, &name),
283 None => literal(text),
284 }
285}
286
287fn tuple_scalar(tuple: &RuntimeTuple, name: &str) -> Option<Scalar> {
288 let (_, v) = tuple.iter().find(|(n, _)| n == name)?;
289 match v {
290 Value::U64(n) => Some(Scalar::Int(*n as i128)),
291 Value::F64(f) => Some(Scalar::Float(*f)),
292 Value::Str(s) => Some(Scalar::Str(s.to_string())),
293 Value::Bool(b) => Some(Scalar::Bool(*b)),
294 Value::Json(j) => match j.as_ref() {
296 serde_json::Value::Number(n) if n.is_i64() => Some(Scalar::Int(n.as_i64()? as i128)),
297 serde_json::Value::Number(n) if n.is_u64() => Some(Scalar::Int(n.as_u64()? as i128)),
298 serde_json::Value::Number(n) => Some(Scalar::Float(n.as_f64()?)),
299 serde_json::Value::String(s) => Some(Scalar::Str(s.clone())),
300 serde_json::Value::Bool(b) => Some(Scalar::Bool(*b)),
301 _ => None,
302 },
303 _ => None,
304 }
305}
306
307fn literal(text: &str) -> Option<Scalar> {
308 if let Some(s) = text.strip_prefix('"').and_then(|s| s.strip_suffix('"')) {
309 return Some(Scalar::Str(s.to_string()));
310 }
311 if let Some(s) = text.strip_prefix('\'').and_then(|s| s.strip_suffix('\'')) {
312 return Some(Scalar::Str(s.to_string()));
313 }
314 match text {
315 "true" => return Some(Scalar::Bool(true)),
316 "false" => return Some(Scalar::Bool(false)),
317 _ => {}
318 }
319 if let Ok(i) = text.parse::<i128>() {
320 return Some(Scalar::Int(i));
321 }
322 if let Ok(f) = text.parse::<f64>() {
323 return Some(Scalar::Float(f));
324 }
325 if !text.is_empty()
328 && text
329 .chars()
330 .all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-')
331 {
332 return Some(Scalar::Str(text.to_string()));
333 }
334 None
335}
336
337fn scalar_eq(a: &Scalar, b: &Scalar) -> bool {
338 match (a, b) {
339 (Scalar::Int(x), Scalar::Float(y)) | (Scalar::Float(y), Scalar::Int(x)) => {
340 (*x as f64) == *y
341 }
342 _ => a == b,
343 }
344}
345
346fn scalar_cmp(a: &Scalar, b: &Scalar) -> Option<std::cmp::Ordering> {
347 match (a, b) {
348 (Scalar::Int(x), Scalar::Int(y)) => Some(x.cmp(y)),
349 (Scalar::Float(x), Scalar::Float(y)) => x.partial_cmp(y),
350 (Scalar::Int(x), Scalar::Float(y)) => (*x as f64).partial_cmp(y),
351 (Scalar::Float(x), Scalar::Int(y)) => x.partial_cmp(&(*y as f64)),
352 (Scalar::Str(x), Scalar::Str(y)) => Some(x.cmp(y)),
353 _ => None,
354 }
355}
356
357fn curly(text: &str) -> Option<String> {
358 let inner = text.strip_prefix('{')?.strip_suffix('}')?;
359 (!inner.is_empty() && inner.chars().all(|c| c.is_ascii_alphanumeric() || c == '_'))
360 .then(|| inner.to_string())
361}
362
363fn split_top(s: &str, sep: &str) -> Option<Vec<String>> {
365 let mut parts = Vec::new();
366 let mut depth = 0i32;
367 let mut quote: Option<char> = None;
368 let mut start = 0;
369 let bytes: Vec<char> = s.chars().collect();
370 let sepc: Vec<char> = sep.chars().collect();
371 let mut i = 0;
372 while i < bytes.len() {
373 let c = bytes[i];
374 if let Some(q) = quote {
375 if c == q {
376 quote = None;
377 }
378 } else {
379 match c {
380 '"' | '\'' => quote = Some(c),
381 '(' | '[' | '{' => depth += 1,
382 ')' | ']' | '}' => depth -= 1,
383 _ => {}
384 }
385 if depth == 0 && bytes[i..].starts_with(&sepc) {
386 parts.push(bytes[start..i].iter().collect::<String>());
387 i += sepc.len();
388 start = i;
389 continue;
390 }
391 }
392 i += 1;
393 }
394 if parts.is_empty() {
395 return None;
396 }
397 parts.push(bytes[start..].iter().collect::<String>());
398 Some(parts)
399}
400
401fn split_op<'a>(s: &'a str, op: &str) -> Option<(&'a str, &'a str)> {
402 let mut depth = 0i32;
403 let mut quote: Option<char> = None;
404 let chars: Vec<(usize, char)> = s.char_indices().collect();
405 for (k, &(idx, c)) in chars.iter().enumerate() {
406 if let Some(q) = quote {
407 if c == q {
408 quote = None;
409 }
410 continue;
411 }
412 match c {
413 '"' | '\'' => {
414 quote = Some(c);
415 continue;
416 }
417 '(' | '[' | '{' => depth += 1,
418 ')' | ']' | '}' => depth -= 1,
419 _ => {}
420 }
421 if depth == 0 && s[idx..].starts_with(op) {
422 let next = chars.get(k + op.len()).map(|(_, c)| *c);
424 if (op == "<" || op == ">") && next == Some('=') {
425 continue;
426 }
427 let prev = if k > 0 { Some(chars[k - 1].1) } else { None };
428 if (op == "<" || op == ">") && matches!(prev, Some('<') | Some('>')) {
429 continue;
430 }
431 return Some((&s[..idx], &s[idx + op.len()..]));
432 }
433 }
434 None
435}
436
437#[derive(Debug, Clone, PartialEq, Eq)]
453pub struct ClauseYield {
454 pub var: String,
456 pub source: Option<String>,
458 pub evaluations: usize,
460 pub values: usize,
462}
463
464#[derive(Debug, Clone)]
467pub struct EvaluatedIteration {
468 pub tuples: Vec<RuntimeTuple>,
470 pub clauses: Vec<ClauseYield>,
472}
473
474struct EvalState<'a> {
475 scope: &'a dyn Lookup,
478 yields: Vec<ClauseYield>,
482 by_leaf: std::collections::HashMap<usize, usize>,
487}
488
489impl<'a> EvalState<'a> {
490 fn new(comp: &Comprehension, scope: &'a dyn Lookup) -> Self {
492 let mut state = EvalState {
493 scope,
494 yields: Vec::new(),
495 by_leaf: std::collections::HashMap::new(),
496 };
497 state.enumerate_leaves(comp);
498 state
499 }
500
501 fn enumerate_leaves(&mut self, node: &Comprehension) {
502 match node {
503 Comprehension::Clause { name, source } => {
504 self.by_leaf
505 .insert(std::ptr::from_ref(source) as usize, self.yields.len());
506 self.yields.push(ClauseYield {
507 var: name.clone(),
508 source: source.to_text(),
509 evaluations: 0,
510 values: 0,
511 });
512 }
513 Comprehension::Cartesian { children }
514 | Comprehension::Zip { children, .. }
515 | Comprehension::Union { children } => {
516 for child in children {
517 self.enumerate_leaves(child);
518 }
519 }
520 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
521 self.enumerate_leaves(child);
522 }
523 }
524 }
525
526 fn record_yield(&mut self, source: &Source, values: usize) {
528 if let Some(&i) = self.by_leaf.get(&(std::ptr::from_ref(source) as usize)) {
529 self.yields[i].evaluations += 1;
530 self.yields[i].values += values;
531 }
532 }
533}
534
535impl EvalState<'_> {
536 fn evaluate_node(
537 &mut self,
538 node: &Comprehension,
539 prefix: &[(String, Value)],
540 ) -> Result<EvaluatedNode, RuntimeError> {
541 match node {
542 Comprehension::Clause { name, source } => self.evaluate_clause(name, source, prefix),
543 Comprehension::Cartesian { children } => self.evaluate_cartesian(children, prefix),
544 Comprehension::Zip { children, mode } => self.evaluate_zip(children, *mode, prefix),
545 Comprehension::Union { children } => self.evaluate_union(children, prefix),
546 Comprehension::Filter { child, predicate } => {
547 let inner = self.evaluate_node(child, prefix)?;
548 self.apply_filter(inner, predicate)
549 }
550 Comprehension::Order {
551 child,
552 strategy,
553 truncation,
554 seed,
555 } => {
556 if has_continuous_axis(child) {
561 return self.sample_space(child, prefix, *strategy, *truncation, *seed);
562 }
563 let inner = self.evaluate_node(child, prefix)?;
564 self.apply_order(inner, *strategy, *truncation, *seed)
565 }
566 }
567 }
568
569 fn evaluate_clause(
570 &mut self,
571 name: &str,
572 source: &Source,
573 prefix: &[(String, Value)],
574 ) -> Result<EvaluatedNode, RuntimeError> {
575 let ctx = EvalContext {
576 var_name: name,
577 scope: self.scope,
578 prefix,
579 };
580 let evaluated = source.evaluate(Some(&ctx)).map_err(|e| match e {
581 crate::iteration::comprehension::eval_source::EvalError::EvalFailed {
582 var,
583 source,
584 message,
585 } => RuntimeError::SourceEval {
586 var,
587 source,
588 message,
589 },
590 crate::iteration::comprehension::eval_source::EvalError::NeedsContext => {
591 RuntimeError::UnsupportedShape(format!(
592 "clause '{name}': source requires kernel context but evaluator \
593 provided none — internal bug in runtime walker"
594 ))
595 }
596 })?;
597
598 self.record_yield(source, evaluated.values.len());
599
600 if evaluated.values.is_empty() {
601 return Ok(EvaluatedNode {
607 tuples: Vec::new(),
608 index_fn: Some(evaluated.index_fn),
609 });
610 }
611 let tuples: Vec<RuntimeTuple> = evaluated
612 .values
613 .into_iter()
614 .map(|v| vec![(name.to_string(), v)])
615 .collect();
616 Ok(EvaluatedNode {
617 tuples,
618 index_fn: Some(evaluated.index_fn),
619 })
620 }
621
622 fn evaluate_cartesian(
623 &mut self,
624 children: &[Comprehension],
625 prefix: &[(String, Value)],
626 ) -> Result<EvaluatedNode, RuntimeError> {
627 if children.is_empty() {
628 return Ok(EvaluatedNode {
629 tuples: vec![Vec::new()],
630 index_fn: Some(IndexFn::Lattice {
631 axis_sizes: vec![1],
632 }),
633 });
634 }
635 let mut child_index_fns: Vec<Option<IndexFn>> = Vec::with_capacity(children.len());
636 let mut dependent_observed = false;
637 let result_tuples = self.evaluate_cartesian_rec(
638 children,
639 prefix,
640 &mut child_index_fns,
641 &mut dependent_observed,
642 )?;
643
644 let combined = if dependent_observed {
650 None
651 } else {
652 combine_cartesian_index_fn(&child_index_fns)
653 };
654 Ok(EvaluatedNode {
655 tuples: result_tuples,
656 index_fn: combined,
657 })
658 }
659
660 fn evaluate_cartesian_rec(
661 &mut self,
662 children: &[Comprehension],
663 prefix: &[(String, Value)],
664 child_index_fns: &mut Vec<Option<IndexFn>>,
665 dependent_observed: &mut bool,
666 ) -> Result<Vec<RuntimeTuple>, RuntimeError> {
667 if children.is_empty() {
668 return Ok(vec![Vec::new()]);
669 }
670 let (head, tail) = children.split_first().unwrap();
671 let head_eval = self.evaluate_node(head, prefix)?;
672 let head_axis_len = head_eval.tuples.len() as u64;
673 if child_index_fns.len() <= prefix_depth(prefix, child_index_fns) {
675 child_index_fns.push(head_eval.index_fn.clone());
676 } else if let Some(prev) = child_index_fns
677 .get(prefix_depth(prefix, child_index_fns))
678 .cloned()
679 .flatten()
680 {
681 if axis_size_of(&prev) != Some(head_axis_len) {
685 *dependent_observed = true;
686 }
687 }
688
689 if tail.is_empty() {
690 return Ok(head_eval.tuples);
691 }
692 let mut out = Vec::new();
693 for head_tuple in head_eval.tuples {
694 let mut extended_prefix: Vec<(String, Value)> = prefix.to_vec();
695 extended_prefix.extend(head_tuple.iter().cloned());
696 let tail_tuples = self.evaluate_cartesian_rec(
697 tail,
698 &extended_prefix,
699 child_index_fns,
700 dependent_observed,
701 )?;
702 for tail_tuple in tail_tuples {
703 let mut merged = head_tuple.clone();
704 merged.extend(tail_tuple);
705 out.push(merged);
706 }
707 }
708 Ok(out)
709 }
710
711 fn evaluate_zip(
712 &mut self,
713 children: &[Comprehension],
714 mode: crate::iteration::comprehension::strategy::ZipMode,
715 prefix: &[(String, Value)],
716 ) -> Result<EvaluatedNode, RuntimeError> {
717 use crate::iteration::comprehension::strategy::ZipMode;
718 if children.is_empty() {
719 return Ok(EvaluatedNode {
720 tuples: vec![Vec::new()],
721 index_fn: Some(IndexFn::Lockstep { length: 1 }),
722 });
723 }
724 let per_child: Vec<EvaluatedNode> = children
725 .iter()
726 .map(|c| self.evaluate_node(c, prefix))
727 .collect::<Result<_, _>>()?;
728 let lengths: Vec<usize> = per_child.iter().map(|n| n.tuples.len()).collect();
729 let iter_count = match mode {
730 ZipMode::Strict => {
731 let first = lengths.first().copied().unwrap_or(0);
732 if lengths.iter().any(|&n| n != first) {
733 return Err(RuntimeError::UnsupportedShape(format!(
734 "zip strict: child lengths differ ({lengths:?})"
735 )));
736 }
737 first
738 }
739 ZipMode::Truncate => lengths.iter().copied().min().unwrap_or(0),
740 ZipMode::Cycle => lengths.iter().copied().max().unwrap_or(0),
741 };
742 let mut tuples = Vec::with_capacity(iter_count);
743 for i in 0..iter_count {
744 let mut bindings: RuntimeTuple = Vec::new();
745 for (child, &len) in per_child.iter().zip(lengths.iter()) {
746 if len == 0 {
747 continue;
748 }
749 let idx = match mode {
750 ZipMode::Cycle => i % len,
751 _ => i,
752 };
753 bindings.extend(child.tuples[idx].iter().cloned());
754 }
755 tuples.push(bindings);
756 }
757 let index_fn = match mode {
758 ZipMode::Strict | ZipMode::Truncate => Some(IndexFn::Lockstep {
759 length: iter_count as u64,
760 }),
761 ZipMode::Cycle => Some(IndexFn::Modular {
762 axis_sizes: lengths.iter().map(|n| *n as u64).collect(),
763 }),
764 };
765 Ok(EvaluatedNode { tuples, index_fn })
766 }
767
768 fn evaluate_union(
769 &mut self,
770 children: &[Comprehension],
771 prefix: &[(String, Value)],
772 ) -> Result<EvaluatedNode, RuntimeError> {
773 let mut tuples = Vec::new();
774 let mut segment_sizes = Vec::with_capacity(children.len());
775 let mut all_segments_addressable = true;
776 for child in children {
777 let sub = self.evaluate_node(child, prefix)?;
778 segment_sizes.push(sub.tuples.len() as u64);
779 if sub.index_fn.is_none() {
780 all_segments_addressable = false;
781 }
782 tuples.extend(sub.tuples);
783 }
784 let index_fn = if all_segments_addressable {
785 Some(IndexFn::Concatenation { segment_sizes })
786 } else {
787 None
788 };
789 Ok(EvaluatedNode { tuples, index_fn })
790 }
791
792 fn apply_filter(
793 &mut self,
794 input: EvaluatedNode,
795 predicate: &str,
796 ) -> Result<EvaluatedNode, RuntimeError> {
797 let mut out = Vec::with_capacity(input.tuples.len());
798 for tuple in input.tuples {
799 if let Some(keep) = fast_predicate(predicate, &tuple) {
805 if keep {
806 out.push(tuple);
807 }
808 continue;
809 }
810 let scope = Layered {
811 prefix: &tuple,
812 inner: self.scope,
813 };
814 let interpolated = interpolate_via_kernel(predicate, &scope).map_err(|e| {
815 RuntimeError::FilterEval {
816 predicate: predicate.to_string(),
817 message: e.to_string(),
818 }
819 })?;
820 let result = eval_const_expr_for(&interpolated, self.scope.ledger()).map_err(|e| {
821 RuntimeError::FilterEval {
822 predicate: predicate.to_string(),
823 message: e.to_string(),
824 }
825 })?;
826 let keep = match result {
827 Value::Bool(b) => b,
828 Value::U64(n) => n != 0,
829 Value::F64(n) => n != 0.0,
830 other => {
831 return Err(RuntimeError::FilterEval {
832 predicate: predicate.to_string(),
833 message: format!("expected bool/u64/f64, got {other:?}"),
834 });
835 }
836 };
837 if keep {
838 out.push(tuple);
839 }
840 }
841 Ok(EvaluatedNode {
843 tuples: out,
844 index_fn: None,
845 })
846 }
847
848 fn sample_space(
861 &mut self,
862 child: &Comprehension,
863 prefix: &[(String, Value)],
864 strategy: StrategyName,
865 truncation: Option<u64>,
866 seed: Option<u64>,
867 ) -> Result<EvaluatedNode, RuntimeError> {
868 let mut space = SampleSpace::default();
869 self.collect_sample_space(child, prefix, &mut space, &mut Vec::new())?;
870 let discrete_axes: Vec<u64> = space
871 .axes
872 .iter()
873 .filter_map(|a| match a {
874 SampleAxis::Discrete(tuples) => Some(tuples.len() as u64),
875 SampleAxis::Continuous { .. } => None,
876 })
877 .collect();
878 if discrete_axes.contains(&0) {
879 return Ok(EvaluatedNode {
880 tuples: Vec::new(),
881 index_fn: None,
882 });
883 }
884 let (intervals, measures): (Vec<Interval>, Vec<ProductMeasure>) = space
885 .axes
886 .iter()
887 .filter_map(|a| match a {
888 SampleAxis::Continuous {
889 interval, measure, ..
890 } => Some((
891 interval.clone(),
892 match measure {
893 AxisMeasure::Uniform => ProductMeasure::Uniform,
894 AxisMeasure::Named { name, .. } => ProductMeasure::Named(*name),
895 },
896 )),
897 SampleAxis::Discrete(_) => None,
898 })
899 .unzip();
900 let sequence = !matches!(strategy, StrategyName::Extrema);
901 let index_fn = if !sequence {
906 IndexFn::Lattice {
907 axis_sizes: space
908 .axes
909 .iter()
910 .map(|a| match a {
911 SampleAxis::Discrete(tuples) => tuples.len() as u64,
912 SampleAxis::Continuous { .. } => 2,
913 })
914 .collect(),
915 }
916 } else if discrete_axes.is_empty() {
917 IndexFn::Continuous {
918 intervals,
919 measure: ProductMeasure::Product(measures),
920 }
921 } else {
922 IndexFn::Hybrid {
923 discrete_axes,
924 continuous_axes: intervals,
925 measure: ProductMeasure::Product(measures),
926 }
927 };
928
929 let mut want = truncation;
930 let mut rounds = 0;
931 loop {
932 let multi_indices = draw_sample(&index_fn, strategy, want, seed)?;
933 let drawn = multi_indices.len() as u64;
934 let tuples = multi_indices
935 .iter()
936 .map(|mi| space.realize(mi, strategy))
937 .collect();
938 let mut node = EvaluatedNode {
939 tuples,
940 index_fn: None,
941 };
942 for predicate in &space.predicates {
943 node = self.apply_filter(node, predicate)?;
944 }
945 let (Some(n), Some(asked)) = (truncation, want) else {
946 return Ok(node);
947 };
948 let enough = node.tuples.len() as u64 >= n;
949 let exhausted = drawn < asked;
950 if !sequence {
951 return Ok(node);
952 }
953 if enough || exhausted || rounds >= SAMPLE_ROUNDS {
954 node.tuples.truncate(n as usize);
955 return Ok(node);
956 }
957
958 want = Some(asked.saturating_mul(2));
959 rounds += 1;
960 }
961 }
962
963 fn collect_sample_space(
970 &mut self,
971 c: &Comprehension,
972 prefix: &[(String, Value)],
973 space: &mut SampleSpace,
974 bound: &mut Vec<String>,
975 ) -> Result<(), RuntimeError> {
976 let measure_error = |name: &str, message: String| RuntimeError::SourceEval {
977 var: name.to_string(),
978 source: "<continuous>".to_string(),
979 message,
980 };
981 match c {
982 Comprehension::Clause {
983 name,
984 source: Source::ContinuousInterval { interval, measure },
985 } => {
986 let measure =
987 AxisMeasure::from_product(measure, 0).map_err(|m| measure_error(name, m))?;
988 space.axes.push(SampleAxis::Continuous {
989 name: name.clone(),
990 interval: interval.clone(),
991 measure,
992 });
993 bound.push(name.clone());
994 }
995 Comprehension::Clause {
996 name,
997 source:
998 Source::Distribution {
999 distribution,
1000 support,
1001 params,
1002 },
1003 } => {
1004 let measure = AxisMeasure::named(*distribution, params)
1005 .map_err(|m| measure_error(name, m))?;
1006 space.axes.push(SampleAxis::Continuous {
1007 name: name.clone(),
1008 interval: support.clone(),
1009 measure,
1010 });
1011 bound.push(name.clone());
1012 }
1013 Comprehension::Clause { name, source } => {
1014 let references = c.referenced_source_names();
1015 if let Some(dep) = bound.iter().find(|b| references.contains(*b)) {
1016 return Err(RuntimeError::UnsupportedShape(format!(
1017 "clause '{name}' references '{dep}' beside a continuous axis; \
1018 a sampled cartesian is independent (comprehension_forms.md §6.2)"
1019 )));
1020 }
1021 let node = self.evaluate_clause(name, source, prefix)?;
1022 space.axes.push(SampleAxis::Discrete(node.tuples));
1023 bound.push(name.clone());
1024 }
1025 Comprehension::Cartesian { children } => {
1026 for child in children {
1027 self.collect_sample_space(child, prefix, space, bound)?;
1028 }
1029 }
1030 Comprehension::Filter { child, predicate } => {
1031 self.collect_sample_space(child, prefix, space, bound)?;
1032 space.predicates.push(predicate.clone());
1033 }
1034 Comprehension::Zip { .. }
1035 | Comprehension::Union { .. }
1036 | Comprehension::Order { .. } => {
1037 let node = self.evaluate_node(c, prefix)?;
1038 bound.extend(c.coordinate_names());
1039 space.axes.push(SampleAxis::Discrete(node.tuples));
1040 }
1041 }
1042 Ok(())
1043 }
1044
1045 fn apply_order(
1046 &mut self,
1047 input: EvaluatedNode,
1048 strategy: StrategyName,
1049 truncation: Option<u64>,
1050 seed: Option<u64>,
1051 ) -> Result<EvaluatedNode, RuntimeError> {
1052 use crate::iteration::comprehension::strategies::{
1053 Strategy, antidiagonal::Antidiagonal, diagonal::Diagonal, extrema::Extrema,
1054 halton::Halton, lex::Lex, lhs::Lhs, reverse_lex::ReverseLex, shells::Shells,
1055 shuffle::Shuffle, sobol::Sobol,
1056 };
1057 use crate::iteration::comprehension::surfaces::polydat_value_to_tuple_value;
1058
1059 let dispatch: Box<dyn Strategy> = match strategy {
1060 StrategyName::Lex => Box::new(Lex),
1061 StrategyName::ReverseLex => Box::new(ReverseLex),
1062 StrategyName::Diagonal => Box::new(Diagonal),
1063 StrategyName::Antidiagonal => Box::new(Antidiagonal),
1064 StrategyName::Extrema => Box::new(Extrema),
1065 StrategyName::Shells => Box::new(Shells),
1066 StrategyName::Halton => Box::new(Halton),
1067 StrategyName::Sobol => Box::new(Sobol),
1068 StrategyName::Lhs => Box::new(Lhs),
1069 StrategyName::Shuffle => Box::new(Shuffle),
1070 };
1071
1072 if !dispatch.accepts_input(input.index_fn.as_ref()) {
1074 return Err(RuntimeError::StrategyRejectsInput {
1075 strategy,
1076 index_fn: input.index_fn.clone(),
1077 });
1078 }
1079
1080 let algebra_tuples: Vec<Tuple> = input
1087 .tuples
1088 .iter()
1089 .map(|rt| Tuple {
1090 bindings: rt
1091 .iter()
1092 .map(|(n, v)| {
1093 let tv = polydat_value_to_tuple_value(v)
1094 .unwrap_or(TupleValue::Str(v.to_display_string()));
1095 (n.clone(), tv)
1096 })
1097 .collect(),
1098 })
1099 .collect();
1100
1101 let index_fn = input.index_fn.clone().unwrap_or(IndexFn::Lattice {
1110 axis_sizes: vec![algebra_tuples.len() as u64],
1111 });
1112 let cardinality = algebra_tuples.len() as u64;
1113 let evaluated_input = EvaluatedInput {
1114 tuples: algebra_tuples.clone(),
1115 cardinality,
1116 index_fn,
1117 };
1118
1119 let ordered = dispatch.apply_seeded(&evaluated_input, truncation, seed);
1120
1121 let mut consumed = vec![false; algebra_tuples.len()];
1124 let mut out = Vec::with_capacity(ordered.len());
1125 for ordered_tuple in &ordered {
1126 let idx = algebra_tuples
1127 .iter()
1128 .enumerate()
1129 .find(|(i, at)| !consumed[*i] && *at == ordered_tuple)
1130 .map(|(i, _)| i)
1131 .ok_or_else(|| RuntimeError::OrderEval {
1132 strategy,
1133 message: "ordered tuple lost reference to runtime source — \
1134 Strategy::apply must return tuples drawn from \
1135 EvaluatedInput.tuples (per spec §10.7.8)"
1136 .into(),
1137 })?;
1138 consumed[idx] = true;
1139 out.push(input.tuples[idx].clone());
1140 }
1141 Ok(EvaluatedNode {
1145 tuples: out,
1146 index_fn: None,
1147 })
1148 }
1149}
1150
1151const SAMPLE_ROUNDS: u32 = 6;
1154
1155const UNIT_SCALE: f64 = (1u64 << 53) as f64;
1158
1159enum SampleAxis {
1163 Discrete(Vec<RuntimeTuple>),
1164 Continuous {
1165 name: String,
1166 interval: Interval,
1167 measure: AxisMeasure,
1168 },
1169}
1170
1171#[derive(Default)]
1175struct SampleSpace {
1176 axes: Vec<SampleAxis>,
1177 predicates: Vec<String>,
1178}
1179
1180impl SampleSpace {
1181 fn realize(&self, mi: &[u64], strategy: StrategyName) -> RuntimeTuple {
1189 let extrema = matches!(strategy, StrategyName::Extrema);
1190 let discrete_count = self
1191 .axes
1192 .iter()
1193 .filter(|a| matches!(a, SampleAxis::Discrete(_)))
1194 .count();
1195 let (mut d, mut c) = (0, if extrema { 0 } else { discrete_count });
1196 let mut out = RuntimeTuple::new();
1197 for axis in &self.axes {
1198 match axis {
1199 SampleAxis::Discrete(tuples) => {
1200 let pos = mi.get(d).copied().unwrap_or(0) as usize;
1201 d += 1;
1202 if extrema {
1203 c += 1;
1204 }
1205 if let Some(t) = tuples.get(pos) {
1206 out.extend(t.iter().cloned());
1207 }
1208 }
1209 SampleAxis::Continuous {
1210 name,
1211 interval,
1212 measure,
1213 } => {
1214 let code = mi.get(c).copied().unwrap_or(0);
1215 c += 1;
1216 if extrema {
1217 d += 1;
1218 }
1219 let x = if extrema {
1220 measure.endpoint(interval, code == 1)
1221 } else {
1222 measure.map_unit(code as f64 / UNIT_SCALE, interval)
1223 };
1224 out.push((name.clone(), Value::F64(x)));
1225 }
1226 }
1227 }
1228 out
1229 }
1230}
1231
1232fn draw_sample(
1237 index_fn: &IndexFn,
1238 strategy: StrategyName,
1239 count: Option<u64>,
1240 seed: Option<u64>,
1241) -> Result<Vec<Vec<u64>>, RuntimeError> {
1242 use crate::iteration::comprehension::strategies::{
1243 extrema::extrema_multi_indices, halton::try_halton_multi_indices,
1244 lhs::try_lhs_multi_indices, shuffle::try_shuffle_multi_indices,
1245 sobol::try_sobol_multi_indices,
1246 };
1247 if matches!(strategy, StrategyName::Extrema) {
1248 return Ok(extrema_multi_indices(index_fn, count));
1249 }
1250 let Some(n) = count else {
1251 return Err(RuntimeError::OrderEval {
1252 strategy,
1253 message: "a continuous source has no finite tuple set; give the order a count, \
1254 as in `order halton/16`"
1255 .into(),
1256 });
1257 };
1258 let drawn = match strategy {
1262 StrategyName::Halton => try_halton_multi_indices(index_fn, Some(n)),
1263 StrategyName::Sobol => try_sobol_multi_indices(index_fn, Some(n)),
1264 StrategyName::Lhs => try_lhs_multi_indices(index_fn, Some(n), seed),
1265 StrategyName::Shuffle => try_shuffle_multi_indices(index_fn, Some(n), seed),
1266 other => {
1267 return Err(RuntimeError::OrderEval {
1268 strategy: other,
1269 message: "a continuous source needs a sampling strategy: halton, sobol, lhs, \
1270 shuffle, or extrema"
1271 .into(),
1272 });
1273 }
1274 };
1275 drawn.map_err(|message| RuntimeError::OrderEval { strategy, message })
1276}
1277
1278pub(crate) fn has_continuous_axis(c: &Comprehension) -> bool {
1283 match c {
1284 Comprehension::Clause { source, .. } => matches!(
1285 source,
1286 Source::ContinuousInterval { .. } | Source::Distribution { .. }
1287 ),
1288 Comprehension::Cartesian { children } => children.iter().any(has_continuous_axis),
1289 Comprehension::Filter { child, .. } => has_continuous_axis(child),
1290 Comprehension::Zip { .. } | Comprehension::Union { .. } | Comprehension::Order { .. } => {
1291 false
1292 }
1293 }
1294}
1295
1296fn prefix_depth(prefix: &[(String, Value)], _recorded: &[Option<IndexFn>]) -> usize {
1305 prefix.len()
1306}
1307
1308fn combine_cartesian_index_fn(children: &[Option<IndexFn>]) -> Option<IndexFn> {
1309 let mut axis_sizes = Vec::new();
1310 for opt in children {
1311 match opt {
1312 Some(IndexFn::Lattice { axis_sizes: a }) => axis_sizes.extend(a.iter().copied()),
1313 Some(IndexFn::Lockstep { length }) => axis_sizes.push(*length),
1314 _ => return None,
1317 }
1318 }
1319 Some(IndexFn::Lattice { axis_sizes })
1320}
1321
1322fn axis_size_of(idx: &IndexFn) -> Option<u64> {
1323 match idx {
1324 IndexFn::Lattice { axis_sizes } if axis_sizes.len() == 1 => Some(axis_sizes[0]),
1325 IndexFn::Lockstep { length } => Some(*length),
1326 _ => None,
1327 }
1328}
1329
1330#[cfg(test)]
1331mod tests {
1332 use super::*;
1333 use crate::iteration::comprehension::source::LiteralValue;
1334
1335 fn empty_kernel() -> Arc<PolydatKernel> {
1336 Arc::new(crate::dsl::compile_polydat_interpreter("\n").unwrap())
1337 }
1338
1339 fn canonical_with_k() -> Arc<PolydatKernel> {
1344 Arc::new(crate::dsl::compile_polydat_interpreter("extern k: u64\n").unwrap())
1345 }
1346
1347 fn clause(name: &str, source: Source) -> Comprehension {
1348 Comprehension::Clause {
1349 name: name.into(),
1350 source,
1351 }
1352 }
1353
1354 fn empty_literal() -> Source {
1355 Source::Literal { values: Vec::new() }
1356 }
1357
1358 #[test]
1360 fn every_leaf_reports_what_it_yielded() {
1361 let comp = Comprehension::Cartesian {
1362 children: vec![
1363 clause(
1364 "a",
1365 Source::IntRange {
1366 lo: 0,
1367 hi: 3,
1368 step: 1,
1369 },
1370 ),
1371 clause(
1372 "b",
1373 Source::Literal {
1374 values: vec![LiteralValue::Int(7), LiteralValue::Int(8)],
1375 },
1376 ),
1377 ],
1378 };
1379 let scope = empty_kernel();
1380
1381 let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
1382 assert_eq!(out.tuples.len(), 6, "3 x 2");
1383 assert_eq!(out.clauses.len(), 2, "one entry per leaf, in tree order");
1384 assert_eq!(out.clauses[0].var, "a");
1385 assert_eq!(out.clauses[0].values, 3);
1386 assert_eq!(out.clauses[1].var, "b");
1387 assert_eq!(out.clauses[1].evaluations, 3);
1390 assert_eq!(out.clauses[1].values, 6);
1391 }
1392
1393 #[test]
1396 fn an_empty_clause_is_reached_and_yields_nothing() {
1397 let comp = Comprehension::Cartesian {
1398 children: vec![
1399 clause(
1400 "a",
1401 Source::IntRange {
1402 lo: 0,
1403 hi: 2,
1404 step: 1,
1405 },
1406 ),
1407 clause("b", empty_literal()),
1408 ],
1409 };
1410 let scope = empty_kernel();
1411
1412 let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
1413 assert!(out.tuples.is_empty(), "an empty clause empties the product");
1414 let culprits: Vec<&str> = out
1415 .clauses
1416 .iter()
1417 .filter(|c| c.evaluations > 0 && c.values == 0)
1418 .map(|c| c.var.as_str())
1419 .collect();
1420 assert_eq!(culprits, ["b"], "only the empty clause is named");
1421 }
1422
1423 #[test]
1426 fn a_clause_behind_an_empty_one_is_never_reached() {
1427 let comp = Comprehension::Cartesian {
1428 children: vec![
1429 clause("outer", empty_literal()),
1430 clause(
1431 "inner",
1432 Source::IntRange {
1433 lo: 0,
1434 hi: 9,
1435 step: 1,
1436 },
1437 ),
1438 ],
1439 };
1440 let scope = empty_kernel();
1441
1442 let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
1443 assert!(out.tuples.is_empty());
1444 let by = |v: &str| {
1445 out.clauses
1446 .iter()
1447 .find(|c| c.var == v)
1448 .expect("every leaf is present whether reached or not")
1449 };
1450 assert_eq!(by("outer").evaluations, 1);
1451 assert_eq!(by("outer").values, 0);
1452 assert_eq!(
1453 by("inner").evaluations,
1454 0,
1455 "never reached: the cause is `outer`, not this"
1456 );
1457 assert_eq!(by("inner").values, 0);
1458 }
1459
1460 #[test]
1463 fn clauses_sharing_a_name_across_a_union_are_counted_apart() {
1464 let comp = Comprehension::Union {
1465 children: vec![
1466 clause(
1467 "k",
1468 Source::Literal {
1469 values: vec![LiteralValue::Int(1)],
1470 },
1471 ),
1472 clause("k", empty_literal()),
1473 ],
1474 };
1475 let scope = empty_kernel();
1476
1477 let out = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
1478 assert_eq!(out.clauses.len(), 2, "two leaves, one name");
1479 assert_eq!(out.clauses[0].values, 1);
1480 assert_eq!(out.clauses[1].values, 0);
1481 assert_eq!(out.clauses[1].evaluations, 1, "reached, and empty");
1482 }
1483
1484 #[test]
1487 fn the_plain_entry_point_agrees_with_the_reported_one() {
1488 let comp = Comprehension::Cartesian {
1489 children: vec![
1490 clause(
1491 "a",
1492 Source::IntRange {
1493 lo: 1,
1494 hi: 4,
1495 step: 1,
1496 },
1497 ),
1498 clause(
1499 "b",
1500 Source::Literal {
1501 values: vec![LiteralValue::Int(5)],
1502 },
1503 ),
1504 ],
1505 };
1506 let scope = empty_kernel();
1507
1508 let plain = evaluate_for_iteration(&comp, &*scope).unwrap();
1509 let reported = evaluate_for_iteration_reported(&comp, &*scope).unwrap();
1510 assert_eq!(plain, reported.tuples);
1511 }
1512
1513 #[test]
1514 fn int_range_yields_values() {
1515 let comp = Comprehension::Clause {
1516 name: "k".into(),
1517 source: Source::IntRange {
1518 lo: 1,
1519 hi: 5,
1520 step: 1,
1521 },
1522 };
1523 let canonical = empty_kernel();
1524
1525 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1526 assert_eq!(tuples.len(), 4);
1527 assert_eq!(tuples[0][0].1, Value::U64(1));
1528 assert_eq!(tuples[3][0].1, Value::U64(4));
1529 }
1530
1531 #[test]
1532 fn literal_list_yields_values() {
1533 let comp = Comprehension::Clause {
1534 name: "x".into(),
1535 source: Source::Literal {
1536 values: vec![LiteralValue::Int(10), LiteralValue::Int(20)],
1537 },
1538 };
1539 let canonical = empty_kernel();
1540
1541 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1542 assert_eq!(tuples.len(), 2);
1543 }
1544
1545 #[test]
1546 fn cartesian_produces_product() {
1547 let comp = Comprehension::cartesian(vec![
1548 Comprehension::Clause {
1549 name: "x".into(),
1550 source: Source::IntRange {
1551 lo: 1,
1552 hi: 3,
1553 step: 1,
1554 },
1555 },
1556 Comprehension::Clause {
1557 name: "y".into(),
1558 source: Source::IntRange {
1559 lo: 10,
1560 hi: 30,
1561 step: 10,
1562 },
1563 },
1564 ]);
1565 let canonical = empty_kernel();
1566
1567 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1568 assert_eq!(tuples.len(), 4);
1570 }
1571
1572 #[test]
1573 fn union_produces_concatenation() {
1574 let comp = Comprehension::union(vec![
1575 Comprehension::Clause {
1576 name: "k".into(),
1577 source: Source::Literal {
1578 values: vec![LiteralValue::Int(1)],
1579 },
1580 },
1581 Comprehension::Clause {
1582 name: "k".into(),
1583 source: Source::Literal {
1584 values: vec![LiteralValue::Int(10), LiteralValue::Int(20)],
1585 },
1586 },
1587 ]);
1588 let canonical = empty_kernel();
1589
1590 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1591 assert_eq!(tuples.len(), 3);
1592 }
1593
1594 #[test]
1595 fn filter_drops_non_matching() {
1596 let comp = Comprehension::filter(
1597 Comprehension::Clause {
1598 name: "k".into(),
1599 source: Source::IntRange {
1600 lo: 1,
1601 hi: 6,
1602 step: 1,
1603 },
1604 },
1605 "{k} > 3",
1606 );
1607 let canonical = canonical_with_k();
1608
1609 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1610 assert_eq!(tuples.len(), 2);
1612 }
1613
1614 #[test]
1615 fn order_lex_truncate() {
1616 let comp = Comprehension::order(
1617 Comprehension::Clause {
1618 name: "k".into(),
1619 source: Source::IntRange {
1620 lo: 1,
1621 hi: 100,
1622 step: 1,
1623 },
1624 },
1625 StrategyName::Lex,
1626 Some(5),
1627 );
1628 let canonical = empty_kernel();
1629
1630 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1631 assert_eq!(tuples.len(), 5);
1632 }
1633
1634 #[test]
1640 fn extrema_over_cartesian_uses_indexed_form() {
1641 let comp = Comprehension::order(
1642 Comprehension::cartesian(vec![
1643 Comprehension::Clause {
1644 name: "k".into(),
1645 source: Source::Literal {
1646 values: vec![
1647 LiteralValue::Int(1),
1648 LiteralValue::Int(2),
1649 LiteralValue::Int(3),
1650 ],
1651 },
1652 },
1653 Comprehension::Clause {
1654 name: "limit".into(),
1655 source: Source::Literal {
1656 values: vec![
1657 LiteralValue::Int(10),
1658 LiteralValue::Int(20),
1659 LiteralValue::Int(30),
1660 ],
1661 },
1662 },
1663 ]),
1664 StrategyName::Extrema,
1665 Some(1),
1669 );
1670 let canonical = empty_kernel();
1671
1672 let tuples = evaluate_for_iteration(&comp, &*canonical).unwrap();
1673 assert_eq!(tuples.len(), 4);
1675 for t in &tuples {
1677 assert_eq!(t.len(), 2);
1678 let k = match &t[0].1 {
1679 Value::U64(n) => *n,
1680 other => panic!("expected u64 k, got {other:?}"),
1681 };
1682 let lim = match &t[1].1 {
1683 Value::U64(n) => *n,
1684 other => panic!("expected u64 limit, got {other:?}"),
1685 };
1686 assert!(k == 1 || k == 3, "expected extreme k, got {k}");
1687 assert!(lim == 10 || lim == 30, "expected extreme limit, got {lim}");
1688 }
1689 }
1690}