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::validate::{
Mode, ValidationError, ValidationReport, 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> {
let ast = flatten_static_sources(ast, &NoScope::new());
if let Some((name, references)) = first_context_required(&ast) {
return Err(ValidationError::ContextRequired { name, references });
}
if let Some((name, message)) = first_failed_static(&ast) {
return Err(ValidationError::SourceFailed { name, message });
}
let report = validate(&ast, mode)?;
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) -> Option<(String, Vec<String>)> {
match ast {
Comprehension::Clause { name, source } => {
(source.eval_class() == EvalClass::ContextRequired).then(|| {
(
name.clone(),
source.referenced_names().into_iter().collect(),
)
})
}
Comprehension::Cartesian { children }
| Comprehension::Zip { children, .. }
| Comprehension::Union { children } => children.iter().find_map(first_context_required),
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
first_context_required(child)
}
}
}
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(&ast).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}");
}
}