use std::sync::Arc;
use crate::iteration::comprehension::ast::Comprehension;
use crate::iteration::comprehension::eval_source::{EvalClass, SourceEval};
use crate::iteration::comprehension::flatten::flatten_static_sources;
use crate::iteration::comprehension::ir::{Program, compile as compile_to_ir};
use crate::iteration::comprehension::optimize::optimize;
use crate::iteration::comprehension::predicate::recognizers::extract_coord_refs;
use crate::iteration::comprehension::validate::{
Mode, Surface, ValidationError, ValidationReport, ValidationWarning, unresolved_names, validate,
};
use crate::kernel::interp::NoScope;
use super::coord_stream::CoordinateStream;
use super::instance::{KernelScope, ScopedKernelInstance};
use super::scope_once::scope_once_with;
use super::scoped_stream::ScopedKernelStream;
#[derive(Debug, Clone)]
pub struct CompiledComprehension {
program: Arc<Program>,
}
impl CompiledComprehension {
pub fn from_ast(ast: &Comprehension) -> Result<Self, ValidationError> {
Self::from_ast_with(ast, Mode::Permissive).map(|(compiled, _)| compiled)
}
pub fn from_ast_with(
ast: &Comprehension,
mode: Mode,
) -> Result<(Self, ValidationReport), ValidationError> {
Self::from_ast_in(ast, mode, &|_| false)
}
pub fn from_ast_in(
ast: &Comprehension,
mode: Mode,
in_scope: &dyn Fn(&str) -> bool,
) -> Result<(Self, ValidationReport), ValidationError> {
let ast = flatten_static_sources(ast, &NoScope::new());
let unresolved = unresolved_names(&ast, Surface::Traversal(in_scope));
if mode == Mode::Strict && !unresolved.is_empty() {
return Err(ValidationError::V3UnresolvedNames { reads: unresolved });
}
if let Some((name, references)) = first_context_required(&ast, in_scope) {
return Err(ValidationError::ContextRequired { name, references });
}
if let Some((predicate, references)) = first_unbound_predicate(&ast, in_scope) {
return Err(ValidationError::PredicateContextRequired {
predicate,
references,
});
}
if let Some((name, message)) = first_failed_static(&ast) {
return Err(ValidationError::SourceFailed { name, message });
}
let mut report = validate(&ast, mode)?;
if !unresolved.is_empty() {
report
.warnings
.insert(0, ValidationWarning::UnresolvedNames { reads: unresolved });
}
Ok((
Self {
program: Arc::new(compile_to_ir(&optimize(ast))),
},
report,
))
}
pub fn from_program(program: Arc<Program>) -> Self {
Self { program }
}
pub fn program(&self) -> &Program {
&self.program
}
pub(crate) fn program_arc(&self) -> Arc<Program> {
Arc::clone(&self.program)
}
pub fn coordinate_stream(&self) -> CoordinateStream {
CoordinateStream::new(self.program_arc())
}
pub fn scoped_kernel_stream<K: KernelScope>(&self, parent: K) -> ScopedKernelStream<K> {
ScopedKernelStream::new(self.program_arc(), parent)
}
pub fn scope_once<K: KernelScope>(
&self,
parent: &K,
coords: &crate::iteration::comprehension::strategies::Tuple,
) -> ScopedKernelInstance<K::Scoped> {
scope_once_with(parent, coords)
}
}
fn first_context_required(
ast: &Comprehension,
in_scope: &dyn Fn(&str) -> bool,
) -> Option<(String, Vec<String>)> {
fn walk(
c: &Comprehension,
in_scope: &dyn Fn(&str) -> bool,
before: &mut Vec<String>,
) -> Option<(String, Vec<String>)> {
match c {
Comprehension::Clause { name, source } => {
let references = source.referenced_names();
let reads_none = references
.iter()
.any(|n| !before.contains(n) && !in_scope(n));
(source.eval_class() == EvalClass::ContextRequired && !reads_none)
.then(|| (name.clone(), references.into_iter().collect()))
}
Comprehension::Cartesian { children } => {
let depth = before.len();
let mut found = None;
for child in children {
found = walk(child, in_scope, before);
if found.is_some() {
break;
}
before.extend(child.coordinate_names());
}
before.truncate(depth);
found
}
Comprehension::Zip { children, .. } | Comprehension::Union { children } => children
.iter()
.find_map(|child| walk(child, in_scope, before)),
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
walk(child, in_scope, before)
}
}
}
walk(ast, in_scope, &mut Vec::new())
}
fn first_unbound_predicate(
ast: &Comprehension,
in_scope: &dyn Fn(&str) -> bool,
) -> Option<(String, Vec<String>)> {
match ast {
Comprehension::Clause { .. } => None,
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => children
.iter()
.find_map(|child| first_unbound_predicate(child, in_scope)),
Comprehension::Filter { child, predicate } => {
let bound = child.coordinate_names();
let unbound: Vec<String> = extract_coord_refs(predicate)
.into_iter()
.filter(|name| !bound.contains(name) && in_scope(name))
.collect();
if unbound.is_empty() {
first_unbound_predicate(child, in_scope)
} else {
Some((predicate.clone(), unbound))
}
}
Comprehension::Order { child, .. } => first_unbound_predicate(child, in_scope),
}
}
fn first_failed_static(ast: &Comprehension) -> Option<(String, String)> {
use crate::iteration::comprehension::eval_source::EvalContext;
use crate::iteration::comprehension::source::Source;
match ast {
Comprehension::Clause { name, source } => match source {
Source::Generator {
cardinality_hint: None,
..
} if source.eval_class() == EvalClass::Static => {
let scope = NoScope::new();
let ctx = EvalContext {
var_name: name,
scope: &scope,
prefix: &[],
};
source
.evaluate(Some(&ctx))
.err()
.map(|e| (name.clone(), e.to_string()))
}
_ => None,
},
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => children.iter().find_map(first_failed_static),
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
first_failed_static(child)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::iteration::comprehension::source::{LiteralValue, Source};
fn clause(name: &str, vs: &[i64]) -> Comprehension {
Comprehension::clause(
name,
Source::Literal {
values: vs.iter().map(|n| LiteralValue::Int(*n)).collect(),
},
)
}
#[test]
fn from_ast_compiles_once() {
let ast = clause("k", &[1, 2, 3]);
let compiled = CompiledComprehension::from_ast(&ast).unwrap();
assert!(!compiled.program().is_empty());
}
#[test]
fn from_ast_compiles_the_optimized_tree() {
let inner = Comprehension::cartesian(vec![clause("a", &[1, 2]), clause("b", &[3])]);
let ast = Comprehension::cartesian(vec![inner, clause("c", &[4])]);
let compiled = CompiledComprehension::from_ast(&ast).unwrap();
assert_eq!(*compiled.program(), compile_to_ir(&optimize(ast.clone())));
assert_ne!(*compiled.program(), compile_to_ir(&ast));
}
#[test]
fn from_ast_refuses_a_tree_that_violates_a_v_axiom() {
let ast = Comprehension::cartesian(vec![clause("k", &[1, 2]), clause("k", &[3, 4])]);
let err = CompiledComprehension::from_ast(&ast).unwrap_err();
assert!(
matches!(err, ValidationError::V1DuplicateName { ref name, .. } if name == "k"),
"{err}"
);
assert!(err.to_string().starts_with("V1:"), "{err}");
}
#[test]
fn from_ast_with_reports_or_refuses_a_degenerate_composition() {
use crate::iteration::comprehension::strategy::StrategyName;
use crate::iteration::comprehension::validate::ValidationWarning;
let ast = Comprehension::order(clause("k", &[1, 2, 3]), StrategyName::Extrema, Some(1));
let (_, report) = CompiledComprehension::from_ast_with(&ast, Mode::Permissive).unwrap();
assert!(matches!(
report.warnings.as_slice(),
[ValidationWarning::DegenerateGeometric { .. }]
));
let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
assert!(matches!(err, ValidationError::StrictWarning(_)), "{err}");
assert!(err.to_string().starts_with("strict mode:"), "{err}");
}
#[test]
fn cloning_compiled_shares_arc() {
let ast = clause("k", &[1, 2, 3]);
let a = CompiledComprehension::from_ast(&ast).unwrap();
let b = a.clone();
let count = Arc::strong_count(&a.program);
assert!(count >= 2, "expected shared Arc, count = {count}");
drop(b);
}
#[test]
fn two_coordinate_streams_share_program() {
let ast = clause("k", &[1, 2, 3]);
let compiled = CompiledComprehension::from_ast(&ast).unwrap();
let _s1 = compiled.coordinate_stream();
let _s2 = compiled.coordinate_stream();
let count = Arc::strong_count(&compiled.program);
assert!(
count >= 3,
"expected shared program across streamers, count = {count}"
);
}
#[test]
fn from_ast_refuses_a_context_required_source_by_name() {
let ast = Comprehension::cartesian(vec![
clause("k", &[1, 2]),
Comprehension::clause(
"j",
Source::Generator {
expr: "pow2({n})".into(),
cardinality_hint: None,
},
),
]);
let err =
CompiledComprehension::from_ast_in(&ast, Mode::Permissive, &|n| n == "n").unwrap_err();
assert!(
matches!(
err,
ValidationError::ContextRequired { ref name, ref references }
if name == "j" && references == &["n".to_string()]
),
"{err}"
);
assert!(err.to_string().contains("traverse it with `for`"), "{err}");
let (compiled, report) = CompiledComprehension::from_ast_with(&ast, Mode::Permissive)
.unwrap_or_else(|e| panic!("{e}"));
assert!(
matches!(report.warnings.as_slice(),
[ValidationWarning::UnresolvedNames { reads }] if reads.len() == 1),
"{:?}",
report.warnings
);
assert_eq!(compiled.coordinate_stream().count(), 0);
let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
assert!(
matches!(err, ValidationError::V3UnresolvedNames { ref reads } if reads.len() == 1),
"{err}"
);
}
#[test]
fn from_ast_refuses_a_predicate_that_needs_a_scope() {
let ast = Comprehension::filter(clause("k", &[1, 2, 3]), "{k} == 1 || {k} > {limit}");
let err = CompiledComprehension::from_ast_in(&ast, Mode::Permissive, &|n| n == "limit")
.unwrap_err();
assert!(
matches!(
err,
ValidationError::PredicateContextRequired { ref references, .. }
if references == &["limit".to_string()]
),
"{err}"
);
assert!(err.to_string().contains("traverse it with `for`"), "{err}");
let compiled = CompiledComprehension::from_ast(&ast).unwrap_or_else(|e| panic!("{e}"));
let kept: Vec<_> = compiled
.coordinate_stream()
.collect::<Result<_, _>>()
.unwrap();
assert_eq!(kept.len(), 1);
assert_eq!(
kept[0].bindings[0].1,
crate::iteration::comprehension::strategies::TupleValue::I64(1)
);
let err = CompiledComprehension::from_ast_with(&ast, Mode::Strict).unwrap_err();
assert!(
matches!(err, ValidationError::V3UnresolvedNames { .. }),
"{err}"
);
let bound = Comprehension::filter(clause("k", &[1, 2, 3]), "{k} > 1");
assert!(CompiledComprehension::from_ast(&bound).is_ok());
}
}