1use mumu::{
4 parser::interpreter::Interpreter,
5 parser::types::{FunctionValue, Value},
6};
7use std::sync::{Arc, Mutex};
8
9fn to_float2d(v: &Value) -> Option<Vec<Vec<f64>>> {
10 match v {
11 Value::Float2DArray(rows) => Some(rows.clone()),
12 Value::Int2DArray(rows) => Some(rows.iter().map(|r| r.iter().map(|&x| x as f64).collect()).collect()),
13 Value::MixedArray(rows) => {
14 let mut result = Vec::with_capacity(rows.len());
15 for row in rows {
16 match row {
17 Value::FloatArray(xs) => result.push(xs.clone()),
18 Value::IntArray(xs) => result.push(xs.iter().map(|&x| x as f64).collect()),
19 Value::MixedArray(inner) => {
20 let mut vrow = Vec::with_capacity(inner.len());
21 for v in inner {
22 match v {
23 Value::Int(i) => vrow.push(*i as f64),
24 Value::Float(f) => vrow.push(*f),
25 _ => return None,
26 }
27 }
28 result.push(vrow);
29 }
30 _ => return None,
31 }
32 }
33 Some(result)
34 }
35 _ => None,
36 }
37}
38
39fn matrix_subtract(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
40 if args.len() != 2 {
41 return Err("matrix:subtract expects 2 arguments".to_string());
42 }
43 match (&args[0], &args[1]) {
44 (Value::Int2DArray(xs), Value::Int2DArray(ys))
45 if !xs.is_empty()
46 && xs.len() == ys.len()
47 && xs[0].len() == ys[0].len() =>
48 {
49 let result: Vec<Vec<i32>> = xs
50 .iter()
51 .zip(ys.iter())
52 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a - b).collect())
53 .collect();
54 Ok(Value::Int2DArray(result))
55 }
56 (Value::Float2DArray(xs), Value::Float2DArray(ys))
57 if !xs.is_empty()
58 && xs.len() == ys.len()
59 && xs[0].len() == ys[0].len() =>
60 {
61 let result: Vec<Vec<f64>> = xs
62 .iter()
63 .zip(ys.iter())
64 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a - b).collect())
65 .collect();
66 Ok(Value::Float2DArray(result))
67 }
68 (Value::Int2DArray(xs), Value::Float2DArray(ys))
69 if !xs.is_empty()
70 && xs.len() == ys.len()
71 && xs[0].len() == ys[0].len() =>
72 {
73 let result: Vec<Vec<f64>> = xs
74 .iter()
75 .zip(ys.iter())
76 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| *a as f64 - b).collect())
77 .collect();
78 Ok(Value::Float2DArray(result))
79 }
80 (Value::Float2DArray(xs), Value::Int2DArray(ys))
81 if !xs.is_empty()
82 && xs.len() == ys.len()
83 && xs[0].len() == ys[0].len() =>
84 {
85 let result: Vec<Vec<f64>> = xs
86 .iter()
87 .zip(ys.iter())
88 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a - *b as f64).collect())
89 .collect();
90 Ok(Value::Float2DArray(result))
91 }
92 (a, b) => {
93 let a_float = to_float2d(a).ok_or("Not numeric")?;
94 let b_float = to_float2d(b).ok_or("Not numeric")?;
95 if a_float.is_empty()
96 || b_float.is_empty()
97 || a_float.len() != b_float.len()
98 || a_float[0].len() != b_float[0].len()
99 {
100 return Err("Matrix dimensions do not match".to_string());
101 }
102 let result: Vec<Vec<f64>> = a_float
103 .iter()
104 .zip(b_float.iter())
105 .map(|(row_a, row_b)| row_a.iter().zip(row_b.iter()).map(|(a, b)| a - b).collect())
106 .collect();
107 Ok(Value::Float2DArray(result))
108 }
109 }
110}
111
112fn matrix_add(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
113 if args.len() != 2 {
114 return Err("matrix:add expects 2 arguments".to_string());
115 }
116 match (&args[0], &args[1]) {
117 (Value::Int2DArray(xs), Value::Int2DArray(ys))
118 if !xs.is_empty()
119 && xs.len() == ys.len()
120 && xs[0].len() == ys[0].len() =>
121 {
122 let result: Vec<Vec<i32>> = xs
123 .iter()
124 .zip(ys.iter())
125 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a + b).collect())
126 .collect();
127 Ok(Value::Int2DArray(result))
128 }
129 (Value::Float2DArray(xs), Value::Float2DArray(ys))
130 if !xs.is_empty()
131 && xs.len() == ys.len()
132 && xs[0].len() == ys[0].len() =>
133 {
134 let result: Vec<Vec<f64>> = xs
135 .iter()
136 .zip(ys.iter())
137 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a + b).collect())
138 .collect();
139 Ok(Value::Float2DArray(result))
140 }
141 (Value::Int2DArray(xs), Value::Float2DArray(ys))
142 if !xs.is_empty()
143 && xs.len() == ys.len()
144 && xs[0].len() == ys[0].len() =>
145 {
146 let result: Vec<Vec<f64>> = xs
147 .iter()
148 .zip(ys.iter())
149 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| *a as f64 + b).collect())
150 .collect();
151 Ok(Value::Float2DArray(result))
152 }
153 (Value::Float2DArray(xs), Value::Int2DArray(ys))
154 if !xs.is_empty()
155 && xs.len() == ys.len()
156 && xs[0].len() == ys[0].len() =>
157 {
158 let result: Vec<Vec<f64>> = xs
159 .iter()
160 .zip(ys.iter())
161 .map(|(row_x, row_y)| row_x.iter().zip(row_y.iter()).map(|(a, b)| a + *b as f64).collect())
162 .collect();
163 Ok(Value::Float2DArray(result))
164 }
165 (a, b) => {
166 let a_float = to_float2d(a).ok_or("Not numeric")?;
167 let b_float = to_float2d(b).ok_or("Not numeric")?;
168 if a_float.is_empty()
169 || b_float.is_empty()
170 || a_float.len() != b_float.len()
171 || a_float[0].len() != b_float[0].len()
172 {
173 return Err("Matrix dimensions do not match".to_string());
174 }
175 let result: Vec<Vec<f64>> = a_float
176 .iter()
177 .zip(b_float.iter())
178 .map(|(row_a, row_b)| row_a.iter().zip(row_b.iter()).map(|(a, b)| a + b).collect())
179 .collect();
180 Ok(Value::Float2DArray(result))
181 }
182 }
183}
184
185fn matrix_multiply(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
186 if args.len() != 2 {
187 return Err("matrix:multiply expects 2 arguments".to_string());
188 }
189 match (&args[0], &args[1]) {
190 (Value::Int2DArray(a), Value::Int2DArray(b)) => matrix_mul_int(a, b),
191 (Value::Float2DArray(a), Value::Float2DArray(b)) => matrix_mul_float(a, b),
192 (Value::Int2DArray(a), Value::Float2DArray(b)) => matrix_mul_float2d_cast(a, b),
193 (Value::Float2DArray(a), Value::Int2DArray(b)) => matrix_mul_float2d_cast_rev(a, b),
194 (a, b) => {
195 let a_float = to_float2d(a).ok_or("Not numeric")?;
196 let b_float = to_float2d(b).ok_or("Not numeric")?;
197 let n = a_float.len();
198 let m = if n > 0 { a_float[0].len() } else { 0 };
199 let p = if !b_float.is_empty() { b_float[0].len() } else { 0 };
200 if b_float.len() != m {
201 return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b_float.len(), p));
202 }
203 let mut result = vec![vec![0.0f64; p]; n];
204 for i in 0..n {
205 for j in 0..p {
206 for k in 0..m {
207 result[i][j] += a_float[i][k] * b_float[k][j];
208 }
209 }
210 }
211 Ok(Value::Float2DArray(result))
212 }
213 }
214}
215
216fn matrix_mul_int(a: &Vec<Vec<i32>>, b: &Vec<Vec<i32>>) -> Result<Value, String> {
217 let n = a.len();
218 let m = a[0].len();
219 let p = b[0].len();
220 if b.len() != m {
221 return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
222 }
223 let mut result = vec![vec![0i32; p]; n];
224 for i in 0..n {
225 for j in 0..p {
226 for k in 0..m {
227 result[i][j] += a[i][k] * b[k][j];
228 }
229 }
230 }
231 Ok(Value::Int2DArray(result))
232}
233
234fn matrix_mul_float(a: &Vec<Vec<f64>>, b: &Vec<Vec<f64>>) -> Result<Value, String> {
235 let n = a.len();
236 let m = a[0].len();
237 let p = b[0].len();
238 if b.len() != m {
239 return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
240 }
241 let mut result = vec![vec![0.0f64; p]; n];
242 for i in 0..n {
243 for j in 0..p {
244 for k in 0..m {
245 result[i][j] += a[i][k] * b[k][j];
246 }
247 }
248 }
249 Ok(Value::Float2DArray(result))
250}
251
252fn matrix_mul_float2d_cast(a: &Vec<Vec<i32>>, b: &Vec<Vec<f64>>) -> Result<Value, String> {
253 let n = a.len();
254 let m = a[0].len();
255 let p = b[0].len();
256 if b.len() != m {
257 return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
258 }
259 let mut result = vec![vec![0.0f64; p]; n];
260 for i in 0..n {
261 for j in 0..p {
262 for k in 0..m {
263 result[i][j] += a[i][k] as f64 * b[k][j];
264 }
265 }
266 }
267 Ok(Value::Float2DArray(result))
268}
269
270fn matrix_mul_float2d_cast_rev(a: &Vec<Vec<f64>>, b: &Vec<Vec<i32>>) -> Result<Value, String> {
271 let n = a.len();
272 let m = a[0].len();
273 let p = b[0].len();
274 if b.len() != m {
275 return Err(format!("dimension mismatch: left={}x{}, right={}x{}", n, m, b.len(), p));
276 }
277 let mut result = vec![vec![0.0f64; p]; n];
278 for i in 0..n {
279 for j in 0..p {
280 for k in 0..m {
281 result[i][j] += a[i][k] * b[k][j] as f64;
282 }
283 }
284 }
285 Ok(Value::Float2DArray(result))
286}
287
288fn transpose2d<T: Clone>(m: &Vec<Vec<T>>) -> Vec<Vec<T>> {
289 let rows = m.len();
290 let cols = if rows > 0 { m[0].len() } else { 0 };
291 (0..cols)
292 .map(|c| (0..rows).map(|r| m[r][c].clone()).collect())
293 .collect()
294}
295
296fn matrix_transpose(_interp: &mut Interpreter, args: Vec<Value>) -> Result<Value, String> {
297 if args.len() != 1 {
298 return Err("matrix:transpose expects 1 argument".to_string());
299 }
300 match &args[0] {
301 Value::Int2DArray(xs) => Ok(Value::Int2DArray(transpose2d(xs))),
302 Value::Float2DArray(xs) => Ok(Value::Float2DArray(transpose2d(xs))),
303 Value::MixedArray(rows) => {
304 let as_float2d = to_float2d(&Value::MixedArray(rows.clone())).ok_or("Not numeric")?;
305 Ok(Value::Float2DArray(transpose2d(&as_float2d)))
306 }
307 _ => Err("Not numeric".to_string()),
308 }
309}
310
311pub fn register_all(interp: &mut Interpreter) {
315 macro_rules! reg {
316 ($name:expr, $f:expr) => {{
317 let func = Arc::new(Mutex::new($f));
318 interp.register_dynamic_function($name, func);
319 interp.set_variable($name, Value::Function(Box::new(FunctionValue::Named($name.into()))));
320 }};
321 }
322 reg!("matrix:add", matrix_add);
323 reg!("matrix:subtract", matrix_subtract);
324 reg!("matrix:multiply", matrix_multiply);
325 reg!("matrix:transpose", matrix_transpose);
326}
327
328#[cfg(not(target_arch = "wasm32"))]
330#[no_mangle]
331pub unsafe extern "C" fn Cargo_lock(
332 interp_ptr: *mut std::ffi::c_void,
333 _extra_str: *const std::ffi::c_void,
334) -> i32 {
335 let interp = &mut *(interp_ptr as *mut Interpreter);
336 register_all(interp);
337 0
338}