1use std::collections::HashMap;
14
15use crate::model::{BindingsDef, ParsedOp, Workload};
16
17pub fn synthesize_inline_workload(op_template: &str) -> Result<Workload, String> {
46 if op_template.trim().is_empty() {
47 return Err("op= value is empty".into());
48 }
49
50 if let Some(w) = try_polydat_block_workload(op_template) {
57 return Ok(w);
58 }
59
60 let segments = split_ops(op_template);
62
63 let mut inline_exprs: Vec<String> = Vec::new();
66 let mut expr_index: HashMap<String, usize> = HashMap::new();
67
68 for seg in &segments {
72 for expr in extract_inline_exprs(&seg.template) {
74 if !expr_index.contains_key(&expr) {
75 let idx = inline_exprs.len();
76 expr_index.insert(expr.clone(), idx);
77 inline_exprs.push(expr);
78 }
79 }
80 for bp in crate::bindpoints::extract_bind_points(&seg.template) {
82 if let crate::bindpoints::BindPoint::InlineDefinition(expr) = bp
83 && !expr_index.contains_key(&expr)
84 {
85 let idx = inline_exprs.len();
86 expr_index.insert(expr.clone(), idx);
87 inline_exprs.push(expr);
88 }
89 }
90 }
91
92 let mut polydat_source = String::from("input cycle: u64\n");
100 for (i, expr) in inline_exprs.iter().enumerate() {
101 polydat_source.push_str(&format!("__inline_{i} := {expr}\n"));
102 }
103
104 let mut ops = Vec::with_capacity(segments.len());
106 for (i, seg) in segments.iter().enumerate() {
107 let rewritten = rewrite_template(&seg.template, &expr_index);
108
109 let mut op = ParsedOp::simple(&format!("inline_{i}"), &rewritten);
110
111 if seg.ratio != 1 {
112 op.params.insert(
113 "ratio".to_string(),
114 serde_json::Value::Number(serde_json::Number::from(seg.ratio)),
115 );
116 }
117
118 op.tags.insert("name".to_string(), op.name.clone());
119 op.tags.insert("op".to_string(), op.name.clone());
120 op.tags.insert("block".to_string(), "inline".to_string());
121
122 op.bindings = BindingsDef::PolydatSource(polydat_source.clone());
123
124 ops.push(op);
125 }
126
127 Ok(Workload {
128 description: Some("inline workload".into()),
129 scenarios: HashMap::new(),
130 stop_when: Vec::new(),
131 ops,
132 bindings: crate::model::BindingsDef::default(),
133 params: HashMap::new(),
134 phases: HashMap::new(),
135 phase_order: Vec::new(),
136 declared_params: Vec::new(),
137 report: crate::report::Report::default(),
138 report_warnings: Vec::new(),
139 resolution_warnings: Vec::new(),
140 scenario_parse_errors: Vec::new(),
141 status_metrics: Vec::new(),
142 readouts: crate::model::ReadoutsBindings::default(),
143 wrappers: None,
144 implements: None,
145 stick_session: None,
146 })
147}
148
149fn try_polydat_block_workload(op_template: &str) -> Option<Workload> {
161 let source = build_polydat_candidate(op_template);
162 polydat::dsl::compile::compile_polydat(&source).ok()?;
164 let names = binding_wire_names(&source);
169 if names.is_empty() {
170 return None;
171 }
172 let mut op_fields: HashMap<String, serde_json::Value> = HashMap::new();
173 for n in &names {
174 op_fields.insert(n.clone(), serde_json::Value::String(format!("{{{n}}}")));
177 }
178 let mut op = ParsedOp::simple("inline_0", "");
179 op.op = op_fields;
180 op.bindings = BindingsDef::PolydatSource(source);
181 op.tags.insert("name".to_string(), "inline_0".to_string());
182 op.tags.insert("op".to_string(), "inline_0".to_string());
183 op.tags.insert("block".to_string(), "inline".to_string());
184
185 Some(Workload {
186 description: Some("inline polydat workload".into()),
187 scenarios: HashMap::new(),
188 stop_when: Vec::new(),
189 ops: vec![op],
190 bindings: crate::model::BindingsDef::default(),
191 params: HashMap::new(),
192 phases: HashMap::new(),
193 phase_order: Vec::new(),
194 declared_params: Vec::new(),
195 report: crate::report::Report::default(),
196 report_warnings: Vec::new(),
197 resolution_warnings: Vec::new(),
198 scenario_parse_errors: Vec::new(),
199 status_metrics: Vec::new(),
200 readouts: crate::model::ReadoutsBindings::default(),
201 wrappers: None,
202 implements: None,
203 stick_session: None,
204 })
205}
206
207fn build_polydat_candidate(op_template: &str) -> String {
214 let segs: Vec<String> = split_top_level_semicolons(op_template)
215 .into_iter()
216 .map(|s| s.trim().to_string())
217 .filter(|s| !s.is_empty())
218 .collect();
219 let mut lines = vec!["input cycle: u64".to_string()];
220 let last = segs.len().saturating_sub(1);
221 for (i, seg) in segs.iter().enumerate() {
222 if has_top_level_assignment(seg) {
223 lines.push(seg.clone());
224 } else if i == last {
225 lines.push(format!("out := {seg}"));
226 } else {
227 lines.push(format!("__expr_{i} := {seg}"));
228 }
229 }
230 lines.join("\n") + "\n"
231}
232
233pub(crate) fn binding_wire_names(source: &str) -> Vec<String> {
238 source
239 .lines()
240 .filter_map(|line| {
241 let line = line.trim();
242 if line.starts_with("input ") || !line.contains(":=") {
243 return None;
244 }
245 let lhs = line.split(":=").next()?.trim();
246 let name = lhs.split_whitespace().last()?;
247 let is_ident = !name.is_empty()
248 && name.chars().all(|c| c.is_alphanumeric() || c == '_')
249 && name
250 .chars()
251 .next()
252 .is_some_and(|c| c.is_alphabetic() || c == '_');
253 if is_ident && !name.starts_with("__") {
254 Some(name.to_string())
255 } else {
256 None
257 }
258 })
259 .collect()
260}
261
262fn has_top_level_assignment(s: &str) -> bool {
265 let bytes = s.as_bytes();
266 let mut depth = 0i32;
267 let mut i = 0;
268 while i < bytes.len() {
269 match bytes[i] {
270 b'{' => depth += 1,
271 b'}' => depth = (depth - 1).max(0),
272 b':' if depth == 0 && i + 1 < bytes.len() && bytes[i + 1] == b'=' => {
273 return true;
274 }
275 _ => {}
276 }
277 i += 1;
278 }
279 false
280}
281
282fn split_top_level_semicolons(s: &str) -> Vec<String> {
284 let mut out = Vec::new();
285 let mut cur = String::new();
286 let mut depth = 0i32;
287 for c in s.chars() {
288 match c {
289 '{' => {
290 depth += 1;
291 cur.push(c);
292 }
293 '}' => {
294 depth = (depth - 1).max(0);
295 cur.push(c);
296 }
297 ';' if depth == 0 => {
298 out.push(std::mem::take(&mut cur));
299 }
300 _ => cur.push(c),
301 }
302 }
303 if !cur.trim().is_empty() {
304 out.push(cur);
305 }
306 out
307}
308
309struct OpSegment {
312 template: String,
313 ratio: u64,
314}
315
316fn split_ops(input: &str) -> Vec<OpSegment> {
321 let mut segments = Vec::new();
322 let mut current = String::new();
323 let mut in_braces = 0u32;
324
325 for c in input.chars() {
326 match c {
327 '{' => {
328 in_braces += 1;
329 current.push(c);
330 }
331 '}' => {
332 in_braces = in_braces.saturating_sub(1);
333 current.push(c);
334 }
335 ';' if in_braces == 0 => {
336 let seg = current.trim().to_string();
337 if !seg.is_empty() {
338 segments.push(parse_segment(&seg));
339 }
340 current.clear();
341 }
342 _ => current.push(c),
343 }
344 }
345 let seg = current.trim().to_string();
346 if !seg.is_empty() {
347 segments.push(parse_segment(&seg));
348 }
349 segments
350}
351
352fn parse_segment(s: &str) -> OpSegment {
354 if let Some(colon_pos) = s.find(':') {
356 let prefix = &s[..colon_pos];
357 if !prefix.is_empty()
359 && prefix.chars().all(|c| c.is_ascii_digit())
360 && let Ok(ratio) = prefix.parse::<u64>()
361 {
362 return OpSegment {
363 template: s[colon_pos + 1..].trim().to_string(),
364 ratio,
365 };
366 }
367 }
368 OpSegment {
369 template: s.to_string(),
370 ratio: 1,
371 }
372}
373
374fn extract_inline_exprs(template: &str) -> Vec<String> {
376 let mut exprs = Vec::new();
377 let bytes = template.as_bytes();
378 let len = bytes.len();
379 let mut i = 0;
380
381 while i + 1 < len {
382 if bytes[i] == b'{' && bytes[i + 1] == b'{' {
383 let start = i + 2;
385 let mut depth = 1u32;
386 let mut j = start;
387 while j + 1 < len {
388 if bytes[j] == b'{' && bytes[j + 1] == b'{' {
389 depth += 1;
390 j += 2;
391 } else if bytes[j] == b'}' && bytes[j + 1] == b'}' {
392 depth -= 1;
393 if depth == 0 {
394 let expr = template[start..j].trim().to_string();
395 if !expr.is_empty() {
396 exprs.push(expr);
397 }
398 i = j + 2;
399 break;
400 }
401 j += 2;
402 } else {
403 j += 1;
404 }
405 }
406 if depth > 0 {
407 i += 2;
409 }
410 } else {
411 i += 1;
412 }
413 }
414 exprs
415}
416
417fn rewrite_template(template: &str, expr_index: &HashMap<String, usize>) -> String {
421 let after_double = rewrite_double_brace(template, expr_index);
423 rewrite_single_brace_exprs(&after_double, expr_index)
425}
426
427fn rewrite_single_brace_exprs(template: &str, expr_index: &HashMap<String, usize>) -> String {
428 let mut result = String::with_capacity(template.len());
429 let chars: Vec<char> = template.chars().collect();
430 let mut i = 0;
431
432 while i < chars.len() {
433 if chars[i] == '{' && (i + 1 >= chars.len() || chars[i + 1] != '{') {
434 let start = i + 1;
435 let mut depth = 1u32;
436 let mut j = start;
437 while j < chars.len() {
438 if chars[j] == '{' {
439 depth += 1;
440 }
441 if chars[j] == '}' {
442 depth -= 1;
443 if depth == 0 {
444 break;
445 }
446 }
447 j += 1;
448 }
449 if j < chars.len() {
450 let raw: String = chars[start..j].iter().collect();
451 let raw = raw.trim();
452
453 let expr = if let Some(e) = raw.strip_prefix(":=") {
455 Some(e.strip_suffix(":=").unwrap_or(e).trim())
456 } else if crate::bindpoints::is_expression_public(raw) {
457 Some(raw)
458 } else {
459 None
460 };
461
462 if let Some(expr) = expr {
463 if let Some(&idx) = expr_index.get(expr) {
464 result.push_str(&format!("{{__inline_{idx}}}"));
465 } else {
466 result.push('{');
468 result.push_str(raw);
469 result.push('}');
470 }
471 } else {
472 result.push('{');
474 result.push_str(raw);
475 result.push('}');
476 }
477 i = j + 1;
478 } else {
479 result.push(chars[i]);
480 i += 1;
481 }
482 } else {
483 result.push(chars[i]);
484 i += 1;
485 }
486 }
487 result
488}
489
490fn rewrite_double_brace(template: &str, expr_index: &HashMap<String, usize>) -> String {
491 let mut result = String::with_capacity(template.len());
492 let bytes = template.as_bytes();
493 let len = bytes.len();
494 let mut i = 0;
495
496 while i < len {
497 if i + 1 < len && bytes[i] == b'{' && bytes[i + 1] == b'{' {
498 let start = i + 2;
499 let mut depth = 1u32;
500 let mut j = start;
501 while j + 1 < len {
502 if bytes[j] == b'{' && bytes[j + 1] == b'{' {
503 depth += 1;
504 j += 2;
505 } else if bytes[j] == b'}' && bytes[j + 1] == b'}' {
506 depth -= 1;
507 if depth == 0 {
508 let expr = template[start..j].trim().to_string();
509 if let Some(&idx) = expr_index.get(&expr) {
510 result.push_str(&format!("{{__inline_{idx}}}"));
511 } else {
512 result.push_str(&template[i..j + 2]);
514 }
515 i = j + 2;
516 break;
517 }
518 j += 2;
519 } else {
520 j += 1;
521 }
522 }
523 if depth > 0 {
524 result.push_str(&template[i..]);
525 break;
526 }
527 } else {
528 result.push(bytes[i] as char);
529 i += 1;
530 }
531 }
532 result
533}
534
535#[cfg(test)]
538mod tests {
539 use super::*;
540
541 #[test]
542 fn simple_inline_binding() {
543 let w = synthesize_inline_workload("hello {{cycle}}").unwrap();
544 assert_eq!(w.ops.len(), 1);
545 assert_eq!(w.ops[0].name, "inline_0");
546 let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
547 assert_eq!(stmt, "hello {__inline_0}");
548 match &w.ops[0].bindings {
549 BindingsDef::PolydatSource(src) => {
550 assert!(src.contains("input cycle: u64"));
551 assert!(src.contains("__inline_0 := cycle"));
552 }
553 _ => panic!("expected PolydatSource bindings"),
554 }
555 }
556
557 #[test]
558 fn multiple_inline_bindings() {
559 let w = synthesize_inline_workload(
560 "id={{mod(hash(cycle), 100000)}} name={{number_to_words(cycle)}}",
561 )
562 .unwrap();
563 assert_eq!(w.ops.len(), 1);
564 let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
565 assert_eq!(stmt, "id={__inline_0} name={__inline_1}");
566 match &w.ops[0].bindings {
567 BindingsDef::PolydatSource(src) => {
568 assert!(src.contains("__inline_0 := mod(hash(cycle), 100000)"));
569 assert!(src.contains("__inline_1 := number_to_words(cycle)"));
570 }
571 _ => panic!("expected PolydatSource bindings"),
572 }
573 }
574
575 #[test]
576 fn bindings_block_op_becomes_polydat_fields() {
577 let w =
580 synthesize_inline_workload("x := cos(to_f64(cycle)); y := sin(to_f64(cycle))").unwrap();
581 assert_eq!(w.ops.len(), 1);
582 let keys: std::collections::BTreeSet<&str> =
583 w.ops[0].op.keys().map(|s| s.as_str()).collect();
584 assert!(keys.contains("x") && keys.contains("y"), "fields: {keys:?}");
585 assert!(!w.ops[0].op.contains_key("stmt"), "should not be a text op");
586 assert_eq!(w.ops[0].op.get("x").unwrap().as_str().unwrap(), "{x}");
587 assert!(matches!(w.ops[0].bindings, BindingsDef::PolydatSource(_)));
588 }
589
590 #[test]
591 fn bare_polydat_expr_becomes_out_field() {
592 let w = synthesize_inline_workload("cos(to_f64(cycle))").unwrap();
593 assert_eq!(w.ops.len(), 1);
594 assert!(
595 w.ops[0].op.contains_key("out"),
596 "fields: {:?}",
597 w.ops[0].op.keys().collect::<Vec<_>>()
598 );
599 }
600
601 #[test]
602 fn invalid_polydat_falls_back_to_text_template() {
603 let w = synthesize_inline_workload("x := not_a_real_fn(@@@)").unwrap();
606 assert_eq!(w.ops.len(), 1);
607 assert!(w.ops[0].op.contains_key("stmt"));
608 }
609
610 #[test]
611 fn detection_is_compile_driven_not_syntactic() {
612 let w = synthesize_inline_workload("id-{cycle}").unwrap();
615 assert!(w.ops[0].op.contains_key("stmt"));
616 }
617
618 #[test]
619 fn helpers_split_and_name_bindings() {
620 assert!(has_top_level_assignment("x := 1"));
621 assert!(!has_top_level_assignment("hello {{x := 1}}")); assert_eq!(split_top_level_semicolons("a := 1; b := 2").len(), 2);
623 let src = build_polydat_candidate("a := 1; sin(cycle)");
624 assert!(src.contains("a := 1"));
625 assert!(src.contains("out := sin(cycle)")); assert_eq!(
627 binding_wire_names("input cycle: u64\nt := 1\n__expr_0 := 2\nx := 3\n"),
628 vec!["t".to_string(), "x".to_string()]
629 ); }
631
632 #[test]
633 fn no_inline_bindings_plain_text() {
634 let w = synthesize_inline_workload("hello world").unwrap();
635 assert_eq!(w.ops.len(), 1);
636 let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
637 assert_eq!(stmt, "hello world");
638 let bindings = match &w.ops[0].bindings {
643 crate::model::BindingsDef::PolydatSource(s) => s.clone(),
644 _ => panic!("expected PolydatSource"),
645 };
646 assert_eq!(bindings, "input cycle: u64\n");
647 }
648
649 #[test]
650 fn reference_bind_points_preserved() {
651 let w = synthesize_inline_workload("value={cycle}").unwrap();
652 assert_eq!(w.ops.len(), 1);
653 let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
654 assert_eq!(stmt, "value={cycle}");
655 let bindings = match &w.ops[0].bindings {
661 crate::model::BindingsDef::PolydatSource(s) => s.clone(),
662 _ => panic!("expected PolydatSource"),
663 };
664 assert_eq!(bindings, "input cycle: u64\n");
665 }
666
667 #[test]
668 fn semicolon_split_multiple_ops() {
669 let w = synthesize_inline_workload("read {{cycle}};write {{mod(cycle, 100)}}").unwrap();
670 assert_eq!(w.ops.len(), 2);
671 assert_eq!(w.ops[0].name, "inline_0");
672 assert_eq!(w.ops[1].name, "inline_1");
673 }
674
675 #[test]
676 fn ratio_prefix() {
677 let w = synthesize_inline_workload("3:read {{cycle}};1:write {{cycle}}").unwrap();
678 assert_eq!(w.ops.len(), 2);
679 assert_eq!(w.ops[0].params.get("ratio").unwrap().as_u64().unwrap(), 3);
680 assert!(!w.ops[1].params.contains_key("ratio"));
682 }
683
684 #[test]
685 fn ratio_one_not_stored() {
686 let w = synthesize_inline_workload("hello {{cycle}}").unwrap();
687 assert!(!w.ops[0].params.contains_key("ratio"));
688 }
689
690 #[test]
691 fn duplicate_expressions_share_output() {
692 let w = synthesize_inline_workload("a={{hash(cycle)}};b={{hash(cycle)}}").unwrap();
693 let stmt0 = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
695 let stmt1 = w.ops[1].op.get("stmt").unwrap().as_str().unwrap();
696 assert_eq!(stmt0, "a={__inline_0}");
697 assert_eq!(stmt1, "b={__inline_0}");
698 match &w.ops[0].bindings {
699 BindingsDef::PolydatSource(src) => {
700 let count = src.matches("__inline_").count();
702 assert_eq!(count, 1);
703 }
704 _ => panic!("expected PolydatSource"),
705 }
706 }
707
708 #[test]
709 fn empty_op_is_error() {
710 assert!(synthesize_inline_workload("").is_err());
711 assert!(synthesize_inline_workload(" ").is_err());
712 }
713
714 #[test]
715 fn mixed_reference_and_inline() {
716 let w = synthesize_inline_workload("id={{mod(hash(cycle), 1000)}} raw={cycle}").unwrap();
717 let stmt = w.ops[0].op.get("stmt").unwrap().as_str().unwrap();
718 assert_eq!(stmt, "id={__inline_0} raw={cycle}");
719 }
720}