1use std::collections::HashMap;
7use std::fs;
8use std::path::Path;
9
10#[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#[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#[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#[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
54pub 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
62pub 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 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 output.push_str("class TestTorshBindings(unittest.TestCase):\n");
86 output.push_str(" \"\"\"Test suite for ToRSh Python bindings\"\"\"\n\n");
87
88 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 for test_case in test_cases {
95 output.push_str(&self.generate_single_test(test_case));
96 output.push('\n');
97 }
98
99 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 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 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 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 }
185 }
186 }
187
188 output
189 }
190}
191
192pub 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 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 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 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 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 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 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 }
323 }
324 }
325
326 output.push_str(" });\n");
327 output
328 }
329}
330
331pub 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 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 for test_case in test_cases {
359 output.push_str(&self.generate_single_test(test_case));
360 output.push('\n');
361 }
362
363 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 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 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 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 }
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
475pub 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
550pub 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 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
594fn 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}