1use std::collections::{HashMap, HashSet};
19
20use crate::scope_kernel::ScopeKernel;
21
22use nmbrs_workload::bindpoints;
23use nmbrs_workload::model::ParsedOp;
24
25#[derive(Debug, Clone)]
27struct BindingFunc {
28 name: String,
29 args: Vec<String>,
30}
31
32fn parse_binding_chain(expr: &str) -> Vec<BindingFunc> {
36 let mut funcs = Vec::new();
37
38 for segment in expr.split(';') {
39 let segment = segment.trim();
40 if segment.is_empty() {
41 continue;
42 }
43
44 if let Some(paren_pos) = segment.find('(') {
46 let name = segment[..paren_pos].trim().to_string();
47 let args_str = &segment[paren_pos + 1..];
48 let args_str = args_str.trim_end_matches(')').trim();
49
50 let args: Vec<String> = if args_str.is_empty() {
51 Vec::new()
52 } else {
53 split_args(args_str)
55 };
56
57 funcs.push(BindingFunc { name, args });
58 } else {
59 funcs.push(BindingFunc {
61 name: segment.trim().to_string(),
62 args: Vec::new(),
63 });
64 }
65 }
66
67 funcs
68}
69
70fn split_args(s: &str) -> Vec<String> {
72 let mut args = Vec::new();
73 let mut current = String::new();
74 let mut depth = 0;
75 let mut in_quote = false;
76
77 for c in s.chars() {
78 match c {
79 '\'' if !in_quote => {
80 in_quote = true;
81 current.push(c);
82 }
83 '\'' if in_quote => {
84 in_quote = false;
85 current.push(c);
86 }
87 '(' if !in_quote => {
88 depth += 1;
89 current.push(c);
90 }
91 ')' if !in_quote => {
92 depth -= 1;
93 current.push(c);
94 }
95 ',' if depth == 0 && !in_quote => {
96 args.push(current.trim().to_string());
97 current = String::new();
98 }
99 _ => current.push(c),
100 }
101 }
102 if !current.trim().is_empty() {
103 args.push(current.trim().to_string());
104 }
105 args
106}
107
108pub fn probe_compile_level(func_name: &str) -> polydat::ast::CompileLevel {
131 let sig = match polydat::dsl::registry::lookup(func_name) {
132 Some(s) => s,
133 None => return polydat::ast::CompileLevel::Phase1,
134 };
135
136 let mut parts = Vec::new();
140 let mut has_wire = false;
141 for p in sig.params {
142 parts.push(p.example.to_string());
143 if p.slot_type.is_wire() {
144 has_wire = true;
145 }
146 }
147 if !has_wire && parts.is_empty() {
149 parts.push("cycle".to_string());
150 }
151
152 let source = format!("input cycle: u64\nout := {func_name}({})", parts.join(", "));
153
154 let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
159 polydat::dsl::compile::compile_polydat_interpreter(&source)
160 }));
161
162 match result {
163 Ok(Ok(kernel)) => kernel.program().last_node_compile_level(),
164 _ => polydat::ast::CompileLevel::Phase1,
165 }
166}
167
168pub fn compile_bindings(ops: &[ParsedOp]) -> Result<ScopeKernel, String> {
169 compile_bindings_with_path(ops, None)
170}
171
172fn collect_param_bindings(
174 params: &HashMap<String, serde_json::Value>,
175 exclude: &[String],
176 required: &mut Vec<String>,
177) {
178 for (key, value) in params.iter() {
179 if key == "gutter" {
186 continue;
187 }
188 collect_json_bindings(value, exclude, required);
189 }
190}
191
192pub fn collect_param_bindings_into(
194 params: &HashMap<String, serde_json::Value>,
195 exclude: &[String],
196 required: &mut Vec<String>,
197) {
198 collect_param_bindings(params, exclude, required);
199}
200
201fn collect_json_bindings(
202 value: &serde_json::Value,
203 exclude: &[String],
204 required: &mut Vec<String>,
205) {
206 match value {
207 serde_json::Value::String(s) => {
208 for name in bindpoints::referenced_bindings(s) {
209 if !required.contains(&name) && !exclude.contains(&name) {
210 required.push(name);
211 }
212 }
213 }
214 serde_json::Value::Object(map) => {
215 for v in map.values() {
216 collect_json_bindings(v, exclude, required);
217 }
218 }
219 serde_json::Value::Array(arr) => {
220 for v in arr {
221 collect_json_bindings(v, exclude, required);
222 }
223 }
224 _ => {}
225 }
226}
227
228pub fn compile_bindings_with_path(
229 ops: &[ParsedOp],
230 source_dir: Option<&std::path::Path>,
231) -> Result<ScopeKernel, String> {
232 compile_bindings_with_opts(ops, source_dir, false)
233}
234
235pub fn compile_from_scope(
250 scope: &crate::scope::BindingScope,
251 source_dir: Option<&std::path::Path>,
252 polydat_lib_paths: Vec<std::path::PathBuf>,
253 strict: bool,
254 context: &str,
255 cursor_limit: Option<u64>,
256 pragmas: &polydat::dsl::pragmas::PragmaSet,
257) -> Result<ScopeKernel, String> {
258 let (source, options) = scope_source_and_options(
259 scope,
260 source_dir,
261 polydat_lib_paths,
262 strict,
263 context,
264 cursor_limit,
265 pragmas,
266 );
267 compile_scope_kernel(&source, &options)
268}
269
270fn scope_source_and_options(
272 scope: &crate::scope::BindingScope,
273 source_dir: Option<&std::path::Path>,
274 polydat_lib_paths: Vec<std::path::PathBuf>,
275 strict: bool,
276 context: &str,
277 cursor_limit: Option<u64>,
278 pragmas: &polydat::dsl::pragmas::PragmaSet,
279) -> (String, polydat::dsl::compile::CompileOptions) {
280 let body = scope.emit();
281 let required = scope.required_outputs();
282 let source = prepend_effective_pragmas(pragmas, &body);
283 let options = polydat::dsl::compile::CompileOptions {
284 source_dir: source_dir.map(std::path::Path::to_path_buf),
285 lib_paths: polydat_lib_paths,
286 required_outputs: required,
287 strict,
288 context: context.to_string(),
289 cursor_limit,
290 ..Default::default()
291 };
292 (source, options)
293}
294
295pub fn compile_scope_kernel(
307 source: &str,
308 options: &polydat::dsl::compile::CompileOptions,
309) -> Result<ScopeKernel, String> {
310 let options = polydat::dsl::compile::CompileOptions {
311 resources: Some(
312 options
313 .resources
314 .clone()
315 .unwrap_or_else(crate::resource_pool::pool_resources),
316 ),
317 ..options.clone()
318 };
319 let options = &options;
320 let interpreter =
321 polydat::dsl::compile::compile_polydat_interpreter_with_options(source, options, None)
322 .map_err(|e| e.to_string())?;
323 let image =
324 crate::fiber_engine::source_image(interpreter.program(), source, options, &options.context);
325 Ok(ScopeKernel::root(interpreter, image))
326}
327
328pub fn compile_scope_program(
332 source: &str,
333 options: &polydat::dsl::compile::CompileOptions,
334) -> Result<std::sync::Arc<polydat::kernel::PolydatProgram>, String> {
335 polydat::dsl::compile::compile_polydat_interpreter_with_options(source, options, None)
336 .map(|k| k.program().clone())
337 .map_err(|e| e.to_string())
338}
339
340pub(crate) fn prepend_effective_pragmas(
351 pragmas: &polydat::dsl::pragmas::PragmaSet,
352 body: &str,
353) -> String {
354 let mut out = String::new();
355 if pragmas.strict_types() && pragmas.strict_values() {
356 out.push_str("pragma strict\n");
357 } else if pragmas.strict_types() {
358 out.push_str("pragma strict_types\n");
359 } else if pragmas.strict_values() {
360 out.push_str("pragma strict_values\n");
361 }
362 if !out.is_empty() {
363 out.push('\n');
364 }
365 out.push_str(body);
366 out
367}
368
369pub const SESSION_START: &str = "session_start";
372
373#[allow(clippy::too_many_arguments)]
401pub fn build_workload_root_kernel(
402 parent: &ScopeKernel,
403 ops: &[ParsedOp],
404 source_dir: Option<&std::path::Path>,
405 polydat_lib_paths: Vec<std::path::PathBuf>,
406 strict: bool,
407 extra_required: &[String],
408 context: &str,
409 cursor_limit: Option<u64>,
410 workload_params: &std::collections::HashMap<String, String>,
411 workload_level_polydat: Option<&str>,
412) -> Result<ScopeKernel, String> {
413 let mut scope = crate::scope::build_scope(
417 ops,
418 &std::collections::HashMap::new(), &[], workload_params,
421 &std::collections::HashMap::new(), None, &[], None, )?;
426
427 if let Some(extra) = workload_level_polydat
432 && !extra.trim().is_empty()
433 {
434 scope.ingest_polydat_source(extra, crate::scope::BindingOrigin::Inherited);
435 }
436 let session_clock = !workload_params.contains_key(SESSION_START)
446 && !scope.defined_names().contains(SESSION_START);
447 if session_clock {
448 scope.ingest_polydat_source(
449 &format!("const {SESSION_START} := current_epoch_millis()\n"),
450 crate::scope::BindingOrigin::Inherited,
451 );
452 }
453 scope.validate().map_err(|e| format!("{context}: {e}"))?;
454
455 let mut scope_required = scope.required_outputs();
464 for name in extra_required {
465 if !scope_required.contains(name) {
466 scope_required.push(name.clone());
467 }
468 }
469 if session_clock && !scope_required.iter().any(|n| n == SESSION_START) {
470 scope_required.push(SESSION_START.to_string());
471 }
472 let mut param_names: Vec<&String> = workload_params.keys().collect();
473 param_names.sort();
474 for name in param_names {
475 if !scope_required.contains(name) {
476 scope_required.push(name.clone());
477 }
478 }
479
480 let mut source = scope.emit();
484 if !source.lines().any(|l| l.trim_start().starts_with("input ")) {
495 source = format!("input cycle: u64\n{source}");
496 }
497 let opts = polydat::kernel::subcontext::CompileOptions {
498 workload_dir: source_dir.map(|p| p.to_path_buf()),
499 polydat_lib_paths,
500 strict,
501 required_outputs: scope_required,
502 context_label: Some(context.to_string()),
503 cursor_limit,
504 ..Default::default()
505 };
506 let mut inherited_param_names: Vec<String> = workload_params.keys().cloned().collect();
522 inherited_param_names.sort();
523 ScopeKernel::build_under(
524 parent.kernel(),
525 crate::scope_kernel::SourceMatter::source(context, source, opts)
526 .inherited(inherited_param_names),
527 )
528}
529
530pub fn compile_bindings_with_opts(
537 ops: &[ParsedOp],
538 source_dir: Option<&std::path::Path>,
539 strict: bool,
540) -> Result<ScopeKernel, String> {
541 use nmbrs_workload::model::BindingsDef;
542
543 let polydat_source = ops.iter().find_map(|op| {
545 if let BindingsDef::PolydatSource(src) = &op.bindings {
546 if !src.trim().is_empty() {
547 Some(src.clone())
548 } else {
549 None
550 }
551 } else {
552 None
553 }
554 });
555
556 if let Some(source) = polydat_source {
557 let mut required: Vec<String> = Vec::new();
561 for op in ops {
562 for value in op.op.values() {
563 if let Some(s) = value.as_str() {
564 for name in bindpoints::referenced_bindings(s) {
565 if !required.contains(&name) {
566 required.push(name);
567 }
568 }
569 }
570 }
571 }
572 let options = polydat::dsl::compile::CompileOptions {
573 source_dir: source_dir.map(std::path::Path::to_path_buf),
574 required_outputs: required,
575 strict,
576 ..Default::default()
577 };
578 return compile_scope_kernel(&source, &options);
579 }
580
581 let mut all_bindings: HashMap<String, String> = HashMap::new();
586 for op in ops {
587 if let BindingsDef::Map(map) = &op.bindings {
588 for (name, expr) in map {
589 all_bindings
590 .entry(name.clone())
591 .or_insert_with(|| expr.clone());
592 }
593 }
594 }
595
596 let mut required: Vec<String> = Vec::new();
598 for op in ops {
599 for value in op.op.values() {
600 if let Some(s) = value.as_str() {
601 for name in bindpoints::referenced_bindings(s) {
602 if !required.contains(&name) {
603 required.push(name);
604 }
605 }
606 }
607 }
608 }
609
610 let mut polydat_lines: Vec<String> = Vec::new();
612 polydat_lines.push("input cycle: u64".into());
613
614 for (binding_name, expr) in &all_bindings {
615 let chain = parse_binding_chain(expr);
616 if chain.is_empty() {
617 return Err(format!("empty binding expression for '{binding_name}'"));
618 }
619
620 let mut prev_wire = "cycle".to_string();
623
624 for (i, func) in chain.iter().enumerate() {
625 let is_last = i == chain.len() - 1;
626 let target = if is_last {
627 binding_name.clone()
628 } else {
629 format!("__chain_{binding_name}_{i}")
630 };
631
632 let (func_name, extra_args) = translate_legacy_func(&func.name, &func.args);
634 let mut call_args = vec![prev_wire.clone()];
635 for arg in &func.args {
636 call_args.push(strip_java_long_suffix(arg.trim()).to_string());
637 }
638 call_args.extend(extra_args);
639
640 polydat_lines.push(format!(
641 "{target} := {func_name}({args})",
642 args = call_args.join(", ")
643 ));
644
645 prev_wire = target;
646 }
647 }
648
649 let coord_names: HashSet<String> = ["cycle".to_string()].into_iter().collect();
654 let mut missing: Vec<String> = Vec::new();
655 for name in &required {
656 if !all_bindings.contains_key(name) && !coord_names.contains(name) {
657 missing.push(name.clone());
658 }
659 }
660 if !missing.is_empty() {
661 return Err(format!(
662 "undeclared bind point references: {}. Add these to your bindings section.",
663 missing.join(", ")
664 ));
665 }
666
667 let polydat_source = polydat_lines.join("\n");
668 let options = polydat::dsl::compile::CompileOptions {
669 source_dir: source_dir.map(std::path::Path::to_path_buf),
670 required_outputs: required,
671 strict,
672 ..Default::default()
673 };
674 compile_scope_kernel(&polydat_source, &options)
675}
676
677fn translate_legacy_func(name: &str, args: &[String]) -> (String, Vec<String>) {
691 match name.to_lowercase().as_str() {
692 "hash" => ("hash".into(), vec![]),
694 "identity" => ("identity".into(), vec![]),
695 "add" => ("add".into(), vec![]),
696 "mul" => ("mul".into(), vec![]),
697 "div" => ("div".into(), vec![]),
698 "mod" => ("mod".into(), vec![]),
699 "clamp" => ("clamp".into(), vec![]),
700
701 "tostring" | "to_string" => ("format_u64".into(), vec!["10".into()]),
703 "tohexstring" => ("format_u64".into(), vec!["16".into()]),
704 "tooctalstring" => ("format_u64".into(), vec!["8".into()]),
705 "tobinarystring" => ("format_u64".into(), vec!["2".into()]),
706
707 "uniform" => {
712 if args.len() >= 2 {
713 ("hash_range".into(), vec![])
716 } else {
717 ("hash_range".into(), vec![])
718 }
719 }
720
721 "normal" | "gaussian" => ("icd_normal".into(), vec![]),
723 "zipf" => ("dist_zipf".into(), vec![]),
724
725 "hashrange" | "hash_range" => ("hash_range".into(), vec![]),
727 "hashinterval" | "hash_interval" => ("hash_interval".into(), vec![]),
728
729 "format" | "printf" => ("printf".into(), vec![]),
731 "numbernamesto_string" | "numbernames" => ("number_to_words".into(), vec![]),
732
733 "shuffle" => ("shuffle".into(), vec![]),
735
736 _ => {
738 (name.to_lowercase(), vec![])
740 }
741 }
742}
743
744fn strip_java_long_suffix(arg: &str) -> &str {
746 arg.strip_suffix('L')
747 .or_else(|| arg.strip_suffix('l'))
748 .unwrap_or(arg)
749}
750
751pub fn legacy_chain_map_to_polydat_lines(
762 map: &std::collections::HashMap<String, String>,
763) -> Result<String, String> {
764 let mut polydat_lines: Vec<String> = Vec::new();
765 for (binding_name, expr) in map {
766 let chain = parse_binding_chain(expr);
767 if chain.is_empty() {
768 return Err(format!("empty binding expression for '{binding_name}'"));
769 }
770 let mut prev_wire = "cycle".to_string();
771 for (i, func) in chain.iter().enumerate() {
772 let is_last = i == chain.len() - 1;
773 let target = if is_last {
774 binding_name.clone()
775 } else {
776 format!("__chain_{binding_name}_{i}")
777 };
778 let (func_name, extra_args) = translate_legacy_func(&func.name, &func.args);
779 let mut call_args = vec![prev_wire.clone()];
780 for arg in &func.args {
781 call_args.push(strip_java_long_suffix(arg.trim()).to_string());
782 }
783 call_args.extend(extra_args);
784 polydat_lines.push(format!(
785 "{target} := {func_name}({args})",
786 args = call_args.join(", ")
787 ));
788 prev_wire = target;
789 }
790 }
791 Ok(polydat_lines.join("\n"))
792}
793
794#[cfg(test)]
795mod tests {
796 use super::*;
797
798 use polydat::dsl::pragmas::{Pragma, PragmaSet};
799
800 #[test]
801 fn prepend_pragmas_strict_alias() {
802 let pragmas = PragmaSet {
803 entries: vec![Pragma {
804 name: "strict".into(),
805 args: vec![],
806 line: 1,
807 }],
808 };
809 let body = "id := cycle\n";
810 let out = prepend_effective_pragmas(&pragmas, body);
811 assert!(out.starts_with("pragma strict\n"));
813 assert!(out.contains("id := cycle"));
814 }
815
816 #[test]
817 fn prepend_pragmas_individual_modes() {
818 let pragmas = PragmaSet {
819 entries: vec![Pragma {
820 name: "strict_values".into(),
821 args: vec![],
822 line: 1,
823 }],
824 };
825 let out = prepend_effective_pragmas(&pragmas, "x := cycle");
826 assert!(out.starts_with("pragma strict_values\n"));
827 assert!(!out.contains("strict_types"));
828 }
829
830 #[test]
831 fn prepend_pragmas_no_op_when_empty() {
832 let pragmas = PragmaSet::default();
833 let out = prepend_effective_pragmas(&pragmas, "x := cycle");
834 assert_eq!(out, "x := cycle");
835 }
836
837 #[test]
838 fn prepend_pragmas_walks_parent_chain() {
839 let parent = PragmaSet {
844 entries: vec![Pragma {
845 name: "strict_values".into(),
846 args: vec![],
847 line: 1,
848 }],
849 };
850 let child = parent.nested(&[]);
851 let out = prepend_effective_pragmas(&child, "x := cycle");
852 assert!(
853 out.starts_with("pragma strict_values\n"),
854 "expected pragma to flow from parent chain, got:\n{out}"
855 );
856 }
857
858 #[test]
859 fn parse_simple_chain() {
860 let chain = parse_binding_chain("Hash(); Mod(1000000)");
861 assert_eq!(chain.len(), 2);
862 assert_eq!(chain[0].name, "Hash");
863 assert!(chain[0].args.is_empty());
864 assert_eq!(chain[1].name, "Mod");
865 assert_eq!(chain[1].args, vec!["1000000"]);
866 }
867
868 #[test]
869 fn parse_identity() {
870 let chain = parse_binding_chain("Identity()");
871 assert_eq!(chain.len(), 1);
872 assert_eq!(chain[0].name, "Identity");
873 }
874
875 #[test]
876 fn parse_with_string_arg() {
877 let chain = parse_binding_chain("Template('user-{}', ToString())");
878 assert_eq!(chain.len(), 1);
879 assert_eq!(chain[0].name, "Template");
880 assert_eq!(chain[0].args.len(), 2);
881 }
882
883 #[test]
884 fn parse_long_chain() {
885 let chain = parse_binding_chain("Add(10); Hash(); Mod(100); ToString()");
886 assert_eq!(chain.len(), 4);
887 assert_eq!(chain[0].name, "Add");
888 assert_eq!(chain[1].name, "Hash");
889 assert_eq!(chain[2].name, "Mod");
890 assert_eq!(chain[3].name, "ToString");
891 }
892
893 #[test]
894 fn parse_with_long_suffix() {
895 let chain = parse_binding_chain("Mod(1000000000L)");
896 assert_eq!(chain[0].args, vec!["1000000000L"]);
897 }
900
901 #[test]
902 fn compile_identity_binding() {
903 let ops = vec![{
904 let mut op = ParsedOp::simple("test", "{myval}");
905 op.bindings.insert("myval".into(), "Identity()".into());
906 op
907 }];
908 let mut kernel = compile_bindings(&ops).unwrap();
909 kernel.set_inputs(&[42]);
910 assert_eq!(kernel.pull("myval").as_u64(), 42);
911 }
912
913 #[test]
914 fn compile_hash_mod_binding() {
915 let ops = vec![{
916 let mut op = ParsedOp::simple("test", "{id}");
917 op.bindings
918 .insert("id".into(), "Hash(); Mod(1000000)".into());
919 op
920 }];
921 let mut kernel = compile_bindings(&ops).unwrap();
922 kernel.set_inputs(&[42]);
923 let val = kernel.pull("id").as_u64();
924 assert!(val < 1_000_000, "got {val}");
925 }
926
927 #[test]
928 fn compile_hash_mod_deterministic() {
929 let ops = vec![{
930 let mut op = ParsedOp::simple("test", "{id}");
931 op.bindings.insert("id".into(), "Hash(); Mod(100)".into());
932 op
933 }];
934 let mut kernel = compile_bindings(&ops).unwrap();
935 kernel.set_inputs(&[42]);
936 let v1 = kernel.pull("id").as_u64();
937 kernel.set_inputs(&[42]);
938 let v2 = kernel.pull("id").as_u64();
939 assert_eq!(v1, v2);
940 }
941
942 #[test]
943 fn compile_multiple_bindings() {
944 let ops = vec![{
945 let mut op = ParsedOp::simple("test", "{a} {b}");
946 op.bindings.insert("a".into(), "Identity()".into());
947 op.bindings.insert("b".into(), "Hash(); Mod(100)".into());
948 op
949 }];
950 let mut kernel = compile_bindings(&ops).unwrap();
951 kernel.set_inputs(&[5]);
952 assert_eq!(kernel.pull("a").as_u64(), 5);
953 assert!(kernel.pull("b").as_u64() < 100);
954 }
955
956 #[test]
957 fn compile_rejects_undeclared_bind_points() {
958 let ops = vec![ParsedOp::simple("test", "val={mystery}")];
960 let result = compile_bindings(&ops);
961 assert!(result.is_err());
962 assert!(result.unwrap_err().contains("undeclared bind point"));
963 }
964
965 #[test]
966 fn compile_add_chain() {
967 let ops = vec![{
968 let mut op = ParsedOp::simple("test", "{val}");
969 op.bindings
970 .insert("val".into(), "Add(100); Mod(1000)".into());
971 op
972 }];
973 let mut kernel = compile_bindings(&ops).unwrap();
974 kernel.set_inputs(&[5]);
975 assert_eq!(kernel.pull("val").as_u64(), 105);
977 }
978
979 #[test]
980 fn compile_provides_cycle_output() {
981 let ops = vec![ParsedOp::simple("test", "cycle={cycle}")];
982 let mut kernel = compile_bindings(&ops).unwrap();
983 kernel.set_inputs(&[99]);
984 assert_eq!(kernel.pull("cycle").as_u64(), 99);
985 }
986
987 #[test]
988 fn legacy_tostring_translates() {
989 let (name, _) = translate_legacy_func("ToString", &[]);
990 assert_eq!(name, "format_u64");
991 }
992
993 #[test]
994 fn legacy_uniform_translates() {
995 let (name, _) = translate_legacy_func("Uniform", &["0".into(), "1000".into()]);
996 assert_eq!(name, "hash_range");
997 }
998}