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(
200 ast: &Comprehension,
201 in_scope: &dyn Fn(&str) -> bool,
202) -> Option<(String, Vec<String>)> {
203 fn walk(
204 c: &Comprehension,
205 in_scope: &dyn Fn(&str) -> bool,
206 before: &mut Vec<String>,
207 ) -> Option<(String, Vec<String>)> {
208 match c {
209 Comprehension::Clause { name, source } => {
210 let references = source.names_read();
211 let reads_none = references
212 .iter()
213 .any(|n| !before.contains(n) && !in_scope(n));
214 (source.eval_class() == EvalClass::ContextRequired && !reads_none)
215 .then(|| (name.clone(), references.into_iter().collect()))
216 }
217 Comprehension::Cartesian { children } => {
218 let depth = before.len();
219 let mut found = None;
220 for child in children {
221 found = walk(child, in_scope, before);
222 if found.is_some() {
223 break;
224 }
225 before.extend(child.coordinate_names());
226 }
227 before.truncate(depth);
228 found
229 }
230 Comprehension::Zip { children, .. } | Comprehension::Union { children } => children
231 .iter()
232 .find_map(|child| walk(child, in_scope, before)),
233 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
234 walk(child, in_scope, before)
235 }
236 }
237 }
238 walk(ast, in_scope, &mut Vec::new())
239}
240
241fn first_unbound_predicate(
247 ast: &Comprehension,
248 in_scope: &dyn Fn(&str) -> bool,
249) -> Option<(String, Vec<String>)> {
250 match ast {
251 Comprehension::Clause { .. } => None,
252 Comprehension::Cartesian { children }
253 | Comprehension::Zip { children, .. }
254 | Comprehension::Union { children } => children
255 .iter()
256 .find_map(|child| first_unbound_predicate(child, in_scope)),
257 Comprehension::Filter { child, predicate } => {
258 let bound = child.coordinate_names();
259 let unbound: Vec<String> = extract_coord_refs(predicate)
260 .into_iter()
261 .filter(|name| !bound.contains(name) && in_scope(name))
262 .collect();
263 if unbound.is_empty() {
264 first_unbound_predicate(child, in_scope)
265 } else {
266 Some((predicate.clone(), unbound))
267 }
268 }
269 Comprehension::Order { child, .. } => first_unbound_predicate(child, in_scope),
270 }
271}
272
273fn first_failed_static(ast: &Comprehension) -> Option<(String, String)> {
281 use crate::iteration::comprehension::eval_source::EvalContext;
282 use crate::iteration::comprehension::source::Source;
283 match ast {
284 Comprehension::Clause { name, source } => match source {
285 Source::Generator {
286 cardinality_hint: None,
287 ..
288 } if source.eval_class() == EvalClass::Static => {
289 let scope = NoScope::new();
290 let ctx = EvalContext {
291 var_name: name,
292 scope: &scope,
293 prefix: &[],
294 };
295 source
296 .evaluate(Some(&ctx))
297 .err()
298 .map(|e| (name.clone(), e.to_string()))
299 }
300 _ => None,
301 },
302 Comprehension::Cartesian { children }
303 | Comprehension::Zip { children, .. }
304 | Comprehension::Union { children } => children.iter().find_map(first_failed_static),
305 Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
306 first_failed_static(child)
307 }
308 }
309}
310
311#[cfg(test)]
312mod tests {
313 use super::*;
314 use crate::iteration::comprehension::source::{LiteralValue, Source};
315
316 fn clause(name: &str, vs: &[i64]) -> Comprehension {
317 Comprehension::clause(
318 name,
319 Source::Literal {
320 values: vs.iter().map(|n| LiteralValue::Int(*n)).collect(),
321 },
322 )
323 }
324
325 #[test]
326 fn from_ast_compiles_once() {
327 let ast = clause("k", &[1, 2, 3]);
328 let compiled = CompiledComprehension::from_ast(&ast).unwrap();
329 assert!(!compiled.program().is_empty());
330 }
331
332 #[test]
338 fn from_ast_compiles_the_optimized_tree() {
339 let inner = Comprehension::cartesian(vec![clause("a", &[1, 2]), clause("b", &[3])]);
342 let ast = Comprehension::cartesian(vec![inner, clause("c", &[4])]);
343 let compiled = CompiledComprehension::from_ast(&ast).unwrap();
344 assert_eq!(*compiled.program(), compile_to_ir(&optimize(ast.clone())));
345 assert_ne!(*compiled.program(), compile_to_ir(&ast));
346 }
347
348 #[test]
352 fn from_ast_refuses_a_tree_that_violates_a_v_axiom() {
353 let ast = Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("k", &[3, 4])]);
355 let err = CompiledComprehension::from_ast(&ast).unwrap_err();
356 assert!(
357 matches!(err, ValidationError::V1DuplicateName { ref name, .. } if name == "k"),
358 "{err}"
359 );
360 assert!(err.to_string().starts_with("V1:"), "{err}");
361 }
362
363 #[test]
367 fn from_ast_with_reports_or_refuses_a_degenerate_composition() {
368 use crate::iteration::comprehension::strategy::StrategyName;
369 use crate::iteration::comprehension::validate::ValidationWarning;
370 let ast = Comprehension::order(clause("k", &[1, 2, 3]), StrategyName::Extrema, Some(1));
371 let (_, report) = CompiledComprehension::from_ast_with(&ast, Mode::Permissive).unwrap();
372 assert!(matches!(
373 report.warnings.as_slice(),
374 [ValidationWarning::DegenerateGeometric { .. }]
375 ));
376 let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
377 assert!(matches!(err, ValidationError::StrictWarning(_)), "{err}");
378 assert!(err.to_string().starts_with("strict mode:"), "{err}");
379 }
380
381 #[test]
382 fn cloning_compiled_shares_arc() {
383 let ast = clause("k", &[1, 2, 3]);
384 let a = CompiledComprehension::from_ast(&ast).unwrap();
385 let b = a.clone();
386 let count = Arc::strong_count(&a.program);
388 assert!(count >= 2, "expected shared Arc, count = {count}");
389 drop(b);
390 }
391
392 #[test]
393 fn two_coordinate_streams_share_program() {
394 let ast = clause("k", &[1, 2, 3]);
395 let compiled = CompiledComprehension::from_ast(&ast).unwrap();
396 let _s1 = compiled.coordinate_stream();
397 let _s2 = compiled.coordinate_stream();
398 let count = Arc::strong_count(&compiled.program);
401 assert!(
402 count >= 3,
403 "expected shared program across streamers, count = {count}"
404 );
405 }
406
407 #[test]
414 fn from_ast_refuses_a_context_required_source_by_name() {
415 let ast = Comprehension::cartesian(vec![
416 clause("k", &[1, 2]),
417 Comprehension::clause(
418 "j",
419 Source::Generator {
420 expr: "pow2({n})".into(),
421 cardinality_hint: None,
422 },
423 ),
424 ]);
425 let err =
426 CompiledComprehension::from_ast_in(&ast, Mode::Permissive, &|n| n == "n").unwrap_err();
427 assert!(
428 matches!(
429 err,
430 ValidationError::ContextRequired { ref name, ref references }
431 if name == "j" && references == &["n".to_string()]
432 ),
433 "{err}"
434 );
435 assert!(err.to_string().contains("traverse it with `for`"), "{err}");
436 let (compiled, report) = CompiledComprehension::from_ast_with(&ast, Mode::Permissive)
437 .unwrap_or_else(|e| panic!("{e}"));
438 assert!(
439 matches!(report.warnings.as_slice(),
440 [ValidationWarning::UnresolvedNames { reads }] if reads.len() == 1),
441 "{:?}",
442 report.warnings
443 );
444 assert_eq!(compiled.coordinate_stream().count(), 0);
445 let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
446 assert!(
447 matches!(err, ValidationError::V3UnresolvedNames { ref reads } if reads.len() == 1),
448 "{err}"
449 );
450 }
451
452 #[test]
459 fn from_ast_refuses_a_predicate_that_needs_a_scope() {
460 let ast = Comprehension::filter(clause("k", &[1, 2, 3]), "{k} == 1 || {k} > {limit}");
461 let err = CompiledComprehension::from_ast_in(&ast, Mode::Permissive, &|n| n == "limit")
462 .unwrap_err();
463 assert!(
464 matches!(
465 err,
466 ValidationError::PredicateContextRequired { ref references, .. }
467 if references == &["limit".to_string()]
468 ),
469 "{err}"
470 );
471 assert!(err.to_string().contains("traverse it with `for`"), "{err}");
472 let compiled = CompiledComprehension::from_ast(&ast).unwrap_or_else(|e| panic!("{e}"));
473 let kept: Vec<_> = compiled
474 .coordinate_stream()
475 .collect::<Result<_, _>>()
476 .unwrap();
477 assert_eq!(kept.len(), 1);
478 assert_eq!(
479 kept[0].bindings[0].1,
480 crate::iteration::comprehension::strategies::TupleValue::I64(1)
481 );
482 let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
483 assert!(
484 matches!(err, ValidationError::V3UnresolvedNames { .. }),
485 "{err}"
486 );
487 let bound = Comprehension::filter(clause("k", &[1, 2, 3]), "{k} > 1");
488 assert!(CompiledComprehension::from_ast(&bound).is_ok());
489 }
490}