use std::collections::{BTreeSet, HashMap};
use std::sync::{Arc, Mutex};
use crate::ast::PortType;
use crate::iteration::comprehension::source::{LiteralValue, Source};
use crate::iteration::comprehension::{Comprehension, StreamerValue};
use crate::kernel::PolydatProgram;
use super::ast::{
Arg, Binding, BindingModifier, CallExpr, Expr, ExternPort, ForSource, ForSourceKind, ForStmt,
InputDecl, PolydatFile, Statement,
};
use super::lexer::Span;
#[derive(Debug, Clone)]
pub struct Traversal {
pub span: Span,
pub source_text: String,
pub comprehension: Comprehension,
pub elements: Vec<(String, PortType)>,
pub cascade: Vec<(String, PortType)>,
pub program: Arc<PolydatProgram>,
pub body: Arc<BodySource>,
}
pub struct BodySource {
pub(crate) file: PolydatFile,
pub(crate) source_text: String,
pub(crate) source_dir: Option<std::path::PathBuf>,
pub(crate) lib_paths: Vec<std::path::PathBuf>,
pub(crate) strict: bool,
pub(crate) context_label: String,
pub(crate) cursor_limit: Option<u64>,
pub(crate) pragmas: super::pragmas::PragmaSet,
pub(super) modules: HashMap<String, super::modules::ResolvedModule>,
pub(crate) programs: Mutex<HashMap<crate::Engine, Arc<dyn crate::kernel::KernelProgram>>>,
pub(crate) ledger: Arc<crate::kernel::CompileLedger>,
}
impl BodySource {
pub fn source_text(&self) -> &str {
&self.source_text
}
}
impl std::fmt::Debug for BodySource {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BodySource")
.field("context", &self.context_label)
.field("statements", &self.file.statements.len())
.finish()
}
}
impl Traversal {
pub fn program_on(
&self,
engine: crate::Engine,
) -> Result<Arc<dyn crate::kernel::KernelProgram>, crate::KernelError> {
if matches!(engine, crate::Engine::Interpreter(_)) {
return Ok(self.program.clone());
}
let mut programs = self
.body
.programs
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
if let Some(program) = programs.get(&engine) {
return Ok(program.clone());
}
let program = super::compile::Compiler::compile_body_on(&self.body, engine)?.into_program();
programs.insert(engine, program.clone());
Ok(program)
}
}
#[derive(Debug, Clone)]
pub struct Producer {
pub name: String,
pub span: Span,
pub source_text: String,
pub comprehension: Comprehension,
}
pub fn strip_for_forms(
file: &PolydatFile,
) -> Result<(PolydatFile, Vec<ForStmt>, Vec<Producer>), String> {
let mut parent = Vec::with_capacity(file.statements.len());
let mut fors = Vec::new();
let mut producers: Vec<Producer> = Vec::new();
for stmt in &file.statements {
match stmt {
Statement::For(f) => fors.push(f.clone()),
Statement::Binding(b) if matches!(b.value, Expr::For(_)) => {
let Expr::For(source) = &b.value else {
unreachable!()
};
let comprehension = resolve_source(source, &producers)?;
let name = b.targets.join(",");
let value = StreamerValue::new(source.text.clone(), comprehension.clone());
parent.push(Statement::Binding(Binding {
targets: b.targets.clone(),
value: Expr::Call(CallExpr {
func: "streamer".into(),
args: vec![Arg::Positional(Expr::StringLit(value.to_json(), b.span))],
span: b.span,
}),
modifier: BindingModifier::CONST,
type_annotation: None,
span: b.span,
}));
producers.push(Producer {
name,
span: b.span,
source_text: source.text.clone(),
comprehension,
});
}
other => parent.push(other.clone()),
}
}
Ok((PolydatFile { statements: parent }, fors, producers))
}
pub fn resolve_source(source: &ForSource, producers: &[Producer]) -> Result<Comprehension, String> {
let find = |name: &str| -> Result<Comprehension, String> {
producers
.iter()
.rev()
.find(|p| p.name == name)
.map(|p| p.comprehension.clone())
.ok_or_else(|| {
let known: Vec<&str> = producers.iter().map(|p| p.name.as_str()).collect();
format!(
"`for {}` at line {}, col {}: no producer named '{name}' is bound in this scope{}",
source.text,
source.span.line,
source.span.col,
if known.is_empty() { String::new() } else { format!("; producers here: {}", known.join(", ")) }
)
})
};
match &source.kind {
ForSourceKind::Comprehension(c) => Ok(c.clone()),
ForSourceKind::Producer(name) => find(name),
ForSourceKind::Derived {
base,
filter,
order,
} => {
let mut c = find(base)?;
if let Some(pred) = filter {
c = Comprehension::filter(c, pred.clone());
}
if let Some(spec) = order {
let (strategy, truncation) = parse_order(spec).map_err(|e| {
format!(
"`for {}` at line {}, col {}: {e}",
source.text, source.span.line, source.span.col
)
})?;
c = Comprehension::order(c, strategy, truncation);
}
Ok(c)
}
}
}
fn parse_order(
spec: &str,
) -> Result<(crate::iteration::comprehension::StrategyName, Option<u64>), String> {
let carrier = format!("__o in 0..1 order {spec}");
let legacy = crate::iteration::comprehension::parse::parse_comprehension_text(&carrier)?;
let algebra = crate::iteration::comprehension::spec::legacy_to_algebra(&legacy)
.map_err(|e| e.to_string())?;
match algebra {
Comprehension::Order {
strategy,
truncation,
..
} => Ok((strategy, truncation)),
other => Err(format!(
"order spec `{spec}` did not produce an ordering (got {other:?})"
)),
}
}
pub fn element_types(
comprehension: &Comprehension,
probe: &mut dyn FnMut(&str) -> Result<PortType, String>,
) -> Result<Vec<(String, PortType)>, String> {
let mut out = Vec::new();
collect_element_types(comprehension, probe, &mut out)?;
Ok(out)
}
fn collect_element_types(
c: &Comprehension,
probe: &mut dyn FnMut(&str) -> Result<PortType, String>,
out: &mut Vec<(String, PortType)>,
) -> Result<(), String> {
match c {
Comprehension::Clause { name, source } => {
if out.iter().any(|(n, _)| n == name) {
return Ok(());
}
let ty = source_type(name, source, probe)?;
out.push((name.clone(), ty));
Ok(())
}
Comprehension::Cartesian { children } | Comprehension::Zip { children, .. } => {
for child in children {
collect_element_types(child, probe, out)?;
}
Ok(())
}
Comprehension::Union { children } => {
if let Some(first) = children.first() {
collect_element_types(first, probe, out)?;
}
Ok(())
}
Comprehension::Filter { child, .. } | Comprehension::Order { child, .. } => {
collect_element_types(child, probe, out)
}
}
}
fn source_type(
name: &str,
source: &Source,
probe: &mut dyn FnMut(&str) -> Result<PortType, String>,
) -> Result<PortType, String> {
match source {
Source::Literal { values } => {
let mut ty: Option<PortType> = None;
for v in values {
let t = match v {
LiteralValue::Int(_) => PortType::U64,
LiteralValue::Float(_) => PortType::F64,
LiteralValue::String(_) => PortType::Str,
LiteralValue::Bool(_) => PortType::Bool,
LiteralValue::Json(_) => PortType::Json,
};
match ty {
None => ty = Some(t),
Some(PortType::F64) if t == PortType::U64 => {}
Some(PortType::U64) if t == PortType::F64 => ty = Some(PortType::F64),
Some(prev) if prev != t => {
return Err(format!(
"element '{name}': literal list mixes {prev:?} and {t:?} values; a comprehension element has one type"
));
}
Some(_) => {}
}
}
ty.ok_or_else(|| format!("element '{name}': literal list is empty"))
}
Source::IntRange { .. } => Ok(PortType::U64),
Source::ContinuousInterval { .. } | Source::Distribution { .. } => Ok(PortType::F64),
Source::WorkloadParamList { .. } => Ok(PortType::Str),
Source::Generator { expr, .. } => {
let head = expr.trim();
if head.starts_with("partitions(")
|| head.starts_with("subdivide(")
|| head.ends_with(".partitions")
{
return Ok(PortType::Ext);
}
probe(head)
.map_err(|e| format!("element '{name}': cannot type generator `{head}`: {e}"))
}
}
}
fn body_declared(body: &[Statement], elements: &[(String, PortType)]) -> BTreeSet<String> {
let mut names: BTreeSet<String> = elements.iter().map(|(n, _)| n.clone()).collect();
names.insert("cycle".to_string());
for stmt in body {
match stmt {
Statement::Binding(b) => names.extend(b.targets.iter().cloned()),
Statement::InputDecl(d) => {
names.insert(d.name.clone());
}
Statement::ExternPort(p) => {
names.insert(p.name.clone());
}
Statement::ModuleDef(m) => {
names.insert(m.name.clone());
}
Statement::Cursor(c) => {
names.insert(c.name.clone());
}
Statement::Pragma { .. } => {}
Statement::For(f) => {
let _ = f;
}
Statement::Tile(t) => {
names.insert(t.name.clone());
}
}
}
names
}
fn tile_references(pieces: &[super::ast::TilePiece], out: &mut BTreeSet<String>) {
use super::ast::TilePiece;
use super::refs::collect_expr_refs;
for piece in pieces {
match piece {
TilePiece::Static(_) => {}
TilePiece::Hole(h) => collect_expr_refs(&h.expr, out),
TilePiece::Projection { body, .. } => tile_references(body, out),
TilePiece::Branch {
cond,
then,
otherwise,
..
} => {
collect_expr_refs(cond, out);
tile_references(then, out);
if let Some(o) = otherwise {
tile_references(o, out);
}
}
}
}
}
fn body_references(body: &[Statement], out: &mut BTreeSet<String>) {
use super::refs::collect_expr_refs;
for stmt in body {
match stmt {
Statement::Binding(b) => collect_expr_refs(&b.value, out),
Statement::ExternPort(p) => {
if let Some(d) = &p.default {
collect_expr_refs(d, out);
}
}
Statement::Cursor(c) => {
collect_expr_refs(&c.constructor, out);
if let Some(over) = &c.over {
collect_expr_refs(over, out);
}
}
Statement::For(f) => {
let mut inner = BTreeSet::new();
body_references(&f.body, &mut inner);
let own = body_declared(&f.body, &[]);
let elems: BTreeSet<String> = f.source.element_names().into_iter().collect();
for n in inner {
if !own.contains(&n) && !elems.contains(&n) {
out.insert(n);
}
}
}
Statement::Tile(t) => tile_references(&t.pieces, out),
Statement::InputDecl(_) | Statement::ModuleDef(_) | Statement::Pragma { .. } => {}
}
}
}
pub fn child_file(
f: &ForStmt,
comprehension: &Comprehension,
elements: &[(String, PortType)],
type_of: &dyn Fn(&str) -> Option<PortType>,
) -> Result<(PolydatFile, Vec<(String, PortType)>), String> {
for stmt in &f.body {
if let Statement::InputDecl(d) = stmt
&& d.name != "cycle"
{
return Err(format!(
"`for {}` at line {}, col {}: a traversal body cannot declare input '{}'; only `cycle` is a coordinate inside a body, and the comprehension supplies the rest",
f.source.text, f.span.line, f.span.col, d.name
));
}
}
let declared = body_declared(&f.body, elements);
let mut referenced = BTreeSet::new();
body_references(&f.body, &mut referenced);
referenced.extend(comprehension.referenced_source_names());
let mut cascade = Vec::new();
for name in referenced {
if declared.contains(&name) {
continue;
}
let ty = type_of(&name);
if let Some(ty) = ty {
cascade.push((name, ty));
}
}
let span = f.span;
let mut statements = Vec::with_capacity(f.body.len() + elements.len() + cascade.len() + 1);
if !f.body.iter().any(|s| matches!(s, Statement::InputDecl(_))) {
statements.push(Statement::InputDecl(InputDecl {
name: "cycle".into(),
ty: Some("u64".into()),
span,
}));
}
for (name, ty) in elements {
statements.push(Statement::ExternPort(ExternPort {
name: name.clone(),
typ: ty.to_keyword().to_string(),
default: None,
span,
}));
}
for (name, ty) in &cascade {
statements.push(Statement::ExternPort(ExternPort {
name: name.clone(),
typ: ty.to_keyword().to_string(),
default: None,
span,
}));
}
statements.extend(f.body.iter().cloned());
Ok((PolydatFile { statements }, cascade))
}