use std::collections::{HashMap, HashSet};
use crate::scope_kernel::ScopeKernel;
use nmbrs_workload::bindpoints;
use nmbrs_workload::model::ParsedOp;
#[derive(Debug, Clone)]
struct BindingFunc {
name: String,
args: Vec<String>,
}
fn parse_binding_chain(expr: &str) -> Vec<BindingFunc> {
let mut funcs = Vec::new();
for segment in expr.split(';') {
let segment = segment.trim();
if segment.is_empty() {
continue;
}
if let Some(paren_pos) = segment.find('(') {
let name = segment[..paren_pos].trim().to_string();
let args_str = &segment[paren_pos + 1..];
let args_str = args_str.trim_end_matches(')').trim();
let args: Vec<String> = if args_str.is_empty() {
Vec::new()
} else {
split_args(args_str)
};
funcs.push(BindingFunc { name, args });
} else {
funcs.push(BindingFunc {
name: segment.trim().to_string(),
args: Vec::new(),
});
}
}
funcs
}
fn split_args(s: &str) -> Vec<String> {
let mut args = Vec::new();
let mut current = String::new();
let mut depth = 0;
let mut in_quote = false;
for c in s.chars() {
match c {
'\'' if !in_quote => {
in_quote = true;
current.push(c);
}
'\'' if in_quote => {
in_quote = false;
current.push(c);
}
'(' if !in_quote => {
depth += 1;
current.push(c);
}
')' if !in_quote => {
depth -= 1;
current.push(c);
}
',' if depth == 0 && !in_quote => {
args.push(current.trim().to_string());
current = String::new();
}
_ => current.push(c),
}
}
if !current.trim().is_empty() {
args.push(current.trim().to_string());
}
args
}
pub fn probe_compile_level(func_name: &str) -> polydat::ast::CompileLevel {
let sig = match polydat::dsl::registry::lookup(func_name) {
Some(s) => s,
None => return polydat::ast::CompileLevel::Phase1,
};
let mut parts = Vec::new();
let mut has_wire = false;
for p in sig.params {
parts.push(p.example.to_string());
if p.slot_type.is_wire() {
has_wire = true;
}
}
if !has_wire && parts.is_empty() {
parts.push("cycle".to_string());
}
let source = format!("input cycle: u64\nout := {func_name}({})", parts.join(", "));
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
polydat::dsl::compile::compile_polydat_interpreter(&source)
}));
match result {
Ok(Ok(kernel)) => kernel.program().last_node_compile_level(),
_ => polydat::ast::CompileLevel::Phase1,
}
}
pub fn compile_bindings(ops: &[ParsedOp]) -> Result<ScopeKernel, String> {
compile_bindings_with_path(ops, None)
}
fn collect_param_bindings(
params: &HashMap<String, serde_json::Value>,
exclude: &[String],
required: &mut Vec<String>,
) {
for (key, value) in params.iter() {
if key == "gutter" {
continue;
}
collect_json_bindings(value, exclude, required);
}
}
pub fn collect_param_bindings_into(
params: &HashMap<String, serde_json::Value>,
exclude: &[String],
required: &mut Vec<String>,
) {
collect_param_bindings(params, exclude, required);
}
fn collect_json_bindings(
value: &serde_json::Value,
exclude: &[String],
required: &mut Vec<String>,
) {
match value {
serde_json::Value::String(s) => {
for name in bindpoints::referenced_bindings(s) {
if !required.contains(&name) && !exclude.contains(&name) {
required.push(name);
}
}
}
serde_json::Value::Object(map) => {
for v in map.values() {
collect_json_bindings(v, exclude, required);
}
}
serde_json::Value::Array(arr) => {
for v in arr {
collect_json_bindings(v, exclude, required);
}
}
_ => {}
}
}
pub fn compile_bindings_with_path(
ops: &[ParsedOp],
source_dir: Option<&std::path::Path>,
) -> Result<ScopeKernel, String> {
compile_bindings_with_opts(ops, source_dir, false)
}
pub fn compile_from_scope(
scope: &crate::scope::BindingScope,
source_dir: Option<&std::path::Path>,
polydat_lib_paths: Vec<std::path::PathBuf>,
strict: bool,
context: &str,
cursor_limit: Option<u64>,
pragmas: &polydat::dsl::pragmas::PragmaSet,
) -> Result<ScopeKernel, String> {
let (source, options) = scope_source_and_options(
scope,
source_dir,
polydat_lib_paths,
strict,
context,
cursor_limit,
pragmas,
);
compile_scope_kernel(&source, &options)
}
fn scope_source_and_options(
scope: &crate::scope::BindingScope,
source_dir: Option<&std::path::Path>,
polydat_lib_paths: Vec<std::path::PathBuf>,
strict: bool,
context: &str,
cursor_limit: Option<u64>,
pragmas: &polydat::dsl::pragmas::PragmaSet,
) -> (String, polydat::dsl::compile::CompileOptions) {
let body = scope.emit();
let required = scope.required_outputs();
let source = prepend_effective_pragmas(pragmas, &body);
let options = polydat::dsl::compile::CompileOptions {
source_dir: source_dir.map(std::path::Path::to_path_buf),
lib_paths: polydat_lib_paths,
required_outputs: required,
strict,
context: context.to_string(),
cursor_limit,
..Default::default()
};
(source, options)
}
pub fn compile_scope_kernel(
source: &str,
options: &polydat::dsl::compile::CompileOptions,
) -> Result<ScopeKernel, String> {
let options = polydat::dsl::compile::CompileOptions {
resources: Some(
options
.resources
.clone()
.unwrap_or_else(crate::resource_pool::pool_resources),
),
..options.clone()
};
let options = &options;
let interpreter =
polydat::dsl::compile::compile_polydat_interpreter_with_options(source, options, None)
.map_err(|e| e.to_string())?;
let image =
crate::fiber_engine::source_image(interpreter.program(), source, options, &options.context);
Ok(ScopeKernel::root(interpreter, image))
}
pub fn compile_scope_program(
source: &str,
options: &polydat::dsl::compile::CompileOptions,
) -> Result<std::sync::Arc<polydat::kernel::PolydatProgram>, String> {
polydat::dsl::compile::compile_polydat_interpreter_with_options(source, options, None)
.map(|k| k.program().clone())
.map_err(|e| e.to_string())
}
pub(crate) fn prepend_effective_pragmas(
pragmas: &polydat::dsl::pragmas::PragmaSet,
body: &str,
) -> String {
let mut out = String::new();
if pragmas.strict_types() && pragmas.strict_values() {
out.push_str("pragma strict\n");
} else if pragmas.strict_types() {
out.push_str("pragma strict_types\n");
} else if pragmas.strict_values() {
out.push_str("pragma strict_values\n");
}
if !out.is_empty() {
out.push('\n');
}
out.push_str(body);
out
}
pub const SESSION_START: &str = "session_start";
#[allow(clippy::too_many_arguments)]
pub fn build_workload_root_kernel(
parent: &ScopeKernel,
ops: &[ParsedOp],
source_dir: Option<&std::path::Path>,
polydat_lib_paths: Vec<std::path::PathBuf>,
strict: bool,
extra_required: &[String],
context: &str,
cursor_limit: Option<u64>,
workload_params: &std::collections::HashMap<String, String>,
workload_level_polydat: Option<&str>,
) -> Result<ScopeKernel, String> {
let mut scope = crate::scope::build_scope(
ops,
&std::collections::HashMap::new(), &[], workload_params,
&std::collections::HashMap::new(), None, &[], None, )?;
if let Some(extra) = workload_level_polydat
&& !extra.trim().is_empty()
{
scope.ingest_polydat_source(extra, crate::scope::BindingOrigin::Inherited);
}
let session_clock = !workload_params.contains_key(SESSION_START)
&& !scope.defined_names().contains(SESSION_START);
if session_clock {
scope.ingest_polydat_source(
&format!("const {SESSION_START} := current_epoch_millis()\n"),
crate::scope::BindingOrigin::Inherited,
);
}
scope.validate().map_err(|e| format!("{context}: {e}"))?;
let mut scope_required = scope.required_outputs();
for name in extra_required {
if !scope_required.contains(name) {
scope_required.push(name.clone());
}
}
if session_clock && !scope_required.iter().any(|n| n == SESSION_START) {
scope_required.push(SESSION_START.to_string());
}
let mut param_names: Vec<&String> = workload_params.keys().collect();
param_names.sort();
for name in param_names {
if !scope_required.contains(name) {
scope_required.push(name.clone());
}
}
let mut source = scope.emit();
if !source.lines().any(|l| l.trim_start().starts_with("input ")) {
source = format!("input cycle: u64\n{source}");
}
let opts = polydat::kernel::subcontext::CompileOptions {
workload_dir: source_dir.map(|p| p.to_path_buf()),
polydat_lib_paths,
strict,
required_outputs: scope_required,
context_label: Some(context.to_string()),
cursor_limit,
..Default::default()
};
let mut inherited_param_names: Vec<String> = workload_params.keys().cloned().collect();
inherited_param_names.sort();
ScopeKernel::build_under(
parent.kernel(),
crate::scope_kernel::SourceMatter::source(context, source, opts)
.inherited(inherited_param_names),
)
}
pub fn compile_bindings_with_opts(
ops: &[ParsedOp],
source_dir: Option<&std::path::Path>,
strict: bool,
) -> Result<ScopeKernel, String> {
use nmbrs_workload::model::BindingsDef;
let polydat_source = ops.iter().find_map(|op| {
if let BindingsDef::PolydatSource(src) = &op.bindings {
if !src.trim().is_empty() {
Some(src.clone())
} else {
None
}
} else {
None
}
});
if let Some(source) = polydat_source {
let mut required: Vec<String> = Vec::new();
for op in ops {
for value in op.op.values() {
if let Some(s) = value.as_str() {
for name in bindpoints::referenced_bindings(s) {
if !required.contains(&name) {
required.push(name);
}
}
}
}
}
let options = polydat::dsl::compile::CompileOptions {
source_dir: source_dir.map(std::path::Path::to_path_buf),
required_outputs: required,
strict,
..Default::default()
};
return compile_scope_kernel(&source, &options);
}
let mut all_bindings: HashMap<String, String> = HashMap::new();
for op in ops {
if let BindingsDef::Map(map) = &op.bindings {
for (name, expr) in map {
all_bindings
.entry(name.clone())
.or_insert_with(|| expr.clone());
}
}
}
let mut required: Vec<String> = Vec::new();
for op in ops {
for value in op.op.values() {
if let Some(s) = value.as_str() {
for name in bindpoints::referenced_bindings(s) {
if !required.contains(&name) {
required.push(name);
}
}
}
}
}
let mut polydat_lines: Vec<String> = Vec::new();
polydat_lines.push("input cycle: u64".into());
for (binding_name, expr) in &all_bindings {
let chain = parse_binding_chain(expr);
if chain.is_empty() {
return Err(format!("empty binding expression for '{binding_name}'"));
}
let mut prev_wire = "cycle".to_string();
for (i, func) in chain.iter().enumerate() {
let is_last = i == chain.len() - 1;
let target = if is_last {
binding_name.clone()
} else {
format!("__chain_{binding_name}_{i}")
};
let (func_name, extra_args) = translate_legacy_func(&func.name, &func.args);
let mut call_args = vec![prev_wire.clone()];
for arg in &func.args {
call_args.push(strip_java_long_suffix(arg.trim()).to_string());
}
call_args.extend(extra_args);
polydat_lines.push(format!(
"{target} := {func_name}({args})",
args = call_args.join(", ")
));
prev_wire = target;
}
}
let coord_names: HashSet<String> = ["cycle".to_string()].into_iter().collect();
let mut missing: Vec<String> = Vec::new();
for name in &required {
if !all_bindings.contains_key(name) && !coord_names.contains(name) {
missing.push(name.clone());
}
}
if !missing.is_empty() {
return Err(format!(
"undeclared bind point references: {}. Add these to your bindings section.",
missing.join(", ")
));
}
let polydat_source = polydat_lines.join("\n");
let options = polydat::dsl::compile::CompileOptions {
source_dir: source_dir.map(std::path::Path::to_path_buf),
required_outputs: required,
strict,
..Default::default()
};
compile_scope_kernel(&polydat_source, &options)
}
fn translate_legacy_func(name: &str, args: &[String]) -> (String, Vec<String>) {
match name.to_lowercase().as_str() {
"hash" => ("hash".into(), vec![]),
"identity" => ("identity".into(), vec![]),
"add" => ("add".into(), vec![]),
"mul" => ("mul".into(), vec![]),
"div" => ("div".into(), vec![]),
"mod" => ("mod".into(), vec![]),
"clamp" => ("clamp".into(), vec![]),
"tostring" | "to_string" => ("format_u64".into(), vec!["10".into()]),
"tohexstring" => ("format_u64".into(), vec!["16".into()]),
"tooctalstring" => ("format_u64".into(), vec!["8".into()]),
"tobinarystring" => ("format_u64".into(), vec!["2".into()]),
"uniform" => {
if args.len() >= 2 {
("hash_range".into(), vec![])
} else {
("hash_range".into(), vec![])
}
}
"normal" | "gaussian" => ("icd_normal".into(), vec![]),
"zipf" => ("dist_zipf".into(), vec![]),
"hashrange" | "hash_range" => ("hash_range".into(), vec![]),
"hashinterval" | "hash_interval" => ("hash_interval".into(), vec![]),
"format" | "printf" => ("printf".into(), vec![]),
"numbernamesto_string" | "numbernames" => ("number_to_words".into(), vec![]),
"shuffle" => ("shuffle".into(), vec![]),
_ => {
(name.to_lowercase(), vec![])
}
}
}
fn strip_java_long_suffix(arg: &str) -> &str {
arg.strip_suffix('L')
.or_else(|| arg.strip_suffix('l'))
.unwrap_or(arg)
}
pub fn legacy_chain_map_to_polydat_lines(
map: &std::collections::HashMap<String, String>,
) -> Result<String, String> {
let mut polydat_lines: Vec<String> = Vec::new();
for (binding_name, expr) in map {
let chain = parse_binding_chain(expr);
if chain.is_empty() {
return Err(format!("empty binding expression for '{binding_name}'"));
}
let mut prev_wire = "cycle".to_string();
for (i, func) in chain.iter().enumerate() {
let is_last = i == chain.len() - 1;
let target = if is_last {
binding_name.clone()
} else {
format!("__chain_{binding_name}_{i}")
};
let (func_name, extra_args) = translate_legacy_func(&func.name, &func.args);
let mut call_args = vec![prev_wire.clone()];
for arg in &func.args {
call_args.push(strip_java_long_suffix(arg.trim()).to_string());
}
call_args.extend(extra_args);
polydat_lines.push(format!(
"{target} := {func_name}({args})",
args = call_args.join(", ")
));
prev_wire = target;
}
}
Ok(polydat_lines.join("\n"))
}
#[cfg(test)]
mod tests {
use super::*;
use polydat::dsl::pragmas::{Pragma, PragmaSet};
#[test]
fn prepend_pragmas_strict_alias() {
let pragmas = PragmaSet {
entries: vec![Pragma {
name: "strict".into(),
args: vec![],
line: 1,
}],
};
let body = "id := cycle\n";
let out = prepend_effective_pragmas(&pragmas, body);
assert!(out.starts_with("pragma strict\n"));
assert!(out.contains("id := cycle"));
}
#[test]
fn prepend_pragmas_individual_modes() {
let pragmas = PragmaSet {
entries: vec![Pragma {
name: "strict_values".into(),
args: vec![],
line: 1,
}],
};
let out = prepend_effective_pragmas(&pragmas, "x := cycle");
assert!(out.starts_with("pragma strict_values\n"));
assert!(!out.contains("strict_types"));
}
#[test]
fn prepend_pragmas_no_op_when_empty() {
let pragmas = PragmaSet::default();
let out = prepend_effective_pragmas(&pragmas, "x := cycle");
assert_eq!(out, "x := cycle");
}
#[test]
fn prepend_pragmas_walks_parent_chain() {
let parent = PragmaSet {
entries: vec![Pragma {
name: "strict_values".into(),
args: vec![],
line: 1,
}],
};
let child = parent.nested(&[]);
let out = prepend_effective_pragmas(&child, "x := cycle");
assert!(
out.starts_with("pragma strict_values\n"),
"expected pragma to flow from parent chain, got:\n{out}"
);
}
#[test]
fn parse_simple_chain() {
let chain = parse_binding_chain("Hash(); Mod(1000000)");
assert_eq!(chain.len(), 2);
assert_eq!(chain[0].name, "Hash");
assert!(chain[0].args.is_empty());
assert_eq!(chain[1].name, "Mod");
assert_eq!(chain[1].args, vec!["1000000"]);
}
#[test]
fn parse_identity() {
let chain = parse_binding_chain("Identity()");
assert_eq!(chain.len(), 1);
assert_eq!(chain[0].name, "Identity");
}
#[test]
fn parse_with_string_arg() {
let chain = parse_binding_chain("Template('user-{}', ToString())");
assert_eq!(chain.len(), 1);
assert_eq!(chain[0].name, "Template");
assert_eq!(chain[0].args.len(), 2);
}
#[test]
fn parse_long_chain() {
let chain = parse_binding_chain("Add(10); Hash(); Mod(100); ToString()");
assert_eq!(chain.len(), 4);
assert_eq!(chain[0].name, "Add");
assert_eq!(chain[1].name, "Hash");
assert_eq!(chain[2].name, "Mod");
assert_eq!(chain[3].name, "ToString");
}
#[test]
fn parse_with_long_suffix() {
let chain = parse_binding_chain("Mod(1000000000L)");
assert_eq!(chain[0].args, vec!["1000000000L"]);
}
#[test]
fn compile_identity_binding() {
let ops = vec![{
let mut op = ParsedOp::simple("test", "{myval}");
op.bindings.insert("myval".into(), "Identity()".into());
op
}];
let mut kernel = compile_bindings(&ops).unwrap();
kernel.set_inputs(&[42]);
assert_eq!(kernel.pull("myval").as_u64(), 42);
}
#[test]
fn compile_hash_mod_binding() {
let ops = vec![{
let mut op = ParsedOp::simple("test", "{id}");
op.bindings
.insert("id".into(), "Hash(); Mod(1000000)".into());
op
}];
let mut kernel = compile_bindings(&ops).unwrap();
kernel.set_inputs(&[42]);
let val = kernel.pull("id").as_u64();
assert!(val < 1_000_000, "got {val}");
}
#[test]
fn compile_hash_mod_deterministic() {
let ops = vec![{
let mut op = ParsedOp::simple("test", "{id}");
op.bindings.insert("id".into(), "Hash(); Mod(100)".into());
op
}];
let mut kernel = compile_bindings(&ops).unwrap();
kernel.set_inputs(&[42]);
let v1 = kernel.pull("id").as_u64();
kernel.set_inputs(&[42]);
let v2 = kernel.pull("id").as_u64();
assert_eq!(v1, v2);
}
#[test]
fn compile_multiple_bindings() {
let ops = vec![{
let mut op = ParsedOp::simple("test", "{a} {b}");
op.bindings.insert("a".into(), "Identity()".into());
op.bindings.insert("b".into(), "Hash(); Mod(100)".into());
op
}];
let mut kernel = compile_bindings(&ops).unwrap();
kernel.set_inputs(&[5]);
assert_eq!(kernel.pull("a").as_u64(), 5);
assert!(kernel.pull("b").as_u64() < 100);
}
#[test]
fn compile_rejects_undeclared_bind_points() {
let ops = vec![ParsedOp::simple("test", "val={mystery}")];
let result = compile_bindings(&ops);
assert!(result.is_err());
assert!(result.unwrap_err().contains("undeclared bind point"));
}
#[test]
fn compile_add_chain() {
let ops = vec![{
let mut op = ParsedOp::simple("test", "{val}");
op.bindings
.insert("val".into(), "Add(100); Mod(1000)".into());
op
}];
let mut kernel = compile_bindings(&ops).unwrap();
kernel.set_inputs(&[5]);
assert_eq!(kernel.pull("val").as_u64(), 105);
}
#[test]
fn compile_provides_cycle_output() {
let ops = vec![ParsedOp::simple("test", "cycle={cycle}")];
let mut kernel = compile_bindings(&ops).unwrap();
kernel.set_inputs(&[99]);
assert_eq!(kernel.pull("cycle").as_u64(), 99);
}
#[test]
fn legacy_tostring_translates() {
let (name, _) = translate_legacy_func("ToString", &[]);
assert_eq!(name, "format_u64");
}
#[test]
fn legacy_uniform_translates() {
let (name, _) = translate_legacy_func("Uniform", &["0".into(), "1000".into()]);
assert_eq!(name, "hash_range");
}
}