1use std::sync::Arc;
15
16use crate::iteration::comprehension::ast::Comprehension;
17use crate::iteration::comprehension::eval_source::{EvalClass, SourceEval};
18use crate::iteration::comprehension::flatten::flatten_static_sources;
19use crate::iteration::comprehension::ir::{Program, compile as compile_to_ir};
20use crate::iteration::comprehension::optimize::optimize;
21use crate::iteration::comprehension::predicate::recognizers::extract_coord_refs;
22use crate::iteration::comprehension::validate::{
23 Mode, Surface, ValidationError, ValidationReport, ValidationWarning, unresolved_names, validate,
24};
25
26use crate::kernel::interp::NoScope;
27
28use super::coord_stream::CoordinateStream;
29use super::instance::{KernelScope, ScopedKernelInstance};
30use super::scope_once::scope_once_with;
31use super::scoped_stream::ScopedKernelStream;
32
33#[derive(Debug, Clone)]
42pub struct CompiledComprehension {
43 program: Arc<Program>,
44}
45
46impl CompiledComprehension {
47 pub fn from_ast(ast: &Comprehension) -> Result<Self, ValidationError> {
54 Self::from_ast_with(ast, Mode::Permissive).map(|(compiled, _)| compiled)
55 }
56
57 pub fn from_ast_with(
73 ast: &Comprehension,
74 mode: Mode,
75 ) -> Result<(Self, ValidationReport), ValidationError> {
76 Self::from_ast_in(ast, mode, &|_| false)
77 }
78
79 pub fn from_ast_in(
96 ast: &Comprehension,
97 mode: Mode,
98 in_scope: &dyn Fn(&str) -> bool,
99 ) -> Result<(Self, ValidationReport), ValidationError> {
100 let ast = flatten_static_sources(ast, &NoScope::new());
101 let unresolved = unresolved_names(&ast, Surface::Traversal(in_scope));
104 if mode == Mode::Strict && !unresolved.is_empty() {
105 return Err(ValidationError::V3UnresolvedNames { reads: unresolved });
106 }
107 if let Some((name, references)) = first_context_required(&ast, in_scope) {
108 return Err(ValidationError::ContextRequired { name, references });
109 }
110 if let Some((predicate, references)) = first_unbound_predicate(&ast, in_scope) {
111 return Err(ValidationError::PredicateContextRequired {
112 predicate,
113 references,
114 });
115 }
116 if let Some((name, message)) = first_failed_static(&ast) {
117 return Err(ValidationError::SourceFailed { name, message });
118 }
119 let mut report = validate(&ast, mode)?;
120 if !unresolved.is_empty() {
121 report
122 .warnings
123 .insert(0, ValidationWarning::UnresolvedNames { reads: unresolved });
124 }
125 Ok((
126 Self {
127 program: Arc::new(compile_to_ir(&optimize(ast))),
128 },
129 report,
130 ))
131 }
132
133 pub fn from_program(program: Arc<Program>) -> Self {
136 Self { program }
137 }
138
139 pub fn program(&self) -> &Program {
142 &self.program
143 }
144
145 pub(crate) fn program_arc(&self) -> Arc<Program> {
148 Arc::clone(&self.program)
149 }
150
151 pub fn coordinate_stream(&self) -> CoordinateStream {
158 CoordinateStream::new(self.program_arc())
159 }
160
161 pub fn scoped_kernel_stream<K: KernelScope>(&self, parent: K) -> ScopedKernelStream<K> {
173 ScopedKernelStream::new(self.program_arc(), parent)
174 }
175
176 pub fn scope_once<K: KernelScope>(
184 &self,
185 parent: &K,
186 coords: &crate::iteration::comprehension::strategies::Tuple,
187 ) -> ScopedKernelInstance<K::Scoped> {
188 scope_once_with(parent, coords)
189 }
190}
191
192fn first_context_required(
199 ast: &Comprehension,
200 in_scope: &dyn Fn(&str) -> bool,
201) -> Option<(String, Vec<String>)> {
202 fn walk(
203 c: &Comprehension,
204 in_scope: &dyn Fn(&str) -> bool,
205 before: &mut Vec<String>,
206 ) -> Option<(String, Vec<String>)> {
207 match c {
208 Comprehension::Clause { name, source } => {
209 let references = source.referenced_names();
210 let reads_none = references
211 .iter()
212 .any(|n| !before.contains(n) && !in_scope(n));
213 (source.eval_class() == EvalClass::ContextRequired && !reads_none)
214 .then(|| (name.clone(), references.into_iter().collect()))
215 }
216 Comprehension::Cartesian { children } => {
217 let depth = before.len();
218 let mut found = None;
219 for child in children {
220 found = walk(child, in_scope, before);
221 if found.is_some() {
222 break;
223 }
224 before.extend(child.coordinate_names());
225 }
226 before.truncate(depth);
227 found
228 }
229 Comprehension::Zip { children, .. } | Comprehension::Union { children } => children
230 .iter()
231 .find_map(|child| walk(child, in_scope, before)),
232 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
233 walk(child, in_scope, before)
234 }
235 }
236 }
237 walk(ast, in_scope, &mut Vec::new())
238}
239
240fn first_unbound_predicate(
246 ast: &Comprehension,
247 in_scope: &dyn Fn(&str) -> bool,
248) -> Option<(String, Vec<String>)> {
249 match ast {
250 Comprehension::Clause { .. } => None,
251 Comprehension::Cartesian { children }
252 | Comprehension::Zip { children, .. }
253 | Comprehension::Union { children } => children
254 .iter()
255 .find_map(|child| first_unbound_predicate(child, in_scope)),
256 Comprehension::Filter { child, predicate } => {
257 let bound = child.coordinate_names();
258 let unbound: Vec<String> = extract_coord_refs(predicate)
259 .into_iter()
260 .filter(|name| !bound.contains(name) && in_scope(name))
261 .collect();
262 if unbound.is_empty() {
263 first_unbound_predicate(child, in_scope)
264 } else {
265 Some((predicate.clone(), unbound))
266 }
267 }
268 Comprehension::Order { child, .. } => first_unbound_predicate(child, in_scope),
269 }
270}
271
272fn first_failed_static(ast: &Comprehension) -> Option<(String, String)> {
280 use crate::iteration::comprehension::eval_source::EvalContext;
281 use crate::iteration::comprehension::source::Source;
282 match ast {
283 Comprehension::Clause { name, source } => match source {
284 Source::Generator {
285 cardinality_hint: None,
286 ..
287 } if source.eval_class() == EvalClass::Static => {
288 let scope = NoScope::new();
289 let ctx = EvalContext {
290 var_name: name,
291 scope: &scope,
292 prefix: &[],
293 };
294 source
295 .evaluate(Some(&ctx))
296 .err()
297 .map(|e| (name.clone(), e.to_string()))
298 }
299 _ => None,
300 },
301 Comprehension::Cartesian { children }
302 | Comprehension::Zip { children, .. }
303 | Comprehension::Union { children } => children.iter().find_map(first_failed_static),
304 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
305 first_failed_static(child)
306 }
307 }
308}
309
310#[cfg(test)]
311mod tests {
312 use super::*;
313 use crate::iteration::comprehension::source::{LiteralValue, Source};
314
315 fn clause(name: &str, vs: &[i64]) -> Comprehension {
316 Comprehension::clause(
317 name,
318 Source::Literal {
319 values: vs.iter().map(|n| LiteralValue::Int(*n)).collect(),
320 },
321 )
322 }
323
324 #[test]
325 fn from_ast_compiles_once() {
326 let ast = clause("k", &[1, 2, 3]);
327 let compiled = CompiledComprehension::from_ast(&ast).unwrap();
328 assert!(!compiled.program().is_empty());
329 }
330
331 #[test]
337 fn from_ast_compiles_the_optimized_tree() {
338 let inner = Comprehension::cartesian(vec![clause("a", &[1, 2]), clause("b", &[3])]);
341 let ast = Comprehension::cartesian(vec![inner, clause("c", &[4])]);
342 let compiled = CompiledComprehension::from_ast(&ast).unwrap();
343 assert_eq!(*compiled.program(), compile_to_ir(&optimize(ast.clone())));
344 assert_ne!(*compiled.program(), compile_to_ir(&ast));
345 }
346
347 #[test]
351 fn from_ast_refuses_a_tree_that_violates_a_v_axiom() {
352 let ast = Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("k", &[3, 4])]);
354 let err = CompiledComprehension::from_ast(&ast).unwrap_err();
355 assert!(
356 matches!(err, ValidationError::V1DuplicateName { ref name, .. } if name == "k"),
357 "{err}"
358 );
359 assert!(err.to_string().starts_with("V1:"), "{err}");
360 }
361
362 #[test]
366 fn from_ast_with_reports_or_refuses_a_degenerate_composition() {
367 use crate::iteration::comprehension::strategy::StrategyName;
368 use crate::iteration::comprehension::validate::ValidationWarning;
369 let ast = Comprehension::order(clause("k", &[1, 2, 3]), StrategyName::Extrema, Some(1));
370 let (_, report) = CompiledComprehension::from_ast_with(&ast, Mode::Permissive).unwrap();
371 assert!(matches!(
372 report.warnings.as_slice(),
373 [ValidationWarning::DegenerateGeometric { .. }]
374 ));
375 let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
376 assert!(matches!(err, ValidationError::StrictWarning(_)), "{err}");
377 assert!(err.to_string().starts_with("strict mode:"), "{err}");
378 }
379
380 #[test]
381 fn cloning_compiled_shares_arc() {
382 let ast = clause("k", &[1, 2, 3]);
383 let a = CompiledComprehension::from_ast(&ast).unwrap();
384 let b = a.clone();
385 let count = Arc::strong_count(&a.program);
387 assert!(count >= 2, "expected shared Arc, count = {count}");
388 drop(b);
389 }
390
391 #[test]
392 fn two_coordinate_streams_share_program() {
393 let ast = clause("k", &[1, 2, 3]);
394 let compiled = CompiledComprehension::from_ast(&ast).unwrap();
395 let _s1 = compiled.coordinate_stream();
396 let _s2 = compiled.coordinate_stream();
397 let count = Arc::strong_count(&compiled.program);
400 assert!(
401 count >= 3,
402 "expected shared program across streamers, count = {count}"
403 );
404 }
405
406 #[test]
413 fn from_ast_refuses_a_context_required_source_by_name() {
414 let ast = Comprehension::cartesian(vec![
415 clause("k", &[1, 2]),
416 Comprehension::clause(
417 "j",
418 Source::Generator {
419 expr: "pow2({n})".into(),
420 cardinality_hint: None,
421 },
422 ),
423 ]);
424 let err =
425 CompiledComprehension::from_ast_in(&ast, Mode::Permissive, &|n| n == "n").unwrap_err();
426 assert!(
427 matches!(
428 err,
429 ValidationError::ContextRequired { ref name, ref references }
430 if name == "j" && references == &["n".to_string()]
431 ),
432 "{err}"
433 );
434 assert!(err.to_string().contains("traverse it with `for`"), "{err}");
435 let (compiled, report) = CompiledComprehension::from_ast_with(&ast, Mode::Permissive)
436 .unwrap_or_else(|e| panic!("{e}"));
437 assert!(
438 matches!(report.warnings.as_slice(),
439 [ValidationWarning::UnresolvedNames { reads }] if reads.len() == 1),
440 "{:?}",
441 report.warnings
442 );
443 assert_eq!(compiled.coordinate_stream().count(), 0);
444 let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
445 assert!(
446 matches!(err, ValidationError::V3UnresolvedNames { ref reads } if reads.len() == 1),
447 "{err}"
448 );
449 }
450
451 #[test]
458 fn from_ast_refuses_a_predicate_that_needs_a_scope() {
459 let ast = Comprehension::filter(clause("k", &[1, 2, 3]), "{k} == 1 || {k} > {limit}");
460 let err = CompiledComprehension::from_ast_in(&ast, Mode::Permissive, &|n| n == "limit")
461 .unwrap_err();
462 assert!(
463 matches!(
464 err,
465 ValidationError::PredicateContextRequired { ref references, .. }
466 if references == &["limit".to_string()]
467 ),
468 "{err}"
469 );
470 assert!(err.to_string().contains("traverse it with `for`"), "{err}");
471 let compiled = CompiledComprehension::from_ast(&ast).unwrap_or_else(|e| panic!("{e}"));
472 let kept: Vec<_> = compiled
473 .coordinate_stream()
474 .collect::<Result<_, _>>()
475 .unwrap();
476 assert_eq!(kept.len(), 1);
477 assert_eq!(
478 kept[0].bindings[0].1,
479 crate::iteration::comprehension::strategies::TupleValue::I64(1)
480 );
481 let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
482 assert!(
483 matches!(err, ValidationError::V3UnresolvedNames { .. }),
484 "{err}"
485 );
486 let bound = Comprehension::filter(clause("k", &[1, 2, 3]), "{k} > 1");
487 assert!(CompiledComprehension::from_ast(&bound).is_ok());
488 }
489}