use std::collections::{HashMap, HashSet};
use crate::scope_kernel::ScopeKernel;
use polydat::kernel::ManifestEntry;
use super::helpers::{
format_value_as_final_literal, port_type_to_extern_name, workload_param_type_name,
};
pub struct CascadeInputs<'a> {
pub parent_kernel: &'a ScopeKernel,
pub workload_params: &'a HashMap<String, String>,
pub parent_manifest: &'a [ManifestEntry],
pub referenced: &'a HashSet<String>,
pub pre_emitted: &'a HashSet<String>,
pub shadow_names: &'a HashSet<String>,
pub include_referenced_cascade: bool,
}
pub struct CascadeOutputs<'a> {
pub source: &'a mut String,
pub emitted: &'a mut HashSet<String>,
pub inherited_names: &'a mut Vec<String>,
}
pub fn cascade_parent_into_source(inputs: CascadeInputs<'_>, outputs: CascadeOutputs<'_>) {
let CascadeInputs {
parent_kernel,
workload_params,
parent_manifest,
referenced,
pre_emitted,
shadow_names,
include_referenced_cascade,
} = inputs;
outputs.emitted.extend(pre_emitted.iter().cloned());
let parent_program = parent_kernel.program();
let coord_names: HashSet<String> = {
let coord_count = parent_program.coord_count();
parent_program
.input_names()
.into_iter()
.take(coord_count)
.collect()
};
let already_satisfied_for_inclusion: HashSet<String> = pre_emitted
.iter()
.chain(coord_names.iter())
.cloned()
.collect();
{
let mut already_satisfied = already_satisfied_for_inclusion.clone();
let mut refs_sorted: Vec<&String> = referenced.iter().collect();
refs_sorted.sort();
for name in refs_sorted {
if already_satisfied.contains(name.as_str()) {
continue;
}
let modifier = parent_program.output_modifier(name);
if modifier == polydat::dsl::ast::BindingModifier::CONST
|| modifier == polydat::dsl::ast::BindingModifier::SHARED
{
continue;
}
let chain = parent_program.local_inclusion_chain(name, &already_satisfied);
if chain.is_empty() {
continue;
}
for stmt in chain {
let line = polydat::dsl::pprint::pp_statement(stmt);
outputs.source.push_str(&line);
outputs.source.push('\n');
if let polydat::dsl::ast::Statement::Binding(b) = stmt {
for t in &b.targets {
outputs.emitted.insert(t.clone());
already_satisfied.insert(t.clone());
}
}
}
}
}
if include_referenced_cascade {
let manifest_by_name: HashMap<&str, &ManifestEntry> = parent_manifest
.iter()
.map(|e| (e.name.as_str(), e))
.collect();
let mut refs_sorted: Vec<&String> = referenced.iter().collect();
refs_sorted.sort();
for name in refs_sorted {
if outputs.emitted.contains(name) {
continue;
}
if pre_emitted.contains(name) {
continue;
}
if coord_names.contains(name) {
continue;
}
if let Some(entry) = manifest_by_name.get(name.as_str()) {
let type_name = port_type_to_extern_name(entry.port_type);
outputs
.source
.push_str(&format!("extern {name}: {type_name}\n"));
outputs.emitted.insert(name.clone());
outputs.inherited_names.push(name.clone());
} else if let Some(value) = workload_params.get(name) {
super::cascade_emit::emit_workload_param_chain_aware(
name,
value,
parent_kernel,
outputs.source,
outputs.emitted,
None,
);
}
}
} else {
let _ = parent_manifest;
}
for (name, value) in workload_params {
if outputs.emitted.contains(name) {
continue;
}
if shadow_names.contains(name) {
continue;
}
let type_name = workload_param_type_name(value);
outputs
.source
.push_str(&format!("extern {name}: {type_name}\n"));
outputs.emitted.insert(name.clone());
outputs.inherited_names.push(name.clone());
}
let skip_cascade = |emitted: &HashSet<String>, name: &str| -> bool {
if emitted.contains(name) {
return true;
}
if coord_names.contains(name) {
return true;
}
if name.starts_with("__") && !name.starts_with("__cursor_extent_") {
return true;
}
false
};
for name in parent_program.output_names() {
let owned = name.to_string();
if skip_cascade(outputs.emitted, &owned) {
continue;
}
if shadow_names.contains(&owned) {
continue;
}
let Some(output_idx) = parent_program.output_index(&owned) else {
continue;
};
let (node_idx, port_idx) = parent_program.resolve_output_by_index(output_idx);
let is_shared =
parent_program.output_modifier(&owned) == polydat::dsl::ast::BindingModifier::SHARED;
let upstream_is_statically_known = !is_shared
&& parent_program
.input_provenance_for(node_idx)
.is_none_or(|p| p.is_zero());
if upstream_is_statically_known
&& let Some(value) = parent_kernel.lookup(&owned)
&& let Some(literal) = format_value_as_final_literal(&value)
{
outputs
.source
.push_str(&format!("const {owned} := {literal}\n"));
outputs.emitted.insert(owned);
continue;
}
let port_type = parent_program.node_meta(node_idx).outs[port_idx].typ;
let type_name = port_type_to_extern_name(port_type);
outputs
.source
.push_str(&format!("extern {owned}: {type_name}\n"));
outputs.emitted.insert(owned.clone());
outputs.inherited_names.push(owned);
}
for name in parent_program.input_names() {
if skip_cascade(outputs.emitted, &name) {
continue;
}
if shadow_names.contains(&name) {
continue;
}
let port_type = parent_program
.input_port_type(&name)
.expect("input_names() returned a name with no port type");
let type_name = port_type.to_keyword();
outputs
.source
.push_str(&format!("extern {name}: {type_name}\n"));
outputs.emitted.insert(name.clone());
outputs.inherited_names.push(name);
}
}