Skip to main content

torsh_ffi/
test_generator.rs

1//! Automatic test generator for ToRSh FFI bindings
2//!
3//! This module generates comprehensive test suites for all language bindings,
4//! ensuring consistent behavior and coverage across all supported languages.
5
6use std::collections::HashMap;
7use std::fs;
8use std::path::Path;
9
10/// Test case definition
11#[derive(Debug, Clone)]
12pub struct TestCase {
13    pub name: String,
14    pub description: String,
15    pub category: TestCategory,
16    pub inputs: Vec<TestInput>,
17    pub expected_output: TestOutput,
18    pub tolerance: Option<f64>,
19}
20
21/// Test categories
22#[derive(Debug, Clone, PartialEq)]
23pub enum TestCategory {
24    TensorCreation,
25    BasicOperations,
26    MatrixOperations,
27    Activations,
28    Reductions,
29    ShapeOperations,
30    NeuralNetwork,
31    ErrorHandling,
32}
33
34/// Test input types
35#[derive(Debug, Clone)]
36pub enum TestInput {
37    TensorData(Vec<Vec<f32>>),
38    Scalar(f32),
39    Shape(Vec<usize>),
40    Integer(i32),
41    String(String),
42}
43
44/// Expected test output
45#[derive(Debug, Clone)]
46pub enum TestOutput {
47    Tensor(Vec<Vec<f32>>),
48    Shape(Vec<usize>),
49    Scalar(f32),
50    Error(String),
51    Boolean(bool),
52}
53
54/// Language-specific test generators
55pub trait TestGenerator {
56    fn language_name(&self) -> &str;
57    fn file_extension(&self) -> &str;
58    fn generate_test_file(&self, test_cases: &[TestCase]) -> String;
59    fn generate_single_test(&self, test_case: &TestCase) -> String;
60}
61
62/// Python test generator
63pub struct PythonTestGenerator;
64
65impl TestGenerator for PythonTestGenerator {
66    fn language_name(&self) -> &str {
67        "Python"
68    }
69
70    fn file_extension(&self) -> &str {
71        "py"
72    }
73
74    fn generate_test_file(&self, test_cases: &[TestCase]) -> String {
75        let mut output = String::new();
76
77        // Header
78        output.push_str("#!/usr/bin/env python3\n");
79        output.push_str("\"\"\"Automatically generated tests for ToRSh Python bindings\"\"\"\n\n");
80        output.push_str("import unittest\n");
81        output.push_str("import numpy as np\n");
82        output.push_str("import torsh\n\n");
83
84        // Test class
85        output.push_str("class TestTorshBindings(unittest.TestCase):\n");
86        output.push_str("    \"\"\"Test suite for ToRSh Python bindings\"\"\"\n\n");
87
88        // Setup method
89        output.push_str("    def setUp(self):\n");
90        output.push_str("        \"\"\"Set up test fixtures\"\"\"\n");
91        output.push_str("        self.tolerance = 1e-6\n\n");
92
93        // Generate individual tests
94        for test_case in test_cases {
95            output.push_str(&self.generate_single_test(test_case));
96            output.push('\n');
97        }
98
99        // Main block
100        output.push_str("if __name__ == '__main__':\n");
101        output.push_str("    unittest.main()\n");
102
103        output
104    }
105
106    fn generate_single_test(&self, test_case: &TestCase) -> String {
107        let mut output = String::new();
108
109        output.push_str(&format!("    def test_{}(self):\n", test_case.name));
110        output.push_str(&format!(
111            "        \"\"\"Test: {}\"\"\"\n",
112            test_case.description
113        ));
114
115        // Generate input setup
116        for (i, input) in test_case.inputs.iter().enumerate() {
117            match input {
118                TestInput::TensorData(data) => {
119                    output.push_str(&format!(
120                        "        input_{} = torsh.tensor({})\n",
121                        i,
122                        format_tensor_data(data)
123                    ));
124                }
125                TestInput::Scalar(value) => {
126                    output.push_str(&format!("        input_{} = {}\n", i, value));
127                }
128                TestInput::Shape(shape) => {
129                    output.push_str(&format!("        input_{} = {}\n", i, format_shape(shape)));
130                }
131                TestInput::Integer(value) => {
132                    output.push_str(&format!("        input_{} = {}\n", i, value));
133                }
134                TestInput::String(value) => {
135                    output.push_str(&format!("        input_{} = '{}'\n", i, value));
136                }
137            }
138        }
139
140        // Generate test operation based on category
141        let operation = match test_case.category {
142            TestCategory::TensorCreation => "torsh.tensor(input_0)",
143            TestCategory::BasicOperations => "input_0.add(input_1)",
144            TestCategory::MatrixOperations => "input_0.matmul(input_1)",
145            TestCategory::Activations => "input_0.relu()",
146            TestCategory::Reductions => "input_0.sum()",
147            TestCategory::ShapeOperations => "input_0.reshape(*input_1)",
148            TestCategory::NeuralNetwork => "torsh.nn.linear(input_0, input_1, input_2)",
149            TestCategory::ErrorHandling => "# Error handling test",
150        };
151
152        if test_case.category == TestCategory::ErrorHandling {
153            output.push_str("        with self.assertRaises(Exception):\n");
154            output.push_str(&format!("            result = {}\n", operation));
155        } else {
156            output.push_str(&format!("        result = {}\n", operation));
157
158            // Generate assertion
159            match &test_case.expected_output {
160                TestOutput::Tensor(expected) => {
161                    output.push_str(&format!(
162                        "        expected = {}\n",
163                        format_tensor_data(expected)
164                    ));
165                    output.push_str("        np.testing.assert_allclose(result.numpy(), expected, rtol=self.tolerance)\n");
166                }
167                TestOutput::Shape(expected) => {
168                    output.push_str(&format!(
169                        "        self.assertEqual(result.shape, {})\n",
170                        format_shape(expected)
171                    ));
172                }
173                TestOutput::Scalar(expected) => {
174                    output.push_str(&format!(
175                        "        self.assertAlmostEqual(result.item(), {}, places=6)\n",
176                        expected
177                    ));
178                }
179                TestOutput::Boolean(expected) => {
180                    output.push_str(&format!("        self.assertEqual(result, {})\n", expected));
181                }
182                TestOutput::Error(_) => {
183                    // Already handled above
184                }
185            }
186        }
187
188        output
189    }
190}
191
192/// JavaScript/Node.js test generator
193pub struct JavaScriptTestGenerator;
194
195impl TestGenerator for JavaScriptTestGenerator {
196    fn language_name(&self) -> &str {
197        "JavaScript"
198    }
199
200    fn file_extension(&self) -> &str {
201        "js"
202    }
203
204    fn generate_test_file(&self, test_cases: &[TestCase]) -> String {
205        let mut output = String::new();
206
207        // Header
208        output
209            .push_str("/**\n * Automatically generated tests for ToRSh Node.js bindings\n */\n\n");
210        output.push_str("const { Tensor } = require('@torsh/core');\n");
211        output.push_str("const assert = require('assert');\n\n");
212
213        output.push_str("describe('ToRSh Node.js Bindings', function() {\n");
214        output.push_str("  const tolerance = 1e-6;\n\n");
215
216        // Helper function
217        output.push_str("  function assertTensorClose(actual, expected, tol = tolerance) {\n");
218        output.push_str("    const actualData = actual.data();\n");
219        output.push_str("    const expectedFlat = expected.flat();\n");
220        output.push_str("    assert.strictEqual(actualData.length, expectedFlat.length);\n");
221        output.push_str("    for (let i = 0; i < actualData.length; i++) {\n");
222        output.push_str("      assert(Math.abs(actualData[i] - expectedFlat[i]) < tol,\n");
223        output.push_str(
224            "        `Expected ${expectedFlat[i]}, got ${actualData[i]} at index ${i}`);\n",
225        );
226        output.push_str("    }\n");
227        output.push_str("  }\n\n");
228
229        // Generate individual tests
230        for test_case in test_cases {
231            output.push_str(&self.generate_single_test(test_case));
232            output.push('\n');
233        }
234
235        output.push_str("});\n");
236
237        output
238    }
239
240    fn generate_single_test(&self, test_case: &TestCase) -> String {
241        let mut output = String::new();
242
243        output.push_str(&format!(
244            "  it('{}', function() {{\n",
245            test_case.description
246        ));
247
248        // Generate input setup
249        for (i, input) in test_case.inputs.iter().enumerate() {
250            match input {
251                TestInput::TensorData(data) => {
252                    output.push_str(&format!(
253                        "    const input{} = Tensor.tensor({});\n",
254                        i,
255                        format_js_array(data)
256                    ));
257                }
258                TestInput::Scalar(value) => {
259                    output.push_str(&format!("    const input{} = {};\n", i, value));
260                }
261                TestInput::Shape(shape) => {
262                    output.push_str(&format!(
263                        "    const input{} = {};\n",
264                        i,
265                        format_js_array_1d(shape)
266                    ));
267                }
268                TestInput::Integer(value) => {
269                    output.push_str(&format!("    const input{} = {};\n", i, value));
270                }
271                TestInput::String(value) => {
272                    output.push_str(&format!("    const input{} = '{}';\n", i, value));
273                }
274            }
275        }
276
277        // Generate test operation
278        let operation = match test_case.category {
279            TestCategory::TensorCreation => "Tensor.tensor(input0)",
280            TestCategory::BasicOperations => "input0.add(input1)",
281            TestCategory::MatrixOperations => "input0.matmul(input1)",
282            TestCategory::Activations => "input0.relu()",
283            TestCategory::Reductions => "input0.sum()",
284            TestCategory::ShapeOperations => "input0.reshape(...input1)",
285            TestCategory::NeuralNetwork => "nn.linear(input0, input1, input2)",
286            TestCategory::ErrorHandling => "// Error handling test",
287        };
288
289        if test_case.category == TestCategory::ErrorHandling {
290            output.push_str("    assert.throws(() => {\n");
291            output.push_str(&format!("      const result = {};\n", operation));
292            output.push_str("    });\n");
293        } else {
294            output.push_str(&format!("    const result = {};\n", operation));
295
296            // Generate assertion
297            match &test_case.expected_output {
298                TestOutput::Tensor(expected) => {
299                    output.push_str(&format!(
300                        "    const expected = {};\n",
301                        format_js_array(expected)
302                    ));
303                    output.push_str("    assertTensorClose(result, expected);\n");
304                }
305                TestOutput::Shape(expected) => {
306                    output.push_str(&format!(
307                        "    assert.deepStrictEqual(result.shape(), {});\n",
308                        format_js_array_1d(expected)
309                    ));
310                }
311                TestOutput::Scalar(expected) => {
312                    output.push_str(&format!(
313                        "    assert(Math.abs(result.data()[0] - {}) < tolerance);\n",
314                        expected
315                    ));
316                }
317                TestOutput::Boolean(expected) => {
318                    output.push_str(&format!("    assert.strictEqual(result, {});\n", expected));
319                }
320                TestOutput::Error(_) => {
321                    // Already handled above
322                }
323            }
324        }
325
326        output.push_str("  });\n");
327        output
328    }
329}
330
331/// Lua test generator
332pub struct LuaTestGenerator;
333
334impl TestGenerator for LuaTestGenerator {
335    fn language_name(&self) -> &str {
336        "Lua"
337    }
338
339    fn file_extension(&self) -> &str {
340        "lua"
341    }
342
343    fn generate_test_file(&self, test_cases: &[TestCase]) -> String {
344        let mut output = String::new();
345
346        // Header
347        output.push_str("-- Automatically generated tests for ToRSh Lua bindings\n\n");
348        output.push_str("local torsh = require('torsh')\n");
349        output.push_str("local function assert_close(a, b, tol)\n");
350        output.push_str("  tol = tol or 1e-6\n");
351        output.push_str("  return math.abs(a - b) < tol\n");
352        output.push_str("end\n\n");
353
354        output.push_str("local tests_passed = 0\n");
355        output.push_str("local tests_failed = 0\n\n");
356
357        // Generate individual tests
358        for test_case in test_cases {
359            output.push_str(&self.generate_single_test(test_case));
360            output.push('\n');
361        }
362
363        // Summary
364        output.push_str("print(string.format('Tests completed: %d passed, %d failed', tests_passed, tests_failed))\n");
365        output.push_str("if tests_failed > 0 then\n");
366        output.push_str("  os.exit(1)\n");
367        output.push_str("end\n");
368
369        output
370    }
371
372    fn generate_single_test(&self, test_case: &TestCase) -> String {
373        let mut output = String::new();
374
375        output.push_str(&format!("-- Test: {}\n", test_case.description));
376        output.push_str("do\n");
377        output.push_str("  local success, error_msg = pcall(function()\n");
378
379        // Generate input setup
380        for (i, input) in test_case.inputs.iter().enumerate() {
381            match input {
382                TestInput::TensorData(data) => {
383                    output.push_str(&format!(
384                        "    local input{} = torsh.tensor({})\n",
385                        i,
386                        format_lua_table(data)
387                    ));
388                }
389                TestInput::Scalar(value) => {
390                    output.push_str(&format!("    local input{} = {}\n", i, value));
391                }
392                TestInput::Shape(shape) => {
393                    output.push_str(&format!(
394                        "    local input{} = {}\n",
395                        i,
396                        format_lua_table_1d(shape)
397                    ));
398                }
399                TestInput::Integer(value) => {
400                    output.push_str(&format!("    local input{} = {}\n", i, value));
401                }
402                TestInput::String(value) => {
403                    output.push_str(&format!("    local input{} = '{}'\n", i, value));
404                }
405            }
406        }
407
408        // Generate test operation
409        let operation = match test_case.category {
410            TestCategory::TensorCreation => "torsh.tensor(input0)",
411            TestCategory::BasicOperations => "input0:add(input1)",
412            TestCategory::MatrixOperations => "input0:matmul(input1)",
413            TestCategory::Activations => "input0:relu()",
414            TestCategory::Reductions => "input0:sum()",
415            TestCategory::ShapeOperations => "input0:reshape(table.unpack(input1))",
416            TestCategory::NeuralNetwork => "torsh.nn.linear(input0, input1, input2)",
417            TestCategory::ErrorHandling => "-- Error handling test",
418        };
419
420        if test_case.category != TestCategory::ErrorHandling {
421            output.push_str(&format!("    local result = {}\n", operation));
422
423            // Generate assertion based on expected output
424            match &test_case.expected_output {
425                TestOutput::Tensor(_) => {
426                    output.push_str("    -- Tensor comparison would go here\n");
427                    output.push_str("    assert(result ~= nil, 'Result should not be nil')\n");
428                }
429                TestOutput::Shape(expected) => {
430                    output.push_str(&format!(
431                        "    local expected_shape = {}\n",
432                        format_lua_table_1d(expected)
433                    ));
434                    output.push_str("    local actual_shape = result:shape()\n");
435                    output.push_str(
436                        "    assert(#actual_shape == #expected_shape, 'Shape length mismatch')\n",
437                    );
438                }
439                TestOutput::Scalar(expected) => {
440                    output.push_str(&format!(
441                        "    assert(assert_close(result:data()[1], {}), 'Scalar value mismatch')\n",
442                        expected
443                    ));
444                }
445                TestOutput::Boolean(expected) => {
446                    output.push_str(&format!(
447                        "    assert(result == {}, 'Boolean value mismatch')\n",
448                        expected
449                    ));
450                }
451                TestOutput::Error(_) => {
452                    // Should not reach here for non-error tests
453                }
454            }
455        }
456
457        output.push_str("  end)\n");
458        output.push_str("  \n");
459        output.push_str("  if success then\n");
460        output.push_str(&format!("    print('PASS: {}')\n", test_case.description));
461        output.push_str("    tests_passed = tests_passed + 1\n");
462        output.push_str("  else\n");
463        output.push_str(&format!(
464            "    print('FAIL: {} - ' .. tostring(error_msg))\n",
465            test_case.description
466        ));
467        output.push_str("    tests_failed = tests_failed + 1\n");
468        output.push_str("  end\n");
469        output.push_str("end\n");
470
471        output
472    }
473}
474
475/// Generate standard test cases
476pub fn create_standard_test_cases() -> Vec<TestCase> {
477    vec![
478        TestCase {
479            name: "tensor_creation_2d".to_string(),
480            description: "Create 2D tensor from nested array".to_string(),
481            category: TestCategory::TensorCreation,
482            inputs: vec![TestInput::TensorData(vec![vec![1.0, 2.0], vec![3.0, 4.0]])],
483            expected_output: TestOutput::Shape(vec![2, 2]),
484            tolerance: Some(1e-6),
485        },
486        TestCase {
487            name: "tensor_addition".to_string(),
488            description: "Element-wise tensor addition".to_string(),
489            category: TestCategory::BasicOperations,
490            inputs: vec![
491                TestInput::TensorData(vec![vec![1.0, 2.0], vec![3.0, 4.0]]),
492                TestInput::TensorData(vec![vec![5.0, 6.0], vec![7.0, 8.0]]),
493            ],
494            expected_output: TestOutput::Tensor(vec![vec![6.0, 8.0], vec![10.0, 12.0]]),
495            tolerance: Some(1e-6),
496        },
497        TestCase {
498            name: "matrix_multiplication".to_string(),
499            description: "Matrix multiplication operation".to_string(),
500            category: TestCategory::MatrixOperations,
501            inputs: vec![
502                TestInput::TensorData(vec![vec![1.0, 2.0], vec![3.0, 4.0]]),
503                TestInput::TensorData(vec![vec![5.0, 6.0], vec![7.0, 8.0]]),
504            ],
505            expected_output: TestOutput::Tensor(vec![vec![19.0, 22.0], vec![43.0, 50.0]]),
506            tolerance: Some(1e-6),
507        },
508        TestCase {
509            name: "relu_activation".to_string(),
510            description: "ReLU activation function".to_string(),
511            category: TestCategory::Activations,
512            inputs: vec![TestInput::TensorData(vec![vec![-1.0, 0.0, 1.0, 2.0]])],
513            expected_output: TestOutput::Tensor(vec![vec![0.0, 0.0, 1.0, 2.0]]),
514            tolerance: Some(1e-6),
515        },
516        TestCase {
517            name: "tensor_sum".to_string(),
518            description: "Sum all tensor elements".to_string(),
519            category: TestCategory::Reductions,
520            inputs: vec![TestInput::TensorData(vec![vec![1.0, 2.0], vec![3.0, 4.0]])],
521            expected_output: TestOutput::Scalar(10.0),
522            tolerance: Some(1e-6),
523        },
524        TestCase {
525            name: "zeros_creation".to_string(),
526            description: "Create tensor filled with zeros".to_string(),
527            category: TestCategory::TensorCreation,
528            inputs: vec![TestInput::Shape(vec![3, 3])],
529            expected_output: TestOutput::Tensor(vec![
530                vec![0.0, 0.0, 0.0],
531                vec![0.0, 0.0, 0.0],
532                vec![0.0, 0.0, 0.0],
533            ]),
534            tolerance: Some(1e-6),
535        },
536        TestCase {
537            name: "shape_mismatch_error".to_string(),
538            description: "Matrix multiplication with incompatible shapes should error".to_string(),
539            category: TestCategory::ErrorHandling,
540            inputs: vec![
541                TestInput::TensorData(vec![vec![1.0, 2.0]]),
542                TestInput::TensorData(vec![vec![1.0], vec![2.0], vec![3.0]]),
543            ],
544            expected_output: TestOutput::Error("Shape mismatch".to_string()),
545            tolerance: None,
546        },
547    ]
548}
549
550/// Test suite generator
551pub struct TestSuiteGenerator {
552    generators: HashMap<String, Box<dyn TestGenerator>>,
553}
554
555impl TestSuiteGenerator {
556    pub fn new() -> Self {
557        let mut generators: HashMap<String, Box<dyn TestGenerator>> = HashMap::new();
558        generators.insert("python".to_string(), Box::new(PythonTestGenerator));
559        generators.insert("javascript".to_string(), Box::new(JavaScriptTestGenerator));
560        generators.insert("lua".to_string(), Box::new(LuaTestGenerator));
561
562        Self { generators }
563    }
564
565    pub fn generate_all_tests(&self, output_dir: &Path) -> Result<(), Box<dyn std::error::Error>> {
566        let test_cases = create_standard_test_cases();
567
568        for (lang, generator) in &self.generators {
569            let test_content = generator.generate_test_file(&test_cases);
570            let filename = format!("test_torsh_bindings.{}", generator.file_extension());
571            let file_path = output_dir.join(lang).join(filename);
572
573            // Create directory if it doesn't exist
574            if let Some(parent) = file_path.parent() {
575                fs::create_dir_all(parent)?;
576            }
577
578            fs::write(&file_path, test_content)?;
579            println!(
580                "Generated {} tests: {}",
581                generator.language_name(),
582                file_path.display()
583            );
584        }
585
586        Ok(())
587    }
588
589    pub fn add_generator(&mut self, name: String, generator: Box<dyn TestGenerator>) {
590        self.generators.insert(name, generator);
591    }
592}
593
594/// Helper functions for formatting data structures
595
596fn format_tensor_data(data: &[Vec<f32>]) -> String {
597    let formatted_rows: Vec<String> = data
598        .iter()
599        .map(|row| {
600            format!(
601                "[{}]",
602                row.iter()
603                    .map(|x| x.to_string())
604                    .collect::<Vec<_>>()
605                    .join(", ")
606            )
607        })
608        .collect();
609    format!("[{}]", formatted_rows.join(", "))
610}
611
612fn format_shape(shape: &[usize]) -> String {
613    format!(
614        "[{}]",
615        shape
616            .iter()
617            .map(|x| x.to_string())
618            .collect::<Vec<_>>()
619            .join(", ")
620    )
621}
622
623fn format_js_array(data: &[Vec<f32>]) -> String {
624    let formatted_rows: Vec<String> = data
625        .iter()
626        .map(|row| {
627            format!(
628                "[{}]",
629                row.iter()
630                    .map(|x| x.to_string())
631                    .collect::<Vec<_>>()
632                    .join(", ")
633            )
634        })
635        .collect();
636    format!("[{}]", formatted_rows.join(", "))
637}
638
639fn format_js_array_1d(data: &[usize]) -> String {
640    format!(
641        "[{}]",
642        data.iter()
643            .map(|x| x.to_string())
644            .collect::<Vec<_>>()
645            .join(", ")
646    )
647}
648
649fn format_lua_table(data: &[Vec<f32>]) -> String {
650    let formatted_rows: Vec<String> = data
651        .iter()
652        .map(|row| {
653            format!(
654                "{{{}}}",
655                row.iter()
656                    .map(|x| x.to_string())
657                    .collect::<Vec<_>>()
658                    .join(", ")
659            )
660        })
661        .collect();
662    format!("{{{}}}", formatted_rows.join(", "))
663}
664
665fn format_lua_table_1d(data: &[usize]) -> String {
666    format!(
667        "{{{}}}",
668        data.iter()
669            .map(|x| x.to_string())
670            .collect::<Vec<_>>()
671            .join(", ")
672    )
673}
674
675#[cfg(test)]
676mod tests {
677    use super::*;
678
679    #[test]
680    fn test_python_test_generation() {
681        let generator = PythonTestGenerator;
682        let test_cases = create_standard_test_cases();
683        let result = generator.generate_test_file(&test_cases);
684
685        assert!(result.contains("import unittest"));
686        assert!(result.contains("class TestTorshBindings"));
687        assert!(result.contains("def test_tensor_creation_2d"));
688    }
689
690    #[test]
691    fn test_javascript_test_generation() {
692        let generator = JavaScriptTestGenerator;
693        let test_cases = create_standard_test_cases();
694        let result = generator.generate_test_file(&test_cases);
695
696        assert!(result.contains("const { Tensor }"));
697        assert!(result.contains("describe('ToRSh Node.js Bindings'"));
698    }
699
700    #[test]
701    fn test_standard_test_cases() {
702        let test_cases = create_standard_test_cases();
703        assert!(!test_cases.is_empty());
704        assert!(test_cases
705            .iter()
706            .any(|tc| tc.category == TestCategory::TensorCreation));
707        assert!(test_cases
708            .iter()
709            .any(|tc| tc.category == TestCategory::BasicOperations));
710    }
711
712    #[test]
713    fn test_test_suite_generator() {
714        let generator = TestSuiteGenerator::new();
715        assert!(generator.generators.contains_key("python"));
716        assert!(generator.generators.contains_key("javascript"));
717        assert!(generator.generators.contains_key("lua"));
718    }
719}