1use symplex::matrix::{CodegenOptions, MathBackend, Matrix, Precision};
45use symplex::prelude::*;
46use symplex::robotics::DhLink;
47
48use std::fs;
49use std::path::{Path, PathBuf};
50
51pub struct CodeGen {
60 functions: Vec<GeneratedFn>,
61 options: CodegenOptions,
62 preamble: Vec<String>,
63 test_points: Vec<Vec<f64>>,
64 generate_tests: bool,
65}
66
67enum GeneratedFn {
68 Scalar {
69 name: String,
70 expr: Ex,
71 params: Vec<String>,
72 },
73 Matrix {
74 name: String,
75 matrix: Matrix,
76 params: Vec<String>,
77 },
78}
79
80impl CodeGen {
81 pub fn new() -> Self {
83 Self {
84 functions: Vec::new(),
85 options: CodegenOptions::default(),
86 preamble: Vec::new(),
87 test_points: Vec::new(),
88 generate_tests: false,
89 }
90 }
91
92 pub fn options(mut self, options: CodegenOptions) -> Self {
124 self.options = options;
125 self
126 }
127
128 pub fn no_std(mut self, enabled: bool) -> Self {
130 if enabled {
131 self.options.math_backend = MathBackend::CfgGated;
132 } else {
133 self.options.math_backend = MathBackend::Std;
134 }
135 self
136 }
137
138 pub fn precision_f32(mut self) -> Self {
140 self.options.precision = Precision::F32;
141 self
142 }
143
144 pub fn inline(mut self, enabled: bool) -> Self {
146 self.options.inline = enabled;
147 self
148 }
149
150 pub fn add_scalar_fn(mut self, name: &str, expr: &Ex, params: &[&str]) -> Self {
155 self.functions.push(GeneratedFn::Scalar {
156 name: name.to_string(),
157 expr: expr.clone(),
158 params: params.iter().map(|s| s.to_string()).collect(),
159 });
160 self
161 }
162
163 pub fn add_matrix_fn(mut self, name: &str, matrix: &Matrix, params: &[&str]) -> Self {
168 self.functions.push(GeneratedFn::Matrix {
169 name: name.to_string(),
170 matrix: matrix.clone(),
171 params: params.iter().map(|s| s.to_string()).collect(),
172 });
173 self
174 }
175
176 pub fn with_tests(mut self, enabled: bool) -> Self {
181 self.generate_tests = enabled;
182 self
183 }
184
185 pub fn add_test_point(mut self, point: &[f64]) -> Self {
191 self.test_points.push(point.to_vec());
192 self
193 }
194
195 pub fn generate(&self) -> Result<String, Box<dyn std::error::Error>> {
205 let mut output = String::new();
206
207 output.push_str("// Auto-generated by symplex-build. Do not edit.\n\n");
209
210 for line in &self.preamble {
212 output.push_str(line);
213 output.push('\n');
214 }
215
216 let emit_cfg_module = self.options.math_backend == MathBackend::CfgGated;
217
218 if emit_cfg_module {
219 append_cfg_gated_module(&mut output, self.options.precision);
221 output.push('\n');
222 }
223
224 let fn_options = CodegenOptions {
227 emit_runtime: false,
228 ..self.options.clone()
229 };
230
231 let mut functions = String::new();
232 let mut first_fn = true;
233 for gfn in &self.functions {
234 if !first_fn {
235 functions.push('\n');
236 }
237 first_fn = false;
238
239 let code = match gfn {
240 GeneratedFn::Scalar { name, expr, params } => {
241 let param_refs: Vec<&str> = params.iter().map(|s| s.as_str()).collect();
242 expr.to_rust_fn_with_options(name, ¶m_refs, &fn_options)?
243 }
244 GeneratedFn::Matrix {
245 name,
246 matrix,
247 params,
248 } => {
249 let param_refs: Vec<&str> = params.iter().map(|s| s.as_str()).collect();
250 matrix.to_rust_fn_with_options(name, ¶m_refs, &fn_options)?
251 }
252 };
253
254 if emit_cfg_module {
257 let stripped = strip_cfg_gated_module(&code);
258 functions.push_str(&stripped);
259 } else {
260 functions.push_str(&code);
261 }
262 functions.push('\n');
263 }
264
265 if self.options.emit_runtime
266 && let Some(runtime) = self.options.runtime_module_for(&functions)
267 {
268 output.push_str(&runtime);
269 output.push_str("\n\n");
270 }
271 output.push_str(&functions);
272
273 if self.generate_tests && !self.test_points.is_empty() {
275 output.push('\n');
276 output.push_str("#[cfg(test)]\n");
277 output.push_str("mod generated_tests {\n");
278 output.push_str(" use super::*;\n\n");
279
280 for (fn_idx, gfn) in self.functions.iter().enumerate() {
281 let (fn_name, param_count) = match gfn {
282 GeneratedFn::Scalar { name, params, .. } => (name.as_str(), params.len()),
283 GeneratedFn::Matrix { name, params, .. } => (name.as_str(), params.len()),
284 };
285
286 for (pt_idx, point) in self.test_points.iter().enumerate() {
287 if point.len() != param_count {
288 continue;
289 }
290 output.push_str(&format!(
291 " #[test]\n fn test_{fn_name}_point_{pt_idx}() {{\n"
292 ));
293
294 let args: Vec<String> = point
295 .iter()
296 .map(|v| {
297 let float_ty = match self.options.precision {
298 Precision::F64 => "f64",
299 Precision::F32 => "f32",
300 };
301 format!("{v}_{float_ty}")
302 })
303 .collect();
304 let args_str = args.join(", ");
305
306 match &self.functions[fn_idx] {
307 GeneratedFn::Scalar { .. } => {
308 output.push_str(&format!(
309 " let result = {fn_name}({args_str});\n"
310 ));
311 output.push_str(" assert!(result.is_finite(), \"expected finite result, got {}\", result);\n");
312 }
313 GeneratedFn::Matrix { matrix, .. } => {
314 let total = matrix.nrows() * matrix.ncols();
315 output.push_str(&format!(
316 " let result = {fn_name}({args_str});\n"
317 ));
318 output.push_str(&format!(" for i in 0..{total} {{\n"));
319 output.push_str(" assert!(result[i].is_finite(), \"entry {} is not finite: {}\", i, result[i]);\n");
320 output.push_str(" }\n");
321 }
322 }
323
324 output.push_str(" }\n\n");
325 }
326 }
327
328 output.push_str("}\n");
329 }
330
331 Ok(output)
332 }
333
334 pub fn write_to_out_dir(&self, filename: &str) -> Result<(), Box<dyn std::error::Error>> {
339 let out_dir = std::env::var("OUT_DIR")
340 .map_err(|_| "OUT_DIR not set — this function must be called from a build script")?;
341 let path = PathBuf::from(out_dir).join(filename);
342 let code = self.generate()?;
343 fs::write(&path, code)?;
344 println!("cargo:rerun-if-changed=build.rs");
345 Ok(())
346 }
347
348 pub fn write_to_path(&self, path: impl AsRef<Path>) -> Result<(), Box<dyn std::error::Error>> {
350 let code = self.generate()?;
351 if let Some(parent) = path.as_ref().parent() {
352 fs::create_dir_all(parent)?;
353 }
354 fs::write(path, code)?;
355 Ok(())
356 }
357}
358
359impl Default for CodeGen {
360 fn default() -> Self {
361 Self::new()
362 }
363}
364
365fn append_cfg_gated_module(out: &mut String, precision: Precision) {
379 let ft = match precision {
380 Precision::F64 => "f64",
381 Precision::F32 => "f32",
382 };
383
384 let funcs = [
385 "sin", "cos", "tan", "exp", "ln", "abs", "sqrt", "cbrt", "asin", "acos", "atan", "sinh",
386 "cosh", "tanh", "asinh", "acosh", "atanh", "floor", "ceil", "signum",
387 ];
388
389 out.push_str("#[cfg(feature = \"std\")]\n");
391 out.push_str("mod math {\n");
392 for func in &funcs {
393 out.push_str(&format!(
394 " #[inline] pub fn {func}(x: {ft}) -> {ft} {{ x.{func}() }}\n"
395 ));
396 }
397 out.push_str(&format!(
398 " #[inline] pub fn atan2(y: {ft}, x: {ft}) -> {ft} {{ y.atan2(x) }}\n"
399 ));
400 out.push_str(&format!(
401 " #[inline] pub fn powf(base: {ft}, exp: {ft}) -> {ft} {{ base.powf(exp) }}\n"
402 ));
403 out.push_str(&format!(
404 " #[inline] pub fn powi(base: {ft}, exp: i32) -> {ft} {{ base.powi(exp) }}\n"
405 ));
406 out.push_str(&format!(
407 " #[inline] pub fn min(a: {ft}, b: {ft}) -> {ft} {{ a.min(b) }}\n"
408 ));
409 out.push_str(&format!(
410 " #[inline] pub fn max(a: {ft}, b: {ft}) -> {ft} {{ a.max(b) }}\n"
411 ));
412 out.push_str(&format!(
413 " #[inline] pub fn expm1(x: {ft}) -> {ft} {{ x.exp_m1() }}\n"
414 ));
415 out.push_str(&format!(
416 " #[inline] pub fn log1p(x: {ft}) -> {ft} {{ x.ln_1p() }}\n"
417 ));
418 out.push_str(&format!(
419 " #[inline] pub fn log2(x: {ft}) -> {ft} {{ x.log2() }}\n"
420 ));
421 out.push_str(&format!(
422 " #[inline] pub fn exp2(x: {ft}) -> {ft} {{ x.exp2() }}\n"
423 ));
424 out.push_str(&format!(
425 " #[inline] pub fn fma(a: {ft}, b: {ft}, c: {ft}) -> {ft} {{ a.mul_add(b, c) }}\n"
426 ));
427 out.push_str(&format!(
428 " #[inline] pub fn sin_cos(x: {ft}) -> ({ft}, {ft}) {{ x.sin_cos() }}\n"
429 ));
430 out.push_str("}\n\n");
431
432 out.push_str("#[cfg(not(feature = \"std\"))]\n");
436 out.push_str("mod math {\n");
437 let libm_funcs = [
438 "sin", "cos", "tan", "exp", "sqrt", "cbrt", "asin", "acos", "atan", "sinh", "cosh", "tanh",
439 "asinh", "acosh", "atanh", "floor", "ceil",
440 ];
441 for func in &libm_funcs {
442 out.push_str(&format!(
443 " #[inline] pub fn {func}(x: {ft}) -> {ft} {{ libm::{func}(x as f64) as {ft} }}\n"
444 ));
445 }
446 out.push_str(&format!(
447 " #[inline] pub fn abs(x: {ft}) -> {ft} {{ libm::fabs(x as f64) as {ft} }}\n"
448 ));
449 out.push_str(&format!(
450 " #[inline] pub fn ln(x: {ft}) -> {ft} {{ libm::log(x as f64) as {ft} }}\n"
451 ));
452 out.push_str(&format!(
453 " #[inline] pub fn signum(x: {ft}) -> {ft} {{ if x > 0.0 {{ 1.0 }} else if x < 0.0 {{ -1.0 }} else {{ 0.0 }} }}\n"
454 ));
455 out.push_str(&format!(
456 " #[inline] pub fn atan2(y: {ft}, x: {ft}) -> {ft} {{ libm::atan2(y as f64, x as f64) as {ft} }}\n"
457 ));
458 out.push_str(&format!(
459 " #[inline] pub fn powf(base: {ft}, exp: {ft}) -> {ft} {{ libm::pow(base as f64, exp as f64) as {ft} }}\n"
460 ));
461 out.push_str(&format!(
462 " #[inline] pub fn powi(base: {ft}, exp: i32) -> {ft} {{ libm::pow(base as f64, exp as f64) as {ft} }}\n"
463 ));
464 out.push_str(&format!(
465 " #[inline] pub fn min(a: {ft}, b: {ft}) -> {ft} {{ libm::fmin(a as f64, b as f64) as {ft} }}\n"
466 ));
467 out.push_str(&format!(
468 " #[inline] pub fn max(a: {ft}, b: {ft}) -> {ft} {{ libm::fmax(a as f64, b as f64) as {ft} }}\n"
469 ));
470 out.push_str(&format!(
471 " #[inline] pub fn expm1(x: {ft}) -> {ft} {{ libm::expm1(x as f64) as {ft} }}\n"
472 ));
473 out.push_str(&format!(
474 " #[inline] pub fn log1p(x: {ft}) -> {ft} {{ libm::log1p(x as f64) as {ft} }}\n"
475 ));
476 out.push_str(&format!(
477 " #[inline] pub fn log2(x: {ft}) -> {ft} {{ libm::log2(x as f64) as {ft} }}\n"
478 ));
479 out.push_str(&format!(
480 " #[inline] pub fn exp2(x: {ft}) -> {ft} {{ libm::exp2(x as f64) as {ft} }}\n"
481 ));
482 out.push_str(&format!(
483 " #[inline] pub fn fma(a: {ft}, b: {ft}, c: {ft}) -> {ft} {{ libm::fma(a as f64, b as f64, c as f64) as {ft} }}\n"
484 ));
485 out.push_str(&format!(
486 " #[inline] pub fn sin_cos(x: {ft}) -> ({ft}, {ft}) {{ (libm::sin(x as f64) as {ft}, libm::cos(x as f64) as {ft}) }}\n"
487 ));
488 out.push_str("}\n");
489}
490
491fn strip_cfg_gated_module(code: &str) -> String {
496 let mut result = String::new();
497 let mut lines = code.lines().peekable();
498
499 while let Some(line) = lines.next() {
500 if line.starts_with("#[cfg(") && line.contains("feature") {
501 if let Some(&next) = lines.peek()
503 && next.starts_with("mod math {")
504 {
505 lines.next();
507 let mut brace_depth = 1;
508 while brace_depth > 0 {
509 if let Some(inner) = lines.next() {
510 for ch in inner.chars() {
511 if ch == '{' {
512 brace_depth += 1;
513 } else if ch == '}' {
514 brace_depth -= 1;
515 }
516 }
517 } else {
518 break;
519 }
520 }
521 if let Some(&next_after) = lines.peek()
523 && next_after.trim().is_empty()
524 {
525 lines.next();
526 }
527 continue;
528 }
529 }
530
531 result.push_str(line);
532 result.push('\n');
533 }
534
535 let trimmed = result.trim_start_matches('\n');
537 trimmed.to_string()
538}
539
540#[derive(serde::Deserialize)]
546struct RobotConfig {
547 #[allow(dead_code)]
548 robot: RobotInfo,
549 joints: Vec<JointConfig>,
550 generate: GenerateConfig,
551}
552
553#[derive(serde::Deserialize)]
555struct RobotInfo {
556 #[allow(dead_code)]
557 name: String,
558}
559
560#[derive(serde::Deserialize)]
562struct JointConfig {
563 theta: String,
564 #[serde(default)]
565 d: f64,
566 #[serde(default)]
567 a: f64,
568 #[serde(default)]
569 alpha: f64,
570}
571
572fn default_functions() -> Vec<String> {
573 vec!["fk".to_string(), "jacobian".to_string()]
574}
575
576fn default_output() -> String {
577 "robot_math.rs".to_string()
578}
579
580#[derive(serde::Deserialize)]
582struct GenerateConfig {
583 #[serde(default = "default_functions")]
584 functions: Vec<String>,
585 #[serde(default = "default_output")]
586 #[allow(dead_code)]
587 output: String,
588}
589
590pub fn from_toml(path: impl AsRef<Path>) -> Result<CodeGen, Box<dyn std::error::Error>> {
618 let content = fs::read_to_string(path.as_ref())?;
619 let config: RobotConfig = toml::from_str(&content)?;
620
621 let ctx = Context::new();
623
624 let theta_vars: Vec<Ex> = config
627 .joints
628 .iter()
629 .map(|j| joint_symbol(&ctx, &j.theta))
630 .collect::<Result<_, _>>()?;
631
632 let d_vals: Vec<Ex> = config
634 .joints
635 .iter()
636 .map(|j| float_to_expr(&ctx, j.d))
637 .collect::<Result<_, _>>()?;
638 let a_vals: Vec<Ex> = config
639 .joints
640 .iter()
641 .map(|j| float_to_expr(&ctx, j.a))
642 .collect::<Result<_, _>>()?;
643 let alpha_vals: Vec<Ex> = config
644 .joints
645 .iter()
646 .map(|j| float_to_expr(&ctx, j.alpha))
647 .collect::<Result<_, _>>()?;
648
649 let dh_params: Vec<DhLink<'_>> = theta_vars
650 .iter()
651 .enumerate()
652 .map(|(i, theta)| DhLink {
653 theta,
654 d: &d_vals[i],
655 a: &a_vals[i],
656 alpha: &alpha_vals[i],
657 })
658 .collect();
659
660 let theta_names: Vec<&str> = config.joints.iter().map(|j| j.theta.as_str()).collect();
661
662 let mut codegen = CodeGen::new();
663
664 for func in &config.generate.functions {
665 match func.as_str() {
666 "fk" => {
667 let (x, y, z) = symplex::robotics::fk_position(&dh_params);
668 codegen = codegen.add_scalar_fn("fk_x", &x, &theta_names);
669 codegen = codegen.add_scalar_fn("fk_y", &y, &theta_names);
670 codegen = codegen.add_scalar_fn("fk_z", &z, &theta_names);
671 }
672 "jacobian" => {
673 let (x, y, _z) = symplex::robotics::fk_position(&dh_params);
674 let theta_refs: Vec<&Ex> = theta_vars.iter().collect();
675 let j = symplex::matrix::jacobian(&[&x, &y], &theta_refs)?;
676 codegen = codegen.add_matrix_fn("jacobian", &j, &theta_names);
677 }
678 "fk_matrix" => {
679 let t = symplex::robotics::fk_chain(&dh_params);
680 codegen = codegen.add_matrix_fn("fk_matrix", &t, &theta_names);
681 }
682 other => {
683 return Err(format!("unknown generate function: {other}").into());
684 }
685 }
686 }
687
688 Ok(codegen)
689}
690
691fn joint_symbol(ctx: &Context, name: &str) -> Result<Ex, symplex::errors::SymplexError> {
699 ctx.try_symbol(name).map_err(|_| {
700 symplex::errors::SymplexError::invalid_argument(
701 "symplex-build",
702 "a joint's variable name must not be empty",
703 )
704 })
705}
706
707fn float_to_expr(ctx: &Context, v: f64) -> Result<Ex, symplex::errors::SymplexError> {
720 if !v.is_finite() {
721 return Err(symplex::errors::SymplexError::invalid_argument(
722 "symplex-build",
723 format!("a DH parameter must be a finite number, got {v}"),
724 ));
725 }
726 ctx.from_f64_approx(v, 1_000_000)
727}
728
729pub fn robot_arm(joints: &[(&str, f64, f64, f64)]) -> RobotArmBuilder {
750 let owned: Vec<(String, f64, f64, f64)> = joints
751 .iter()
752 .map(|(name, d, a, alpha)| (name.to_string(), *d, *a, *alpha))
753 .collect();
754 RobotArmBuilder::new(owned)
755}
756
757#[derive(Default)]
759struct DhTable {
760 thetas: Vec<Ex>,
761 d: Vec<Ex>,
762 a: Vec<Ex>,
763 alpha: Vec<Ex>,
764}
765
766impl DhTable {
767 fn links(&self) -> Vec<DhLink<'_>> {
769 self.thetas
770 .iter()
771 .zip(&self.d)
772 .zip(&self.a)
773 .zip(&self.alpha)
774 .map(|(((theta, d), a), alpha)| DhLink { theta, d, a, alpha })
775 .collect()
776 }
777}
778
779pub struct RobotArmBuilder {
793 ctx: Context,
794 joints: Vec<(String, f64, f64, f64)>,
795 codegen: CodeGen,
796 generated_fk: bool,
797 generated_jacobian: bool,
798 error: Option<symplex::errors::SymplexError>,
799}
800
801impl RobotArmBuilder {
802 fn new(joints: Vec<(String, f64, f64, f64)>) -> Self {
803 Self {
804 ctx: Context::new(),
805 joints,
806 codegen: CodeGen::new(),
807 generated_fk: false,
808 generated_jacobian: false,
809 error: None,
810 }
811 }
812
813 fn build_dh(&self) -> Result<DhTable, symplex::errors::SymplexError> {
823 let ctx = &self.ctx;
824 let mut table = DhTable::default();
825 for (name, d, a, alpha) in &self.joints {
826 table.thetas.push(joint_symbol(ctx, name)?);
827 table.d.push(float_to_expr(ctx, *d)?);
828 table.a.push(float_to_expr(ctx, *a)?);
829 table.alpha.push(float_to_expr(ctx, *alpha)?);
830 }
831 Ok(table)
832 }
833
834 fn dh_or_record(&mut self) -> Option<DhTable> {
837 match self.build_dh() {
838 Ok(table) => Some(table),
839 Err(e) => {
840 self.error = self.error.take().or(Some(e));
841 None
842 }
843 }
844 }
845
846 fn theta_names_owned(&self) -> Vec<String> {
847 self.joints
848 .iter()
849 .map(|(name, _, _, _)| name.clone())
850 .collect()
851 }
852
853 pub fn generate_fk(mut self, name: &str) -> Self {
855 let Some(table) = self.dh_or_record() else {
856 self.generated_fk = true;
857 return self;
858 };
859 let dh = table.links();
860 let owned_names = self.theta_names_owned();
861 let theta_names: Vec<&str> = owned_names.iter().map(|s| s.as_str()).collect();
862 let (x, y, z) = symplex::robotics::fk_position(&dh);
863
864 let name_x = format!("{name}_x");
865 let name_y = format!("{name}_y");
866 let name_z = format!("{name}_z");
867
868 self.codegen = self.codegen.add_scalar_fn(&name_x, &x, &theta_names);
869 self.codegen = self.codegen.add_scalar_fn(&name_y, &y, &theta_names);
870 self.codegen = self.codegen.add_scalar_fn(&name_z, &z, &theta_names);
871 self.generated_fk = true;
872 self
873 }
874
875 pub fn generate_fk_matrix(mut self, name: &str) -> Self {
883 let Some(table) = self.dh_or_record() else {
884 return self;
885 };
886 let dh = table.links();
887 let owned_names = self.theta_names_owned();
888 let theta_names: Vec<&str> = owned_names.iter().map(|s| s.as_str()).collect();
889 let t = symplex::robotics::fk_chain(&dh);
890 self.codegen = self.codegen.add_matrix_fn(name, &t, &theta_names);
891 self
892 }
893
894 pub fn generate_jacobian(mut self, name: &str) -> Self {
896 let Some(table) = self.dh_or_record() else {
897 self.generated_jacobian = true;
898 return self;
899 };
900 let dh = table.links();
901 let owned_names = self.theta_names_owned();
902 let theta_names: Vec<&str> = owned_names.iter().map(|s| s.as_str()).collect();
903 let (x, y, _z) = symplex::robotics::fk_position(&dh);
904 let theta_refs: Vec<&Ex> = table.thetas.iter().collect();
905 match symplex::matrix::jacobian(&[&x, &y], &theta_refs) {
906 Ok(j) => self.codegen = self.codegen.add_matrix_fn(name, &j, &theta_names),
907 Err(e) => self.error = self.error.or(Some(e)),
908 }
909 self.generated_jacobian = true;
910 self
911 }
912
913 pub fn generate_all(self) -> Self {
915 let s = if !self.generated_fk {
916 self.generate_fk("fk")
917 } else {
918 self
919 };
920 if !s.generated_jacobian {
921 s.generate_jacobian("jacobian")
922 } else {
923 s
924 }
925 }
926
927 pub fn no_std(mut self) -> Self {
929 self.codegen = self.codegen.no_std(true);
930 self
931 }
932
933 pub fn write_to_out_dir(self, filename: &str) -> Result<(), Box<dyn std::error::Error>> {
940 if let Some(e) = self.error {
941 return Err(e.into());
942 }
943 self.codegen.write_to_out_dir(filename)
944 }
945
946 pub fn write_to_path(self, path: impl AsRef<Path>) -> Result<(), Box<dyn std::error::Error>> {
953 if let Some(e) = self.error {
954 return Err(e.into());
955 }
956 self.codegen.write_to_path(path)
957 }
958
959 pub fn into_codegen(self) -> CodeGen {
961 self.codegen
962 }
963}
964
965#[cfg(test)]
970mod tests {
971 use super::*;
972
973 #[test]
974 fn codegen_new_default() {
975 let cg = CodeGen::new();
976 assert!(cg.functions.is_empty());
978 assert!(!cg.generate_tests);
979 }
980
981 #[test]
982 fn codegen_default_trait() {
983 let cg = CodeGen::default();
985 assert!(cg.functions.is_empty());
986 }
987
988 #[test]
989 fn codegen_generate_empty() {
990 let cg = CodeGen::new();
991 let code = cg.generate().unwrap();
992 assert!(
994 code.contains("Auto-generated by symplex-build"),
995 "expected header comment in generated output, got: {code}"
996 );
997 assert!(
999 !code.contains("fn "),
1000 "expected no function definitions in empty codegen"
1001 );
1002 }
1003
1004 #[test]
1005 fn float_to_expr_is_exact_and_reduced() {
1006 let ctx = Context::new();
1007 let show = |v: f64| format!("{}", float_to_expr(&ctx, v).unwrap());
1008 assert_eq!(show(0.0), "0");
1009 assert_eq!(show(2.0), "2");
1010 assert_eq!(show(-3.0), "-3");
1011 assert_eq!(show(0.3), "3/10");
1012 assert_eq!(show(0.25), "1/4");
1013 assert_eq!(show(0.1 + 0.2), "3/10");
1014 assert_eq!(show(1.0 / 3.0), "1/3");
1015 assert_eq!(show(0.123456), "1929/15625");
1016 for bad in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
1018 assert!(float_to_expr(&ctx, bad).is_err(), "{bad}");
1019 }
1020 }
1021
1022 #[test]
1023 fn robot_arm_generates_fk_matrix() {
1024 let code = robot_arm(&[("q1", 0.0, 0.3, 0.0), ("q2", 0.1, 0.25, 0.0)])
1025 .generate_fk_matrix("fk_t")
1026 .into_codegen()
1027 .generate()
1028 .unwrap();
1029 assert!(code.contains("fn fk_t("), "{code}");
1030 assert!(
1031 code.contains("q1: f64") && code.contains("q2: f64"),
1032 "{code}"
1033 );
1034 assert!(code.contains("[f64; 16]"), "expected 4×4 matrix:\n{code}");
1036 assert!(!code.contains("0.30000000000000004"), "{code}");
1038 }
1039
1040 #[test]
1041 fn fk_matrix_last_column_matches_fk_position() {
1042 let ctx = Context::new();
1045 let (q1, q2) = (ctx.symbol("q1"), ctx.symbol("q2"));
1046 let zero = ctx.int(0);
1047 let (l1, l2) = (
1048 float_to_expr(&ctx, 0.3).unwrap(),
1049 float_to_expr(&ctx, 0.25).unwrap(),
1050 );
1051 let dh = [
1052 DhLink {
1053 theta: &q1,
1054 d: &zero,
1055 a: &l1,
1056 alpha: &zero,
1057 },
1058 DhLink {
1059 theta: &q2,
1060 d: &zero,
1061 a: &l2,
1062 alpha: &zero,
1063 },
1064 ];
1065 let t = symplex::robotics::fk_chain(&dh);
1066 let (x, y, z) = symplex::robotics::fk_position(&dh);
1067 assert_eq!(t.get(0, 3).eval(), x);
1068 assert_eq!(t.get(1, 3).eval(), y);
1069 assert_eq!(t.get(2, 3).eval(), z);
1070 }
1071
1072 #[test]
1073 fn from_toml_accepts_fk_matrix() {
1074 let dir = std::env::temp_dir().join(format!("symplex_build_fkm_{}", std::process::id()));
1075 fs::create_dir_all(&dir).unwrap();
1076 let path = dir.join("robot.toml");
1077 fs::write(
1078 &path,
1079 r#"
1080[robot]
1081name = "one_link"
1082[[joints]]
1083theta = "q"
1084a = 0.5
1085[generate]
1086functions = ["fk_matrix"]
1087"#,
1088 )
1089 .unwrap();
1090 let code = from_toml(&path).unwrap().generate().unwrap();
1091 assert!(code.contains("fn fk_matrix("), "{code}");
1092 assert!(code.contains("[f64; 16]"), "{code}");
1093 let _ = fs::remove_dir_all(&dir);
1094 }
1095
1096 #[test]
1097 fn robot_arm_generates_fk_and_jacobian() {
1098 let code = robot_arm(&[("theta1", 0.0, 0.3, 0.0), ("theta2", 0.0, 0.25, 0.0)])
1099 .generate_all()
1100 .into_codegen()
1101 .generate()
1102 .unwrap();
1103
1104 for name in ["fk_x", "fk_y", "fk_z", "jacobian"] {
1105 assert!(
1106 code.contains(&format!("fn {name}(")),
1107 "expected `{name}` in generated code:\n{code}"
1108 );
1109 }
1110 assert!(code.contains("theta1: f64") && code.contains("theta2: f64"));
1111 assert!(code.contains("[f64; 4]"), "expected 2×2 Jacobian:\n{code}");
1113 }
1114
1115 #[test]
1116 fn robot_arm_no_std_emits_single_math_module() {
1117 let code = robot_arm(&[("q", 0.0, 1.0, 0.0)])
1118 .no_std()
1119 .generate_all()
1120 .into_codegen()
1121 .generate()
1122 .unwrap();
1123 assert_eq!(code.matches("mod math {").count(), 2, "{code}");
1125 }
1126
1127 fn two_special_fns(opts: CodegenOptions) -> CodeGen {
1129 let ctx = Context::new();
1130 let x = ctx.symbol("x");
1131 let y = ctx.symbol("y");
1132 CodeGen::new()
1133 .options(opts)
1134 .add_scalar_fn("g", &(x.gamma() + &y), &["x", "y"])
1135 .add_scalar_fn("e", &(x.erf() * &y), &["x", "y"])
1136 }
1137
1138 #[test]
1139 fn generate_emits_runtime_module_once_for_two_special_functions() {
1140 let code = two_special_fns(CodegenOptions::default())
1141 .generate()
1142 .unwrap();
1143 assert_eq!(code.matches("mod symplex_rt {").count(), 1, "{code}");
1144 assert!(
1145 code.contains("pub fn gamma(") && code.contains("pub fn erf("),
1146 "{code}"
1147 );
1148 assert!(code.contains("fn g(") && code.contains("fn e("), "{code}");
1149 assert!(code.find("mod symplex_rt {").unwrap() < code.find("fn g(").unwrap());
1151 assert!(!code.contains("pub fn bessel_k("), "{code}");
1153 }
1154
1155 #[test]
1156 fn generate_emits_runtime_module_once_in_no_std_mode() {
1157 let code = two_special_fns(CodegenOptions::no_std())
1158 .generate()
1159 .unwrap();
1160 assert_eq!(code.matches("mod symplex_rt {").count(), 1, "{code}");
1161 assert_eq!(code.matches("mod math {").count(), 2, "{code}");
1162 let math_pos = code.find("mod math {").unwrap();
1164 let rt_pos = code.find("mod symplex_rt {").unwrap();
1165 let fn_pos = code.find("fn g(").unwrap();
1166 assert!(math_pos < rt_pos && rt_pos < fn_pos, "{code}");
1167 }
1168
1169 #[test]
1170 fn generate_honours_emit_runtime_false() {
1171 let opts = CodegenOptions {
1172 emit_runtime: false,
1173 ..Default::default()
1174 };
1175 let code = two_special_fns(opts).generate().unwrap();
1176 assert_eq!(code.matches("mod symplex_rt {").count(), 0, "{code}");
1177 assert!(code.contains("symplex_rt::gamma("), "{code}");
1178 }
1179
1180 #[test]
1181 fn generate_omits_runtime_when_unused() {
1182 let ctx = Context::new();
1183 let x = ctx.symbol("x");
1184 let code = CodeGen::new()
1185 .add_scalar_fn("f", &(x.sin() + x.powi(2)), &["x"])
1186 .generate()
1187 .unwrap();
1188 assert!(!code.contains("mod symplex_rt"), "{code}");
1189 }
1190
1191 #[test]
1194 fn generated_file_with_two_special_functions_compiles() {
1195 let Ok(out) = std::process::Command::new("rustc")
1196 .arg("--version")
1197 .output()
1198 else {
1199 eprintln!("rustc not available; skipping compile check");
1200 return;
1201 };
1202 if !out.status.success() {
1203 return;
1204 }
1205 let code = two_special_fns(CodegenOptions::default())
1206 .generate()
1207 .unwrap();
1208 let dir = std::env::temp_dir().join(format!("symplex_build_rt_{}", std::process::id()));
1209 fs::create_dir_all(&dir).unwrap();
1210 let src = dir.join("gen.rs");
1211 fs::write(&src, format!("#![allow(dead_code)]\n{code}")).unwrap();
1212 let out = std::process::Command::new("rustc")
1213 .args(["--crate-type", "lib", "--edition", "2024", "-o"])
1214 .arg(dir.join("gen.rlib"))
1215 .arg(&src)
1216 .output()
1217 .unwrap();
1218 let stderr = String::from_utf8_lossy(&out.stderr).into_owned();
1219 let _ = fs::remove_dir_all(&dir);
1220 assert!(
1221 out.status.success(),
1222 "generated file failed to compile:\n{stderr}\n{code}"
1223 );
1224 }
1225
1226 #[test]
1227 fn from_toml_round_trip() {
1228 let dir = std::env::temp_dir().join(format!("symplex_build_{}", std::process::id()));
1229 fs::create_dir_all(&dir).unwrap();
1230 let path = dir.join("robot.toml");
1231 fs::write(
1232 &path,
1233 r#"
1234[robot]
1235name = "two_link"
1236
1237[[joints]]
1238theta = "theta1"
1239a = 0.3
1240
1241[[joints]]
1242theta = "theta2"
1243a = 0.25
1244
1245[generate]
1246functions = ["fk", "jacobian"]
1247"#,
1248 )
1249 .unwrap();
1250
1251 let code = from_toml(&path).unwrap().generate().unwrap();
1252 assert!(code.contains("fn fk_x("), "{code}");
1253 assert!(code.contains("fn jacobian("), "{code}");
1254
1255 let _ = fs::remove_dir_all(&dir);
1256 }
1257
1258 #[test]
1259 fn from_toml_rejects_unknown_function() {
1260 let dir = std::env::temp_dir().join(format!("symplex_build_bad_{}", std::process::id()));
1261 fs::create_dir_all(&dir).unwrap();
1262 let path = dir.join("robot.toml");
1263 fs::write(
1264 &path,
1265 r#"
1266[robot]
1267name = "r"
1268[[joints]]
1269theta = "q"
1270[generate]
1271functions = ["dynamics"]
1272"#,
1273 )
1274 .unwrap();
1275 let err = from_toml(&path)
1276 .err()
1277 .expect("unknown function should error");
1278 assert!(err.to_string().contains("dynamics"), "{err}");
1279 let _ = fs::remove_dir_all(&dir);
1280 }
1281
1282 #[test]
1287 fn robot_arm_reports_an_empty_joint_name_instead_of_panicking() {
1288 let dir = std::env::temp_dir().join(format!("symplex_build_empty_{}", std::process::id()));
1289 let path = dir.join("arm.rs");
1290 let err = robot_arm(&[("q1", 0.0, 0.3, 0.0), ("", 0.0, 0.25, 0.0)])
1291 .generate_all()
1292 .generate_fk_matrix("fk_t")
1293 .write_to_path(&path)
1294 .expect_err("an empty joint name is an error");
1295 assert!(err.to_string().contains("name must not be empty"), "{err}");
1296 assert!(!path.exists(), "nothing is written");
1297 let generated = robot_arm(&[("", 0.0, 0.3, 0.0)])
1298 .generate_fk("fk")
1299 .into_codegen();
1300 assert!(generated.functions.is_empty());
1301 let err = robot_arm(&[("q1", f64::NAN, 0.3, 0.0)])
1302 .generate_jacobian("j")
1303 .write_to_path(&path)
1304 .expect_err("a NaN DH parameter is an error");
1305 assert!(err.to_string().contains("finite"), "{err}");
1306 assert!(!path.exists(), "nothing is written");
1307 }
1308
1309 #[test]
1312 fn from_toml_rejects_an_empty_joint_name_and_a_nan_parameter() {
1313 let dir =
1314 std::env::temp_dir().join(format!("symplex_build_empty_toml_{}", std::process::id()));
1315 fs::create_dir_all(&dir).unwrap();
1316 let path = dir.join("robot.toml");
1317 for (joint, needle) in [
1318 ("theta = \"\"", "name must not be empty"),
1319 ("theta = \"q\"\nd = nan", "finite"),
1320 ] {
1321 fs::write(
1322 &path,
1323 format!(
1324 "[robot]\nname = \"r\"\n[[joints]]\n{joint}\n[generate]\nfunctions = [\"fk\"]\n"
1325 ),
1326 )
1327 .unwrap();
1328 let err = from_toml(&path).err().expect("a bad joint is an error");
1329 assert!(err.to_string().contains(needle), "{err}");
1330 }
1331 let _ = fs::remove_dir_all(&dir);
1332 }
1333}