use std::collections::HashMap;
use crate::model::{BindingsDef, ParsedOp, Workload};
pub fn synthesize_inline_workload(op_template: &str) -> Result<Workload, String> {
if op_template.trim().is_empty() {
return Err("op= value is empty".into());
}
if let Some(w) = try_polydat_block_workload(op_template) {
return Ok(w);
}
let segments = split_ops(op_template);
let mut inline_exprs: Vec<String> = Vec::new();
let mut expr_index: HashMap<String, usize> = HashMap::new();
for seg in &segments {
for expr in extract_inline_exprs(&seg.template) {
if !expr_index.contains_key(&expr) {
let idx = inline_exprs.len();
expr_index.insert(expr.clone(), idx);
inline_exprs.push(expr);
}
}
for bp in crate::bindpoints::extract_bind_points(&seg.template) {
if let crate::bindpoints::BindPoint::InlineDefinition(expr) = bp
&& !expr_index.contains_key(&expr)
{
let idx = inline_exprs.len();
expr_index.insert(expr.clone(), idx);
inline_exprs.push(expr);
}
}
}
let mut polydat_source = String::from("input cycle: u64\n");
for (i, expr) in inline_exprs.iter().enumerate() {
polydat_source.push_str(&format!("__inline_{i} := {expr}\n"));
}
let mut ops = Vec::with_capacity(segments.len());
for (i, seg) in segments.iter().enumerate() {
let rewritten = rewrite_template(&seg.template, &expr_index);
let mut op = ParsedOp::simple(&format!("inline_{i}"), &rewritten);
if seg.ratio != 1 {
op.params.insert(
"ratio".to_string(),
serde_json::Value::Number(serde_json::Number::from(seg.ratio)),
);
}
op.tags.insert("name".to_string(), op.name.clone());
op.tags.insert("op".to_string(), op.name.clone());
op.tags.insert("block".to_string(), "inline".to_string());
op.bindings = BindingsDef::PolydatSource(polydat_source.clone());
ops.push(op);
}
Ok(Workload {
description: Some("inline workload".into()),
scenarios: HashMap::new(),
stop_when: Vec::new(),
ops,
bindings: crate::model::BindingsDef::default(),
params: HashMap::new(),
phases: HashMap::new(),
phase_order: Vec::new(),
declared_params: Vec::new(),
report: crate::report::Report::default(),
report_warnings: Vec::new(),
resolution_warnings: Vec::new(),
scenario_parse_errors: Vec::new(),
status_metrics: Vec::new(),
readouts: crate::model::ReadoutsBindings::default(),
wrappers: None,
implements: None,
stick_session: None,
})
}
fn try_polydat_block_workload(op_template: &str) -> Option<Workload> {
let source = build_polydat_candidate(op_template);
polydat::dsl::compile::compile_polydat(&source).ok()?;
let names = binding_wire_names(&source);
if names.is_empty() {
return None;
}
let mut op_fields: HashMap<String, serde_json::Value> = HashMap::new();
for n in &names {
op_fields.insert(n.clone(), serde_json::Value::String(format!("{{{n}}}")));
}
let mut op = ParsedOp::simple("inline_0", "");
op.op = op_fields;
op.bindings = BindingsDef::PolydatSource(source);
op.tags.insert("name".to_string(), "inline_0".to_string());
op.tags.insert("op".to_string(), "inline_0".to_string());
op.tags.insert("block".to_string(), "inline".to_string());
Some(Workload {
description: Some("inline polydat workload".into()),
scenarios: HashMap::new(),
stop_when: Vec::new(),
ops: vec![op],
bindings: crate::model::BindingsDef::default(),
params: HashMap::new(),
phases: HashMap::new(),
phase_order: Vec::new(),
declared_params: Vec::new(),
report: crate::report::Report::default(),
report_warnings: Vec::new(),
resolution_warnings: Vec::new(),
scenario_parse_errors: Vec::new(),
status_metrics: Vec::new(),
readouts: crate::model::ReadoutsBindings::default(),
wrappers: None,
implements: None,
stick_session: None,
})
}
fn build_polydat_candidate(op_template: &str) -> String {
let segs: Vec<String> = split_top_level_semicolons(op_template)
.into_iter()
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
.collect();
let mut lines = vec!["input cycle: u64".to_string()];
let last = segs.len().saturating_sub(1);
for (i, seg) in segs.iter().enumerate() {
if has_top_level_assignment(seg) {
lines.push(seg.clone());
} else if i == last {
lines.push(format!("out := {seg}"));
} else {
lines.push(format!("__expr_{i} := {seg}"));
}
}
lines.join("\n") + "\n"
}
pub(crate) fn binding_wire_names(source: &str) -> Vec<String> {
source
.lines()
.filter_map(|line| {
let line = line.trim();
if line.starts_with("input ") || !line.contains(":=") {
return None;
}
let lhs = line.split(":=").next()?.trim();
let name = lhs.split_whitespace().last()?;
let is_ident = !name.is_empty()
&& name.chars().all(|c| c.is_alphanumeric() || c == '_')
&& name
.chars()
.next()
.is_some_and(|c| c.is_alphabetic() || c == '_');
if is_ident && !name.starts_with("__") {
Some(name.to_string())
} else {
None
}
})
.collect()
}
fn has_top_level_assignment(s: &str) -> bool {
let bytes = s.as_bytes();
let mut depth = 0i32;
let mut i = 0;
while i < bytes.len() {
match bytes[i] {
b'{' => depth += 1,
b'}' => depth = (depth - 1).max(0),
b':' if depth == 0 && i + 1 < bytes.len() && bytes[i + 1] == b'=' => {
return true;
}
_ => {}
}
i += 1;
}
false
}
fn split_top_level_semicolons(s: &str) -> Vec<String> {
let mut out = Vec::new();
let mut cur = String::new();
let mut depth = 0i32;
for c in s.chars() {
match c {
'{' => {
depth += 1;
cur.push(c);
}
'}' => {
depth = (depth - 1).max(0);
cur.push(c);
}
';' if depth == 0 => {
out.push(std::mem::take(&mut cur));
}
_ => cur.push(c),
}
}
if !cur.trim().is_empty() {
out.push(cur);
}
out
}
struct OpSegment {
template: String,
ratio: u64,
}
fn split_ops(input: &str) -> Vec<OpSegment> {
let mut segments = Vec::new();
let mut current = String::new();
let mut in_braces = 0u32;
for c in input.chars() {
match c {
'{' => {
in_braces += 1;
current.push(c);
}
'}' => {
in_braces = in_braces.saturating_sub(1);
current.push(c);
}
';' if in_braces == 0 => {
let seg = current.trim().to_string();
if !seg.is_empty() {
segments.push(parse_segment(&seg));
}
current.clear();
}
_ => current.push(c),
}
}
let seg = current.trim().to_string();
if !seg.is_empty() {
segments.push(parse_segment(&seg));
}
segments
}
fn parse_segment(s: &str) -> OpSegment {
if let Some(colon_pos) = s.find(':') {
let prefix = &s[..colon_pos];
if !prefix.is_empty()
&& prefix.chars().all(|c| c.is_ascii_digit())
&& let Ok(ratio) = prefix.parse::<u64>()
{
return OpSegment {
template: s[colon_pos + 1..].trim().to_string(),
ratio,
};
}
}
OpSegment {
template: s.to_string(),
ratio: 1,
}
}
fn extract_inline_exprs(template: &str) -> Vec<String> {
let mut exprs = Vec::new();
let bytes = template.as_bytes();
let len = bytes.len();
let mut i = 0;
while i + 1 < len {
if bytes[i] == b'{' && bytes[i + 1] == b'{' {
let start = i + 2;
let mut depth = 1u32;
let mut j = start;
while j + 1 < len {
if bytes[j] == b'{' && bytes[j + 1] == b'{' {
depth += 1;
j += 2;
} else if bytes[j] == b'}' && bytes[j + 1] == b'}' {
depth -= 1;
if depth == 0 {
let expr = template[start..j].trim().to_string();
if !expr.is_empty() {
exprs.push(expr);
}
i = j + 2;
break;
}
j += 2;
} else {
j += 1;
}
}
if depth > 0 {
i += 2;
}
} else {
i += 1;
}
}
exprs
}
fn rewrite_template(template: &str, expr_index: &HashMap<String, usize>) -> String {
let after_double = rewrite_double_brace(template, expr_index);
rewrite_single_brace_exprs(&after_double, expr_index)
}
fn rewrite_single_brace_exprs(template: &str, expr_index: &HashMap<String, usize>) -> String {
let mut result = String::with_capacity(template.len());
let chars: Vec<char> = template.chars().collect();
let mut i = 0;
while i < chars.len() {
if chars[i] == '{' && (i + 1 >= chars.len() || chars[i + 1] != '{') {
let start = i + 1;
let mut depth = 1u32;
let mut j = start;
while j < chars.len() {
if chars[j] == '{' {
depth += 1;
}
if chars[j] == '}' {
depth -= 1;
if depth == 0 {
break;
}
}
j += 1;
}
if j < chars.len() {
let raw: String = chars[start..j].iter().collect();
let raw = raw.trim();
let expr = if let Some(e) = raw.strip_prefix(":=") {
Some(e.strip_suffix(":=").unwrap_or(e).trim())
} else if crate::bindpoints::is_expression_public(raw) {
Some(raw)
} else {
None
};
if let Some(expr) = expr {
if let Some(&idx) = expr_index.get(expr) {
result.push_str(&format!("{{__inline_{idx}}}"));
} else {
result.push('{');
result.push_str(raw);
result.push('}');
}
} else {
result.push('{');
result.push_str(raw);
result.push('}');
}
i = j + 1;
} else {
result.push(chars[i]);
i += 1;
}
} else {
result.push(chars[i]);
i += 1;
}
}
result
}
fn rewrite_double_brace(template: &str, expr_index: &HashMap<String, usize>) -> String {
let mut result = String::with_capacity(template.len());
let bytes = template.as_bytes();
let len = bytes.len();
let mut i = 0;
while i < len {
if i + 1 < len && bytes[i] == b'{' && bytes[i + 1] == b'{' {
let start = i + 2;
let mut depth = 1u32;
let mut j = start;
while j + 1 < len {
if bytes[j] == b'{' && bytes[j + 1] == b'{' {
depth += 1;
j += 2;
} else if bytes[j] == b'}' && bytes[j + 1] == b'}' {
depth -= 1;
if depth == 0 {
let expr = template[start..j].trim().to_string();
if let Some(&idx) = expr_index.get(&expr) {
result.push_str(&format!("{{__inline_{idx}}}"));
} else {
result.push_str(&template[i..j + 2]);
}
i = j + 2;
break;
}
j += 2;
} else {
j += 1;
}
}
if depth > 0 {
result.push_str(&template[i..]);
break;
}
} else {
result.push(bytes[i] as char);
i += 1;
}
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn simple_inline_binding() {
let w = synthesize_inline_workload("hello {{cycle}}").unwrap();
assert_eq!(w.ops.len(), 1);
assert_eq!(w.ops[0].name, "inline_0");
let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
assert_eq!(stmt, "hello {__inline_0}");
match &w.ops[0].bindings {
BindingsDef::PolydatSource(src) => {
assert!(src.contains("input cycle: u64"));
assert!(src.contains("__inline_0 := cycle"));
}
_ => panic!("expected PolydatSource bindings"),
}
}
#[test]
fn multiple_inline_bindings() {
let w = synthesize_inline_workload(
"id={{mod(hash(cycle), 100000)}} name={{number_to_words(cycle)}}",
)
.unwrap();
assert_eq!(w.ops.len(), 1);
let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
assert_eq!(stmt, "id={__inline_0} name={__inline_1}");
match &w.ops[0].bindings {
BindingsDef::PolydatSource(src) => {
assert!(src.contains("__inline_0 := mod(hash(cycle), 100000)"));
assert!(src.contains("__inline_1 := number_to_words(cycle)"));
}
_ => panic!("expected PolydatSource bindings"),
}
}
#[test]
fn bindings_block_op_becomes_polydat_fields() {
let w =
synthesize_inline_workload("x := cos(to_f64(cycle)); y := sin(to_f64(cycle))").unwrap();
assert_eq!(w.ops.len(), 1);
let keys: std::collections::BTreeSet<&str> =
w.ops[0].op.keys().map(|s| s.as_str()).collect();
assert!(keys.contains("x") && keys.contains("y"), "fields: {keys:?}");
assert!(!w.ops[0].op.contains_key("stmt"), "should not be a text op");
assert_eq!(w.ops[0].op.get("x").unwrap().as_str().unwrap(), "{x}");
assert!(matches!(w.ops[0].bindings, BindingsDef::PolydatSource(_)));
}
#[test]
fn bare_polydat_expr_becomes_out_field() {
let w = synthesize_inline_workload("cos(to_f64(cycle))").unwrap();
assert_eq!(w.ops.len(), 1);
assert!(
w.ops[0].op.contains_key("out"),
"fields: {:?}",
w.ops[0].op.keys().collect::<Vec<_>>()
);
}
#[test]
fn invalid_polydat_falls_back_to_text_template() {
let w = synthesize_inline_workload("x := not_a_real_fn(@@@)").unwrap();
assert_eq!(w.ops.len(), 1);
assert!(w.ops[0].op.contains_key("stmt"));
}
#[test]
fn detection_is_compile_driven_not_syntactic() {
let w = synthesize_inline_workload("id-{cycle}").unwrap();
assert!(w.ops[0].op.contains_key("stmt"));
}
#[test]
fn helpers_split_and_name_bindings() {
assert!(has_top_level_assignment("x := 1"));
assert!(!has_top_level_assignment("hello {{x := 1}}")); assert_eq!(split_top_level_semicolons("a := 1; b := 2").len(), 2);
let src = build_polydat_candidate("a := 1; sin(cycle)");
assert!(src.contains("a := 1"));
assert!(src.contains("out := sin(cycle)")); assert_eq!(
binding_wire_names("input cycle: u64\nt := 1\n__expr_0 := 2\nx := 3\n"),
vec!["t".to_string(), "x".to_string()]
); }
#[test]
fn no_inline_bindings_plain_text() {
let w = synthesize_inline_workload("hello world").unwrap();
assert_eq!(w.ops.len(), 1);
let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
assert_eq!(stmt, "hello world");
let bindings = match &w.ops[0].bindings {
crate::model::BindingsDef::PolydatSource(s) => s.clone(),
_ => panic!("expected PolydatSource"),
};
assert_eq!(bindings, "input cycle: u64\n");
}
#[test]
fn reference_bind_points_preserved() {
let w = synthesize_inline_workload("value={cycle}").unwrap();
assert_eq!(w.ops.len(), 1);
let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
assert_eq!(stmt, "value={cycle}");
let bindings = match &w.ops[0].bindings {
crate::model::BindingsDef::PolydatSource(s) => s.clone(),
_ => panic!("expected PolydatSource"),
};
assert_eq!(bindings, "input cycle: u64\n");
}
#[test]
fn semicolon_split_multiple_ops() {
let w = synthesize_inline_workload("read {{cycle}};write {{mod(cycle, 100)}}").unwrap();
assert_eq!(w.ops.len(), 2);
assert_eq!(w.ops[0].name, "inline_0");
assert_eq!(w.ops[1].name, "inline_1");
}
#[test]
fn ratio_prefix() {
let w = synthesize_inline_workload("3:read {{cycle}};1:write {{cycle}}").unwrap();
assert_eq!(w.ops.len(), 2);
assert_eq!(w.ops[0].params.get("ratio").unwrap().as_u64().unwrap(), 3);
assert!(!w.ops[1].params.contains_key("ratio"));
}
#[test]
fn ratio_one_not_stored() {
let w = synthesize_inline_workload("hello {{cycle}}").unwrap();
assert!(!w.ops[0].params.contains_key("ratio"));
}
#[test]
fn duplicate_expressions_share_output() {
let w = synthesize_inline_workload("a={{hash(cycle)}};b={{hash(cycle)}}").unwrap();
let stmt0 = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
let stmt1 = w.ops[1].op.get("stmt").unwrap().as_str().unwrap();
assert_eq!(stmt0, "a={__inline_0}");
assert_eq!(stmt1, "b={__inline_0}");
match &w.ops[0].bindings {
BindingsDef::PolydatSource(src) => {
let count = src.matches("__inline_").count();
assert_eq!(count, 1);
}
_ => panic!("expected PolydatSource"),
}
}
#[test]
fn empty_op_is_error() {
assert!(synthesize_inline_workload("").is_err());
assert!(synthesize_inline_workload(" ").is_err());
}
#[test]
fn mixed_reference_and_inline() {
let w = synthesize_inline_workload("id={{mod(hash(cycle), 1000)}} raw={cycle}").unwrap();
let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
assert_eq!(stmt, "id={__inline_0} raw={cycle}");
}
}