Skip to main content

mumumatrix/
lib.rs

1// src/lib.rs — matrix-mumu (WASM- and Lava-compatible)
2
3use 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
311/* ───────────────────────── Public registration (static) ────────────────── */
312/// Register all `matrix:*` bridges into the provided interpreter.
313/// Used by embedded/static (WASM) builds and can also be used by hosts.
314pub 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/* ───────────── Host/dynamic loader entrypoint (extend("matrix")) ───────── */
329#[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}